Skip to main content

computer_transfer/
pipe.rs

1//! Blocking `Read` and `Write` ends of async streams.
2//!
3//! The tar code is synchronous and runs in `spawn_blocking`. These types move its bytes through
4//! async streams one chunk at a time, so memory stays bounded. Every wait ends when the transfer
5//! is cancelled or when no bytes moved for the stall time.
6
7use 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
16/// Longest a transfer waits for bytes to move before it fails.
17pub const STALL: Duration = Duration::from_secs(60);
18
19/// Size of the chunks a [`SyncWriter`] sends.
20pub const CHUNK: usize = 256 * 1024;
21
22/// Chunks a transfer channel holds before the writer waits.
23pub const CHANNEL_CHUNKS: usize = 4;
24
25/// `io::copy` and `write_all` retry on `ErrorKind::Interrupted`, so a cancelled transfer must not use it.
26fn 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
37/// Reads the chunks of an async stream.
38pub 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    /// Call from async code, then move the reader into `spawn_blocking`.
53    #[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
102/// Sends what is written to a bounded channel, in chunks.
103pub struct SyncWriter {
104    handle: Handle,
105    tx: mpsc::Sender<io::Result<Vec<u8>>>,
106    cancel: CancellationToken,
107    stall: Duration,
108}
109
110impl SyncWriter {
111    /// Call from async code, then move the writer into `spawn_blocking`.
112    #[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    /// Sends a failure after the last chunk so the receiver's stream ends with an error, not early.
127    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}