1use std::{
8 io::{self, Read, Write},
9 time::Duration,
10};
11
12use futures_util::{Stream, StreamExt};
13use tokio::{runtime::Handle, sync::mpsc, time::timeout};
14use tokio_util::sync::CancellationToken;
15
16pub const STALL: Duration = Duration::from_secs(60);
18
19pub const CHUNK: usize = 256 * 1024;
21
22pub const CHANNEL_CHUNKS: usize = 4;
24
25fn cancelled() -> io::Error {
27 io::Error::other("the transfer was cancelled")
28}
29
30fn stalled(stall: Duration) -> io::Error {
31 io::Error::new(
32 io::ErrorKind::TimedOut,
33 format!("no data moved for {} s", stall.as_secs()),
34 )
35}
36
37pub struct SyncReader<S, B> {
39 handle: Handle,
40 stream: S,
41 chunk: Option<B>,
42 pos: usize,
43 cancel: CancellationToken,
44 stall: Duration,
45}
46
47impl<S, B> SyncReader<S, B>
48where
49 S: Stream<Item = io::Result<B>> + Unpin,
50 B: AsRef<[u8]>,
51{
52 #[must_use]
54 pub fn new(stream: S, cancel: CancellationToken, stall: Duration) -> Self {
55 Self {
56 handle: Handle::current(),
57 stream,
58 chunk: None,
59 pos: 0,
60 cancel,
61 stall,
62 }
63 }
64}
65
66impl<S, B> Read for SyncReader<S, B>
67where
68 S: Stream<Item = io::Result<B>> + Unpin,
69 B: AsRef<[u8]>,
70{
71 fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
72 loop {
73 if let Some(chunk) = &self.chunk {
74 let rest = &chunk.as_ref()[self.pos..];
75 if !rest.is_empty() {
76 let n = rest.len().min(buf.len());
77 buf[..n].copy_from_slice(&rest[..n]);
78 self.pos += n;
79 return Ok(n);
80 }
81 }
82 let next = self.handle.block_on(async {
83 tokio::select! {
84 biased;
85 () = self.cancel.cancelled() => Err(cancelled()),
86 next = timeout(self.stall, self.stream.next()) => {
87 next.map_err(|_| stalled(self.stall))
88 }
89 }
90 })?;
91 match next {
92 Some(chunk) => {
93 self.chunk = Some(chunk?);
94 self.pos = 0;
95 }
96 None => return Ok(0),
97 }
98 }
99 }
100}
101
102pub struct SyncWriter {
104 handle: Handle,
105 tx: mpsc::Sender<io::Result<Vec<u8>>>,
106 cancel: CancellationToken,
107 stall: Duration,
108}
109
110impl SyncWriter {
111 #[must_use]
113 pub fn new(
114 tx: mpsc::Sender<io::Result<Vec<u8>>>,
115 cancel: CancellationToken,
116 stall: Duration,
117 ) -> Self {
118 Self {
119 handle: Handle::current(),
120 tx,
121 cancel,
122 stall,
123 }
124 }
125
126 pub fn fail(&self, error: io::Error) {
128 self.handle.block_on(async {
129 tokio::select! {
130 biased;
131 () = self.cancel.cancelled() => {}
132 _ = timeout(self.stall, self.tx.send(Err(error))) => {}
133 }
134 });
135 }
136}
137
138impl Write for SyncWriter {
139 fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
140 if buf.is_empty() {
141 return Ok(0);
142 }
143 self.handle.block_on(async {
144 tokio::select! {
145 biased;
146 () = self.cancel.cancelled() => Err(cancelled()),
147 sent = timeout(self.stall, self.tx.send(Ok(buf.to_vec()))) => match sent {
148 Err(_) => Err(stalled(self.stall)),
149 Ok(Err(_)) => Err(io::Error::new(
150 io::ErrorKind::BrokenPipe,
151 "the other side stopped reading",
152 )),
153 Ok(Ok(())) => Ok(()),
154 },
155 }
156 })?;
157 Ok(buf.len())
158 }
159
160 fn flush(&mut self) -> io::Result<()> {
161 Ok(())
162 }
163}
164
165#[cfg(test)]
166mod tests {
167 use futures_util::stream;
168
169 use super::*;
170
171 #[tokio::test]
172 async fn reader_hands_over_chunks_in_order_across_small_reads() {
173 let chunks = vec![
174 Ok(b"hello ".to_vec()),
175 Ok(b"wor".to_vec()),
176 Ok(b"ld".to_vec()),
177 ];
178 let reader = SyncReader::new(
179 Box::pin(stream::iter(chunks)),
180 CancellationToken::new(),
181 STALL,
182 );
183 let text = tokio::task::spawn_blocking(move || {
184 let mut text = String::new();
185 let mut reader = io::BufReader::with_capacity(4, reader);
186 reader.read_to_string(&mut text).unwrap();
187 text
188 })
189 .await
190 .unwrap();
191 assert_eq!(text, "hello world");
192 }
193
194 #[tokio::test]
195 async fn reader_fails_when_no_bytes_arrive_for_the_stall_time() {
196 let pending = stream::pending::<io::Result<Vec<u8>>>();
197 let mut reader = SyncReader::new(
198 Box::pin(pending),
199 CancellationToken::new(),
200 Duration::from_millis(50),
201 );
202 let error = tokio::task::spawn_blocking(move || reader.read(&mut [0; 8]).unwrap_err())
203 .await
204 .unwrap();
205 assert_eq!(error.kind(), io::ErrorKind::TimedOut);
206 }
207
208 #[tokio::test]
209 async fn reader_stops_when_cancelled_while_waiting() {
210 let cancel = CancellationToken::new();
211 let mut reader = SyncReader::new(
212 Box::pin(stream::pending::<io::Result<Vec<u8>>>()),
213 cancel.clone(),
214 STALL,
215 );
216 let read = tokio::task::spawn_blocking(move || reader.read(&mut [0; 8]).unwrap_err());
217 cancel.cancel();
218 let error = read.await.unwrap();
219 assert_eq!(error.kind(), io::ErrorKind::Other);
220 assert_eq!(error.to_string(), "the transfer was cancelled");
221 }
222
223 #[tokio::test]
224 async fn io_copy_gives_up_on_a_cancelled_reader_instead_of_retrying() {
225 let cancel = CancellationToken::new();
226 cancel.cancel();
227 let mut reader = SyncReader::new(
228 Box::pin(stream::pending::<io::Result<Vec<u8>>>()),
229 cancel,
230 STALL,
231 );
232 let copy = tokio::task::spawn_blocking(move || io::copy(&mut reader, &mut io::sink()));
233 let finished = timeout(Duration::from_secs(5), copy).await;
234 assert!(finished.unwrap().unwrap().is_err());
235 }
236
237 #[tokio::test]
238 async fn writer_fails_when_nobody_reads_for_the_stall_time() {
239 let (tx, _rx) = mpsc::channel(1);
240 let mut writer = SyncWriter::new(tx, CancellationToken::new(), Duration::from_millis(50));
241 let error = tokio::task::spawn_blocking(move || {
242 writer.write_all(b"one").unwrap();
243 writer.write_all(b"two").unwrap_err()
244 })
245 .await
246 .unwrap();
247 assert_eq!(error.kind(), io::ErrorKind::TimedOut);
248 }
249
250 #[tokio::test]
251 async fn writer_stops_when_the_receiver_is_gone() {
252 let (tx, rx) = mpsc::channel(1);
253 drop(rx);
254 let mut writer = SyncWriter::new(tx, CancellationToken::new(), STALL);
255 let error = tokio::task::spawn_blocking(move || writer.write_all(b"x").unwrap_err())
256 .await
257 .unwrap();
258 assert_eq!(error.kind(), io::ErrorKind::BrokenPipe);
259 }
260}