Skip to main content

s3_wire/stream/
upload.rs

1use std::fmt;
2use std::io;
3use std::path::{Path, PathBuf};
4use std::pin::Pin;
5use std::task::{Context, Poll};
6
7use base64::Engine as _;
8use bytes::Bytes;
9use futures_core::Stream;
10use hyper::body::{Body, Frame, SizeHint};
11use sha2::{Digest, Sha256};
12use tokio::io::AsyncReadExt;
13use tokio::sync::Mutex;
14use tokio_util::io::ReaderStream;
15
16use crate::error::S3Error;
17use crate::stream::FileSnapshot;
18
19type UploadItems = Pin<Box<dyn Stream<Item = Result<Bytes, io::Error>> + Send + 'static>>;
20
21/// An upload body with explicit length, integrity, and replay semantics.
22///
23/// Byte and file bodies are replayable. A caller-provided stream is one-shot
24/// and requires both its exact content length and SHA-256 digest so the client
25/// never buffers it or silently switches to unsigned payloads.
26pub struct ByteStream {
27    source: UploadSource,
28}
29
30enum UploadSource {
31    Bytes(Bytes),
32    File(PathBuf),
33    Stream {
34        stream: UploadItems,
35        length: u64,
36        sha256: [u8; 32],
37    },
38}
39
40impl ByteStream {
41    /// Creates a replayable in-memory body.
42    pub fn from_bytes(bytes: impl Into<Bytes>) -> Self {
43        Self {
44            source: UploadSource::Bytes(bytes.into()),
45        }
46    }
47
48    /// Creates a replayable file body.
49    ///
50    /// The explicitly supplied file is copied into a private disk-backed snapshot
51    /// while it is hashed. Every retry reads that immutable snapshot, so changes
52    /// to the original path cannot alter the signed upload bytes.
53    pub fn from_path(path: impl AsRef<Path>) -> Self {
54        Self {
55            source: UploadSource::File(path.as_ref().to_owned()),
56        }
57    }
58
59    /// Creates a one-shot stream with a caller-supplied exact digest.
60    ///
61    /// The stream is never retried. `sha256` is the digest of exactly `length`
62    /// bytes and is used for `SigV4` payload signing.
63    pub fn from_stream<S>(stream: S, length: u64, sha256: [u8; 32]) -> Self
64    where
65        S: Stream<Item = Result<Bytes, io::Error>> + Send + 'static,
66    {
67        Self {
68            source: UploadSource::Stream {
69                stream: Box::pin(stream),
70                length,
71                sha256,
72            },
73        }
74    }
75
76    pub(crate) async fn prepare(self) -> Result<PreparedBody, S3Error> {
77        match self.source {
78            UploadSource::Bytes(bytes) => {
79                let sha256 = cooperative_sha256(&bytes).await;
80                let length = u64::try_from(bytes.len()).map_err(|_| {
81                    S3Error::configuration("in-memory body length does not fit in u64")
82                })?;
83                Ok(PreparedBody {
84                    source: PreparedSource::Bytes(bytes),
85                    length,
86                    sha256,
87                })
88            }
89            UploadSource::File(path) => prepare_file(path).await,
90            UploadSource::Stream {
91                stream,
92                length,
93                sha256,
94            } => Ok(PreparedBody {
95                source: PreparedSource::OneShot(Mutex::new(Some(stream))),
96                length,
97                sha256,
98            }),
99        }
100    }
101}
102
103async fn cooperative_sha256(bytes: &[u8]) -> [u8; 32] {
104    const CHUNK_SIZE: usize = 1024 * 1024;
105
106    let mut hasher = Sha256::new();
107    for chunk in bytes.chunks(CHUNK_SIZE) {
108        hasher.update(chunk);
109        tokio::task::yield_now().await;
110    }
111    hasher.finalize().into()
112}
113
114impl From<Bytes> for ByteStream {
115    fn from(value: Bytes) -> Self {
116        Self::from_bytes(value)
117    }
118}
119
120impl From<Vec<u8>> for ByteStream {
121    fn from(value: Vec<u8>) -> Self {
122        Self::from_bytes(value)
123    }
124}
125
126impl From<&'static [u8]> for ByteStream {
127    fn from(value: &'static [u8]) -> Self {
128        Self::from_bytes(Bytes::from_static(value))
129    }
130}
131
132impl fmt::Debug for ByteStream {
133    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
134        let (kind, replayable) = match self.source {
135            UploadSource::Bytes(_) => ("bytes", true),
136            UploadSource::File(_) => ("file", true),
137            UploadSource::Stream { .. } => ("stream", false),
138        };
139        formatter
140            .debug_struct("ByteStream")
141            .field("kind", &kind)
142            .field("replayable", &replayable)
143            .finish_non_exhaustive()
144    }
145}
146
147async fn prepare_file(path: PathBuf) -> Result<PreparedBody, S3Error> {
148    // Every attempt opens the same private snapshot, so later changes to the
149    // caller's path cannot change already signed bytes.
150    let snapshot = FileSnapshot::create(path, true).await?;
151    let length = snapshot.length();
152    let sha256 = snapshot
153        .sha256()
154        .ok_or_else(|| S3Error::integrity("upload snapshot digest was not calculated"))?;
155
156    Ok(PreparedBody {
157        source: PreparedSource::FileSnapshot(snapshot),
158        length,
159        sha256,
160    })
161}
162
163pub(crate) struct PreparedBody {
164    source: PreparedSource,
165    length: u64,
166    sha256: [u8; 32],
167}
168
169enum PreparedSource {
170    Bytes(Bytes),
171    FileSnapshot(FileSnapshot),
172    OneShot(Mutex<Option<UploadItems>>),
173}
174
175impl PreparedBody {
176    pub(crate) fn length(&self) -> u64 {
177        self.length
178    }
179
180    pub(crate) fn sha256_hex(&self) -> String {
181        encode_hex(&self.sha256)
182    }
183
184    pub(crate) fn sha256_base64(&self) -> String {
185        base64::engine::general_purpose::STANDARD.encode(self.sha256)
186    }
187
188    pub(crate) fn is_replayable(&self) -> bool {
189        !matches!(self.source, PreparedSource::OneShot(_))
190    }
191
192    pub(crate) async fn request_body(&self) -> Result<TransportBody, S3Error> {
193        match &self.source {
194            PreparedSource::Bytes(bytes) => Ok(TransportBody::new(
195                Box::pin(futures_util::stream::once(std::future::ready(Ok(
196                    bytes.clone()
197                )))),
198                self.length,
199                self.sha256,
200            )),
201            PreparedSource::FileSnapshot(snapshot) => {
202                let file = snapshot.open().await?;
203                let stream = ReaderStream::with_capacity(file.take(self.length), 64 * 1024);
204                Ok(TransportBody::new(
205                    Box::pin(stream),
206                    self.length,
207                    self.sha256,
208                ))
209            }
210            PreparedSource::OneShot(stream) => {
211                let stream = stream.lock().await.take().ok_or_else(|| {
212                    S3Error::unsupported("a non-replayable upload body cannot be sent twice")
213                })?;
214                Ok(TransportBody::new(stream, self.length, self.sha256))
215            }
216        }
217    }
218}
219
220/// The crate-private HTTP body used by the transport.
221pub(crate) struct TransportBody {
222    stream: UploadItems,
223    remaining: u64,
224    expected_sha256: [u8; 32],
225    hasher: Sha256,
226    final_chunk: Option<Bytes>,
227    finished: bool,
228}
229
230impl TransportBody {
231    fn new(stream: UploadItems, length: u64, expected_sha256: [u8; 32]) -> Self {
232        Self {
233            stream,
234            remaining: length,
235            expected_sha256,
236            hasher: Sha256::new(),
237            final_chunk: None,
238            finished: false,
239        }
240    }
241
242    pub(crate) fn empty() -> Self {
243        let digest = Sha256::digest([]).into();
244        Self::new(Box::pin(futures_util::stream::empty()), 0, digest)
245    }
246}
247
248impl Body for TransportBody {
249    type Data = Bytes;
250    type Error = S3Error;
251
252    fn poll_frame(
253        mut self: Pin<&mut Self>,
254        context: &mut Context<'_>,
255    ) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
256        if self.finished {
257            return Poll::Ready(None);
258        }
259
260        loop {
261            if self.final_chunk.is_some() {
262                match self.stream.as_mut().poll_next(context) {
263                    Poll::Ready(Some(Ok(chunk))) if chunk.is_empty() => continue,
264                    Poll::Ready(Some(Ok(_))) => {
265                        self.finished = true;
266                        self.remaining = 0;
267                        self.final_chunk = None;
268                        return Poll::Ready(Some(Err(S3Error::integrity(
269                            "upload exceeded its declared content length",
270                        ))));
271                    }
272                    Poll::Ready(Some(Err(error))) => {
273                        self.finished = true;
274                        self.remaining = 0;
275                        self.final_chunk = None;
276                        return Poll::Ready(Some(Err(S3Error::transport(error))));
277                    }
278                    Poll::Ready(None) => {
279                        let digest: [u8; 32] = self.hasher.clone().finalize().into();
280                        if digest != self.expected_sha256 {
281                            self.finished = true;
282                            self.remaining = 0;
283                            self.final_chunk = None;
284                            return Poll::Ready(Some(Err(S3Error::integrity(
285                                "upload bytes did not match the signed payload digest",
286                            ))));
287                        }
288                        self.finished = true;
289                        self.remaining = 0;
290                        let chunk = self
291                            .final_chunk
292                            .take()
293                            .expect("final upload chunk was established");
294                        return Poll::Ready(Some(Ok(Frame::data(chunk))));
295                    }
296                    Poll::Pending => return Poll::Pending,
297                }
298            }
299
300            match self.stream.as_mut().poll_next(context) {
301                Poll::Ready(Some(Ok(chunk))) => {
302                    if chunk.is_empty() {
303                        continue;
304                    }
305                    let Ok(chunk_length) = u64::try_from(chunk.len()) else {
306                        self.finished = true;
307                        self.remaining = 0;
308                        return Poll::Ready(Some(Err(S3Error::integrity(
309                            "upload chunk length does not fit in u64",
310                        ))));
311                    };
312                    if chunk_length > self.remaining {
313                        self.finished = true;
314                        self.remaining = 0;
315                        return Poll::Ready(Some(Err(S3Error::integrity(
316                            "upload exceeded its declared content length",
317                        ))));
318                    }
319                    self.hasher.update(&chunk);
320                    if chunk_length == self.remaining {
321                        self.final_chunk = Some(chunk);
322                        continue;
323                    }
324                    self.remaining -= chunk_length;
325                    return Poll::Ready(Some(Ok(Frame::data(chunk))));
326                }
327                Poll::Ready(Some(Err(error))) => {
328                    self.finished = true;
329                    self.remaining = 0;
330                    return Poll::Ready(Some(Err(S3Error::transport(error))));
331                }
332                Poll::Ready(None) => {
333                    self.finished = true;
334                    if self.remaining != 0 {
335                        self.remaining = 0;
336                        return Poll::Ready(Some(Err(S3Error::integrity(
337                            "upload ended before its declared content length",
338                        ))));
339                    }
340                    let digest: [u8; 32] = self.hasher.clone().finalize().into();
341                    return if digest == self.expected_sha256 {
342                        Poll::Ready(None)
343                    } else {
344                        Poll::Ready(Some(Err(S3Error::integrity(
345                            "upload bytes did not match the signed payload digest",
346                        ))))
347                    };
348                }
349                Poll::Pending => return Poll::Pending,
350            }
351        }
352    }
353
354    fn is_end_stream(&self) -> bool {
355        self.finished
356    }
357
358    fn size_hint(&self) -> SizeHint {
359        SizeHint::with_exact(self.remaining)
360    }
361}
362
363fn encode_hex(bytes: &[u8]) -> String {
364    const HEX: &[u8; 16] = b"0123456789abcdef";
365    let mut output = String::with_capacity(bytes.len() * 2);
366    for byte in bytes {
367        output.push(char::from(HEX[usize::from(byte >> 4)]));
368        output.push(char::from(HEX[usize::from(byte & 0x0f)]));
369    }
370    output
371}
372
373#[cfg(test)]
374mod tests {
375    use std::future::{Future as _, poll_fn};
376    use std::io;
377    use std::task::Poll;
378
379    use bytes::Bytes;
380    use futures_util::stream;
381    use http_body_util::BodyExt;
382    use sha2::{Digest, Sha256};
383    use tokio::io::AsyncWriteExt;
384
385    use super::ByteStream;
386
387    #[tokio::test]
388    async fn bytes_are_replayable_and_hashed() {
389        let prepared = ByteStream::from_bytes(Bytes::from_static(b"hello"))
390            .prepare()
391            .await
392            .expect("body prepares");
393        assert!(prepared.is_replayable());
394        assert_eq!(prepared.length(), 5);
395        assert_eq!(
396            prepared.sha256_hex(),
397            "2cf24dba5fb0a30e26e83b2ac5b9e29e1b161e5c1fa7425e73043362938b9824"
398        );
399        let first = prepared
400            .request_body()
401            .await
402            .expect("first body")
403            .collect()
404            .await
405            .expect("first body is valid")
406            .to_bytes();
407        let second = prepared
408            .request_body()
409            .await
410            .expect("second body")
411            .collect()
412            .await
413            .expect("second body is valid")
414            .to_bytes();
415        assert_eq!(first, Bytes::from_static(b"hello"));
416        assert_eq!(second, first);
417    }
418
419    #[tokio::test]
420    async fn hashing_large_in_memory_bodies_is_cooperative() {
421        let mut preparation =
422            Box::pin(ByteStream::from_bytes(vec![0_u8; 2 * 1024 * 1024]).prepare());
423
424        poll_fn(|context| match preparation.as_mut().poll(context) {
425            Poll::Pending => Poll::Ready(()),
426            Poll::Ready(_) => panic!("large body hashing completed in one cooperative poll"),
427        })
428        .await;
429        preparation.await.expect("body finishes preparing");
430    }
431
432    #[tokio::test]
433    async fn one_shot_stream_rejects_second_attempt() {
434        let sha256 = Sha256::digest(b"x").into();
435        let source = stream::iter([Ok::<_, io::Error>(Bytes::from_static(b"x"))]);
436        let prepared = ByteStream::from_stream(source, 1, sha256)
437            .prepare()
438            .await
439            .expect("body prepares");
440        assert!(!prepared.is_replayable());
441        prepared
442            .request_body()
443            .await
444            .expect("first body")
445            .collect()
446            .await
447            .expect("first body is valid");
448        assert!(prepared.request_body().await.is_err());
449    }
450
451    #[tokio::test]
452    async fn file_retries_use_the_immutable_hashed_snapshot() {
453        let mut original = tempfile::NamedTempFile::new().unwrap();
454        std::io::Write::write_all(&mut original, b"signed bytes").unwrap();
455        std::io::Write::flush(&mut original).unwrap();
456        let prepared = ByteStream::from_path(original.path())
457            .prepare()
458            .await
459            .expect("file body prepares");
460
461        let mut changed = tokio::fs::File::create(original.path()).await.unwrap();
462        changed.write_all(b"different bytes").await.unwrap();
463        changed.flush().await.unwrap();
464
465        for _ in 0..2 {
466            let body = prepared
467                .request_body()
468                .await
469                .expect("snapshot opens")
470                .collect()
471                .await
472                .expect("snapshot body is valid")
473                .to_bytes();
474            assert_eq!(body, Bytes::from_static(b"signed bytes"));
475        }
476    }
477
478    #[tokio::test]
479    async fn one_shot_stream_enforces_declared_length_and_digest() {
480        let too_long = stream::iter([Ok::<_, io::Error>(Bytes::from_static(b"ab"))]);
481        let prepared = ByteStream::from_stream(too_long, 1, Sha256::digest(b"a").into())
482            .prepare()
483            .await
484            .unwrap();
485        assert!(
486            prepared
487                .request_body()
488                .await
489                .unwrap()
490                .collect()
491                .await
492                .is_err()
493        );
494
495        let wrong_digest = stream::iter([Ok::<_, io::Error>(Bytes::from_static(b"x"))]);
496        let prepared = ByteStream::from_stream(wrong_digest, 1, [0_u8; 32])
497            .prepare()
498            .await
499            .unwrap();
500        assert!(
501            prepared
502                .request_body()
503                .await
504                .unwrap()
505                .collect()
506                .await
507                .is_err()
508        );
509
510        let trailing_chunk = stream::iter([
511            Ok::<_, io::Error>(Bytes::from_static(b"a")),
512            Ok::<_, io::Error>(Bytes::from_static(b"b")),
513        ]);
514        let prepared = ByteStream::from_stream(trailing_chunk, 1, Sha256::digest(b"a").into())
515            .prepare()
516            .await
517            .unwrap();
518        let error = prepared
519            .request_body()
520            .await
521            .unwrap()
522            .collect()
523            .await
524            .unwrap_err();
525        assert_eq!(error.category(), crate::error::ErrorCategory::Integrity);
526    }
527}