diff --git a/src/database.rs b/src/database.rs index ad62db5..d142d49 100644 --- a/src/database.rs +++ b/src/database.rs @@ -1,20 +1,26 @@ use crate::response::Server; use crate::utils; +use futures_util::{future, FutureExt}; +use indicatif::{ProgressBar, ProgressStyle}; use sqlx::postgres::{PgQueryResult, PgRow}; use sqlx::types::ipnet::IpNet; use sqlx::types::Uuid; -use sqlx::{PgConnection, Pool, Postgres}; +use sqlx::{PgConnection, Pool, Postgres, Transaction}; use std::str::FromStr; +use std::sync::Arc; use std::time::{SystemTime, UNIX_EPOCH}; +use tokio::sync::Mutex; use tracing::info; +/// Returns all servers from the database pub async fn fetch_servers(pool: &Pool) -> Vec { - sqlx::query("SELECT address::text, port FROM servers ORDER BY last_seen ASC") + sqlx::query("SELECT address, port FROM servers ORDER BY last_seen ASC") .fetch_all(pool) .await .expect("failed to fetch servers from database") } +/// Deletes a server from the database pub async fn delete_server( address: String, transaction: &mut PgConnection, @@ -25,7 +31,56 @@ pub async fn delete_server( .await } -pub async fn update(server: Server, transaction: &mut PgConnection) -> anyhow::Result<()> { +/// Takes in a list of completed servers and a database connection +/// and handles all joining of tasks, progress bar updates and updating +pub async fn update_servers_from_vec(vec: Vec, pool: Pool) { + // Transactions add multiple SQL statements into one big query + let transaction = Arc::new(Mutex::new( + pool.begin().await.expect("failed to create transaction"), + )); + + let style = ProgressStyle::with_template( + "[{elapsed}] [{bar:40.white/blue}] {pos:>7}/{len:7} ETA {eta}", + ) + .expect("failed to create progress bar style") + .progress_chars("=>-"); + + let bar = ProgressBar::new(vec.len() as u64).with_style(style); + + info!("Commiting {} servers to database", vec.len()); + + // Create all handles + let handles = vec + .into_iter() + .map(|s| { + update_server(s, transaction.clone()).map(|r| { + bar.inc(1); + r + }) + }) + .collect::>(); + + future::join_all(handles).await; + + bar.finish_and_clear(); + + Arc::try_unwrap(transaction) + .unwrap() + .into_inner() + .commit() + .await + .expect("error while commiting to database"); +} + +/// Updates a single server in the database, this includes all mods +/// and players that come with it. Will also remove a server from the +/// database if it has requested to be removed +pub async fn update_server( + server: Server, + transaction: Arc>>, +) -> anyhow::Result<()> { + let conn = &mut **transaction.lock().await; + // SQLx requires each IP address to be in CIDR notation to add to Postgres let address = IpNet::from_str(&(server.address.to_string() + "/32"))?; let timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs() as i32; @@ -38,7 +93,7 @@ pub async fn update(server: Server, transaction: &mut PgConnection) -> anyhow::R let description_formatted = server.build_server_description(description_raw); if server.check_opt_out() { - let modified_rows = delete_server(address.to_string(), transaction) + let modified_rows = delete_server(address.to_string(), conn) .await? .rows_affected(); info!( @@ -96,7 +151,7 @@ pub async fn update(server: Server, transaction: &mut PgConnection) -> anyhow::R .bind(timestamp) .bind(server.players.online) .bind(server.players.max) - .execute(&mut *transaction) + .execute(&mut *conn) .await?; if let Some(sample) = server.players.sample { @@ -105,14 +160,14 @@ pub async fn update(server: Server, transaction: &mut PgConnection) -> anyhow::R sqlx::query("INSERT INTO players (address, port, uuid, name, first_seen, last_seen) VALUES ($1, $2, $3, $4, $5, $6) ON CONFLICT (address, port, uuid) DO UPDATE SET last_seen = EXCLUDED.last_seen") - .bind(&address) - .bind(server.port as i32) - .bind(uuid) - .bind(player.name) - .bind(timestamp) - .bind(timestamp) - .execute(&mut *transaction) - .await?; + .bind(&address) + .bind(server.port as i32) + .bind(uuid) + .bind(player.name) + .bind(timestamp) + .bind(timestamp) + .execute(&mut *conn) + .await?; } } } @@ -120,14 +175,14 @@ pub async fn update(server: Server, transaction: &mut PgConnection) -> anyhow::R if let Some(mods_sample) = server.forge_data { for mods in mods_sample.mods { sqlx::query("INSERT INTO mods (address, port, id, mod_marker) VALUES ($1, $2, $3, $4) ON CONFLICT (address, port, id) DO NOTHING") - .bind(&address) - .bind(server.port as i32) - .bind(mods.id) - .bind(mods.version) - .bind(timestamp) - .bind(timestamp) - .execute(&mut *transaction) - .await?; + .bind(&address) + .bind(server.port as i32) + .bind(mods.id) + .bind(mods.version) + .bind(timestamp) + .bind(timestamp) + .execute(&mut *conn) + .await?; } } diff --git a/src/scanner.rs b/src/scanner.rs index 24ee58d..1b53465 100644 --- a/src/scanner.rs +++ b/src/scanner.rs @@ -2,20 +2,20 @@ use crate::config::Config; use crate::response::Server; use crate::utils::RunError; use crate::{database, ping}; -use futures_util::future; +use futures_util::{future, FutureExt}; use indicatif::{ProgressBar, ProgressStyle}; +use sqlx::types::ipnet::IpNet; use sqlx::{Pool, Postgres, Row}; use std::fmt::Debug; +use std::fs::OpenOptions; use std::io::{BufRead, BufReader}; -use std::sync::Arc; use std::time::Duration; use tokio::process::Command; -use tokio::sync::{Mutex, Semaphore}; -use tracing::log::warn; -use tracing::{debug, error, info}; +use tokio::sync::Semaphore; +use tracing::{debug, error, info, warn}; static PERMITS: Semaphore = Semaphore::const_new(1000); -const TIMEOUT_SECS: Duration = Duration::from_secs(5); +const TIMEOUT_SECS: Duration = Duration::from_secs(3); #[derive(Debug, Default)] pub struct ScanBuilder { @@ -78,42 +78,46 @@ impl Scanner { } /// Starts the scanner based on the selected mode - pub async fn start(self) { + pub async fn start(&self) { match self.mode { - Mode::Discovery => Self::masscan(self).await, - Mode::Rescanner => Self::rescan(self).await, + Mode::Discovery => self.masscan().await, + Mode::Rescanner => self.rescan().await, } } - /// Takes a vector of pingable servers - pub async fn scan_servers_from_vec(servers: Vec<(String, u16)>) -> Vec { + /// Takes a list of ip addresses and ports, either from postgres or masscan + /// and creates a ping job for every item, then joins on them + pub async fn scan_servers_from_vec( + &self, + servers: Vec<(String, u16)>, + ) -> Vec> { let style = ProgressStyle::with_template( "[{elapsed}] [{bar:40.white/blue}] {pos:>7}/{len:7} ETA {eta}", ) - .unwrap() + .expect("failed to create progress bar style") .progress_chars("=>-"); - let bar = Arc::new(ProgressBar::new(servers.len() as u64).with_style(style)); + let bar = ProgressBar::new(servers.len() as u64).with_style(style); let handles = servers .into_iter() - .map(|s| Scanner::run(s.0, s.1, bar.clone())) + .map(|s| { + self.run(s.0, s.1).map(|r| { + bar.inc(1); + r + }) + }) .collect::>(); - // Wait for all tasks to finish let results = future::join_all(handles).await; bar.finish_and_clear(); - results - .into_iter() - // TODO! Don't flatten here, errors are needed later - .flatten() - .collect::>() + results.into_iter().collect() } /// Rescan servers already found in the database - async fn rescan(self) { + async fn rescan(&self) { let port_start = self.config.scanner.port_range_start; let port_end = self.config.scanner.port_range_end; let total_ports = self.config.scanner.total_ports(); @@ -132,10 +136,18 @@ impl Scanner { let servers = database::fetch_servers(&self.pool) .await .into_iter() - .map(|row| (row.get::(0), row.get::(1) as u16)) + .map(|row| { + ( + row.get::(0).addr().to_string(), + row.get::(1) as u16, + ) + }) .collect::>(); - Self::scan_servers_from_vec(servers).await; + let scan_results = self.scan_servers_from_vec(servers).await; + let completed_servers = self.print_scan_results(scan_results).await; + + database::update_servers_from_vec(completed_servers, self.pool.clone()).await; // Quit if only one scan is requested in config if !self.config.scanner.repeat { @@ -155,7 +167,7 @@ impl Scanner { } /// Starts an instance of masscan to find new servers - async fn masscan(self) { + async fn masscan(&self) { let masscan_config = &self.config.masscan.config_file; let masscan_output = &self.config.masscan.output_file; @@ -164,7 +176,10 @@ impl Scanner { } loop { - let output_file = std::fs::File::open(masscan_output).unwrap(); + let file = OpenOptions::new() + .read(true) + .open(masscan_output) + .expect("couldn't open file"); let _ = Command::new("sudo") .args(["masscan", "-c", masscan_config]) @@ -175,7 +190,7 @@ impl Scanner { info!("Masscan has completed!"); - let reader = BufReader::new(output_file); + let reader = BufReader::new(file); let mut servers = Vec::new(); for line in reader.lines() { @@ -199,7 +214,10 @@ impl Scanner { servers.push((address, port)); } - Self::scan_servers_from_vec(servers).await; + let scan_results = self.scan_servers_from_vec(servers).await; + let completed_servers = self.print_scan_results(scan_results).await; + + database::update_servers_from_vec(completed_servers, self.pool.clone()).await; // Quit if only one scan is requested in config if !self.config.scanner.repeat { @@ -218,42 +236,32 @@ impl Scanner { } } - async fn run(address: String, port: u16, bar: Arc) -> anyhow::Result { - async fn inner(address: String, port: u16) -> anyhow::Result { - // Ping server - let permit = PERMITS.acquire().await?; - let pinged_server = tokio::time::timeout( - TIMEOUT_SECS, - // TODO! Fix - ping::ping_server((&*address.split('/').nth(0).unwrap(), port)), - ) - .await??; - drop(permit); + async fn run(&self, address: String, port: u16) -> Result { + // Ping server + let permit = PERMITS + .acquire() + .await + .expect("failed to acquire a semaphore"); + let pinged_server = + tokio::time::timeout(TIMEOUT_SECS, ping::ping_server((&*address, port))).await??; + drop(permit); - // Parse response - let mut server = serde_json::from_str::(&pinged_server)?; - server.address = address; - server.port = port; + // Parse response + let mut server = serde_json::from_str::(&pinged_server)?; + server.address = address; + server.port = port; - Ok(server) - } - - let server = inner(address, port).await; - bar.inc(1); - server + Ok(server) } - pub async fn complete_scan(&self, results: Vec>) { - let results_len = results.len(); - debug!("results_len = {}", results_len); + pub async fn print_scan_results(&self, results: Vec>) -> Vec { + debug!("results_len = {}", results.len()); let (servers, errors): (Vec<_>, Vec<_>) = results.into_iter().partition(Result::is_ok); - let errors_len = errors.len(); - - // Print scan errors + // Print errors if !errors.is_empty() { - warn!("Scan returned {} total errors!", errors_len); + warn!("Scan returned {} total errors!", errors.len()); let mut counts = [0u32; 6]; for e in errors.into_iter().filter_map(Result::err) { let i: usize = e.into(); @@ -268,29 +276,9 @@ impl Scanner { info!("{} servers removed due to opting out", counts[5]) } - // Transactions allow adding multiple statements to a single query - let transaction = Arc::new(Mutex::new( - self.pool - .begin() - .await - .expect("failed to create transaction"), - )); - - let completed_servers = servers + servers .into_iter() .filter_map(Result::ok) - .collect::>(); - - info!( - "Commiting {} servers to database...", - completed_servers.len() - ); - - Arc::try_unwrap(transaction) - .unwrap() - .into_inner() - .commit() - .await - .expect("error while commiting to database"); + .collect::>() } } diff --git a/src/utils.rs b/src/utils.rs index e589020..bd8f26b 100644 --- a/src/utils.rs +++ b/src/utils.rs @@ -1,4 +1,3 @@ -use crate::response::Server; use thiserror::Error; #[derive(Debug, Error)] @@ -18,13 +17,15 @@ pub enum RunError { } impl Into for RunError { fn into(self) -> usize { + use RunError::*; + match self { - Self::AddressParseError(_) => 0, - Self::IOError(_) => 1, - Self::MalformedResponse => 2, - Self::ParseResponse(_) => 3, - Self::TimedOut(_) => 4, - Self::ServerOptOut => 5, + AddressParseError(_) => 0, + IOError(_) => 1, + MalformedResponse => 2, + ParseResponse(_) => 3, + TimedOut(_) => 4, + ServerOptOut => 5, } } } @@ -105,10 +106,3 @@ impl MinecraftColorCodes { } } } - -#[derive(Debug)] -pub struct CompletedServer { - pub ip: String, - pub port: u16, - pub server: Server, -}