use std::collections::TryReserveError;
use std::future::poll_fn;
use std::io;
use std::io::Error;
use std::io::ErrorKind;
use std::pin::Pin;
use std::task::Context;
use std::task::Poll;
use crate::AsyncClose;
use crate::AsyncOutput;
use crate::Buffer;
use crate::async_io::MAX_READY_OPERATIONS_PER_POLL;
use crate::buffered::DEFAULT_BUFFER_CAPACITY;
use crate::traits::normalize_async_error;
#[must_use]
#[derive(Debug)]
pub struct AsyncBufferedOutput<O>
where
O: AsyncOutput,
O::Item: Clone + Default,
{
inner: O,
buffer: Buffer<O::Item>,
}
impl<O> AsyncBufferedOutput<O>
where
O: AsyncOutput,
O::Item: Clone + Default,
{
#[inline(always)]
pub fn new(inner: O) -> Self {
Self::with_capacity(inner, DEFAULT_BUFFER_CAPACITY)
}
#[inline]
pub fn with_capacity(inner: O, capacity: usize) -> Self {
Self {
inner,
buffer: Buffer::with_capacity(capacity),
}
}
#[inline]
pub fn try_with_capacity(inner: O, capacity: usize) -> Result<Self, TryReserveError> {
Ok(Self {
inner,
buffer: Buffer::try_with_capacity(capacity)?,
})
}
#[inline(always)]
#[must_use]
pub const fn inner(&self) -> &O {
&self.inner
}
#[inline(always)]
#[must_use]
pub fn inner_mut(&mut self) -> &mut O {
&mut self.inner
}
#[inline(always)]
#[must_use = "the returned inner output and pending buffer must be handled"]
pub fn into_parts(self) -> (O, Buffer<O::Item>) {
(self.inner, self.buffer)
}
#[inline(always)]
#[must_use]
pub fn capacity(&self) -> usize {
self.buffer.capacity()
}
#[inline(always)]
#[must_use]
pub const fn pending_len(&self) -> usize {
self.buffer.available()
}
#[inline(always)]
#[must_use]
pub fn pending(&self) -> &[O::Item] {
self.buffer.readable()
}
#[inline(always)]
pub fn try_reserve_capacity(&mut self, capacity: usize) -> Result<(), TryReserveError> {
self.buffer.try_reserve_capacity(capacity)
}
#[inline(always)]
#[must_use]
pub fn spare_capacity(&self) -> usize {
self.buffer.spare_capacity()
}
#[inline(always)]
#[must_use]
pub fn spare_raw_parts_mut(&mut self) -> (&mut [O::Item], usize, usize) {
self.buffer.spare_raw_parts_mut()
}
#[inline(always)]
pub unsafe fn advance(&mut self, count: usize) {
unsafe {
self.buffer.advance(count);
}
}
pub fn poll_ensure_spare_capacity(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
count: usize,
) -> Poll<io::Result<()>> {
if count > self.as_ref().get_ref().buffer.capacity() {
return Poll::Ready(Err(Error::new(
ErrorKind::InvalidInput,
"requested spare capacity exceeds buffered output capacity",
)));
}
if self.as_ref().get_ref().buffer.spare_capacity() < count {
return self.as_mut().poll_drain_buffer(cx);
}
Poll::Ready(Ok(()))
}
fn poll_drain_buffer(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
let this = unsafe { self.as_mut().get_unchecked_mut() };
let mut ready_operations = 0;
while !this.buffer.is_empty() {
let result = {
let pending = this.buffer.readable();
let inner = unsafe { Pin::new_unchecked(&mut this.inner) };
inner.poll_write(cx, pending)
};
match result {
Poll::Ready(Ok(0)) => {
return Poll::Ready(Err(io::ErrorKind::WriteZero.into()));
}
Poll::Ready(Ok(written)) => {
unsafe {
this.buffer.consume(written);
}
ready_operations += 1;
if !this.buffer.is_empty() && ready_operations >= MAX_READY_OPERATIONS_PER_POLL {
cx.waker().wake_by_ref();
return Poll::Pending;
}
}
Poll::Ready(Err(error)) => return Poll::Ready(Err(error)),
Poll::Pending => return Poll::Pending,
}
}
this.buffer.clear();
Poll::Ready(Ok(()))
}
}
impl<O> AsyncBufferedOutput<O>
where
O: AsyncOutput + Unpin,
O::Item: Clone + Default + Unpin,
{
pub async fn ensure_spare_capacity_async(&mut self, count: usize) -> io::Result<()> {
poll_fn(|cx| Pin::new(&mut *self).poll_ensure_spare_capacity(cx, count)).await
}
}
impl<O> AsyncOutput for AsyncBufferedOutput<O>
where
O: AsyncOutput,
O::Item: Clone + Default,
{
type Item = O::Item;
#[inline(always)]
fn is_buffered(&self) -> bool {
true
}
unsafe fn poll_write_unchecked(
mut self: Pin<&mut Self>,
cx: &mut Context<'_>,
input: &[Self::Item],
index: usize,
count: usize,
) -> Poll<io::Result<usize>> {
if count == 0 {
return Poll::Ready(Ok(0));
}
let (spare, capacity) = unsafe {
let this = self.as_mut().get_unchecked_mut();
(this.buffer.spare_capacity(), this.buffer.capacity())
};
if count <= spare && count < capacity {
unsafe {
self.as_mut().get_unchecked_mut().buffer.copy_from(input, index, count);
}
return Poll::Ready(Ok(count));
}
match self.as_mut().poll_drain_buffer(cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(error)) => return Poll::Ready(Err(error)),
Poll::Pending => return Poll::Pending,
}
let this = unsafe { self.as_mut().get_unchecked_mut() };
if count >= capacity {
let source = &input[index..index + count];
let inner = unsafe { Pin::new_unchecked(&mut this.inner) };
return inner.poll_write(cx, source);
}
unsafe {
this.buffer.copy_from(input, index, count);
}
Poll::Ready(Ok(count))
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.as_mut().poll_drain_buffer(cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(error)) => return Poll::Ready(Err(error)),
Poll::Pending => return Poll::Pending,
}
let this = unsafe { self.get_unchecked_mut() };
unsafe { Pin::new_unchecked(&mut this.inner) }
.poll_flush(cx)
.map(|result| result.map_err(normalize_async_error))
}
}
impl<O> AsyncClose for AsyncBufferedOutput<O>
where
O: AsyncClose,
O::Item: Clone + Default,
{
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<io::Result<()>> {
match self.as_mut().poll_drain_buffer(cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(error)) => return Poll::Ready(Err(error)),
Poll::Pending => return Poll::Pending,
}
let this = unsafe { self.get_unchecked_mut() };
unsafe { Pin::new_unchecked(&mut this.inner) }
.poll_close(cx)
.map(|result| result.map_err(normalize_async_error))
}
}