use anyhow::Result;
use nntp_proxy::config::Server;
use nntp_proxy::types::{MaxConnections, Port};
use nntp_proxy::{Config, NntpProxy, RoutingMode};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
use tokio::net::{TcpListener, TcpStream};
use tokio::time::Duration;
use crate::test_helpers::wait_for_server;
struct CompressionMockServer {
compress_received: Arc<AtomicBool>,
compress_response: String,
}
impl CompressionMockServer {
fn new(compress_response: &str) -> Self {
Self {
compress_received: Arc::new(AtomicBool::new(false)),
compress_response: compress_response.to_string(),
}
}
fn compress_was_received(&self) -> Arc<AtomicBool> {
Arc::clone(&self.compress_received)
}
async fn run_on_listener(
listener: TcpListener,
compress_received: Arc<AtomicBool>,
compress_response: String,
) {
while let Ok((mut stream, _)) = listener.accept().await {
let compress_received = compress_received.clone();
let compress_response = compress_response.clone();
tokio::spawn(async move {
if stream
.write_all(b"200 CompressionTest Ready\r\n")
.await
.is_err()
{
return;
}
let mut buffer = [0u8; 4096];
loop {
let n = match stream.read(&mut buffer).await {
Ok(0) | Err(_) => break,
Ok(n) => n,
};
let cmd = String::from_utf8_lossy(&buffer[..n]);
let cmd_upper = cmd.trim().to_uppercase();
if cmd_upper.starts_with("COMPRESS") {
compress_received.store(true, Ordering::SeqCst);
let _ = stream.write_all(compress_response.as_bytes()).await;
} else if cmd_upper.starts_with("QUIT") {
let _ = stream.write_all(b"205 Goodbye\r\n").await;
break;
} else if cmd_upper.starts_with("STAT") {
let _ = stream
.write_all(b"223 0 <test@example.com> exists\r\n")
.await;
} else {
let _ = stream.write_all(b"200 OK\r\n").await;
}
}
});
}
}
fn spawn_on_listener(self, listener: TcpListener) -> tokio::task::AbortHandle {
let compress_received = self.compress_received;
let compress_response = self.compress_response;
tokio::spawn(Self::run_on_listener(
listener,
compress_received,
compress_response,
))
.abort_handle()
}
}
fn build_server_config(port: u16, compress: Option<bool>) -> Server {
Server::builder("127.0.0.1", Port::try_new(port).unwrap())
.name("CompressionTest")
.max_connections(MaxConnections::try_new(5).unwrap())
.compress(compress)
.build()
.expect("Valid server config")
}
#[tokio::test]
async fn test_compression_disabled_skips_negotiation() -> Result<()> {
let mock_listener = TcpListener::bind("127.0.0.1:0").await?;
let mock_port = mock_listener.local_addr()?.port();
let proxy_listener = TcpListener::bind("127.0.0.1:0").await?;
let proxy_port = proxy_listener.local_addr()?.port();
let mock = CompressionMockServer::new("206 Compression active\r\n");
let compress_received = mock.compress_was_received();
let _handle = mock.spawn_on_listener(mock_listener);
wait_for_server(&format!("127.0.0.1:{mock_port}"), 20).await?;
let config = Config {
servers: vec![build_server_config(mock_port, Some(false))],
..Default::default()
};
let proxy = NntpProxy::new(config, RoutingMode::PerCommand).await?;
let proxy_addr = format!("127.0.0.1:{proxy_port}");
tokio::spawn({
let proxy = proxy.clone();
async move {
while let Ok((stream, addr)) = proxy_listener.accept().await {
let p = proxy.clone();
tokio::spawn(async move {
let _ = p
.handle_client_per_command_routing(stream, addr.into())
.await;
});
}
}
});
wait_for_server(&proxy_addr, 20).await?;
let mut client = TcpStream::connect(&proxy_addr).await?;
let mut reader = BufReader::new(&mut client);
let mut line = String::new();
reader.read_line(&mut line).await?;
assert!(line.starts_with("201"));
line.clear();
reader.get_mut().write_all(b"STAT <test@test>\r\n").await?;
reader.read_line(&mut line).await?;
assert!(line.starts_with("223"), "Expected 223, got: {line}");
assert!(
!compress_received.load(Ordering::SeqCst),
"COMPRESS DEFLATE should not be sent when compression is disabled"
);
reader.get_mut().write_all(b"QUIT\r\n").await?;
Ok(())
}
#[tokio::test]
async fn test_compression_auto_fallback_on_unsupported() -> Result<()> {
let mock_listener = TcpListener::bind("127.0.0.1:0").await?;
let mock_port = mock_listener.local_addr()?.port();
let proxy_listener = TcpListener::bind("127.0.0.1:0").await?;
let proxy_port = proxy_listener.local_addr()?.port();
let mock = CompressionMockServer::new("500 Command not recognized\r\n");
let compress_received = mock.compress_was_received();
let _handle = mock.spawn_on_listener(mock_listener);
wait_for_server(&format!("127.0.0.1:{mock_port}"), 20).await?;
let config = Config {
servers: vec![build_server_config(mock_port, None)],
..Default::default()
};
let proxy = NntpProxy::new(config, RoutingMode::PerCommand).await?;
let proxy_addr = format!("127.0.0.1:{proxy_port}");
tokio::spawn({
let proxy = proxy.clone();
async move {
while let Ok((stream, addr)) = proxy_listener.accept().await {
let p = proxy.clone();
tokio::spawn(async move {
let _ = p
.handle_client_per_command_routing(stream, addr.into())
.await;
});
}
}
});
wait_for_server(&proxy_addr, 20).await?;
let mut client = TcpStream::connect(&proxy_addr).await?;
let mut reader = BufReader::new(&mut client);
let mut line = String::new();
reader.read_line(&mut line).await?;
assert!(line.starts_with("201"));
line.clear();
reader.get_mut().write_all(b"STAT <test@test>\r\n").await?;
reader.read_line(&mut line).await?;
assert!(line.starts_with("223"), "Expected 223, got: {line}");
assert!(
compress_received.load(Ordering::SeqCst),
"COMPRESS DEFLATE should be sent in auto mode"
);
reader.get_mut().write_all(b"QUIT\r\n").await?;
Ok(())
}
#[tokio::test]
async fn test_compression_required_fails_on_unsupported() -> Result<()> {
let mock_listener = TcpListener::bind("127.0.0.1:0").await?;
let mock_port = mock_listener.local_addr()?.port();
let proxy_listener = TcpListener::bind("127.0.0.1:0").await?;
let proxy_port = proxy_listener.local_addr()?.port();
let mock = CompressionMockServer::new("500 Command not recognized\r\n");
let _handle = mock.spawn_on_listener(mock_listener);
wait_for_server(&format!("127.0.0.1:{mock_port}"), 20).await?;
let config = Config {
servers: vec![build_server_config(mock_port, Some(true))],
..Default::default()
};
let proxy = NntpProxy::new(config, RoutingMode::PerCommand).await?;
let proxy_addr = format!("127.0.0.1:{proxy_port}");
tokio::spawn({
let proxy = proxy.clone();
async move {
while let Ok((stream, addr)) = proxy_listener.accept().await {
let p = proxy.clone();
tokio::spawn(async move {
let _ = p
.handle_client_per_command_routing(stream, addr.into())
.await;
});
}
}
});
wait_for_server(&proxy_addr, 20).await?;
let mut client = TcpStream::connect(&proxy_addr).await?;
let mut reader = BufReader::new(&mut client);
let mut line = String::new();
reader.read_line(&mut line).await?;
assert!(line.starts_with("201"));
line.clear();
reader.get_mut().write_all(b"STAT <test@test>\r\n").await?;
let result = tokio::time::timeout(Duration::from_secs(2), reader.read_line(&mut line)).await;
match result {
Ok(Ok(n)) if n > 0 => {
assert!(
!line.starts_with("223"),
"Should not get successful response when compression is required but unsupported, got: {line}"
);
}
Ok(Ok(_) | Err(_)) | Err(_) => {}
}
Ok(())
}