1use 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
93pub trait AsyncStream: AsyncRead + AsyncWrite + Unpin + Send {}
99
100impl<T> AsyncStream for T where T: AsyncRead + AsyncWrite + Unpin + Send {}
102
103#[derive(Debug)]
104enum ConnectionTransport {
105 Plain(TcpStream),
107 Tls(Box<TlsStream<TcpStream>>),
109 CompressedPlain(Box<DecompressStream<TcpStream>>),
111 CompressedTls(Box<DecompressStream<TlsStream<TcpStream>>>),
113}
114
115#[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 pub fn plain(stream: TcpStream) -> Self {
138 Self::new(ConnectionTransport::Plain(stream))
139 }
140
141 pub fn tls(stream: TlsStream<TcpStream>) -> Self {
143 Self::new(ConnectionTransport::Tls(Box::new(stream)))
144 }
145
146 pub fn compressed_plain(stream: TcpStream) -> Self {
148 Self::new(ConnectionTransport::CompressedPlain(Box::new(
149 DecompressStream::new(stream),
150 )))
151 }
152
153 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 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 #[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 #[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 #[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 #[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 #[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 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 #[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 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 #[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 pub fn queue_pending_bytes(&mut self, bytes: &[u8]) -> anyhow::Result<()> {
288 self.queue_pending_bytes_ordered(bytes, PendingByteOrder::Back)
289 }
290
291 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 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 let debug_str = format!("{conn:?}");
795 assert!(
796 debug_str.contains("Plain"),
797 "Debug output should indicate Plain TCP"
798 );
799 }
800}