use hyper::rt::ReadBufCursor;
use std::fmt::Debug;
use std::pin::Pin;
use std::task::{Context, Poll};
pub struct FuturesIo<S> {
inner: S,
scratch: Box<[u8]>,
}
const SCRATCH: usize = 8 * 1024;
impl<S> FuturesIo<S> {
pub fn new(inner: S) -> Self {
Self {
inner,
scratch: vec![0u8; SCRATCH].into_boxed_slice(),
}
}
pub fn into_inner(self) -> S {
self.inner
}
pub fn get_ref(&self) -> &S {
&self.inner
}
}
impl<S: Debug> Debug for FuturesIo<S> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FuturesIo")
.field("inner", &self.inner)
.field("scratch_len", &self.scratch.len())
.finish()
}
}
impl<S: futures_io::AsyncRead + Unpin> hyper::rt::Read for FuturesIo<S> {
fn poll_read(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
mut buf: ReadBufCursor<'_>,
) -> Poll<std::io::Result<()>> {
let want = buf.remaining().min(self.scratch.len());
if want == 0 {
return Poll::Ready(Ok(()));
}
let Self { inner, scratch } = &mut *self;
let n = std::task::ready!(Pin::new(inner).poll_read(cx, &mut scratch[..want]))?;
buf.put_slice(&scratch[..n]);
Poll::Ready(Ok(()))
}
}
impl<S: futures_io::AsyncWrite + Unpin> hyper::rt::Write for FuturesIo<S> {
fn poll_write(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.inner).poll_write(cx, buf)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_flush(cx)
}
fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Pin::new(&mut self.inner).poll_close(cx)
}
fn poll_write_vectored(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
bufs: &[std::io::IoSlice<'_>],
) -> Poll<std::io::Result<usize>> {
Pin::new(&mut self.inner).poll_write_vectored(cx, bufs)
}
fn is_write_vectored(&self) -> bool {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures_executor::block_on;
use std::future::poll_fn;
use std::pin::Pin;
use std::task::Context;
use std::task::Poll;
struct Chunked {
data: Vec<u8>,
at: usize,
step: usize,
}
impl futures_io::AsyncRead for Chunked {
fn poll_read(
mut self: Pin<&mut Self>,
_: &mut Context<'_>,
buf: &mut [u8],
) -> Poll<std::io::Result<usize>> {
let n = self.step.min(buf.len()).min(self.data.len() - self.at);
buf[..n].copy_from_slice(&self.data[self.at..self.at + n]);
self.at += n;
Poll::Ready(Ok(n))
}
}
fn read_all(mut io: FuturesIo<Chunked>) -> Vec<u8> {
let mut out = Vec::new();
let mut store = [0u8; 8];
loop {
let mut rb = hyper::rt::ReadBuf::new(&mut store);
let poll = block_on(poll_fn(|cx| {
hyper::rt::Read::poll_read(Pin::new(&mut io), cx, rb.unfilled())
}));
poll.unwrap();
let filled = rb.filled().to_vec();
if filled.is_empty() {
return out;
}
out.extend_from_slice(&filled);
}
}
#[test]
fn forwards_bytes_through_partial_reads() {
let io = FuturesIo::new(Chunked {
data: b"hello world".to_vec(),
at: 0,
step: 3,
});
assert_eq!(read_all(io), b"hello world");
}
#[test]
fn never_writes_more_than_remaining() {
let io = FuturesIo::new(Chunked {
data: vec![7u8; 64],
at: 0,
step: 64,
});
assert_eq!(read_all(io).len(), 64);
}
#[test]
fn into_inner_round_trips() {
let io = FuturesIo::new(Chunked {
data: vec![],
at: 0,
step: 1,
});
let c = io.into_inner();
assert_eq!(c.step, 1);
}
#[test]
fn read_request_larger_than_scratch_buffer_does_not_panic() {
let len = super::SCRATCH + 137;
let data: Vec<u8> = (0..len).map(|i| (i % 251) as u8).collect();
let mut io = FuturesIo::new(Chunked {
data: data.clone(),
at: 0,
step: len,
});
let mut out = Vec::new();
let mut store = vec![0u8; len];
loop {
let mut rb = hyper::rt::ReadBuf::new(&mut store);
let poll = block_on(poll_fn(|cx| {
hyper::rt::Read::poll_read(Pin::new(&mut io), cx, rb.unfilled())
}));
poll.unwrap();
let filled = rb.filled().to_vec();
if filled.is_empty() {
break;
}
out.extend_from_slice(&filled);
}
assert_eq!(out, data);
}
struct NullWrite;
impl futures_io::AsyncWrite for NullWrite {
fn poll_write(
self: Pin<&mut Self>,
_: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
}
#[test]
fn is_write_vectored_reports_false_not_an_unverifiable_claim() {
let io = FuturesIo::new(NullWrite);
assert!(!hyper::rt::Write::is_write_vectored(&io));
}
}