Skip to main content

s3_wire/stream/
download.rs

1use std::fmt;
2use std::pin::Pin;
3use std::task::{Context, Poll};
4use std::time::Duration;
5
6use bytes::Bytes;
7use futures_core::Stream;
8use futures_util::StreamExt;
9use sha2::{Digest, Sha256};
10use tokio::io::{AsyncWrite, AsyncWriteExt};
11use tokio::time::{Instant, Sleep};
12
13use crate::error::S3Error;
14
15type DownloadItems = Pin<Box<dyn Stream<Item = Result<Bytes, S3Error>> + Send + 'static>>;
16
17/// A download stream that enforces idle timeout, content length, and any
18/// supported full-object checksum returned by S3.
19pub struct ResponseStream {
20    inner: DownloadItems,
21    expected_length: Option<u64>,
22    expected_sha256: Option<[u8; 32]>,
23    sha256: Option<Sha256>,
24    received: u64,
25    idle_timeout: Duration,
26    idle: Pin<Box<Sleep>>,
27    deadline: Option<Pin<Box<Sleep>>>,
28    finished: bool,
29}
30
31impl ResponseStream {
32    pub(crate) fn new<S>(stream: S, expected_length: Option<u64>, idle_timeout: Duration) -> Self
33    where
34        S: Stream<Item = Result<Bytes, S3Error>> + Send + 'static,
35    {
36        Self {
37            inner: Box::pin(stream),
38            expected_length,
39            expected_sha256: None,
40            sha256: None,
41            received: 0,
42            idle_timeout,
43            idle: Box::pin(tokio::time::sleep(idle_timeout)),
44            deadline: None,
45            finished: false,
46        }
47    }
48
49    pub(crate) fn with_deadline<S>(
50        stream: S,
51        expected_length: Option<u64>,
52        expected_sha256: Option<[u8; 32]>,
53        idle_timeout: Duration,
54        deadline: Instant,
55    ) -> Self
56    where
57        S: Stream<Item = Result<Bytes, S3Error>> + Send + 'static,
58    {
59        let mut response = Self::new(stream, expected_length, idle_timeout);
60        response.expected_sha256 = expected_sha256;
61        response.sha256 = expected_sha256.map(|_| Sha256::new());
62        response.deadline = Some(Box::pin(tokio::time::sleep_until(deadline)));
63        response
64    }
65
66    /// Streams the remaining body into `writer` with backpressure.
67    ///
68    /// # Errors
69    ///
70    /// Returns an error for transport, timeout, output, length, or integrity
71    /// failures encountered while streaming the response.
72    pub async fn write_to<W>(mut self, writer: &mut W) -> Result<u64, S3Error>
73    where
74        W: AsyncWrite + Unpin,
75    {
76        let mut written = 0_u64;
77        while let Some(chunk) = self.next().await {
78            let chunk = chunk?;
79            writer.write_all(&chunk).await.map_err(S3Error::transport)?;
80            written =
81                written
82                    .checked_add(u64::try_from(chunk.len()).map_err(|_| {
83                        S3Error::integrity("download chunk length does not fit in u64")
84                    })?)
85                    .ok_or_else(|| S3Error::integrity("download length overflow"))?;
86        }
87        writer.flush().await.map_err(S3Error::transport)?;
88        Ok(written)
89    }
90
91    /// Returns the number of bytes yielded so far.
92    pub fn bytes_received(&self) -> u64 {
93        self.received
94    }
95}
96
97impl fmt::Debug for ResponseStream {
98    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
99        formatter
100            .debug_struct("ResponseStream")
101            .field("expected_length", &self.expected_length)
102            .field("verifies_sha256", &self.expected_sha256.is_some())
103            .field("received", &self.received)
104            .field("idle_timeout", &self.idle_timeout)
105            .field("has_deadline", &self.deadline.is_some())
106            .field("finished", &self.finished)
107            .finish_non_exhaustive()
108    }
109}
110
111impl Stream for ResponseStream {
112    type Item = Result<Bytes, S3Error>;
113
114    fn poll_next(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
115        if self.finished {
116            return Poll::Ready(None);
117        }
118
119        if self
120            .deadline
121            .as_mut()
122            .is_some_and(|deadline| deadline.as_mut().poll(context).is_ready())
123        {
124            self.finished = true;
125            return Poll::Ready(Some(Err(S3Error::timeout(
126                crate::error::TimeoutPhase::Operation,
127                "operation deadline expired while streaming the response body",
128            ))));
129        }
130
131        if self.idle.as_mut().poll(context).is_ready() {
132            self.finished = true;
133            return Poll::Ready(Some(Err(S3Error::timeout(
134                crate::error::TimeoutPhase::ResponseBody,
135                "download body was idle past its configured timeout",
136            ))));
137        }
138
139        match self.inner.as_mut().poll_next(context) {
140            Poll::Ready(Some(Ok(chunk))) => {
141                if chunk.is_empty() {
142                    context.waker().wake_by_ref();
143                    return Poll::Pending;
144                }
145                let idle_timeout = self.idle_timeout;
146                self.idle.as_mut().reset(Instant::now() + idle_timeout);
147                let Ok(chunk_length) = u64::try_from(chunk.len()) else {
148                    self.finished = true;
149                    return Poll::Ready(Some(Err(S3Error::integrity(
150                        "download chunk length does not fit in u64",
151                    ))));
152                };
153                let Some(received) = self.received.checked_add(chunk_length) else {
154                    self.finished = true;
155                    return Poll::Ready(Some(Err(S3Error::integrity("download length overflow"))));
156                };
157                self.received = received;
158                if self
159                    .expected_length
160                    .is_some_and(|expected| self.received > expected)
161                {
162                    self.finished = true;
163                    return Poll::Ready(Some(Err(S3Error::integrity(
164                        "download exceeded the declared content length",
165                    ))));
166                }
167                if let Some(hasher) = &mut self.sha256 {
168                    hasher.update(&chunk);
169                }
170                Poll::Ready(Some(Ok(chunk)))
171            }
172            Poll::Ready(Some(Err(error))) => {
173                self.finished = true;
174                Poll::Ready(Some(Err(error)))
175            }
176            Poll::Ready(None) => {
177                self.finished = true;
178                if self
179                    .expected_length
180                    .is_some_and(|expected| self.received != expected)
181                {
182                    Poll::Ready(Some(Err(S3Error::integrity(
183                        "download ended before the declared content length",
184                    ))))
185                } else if let Some(expected) = self.expected_sha256 {
186                    let actual: [u8; 32] = self
187                        .sha256
188                        .take()
189                        .expect("SHA-256 state accompanies an expected digest")
190                        .finalize()
191                        .into();
192                    if actual == expected {
193                        Poll::Ready(None)
194                    } else {
195                        Poll::Ready(Some(Err(S3Error::integrity(
196                            "download bytes did not match the returned SHA-256 checksum",
197                        ))))
198                    }
199                } else {
200                    Poll::Ready(None)
201                }
202            }
203            Poll::Pending => Poll::Pending,
204        }
205    }
206}
207
208#[cfg(test)]
209mod tests {
210    use bytes::Bytes;
211    use futures_util::{StreamExt, stream};
212    use sha2::{Digest, Sha256};
213
214    use super::ResponseStream;
215
216    #[tokio::test]
217    async fn response_detects_truncation() {
218        let source = stream::iter([Ok(Bytes::from_static(b"abc"))]);
219        let mut body = ResponseStream::new(source, Some(4), std::time::Duration::from_secs(1));
220        assert!(body.next().await.expect("chunk").is_ok());
221        assert!(body.next().await.expect("integrity result").is_err());
222        assert!(body.next().await.is_none());
223    }
224
225    #[tokio::test(start_paused = true)]
226    async fn empty_download_chunks_do_not_count_as_progress() {
227        let source = stream::once(async { Ok(Bytes::new()) })
228            .chain(stream::pending::<Result<Bytes, crate::error::S3Error>>());
229        let mut body = ResponseStream::new(source, None, std::time::Duration::from_secs(2));
230        let next = tokio::spawn(async move { body.next().await });
231
232        tokio::time::advance(std::time::Duration::from_secs(2)).await;
233        let error = next.await.unwrap().unwrap().unwrap_err();
234        assert_eq!(
235            error.timeout_phase(),
236            Some(crate::error::TimeoutPhase::ResponseBody)
237        );
238    }
239
240    #[tokio::test(start_paused = true)]
241    async fn response_enforces_operation_deadline() {
242        let source = stream::pending::<Result<Bytes, crate::error::S3Error>>();
243        let mut body = ResponseStream::with_deadline(
244            source,
245            None,
246            None,
247            std::time::Duration::from_secs(60),
248            tokio::time::Instant::now() + std::time::Duration::from_secs(2),
249        );
250        let next = tokio::spawn(async move { body.next().await });
251        tokio::time::advance(std::time::Duration::from_secs(2)).await;
252        let error = next.await.unwrap().unwrap().unwrap_err();
253        assert_eq!(
254            error.timeout_phase(),
255            Some(crate::error::TimeoutPhase::Operation)
256        );
257    }
258
259    #[tokio::test]
260    async fn response_verifies_expected_sha256_at_end_of_stream() {
261        let source = stream::iter([Ok(Bytes::from_static(b"abc"))]);
262        let mut body = ResponseStream::with_deadline(
263            source,
264            Some(3),
265            Some(Sha256::digest(b"abc").into()),
266            std::time::Duration::from_secs(1),
267            tokio::time::Instant::now() + std::time::Duration::from_secs(1),
268        );
269        assert_eq!(
270            body.next().await.unwrap().unwrap(),
271            Bytes::from_static(b"abc")
272        );
273        assert!(body.next().await.is_none());
274
275        let source = stream::iter([Ok(Bytes::from_static(b"abc"))]);
276        let mut body = ResponseStream::with_deadline(
277            source,
278            Some(3),
279            Some(Sha256::digest(b"different").into()),
280            std::time::Duration::from_secs(1),
281            tokio::time::Instant::now() + std::time::Duration::from_secs(1),
282        );
283        assert!(body.next().await.unwrap().is_ok());
284        let error = body.next().await.unwrap().unwrap_err();
285        assert_eq!(error.category(), crate::error::ErrorCategory::Integrity);
286        assert_eq!(
287            error.retry_classification(),
288            crate::error::RetryClassification::Never
289        );
290    }
291}