use std::io;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
use crate::BoxStream;
const DEFAULT_MAX_BUFFER: usize = 8 * 1024;
const SNIFF_READ_CHUNK: usize = 2048;
pub struct ReplayStream {
inner: BoxStream,
buffer: Vec<u8>,
read_pos: usize,
sniffing: bool,
max_buffer: usize,
}
impl std::fmt::Debug for ReplayStream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ReplayStream")
.field("buffer_len", &self.buffer.len())
.field("read_pos", &self.read_pos)
.field("sniffing", &self.sniffing)
.field("max_buffer", &self.max_buffer)
.finish()
}
}
impl ReplayStream {
pub fn new(stream: BoxStream) -> Self {
Self {
inner: stream,
buffer: Vec::new(),
read_pos: 0,
sniffing: true,
max_buffer: DEFAULT_MAX_BUFFER,
}
}
pub fn with_max_buffer(stream: BoxStream, max_buffer: usize) -> Self {
Self {
inner: stream,
buffer: Vec::new(),
read_pos: 0,
sniffing: true,
max_buffer,
}
}
pub fn buffer(&self) -> &[u8] {
&self.buffer
}
pub fn into_inner(mut self) -> BoxStream {
self.finish_sniff();
if self.buffered_remaining() == 0 {
self.inner
} else {
Box::new(self)
}
}
pub fn finish_sniff(&mut self) {
self.sniffing = false;
}
pub fn buffered_remaining(&self) -> usize {
self.buffer.len().saturating_sub(self.read_pos)
}
}
impl AsyncRead for ReplayStream {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let remaining = self.buffer.len().saturating_sub(self.read_pos);
if remaining > 0 {
let to_copy = remaining.min(buf.remaining());
buf.put_slice(&self.buffer[self.read_pos..self.read_pos + to_copy]);
self.read_pos += to_copy;
return Poll::Ready(Ok(()));
}
if self.sniffing {
if self.buffer.len() >= self.max_buffer {
return Poll::Ready(Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"sniff buffer full",
)));
}
let space = self.max_buffer - self.buffer.len();
let mut temp = [0u8; SNIFF_READ_CHUNK];
let read_size = space.min(buf.remaining()).min(temp.len());
let mut temp_buf = ReadBuf::new(&mut temp[..read_size]);
match Pin::new(&mut self.inner).poll_read(cx, &mut temp_buf) {
Poll::Ready(Ok(())) => {
let filled = temp_buf.filled().len();
if filled == 0 {
return Poll::Ready(Ok(()));
}
self.buffer.extend_from_slice(temp_buf.filled());
let to_copy = filled.min(buf.remaining());
buf.put_slice(&self.buffer[self.read_pos..self.read_pos + to_copy]);
self.read_pos += to_copy;
Poll::Ready(Ok(()))
}
Poll::Ready(Err(e)) => Poll::Ready(Err(e)),
Poll::Pending => {
let filled = temp_buf.filled().len();
if filled > 0 {
self.buffer.extend_from_slice(temp_buf.filled());
let to_copy = filled.min(buf.remaining());
buf.put_slice(&self.buffer[self.read_pos..self.read_pos + to_copy]);
self.read_pos += to_copy;
Poll::Ready(Ok(()))
} else {
Poll::Pending
}
}
}
} else {
Pin::new(&mut self.inner).poll_read(cx, buf)
}
}
}
impl AsyncWrite for ReplayStream {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Pin::new(&mut self.inner).poll_shutdown(cx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
#[tokio::test]
async fn test_replay_stream_buffers_during_sniff() {
let (mut tx, rx) = tokio::io::duplex(1024);
let replay = ReplayStream::new(Box::new(rx));
let mut replay = Box::pin(replay);
tx.write_all(b"hello").await.unwrap();
tx.shutdown().await.unwrap();
let mut buf = [0u8; 1024];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"hello");
assert_eq!(replay.buffer(), b"hello");
}
#[tokio::test]
async fn test_replay_stream_preserves_all_bytes() {
let (mut tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
tx.write_all(b"abcdef").await.unwrap();
let mut buf = [0u8; 3];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"abc");
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"def");
assert_eq!(replay.buffer(), b"abcdef");
}
#[tokio::test]
async fn test_replay_stream_into_inner_after_partial_read() {
let (mut tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
tx.write_all(b"abcdefghij").await.unwrap();
let mut buf = [0u8; 5];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"abcde");
assert_eq!(replay.buffer(), b"abcde");
drop(tx);
let mut inner = replay.into_inner();
let mut remaining = Vec::new();
inner.read_to_end(&mut remaining).await.unwrap();
assert_eq!(&remaining[..], b"fghij");
}
#[tokio::test]
async fn test_replay_stream_delegates_writes() {
let (rx, mut tx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
replay.write_all(b"test").await.unwrap();
let mut buf = [0u8; 4];
tx.read_exact(&mut buf).await.unwrap();
assert_eq!(&buf, b"test");
}
#[tokio::test]
async fn test_replay_stream_finish_sniff_delegates_to_inner() {
let (mut tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
tx.write_all(b"hello").await.unwrap();
let mut buf = [0u8; 1024];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"hello");
replay.finish_sniff();
assert!(!replay.sniffing);
tx.write_all(b"world").await.unwrap();
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"world");
}
#[tokio::test]
async fn test_replay_stream_custom_max_buffer() {
let (tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::with_max_buffer(Box::new(rx), 4);
let write_jh = tokio::spawn(async move {
let mut stream = tx;
stream.write_all(b"abcdef").await.unwrap();
stream.shutdown().await.unwrap();
});
let mut buf = [0u8; 4];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"abcd");
let result = replay.read(&mut buf).await;
assert!(result.is_err());
write_jh.await.unwrap();
}
#[tokio::test]
async fn test_replay_stream_empty_read() {
let (tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
drop(tx);
let mut buf = [0u8; 1024];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(n, 0);
}
#[tokio::test]
async fn test_replay_stream_reads_after_sniff_continue_from_inner() {
let (mut tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
tx.write_all(b"first").await.unwrap();
let mut buf = [0u8; 1024];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"first");
assert_eq!(replay.buffer(), b"first");
assert_eq!(replay.buffered_remaining(), 0);
replay.finish_sniff();
tx.write_all(b"second").await.unwrap();
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"second");
}
#[tokio::test]
async fn test_finish_sniff_does_not_replay_consumed_prefix() {
let (mut tx, rx) = tokio::io::duplex(1024);
let mut replay = ReplayStream::new(Box::new(rx));
tx.write_all(b"prefix").await.unwrap();
let mut buf = [0u8; 1024];
let n = replay.read(&mut buf).await.unwrap();
assert_eq!(&buf[..n], b"prefix");
replay.finish_sniff();
tx.write_all(b"next").await.unwrap();
drop(tx);
let mut rest = Vec::new();
replay.read_to_end(&mut rest).await.unwrap();
assert_eq!(&rest, b"next");
}
#[tokio::test]
async fn test_finish_sniff_preserves_unread_prefix() {
let (mut tx, rx) = tokio::io::duplex(1024);
tx.write_all(b"abcdef").await.unwrap();
drop(tx);
let mut replay = ReplayStream::new(Box::new(rx));
let mut sniffed = [0u8; 2];
replay.read_exact(&mut sniffed).await.unwrap();
replay.finish_sniff();
let mut rest = Vec::new();
replay.read_to_end(&mut rest).await.unwrap();
assert_eq!(rest, b"cdef");
}
}