use std::{
fmt, io,
pin::Pin,
task::{Context, Poll},
};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
pub struct BiStream<R, W> {
pub read: R,
pub write: W,
pub name: String,
}
impl<R, W> BiStream<R, W> {
pub fn new(read: R, write: W, name: String) -> Self {
Self { read, write, name }
}
}
impl<R: AsyncRead + Unpin, W: Unpin> AsyncRead for BiStream<R, W> {
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let this = self.get_mut();
Pin::new(&mut this.read).poll_read(cx, buf)
}
}
impl<R: Unpin, W: AsyncWrite + Unpin> AsyncWrite for BiStream<R, W> {
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
let this = self.get_mut();
Pin::new(&mut this.write).poll_write(cx, buf)
}
fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
Pin::new(&mut this.write).poll_flush(cx)
}
fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = self.get_mut();
Pin::new(&mut this.write).poll_shutdown(cx)
}
}
impl<R, W> fmt::Display for BiStream<R, W> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "BiStream{{{}}}", self.name)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
struct MemoryStream {
inner: Cursor<Vec<u8>>,
}
impl MemoryStream {
fn new(data: Vec<u8>) -> Self {
Self {
inner: Cursor::new(data),
}
}
}
impl Unpin for MemoryStream {}
impl AsyncRead for MemoryStream {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<io::Result<()>> {
let pos = self.inner.position() as usize;
let data = self.inner.get_ref();
let available = &data[pos..];
if available.is_empty() {
return Poll::Ready(Ok(()));
}
let to_read = available.len().min(buf.remaining());
buf.put_slice(&available[..to_read]);
self.inner.set_position((pos + to_read) as u64);
Poll::Ready(Ok(()))
}
}
impl AsyncWrite for MemoryStream {
fn poll_write(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<io::Result<usize>> {
self.inner.get_mut().extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn test_bistream_read() {
let data = b"hello world".to_vec();
let reader = MemoryStream::new(data.clone());
let writer = MemoryStream::new(Vec::new());
let mut bistream = BiStream::new(reader, writer, "test_read".to_string());
let mut buffer = vec![0u8; data.len()];
bistream.read.read_exact(&mut buffer).await.unwrap();
assert_eq!(buffer, data);
}
#[tokio::test]
async fn test_bistream_write() {
let reader = MemoryStream::new(Vec::new());
let writer = MemoryStream::new(Vec::new());
let mut bistream = BiStream::new(reader, writer, "test_write".to_string());
let data = b"test data";
let bytes_written = bistream.write.write(data).await.unwrap();
assert_eq!(bytes_written, data.len());
}
#[test]
fn test_display() {
let reader = MemoryStream::new(Vec::new());
let writer = MemoryStream::new(Vec::new());
let bistream = BiStream::new(reader, writer, "DisplayTest".to_string());
assert_eq!(format!("{}", bistream), "BiStream{DisplayTest}");
}
}