Skip to main content

moirai_async/io/
ext.rs

1use std::future::Future;
2use std::io;
3use std::pin::Pin;
4use std::task::{Context, Poll};
5
6use crate::io::traits::{AsyncRead, AsyncWrite};
7
8/// Extension methods for types implementing [`AsyncRead`].
9pub trait AsyncReadExt: AsyncRead {
10    /// Read some bytes asynchronously from the reader.
11    fn read<'a>(&'a mut self, buf: &'a mut [u8]) -> Read<'a, Self>
12    where
13        Self: Unpin,
14    {
15        Read { reader: self, buf }
16    }
17
18    /// Read exactly enough bytes to fill the provided buffer.
19    fn read_exact<'a>(&'a mut self, buf: &'a mut [u8]) -> ReadExact<'a, Self>
20    where
21        Self: Unpin,
22    {
23        ReadExact {
24            reader: self,
25            buf,
26            filled: 0,
27        }
28    }
29}
30
31impl<R: AsyncRead + ?Sized> AsyncReadExt for R {}
32
33/// Future returned by [`AsyncReadExt::read`].
34pub struct Read<'a, R: ?Sized> {
35    reader: &'a mut R,
36    buf: &'a mut [u8],
37}
38
39impl<'a, R: AsyncRead + ?Sized + Unpin> Future for Read<'a, R> {
40    type Output = io::Result<usize>;
41
42    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
43        let this = &mut *self;
44        Pin::new(&mut *this.reader).poll_read(cx, this.buf)
45    }
46}
47
48/// Future returned by [`AsyncReadExt::read_exact`].
49pub struct ReadExact<'a, R: ?Sized> {
50    reader: &'a mut R,
51    buf: &'a mut [u8],
52    filled: usize,
53}
54
55impl<'a, R: AsyncRead + ?Sized + Unpin> Future for ReadExact<'a, R> {
56    type Output = io::Result<()>;
57
58    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
59        let this = &mut *self;
60        while this.filled < this.buf.len() {
61            let filled = this.filled;
62            match Pin::new(&mut *this.reader).poll_read(cx, &mut this.buf[filled..]) {
63                Poll::Ready(Ok(0)) => {
64                    return Poll::Ready(Err(io::Error::new(
65                        io::ErrorKind::UnexpectedEof,
66                        "failed to fill whole buffer",
67                    )));
68                }
69                Poll::Ready(Ok(n)) => {
70                    this.filled += n;
71                }
72                Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
73                Poll::Pending => return Poll::Pending,
74            }
75        }
76        Poll::Ready(Ok(()))
77    }
78}
79
80/// Extension methods for types implementing [`AsyncWrite`].
81pub trait AsyncWriteExt: AsyncWrite {
82    /// Write some bytes asynchronously.
83    fn write<'a>(&'a mut self, buf: &'a [u8]) -> Write<'a, Self>
84    where
85        Self: Unpin,
86    {
87        Write { writer: self, buf }
88    }
89
90    /// Write all bytes asynchronously.
91    fn write_all<'a>(&'a mut self, buf: &'a [u8]) -> WriteAll<'a, Self>
92    where
93        Self: Unpin,
94    {
95        WriteAll {
96            writer: self,
97            buf,
98            written: 0,
99        }
100    }
101
102    /// Flush pending writes asynchronously.
103    fn flush(&mut self) -> Flush<'_, Self>
104    where
105        Self: Unpin,
106    {
107        Flush { writer: self }
108    }
109
110    /// Shutdown the write side of the stream.
111    fn shutdown(&mut self) -> Shutdown<'_, Self>
112    where
113        Self: Unpin,
114    {
115        Shutdown { writer: self }
116    }
117}
118
119impl<W: AsyncWrite + ?Sized> AsyncWriteExt for W {}
120
121/// Future returned by [`AsyncWriteExt::write`].
122pub struct Write<'a, W: ?Sized> {
123    writer: &'a mut W,
124    buf: &'a [u8],
125}
126
127impl<'a, W: AsyncWrite + ?Sized + Unpin> Future for Write<'a, W> {
128    type Output = io::Result<usize>;
129
130    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
131        let this = &mut *self;
132        Pin::new(&mut *this.writer).poll_write(cx, this.buf)
133    }
134}
135
136/// Future returned by [`AsyncWriteExt::write_all`].
137pub struct WriteAll<'a, W: ?Sized> {
138    writer: &'a mut W,
139    buf: &'a [u8],
140    written: usize,
141}
142
143impl<'a, W: AsyncWrite + ?Sized + Unpin> Future for WriteAll<'a, W> {
144    type Output = io::Result<()>;
145
146    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
147        let this = &mut *self;
148        while this.written < this.buf.len() {
149            match Pin::new(&mut *this.writer).poll_write(cx, &this.buf[this.written..]) {
150                Poll::Ready(Ok(0)) => {
151                    return Poll::Ready(Err(io::Error::new(
152                        io::ErrorKind::WriteZero,
153                        "failed to write whole buffer",
154                    )));
155                }
156                Poll::Ready(Ok(n)) => {
157                    this.written += n;
158                }
159                Poll::Ready(Err(e)) => return Poll::Ready(Err(e)),
160                Poll::Pending => return Poll::Pending,
161            }
162        }
163        Poll::Ready(Ok(()))
164    }
165}
166
167/// Future returned by [`AsyncWriteExt::flush`].
168pub struct Flush<'a, W: ?Sized> {
169    writer: &'a mut W,
170}
171
172impl<'a, W: AsyncWrite + ?Sized + Unpin> Future for Flush<'a, W> {
173    type Output = io::Result<()>;
174
175    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
176        let this = &mut *self;
177        Pin::new(&mut *this.writer).poll_flush(cx)
178    }
179}
180
181/// Future returned by [`AsyncWriteExt::shutdown`].
182pub struct Shutdown<'a, W: ?Sized> {
183    writer: &'a mut W,
184}
185
186impl<'a, W: AsyncWrite + ?Sized + Unpin> Future for Shutdown<'a, W> {
187    type Output = io::Result<()>;
188
189    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
190        let this = &mut *self;
191        Pin::new(&mut *this.writer).poll_shutdown(cx)
192    }
193}