Update src/main.rs

This commit is contained in:
gigirassy
2026-07-14 05:06:18 +02:00
parent b29df807fb
commit 3e1f2a1717
+82 -27
View File
@@ -2,17 +2,21 @@
use mimalloc::MiMalloc; use mimalloc::MiMalloc;
use std::{ use std::{
collections::HashMap,
env, env,
hash::{Hash, Hasher}, hash::{Hash, Hasher},
io, io,
net::SocketAddr, net::{IpAddr, SocketAddr},
sync::Arc, sync::{
time::{SystemTime, UNIX_EPOCH}, atomic::{AtomicUsize, Ordering},
Arc,
},
time::{Duration as StdDuration, SystemTime, UNIX_EPOCH},
}; };
use tokio::{ use tokio::{
io::AsyncWriteExt, io::AsyncWriteExt,
net::{TcpListener, TcpStream}, net::{TcpListener, TcpStream},
sync::{OwnedSemaphorePermit, Semaphore}, sync::{Mutex, OwnedSemaphorePermit, Semaphore},
time::{self, Duration, Instant, MissedTickBehavior}, time::{self, Duration, Instant, MissedTickBehavior},
}; };
@@ -53,17 +57,15 @@ impl Config {
)); ));
} }
if self.max_connections == 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"TARPIT_MAX_CONNECTIONS must be > 0",
));
}
Ok(()) Ok(())
} }
} }
struct State {
connections: AtomicUsize,
totals: Mutex<HashMap<IpAddr, StdDuration>>,
}
#[derive(Clone, Copy)] #[derive(Clone, Copy)]
struct XorShift64 { struct XorShift64 {
state: u64, state: u64,
@@ -71,9 +73,12 @@ struct XorShift64 {
impl XorShift64 { impl XorShift64 {
fn new(seed: u64) -> Self { fn new(seed: u64) -> Self {
// Avoid the all-zero state.
Self { Self {
state: if seed == 0 { 0x9e3779b97f4a7c15 } else { seed }, state: if seed == 0 {
0x9e3779b97f4a7c15
} else {
seed
},
} }
} }
@@ -86,9 +91,8 @@ impl XorShift64 {
x.wrapping_mul(0x2545F4914F6CDD1D) x.wrapping_mul(0x2545F4914F6CDD1D)
} }
fn gen_range_inclusive(&mut self, start: usize, end_inclusive: usize) -> usize { fn gen_range_inclusive(&mut self, start: usize, end: usize) -> usize {
debug_assert!(start <= end_inclusive); let span = (end - start + 1) as u64;
let span = (end_inclusive - start + 1) as u64;
(start as u64 + (self.next_u64() % span)) as usize (start as u64 + (self.next_u64() % span)) as usize
} }
} }
@@ -108,12 +112,18 @@ fn seed_for_peer(peer: SocketAddr) -> u64 {
#[tokio::main] #[tokio::main]
async fn main() -> io::Result<()> { async fn main() -> io::Result<()> {
let bind_addr = env::var("TARPIT_BIND").unwrap_or_else(|_| "0.0.0.0:2222".to_string()); let bind_addr = env::var("TARPIT_BIND").unwrap_or_else(|_| "0.0.0.0:2222".to_string());
let cfg = Config::from_env(); let cfg = Config::from_env();
cfg.validate()?; cfg.validate()?;
let listener = TcpListener::bind(&bind_addr).await?; let listener = TcpListener::bind(&bind_addr).await?;
let limiter = Arc::new(Semaphore::new(cfg.max_connections)); let limiter = Arc::new(Semaphore::new(cfg.max_connections));
let state = Arc::new(State {
connections: AtomicUsize::new(0),
totals: Mutex::new(HashMap::new()),
});
eprintln!( eprintln!(
"listening on {bind_addr}, delay={}ms, line={}..{} bytes, max_conn={}", "listening on {bind_addr}, delay={}ms, line={}..{} bytes, max_conn={}",
cfg.delay.as_millis(), cfg.delay.as_millis(),
@@ -122,12 +132,13 @@ async fn main() -> io::Result<()> {
cfg.max_connections cfg.max_connections
); );
accept_loop(listener, limiter, cfg).await accept_loop(listener, limiter, state, cfg).await
} }
async fn accept_loop( async fn accept_loop(
listener: TcpListener, listener: TcpListener,
limiter: Arc<Semaphore>, limiter: Arc<Semaphore>,
state: Arc<State>,
cfg: Config, cfg: Config,
) -> io::Result<()> { ) -> io::Result<()> {
let shutdown = shutdown_signal(); let shutdown = shutdown_signal();
@@ -154,8 +165,10 @@ async fn accept_loop(
Err(_) => break, Err(_) => break,
}; };
let state = state.clone();
tokio::spawn(async move { tokio::spawn(async move {
if let Err(e) = handle_client(stream, peer, cfg, permit).await { if let Err(e) = handle_client(stream, peer, state, cfg, permit).await {
eprintln!("{peer}: {e}"); eprintln!("{peer}: {e}");
} }
}); });
@@ -168,28 +181,70 @@ async fn accept_loop(
async fn handle_client( async fn handle_client(
mut stream: TcpStream, mut stream: TcpStream,
_peer: SocketAddr, peer: SocketAddr,
state: Arc<State>,
cfg: Config, cfg: Config,
_permit: OwnedSemaphorePermit, _permit: OwnedSemaphorePermit,
) -> io::Result<()> { ) -> io::Result<()> {
let _ = stream.set_nodelay(true); let _ = stream.set_nodelay(true);
let mut rng = XorShift64::new(seed_for_peer(_peer)); let ip = peer.ip();
let current = state.connections.fetch_add(1, Ordering::SeqCst) + 1;
eprintln!(
"[+] {} connected ({} active)",
peer,
current,
);
let connected_at = Instant::now();
let mut rng = XorShift64::new(seed_for_peer(peer));
let mut buf = vec![0u8; cfg.max_line_len + 2]; let mut buf = vec![0u8; cfg.max_line_len + 2];
let mut interval = time::interval_at(Instant::now() + cfg.delay, cfg.delay); let mut interval = time::interval_at(Instant::now() + cfg.delay, cfg.delay);
interval.set_missed_tick_behavior(MissedTickBehavior::Delay); interval.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop { let result = async {
interval.tick().await; loop {
interval.tick().await;
let line_len = rng.gen_range_inclusive(cfg.min_line_len, cfg.max_line_len); let line_len =
fill_line(&mut buf[..line_len], &mut rng); rng.gen_range_inclusive(cfg.min_line_len, cfg.max_line_len);
buf[line_len] = b'\r';
buf[line_len + 1] = b'\n';
stream.write_all(&buf[..line_len + 2]).await?; fill_line(&mut buf[..line_len], &mut rng);
buf[line_len] = b'\r';
buf[line_len + 1] = b'\n';
stream.write_all(&buf[..line_len + 2]).await?;
}
#[allow(unreachable_code)]
Ok::<(), io::Error>(())
} }
.await;
let elapsed = connected_at.elapsed();
let total = {
let mut totals = state.totals.lock().await;
let entry = totals.entry(ip).or_insert(StdDuration::ZERO);
*entry += elapsed;
*entry
};
let current = state.connections.fetch_sub(1, Ordering::SeqCst) - 1;
eprintln!(
"[-] {} disconnected ({} active) session={:.1?} total={:.1?}",
peer,
current,
elapsed,
total
);
result
} }
fn fill_line(dst: &mut [u8], rng: &mut XorShift64) { fn fill_line(dst: &mut [u8], rng: &mut XorShift64) {