use std::pin::Pin;
use std::task::{
Context,
Poll,
};
use tokio::io::AsyncWrite;
use crate::{
AsyncClose,
traits::normalize_async_error,
};
#[must_use]
#[repr(transparent)]
pub struct TokioAsyncWrite<O> {
inner: O,
}
impl<O> TokioAsyncWrite<O> {
#[inline(always)]
pub const fn new(inner: O) -> Self {
Self { inner }
}
#[inline(always)]
#[must_use]
pub const fn get_ref(&self) -> &O {
&self.inner
}
#[inline(always)]
#[must_use]
pub const fn get_mut(&mut self) -> &mut O {
&mut self.inner
}
#[inline(always)]
#[must_use]
pub fn get_pin_mut(self: Pin<&mut Self>) -> Pin<&mut O> {
unsafe { self.map_unchecked_mut(|this| &mut this.inner) }
}
#[inline(always)]
#[must_use]
pub fn into_inner(self) -> O {
self.inner
}
}
impl<O> AsyncWrite for TokioAsyncWrite<O>
where
O: AsyncClose<Item = u8>,
{
#[inline(always)]
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
input: &[u8],
) -> Poll<std::io::Result<usize>> {
self.get_pin_mut().poll_write(cx, input)
}
#[inline(always)]
fn poll_flush(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<std::io::Result<()>> {
self.get_pin_mut()
.poll_flush(cx)
.map(|result| result.map_err(normalize_async_error))
}
#[inline(always)]
fn poll_shutdown(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<std::io::Result<()>> {
self.get_pin_mut()
.poll_close(cx)
.map(|result| result.map_err(normalize_async_error))
}
}