use std::{
io,
pin::Pin,
task::{Context, Poll},
};
use tokio::io::AsyncReadExt;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
pub trait AsyncPeek {
fn peek_u8(&mut self) -> impl Future<Output = anyhow::Result<u8>>;
fn peek(&mut self, buf: &mut [u8]) -> impl Future<Output = anyhow::Result<usize>>;
fn drain(&mut self) -> Option<Vec<u8>>;
}
#[derive(Debug)]
pub struct PeekableStream<S> {
inner: S,
buf: Option<Vec<u8>>,
}
impl<S> AsyncPeek for PeekableStream<S>
where
S: AsyncRead + Unpin + Send,
{
async fn peek_u8(&mut self) -> anyhow::Result<u8> {
let u8 = self.inner.read_u8().await?;
if let Some(ref mut peek_buf) = self.buf {
peek_buf.push(u8)
} else {
self.buf = Some(vec![u8]);
}
Ok(u8)
}
async fn peek(&mut self, buf: &mut [u8]) -> anyhow::Result<usize> {
let n = self.inner.read(buf).await?;
let mut buf = buf[0..n].to_vec();
if let Some(ref mut peek_buf) = self.buf {
peek_buf.append(&mut buf)
} else {
self.buf = Some(buf);
}
Ok(n)
}
fn drain(&mut self) -> Option<Vec<u8>> {
self.buf.take()
}
}
impl<S> AsyncRead for PeekableStream<S>
where
S: AsyncRead + Unpin,
{
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf,
) -> Poll<io::Result<()>> {
let me = self.get_mut();
if let Some(peek_buf) = me.buf.take() {
buf.put_slice(&peek_buf);
Poll::Ready(Ok(()))
} else {
Pin::new(&mut me.inner).poll_read(cx, buf)
}
}
}
impl<S> AsyncWrite for PeekableStream<S>
where
S: AsyncWrite + Unpin,
{
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)
}
}
impl<S> PeekableStream<S> {
pub fn new(inner: S) -> Self {
PeekableStream { inner, buf: None }
}
}
#[cfg(test)]
mod tests {
use std::io::Cursor;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use super::{AsyncPeek, PeekableStream};
#[tokio::test]
async fn peek_u8_is_replayed_on_next_read() {
let mut stream = PeekableStream::new(Cursor::new(b"abc".to_vec()));
let first = stream.peek_u8().await.unwrap();
let mut buf = [0u8; 3];
let n = stream.read(&mut buf).await.unwrap();
assert_eq!(first, b'a');
assert_eq!(n, 1);
assert_eq!(&buf[..n], b"a");
}
#[tokio::test]
async fn peek_collects_multiple_reads_and_drain_returns_them() {
let mut stream = PeekableStream::new(Cursor::new(b"abcdef".to_vec()));
let mut buf = [0u8; 2];
let mut buf2 = [0u8; 3];
let n1 = stream.peek(&mut buf).await.unwrap();
let n2 = stream.peek(&mut buf2).await.unwrap();
assert_eq!(n1, 2);
assert_eq!(n2, 3);
assert_eq!(stream.drain().unwrap(), b"abcde");
}
#[tokio::test]
async fn async_write_delegates_to_inner_stream() {
let (client, mut server) = tokio::io::duplex(16);
let mut stream = PeekableStream::new(client);
stream.write_all(b"ping").await.unwrap();
let mut received = [0u8; 4];
server.read_exact(&mut received).await.unwrap();
assert_eq!(&received, b"ping");
}
}