diff --git a/src/database.rs b/src/database.rs index a232b2a..0fe6c0c 100644 --- a/src/database.rs +++ b/src/database.rs @@ -1,14 +1,13 @@ use crate::response::Server; -use crate::utils; +use crate::utils::RunError; use futures_util::stream::BoxStream; use sqlx::postgres::{PgQueryResult, PgRow}; use sqlx::types::ipnet::IpNet; use sqlx::types::Uuid; -use sqlx::{PgConnection, Pool, Postgres, Transaction}; +use sqlx::{Pool, Postgres}; 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 @@ -19,23 +18,18 @@ pub async fn fetch_servers(pool: &Pool) -> BoxStream>, ) -> Result { sqlx::query("DELETE FROM servers WHERE address = $1") .bind(address) - .execute(transaction) + .execute(&*conn) .await } /// 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; - +pub async fn update_server(server: Server, conn: Arc>) -> anyhow::Result<()> { // SQLx requires each IP address to be in CIDR notation to add to Postgres let address = IpNet::from_str(&(server.address.clone() + "/32"))?; let timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs() as i32; @@ -44,19 +38,19 @@ pub async fn update_server( let description_raw = server .description_raw .as_ref() - .ok_or(utils::RunError::MalformedResponse)?; + .ok_or(RunError::MalformedResponse)?; let description_formatted = server.build_server_description(description_raw); + // Delete server if they opted out if server.check_opt_out() { - let modified_rows = delete_server(address.to_string(), conn) + let modified_rows = delete_server(address.addr().to_string(), conn) .await? .rows_affected(); info!( - "Removing {} from the database due to opt out! ({} rows modified)", - address, modified_rows + "Removing {address} from the database due to opt out! ({modified_rows} rows modified)" ); - return Err(utils::RunError::ServerOptOut)?; + return Err(RunError::ServerOptOut)?; } // description_raw is for storing raw JSON descriptions @@ -106,7 +100,7 @@ pub async fn update_server( .bind(timestamp) .bind(server.players.online) .bind(server.players.max) - .execute(&mut *conn) + .execute(&*conn) .await?; if let Some(sample) = server.players.sample { @@ -121,7 +115,7 @@ pub async fn update_server( .bind(player.name) .bind(timestamp) .bind(timestamp) - .execute(&mut *conn) + .execute(&*conn) .await?; } } @@ -136,7 +130,7 @@ pub async fn update_server( .bind(mods.version) .bind(timestamp) .bind(timestamp) - .execute(&mut *conn) + .execute(&*conn) .await?; } } diff --git a/src/scanner.rs b/src/scanner.rs index 0983e1b..8c8ceec 100644 --- a/src/scanner.rs +++ b/src/scanner.rs @@ -1,10 +1,9 @@ use crate::config::Config; +use crate::database; use crate::protocol::MinecraftServer; use crate::response::Server; use crate::utils::RunError; -use crate::{database, utils}; use futures_util::StreamExt; -use indicatif::{ProgressBar, ProgressStyle}; use sqlx::types::ipnet::IpNet; use sqlx::{Pool, Postgres, Row}; use std::fmt::Debug; @@ -12,7 +11,6 @@ use std::sync::Arc; use std::time::{Duration, SystemTime}; use tokio::io::{AsyncBufReadExt, BufReader}; use tokio::process::Command; -use tokio::sync::Mutex; use tokio::sync::Semaphore; use tracing::{debug, error, info, warn}; @@ -95,36 +93,13 @@ impl Scanner { async fn rescan(&self) { loop { let mut servers = database::fetch_servers(&self.pool).await; - - // Create transaction to update each server - let transaction = Arc::new(Mutex::new( - self.pool - .begin() - .await - .expect("failed to create transaction"), - )); - - // Fetch how many rows are returned from database - let count: i64 = sqlx::query("SELECT count(address) FROM servers") - .fetch_one(&self.pool) - .await - .expect("Failed to get a valid row count from the database!") - .get(0); + let pool = Arc::new(self.pool.clone()); let start_time = match SystemTime::now().duration_since(SystemTime::UNIX_EPOCH) { Ok(n) => n.as_secs(), Err(_) => panic!("SystemTime before UNIX EPOCH!"), }; - 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("=>-"); - - info!("Scanning {count} servers from database"); - - let bar = Arc::new(ProgressBar::new(count as u64).with_style(style)); let mut handles = Vec::new(); // Iterate over results from Postgres as they become available @@ -132,39 +107,12 @@ impl Scanner { let address = row.get::("address").addr().to_string(); let port = row.get::("port") as u16; - let bar = bar.clone(); - let conn = transaction.clone(); - // Wait to acquire permit before spawning a new task let permit = PERMITS.acquire().await; - handles.push(tokio::task::spawn(async move { - let minecraft_server = MinecraftServer::new(address, port); - - async fn run(minecraft_server: MinecraftServer) -> Result { - // In the future there will be a config option to specify the ping type for - // a server, as some servers require the hostname and port to be set. - // Something that this program isn't doing yet. - let response = minecraft_server.simple_ping().await?; - let mut server: Server = serde_json::from_str(&response)?; - server.address = minecraft_server.address; - server.port = minecraft_server.port; - - Ok(server) - } - - match run(minecraft_server).await { - Ok(server) => { - let _ = database::update_server(server, conn).await; - } - Err(error) => debug!("Error occurred while pinging server: {error}"), - } - - bar.inc(1) - })); + handles.push(task_wrapper(address, port, pool.clone())); } futures_util::future::join_all(handles).await; - bar.finish_and_clear(); let end_time = match SystemTime::now().duration_since(SystemTime::UNIX_EPOCH) { Ok(d) => d.as_secs(), @@ -193,13 +141,7 @@ impl Scanner { /// Starts an instance of masscan to find new servers async fn masscan(&self) { loop { - // Create transaction to update each server - let transaction = Arc::new(Mutex::new( - self.pool - .begin() - .await - .expect("failed to create transaction"), - )); + let pool = Arc::new(self.pool.clone()); // Spawn masscan let mut command = Command::new("sudo") @@ -240,38 +182,10 @@ impl Scanner { None => continue, }; - let transaction = transaction.clone(); - - // Ping and update server in database - tokio::task::spawn(async move { - let minecraft_server = MinecraftServer::new(address, port); - - async fn run(minecraft_server: MinecraftServer) -> Result { - let response = minecraft_server.simple_ping().await?; - let mut server: Server = serde_json::from_str(&response)?; - server.address = minecraft_server.address; - server.port = minecraft_server.port; - - Ok(server) - } - - match run(minecraft_server).await { - Ok(server) => { - let _ = database::update_server(server, transaction).await; - } - Err(_) => (), - } - }); + // Spawn a pinging task for each server found + task_wrapper(address, port, pool.clone()); } - // Commit transaction to database - Arc::try_unwrap(transaction) - .unwrap() - .into_inner() - .commit() - .await - .expect("error while commiting to database"); - // Quit if only one scan is requested in config if !self.config.scanner.repeat { info!("Exiting"); @@ -289,3 +203,34 @@ impl Scanner { } } } + +fn task_wrapper( + address: String, + port: u16, + conn: Arc>, +) -> tokio::task::JoinHandle<()> { + tokio::task::spawn(async move { + let minecraft_server = MinecraftServer::new(address, port); + + async fn run(minecraft_server: MinecraftServer) -> Result { + // In the future there will be a config option to specify the ping type for + // a server, as some servers require the hostname and port to be set. + // Something that this program isn't doing yet. + let response = minecraft_server.simple_ping().await?; + + // Assign address and port to the server struct + let mut server: Server = serde_json::from_str(&response)?; + server.address = minecraft_server.address; + server.port = minecraft_server.port; + + Ok(server) + } + + match run(minecraft_server).await { + Ok(server) => { + let _ = database::update_server(server, conn).await; + } + Err(error) => debug!("Error occurred while pinging server: {error}"), + } + }) +} diff --git a/src/utils.rs b/src/utils.rs index bd8f26b..844a1a0 100644 --- a/src/utils.rs +++ b/src/utils.rs @@ -14,7 +14,10 @@ pub enum RunError { TimedOut(#[from] tokio::time::error::Elapsed), #[error("Server opted out of scanning")] ServerOptOut, + #[error("Error while updating server in database")] + DatabaseError(#[from] sqlx::Error), } + impl Into for RunError { fn into(self) -> usize { use RunError::*; @@ -26,6 +29,7 @@ impl Into for RunError { ParseResponse(_) => 3, TimedOut(_) => 4, ServerOptOut => 5, + DatabaseError(_) => 6, } } }