Skip to main content

nntp_proxy/stream/
connection_stream.rs

1//! Stream abstraction for supporting multiple connection types
2//!
3//! This module provides abstractions for handling different stream types (TCP, TLS, etc.)
4//! in a unified way. This is preparation for adding SSL/TLS support to backend connections.
5
6use smallvec::SmallVec;
7
8use crate::compression::DecompressStream;
9use crate::constants::buffer::MAX_PENDING_BACKEND_BYTES;
10use crate::tls::TlsStream;
11use std::io::{self, IoSlice};
12use std::ops::Range;
13use std::pin::Pin;
14use std::task::{Context, Poll};
15use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
16use tokio::net::TcpStream;
17
18const PENDING_BACKEND_INLINE_BYTES: usize = 1024;
19const PENDING_BACKEND_INLINE_SEGMENTS: usize = 2;
20
21#[derive(Clone, Copy)]
22enum PendingByteOrder {
23    Front,
24    Back,
25}
26
27#[derive(Debug)]
28#[allow(clippy::large_enum_variant)]
29enum PendingBackendInput {
30    Copied {
31        bytes: SmallVec<[u8; PENDING_BACKEND_INLINE_BYTES]>,
32        pos: usize,
33    },
34    Pooled {
35        buffer: crate::pool::PooledBuffer,
36        range: Range<usize>,
37        pos: usize,
38    },
39}
40
41impl PendingBackendInput {
42    fn copied(bytes: &[u8]) -> Self {
43        let mut copied = SmallVec::new();
44        if bytes.len() > PENDING_BACKEND_INLINE_BYTES {
45            crate::pool::buffer::record_pending_backend_byte_heap_fallback();
46        }
47        copied.extend_from_slice(bytes);
48        Self::Copied {
49            bytes: copied,
50            pos: 0,
51        }
52    }
53
54    const fn pooled(buffer: crate::pool::PooledBuffer, range: Range<usize>) -> Self {
55        Self::Pooled {
56            buffer,
57            range,
58            pos: 0,
59        }
60    }
61
62    fn remaining(&self) -> usize {
63        match self {
64            Self::Copied { bytes, pos } => bytes.len().saturating_sub(*pos),
65            Self::Pooled { range, pos, .. } => range.len().saturating_sub(*pos),
66        }
67    }
68
69    fn is_empty(&self) -> bool {
70        self.remaining() == 0
71    }
72
73    fn read_into(&mut self, buf: &mut ReadBuf<'_>) -> usize {
74        match self {
75            Self::Copied { bytes, pos } => {
76                let available = &bytes[*pos..];
77                let n = available.len().min(buf.remaining());
78                buf.put_slice(&available[..n]);
79                *pos += n;
80                n
81            }
82            Self::Pooled { buffer, range, pos } => {
83                let available = &buffer.as_ref()[range.start + *pos..range.end];
84                let n = available.len().min(buf.remaining());
85                buf.put_slice(&available[..n]);
86                *pos += n;
87                n
88            }
89        }
90    }
91}
92
93/// Trait for async streams that can be used for NNTP connections
94///
95/// This trait is automatically implemented for any type that implements
96/// `AsyncRead` + `AsyncWrite` + Unpin + Send, making it easy to support
97/// different connection types (TCP, TLS, etc.).
98pub trait AsyncStream: AsyncRead + AsyncWrite + Unpin + Send {}
99
100// Blanket implementation for all types that meet the requirements
101impl<T> AsyncStream for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
102
103#[derive(Debug)]
104enum ConnectionTransport {
105    /// Plain TCP connection
106    Plain(TcpStream),
107    /// TLS-encrypted connection
108    Tls(Box<TlsStream<TcpStream>>),
109    /// Compressed plain TCP connection (RFC 8054 / XFEATURE COMPRESS GZIP)
110    CompressedPlain(Box<DecompressStream<TcpStream>>),
111    /// Compressed TLS connection (RFC 8054 / XFEATURE COMPRESS GZIP)
112    CompressedTls(Box<DecompressStream<TlsStream<TcpStream>>>),
113}
114
115/// Unified stream type that can represent different connection types.
116///
117/// The stream also owns pending backend input that was already consumed from
118/// the socket but belongs to the next NNTP response.
119#[derive(Debug)]
120pub struct ConnectionStream {
121    transport: ConnectionTransport,
122    pending_input: SmallVec<[PendingBackendInput; PENDING_BACKEND_INLINE_SEGMENTS]>,
123}
124
125fn ensure_pending_bytes_capacity(current: usize, additional: usize) -> anyhow::Result<()> {
126    anyhow::ensure!(
127        current + additional <= MAX_PENDING_BACKEND_BYTES,
128        "Pending bytes exceeds {} bytes ({} bytes): probable protocol desync",
129        MAX_PENDING_BACKEND_BYTES,
130        current + additional
131    );
132    Ok(())
133}
134
135impl ConnectionStream {
136    /// Create a new plain TCP connection stream
137    pub fn plain(stream: TcpStream) -> Self {
138        Self::new(ConnectionTransport::Plain(stream))
139    }
140
141    /// Create a new TLS-encrypted connection stream
142    pub fn tls(stream: TlsStream<TcpStream>) -> Self {
143        Self::new(ConnectionTransport::Tls(Box::new(stream)))
144    }
145
146    /// Create a compressed plain TCP connection stream
147    pub fn compressed_plain(stream: TcpStream) -> Self {
148        Self::new(ConnectionTransport::CompressedPlain(Box::new(
149            DecompressStream::new(stream),
150        )))
151    }
152
153    /// Create a compressed TLS connection stream
154    pub fn compressed_tls(stream: TlsStream<TcpStream>) -> Self {
155        Self::new(ConnectionTransport::CompressedTls(Box::new(
156            DecompressStream::new(stream),
157        )))
158    }
159
160    fn new(transport: ConnectionTransport) -> Self {
161        Self {
162            transport,
163            pending_input: SmallVec::new(),
164        }
165    }
166
167    /// Wrap the current transport in a decompressor, preserving any pending bytes.
168    ///
169    /// Returns an error if compression is requested for a stream that is already compressed.
170    /// This keeps the state transition explicit instead of panicking on an invalid call.
171    pub(crate) fn into_compressed(self, level: u32) -> io::Result<Self> {
172        let transport = match self.transport {
173            ConnectionTransport::Plain(tcp) => ConnectionTransport::CompressedPlain(Box::new(
174                DecompressStream::with_level(tcp, level),
175            )),
176            ConnectionTransport::Tls(tls) => ConnectionTransport::CompressedTls(Box::new(
177                DecompressStream::with_level(*tls, level),
178            )),
179            ConnectionTransport::CompressedPlain(_) | ConnectionTransport::CompressedTls(_) => {
180                return Err(io::Error::new(
181                    io::ErrorKind::InvalidInput,
182                    "cannot enable compression on an already-compressed connection",
183                ));
184            }
185        };
186
187        Ok(Self {
188            transport,
189            pending_input: self.pending_input,
190        })
191    }
192
193    /// Returns the connection type as a string for logging/debugging
194    #[must_use]
195    pub const fn connection_type(&self) -> &'static str {
196        match &self.transport {
197            ConnectionTransport::Plain(_) => "TCP",
198            ConnectionTransport::Tls(_) => "TLS",
199            ConnectionTransport::CompressedPlain(_) => "TCP+COMPRESS",
200            ConnectionTransport::CompressedTls(_) => "TLS+COMPRESS",
201        }
202    }
203
204    /// Returns true if this connection uses encryption (TLS/SSL)
205    #[inline]
206    #[must_use]
207    pub const fn is_encrypted(&self) -> bool {
208        matches!(
209            &self.transport,
210            ConnectionTransport::Tls(_) | ConnectionTransport::CompressedTls(_)
211        )
212    }
213
214    /// Returns true if this connection is unencrypted (plain TCP)
215    #[inline]
216    #[must_use]
217    pub const fn is_unencrypted(&self) -> bool {
218        matches!(
219            &self.transport,
220            ConnectionTransport::Plain(_) | ConnectionTransport::CompressedPlain(_)
221        )
222    }
223
224    /// Returns true if this connection uses wire compression
225    #[inline]
226    #[must_use]
227    pub const fn is_compressed(&self) -> bool {
228        matches!(
229            &self.transport,
230            ConnectionTransport::CompressedPlain(_) | ConnectionTransport::CompressedTls(_)
231        )
232    }
233
234    /// Get a reference to the underlying TCP stream (if plain, uncompressed TCP)
235    ///
236    /// Returns None for TLS or compressed streams.
237    /// Useful for socket optimization that requires direct TCP access.
238    #[must_use]
239    pub const fn as_tcp_stream(&self) -> Option<&TcpStream> {
240        match &self.transport {
241            ConnectionTransport::Plain(tcp) => Some(tcp),
242            _ => None,
243        }
244    }
245
246    /// Get a mutable reference to the underlying TCP stream (if plain, uncompressed TCP)
247    pub const fn as_tcp_stream_mut(&mut self) -> Option<&mut TcpStream> {
248        match &mut self.transport {
249            ConnectionTransport::Plain(tcp) => Some(tcp),
250            _ => None,
251        }
252    }
253
254    /// Get a reference to the TLS stream (if uncompressed TLS connection)
255    #[must_use]
256    pub fn as_tls_stream(&self) -> Option<&TlsStream<TcpStream>> {
257        match &self.transport {
258            ConnectionTransport::Tls(tls) => Some(tls.as_ref()),
259            _ => None,
260        }
261    }
262
263    /// Get a mutable reference to the TLS stream (if uncompressed TLS connection)
264    pub fn as_tls_stream_mut(&mut self) -> Option<&mut TlsStream<TcpStream>> {
265        match &mut self.transport {
266            ConnectionTransport::Tls(tls) => Some(tls.as_mut()),
267            _ => None,
268        }
269    }
270
271    /// Get the underlying TCP stream reference regardless of connection type
272    ///
273    /// For plain TCP, returns the stream directly.
274    /// For TLS, returns the underlying TCP stream within the TLS wrapper.
275    /// For compressed streams, returns the TCP stream from within the wrapper.
276    #[must_use]
277    pub fn underlying_tcp_stream(&self) -> &TcpStream {
278        match &self.transport {
279            ConnectionTransport::Plain(tcp) => tcp,
280            ConnectionTransport::Tls(tls) => tls.get_ref().0,
281            ConnectionTransport::CompressedPlain(cs) => cs.get_ref(),
282            ConnectionTransport::CompressedTls(cs) => cs.get_ref().get_ref().0,
283        }
284    }
285
286    /// Queue bytes that were already read from the backend for the next response read.
287    pub fn queue_pending_bytes(&mut self, bytes: &[u8]) -> anyhow::Result<()> {
288        self.queue_pending_bytes_ordered(bytes, PendingByteOrder::Back)
289    }
290
291    /// Queue bytes ahead of any bytes already retained by this connection.
292    pub fn queue_pending_bytes_first(&mut self, bytes: &[u8]) -> anyhow::Result<()> {
293        self.queue_pending_bytes_ordered(bytes, PendingByteOrder::Front)
294    }
295
296    fn queue_pending_bytes_ordered(
297        &mut self,
298        bytes: &[u8],
299        order: PendingByteOrder,
300    ) -> anyhow::Result<()> {
301        let Some(bytes) = (!bytes.is_empty()).then_some(bytes) else {
302            return Ok(());
303        };
304
305        ensure_pending_bytes_capacity(self.pending_bytes_len(), bytes.len())?;
306
307        if self.pending_input.spilled()
308            || self.pending_input.len() >= PENDING_BACKEND_INLINE_SEGMENTS
309        {
310            crate::pool::buffer::record_pending_backend_byte_heap_fallback();
311        }
312        let segment = PendingBackendInput::copied(bytes);
313        match order {
314            PendingByteOrder::Back => self.pending_input.push(segment),
315            PendingByteOrder::Front => self.pending_input.insert(0, segment),
316        }
317        Ok(())
318    }
319
320    /// Queue bytes already read into a pooled backend buffer ahead of any
321    /// retained bytes. The byte range is opaque to this type; response framing
322    /// decisions remain owned by the caller that installs the segment.
323    pub(crate) fn queue_pooled_pending_bytes_first(
324        &mut self,
325        buffer: crate::pool::PooledBuffer,
326        range: Range<usize>,
327    ) -> anyhow::Result<()> {
328        anyhow::ensure!(
329            range.start <= range.end,
330            "pending pooled range start exceeds end"
331        );
332        let len = range.len();
333        if len == 0 {
334            return Ok(());
335        }
336        anyhow::ensure!(
337            range.end <= buffer.initialized(),
338            "pending pooled range exceeds initialized buffer"
339        );
340        ensure_pending_bytes_capacity(self.pending_bytes_len(), len)?;
341        if self.pending_input.spilled()
342            || self.pending_input.len() >= PENDING_BACKEND_INLINE_SEGMENTS
343        {
344            crate::pool::buffer::record_pending_backend_byte_heap_fallback();
345        }
346        self.pending_input
347            .insert(0, PendingBackendInput::pooled(buffer, range));
348        Ok(())
349    }
350
351    #[must_use]
352    pub fn has_pending_bytes(&self) -> bool {
353        self.pending_input.iter().any(|segment| !segment.is_empty())
354    }
355
356    #[must_use]
357    pub fn pending_bytes_len(&self) -> usize {
358        self.pending_input
359            .iter()
360            .map(PendingBackendInput::remaining)
361            .sum()
362    }
363
364    pub fn clear_pending_bytes(&mut self) {
365        self.pending_input.clear();
366    }
367}
368
369impl AsyncRead for ConnectionStream {
370    fn poll_read(
371        mut self: Pin<&mut Self>,
372        cx: &mut Context<'_>,
373        buf: &mut ReadBuf<'_>,
374    ) -> Poll<io::Result<()>> {
375        if let Some(front) = self.pending_input.first_mut()
376            && buf.remaining() > 0
377        {
378            front.read_into(buf);
379            if front.is_empty() {
380                self.pending_input.remove(0);
381            }
382            return Poll::Ready(Ok(()));
383        }
384
385        match &mut self.transport {
386            ConnectionTransport::Plain(stream) => Pin::new(stream).poll_read(cx, buf),
387            ConnectionTransport::Tls(stream) => Pin::new(stream.as_mut()).poll_read(cx, buf),
388            ConnectionTransport::CompressedPlain(stream) => {
389                Pin::new(stream.as_mut()).poll_read(cx, buf)
390            }
391            ConnectionTransport::CompressedTls(stream) => {
392                Pin::new(stream.as_mut()).poll_read(cx, buf)
393            }
394        }
395    }
396}
397
398impl AsyncWrite for ConnectionStream {
399    fn poll_write(
400        mut self: Pin<&mut Self>,
401        cx: &mut Context<'_>,
402        buf: &[u8],
403    ) -> Poll<io::Result<usize>> {
404        match &mut self.transport {
405            ConnectionTransport::Plain(stream) => Pin::new(stream).poll_write(cx, buf),
406            ConnectionTransport::Tls(stream) => Pin::new(stream.as_mut()).poll_write(cx, buf),
407            ConnectionTransport::CompressedPlain(stream) => {
408                Pin::new(stream.as_mut()).poll_write(cx, buf)
409            }
410            ConnectionTransport::CompressedTls(stream) => {
411                Pin::new(stream.as_mut()).poll_write(cx, buf)
412            }
413        }
414    }
415
416    fn poll_write_vectored(
417        mut self: Pin<&mut Self>,
418        cx: &mut Context<'_>,
419        bufs: &[IoSlice<'_>],
420    ) -> Poll<io::Result<usize>> {
421        match &mut self.transport {
422            ConnectionTransport::Plain(stream) => Pin::new(stream).poll_write_vectored(cx, bufs),
423            ConnectionTransport::Tls(stream) => {
424                Pin::new(stream.as_mut()).poll_write_vectored(cx, bufs)
425            }
426            ConnectionTransport::CompressedPlain(stream) => {
427                Pin::new(stream.as_mut()).poll_write_vectored(cx, bufs)
428            }
429            ConnectionTransport::CompressedTls(stream) => {
430                Pin::new(stream.as_mut()).poll_write_vectored(cx, bufs)
431            }
432        }
433    }
434
435    fn is_write_vectored(&self) -> bool {
436        match &self.transport {
437            ConnectionTransport::Plain(stream) => stream.is_write_vectored(),
438            ConnectionTransport::Tls(stream) => stream.is_write_vectored(),
439            ConnectionTransport::CompressedPlain(stream) => stream.is_write_vectored(),
440            ConnectionTransport::CompressedTls(stream) => stream.is_write_vectored(),
441        }
442    }
443
444    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
445        match &mut self.transport {
446            ConnectionTransport::Plain(stream) => Pin::new(stream).poll_flush(cx),
447            ConnectionTransport::Tls(stream) => Pin::new(stream.as_mut()).poll_flush(cx),
448            ConnectionTransport::CompressedPlain(stream) => {
449                Pin::new(stream.as_mut()).poll_flush(cx)
450            }
451            ConnectionTransport::CompressedTls(stream) => Pin::new(stream.as_mut()).poll_flush(cx),
452        }
453    }
454
455    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
456        match &mut self.transport {
457            ConnectionTransport::Plain(stream) => Pin::new(stream).poll_shutdown(cx),
458            ConnectionTransport::Tls(stream) => Pin::new(stream.as_mut()).poll_shutdown(cx),
459            ConnectionTransport::CompressedPlain(stream) => {
460                Pin::new(stream.as_mut()).poll_shutdown(cx)
461            }
462            ConnectionTransport::CompressedTls(stream) => {
463                Pin::new(stream.as_mut()).poll_shutdown(cx)
464            }
465        }
466    }
467}
468
469#[cfg(test)]
470mod tests {
471    use super::*;
472    use tokio::io::{AsyncReadExt, AsyncWrite, AsyncWriteExt};
473
474    #[tokio::test]
475    async fn test_connection_stream_plain_tcp() {
476        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
477        let addr = listener.local_addr().unwrap();
478
479        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
480
481        let (server_stream, _) = listener.accept().await.unwrap();
482        let client_stream = client_handle.await.unwrap();
483
484        let mut server_conn = ConnectionStream::plain(server_stream);
485        let mut client_conn = ConnectionStream::plain(client_stream);
486
487        client_conn.write_all(b"Hello").await.unwrap();
488
489        let mut buf = [0u8; 5];
490        server_conn.read_exact(&mut buf).await.unwrap();
491        assert_eq!(&buf, b"Hello");
492
493        assert!(client_conn.is_unencrypted());
494        assert!(!client_conn.is_encrypted());
495        assert_eq!(client_conn.connection_type(), "TCP");
496        assert!(client_conn.as_tcp_stream().is_some());
497    }
498
499    #[test]
500    fn test_async_stream_trait() {
501        fn assert_async_stream<T: AsyncStream>() {}
502        assert_async_stream::<TcpStream>();
503        assert_async_stream::<ConnectionStream>();
504    }
505
506    #[tokio::test]
507    async fn test_plain_connection_stream_preserves_vectored_write_support() {
508        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
509        let addr = listener.local_addr().unwrap();
510
511        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
512        let (_server_stream, _) = listener.accept().await.unwrap();
513        let client_stream = client_handle.await.unwrap();
514
515        let conn = ConnectionStream::plain(client_stream);
516
517        assert!(AsyncWrite::is_write_vectored(&conn));
518    }
519
520    #[tokio::test]
521    async fn test_connection_stream_tcp_access() {
522        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
523        let addr = listener.local_addr().unwrap();
524
525        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
526        let (server_stream, _) = listener.accept().await.unwrap();
527        let _client_stream = client_handle.await.unwrap();
528
529        let mut conn_stream = ConnectionStream::plain(server_stream);
530
531        assert!(conn_stream.is_unencrypted());
532        assert!(conn_stream.as_tcp_stream().is_some());
533        assert!(conn_stream.as_tls_stream().is_none());
534
535        let _underlying = conn_stream.underlying_tcp_stream();
536
537        let tcp_mut = conn_stream.as_tcp_stream_mut().unwrap();
538        tcp_mut.set_nodelay(true).unwrap();
539    }
540
541    #[tokio::test]
542    async fn test_plain_connection_type_checks() {
543        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
544        let addr = listener.local_addr().unwrap();
545
546        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
547        let (server_stream, _) = listener.accept().await.unwrap();
548        let _client = client_handle.await.unwrap();
549
550        let conn = ConnectionStream::plain(server_stream);
551
552        assert_eq!(conn.connection_type(), "TCP");
553        assert!(conn.is_unencrypted());
554        assert!(!conn.is_encrypted());
555        assert!(!conn.is_compressed());
556    }
557
558    #[tokio::test]
559    async fn test_tcp_access_methods_work() {
560        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
561        let addr = listener.local_addr().unwrap();
562
563        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
564        let (server_stream, _) = listener.accept().await.unwrap();
565        let _client = client_handle.await.unwrap();
566
567        let conn = ConnectionStream::plain(server_stream);
568
569        assert!(conn.as_tcp_stream().is_some());
570        assert!(conn.as_tls_stream().is_none());
571    }
572
573    #[tokio::test]
574    async fn test_mutable_tcp_access_methods_work() {
575        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
576        let addr = listener.local_addr().unwrap();
577
578        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
579        let (server_stream, _) = listener.accept().await.unwrap();
580        let _client = client_handle.await.unwrap();
581
582        let mut conn = ConnectionStream::plain(server_stream);
583
584        assert!(conn.as_tcp_stream_mut().is_some());
585        assert!(conn.as_tls_stream_mut().is_none());
586        assert!(conn.as_tcp_stream().is_some());
587
588        let _underlying = conn.underlying_tcp_stream();
589    }
590
591    #[tokio::test]
592    async fn test_constructor_creates_plain_variant() {
593        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
594        let addr = listener.local_addr().unwrap();
595
596        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
597        let (server_stream, _) = listener.accept().await.unwrap();
598        let _client = client_handle.await.unwrap();
599
600        let plain_conn = ConnectionStream::plain(server_stream);
601
602        assert_eq!(plain_conn.connection_type(), "TCP");
603        assert!(plain_conn.is_unencrypted());
604    }
605
606    #[tokio::test]
607    async fn test_pending_bytes_is_read_before_socket() {
608        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
609        let addr = listener.local_addr().unwrap();
610
611        let client_handle = tokio::spawn(async move {
612            let mut client = TcpStream::connect(addr).await.unwrap();
613            client.write_all(b"socket").await.unwrap();
614            client
615        });
616
617        let (server_stream, _) = listener.accept().await.unwrap();
618        let _client = client_handle.await.unwrap();
619
620        let mut conn = ConnectionStream::plain(server_stream);
621        conn.queue_pending_bytes(b"left").unwrap();
622
623        let mut buf = [0u8; 4];
624        conn.read_exact(&mut buf).await.unwrap();
625        assert_eq!(&buf, b"left");
626
627        let mut buf = [0u8; 6];
628        conn.read_exact(&mut buf).await.unwrap();
629        assert_eq!(&buf, b"socket");
630    }
631
632    #[tokio::test]
633    async fn test_pooled_pending_bytes_are_read_before_socket_without_copy_metric() {
634        crate::pool::buffer::reset_hot_path_allocation_metrics();
635        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
636        let addr = listener.local_addr().unwrap();
637
638        let client_handle = tokio::spawn(async move {
639            let mut client = TcpStream::connect(addr).await.unwrap();
640            client.write_all(b"socket").await.unwrap();
641            client
642        });
643
644        let (server_stream, _) = listener.accept().await.unwrap();
645        let _client = client_handle.await.unwrap();
646
647        let pool =
648            crate::pool::BufferPool::new(crate::types::BufferSize::try_new(1024).unwrap(), 1);
649        let mut pending = pool.acquire();
650        pending.copy_from_slice(b"xxpooledyy");
651        assert_eq!(pool.available_buffers(), 0);
652
653        let mut conn = ConnectionStream::plain(server_stream);
654        conn.queue_pooled_pending_bytes_first(pending, 2..8)
655            .unwrap();
656        assert_eq!(conn.pending_bytes_len(), 6);
657        assert_eq!(
658            pool.available_buffers(),
659            0,
660            "queued pending input should retain the pooled allocation"
661        );
662
663        let mut buf = [0u8; 6];
664        conn.read_exact(&mut buf).await.unwrap();
665        assert_eq!(&buf, b"pooled");
666        assert!(!conn.has_pending_bytes());
667        assert_eq!(
668            pool.available_buffers(),
669            1,
670            "fully consumed pending input should return the pooled allocation"
671        );
672
673        let metrics = crate::pool::buffer::hot_path_allocation_metrics_snapshot();
674        assert_eq!(metrics.pending_backend_byte_heap_fallbacks, 0);
675
676        let mut buf = [0u8; 6];
677        conn.read_exact(&mut buf).await.unwrap();
678        assert_eq!(&buf, b"socket");
679    }
680
681    #[tokio::test]
682    async fn test_queue_pending_bytes_rejects_oversized_buffers() {
683        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
684        let addr = listener.local_addr().unwrap();
685
686        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
687        let (server_stream, _) = listener.accept().await.unwrap();
688        let _client = client_handle.await.unwrap();
689
690        let mut conn = ConnectionStream::plain(server_stream);
691        let large = vec![b'x'; MAX_PENDING_BACKEND_BYTES + 1];
692        let err = conn.queue_pending_bytes(&large).unwrap_err();
693        assert!(
694            err.to_string()
695                .contains(&MAX_PENDING_BACKEND_BYTES.to_string())
696        );
697        assert_eq!(conn.pending_bytes_len(), 0);
698    }
699
700    #[allow(clippy::reversed_empty_ranges)]
701    #[tokio::test]
702    async fn test_queue_pooled_pending_bytes_rejects_backwards_range() {
703        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
704        let addr = listener.local_addr().unwrap();
705
706        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
707        let (server_stream, _) = listener.accept().await.unwrap();
708        let _client = client_handle.await.unwrap();
709
710        let pool =
711            crate::pool::BufferPool::new(crate::types::BufferSize::try_new(1024).unwrap(), 1);
712        let mut pending = pool.acquire();
713        pending.copy_from_slice(b"pooled");
714
715        let mut conn = ConnectionStream::plain(server_stream);
716        let err = conn
717            .queue_pooled_pending_bytes_first(pending, 5..2)
718            .unwrap_err();
719        assert!(err.to_string().contains("range start exceeds end"));
720        assert_eq!(conn.pending_bytes_len(), 0);
721    }
722
723    #[tokio::test]
724    async fn test_queue_pending_bytes_ignores_empty_buffers() {
725        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
726        let addr = listener.local_addr().unwrap();
727
728        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
729        let (server_stream, _) = listener.accept().await.unwrap();
730        let _client = client_handle.await.unwrap();
731
732        let mut conn = ConnectionStream::plain(server_stream);
733        conn.queue_pending_bytes(b"").unwrap();
734        assert!(!conn.has_pending_bytes());
735        assert_eq!(conn.pending_bytes_len(), 0);
736    }
737
738    #[tokio::test]
739    async fn test_into_compressed_preserves_pending_bytes() {
740        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
741        let addr = listener.local_addr().unwrap();
742
743        let client_handle = tokio::spawn(async move {
744            let mut client = TcpStream::connect(addr).await.unwrap();
745            client.write_all(b"socket").await.unwrap();
746            client
747        });
748
749        let (server_stream, _) = listener.accept().await.unwrap();
750        let _client = client_handle.await.unwrap();
751
752        let mut conn = ConnectionStream::plain(server_stream);
753        conn.queue_pending_bytes(b"left").unwrap();
754
755        let mut conn = conn.into_compressed(1).unwrap();
756        assert!(conn.is_compressed());
757        assert_eq!(conn.pending_bytes_len(), 4);
758
759        let mut buf = [0u8; 4];
760        conn.read_exact(&mut buf).await.unwrap();
761        assert_eq!(&buf, b"left");
762    }
763
764    #[tokio::test]
765    async fn test_into_compressed_rejects_already_compressed_stream() {
766        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
767        let addr = listener.local_addr().unwrap();
768
769        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
770        let (server_stream, _) = listener.accept().await.unwrap();
771        let _client = client_handle.await.unwrap();
772
773        let conn = ConnectionStream::compressed_plain(server_stream);
774        let err = conn.into_compressed(1).unwrap_err();
775
776        assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
777        assert_eq!(
778            err.to_string(),
779            "cannot enable compression on an already-compressed connection"
780        );
781    }
782
783    #[tokio::test]
784    async fn test_plain_connection_debug_format() {
785        let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
786        let addr = listener.local_addr().unwrap();
787
788        let client_handle = tokio::spawn(async move { TcpStream::connect(addr).await.unwrap() });
789        let (server_stream, _) = listener.accept().await.unwrap();
790        let _client = client_handle.await.unwrap();
791
792        let conn = ConnectionStream::plain(server_stream);
793        // Test Debug implementation
794        let debug_str = format!("{conn:?}");
795        assert!(
796            debug_str.contains("Plain"),
797            "Debug output should indicate Plain TCP"
798        );
799    }
800}