Skip to main content

sfo_io/
blocking.rs

1use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
2
3use std::cmp;
4use std::future::Future;
5use std::io;
6use std::io::prelude::*;
7use std::pin::Pin;
8use std::task::{ready, Context, Poll};
9use tokio::task::{spawn_blocking, JoinHandle};
10use crate::buf::{Buf, DEFAULT_MAX_BUF_SIZE};
11
12#[derive(Debug)]
13pub struct ReadBlocking<T> {
14    inner: Option<T>,
15    state: State<T>,
16}
17
18#[derive(Debug)]
19pub struct WriteBlocking<T> {
20    inner: Option<T>,
21    state: State<T>,
22    flush_state: FlushState<T>,
23}
24
25#[derive(Debug)]
26enum State<T> {
27    Idle(Option<Buf>),
28    Busy(JoinHandle<(io::Result<usize>, Buf, T)>),
29}
30
31#[derive(Debug)]
32enum FlushState<T> {
33    Idle,
34    Busy(JoinHandle<(io::Result<()>, T)>),
35}
36
37impl<T> ReadBlocking<T> {
38    pub fn new(inner: T) -> ReadBlocking<T> {
39        ReadBlocking {
40            inner: Some(inner),
41            state: State::Idle(Some(Buf::with_capacity(0))),
42        }
43    }
44}
45
46impl<T> WriteBlocking<T> {
47    pub fn new(inner: T) -> WriteBlocking<T> {
48        WriteBlocking {
49            inner: Some(inner),
50            state: State::Idle(Some(Buf::with_capacity(0))),
51            flush_state: FlushState::Idle,
52        }
53    }
54}
55
56impl<T> AsyncRead for ReadBlocking<T>
57where
58    T: Read + Unpin + Send + 'static,
59{
60    fn poll_read(
61        mut self: Pin<&mut Self>,
62        cx: &mut Context<'_>,
63        dst: &mut ReadBuf<'_>,
64    ) -> Poll<io::Result<()>> {
65        loop {
66            match self.state {
67                State::Idle(ref mut buf_cell) => {
68                    let mut buf = buf_cell.take().unwrap();
69
70                    if !buf.is_empty() {
71                        buf.copy_to(dst);
72                        *buf_cell = Some(buf);
73                        return Poll::Ready(Ok(()));
74                    }
75
76                    let mut inner = self.inner.take().unwrap();
77
78                    let max_buf_size = cmp::min(dst.remaining(), DEFAULT_MAX_BUF_SIZE);
79                    self.state = State::Busy(spawn_blocking(move || {
80                        // SAFETY: the requirements are satisfied by `Blocking::new`.
81                        let res = unsafe { buf.read_from(&mut inner, max_buf_size) };
82                        (res, buf, inner)
83                    }));
84                }
85                State::Busy(ref mut rx) => {
86                    let (res, mut buf, inner) = ready!(Pin::new(rx).poll(cx))?;
87                    self.inner = Some(inner);
88
89                    match res {
90                        Ok(_) => {
91                            buf.copy_to(dst);
92                            self.state = State::Idle(Some(buf));
93                            return Poll::Ready(Ok(()));
94                        }
95                        Err(e) => {
96                            assert!(buf.is_empty());
97
98                            self.state = State::Idle(Some(buf));
99                            return Poll::Ready(Err(e));
100                        }
101                    }
102                }
103            }
104        }
105    }
106}
107
108impl<T> AsyncWrite for WriteBlocking<T>
109where
110    T: Write + Unpin + Send + 'static,
111{
112    fn poll_write(
113        mut self: Pin<&mut Self>,
114        cx: &mut Context<'_>,
115        src: &[u8],
116    ) -> Poll<io::Result<usize>> {
117        loop {
118            match self.state {
119                State::Idle(ref mut buf_cell) => {
120                    let mut buf = buf_cell.take().unwrap();
121
122                    assert!(buf.is_empty());
123
124                    let _n = buf.copy_from(src, DEFAULT_MAX_BUF_SIZE);
125                    let mut inner = self.inner.take().unwrap();
126
127                    self.state = State::Busy(spawn_blocking(move || {
128                        let n = buf.len();
129                        let res = buf.write_to(&mut inner).map(|()| n);
130
131                        (res, buf, inner)
132                    }));
133                }
134                State::Busy(ref mut rx) => {
135                    let (res, buf, inner) = ready!(Pin::new(rx).poll(cx))?;
136                    self.state = State::Idle(Some(buf));
137                    self.inner = Some(inner);
138
139                    // If error, return
140                    return Poll::Ready(res);
141                }
142            }
143        }
144    }
145
146    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
147        loop {
148            match self.flush_state {
149                // The buffer is not used here
150                FlushState::Idle => {
151                    let mut inner = self.inner.take().unwrap();
152
153                    self.flush_state = FlushState::Busy(spawn_blocking(move || {
154                        let res = inner.flush();
155                        (res, inner)
156                    }));
157                }
158                FlushState::Busy(ref mut rx) => {
159                    let (res, inner) = ready!(Pin::new(rx).poll(cx))?;
160                    self.flush_state = FlushState::Idle;
161                    self.inner = Some(inner);
162
163                    return Poll::Ready(res);
164                }
165            }
166        }
167    }
168
169    fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
170        Poll::Ready(Ok(()))
171    }
172}