Small refactor to protocol.rs

This commit is contained in:
funtimes909
2025-05-22 23:41:30 +12:00
parent b6ddf9d4c2
commit aa11099dfd
4 changed files with 178 additions and 155 deletions
+1 -3
View File
@@ -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
+80 -56
View File
@@ -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<String, RunError> {
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<String, RunError> {
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
+97 -55
View File
@@ -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::<IpNet, _>("address").addr().to_string();
let port = row.get::<i32, _>("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<Server, RunError> {
// 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::<u16>().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::<u16>().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<Server, RunError> {
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()
-41
View File
@@ -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<Mutex<Transaction<'_, Postgres>>>,
) -> 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::<Server>(&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<Mutex<Transaction<'_, Postgres>>>,
bar: Arc<ProgressBar>,
) {
let _ = run_and_update(address, port, conn).await;
bar.inc(1);
}