Skip to main content

sfo_io/
simple_async_io.rs

1use std::{cmp, io};
2use std::ops::DerefMut;
3use std::pin::Pin;
4use std::sync::Mutex;
5use std::task::{ready, Context, Poll};
6use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
7use crate::buf::{Buf, DEFAULT_MAX_BUF_SIZE};
8
9#[async_trait::async_trait]
10pub trait SimpleAsyncRead: Send + 'static + Unpin {
11    async fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize>;
12}
13
14#[async_trait::async_trait]
15pub trait SimpleAsyncWrite: Send + 'static + Unpin {
16    async fn write(&mut self, buf: &[u8]) -> std::io::Result<usize>;
17    async fn flush(&mut self) -> std::io::Result<()>;
18    async fn shutdown(&mut self) -> std::io::Result<()>;
19}
20
21enum State<T> {
22    Idle(Option<Buf>),
23    Busy(Pin<Box<dyn Future<Output=(std::io::Result<usize>, Buf, T)> + Send>>),
24}
25
26enum FlushState<T> {
27    Idle,
28    Busy(Pin<Box<dyn Future<Output=(std::io::Result<()>, T)> + Send>>),
29}
30
31enum ShutdownState<T> {
32    Idle,
33    Busy(Pin<Box<dyn Future<Output=(std::io::Result<()>, T)> + Send>>),
34}
35
36struct ReadHolderInner<T: SimpleAsyncRead> {
37    inner: Option<T>,
38    state: State<T>,
39}
40pub struct SimpleAsyncReadHolder<T: SimpleAsyncRead> {
41    inner: Mutex<ReadHolderInner<T>>,
42}
43
44impl<T: SimpleAsyncRead> SimpleAsyncReadHolder<T> {
45    pub fn new(inner: T) -> SimpleAsyncReadHolder<T> {
46        SimpleAsyncReadHolder {
47            inner: Mutex::new(ReadHolderInner {
48                inner: Some(inner),
49                state: State::Idle(Some(Buf::with_capacity(0))),
50            }),
51        }
52    }
53
54    pub fn with_lock_read<R>(&self, f: impl FnOnce(Option<&mut T>) -> R) -> R {
55        let mut state = self.inner.lock().unwrap();
56        f(state.inner.as_mut())
57    }
58
59    pub fn into_read(self) -> Option<T> {
60        let mut state = self.inner.lock().unwrap();
61        return state.inner.take();
62    }
63}
64
65impl <T: SimpleAsyncRead> AsyncRead for SimpleAsyncReadHolder<T> {
66    fn poll_read(
67        self: Pin<&mut Self>,
68        cx: &mut Context<'_>,
69        dst: &mut ReadBuf<'_>,
70    ) -> Poll<io::Result<()>> {
71        let mut state = self.inner.lock().unwrap();
72        loop {
73            match &mut state.deref_mut().state {
74                State::Idle(buf_cell) => {
75                    let mut buf = buf_cell.take().unwrap();
76
77                    if !buf.is_empty() {
78                        buf.copy_to(dst);
79                        *buf_cell = Some(buf);
80                        return Poll::Ready(Ok(()));
81                    }
82
83                    let mut inner = state.inner.take().unwrap();
84                    let max_buf_size = cmp::min(dst.remaining(), DEFAULT_MAX_BUF_SIZE);
85                    state.state = State::Busy(Box::pin(async move {
86                        let ret = unsafe {buf.read_from_async(&mut inner, max_buf_size).await };
87                        (ret, buf, inner)
88                    }));
89                }
90                State::Busy(rx) => {
91                    let (res, mut buf, inner) = ready!(Pin::new(rx).poll(cx));
92                    state.inner = Some(inner);
93
94                    match res {
95                        Ok(_) => {
96                            buf.copy_to(dst);
97                            state.state = State::Idle(Some(buf));
98                            return Poll::Ready(Ok(()));
99                        }
100                        Err(e) => {
101                            assert!(buf.is_empty());
102
103                            state.state = State::Idle(Some(buf));
104                            return Poll::Ready(Err(e));
105                        }
106                    }
107                }
108            }
109        }
110    }
111}
112
113struct SimpleAsyncWriteHolderInner<T: SimpleAsyncWrite> {
114    inner: Option<T>,
115    state: State<T>,
116    flush_state: FlushState<T>,
117    shutdown_state: ShutdownState<T>,
118}
119pub struct SimpleAsyncWriteHolder<T: SimpleAsyncWrite> {
120    inner: Mutex<SimpleAsyncWriteHolderInner<T>>,
121}
122
123impl<T: SimpleAsyncWrite> SimpleAsyncWriteHolder<T> {
124    pub fn new(inner: T) -> SimpleAsyncWriteHolder<T> {
125        Self {
126            inner: Mutex::new(SimpleAsyncWriteHolderInner {
127                inner: Some(inner),
128                state: State::Idle(Some(Buf::with_capacity(0))),
129                flush_state: FlushState::Idle,
130                shutdown_state: ShutdownState::Idle,
131            }),
132        }
133    }
134
135    pub fn with_lock_write<R>(&self, f: impl FnOnce(Option<&mut T>) -> R) -> R {
136        let mut state = self.inner.lock().unwrap();
137        f(state.inner.as_mut())
138    }
139
140    pub fn into_write(self) -> Option<T> {
141        let mut state = self.inner.lock().unwrap();
142        return state.inner.take();
143    }
144}
145
146impl<T: SimpleAsyncWrite> AsyncWrite for SimpleAsyncWriteHolder<T> {
147    fn poll_write(
148        self: Pin<&mut Self>,
149        cx: &mut Context<'_>,
150        src: &[u8],
151    ) -> Poll<io::Result<usize>> {
152        let mut state = self.inner.lock().unwrap();
153        loop {
154            match state.deref_mut().state {
155                State::Idle(ref mut buf_cell) => {
156                    let mut buf = buf_cell.take().unwrap();
157
158                    assert!(buf.is_empty());
159
160                    buf.copy_from(src, DEFAULT_MAX_BUF_SIZE);
161                    let mut inner = state.inner.take().unwrap();
162
163                    state.state = State::Busy(Box::pin(async move {
164                        let res = buf.write_to_async(&mut inner).await;
165
166                        (res, buf, inner)
167                    }));
168                }
169                State::Busy(ref mut rx) => {
170                    let (res, buf, inner) = ready!(Pin::new(rx).poll(cx));
171                    state.state = State::Idle(Some(buf));
172                    state.inner = Some(inner);
173
174                    // If error, return
175                    return Poll::Ready(res);
176                }
177            }
178        }
179    }
180
181    fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
182        let mut state = self.inner.lock().unwrap();
183        loop {
184            match state.deref_mut().flush_state {
185                // The buffer is not used here
186                FlushState::Idle => {
187                        let mut inner = state.inner.take().unwrap();
188
189                        state.flush_state = FlushState::Busy(Box::pin(async move {
190                            let res = inner.flush().await;
191                            (res, inner)
192                        }));
193                }
194                FlushState::Busy(ref mut rx) => {
195                    let (res, inner) = ready!(Pin::new(rx).poll(cx));
196                    state.flush_state = FlushState::Idle;
197                    state.inner = Some(inner);
198
199                    // If error, return
200                    return Poll::Ready(res);
201                }
202            }
203        }
204    }
205
206    fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
207        let mut state = self.inner.lock().unwrap();
208        loop {
209            match state.deref_mut().shutdown_state {
210                ShutdownState::Idle => {
211                        let mut inner = state.inner.take().unwrap();
212
213                        state.shutdown_state = ShutdownState::Busy(Box::pin(async move {
214                            let res = inner.shutdown().await;
215                            (res, inner)
216                        }));
217                }
218                ShutdownState::Busy(ref mut rx) => {
219                    let (res, inner) = ready!(Pin::new(rx).poll(cx));
220                    state.shutdown_state = ShutdownState::Idle;
221                    state.inner = Some(inner);
222                    return Poll::Ready(res);
223                }
224            }
225        }
226    }
227}
228
229#[cfg(test)]
230mod test {
231    use std::io;
232    use std::sync::{Arc, Mutex};
233    use crate::{SimpleAsyncRead, SimpleAsyncReadHolder, SimpleAsyncWrite, SimpleAsyncWriteHolder};
234    use tokio::io::AsyncWriteExt;
235    use tokio::io::AsyncReadExt;
236
237    pub struct TestSimpleAsyncWrite {
238        buf: Arc<Mutex<Vec<u8>>>
239    }
240
241    impl TestSimpleAsyncWrite {
242        pub fn new(buf: Arc<Mutex<Vec<u8>>>) -> TestSimpleAsyncWrite {
243            TestSimpleAsyncWrite {
244                buf,
245            }
246        }
247    }
248
249    #[async_trait::async_trait]
250    impl SimpleAsyncWrite for TestSimpleAsyncWrite {
251        async fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
252            {
253                let mut buffer = self.buf.lock().unwrap();
254                buffer.extend_from_slice(buf);
255            }
256            tokio::time::sleep(tokio::time::Duration::from_millis(1)).await;
257            Ok(buf.len())
258        }
259
260        async fn flush(&mut self) -> io::Result<()> {
261            Ok(())
262        }
263
264        async fn shutdown(&mut self) -> io::Result<()> {
265            Ok(())
266        }
267    }
268
269    pub struct TestSimpleAsyncRead {
270        buf: Arc<Mutex<Vec<u8>>>
271    }
272
273    impl TestSimpleAsyncRead {
274        pub fn new(buf: Arc<Mutex<Vec<u8>>>) -> TestSimpleAsyncRead {
275            TestSimpleAsyncRead { buf }
276        }
277    }
278
279    #[async_trait::async_trait]
280    impl SimpleAsyncRead for TestSimpleAsyncRead {
281        async fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
282            tokio::time::sleep(tokio::time::Duration::from_millis(1)).await;
283            let buffer = self.buf.lock().unwrap();
284            let len = buffer.len();
285            buf.copy_from_slice(&buffer[..len]);
286            Ok(buf.len())
287        }
288    }
289
290    #[tokio::test]
291    async fn test_simple_async_io() {
292        let buf = Arc::new(Mutex::new(Vec::new()));
293        let mut write = SimpleAsyncWriteHolder::new(TestSimpleAsyncWrite::new(buf.clone()));
294        let data = "tttt".to_string();
295        write.write_all(data.as_bytes()).await.unwrap();
296        write.flush().await.unwrap();
297        write.shutdown().await.unwrap();
298
299        let mut read = SimpleAsyncReadHolder::new(TestSimpleAsyncRead::new(buf.clone()));
300        let mut buf = [0u8; 4];
301
302        read.read(&mut buf).await.unwrap();
303        assert_eq!(String::from_utf8_lossy(buf.as_slice()), data);
304    }
305}