Skip to main content

proxy_header/
io.rs

1//! IO wrapper for proxied streams.
2//!
3//! PROXY protocol header is variable length so it is not possible to read a fixed number of bytes
4//! directly from the stream and reading it byte-by-byte can be inefficient. [`ProxiedStream`] reads
5//! enough bytes to parse the header and retains any extra bytes that may have been read.
6//!
7//! If the underlying stream is already buffered (i.e. [`std::io::BufRead`] or equivalent), it is
8//! probably a better idea to just decode the header directly instead of using [`ProxiedStream`].
9//!
10//! The wrapper is usable both with standard ([`std::io::Read`]) and Tokio streams ([`tokio::io::AsyncRead`]).
11//!
12//! ## Example (Tokio)
13//!
14//! ```no_run
15//! # #[cfg(feature = "tokio")]
16//! # #[tokio::main] async fn main() -> Result<(), Box<dyn std::error::Error>> {
17//! use tokio::io::{AsyncReadExt, AsyncWriteExt};
18//! use tokio::net::TcpListener;
19//! use proxy_header::io::ProxiedStream;
20//!
21//! let listener = TcpListener::bind("[::]:1234").await?;
22//!
23//! loop {
24//!     let (mut socket, _) = listener.accept().await?;
25//!     tokio::spawn(async move {
26//!         // Read the proxy header first
27//!         let mut socket = ProxiedStream::create_from_tokio(socket, Default::default())
28//!             .await
29//!             .expect("failed to create proxied stream");
30//!
31//!         // We can now inspect the address
32//!         println!("proxy header: {:?}", socket.proxy_header());
33//!
34//!         // Then process the protocol
35//!         let mut buf = vec![0; 1024];
36//!         loop {
37//!             let n = socket.read(&mut buf).await.unwrap();
38//!             if n == 0 {
39//!                 return;
40//!             }
41//!             socket.write_all(&buf[0..n]).await.unwrap();
42//!         }
43//!     });
44//! }
45//! # }
46//! # #[cfg(not(feature = "tokio"))]
47//! # fn main() {}
48//! ```
49use std::io::{self, BufRead, Read, Write};
50
51#[cfg(any(unix, target_os = "wasi"))]
52use std::os::fd::{AsFd, AsRawFd, BorrowedFd, RawFd};
53
54#[cfg(feature = "tokio")]
55use std::{
56    pin::Pin,
57    task::{Context, Poll},
58};
59
60#[cfg(feature = "tokio")]
61use pin_project_lite::pin_project;
62
63#[cfg(feature = "tokio")]
64use tokio::io::{AsyncBufRead, AsyncRead, AsyncWrite, ReadBuf};
65
66use crate::{Error, ParseConfig, ProxyHeader};
67
68#[cfg(feature = "tokio")]
69pin_project! {
70    /// Wrapper around a stream that starts with a proxy header.
71    ///
72    /// See [module level documentation](`crate::io`)
73    #[derive(Debug)]
74    pub struct ProxiedStream<IO> {
75        #[pin]
76        io: IO,
77        remaining: Vec<u8>,
78        header: ProxyHeader<'static>,
79    }
80}
81
82/// Wrapper around a stream that starts with a proxy header.
83///
84/// See [module level documentation](`crate::io`)
85#[cfg(not(feature = "tokio"))]
86#[derive(Debug)]
87pub struct ProxiedStream<IO> {
88    io: IO,
89    remaining: Vec<u8>,
90    header: ProxyHeader<'static>,
91}
92
93impl<IO> ProxiedStream<IO> {
94    /// Create a new proxied stream from an stream that does not have a proxy header.
95    ///
96    /// This is useful if you want to use the same stream type for proxied and unproxied
97    /// connections.
98    pub fn unproxied(io: IO) -> Self {
99        Self {
100            io,
101            remaining: vec![],
102            header: Default::default(),
103        }
104    }
105
106    /// Get the proxy header.
107    pub fn proxy_header(&self) -> &ProxyHeader<'_> {
108        &self.header
109    }
110
111    /// Gets a reference to the underlying stream.
112    pub fn get_ref(&self) -> &IO {
113        &self.io
114    }
115
116    /// Gets a mutable reference to the underlying stream.
117    pub fn get_mut(&mut self) -> &mut IO {
118        &mut self.io
119    }
120
121    /// Gets a pinned mutable reference to the underlying stream.
122    #[cfg(feature = "tokio")]
123    pub fn get_pin_mut(self: Pin<&mut Self>) -> Pin<&mut IO> {
124        self.project().io
125    }
126
127    /// Consumes this wrapper, returning the underlying stream.
128    pub fn into_inner(self) -> IO {
129        self.io
130    }
131}
132
133#[cfg(feature = "tokio")]
134#[cfg_attr(docsrs, doc(cfg(feature = "tokio")))]
135impl<IO> ProxiedStream<IO>
136where
137    IO: AsyncRead + Unpin,
138{
139    /// Reads the proxy header from an [`tokio::io::AsyncRead`] stream and returns a new [`ProxiedStream`].
140    ///
141    /// This method will read from the stream until a proxy header is found, or the
142    /// stream is closed. If the stream is closed before a proxy header is found,
143    /// this method will return an [`io::Error`] with [`io::ErrorKind::UnexpectedEof`].
144    ///
145    /// If the stream contains invalid data, this method will return an [`io::Error`]
146    /// with [`io::ErrorKind::InvalidData`]. In case of an error, the stream is dropped,
147    /// and any remaining bytes are discarded (which usually means the connection
148    /// is closed).
149    pub async fn create_from_tokio(mut io: IO, config: ParseConfig) -> io::Result<Self> {
150        use tokio::io::AsyncReadExt;
151
152        // 256 bytes should be enough for the longest realistic header with
153        // all extensions. If not, we'll just reallocate. theoretical maximum
154        // is 12 + 4 + 65535 = 65551 bytes, though that would be very silly.
155        //
156        // Maybe we should just error out if we get more than 512 bytes?
157        let mut bytes = Vec::with_capacity(256);
158
159        loop {
160            let bytes_read = io.read_buf(&mut bytes).await?;
161            if bytes_read == 0 {
162                return Err(io::Error::new(
163                    io::ErrorKind::UnexpectedEof,
164                    "end of stream",
165                ));
166            }
167
168            match ProxyHeader::parse(&bytes, config) {
169                Ok((ret, consumed)) => {
170                    let ret = ret.into_owned();
171                    bytes.drain(..consumed);
172
173                    return Ok(Self {
174                        io,
175                        remaining: bytes,
176                        header: ret,
177                    });
178                }
179                Err(Error::BufferTooShort) => continue,
180                Err(_) => {
181                    return Err(io::Error::new(
182                        io::ErrorKind::InvalidData,
183                        "invalid proxy header",
184                    ))
185                }
186            }
187        }
188    }
189}
190
191impl<IO> ProxiedStream<IO>
192where
193    IO: Read,
194{
195    /// Reads the proxy header from a [`Read`] stream and returns a new `ProxiedStream`.
196    ///
197    /// Other than the fact that this method is synchronous, it is identical to [`create_from_tokio`](Self::create_from_tokio).
198    pub fn create_from_std(mut io: IO, config: ParseConfig) -> io::Result<Self> {
199        let mut bytes = Vec::with_capacity(256);
200
201        loop {
202            if bytes.capacity() == bytes.len() {
203                bytes.reserve(32);
204            }
205
206            // Read into the zero-initialized spare capacity, then trim to what was read.
207            let filled = bytes.len();
208            bytes.resize(bytes.capacity(), 0);
209
210            let bytes_read = io.read(&mut bytes[filled..])?;
211            assert!(filled + bytes_read <= bytes.len());
212            bytes.truncate(filled + bytes_read);
213
214            if bytes_read == 0 {
215                return Err(io::Error::new(
216                    io::ErrorKind::UnexpectedEof,
217                    "end of stream",
218                ));
219            }
220
221            match ProxyHeader::parse(&bytes, config) {
222                Ok((ret, consumed)) => {
223                    let ret = ret.into_owned();
224                    bytes.drain(..consumed);
225
226                    return Ok(Self {
227                        io,
228                        remaining: bytes,
229                        header: ret,
230                    });
231                }
232                Err(Error::BufferTooShort) => continue,
233                Err(_) => {
234                    return Err(io::Error::new(
235                        io::ErrorKind::InvalidData,
236                        "invalid proxy header",
237                    ))
238                }
239            }
240        }
241    }
242}
243
244impl<IO> Read for ProxiedStream<IO>
245where
246    IO: Read,
247{
248    fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
249        if !self.remaining.is_empty() {
250            let len = std::cmp::min(self.remaining.len(), buf.len());
251
252            buf[..len].copy_from_slice(&self.remaining[..len]);
253            self.remaining.drain(..len);
254
255            return Ok(len);
256        }
257
258        self.io.read(buf)
259    }
260}
261
262impl<IO> BufRead for ProxiedStream<IO>
263where
264    IO: BufRead,
265{
266    fn fill_buf(&mut self) -> io::Result<&[u8]> {
267        if !self.remaining.is_empty() {
268            return Ok(&self.remaining);
269        }
270        self.io.fill_buf()
271    }
272
273    fn consume(&mut self, mut amt: usize) {
274        if !self.remaining.is_empty() {
275            let len = std::cmp::min(self.remaining.len(), amt);
276            self.remaining.drain(..len);
277            amt -= len;
278        }
279        self.io.consume(amt);
280    }
281}
282
283impl<IO> Write for ProxiedStream<IO>
284where
285    IO: Write,
286{
287    #[inline]
288    fn write_vectored(&mut self, bufs: &[io::IoSlice<'_>]) -> io::Result<usize> {
289        self.io.write_vectored(bufs)
290    }
291
292    #[inline]
293    fn write_all(&mut self, buf: &[u8]) -> io::Result<()> {
294        self.io.write_all(buf)
295    }
296
297    #[inline]
298    fn write_fmt(&mut self, fmt: std::fmt::Arguments<'_>) -> io::Result<()> {
299        self.io.write_fmt(fmt)
300    }
301
302    #[inline]
303    fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
304        self.io.write(buf)
305    }
306
307    #[inline]
308    fn flush(&mut self) -> io::Result<()> {
309        self.io.flush()
310    }
311}
312
313#[cfg(feature = "tokio")]
314#[cfg_attr(docsrs, doc(cfg(feature = "tokio")))]
315impl<IO> AsyncBufRead for ProxiedStream<IO>
316where
317    IO: AsyncBufRead,
318{
319    fn poll_fill_buf(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<&[u8]>> {
320        let me = self.project();
321
322        if !me.remaining.is_empty() {
323            return Poll::Ready(Ok(&me.remaining[..]));
324        }
325
326        me.io.poll_fill_buf(cx)
327    }
328
329    fn consume(self: Pin<&mut Self>, mut amt: usize) {
330        let me = self.project();
331
332        if !me.remaining.is_empty() {
333            let len = std::cmp::min(me.remaining.len(), amt);
334            me.remaining.drain(..len);
335            amt -= len;
336        }
337
338        me.io.consume(amt);
339    }
340}
341
342#[cfg(feature = "tokio")]
343#[cfg_attr(docsrs, doc(cfg(feature = "tokio")))]
344impl<IO> AsyncRead for ProxiedStream<IO>
345where
346    IO: AsyncRead,
347{
348    fn poll_read(
349        self: Pin<&mut Self>,
350        cx: &mut Context<'_>,
351        buf: &mut ReadBuf<'_>,
352    ) -> Poll<io::Result<()>> {
353        let me = self.project();
354
355        if !me.remaining.is_empty() {
356            let len = std::cmp::min(me.remaining.len(), buf.remaining());
357
358            buf.put_slice(&me.remaining[..len]);
359            me.remaining.drain(..len);
360
361            return Poll::Ready(Ok(()));
362        }
363
364        me.io.poll_read(cx, buf)
365    }
366}
367
368#[cfg(feature = "tokio")]
369#[cfg_attr(docsrs, doc(cfg(feature = "tokio")))]
370impl<IO> AsyncWrite for ProxiedStream<IO>
371where
372    IO: AsyncWrite,
373{
374    #[inline]
375    fn poll_write(
376        self: Pin<&mut Self>,
377        cx: &mut Context<'_>,
378        buf: &[u8],
379    ) -> Poll<io::Result<usize>> {
380        self.project().io.poll_write(cx, buf)
381    }
382
383    #[inline]
384    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
385        self.project().io.poll_flush(cx)
386    }
387
388    #[inline]
389    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
390        self.project().io.poll_shutdown(cx)
391    }
392
393    #[inline]
394    fn poll_write_vectored(
395        self: Pin<&mut Self>,
396        cx: &mut Context<'_>,
397        bufs: &[io::IoSlice<'_>],
398    ) -> Poll<Result<usize, io::Error>> {
399        self.project().io.poll_write_vectored(cx, bufs)
400    }
401
402    #[inline]
403    fn is_write_vectored(&self) -> bool {
404        self.io.is_write_vectored()
405    }
406}
407
408#[cfg(any(unix, target_os = "wasi"))]
409#[cfg_attr(docsrs, doc(cfg(any(unix, target_os = "wasi"))))]
410impl<IO> AsRawFd for ProxiedStream<IO>
411where
412    IO: AsRawFd,
413{
414    fn as_raw_fd(&self) -> RawFd {
415        self.io.as_raw_fd()
416    }
417}
418
419#[cfg(any(unix, target_os = "wasi"))]
420#[cfg_attr(docsrs, doc(cfg(any(unix, target_os = "wasi"))))]
421impl<IO> AsFd for ProxiedStream<IO>
422where
423    IO: AsFd,
424{
425    fn as_fd(&self) -> BorrowedFd<'_> {
426        self.io.as_fd()
427    }
428}
429
430#[cfg(test)]
431mod tests {
432    use super::*;
433
434    use crate::{Protocol, ProxiedAddress, ProxyHeader, Tlv};
435    use std::{
436        io::Cursor,
437        net::{Ipv4Addr, SocketAddr, SocketAddrV4},
438    };
439
440    /// Payload long enough that some of it is always read along with the header
441    /// (header reads use a 256 byte buffer) and some of it is left in the inner stream.
442    const PAYLOAD_LEN: usize = 4000;
443
444    fn payload() -> Vec<u8> {
445        (0..PAYLOAD_LEN).map(|i| (i % 251) as u8).collect()
446    }
447
448    fn address() -> ProxiedAddress {
449        ProxiedAddress::stream(
450            "127.0.0.1:1234".parse().unwrap(),
451            "10.0.0.1:443".parse().unwrap(),
452        )
453    }
454
455    /// Returns (description, expected header, encoded header followed by payload)
456    fn cases() -> Vec<(&'static str, ProxyHeader<'static>, Vec<u8>)> {
457        let v4 = ProxyHeader::with_address(ProxiedAddress::stream(
458            "127.0.0.1:1234".parse().unwrap(),
459            "8.8.4.4:5678".parse().unwrap(),
460        ));
461        let v2_tlvs = ProxyHeader::with_tlvs(
462            Some(address()),
463            [
464                Tlv::Authority("example.com".into()),
465                Tlv::UniqueId(b"0123456789"[..].into()),
466            ],
467        );
468
469        let mut ret = Vec::new();
470
471        let mut buf = Vec::new();
472        v4.encode_v1(&mut buf).unwrap();
473        ret.push(("v1", v4.clone(), buf));
474
475        let mut buf = Vec::new();
476        ProxyHeader::with_local().encode_v1(&mut buf).unwrap();
477        ret.push(("v1 local", ProxyHeader::with_local(), buf));
478
479        let mut buf = Vec::new();
480        v4.encode_v2(&mut buf).unwrap();
481        ret.push(("v2", v4, buf));
482
483        let mut buf = Vec::new();
484        ProxyHeader::with_local().encode_v2(&mut buf).unwrap();
485        ret.push(("v2 local", ProxyHeader::with_local(), buf));
486
487        let mut buf = Vec::new();
488        v2_tlvs.encode_v2(&mut buf).unwrap();
489        ret.push(("v2 with TLVs", v2_tlvs, buf));
490
491        for (_, _, buf) in ret.iter_mut() {
492            buf.extend_from_slice(&payload());
493        }
494
495        ret
496    }
497
498    /// A reader that returns at most one byte per read call.
499    struct Trickle<R>(R);
500
501    impl<R: Read> Read for Trickle<R> {
502        fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
503            let len = buf.len().min(1);
504            self.0.read(&mut buf[..len])
505        }
506    }
507
508    /// Drains a `BufRead` through `fill_buf` / `consume`, consuming at most `step` bytes at a time.
509    fn drain_buf_read(mut r: impl BufRead, step: usize) -> Vec<u8> {
510        let mut out = Vec::new();
511        loop {
512            let buf = r.fill_buf().unwrap();
513            if buf.is_empty() {
514                return out;
515            }
516            let n = buf.len().min(step);
517            out.extend_from_slice(&buf[..n]);
518            r.consume(n);
519        }
520    }
521
522    #[test]
523    fn test_sync() {
524        let mut buf = [0; 1024];
525
526        let header = ProxyHeader::with_address(ProxiedAddress {
527            protocol: Protocol::Stream,
528            source: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 1), 1234)),
529            destination: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(8, 8, 4, 4), 5678)),
530        });
531
532        let written_len = header.encode_to_slice_v2(&mut buf).unwrap();
533        buf[written_len..].fill(255);
534
535        let mut stream = Cursor::new(&buf);
536
537        let mut proxied = ProxiedStream::create_from_std(&mut stream, Default::default()).unwrap();
538        assert_eq!(proxied.proxy_header(), &header);
539
540        let mut buf = Vec::new();
541        proxied.read_to_end(&mut buf).unwrap();
542
543        assert_eq!(buf.len(), 1024 - written_len);
544        assert!(buf.into_iter().all(|b| b == 255));
545    }
546
547    #[test]
548    fn test_sync_read() {
549        for (name, header, data) in cases() {
550            for read_size in [1, 7, 100, 8192] {
551                let mut proxied =
552                    ProxiedStream::create_from_std(Cursor::new(&data), Default::default()).unwrap();
553                assert_eq!(proxied.proxy_header(), &header, "{name}");
554
555                let mut out = Vec::new();
556                let mut buf = vec![0; read_size];
557                loop {
558                    let n = proxied.read(&mut buf).unwrap();
559                    if n == 0 {
560                        break;
561                    }
562                    out.extend_from_slice(&buf[..n]);
563                }
564                assert_eq!(out, payload(), "{name}, read size {read_size}");
565            }
566        }
567    }
568
569    #[test]
570    fn test_sync_trickle() {
571        // Header arrives one byte at a time, so parsing has to retry on BufferTooShort
572        // and nothing past the header may be over-read.
573        for (name, header, data) in cases() {
574            let mut proxied =
575                ProxiedStream::create_from_std(Trickle(Cursor::new(&data)), Default::default())
576                    .unwrap();
577            assert_eq!(proxied.proxy_header(), &header, "{name}");
578
579            let mut out = Vec::new();
580            proxied.read_to_end(&mut out).unwrap();
581            assert_eq!(out, payload(), "{name}");
582        }
583    }
584
585    #[test]
586    fn test_sync_buf_read() {
587        for (name, header, data) in cases() {
588            for capacity in [16, 300, 8192] {
589                for step in [1, 13, usize::MAX] {
590                    let inner = io::BufReader::with_capacity(capacity, Cursor::new(&data));
591                    let proxied =
592                        ProxiedStream::create_from_std(inner, Default::default()).unwrap();
593                    assert_eq!(proxied.proxy_header(), &header, "{name}");
594
595                    assert_eq!(
596                        drain_buf_read(proxied, step),
597                        payload(),
598                        "{name}, capacity {capacity}, step {step}"
599                    );
600                }
601            }
602        }
603    }
604
605    #[test]
606    fn test_sync_mixed_read_and_buf_read() {
607        for (name, _, data) in cases() {
608            let inner = io::BufReader::new(Cursor::new(&data));
609            let mut proxied = ProxiedStream::create_from_std(inner, Default::default()).unwrap();
610
611            let mut out = vec![0; 5];
612            proxied.read_exact(&mut out).unwrap();
613
614            let buf = proxied.fill_buf().unwrap();
615            let n = buf.len().min(3);
616            out.extend_from_slice(&buf[..n]);
617            proxied.consume(n);
618
619            proxied.read_to_end(&mut out).unwrap();
620            assert_eq!(out, payload(), "{name}");
621        }
622    }
623
624    #[test]
625    fn test_sync_errors() {
626        let err =
627            ProxiedStream::create_from_std(Cursor::new(b"PROXY TCP4 1.2."), Default::default())
628                .unwrap_err();
629        assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
630
631        let err = ProxiedStream::create_from_std(Cursor::new(b""), Default::default()).unwrap_err();
632        assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
633
634        let err = ProxiedStream::create_from_std(
635            Cursor::new(b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n"),
636            Default::default(),
637        )
638        .unwrap_err();
639        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
640
641        let err = ProxiedStream::create_from_std(
642            Cursor::new(b"PROXY TCP4 1.2.3.4 5.6.7.8 1 2\r\n"),
643            ParseConfig {
644                allow_v1: false,
645                ..Default::default()
646            },
647        )
648        .unwrap_err();
649        assert_eq!(err.kind(), io::ErrorKind::InvalidData);
650    }
651
652    #[test]
653    fn test_sync_unproxied() {
654        let data = payload();
655        let mut proxied = ProxiedStream::unproxied(io::BufReader::new(Cursor::new(&data)));
656        assert_eq!(proxied.proxy_header(), &ProxyHeader::with_local());
657
658        let mut out = Vec::new();
659        proxied.read_to_end(&mut out).unwrap();
660        assert_eq!(out, data);
661
662        let proxied = ProxiedStream::unproxied(io::BufReader::new(Cursor::new(&data)));
663        assert_eq!(drain_buf_read(proxied, 17), data);
664    }
665
666    #[test]
667    fn test_sync_write_passthrough() {
668        let mut proxied = ProxiedStream::unproxied(Vec::new());
669        proxied.write_all(b"hello ").unwrap();
670        write!(proxied, "{}", 42).unwrap();
671        proxied.flush().unwrap();
672        assert_eq!(proxied.into_inner(), b"hello 42");
673    }
674
675    #[cfg(feature = "tokio")]
676    mod tokio_tests {
677        use super::*;
678        use tokio::io::{AsyncBufReadExt, AsyncReadExt, AsyncWriteExt, BufReader};
679
680        /// Drains an `AsyncBufRead` through `fill_buf` / `consume`, consuming at most
681        /// `step` bytes at a time.
682        async fn drain_async_buf_read(mut r: impl AsyncBufRead + Unpin, step: usize) -> Vec<u8> {
683            let mut out = Vec::new();
684            loop {
685                let buf = r.fill_buf().await.unwrap();
686                if buf.is_empty() {
687                    return out;
688                }
689                let n = buf.len().min(step);
690                out.extend_from_slice(&buf[..n]);
691                Pin::new(&mut r).consume(n);
692            }
693        }
694
695        #[tokio::test]
696        async fn test_tokio() {
697            let mut buf = [0; 1024];
698
699            let header = ProxyHeader::with_address(ProxiedAddress {
700                protocol: Protocol::Stream,
701                source: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(127, 0, 0, 1), 1234)),
702                destination: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(8, 8, 4, 4), 5678)),
703            });
704
705            let written_len = header.encode_to_slice_v2(&mut buf).unwrap();
706            buf[written_len..].fill(255);
707
708            let mut stream = Cursor::new(&buf);
709
710            let mut proxied = ProxiedStream::create_from_tokio(&mut stream, Default::default())
711                .await
712                .unwrap();
713            assert_eq!(proxied.proxy_header(), &header);
714
715            let mut buf = Vec::new();
716            AsyncReadExt::read_to_end(&mut proxied, &mut buf)
717                .await
718                .unwrap();
719
720            assert_eq!(buf.len(), 1024 - written_len);
721            assert!(buf.into_iter().all(|b| b == 255));
722        }
723
724        #[tokio::test]
725        async fn test_tokio_read() {
726            for (name, header, data) in cases() {
727                for read_size in [1, 7, 100, 8192] {
728                    let mut proxied =
729                        ProxiedStream::create_from_tokio(Cursor::new(&data), Default::default())
730                            .await
731                            .unwrap();
732                    assert_eq!(proxied.proxy_header(), &header, "{name}");
733
734                    let mut out = Vec::new();
735                    let mut buf = vec![0; read_size];
736                    loop {
737                        let n = AsyncReadExt::read(&mut proxied, &mut buf).await.unwrap();
738                        if n == 0 {
739                            break;
740                        }
741                        out.extend_from_slice(&buf[..n]);
742                    }
743                    assert_eq!(out, payload(), "{name}, read size {read_size}");
744                }
745            }
746        }
747
748        #[tokio::test]
749        async fn test_tokio_trickle() {
750            // A duplex pipe with a 1 byte buffer delivers the header one byte at a time.
751            for (name, header, data) in cases() {
752                let (mut tx, rx) = tokio::io::duplex(1);
753                let writer = tokio::spawn(async move {
754                    tx.write_all(&data).await.unwrap();
755                });
756
757                let mut proxied = ProxiedStream::create_from_tokio(rx, Default::default())
758                    .await
759                    .unwrap();
760                assert_eq!(proxied.proxy_header(), &header, "{name}");
761
762                let mut out = Vec::new();
763                AsyncReadExt::read_to_end(&mut proxied, &mut out)
764                    .await
765                    .unwrap();
766                assert_eq!(out, payload(), "{name}");
767
768                writer.await.unwrap();
769            }
770        }
771
772        #[tokio::test]
773        async fn test_tokio_buf_read() {
774            // Regression test: consuming the bytes that were over-read along with the header
775            // used to consume the same amount from the inner reader as well, losing data.
776            for (name, header, data) in cases() {
777                for capacity in [16, 300, 8192] {
778                    for step in [1, 13, usize::MAX] {
779                        let inner = BufReader::with_capacity(capacity, Cursor::new(data.clone()));
780                        let proxied = ProxiedStream::create_from_tokio(inner, Default::default())
781                            .await
782                            .unwrap();
783                        assert_eq!(proxied.proxy_header(), &header, "{name}");
784
785                        assert_eq!(
786                            drain_async_buf_read(proxied, step).await,
787                            payload(),
788                            "{name}, capacity {capacity}, step {step}"
789                        );
790                    }
791                }
792            }
793        }
794
795        #[tokio::test]
796        async fn test_tokio_mixed_read_and_buf_read() {
797            for (name, _, data) in cases() {
798                let inner = BufReader::new(Cursor::new(data));
799                let mut proxied = ProxiedStream::create_from_tokio(inner, Default::default())
800                    .await
801                    .unwrap();
802
803                let mut out = vec![0; 5];
804                AsyncReadExt::read_exact(&mut proxied, &mut out)
805                    .await
806                    .unwrap();
807
808                let buf = proxied.fill_buf().await.unwrap();
809                let n = buf.len().min(3);
810                out.extend_from_slice(&buf[..n]);
811                Pin::new(&mut proxied).consume(n);
812
813                AsyncReadExt::read_to_end(&mut proxied, &mut out)
814                    .await
815                    .unwrap();
816                assert_eq!(out, payload(), "{name}");
817            }
818        }
819
820        #[tokio::test]
821        async fn test_tokio_errors() {
822            let err = ProxiedStream::create_from_tokio(
823                Cursor::new(b"PROXY TCP4 1.2."),
824                Default::default(),
825            )
826            .await
827            .unwrap_err();
828            assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
829
830            let err = ProxiedStream::create_from_tokio(
831                Cursor::new(b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n"),
832                Default::default(),
833            )
834            .await
835            .unwrap_err();
836            assert_eq!(err.kind(), io::ErrorKind::InvalidData);
837        }
838
839        #[tokio::test]
840        async fn test_tokio_unproxied() {
841            let data = payload();
842            let proxied = ProxiedStream::unproxied(BufReader::new(Cursor::new(data.clone())));
843            assert_eq!(proxied.proxy_header(), &ProxyHeader::with_local());
844            assert_eq!(drain_async_buf_read(proxied, 17).await, data);
845        }
846
847        #[tokio::test]
848        async fn test_tokio_write_passthrough() {
849            let (a, mut b) = tokio::io::duplex(64);
850            let mut proxied = ProxiedStream::unproxied(a);
851            proxied.write_all(b"hello world").await.unwrap();
852            proxied.shutdown().await.unwrap();
853
854            let mut out = Vec::new();
855            b.read_to_end(&mut out).await.unwrap();
856            assert_eq!(out, b"hello world");
857        }
858    }
859}