sfo_io/
simple_async_io.rs1use 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 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 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 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}