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
21pub 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 pub fn from_bytes(bytes: impl Into<Bytes>) -> Self {
43 Self {
44 source: UploadSource::Bytes(bytes.into()),
45 }
46 }
47
48 pub fn from_path(path: impl AsRef<Path>) -> Self {
54 Self {
55 source: UploadSource::File(path.as_ref().to_owned()),
56 }
57 }
58
59 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 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
220pub(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}