1use anyhow::{Context, Result};
41use deadpool::managed::Object;
42
43use crate::pool::deadpool_connection::TcpManager;
44use crate::pool::{BufferPool, DeadpoolConnectionProvider, PooledBuffer};
45use crate::protocol::{RequestContext, article_request, body_request, head_request, stat_request};
46use crate::session::backend::send_request;
47
48#[derive(Clone)]
54pub struct NntpClient {
55 conn_pool: DeadpoolConnectionProvider,
56 buffer_pool: BufferPool,
57}
58
59impl NntpClient {
60 #[must_use]
64 pub const fn new(conn_pool: DeadpoolConnectionProvider, buffer_pool: BufferPool) -> Self {
65 Self {
66 conn_pool,
67 buffer_pool,
68 }
69 }
70
71 #[inline]
83 pub fn fetch_body(
84 &self,
85 message_id: &crate::types::MessageId<'_>,
86 ) -> impl std::future::Future<Output = Result<PooledBuffer>> + '_ {
87 self.fetch_response(body_request(message_id))
88 }
89
90 #[inline]
102 pub fn fetch_head(
103 &self,
104 message_id: &crate::types::MessageId<'_>,
105 ) -> impl std::future::Future<Output = Result<PooledBuffer>> + '_ {
106 self.fetch_response(head_request(message_id))
107 }
108
109 #[inline]
121 pub fn fetch_article(
122 &self,
123 message_id: &crate::types::MessageId<'_>,
124 ) -> impl std::future::Future<Output = Result<PooledBuffer>> + '_ {
125 self.fetch_response(article_request(message_id))
126 }
127
128 pub async fn stat(&self, message_id: &crate::types::MessageId<'_>) -> Result<bool> {
140 let request = stat_request(message_id);
141 let mut conn = self
142 .conn_pool
143 .get_pooled_connection()
144 .await
145 .context("Failed to get connection from pool")?;
146 let mut buffer = self.buffer_pool.acquire();
147
148 let response = send_request(&mut *conn, &request, &mut buffer).await?;
149 let Some(status_code) = response.status_code() else {
150 anyhow::bail!("Invalid STAT response");
151 };
152
153 Self::parse_stat_response(status_code)
154 }
155
156 #[inline]
158 fn parse_stat_response(status_code: crate::protocol::StatusCode) -> Result<bool> {
159 match status_code.as_u16() {
160 223 => Ok(true), 430 => Ok(false), code => anyhow::bail!("Unexpected STAT response: {code}"),
163 }
164 }
165
166 async fn fetch_response(&self, request: RequestContext) -> Result<PooledBuffer> {
172 let mut conn = self
173 .conn_pool
174 .get_pooled_connection()
175 .await
176 .context("Failed to get connection from pool")?;
177 let mut io_buffer = self.buffer_pool.acquire();
178
179 let response = send_request(&mut *conn, &request, &mut io_buffer).await?;
180 let Some(status_code) = response.status_code() else {
181 anyhow::bail!("Invalid response from server");
182 };
183
184 Self::validate_response(status_code)?;
185
186 if request.has_response_body(status_code) {
187 return self
188 .fetch_captured_multiline_response(conn, io_buffer)
189 .await;
190 }
191
192 Ok(io_buffer)
193 }
194
195 async fn fetch_captured_multiline_response(
196 &self,
197 mut conn: Object<TcpManager>,
198 mut io_buffer: PooledBuffer,
199 ) -> Result<PooledBuffer> {
200 let mut capture = self.buffer_pool.acquire_capture();
204 if let Err(err) = crate::session::backend::capture_complete_multiline_response(
205 &mut conn,
206 &mut io_buffer,
207 &mut capture,
208 )
209 .await
210 {
211 self.conn_pool.remove_with_cooldown(conn);
212 return Err(err);
213 }
214 Ok(capture)
215 }
216
217 #[inline]
219 fn validate_response(status_code: crate::protocol::StatusCode) -> Result<()> {
220 match status_code.as_u16() {
221 430 => anyhow::bail!("Article not found (430)"),
222 code if code >= 400 => anyhow::bail!("Server error: {code}"),
223 _ => Ok(()),
224 }
225 }
226}
227
228#[cfg(test)]
229mod tests {
230 use super::*;
231
232 #[test]
233 fn test_parse_stat_response_success() {
234 use crate::protocol::StatusCode;
235 assert!(NntpClient::parse_stat_response(StatusCode::parse(b"223").unwrap()).unwrap());
237
238 assert!(!NntpClient::parse_stat_response(StatusCode::parse(b"430").unwrap()).unwrap());
240 }
241
242 #[test]
243 fn test_parse_stat_response_errors() {
244 use crate::protocol::StatusCode;
245 assert!(NntpClient::parse_stat_response(StatusCode::parse(b"500").unwrap()).is_err());
247 assert!(NntpClient::parse_stat_response(StatusCode::parse(b"200").unwrap()).is_err());
248 assert!(NntpClient::parse_stat_response(StatusCode::parse(b"400").unwrap()).is_err());
249 }
250
251 async fn spawn_fetch_test_server(
252 expected_command: &'static str,
253 response: &'static [u8],
254 ) -> std::net::SocketAddr {
255 use tokio::io::{AsyncReadExt, AsyncWriteExt};
256 use tokio::net::TcpListener;
257
258 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
259 let addr = listener.local_addr().unwrap();
260
261 tokio::spawn(async move {
262 loop {
263 if let Ok((mut stream, _)) = listener.accept().await {
264 tokio::spawn(async move {
265 let _ = stream.write_all(b"200 mock\r\n").await;
266 let mut cmd_buf = [0u8; 1024];
267 loop {
268 let Ok(n) = stream.read(&mut cmd_buf).await else {
269 return;
270 };
271 if n == 0 {
272 return;
273 }
274
275 let command = std::str::from_utf8(&cmd_buf[..n]).unwrap();
276 if command.starts_with(expected_command) {
277 let _ = stream.write_all(response).await;
278 tokio::time::sleep(std::time::Duration::from_secs(30)).await;
279 return;
280 }
281
282 let _ = stream.write_all(b"200 OK\r\n").await;
283 }
284 });
285 }
286 }
287 });
288
289 addr
290 }
291
292 async fn spawn_test_server(
300 article_data: &'static [u8],
301 ) -> (std::net::SocketAddr, std::sync::Arc<tokio::sync::Notify>) {
302 use std::sync::Arc;
303 use tokio::io::AsyncWriteExt;
304 use tokio::net::TcpListener;
305 use tokio::sync::Notify;
306
307 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
308 let addr = listener.local_addr().unwrap();
309 let notify = Arc::new(Notify::new());
310 let n = Arc::clone(¬ify);
311
312 tokio::spawn(async move {
313 loop {
314 if let Ok((mut stream, _)) = listener.accept().await {
315 let wake = Arc::clone(&n);
316 tokio::spawn(async move {
317 use tokio::io::AsyncReadExt;
318 let _ = stream.write_all(b"200 mock\r\n").await;
319 let mut cmd_buf = vec![0u8; 256];
322 loop {
323 tokio::select! {
324 () = wake.notified() => break,
325 result = stream.read(&mut cmd_buf) => {
326 match result {
327 Ok(n) if n > 0 => { let _ = stream.write_all(b"200 OK\r\n").await; }
328 _ => break,
329 }
330 }
331 }
332 }
333 let _ = stream.write_all(article_data).await;
334 tokio::time::sleep(std::time::Duration::from_secs(30)).await;
336 });
337 }
338 }
339 });
340
341 (addr, notify)
342 }
343
344 async fn spawn_truncated_test_server(
345 article_prefix: &'static [u8],
346 ) -> (std::net::SocketAddr, std::sync::Arc<tokio::sync::Notify>) {
347 use std::sync::Arc;
348 use tokio::io::AsyncWriteExt;
349 use tokio::net::TcpListener;
350 use tokio::sync::Notify;
351
352 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
353 let addr = listener.local_addr().unwrap();
354 let notify = Arc::new(Notify::new());
355 let n = Arc::clone(¬ify);
356
357 tokio::spawn(async move {
358 loop {
359 if let Ok((mut stream, _)) = listener.accept().await {
360 let wake = Arc::clone(&n);
361 tokio::spawn(async move {
362 use tokio::io::AsyncReadExt;
363 let _ = stream.write_all(b"200 mock\r\n").await;
364 let mut cmd_buf = vec![0u8; 256];
367 loop {
368 tokio::select! {
369 () = wake.notified() => break,
370 result = stream.read(&mut cmd_buf) => {
371 match result {
372 Ok(n) if n > 0 => { let _ = stream.write_all(b"200 OK\r\n").await; }
373 _ => break,
374 }
375 }
376 }
377 }
378 let _ = stream.write_all(article_prefix).await;
379 let _ = stream.shutdown().await;
380 });
381 }
382 }
383 });
384
385 (addr, notify)
386 }
387
388 fn make_test_pool(addr: std::net::SocketAddr) -> crate::pool::deadpool_connection::Pool {
389 let manager = crate::pool::deadpool_connection::TcpManager::new(
390 addr.ip().to_string(),
391 addr.port(),
392 "test".to_string(),
393 crate::pool::deadpool_connection::TcpManagerOptions {
394 compress: Some(false), ..crate::pool::deadpool_connection::TcpManagerOptions::default()
396 },
397 )
398 .unwrap();
399 crate::pool::deadpool_connection::Pool::builder(manager)
400 .max_size(2)
401 .build()
402 .unwrap()
403 }
404
405 fn make_test_client(addr: std::net::SocketAddr) -> NntpClient {
406 use crate::pool::BufferPool;
407 use crate::types::BufferSize;
408
409 let provider = DeadpoolConnectionProvider::builder(addr.ip().to_string(), addr.port())
410 .name("test")
411 .max_connections(2)
412 .build()
413 .unwrap();
414 let buffer_pool = BufferPool::new(BufferSize::try_new(4096).unwrap(), 2);
415 NntpClient::new(provider, buffer_pool)
416 }
417
418 async fn capture_multiline_response_for_test(
419 conn: &mut crate::stream::ConnectionStream,
420 io_buffer: &mut PooledBuffer,
421 capture: &mut PooledBuffer,
422 ) -> Result<()> {
423 crate::session::backend::capture_complete_multiline_response(conn, io_buffer, capture).await
424 }
425
426 #[tokio::test]
429 async fn test_multiline_response_capture_single_read() {
430 use crate::pool::BufferPool;
431 use crate::types::BufferSize;
432
433 let article = b"220 body follows\r\nHello world\r\n.\r\n";
434 let (addr, notify) = spawn_test_server(article).await;
435 let pool = make_test_pool(addr);
436 let buffer_pool = BufferPool::new(BufferSize::try_new(4096).unwrap(), 2);
437
438 let mut conn = pool.get().await.unwrap();
439 notify.notify_one();
441
442 let mut io_buffer = buffer_pool.acquire();
443 let mut capture = buffer_pool.acquire_capture();
444
445 io_buffer.read_from(&mut *conn).await.unwrap();
447
448 capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
449 .await
450 .unwrap();
451
452 assert_eq!(&capture[..], article as &[u8]);
453 }
454
455 #[tokio::test]
462 async fn test_multiline_response_capture_multi_read_spanning_body_end() {
463 use crate::pool::BufferPool;
464 use crate::types::BufferSize;
465
466 let article = b"220 article\r\nLine one\r\nLine two\r\n.\r\n";
469 let (addr, notify) = spawn_test_server(article).await;
470 let pool = make_test_pool(addr);
471 let buffer_pool = BufferPool::new(BufferSize::try_new(8).unwrap(), 4);
473
474 let mut conn = pool.get().await.unwrap();
475 notify.notify_one();
476
477 let mut io_buffer = buffer_pool.acquire();
478 let mut capture = buffer_pool.acquire_capture();
479
480 capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
482 .await
483 .unwrap();
484
485 assert_eq!(&capture[..], article as &[u8]);
486 }
487
488 #[tokio::test]
489 async fn test_multiline_response_capture_errors_on_truncated_response() {
490 use crate::pool::BufferPool;
491 use crate::types::BufferSize;
492
493 let article_prefix = b"220 body follows\r\npartial article";
494 let (addr, notify) = spawn_truncated_test_server(article_prefix).await;
495 let pool = make_test_pool(addr);
496 let buffer_pool = BufferPool::new(BufferSize::try_new(8).unwrap(), 4);
497
498 let mut conn = pool.get().await.unwrap();
499 notify.notify_one();
500
501 let mut io_buffer = buffer_pool.acquire();
502 let mut capture = buffer_pool.acquire_capture();
503
504 let err = capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
505 .await
506 .unwrap_err();
507
508 assert!(
509 err.to_string()
510 .contains("Backend closed connection before complete"),
511 "unexpected error: {err:#}"
512 );
513 }
514
515 #[tokio::test]
516 async fn test_multiline_response_capture_errors_on_extra_response_bytes() {
517 use crate::pool::BufferPool;
518 use crate::types::BufferSize;
519
520 let article = b"220 body follows\r\nHello world\r\n.\r\n";
521 let extra_response = [article.as_slice(), b"430 No such article\r\n"].concat();
522 let extra_response: &'static [u8] = Box::leak(extra_response.into_boxed_slice());
523 let (addr, notify) = spawn_test_server(extra_response).await;
524 let pool = make_test_pool(addr);
525 let buffer_pool = BufferPool::new(BufferSize::try_new(4096).unwrap(), 2);
526
527 let mut conn = pool.get().await.unwrap();
528 notify.notify_one();
529
530 let mut io_buffer = buffer_pool.acquire();
531 let mut capture = buffer_pool.acquire_capture();
532 io_buffer.read_from(&mut *conn).await.unwrap();
533
534 let err = capture_multiline_response_for_test(&mut conn, &mut io_buffer, &mut capture)
535 .await
536 .unwrap_err();
537
538 assert!(err.to_string().contains("unexpected"));
539 }
540
541 #[tokio::test]
542 async fn fetch_head_reads_multiline_response() {
543 let response = b"221 0 <test@example.com>\r\nSubject: test\r\nFrom: tester\r\n\r\n.\r\n";
544 let addr = spawn_fetch_test_server("HEAD <test@example.com>", response).await;
545 let client = make_test_client(addr);
546 let msg_id = crate::types::MessageId::new("<test@example.com>".to_string()).unwrap();
547
548 let buffer = client.fetch_head(&msg_id).await.unwrap();
549
550 assert_eq!(&buffer[..], response);
551 }
552
553 #[tokio::test]
554 async fn fetch_body_reads_multiline_response() {
555 let response = b"222 0 <test@example.com>\r\nhello world\r\n.\r\n";
556 let addr = spawn_fetch_test_server("BODY <test@example.com>", response).await;
557 let client = make_test_client(addr);
558 let msg_id = crate::types::MessageId::new("<test@example.com>".to_string()).unwrap();
559
560 let buffer = client.fetch_body(&msg_id).await.unwrap();
561
562 assert_eq!(&buffer[..], response);
563 }
564
565 #[tokio::test]
566 async fn fetch_body_reads_multiline_response_above_retention_limit() {
567 let mut response = Vec::with_capacity((4 * 1024 * 1024) + 64);
568 response.extend_from_slice(b"222 0 <large@example.com>\r\n");
569 response.extend(std::iter::repeat_n(b'x', 4 * 1024 * 1024));
570 response.extend_from_slice(b"\r\n.\r\n");
571 let response: &'static [u8] = Box::leak(response.into_boxed_slice());
572 let addr = spawn_fetch_test_server("BODY <large@example.com>", response).await;
573 let client = make_test_client(addr);
574 let msg_id = crate::types::MessageId::new("<large@example.com>".to_string()).unwrap();
575
576 let buffer = client.fetch_body(&msg_id).await.unwrap();
577
578 assert_eq!(&buffer[..], response);
579 }
580}