use tokio::io::AsyncWrite;
use pin_project_lite::pin_project;
use std::marker::PhantomPinned;
use std::pin::Pin;
use std::task::{ready, Context, Poll};
use std::{future::Future, io::IoSlice};
use std::{io, mem};
pin_project! {
#[derive(Debug)]
#[must_use = "futures do nothing unless you `.await` or poll them"]
pub struct WriteAllVectored<'a, 'b, W: ?Sized> {
writer: &'a mut W,
bufs: &'a mut [IoSlice<'b>],
#[pin]
_pin: PhantomPinned,
}
}
pub fn write_all_vectored<'a, 'b, W>(
writer: &'a mut W,
bufs: &'a mut [IoSlice<'b>],
) -> WriteAllVectored<'a, 'b, W>
where
W: AsyncWrite + Unpin + ?Sized,
{
WriteAllVectored {
writer,
bufs,
_pin: PhantomPinned,
}
}
impl<W> Future for WriteAllVectored<'_, '_, W>
where
W: AsyncWrite + Unpin + ?Sized,
{
type Output = io::Result<()>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let me = self.project();
while !me.bufs.is_empty() {
let non_empty = match me.bufs.iter().position(|b| !b.is_empty()) {
Some(pos) => pos,
None => return Poll::Ready(Ok(())),
};
*me.bufs = &mut mem::take(me.bufs)[non_empty..];
let n = ready!(Pin::new(&mut *me.writer).poll_write_vectored(cx, me.bufs))?;
if n == 0 {
return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
}
self::advance_slices(me.bufs, n);
}
Poll::Ready(Ok(()))
}
}
fn advance_slices<'a>(bufs: &mut &mut [IoSlice<'a>], n: usize) {
let mut remove = 0;
let mut left = n;
for buf in bufs.iter() {
if let Some(remainder) = left.checked_sub(buf.len()) {
left = remainder;
remove += 1;
} else {
break;
}
}
*bufs = &mut std::mem::take(bufs)[remove..];
if let Some(first) = bufs.first_mut() {
let buf = &first[left..];
unsafe {
*first = IoSlice::new(std::mem::transmute::<&[u8], &'a [u8]>(buf));
}
} else {
assert!(left == 0, "advancing io slices beyond their length");
}
}