Skip to main content

eggress_core/
replay.rs

1use std::io;
2use std::pin::Pin;
3use std::task::{Context, Poll};
4
5use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
6
7use crate::BoxStream;
8
9const DEFAULT_MAX_BUFFER: usize = 8 * 1024;
10
11/// Stack-allocated read chunk used while sniffing; avoids a heap allocation
12/// on every poll_read during protocol detection.
13const SNIFF_READ_CHUNK: usize = 2048;
14
15/// A wrapper around a `BoxStream` that provides bounded sniff buffering for
16/// protocol detection. All bytes read during detection are preserved in an
17/// internal buffer. After detection completes, subsequent reads are served
18/// directly from the underlying stream.
19pub struct ReplayStream {
20    inner: BoxStream,
21    buffer: Vec<u8>,
22    read_pos: usize,
23    sniffing: bool,
24    max_buffer: usize,
25}
26
27impl std::fmt::Debug for ReplayStream {
28    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
29        f.debug_struct("ReplayStream")
30            .field("buffer_len", &self.buffer.len())
31            .field("read_pos", &self.read_pos)
32            .field("sniffing", &self.sniffing)
33            .field("max_buffer", &self.max_buffer)
34            .finish()
35    }
36}
37
38impl ReplayStream {
39    /// Creates a new `ReplayStream` with the default 8 KiB buffer size.
40    pub fn new(stream: BoxStream) -> Self {
41        Self {
42            inner: stream,
43            buffer: Vec::new(),
44            read_pos: 0,
45            sniffing: true,
46            max_buffer: DEFAULT_MAX_BUFFER,
47        }
48    }
49
50    /// Creates a new `ReplayStream` with a custom maximum buffer size.
51    pub fn with_max_buffer(stream: BoxStream, max_buffer: usize) -> Self {
52        Self {
53            inner: stream,
54            buffer: Vec::new(),
55            read_pos: 0,
56            sniffing: true,
57            max_buffer,
58        }
59    }
60
61    /// Returns a reference to the bytes that have been sniffed so far.
62    pub fn buffer(&self) -> &[u8] {
63        &self.buffer
64    }
65
66    /// Consumes this `ReplayStream` and returns the underlying stream.
67    ///
68    /// If detection left buffered bytes unread, they remain available through
69    /// the returned stream before reads reach the underlying stream.
70    pub fn into_inner(mut self) -> BoxStream {
71        self.finish_sniff();
72        if self.buffered_remaining() == 0 {
73            self.inner
74        } else {
75            Box::new(self)
76        }
77    }
78
79    /// Disables sniffing mode. After this call, reads are served directly
80    /// from the underlying stream after any unread buffered bytes have been
81    /// delivered.
82    pub fn finish_sniff(&mut self) {
83        self.sniffing = false;
84    }
85
86    /// Returns the number of bytes remaining in the sniff buffer that have
87    /// not yet been returned to a caller.
88    pub fn buffered_remaining(&self) -> usize {
89        self.buffer.len().saturating_sub(self.read_pos)
90    }
91}
92
93impl AsyncRead for ReplayStream {
94    fn poll_read(
95        mut self: Pin<&mut Self>,
96        cx: &mut Context<'_>,
97        buf: &mut ReadBuf<'_>,
98    ) -> Poll<io::Result<()>> {
99        // Serve buffered bytes first, including after `finish_sniff`.
100        let remaining = self.buffer.len().saturating_sub(self.read_pos);
101        if remaining > 0 {
102            let to_copy = remaining.min(buf.remaining());
103            buf.put_slice(&self.buffer[self.read_pos..self.read_pos + to_copy]);
104            self.read_pos += to_copy;
105            return Poll::Ready(Ok(()));
106        }
107
108        if self.sniffing {
109            // Buffer exhausted while sniffing: read more from the underlying
110            // stream via a temporary buffer, then copy into our internal buffer.
111            if self.buffer.len() >= self.max_buffer {
112                return Poll::Ready(Err(io::Error::new(
113                    io::ErrorKind::UnexpectedEof,
114                    "sniff buffer full",
115                )));
116            }
117
118            let space = self.max_buffer - self.buffer.len();
119            let mut temp = [0u8; SNIFF_READ_CHUNK];
120            let read_size = space.min(buf.remaining()).min(temp.len());
121            let mut temp_buf = ReadBuf::new(&mut temp[..read_size]);
122
123            match Pin::new(&mut self.inner).poll_read(cx, &mut temp_buf) {
124                Poll::Ready(Ok(())) => {
125                    let filled = temp_buf.filled().len();
126                    if filled == 0 {
127                        // Underlying stream closed; return EOF.
128                        return Poll::Ready(Ok(()));
129                    }
130                    self.buffer.extend_from_slice(temp_buf.filled());
131                    let to_copy = filled.min(buf.remaining());
132                    buf.put_slice(&self.buffer[self.read_pos..self.read_pos + to_copy]);
133                    self.read_pos += to_copy;
134                    Poll::Ready(Ok(()))
135                }
136                Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
137                Poll::Pending => {
138                    // AsyncRead allows `poll_read` to return `Pending` after
139                    // partially filling the buffer. Propagate buffered bytes
140                    // instead of discarding them.
141                    let filled = temp_buf.filled().len();
142                    if filled > 0 {
143                        self.buffer.extend_from_slice(temp_buf.filled());
144                        let to_copy = filled.min(buf.remaining());
145                        buf.put_slice(&self.buffer[self.read_pos..self.read_pos + to_copy]);
146                        self.read_pos += to_copy;
147                        Poll::Ready(Ok(()))
148                    } else {
149                        Poll::Pending
150                    }
151                }
152            }
153        } else {
154            // Not sniffing: delegate directly to the underlying stream. All
155            // sniffed bytes were already delivered through detection reads,
156            // so the buffer is fully consumed here (`read_pos` always
157            // reaches `buffer.len()`); use `buffer()` for inspection only.
158            Pin::new(&mut self.inner).poll_read(cx, buf)
159        }
160    }
161}
162
163impl AsyncWrite for ReplayStream {
164    fn poll_write(
165        mut self: Pin<&mut Self>,
166        cx: &mut Context<'_>,
167        buf: &[u8],
168    ) -> Poll<io::Result<usize>> {
169        Pin::new(&mut self.inner).poll_write(cx, buf)
170    }
171
172    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
173        Pin::new(&mut self.inner).poll_flush(cx)
174    }
175
176    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
177        Pin::new(&mut self.inner).poll_shutdown(cx)
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184    use tokio::io::{AsyncReadExt, AsyncWriteExt};
185
186    #[tokio::test]
187    async fn test_replay_stream_buffers_during_sniff() {
188        let (mut tx, rx) = tokio::io::duplex(1024);
189        let replay = ReplayStream::new(Box::new(rx));
190        let mut replay = Box::pin(replay);
191
192        tx.write_all(b"hello").await.unwrap();
193        tx.shutdown().await.unwrap();
194
195        let mut buf = [0u8; 1024];
196        let n = replay.read(&mut buf).await.unwrap();
197        assert_eq!(&buf[..n], b"hello");
198        assert_eq!(replay.buffer(), b"hello");
199    }
200
201    #[tokio::test]
202    async fn test_replay_stream_preserves_all_bytes() {
203        let (mut tx, rx) = tokio::io::duplex(1024);
204        let mut replay = ReplayStream::new(Box::new(rx));
205
206        tx.write_all(b"abcdef").await.unwrap();
207
208        // Read 3 bytes at a time
209        let mut buf = [0u8; 3];
210        let n = replay.read(&mut buf).await.unwrap();
211        assert_eq!(&buf[..n], b"abc");
212
213        let n = replay.read(&mut buf).await.unwrap();
214        assert_eq!(&buf[..n], b"def");
215
216        assert_eq!(replay.buffer(), b"abcdef");
217    }
218
219    #[tokio::test]
220    async fn test_replay_stream_into_inner_after_partial_read() {
221        let (mut tx, rx) = tokio::io::duplex(1024);
222        let mut replay = ReplayStream::new(Box::new(rx));
223
224        tx.write_all(b"abcdefghij").await.unwrap();
225
226        // Read only 5 bytes — the ReplayStream reads 5 from inner into its
227        // buffer (limited by caller's buf size), then returns 5 to caller.
228        let mut buf = [0u8; 5];
229        let n = replay.read(&mut buf).await.unwrap();
230        assert_eq!(&buf[..n], b"abcde");
231        assert_eq!(replay.buffer(), b"abcde");
232
233        // Drop tx so the inner stream sees EOF after remaining bytes.
234        drop(tx);
235
236        // into_inner returns the underlying stream, which still has the
237        // remaining 5 bytes unread.
238        let mut inner = replay.into_inner();
239        let mut remaining = Vec::new();
240        inner.read_to_end(&mut remaining).await.unwrap();
241        assert_eq!(&remaining[..], b"fghij");
242    }
243
244    #[tokio::test]
245    async fn test_replay_stream_delegates_writes() {
246        let (rx, mut tx) = tokio::io::duplex(1024);
247        let mut replay = ReplayStream::new(Box::new(rx));
248
249        replay.write_all(b"test").await.unwrap();
250
251        let mut buf = [0u8; 4];
252        tx.read_exact(&mut buf).await.unwrap();
253        assert_eq!(&buf, b"test");
254    }
255
256    #[tokio::test]
257    async fn test_replay_stream_finish_sniff_delegates_to_inner() {
258        let (mut tx, rx) = tokio::io::duplex(1024);
259        let mut replay = ReplayStream::new(Box::new(rx));
260
261        tx.write_all(b"hello").await.unwrap();
262
263        let mut buf = [0u8; 1024];
264        let n = replay.read(&mut buf).await.unwrap();
265        assert_eq!(&buf[..n], b"hello");
266
267        replay.finish_sniff();
268        assert!(!replay.sniffing);
269
270        tx.write_all(b"world").await.unwrap();
271
272        let n = replay.read(&mut buf).await.unwrap();
273        assert_eq!(&buf[..n], b"world");
274    }
275
276    #[tokio::test]
277    async fn test_replay_stream_custom_max_buffer() {
278        let (tx, rx) = tokio::io::duplex(1024);
279        let mut replay = ReplayStream::with_max_buffer(Box::new(rx), 4);
280
281        // Write 6 bytes
282        let write_jh = tokio::spawn(async move {
283            let mut stream = tx;
284            stream.write_all(b"abcdef").await.unwrap();
285            stream.shutdown().await.unwrap();
286        });
287
288        let mut buf = [0u8; 4];
289        let n = replay.read(&mut buf).await.unwrap();
290        assert_eq!(&buf[..n], b"abcd");
291
292        // Buffer is now full (6 bytes > max_buffer of 4), next read should error
293        let result = replay.read(&mut buf).await;
294        assert!(result.is_err());
295
296        write_jh.await.unwrap();
297    }
298
299    #[tokio::test]
300    async fn test_replay_stream_empty_read() {
301        let (tx, rx) = tokio::io::duplex(1024);
302        let mut replay = ReplayStream::new(Box::new(rx));
303
304        // Close the write half
305        drop(tx);
306
307        let mut buf = [0u8; 1024];
308        let n = replay.read(&mut buf).await.unwrap();
309        assert_eq!(n, 0);
310    }
311
312    #[tokio::test]
313    async fn test_replay_stream_reads_after_sniff_continue_from_inner() {
314        let (mut tx, rx) = tokio::io::duplex(1024);
315        let mut replay = ReplayStream::new(Box::new(rx));
316
317        tx.write_all(b"first").await.unwrap();
318
319        // Read all data (triggers read from inner into buffer)
320        let mut buf = [0u8; 1024];
321        let n = replay.read(&mut buf).await.unwrap();
322        assert_eq!(&buf[..n], b"first");
323        assert_eq!(replay.buffer(), b"first");
324        assert_eq!(replay.buffered_remaining(), 0);
325
326        // Finish sniff - subsequent reads go to inner
327        replay.finish_sniff();
328
329        tx.write_all(b"second").await.unwrap();
330        let n = replay.read(&mut buf).await.unwrap();
331        assert_eq!(&buf[..n], b"second");
332    }
333
334    #[tokio::test]
335    async fn test_finish_sniff_does_not_replay_consumed_prefix() {
336        // Pins the documented contract: sniffed bytes are delivered during
337        // detection reads only. After `finish_sniff`, the stream continues
338        // from the underlying position — the prefix is NOT replayed again.
339        let (mut tx, rx) = tokio::io::duplex(1024);
340        let mut replay = ReplayStream::new(Box::new(rx));
341
342        tx.write_all(b"prefix").await.unwrap();
343        let mut buf = [0u8; 1024];
344        let n = replay.read(&mut buf).await.unwrap();
345        assert_eq!(&buf[..n], b"prefix");
346
347        replay.finish_sniff();
348
349        tx.write_all(b"next").await.unwrap();
350        drop(tx);
351        let mut rest = Vec::new();
352        replay.read_to_end(&mut rest).await.unwrap();
353        assert_eq!(&rest, b"next");
354    }
355
356    #[tokio::test]
357    async fn test_finish_sniff_preserves_unread_prefix() {
358        let (mut tx, rx) = tokio::io::duplex(1024);
359        tx.write_all(b"abcdef").await.unwrap();
360        drop(tx);
361
362        let mut replay = ReplayStream::new(Box::new(rx));
363        let mut sniffed = [0u8; 2];
364        replay.read_exact(&mut sniffed).await.unwrap();
365        replay.finish_sniff();
366
367        let mut rest = Vec::new();
368        replay.read_to_end(&mut rest).await.unwrap();
369        assert_eq!(rest, b"cdef");
370    }
371}