diff --git a/Cargo.lock b/Cargo.lock index ecef40b..6591e89 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -6,6 +6,8 @@ version = 4 name = "ServerSeekerV2-Rust" version = "0.1.0" dependencies = [ + "futures-core", + "futures-util", "indicatif", "serde", "serde_json", @@ -366,6 +368,17 @@ version = "0.3.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9e5c1b78ca4aae1ac06c48a526a655760685149f0d465d21f37abfe57ce075c6" +[[package]] +name = "futures-macro" +version = "0.3.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "futures-sink" version = "0.3.31" @@ -386,6 +399,7 @@ checksum = "9fa08315bb612088cc391249efdc3bc77536f16c91f6cf495e6fbe85b20a4a81" dependencies = [ "futures-core", "futures-io", + "futures-macro", "futures-sink", "futures-task", "memchr", diff --git a/Cargo.toml b/Cargo.toml index 186a159..a5149fb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -13,10 +13,12 @@ toml = "0.8.19" tokio = { version = "1.0.0", features = ["full"] } sqlx = { version = "0.8.3", features = ["postgres", "runtime-tokio"] } indicatif = { version = "0.17.11" } +futures-core = "0.3.31" +futures-util = "0.3.31" #anyhow = "1.0.97" #futures = "0.3.31" [profile.release] strip = true lto = "fat" -opt-level = 3 \ No newline at end of file +opt-level = 3 diff --git a/src/database.rs b/src/database.rs index 7f07e45..49849fb 100644 --- a/src/database.rs +++ b/src/database.rs @@ -1,5 +1,7 @@ use crate::colors::{RED, RESET}; use crate::response::Server; +use futures_core::stream::BoxStream; +use sqlx::postgres::PgRow; use sqlx::{Error, PgPool, Pool, Postgres, Row}; use std::time::{SystemTime, UNIX_EPOCH}; @@ -10,15 +12,16 @@ pub async fn connect(url: &str) -> Pool { } } -// TODO! Return a stream of results instead of a Vec for performance -pub async fn fetch_servers(pool: &Pool) -> Result, Error> { - // Sort results by oldest servers first - sqlx::query("SELECT address FROM servers ORDER BY lastseen DESC") - .fetch_all(pool) - .await? - .into_iter() - .map(|row| row.try_get(0)) - .collect() +pub async fn fetch_servers(pool: &Pool) -> BoxStream> { + sqlx::query("SELECT address FROM servers ORDER BY lastseen DESC").fetch(pool) +} + +pub async fn fetch_count(pool: &Pool) -> i64 { + sqlx::query("SELECT COUNT(address) FROM servers") + .fetch_one(pool) + .await + .unwrap() + .get(0) } pub async fn update(server: Server, conn: &PgPool) -> Result<(), Error> { diff --git a/src/main.rs b/src/main.rs index c59bcd4..00f9176 100644 --- a/src/main.rs +++ b/src/main.rs @@ -4,11 +4,13 @@ mod database; mod ping; mod response; +use crate::database::fetch_count; use colors::{GREEN, RED, RESET, YELLOW}; use config::load_config; use database::{connect, fetch_servers}; +use futures_util::TryStreamExt; use indicatif::{ProgressBar, ProgressStyle}; -use sqlx::{Pool, Postgres}; +use sqlx::{Pool, Postgres, Row}; use std::time::{SystemTime, UNIX_EPOCH}; use std::{sync::Arc, time::Duration}; use thiserror::Error; @@ -56,21 +58,8 @@ async fn main() { loop { let pool = connect(database_url.as_str()).await; - - let servers = match fetch_servers(&pool).await { - Ok(servers) => { - println!( - "{GREEN}[INFO] Found {} servers to rescan!{RESET}", - servers.len() - ); - servers - } - Err(_) => { - println!("{RED}[ERROR] Failed to fetch servers! Waiting 30 seconds and retrying...{RESET}"); - tokio::time::sleep(Duration::from_secs(30)).await; - continue; - } - }; + let mut servers = fetch_servers(&pool).await; + let length = fetch_count(&pool).await as u64; let style = ProgressStyle::with_template("[{elapsed}] [{bar:40.white/blue}] {pos:>7}/{len:7}") @@ -79,8 +68,8 @@ async fn main() { // Create state to be passed to each task let state = Arc::new(State { - pool, - progress_bar: ProgressBar::new(servers.len() as u64).with_style(style), + pool: pool.clone(), + progress_bar: ProgressBar::new(length).with_style(style), }); let mut ping_set = JoinSet::new(); @@ -93,9 +82,11 @@ async fn main() { .as_secs() as i64; // Spawn a new task for every result - for ip in servers { + while let Some(row) = servers.try_next().await.unwrap() { + let address: String = row.get(0); + for port in port_start..=port_end { - ping_set.spawn(run((ip.to_owned(), port), state.clone())); + ping_set.spawn(run((address.to_owned(), port), state.clone())); } }