s3_wire/stream/
download.rs1use 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
17pub 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 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 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}