borer-core 0.5.9

network borer
Documentation
use std::{
    io,
    pin::Pin,
    task::{Context, Poll},
};
use tokio::io::AsyncReadExt;
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};

/// Async stream extension that can peek bytes without losing them.
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)]
/// Async stream wrapper that replays previously peeked bytes on the next read.
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");
    }
}