From aa11099dfdb6cb377108840fb29cb4957b077e97 Mon Sep 17 00:00:00 2001 From: funtimes909 Date: Thu, 22 May 2025 23:41:30 +1200 Subject: [PATCH] Small refactor to protocol.rs --- src/database.rs | 4 +- src/protocol.rs | 136 +++++++++++++++++++++++++------------------ src/scanner.rs | 152 ++++++++++++++++++++++++++++++------------------ src/utils.rs | 41 ------------- 4 files changed, 178 insertions(+), 155 deletions(-) diff --git a/src/database.rs b/src/database.rs index 8f1689f..a232b2a 100644 --- a/src/database.rs +++ b/src/database.rs @@ -1,8 +1,6 @@ use crate::response::Server; use crate::utils; use futures_util::stream::BoxStream; -use futures_util::{future, FutureExt}; -use indicatif::{ProgressBar, ProgressStyle}; use sqlx::postgres::{PgQueryResult, PgRow}; use sqlx::types::ipnet::IpNet; use sqlx::types::Uuid; @@ -39,7 +37,7 @@ pub async fn update_server( 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 address = IpNet::from_str(&(server.address.clone() + "/32"))?; let timestamp = SystemTime::now().duration_since(UNIX_EPOCH)?.as_secs() as i32; // Handle server descriptions diff --git a/src/protocol.rs b/src/protocol.rs index 8d1aff4..1036c13 100644 --- a/src/protocol.rs +++ b/src/protocol.rs @@ -5,7 +5,7 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpStream; use tracing::debug; -const PAYLOAD: [u8; 9] = [ +const SIMPLE_PAYLOAD: [u8; 9] = [ 6, // Size: Amount of bytes in the message 0, // ID: Has to be 0 0, // Protocol Version: Can be anything as long as it's a valid varint @@ -16,75 +16,99 @@ const PAYLOAD: [u8; 9] = [ 0, // ID ]; -pub async fn ping_server((address, port): (&str, u16)) -> Result { - let socket = SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::from_str(address)?, port)); +#[derive(Debug)] +pub struct MinecraftServer { + pub address: String, + pub port: u16, +} - let mut stream = TcpStream::connect(&socket).await?; - stream.write_all(&PAYLOAD).await?; - let mut response = [0; 1024]; +impl From<(String, u16)> for MinecraftServer { + fn from(value: (String, u16)) -> Self { + MinecraftServer::new(value.0, value.1) + } +} - // The index is used to point to the position at the start of the string. - // It gets increased by the amount of bytes read to decode the packet ID, Packet length - // And string length - let mut index = 0; - - // Returns how many bytes were read from the stream into the buffer - let total_read_bytes = stream.read(&mut response).await?; - - if total_read_bytes == 0 { - debug!("[{address}] Total read bytes is 0"); - return Err(RunError::MalformedResponse); +impl MinecraftServer { + pub fn new(address: String, port: u16) -> Self { + Self { address, port } } - // Decode Packet length - index += decode_varint(&response).1; + pub async fn simple_ping(&self) -> Result { + let socket = SocketAddr::V4(SocketAddrV4::new( + Ipv4Addr::from_str(&self.address)?, + self.port, + )); - // Since Packet ID should always be 0 and will never take more than 1 byte to encode - // We can ignore it entirely and just advance the index by 1 - index += 1; + let mut stream = + tokio::time::timeout(crate::scanner::TIMEOUT_SECS, TcpStream::connect(&socket)) + .await??; + stream.write_all(&SIMPLE_PAYLOAD).await?; + let mut response = [0; 1024]; - // Decode the string length - let (string_length, string_length_bytes) = decode_varint(&response[index as usize..]); - index += string_length_bytes; - if string_length == 0 { - debug!("[{address}] String length is 0"); - return Err(RunError::MalformedResponse); - } + // The index is used to point to the position at the start of the string. + // It gets increased by the amount of bytes read to decode the packet ID, Packet length + // And string length + let mut index = 0; - if string_length > 32767 { - debug!("[{address}] Received abnormally large string length: {string_length}"); - } + // Returns how many bytes were read from the stream into the buffer + let total_read_bytes = stream.read(&mut response).await?; - // Error checking - if index as usize > total_read_bytes { - debug!("[{address}] Index: {index} is bigger than total read bytes: {total_read_bytes}"); - return Err(RunError::MalformedResponse); - } + if total_read_bytes == 0 { + debug!("[{}] Total read bytes is 0", &self.address); + return Err(RunError::MalformedResponse); + } - // WARNING: Don't allocate vec size based on what the server says it needs from the varint. - // Allocate size based on what the server *actually* sends back, some servers can crash the - // program by attempting to allocate insane amounts of memory this way. - // - // Adds everything we have read so far minus the packet ID and packet length to a new vec - let mut output = Vec::from(&response[index as usize..total_read_bytes]); - let string_length = string_length + index as usize; + // Packet length + index += decode_varint(&response).1; - if total_read_bytes > string_length { - debug!( - "[{address}] Total read bytes: {total_read_bytes} is larger than string length: {string_length}" + // Since Packet ID should always be 0 and will never take more than 1 byte to encode + // We can ignore it entirely and just advance the index by 1 + index += 1; + + // Decode the string length + let (string_length, string_length_bytes) = decode_varint(&response[index as usize..]); + index += string_length_bytes; + + // Error checking + if string_length == 0 || string_length > 32767 { + debug!( + "[{}] String length: {string_length} was either 0 or too long", + &self.address + ); + return Err(RunError::MalformedResponse); + } + + // WARNING: Don't allocate vec size based on what the server says it needs from the varint. + // Allocate size based on what the server *actually* sends back, some servers can crash the + // program by attempting to allocate insane amounts of memory this way. + // + // Adds everything we have read so far minus the packet ID and packet length to a new vec + let mut output = Vec::from(&response[index as usize..total_read_bytes]); + let string_length = string_length + index as usize; + + if total_read_bytes > string_length { + debug!( + "[{}] Total read bytes: {total_read_bytes} is larger than string length: {string_length}", &self.address ); - return Err(RunError::MalformedResponse); + return Err(RunError::MalformedResponse); + } + + // Read the rest of the servers JSON + stream + // Takes everything after the end of the data we already have in the buffer + // Up until the end of the strings length + .take((string_length - total_read_bytes) as u64) + .read_to_end(&mut output) + .await?; + + Ok(String::from_utf8_lossy(&output).to_string()) } - // Read the rest of the servers JSON - stream - // Takes everything after the end of the data we already have in the buffer - // Up until the end of the strings length - .take((string_length - total_read_bytes) as u64) - .read_to_end(&mut output) - .await?; + // TODO + pub async fn legacy_ping() {} - Ok(String::from_utf8_lossy(&output).to_string()) + // TODO + pub async fn proper_ping() {} } // returns the decoded varint and how many bytes were read diff --git a/src/scanner.rs b/src/scanner.rs index fd025a3..5c94ff5 100644 --- a/src/scanner.rs +++ b/src/scanner.rs @@ -1,4 +1,7 @@ use crate::config::Config; +use crate::protocol::MinecraftServer; +use crate::response::Server; +use crate::utils::RunError; use crate::{database, utils}; use futures_util::StreamExt; use indicatif::{ProgressBar, ProgressStyle}; @@ -6,12 +9,12 @@ use sqlx::types::ipnet::IpNet; use sqlx::{Pool, Postgres, Row}; use std::fmt::Debug; use std::sync::Arc; -use std::time::Duration; +use std::time::{Duration, SystemTime}; use tokio::io::{AsyncBufReadExt, BufReader}; use tokio::process::Command; use tokio::sync::Mutex; use tokio::sync::Semaphore; -use tracing::{error, info, warn}; +use tracing::{debug, error, info, warn}; pub static PERMITS: Semaphore = Semaphore::const_new(10000); pub const TIMEOUT_SECS: Duration = Duration::from_secs(3); @@ -90,16 +93,6 @@ impl Scanner { /// Rescan servers already found in the database 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(); - - if total_ports > 10 { - warn!("Large amount of ports! Each extra port to scan doubles the total time taken!"); - } - - info!("Scanning port range {port_start} - {port_end} ({total_ports} port(s) per host)"); - loop { let mut servers = database::fetch_servers(&self.pool).await; @@ -118,32 +111,66 @@ impl Scanner { .expect("Failed to get a valid row count from the database!") .get(0); + 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("=>-"); - let bar = Arc::new(ProgressBar::new(count as u64).with_style(style)); - info!("Scanning {count} servers from database"); - let mut handles = Vec::with_capacity(count as usize); + 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 while let Some(Ok(row)) = servers.next().await { let address = row.get::("address").addr().to_string(); let port = row.get::("port") as u16; - handles.push(tokio::task::spawn(utils::run_and_update_with_progress( - address, - port, - transaction.clone(), - bar.clone(), - ))); + 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 server: Server = serde_json::from_str(&response)?; + + 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) + })); } 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(), + Err(_) => panic!("SystemTime before UNIX EPOCH!"), + }; + + info!("Scan completed in {} seconds", end_time - start_time); + // Quit if only one scan is requested in config if !self.config.scanner.repeat { info!("Exiting"); @@ -179,44 +206,59 @@ impl Scanner { .spawn() .expect("error while executing masscan"); - let mut count = 0; - - // Get output from masscan - if let Some(stdout) = command.stdout.take() { - let mut reader = BufReader::new(stdout).lines(); - - // Iterate over the lines of output from masscan - while let Ok(Some(line)) = reader.next_line().await { - let mut line = line.split_whitespace(); - - let port = match line - .nth(3) - // Split on port/tcp - .and_then(|p| p.split('/').nth(0)) - // Parse as u16 - .and_then(|s| s.parse::().ok()) - { - Some(port) => port, - None => continue, - }; - - // .nth() consumes all preceding elements so address will be the 2nd - let address = match line.nth(1) { - Some(address) => address.to_owned(), - None => continue, - }; - - // Ping and update server in database - tokio::task::spawn(utils::run_and_update(address, port, transaction.clone())); - count += 1; + // Verify stdout is valid + let stdout = match command.stdout.take() { + Some(o) => o, + None => { + error!("Failed to get stdout from masscan!"); + std::process::exit(1); } - } else { - error!("Failed to get stdout from masscan!"); - std::process::exit(1); + }; + + let mut reader = BufReader::new(stdout).lines(); + + // Iterate over the lines of output from masscan + while let Ok(Some(line)) = reader.next_line().await { + let mut line = line.split_whitespace(); + + let port = match line + .nth(3) + // Split on port/tcp + .and_then(|p| p.split('/').nth(0)) + // Parse as u16 + .and_then(|s| s.parse::().ok()) + { + Some(port) => port, + None => continue, + }; + + // .nth() consumes all preceding elements so address will be the 2nd + let address = match line.nth(1) { + Some(address) => address.to_owned(), + 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 server: Server = serde_json::from_str(&response)?; + + Ok(server) + } + + match run(minecraft_server).await { + Ok(server) => database::update_server(server, transaction).await.unwrap(), + Err(_) => (), + } + }); } // Commit transaction to database - info!("Inserting {count} servers to database"); Arc::try_unwrap(transaction) .unwrap() .into_inner() diff --git a/src/utils.rs b/src/utils.rs index 8fc08f0..bd8f26b 100644 --- a/src/utils.rs +++ b/src/utils.rs @@ -1,11 +1,4 @@ -use crate::response::Server; -use crate::scanner::{PERMITS, TIMEOUT_SECS}; -use crate::{database, protocol}; -use indicatif::ProgressBar; -use sqlx::{Postgres, Transaction}; -use std::sync::Arc; use thiserror::Error; -use tokio::sync::Mutex; #[derive(Debug, Error)] pub enum RunError { @@ -113,37 +106,3 @@ impl MinecraftColorCodes { } } } - -pub async fn run_and_update( - address: String, - port: u16, - conn: Arc>>, -) -> Result<(), RunError> { - // Ping server - let permit = PERMITS - .acquire() - .await - .expect("failed to acquire a semaphore"); - - let pinged_server = - tokio::time::timeout(TIMEOUT_SECS, protocol::ping_server((&*address, port))).await??; - drop(permit); - - // Parse response - let mut server = serde_json::from_str::(&pinged_server)?; - server.address = address; - server.port = port; - - let _ = database::update_server(server, conn).await; - Ok(()) -} - -pub async fn run_and_update_with_progress( - address: String, - port: u16, - conn: Arc>>, - bar: Arc, -) { - let _ = run_and_update(address, port, conn).await; - bar.inc(1); -}