feature(database): Replace transaction with a Pool<Postgres> allows updating as servers are found rather than at the end of each scan

This commit is contained in:
funtimes909
2025-05-23 14:06:04 +12:00
parent dd26eafbb1
commit 90400bdc47
3 changed files with 54 additions and 111 deletions
+13 -19
View File
@@ -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<Postgres>) -> BoxStream<Result<PgRow, sql
/// Deletes a server from the database
pub async fn delete_server(
address: String,
transaction: &mut PgConnection,
conn: Arc<Pool<Postgres>>,
) -> Result<PgQueryResult, sqlx::Error> {
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<Mutex<Transaction<'_, Postgres>>>,
) -> anyhow::Result<()> {
let conn = &mut **transaction.lock().await;
pub async fn update_server(server: Server, conn: Arc<Pool<Postgres>>) -> 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?;
}
}
+37 -92
View File
@@ -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::<IpNet, _>("address").addr().to_string();
let port = row.get::<i32, _>("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<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 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<Server, RunError> {
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<Pool<Postgres>>,
) -> tokio::task::JoinHandle<()> {
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?;
// 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}"),
}
})
}
+4
View File
@@ -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<usize> for RunError {
fn into(self) -> usize {
use RunError::*;
@@ -26,6 +29,7 @@ impl Into<usize> for RunError {
ParseResponse(_) => 3,
TimedOut(_) => 4,
ServerOptOut => 5,
DatabaseError(_) => 6,
}
}
}