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
11const SNIFF_READ_CHUNK: usize = 2048;
14
15pub 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 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 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 pub fn buffer(&self) -> &[u8] {
63 &self.buffer
64 }
65
66 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 pub fn finish_sniff(&mut self) {
83 self.sniffing = false;
84 }
85
86 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 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 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 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 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 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 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 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);
235
236 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 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 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 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 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 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 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}