use anyhow::{Context, Result};
use deadpool::managed::Object;
use crate::pool::deadpool_connection::TcpManager;
use crate::pool::{BufferPool, DeadpoolConnectionProvider, PooledBuffer};
use crate::protocol::{RequestContext, article_request, body_request, head_request, stat_request};
use crate::session::backend::send_request;
#[derive(Clone)]
pub struct NntpClient {
conn_pool: DeadpoolConnectionProvider,
buffer_pool: BufferPool,
}
impl NntpClient {
#[must_use]
pub const fn new(conn_pool: DeadpoolConnectionProvider, buffer_pool: BufferPool) -> Self {
Self {
conn_pool,
buffer_pool,
}
}
#[inline]
pub fn fetch_body(
&self,
message_id: &crate::types::MessageId<'_>,
) -> impl std::future::Future<Output = Result<PooledBuffer>> + '_ {
self.fetch_response(body_request(message_id))
}
#[inline]
pub fn fetch_head(
&self,
message_id: &crate::types::MessageId<'_>,
) -> impl std::future::Future<Output = Result<PooledBuffer>> + '_ {
self.fetch_response(head_request(message_id))
}
#[inline]
pub fn fetch_article(
&self,
message_id: &crate::types::MessageId<'_>,
) -> impl std::future::Future<Output = Result<PooledBuffer>> + '_ {
self.fetch_response(article_request(message_id))
}
pub async fn stat(&self, message_id: &crate::types::MessageId<'_>) -> Result<bool> {
let request = stat_request(message_id);
let mut conn = self
.conn_pool
.get_pooled_connection()
.await
.context("Failed to get connection from pool")?;
let mut buffer = self.buffer_pool.acquire();
let response = send_request(&mut *conn, &request, &mut buffer).await?;
let Some(status_code) = response.status_code() else {
anyhow::bail!("Invalid STAT response");
};
Self::parse_stat_response(status_code)
}
#[inline]
fn parse_stat_response(status_code: crate::protocol::StatusCode) -> Result<bool> {
match status_code.as_u16() {
223 => Ok(true), 430 => Ok(false), code => anyhow::bail!("Unexpected STAT response: {code}"),
}
}
async fn fetch_response(&self, request: RequestContext) -> Result<PooledBuffer> {
let mut conn = self
.conn_pool
.get_pooled_connection()
.await
.context("Failed to get connection from pool")?;
let mut io_buffer = self.buffer_pool.acquire();
let response = send_request(&mut *conn, &request, &mut io_buffer).await?;
let Some(status_code) = response.status_code() else {
anyhow::bail!("Invalid response from server");
};
Self::validate_response(status_code)?;
if request.has_response_body(status_code) {
return self
.fetch_captured_multiline_response(conn, io_buffer)
.await;
}
Ok(io_buffer)
}
async fn fetch_captured_multiline_response(
&self,
mut conn: Object<TcpManager>,
mut io_buffer: PooledBuffer,
) -> Result<PooledBuffer> {
let mut capture = self.buffer_pool.acquire_capture();
if let Err(err) = crate::session::backend::capture_complete_multiline_response(
&mut conn,
&mut io_buffer,
&mut capture,
)
.await
{
self.conn_pool.remove_with_cooldown(conn);
return Err(err);
}
Ok(capture)
}
#[inline]
fn validate_response(status_code: crate::protocol::StatusCode) -> Result<()> {
match status_code.as_u16() {
430 => anyhow::bail!("Article not found (430)"),
code if code >= 400 => anyhow::bail!("Server error: {code}"),
_ => Ok(()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_stat_response_success() {
use crate::protocol::StatusCode;
assert!(NntpClient::parse_stat_response(StatusCode::parse(b"223").unwrap()).unwrap());
assert!(!NntpClient::parse_stat_response(StatusCode::parse(b"430").unwrap()).unwrap());
}
#[test]
fn test_parse_stat_response_errors() {
use crate::protocol::StatusCode;
assert!(NntpClient::parse_stat_response(StatusCode::parse(b"500").unwrap()).is_err());
assert!(NntpClient::parse_stat_response(StatusCode::parse(b"200").unwrap()).is_err());
assert!(NntpClient::parse_stat_response(StatusCode::parse(b"400").unwrap()).is_err());
}
async fn spawn_fetch_test_server(
expected_command: &'static str,
response: &'static [u8],
) -> std::net::SocketAddr {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
loop {
if let Ok((mut stream, _)) = listener.accept().await {
tokio::spawn(async move {
let _ = stream.write_all(b"200 mock\r\n").await;
let mut cmd_buf = [0u8; 1024];
loop {
let Ok(n) = stream.read(&mut cmd_buf).await else {
return;
};
if n == 0 {
return;
}
let command = std::str::from_utf8(&cmd_buf[..n]).unwrap();
if command.starts_with(expected_command) {
let _ = stream.write_all(response).await;
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
return;
}
let _ = stream.write_all(b"200 OK\r\n").await;
}
});
}
}
});
addr
}
async fn spawn_test_server(
article_data: &'static [u8],
) -> (std::net::SocketAddr, std::sync::Arc<tokio::sync::Notify>) {
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
use tokio::net::TcpListener;
use tokio::sync::Notify;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let notify = Arc::new(Notify::new());
let n = Arc::clone(¬ify);
tokio::spawn(async move {
loop {
if let Ok((mut stream, _)) = listener.accept().await {
let wake = Arc::clone(&n);
tokio::spawn(async move {
use tokio::io::AsyncReadExt;
let _ = stream.write_all(b"200 mock\r\n").await;
let mut cmd_buf = vec![0u8; 256];
loop {
tokio::select! {
() = wake.notified() => break,
result = stream.read(&mut cmd_buf) => {
match result {
Ok(n) if n > 0 => { let _ = stream.write_all(b"200 OK\r\n").await; }
_ => break,
}
}
}
}
let _ = stream.write_all(article_data).await;
tokio::time::sleep(std::time::Duration::from_secs(30)).await;
});
}
}
});
(addr, notify)
}
async fn spawn_truncated_test_server(
article_prefix: &'static [u8],
) -> (std::net::SocketAddr, std::sync::Arc<tokio::sync::Notify>) {
use std::sync::Arc;
use tokio::io::AsyncWriteExt;
use tokio::net::TcpListener;
use tokio::sync::Notify;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let notify = Arc::new(Notify::new());
let n = Arc::clone(¬ify);
tokio::spawn(async move {
loop {
if let Ok((mut stream, _)) = listener.accept().await {
let wake = Arc::clone(&n);
tokio::spawn(async move {
use tokio::io::AsyncReadExt;
let _ = stream.write_all(b"200 mock\r\n").await;
let mut cmd_buf = vec![0u8; 256];
loop {
tokio::select! {
() = wake.notified() => break,
result = stream.read(&mut cmd_buf) => {
match result {
Ok(n) if n > 0 => { let _ = stream.write_all(b"200 OK\r\n").await; }
_ => break,
}
}
}
}
let _ = stream.write_all(article_prefix).await;
let _ = stream.shutdown().await;
});
}
}
});
(addr, notify)
}
fn make_test_pool(addr: std::net::SocketAddr) -> crate::pool::deadpool_connection::Pool {
let manager = crate::pool::deadpool_connection::TcpManager::new(
addr.ip().to_string(),
addr.port(),
"test".to_string(),
crate::pool::deadpool_connection::TcpManagerOptions {
compress: Some(false), ..crate::pool::deadpool_connection::TcpManagerOptions::default()
},
)
.unwrap();
crate::pool::deadpool_connection::Pool::builder(manager)
.max_size(2)
.build()
.unwrap()
}
fn make_test_client(addr: std::net::SocketAddr) -> NntpClient {
use crate::pool::BufferPool;
use crate::types::BufferSize;
let provider = DeadpoolConnectionProvider::builder(addr.ip().to_string(), addr.port())
.name("test")
.max_connections(2)
.build()
.unwrap();
let buffer_pool = BufferPool::new(BufferSize::try_new(4096).unwrap(), 2);
NntpClient::new(provider, buffer_pool)
}
async fn capture_multiline_response_for_test(
conn: &mut crate::stream::ConnectionStream,
io_buffer: &mut PooledBuffer,
capture: &mut PooledBuffer,
) -> Result<()> {
crate::session::backend::capture_complete_multiline_response(conn, io_buffer, capture).await
}
#[tokio::test]
async fn test_multiline_response_capture_single_read() {
use crate::pool::BufferPool;
use crate::types::BufferSize;
let article = b"220 body follows\r\nHello world\r\n.\r\n";
let (addr, notify) = spawn_test_server(article).await;
let pool = make_test_pool(addr);
let buffer_pool = BufferPool::new(BufferSize::try_new(4096).unwrap(), 2);
let mut conn = pool.get().await.unwrap();
notify.notify_one();
let mut io_buffer = buffer_pool.acquire();
let mut capture = buffer_pool.acquire_capture();
io_buffer.read_from(&mut *conn).await.unwrap();
capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
.await
.unwrap();
assert_eq!(&capture[..], article as &[u8]);
}
#[tokio::test]
async fn test_multiline_response_capture_multi_read_spanning_body_end() {
use crate::pool::BufferPool;
use crate::types::BufferSize;
let article = b"220 article\r\nLine one\r\nLine two\r\n.\r\n";
let (addr, notify) = spawn_test_server(article).await;
let pool = make_test_pool(addr);
let buffer_pool = BufferPool::new(BufferSize::try_new(8).unwrap(), 4);
let mut conn = pool.get().await.unwrap();
notify.notify_one();
let mut io_buffer = buffer_pool.acquire();
let mut capture = buffer_pool.acquire_capture();
capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
.await
.unwrap();
assert_eq!(&capture[..], article as &[u8]);
}
#[tokio::test]
async fn test_multiline_response_capture_errors_on_truncated_response() {
use crate::pool::BufferPool;
use crate::types::BufferSize;
let article_prefix = b"220 body follows\r\npartial article";
let (addr, notify) = spawn_truncated_test_server(article_prefix).await;
let pool = make_test_pool(addr);
let buffer_pool = BufferPool::new(BufferSize::try_new(8).unwrap(), 4);
let mut conn = pool.get().await.unwrap();
notify.notify_one();
let mut io_buffer = buffer_pool.acquire();
let mut capture = buffer_pool.acquire_capture();
let err = capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
.await
.unwrap_err();
assert!(
err.to_string()
.contains("Backend closed connection before complete"),
"unexpected error: {err:#}"
);
}
#[tokio::test]
async fn test_multiline_response_capture_errors_on_extra_response_bytes() {
use crate::pool::BufferPool;
use crate::types::BufferSize;
let article = b"220 body follows\r\nHello world\r\n.\r\n";
let extra_response = [article.as_slice(), b"430 No such article\r\n"].concat();
let extra_response: &'static [u8] = Box::leak(extra_response.into_boxed_slice());
let (addr, notify) = spawn_test_server(extra_response).await;
let pool = make_test_pool(addr);
let buffer_pool = BufferPool::new(BufferSize::try_new(4096).unwrap(), 2);
let mut conn = pool.get().await.unwrap();
notify.notify_one();
let mut io_buffer = buffer_pool.acquire();
let mut capture = buffer_pool.acquire_capture();
io_buffer.read_from(&mut *conn).await.unwrap();
let err = capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
.await
.unwrap_err();
assert!(err.to_string().contains("unexpected"));
}
#[tokio::test]
async fn fetch_head_reads_multiline_response() {
let response = b"221 0 <test@example.com>\r\nSubject: test\r\nFrom: tester\r\n\r\n.\r\n";
let addr = spawn_fetch_test_server("HEAD <test@example.com>", response).await;
let client = make_test_client(addr);
let msg_id = crate::types::MessageId::new("<test@example.com>".to_string()).unwrap();
let buffer = client.fetch_head(&msg_id).await.unwrap();
assert_eq!(&buffer[..], response);
}
#[tokio::test]
async fn fetch_body_reads_multiline_response() {
let response = b"222 0 <test@example.com>\r\nhello world\r\n.\r\n";
let addr = spawn_fetch_test_server("BODY <test@example.com>", response).await;
let client = make_test_client(addr);
let msg_id = crate::types::MessageId::new("<test@example.com>".to_string()).unwrap();
let buffer = client.fetch_body(&msg_id).await.unwrap();
assert_eq!(&buffer[..], response);
}
#[tokio::test]
async fn fetch_body_reads_multiline_response_above_retention_limit() {
let mut response = Vec::with_capacity((4 * 1024 * 1024) + 64);
response.extend_from_slice(b"222 0 <large@example.com>\r\n");
response.extend(std::iter::repeat_n(b'x', 4 * 1024 * 1024));
response.extend_from_slice(b"\r\n.\r\n");
let response: &'static [u8] = Box::leak(response.into_boxed_slice());
let addr = spawn_fetch_test_server("BODY <large@example.com>", response).await;
let client = make_test_client(addr);
let msg_id = crate::types::MessageId::new("<large@example.com>".to_string()).unwrap();
let buffer = client.fetch_body(&msg_id).await.unwrap();
assert_eq!(&buffer[..], response);
}
}