Skip to main content

s3_wire/client/
multipart_upload.rs

1//! Bounded orchestration for managed multipart uploads.
2
3use std::future::Future;
4use std::path::PathBuf;
5use std::pin::Pin;
6use std::sync::Arc;
7use std::task::{Context, Poll};
8use std::time::Duration;
9
10use bytes::Bytes;
11use futures_util::stream::{FuturesUnordered, StreamExt as _};
12use http::HeaderMap;
13use tokio::sync::watch;
14use tokio::task::JoinHandle;
15
16use super::S3Client;
17use super::request::OperationDeadline;
18use crate::error::S3Error;
19use crate::operation::{
20    AbortMultipartUploadRequest, CompleteMultipartUploadOutput, CompleteMultipartUploadRequest,
21    CompletedPart, CreateMultipartUploadRequest, ManagedMultipartUploadRequest, MultipartUpload,
22    MultipartUploadSource, ObjectKey, PartNumber, UploadId, UploadPartRequest,
23};
24use crate::stream::{ByteStream, FileSnapshot};
25
26impl S3Client {
27    /// Uploads a replayable bytes or file source using bounded multipart requests.
28    ///
29    /// Part size, concurrency, total in-flight part bytes, and deadlines are
30    /// derived from the request's [`crate::MultipartOptions`]. File input is
31    /// snapshotted before an upload is created. Dropping this future cancels
32    /// outstanding parts and leaves an owned cleanup task to quiesce transmitted
33    /// requests before aborting the upload. A normal failure waits for cleanup;
34    /// if cleanup cannot be confirmed,
35    /// [`S3Error::cleanup_failure`] exposes that error while preserving the
36    /// original failure.
37    ///
38    /// # Errors
39    ///
40    /// Returns an error if the source cannot be prepared, exceeds S3's 10,000
41    /// part limit, a part fails, cancellation occurs, completion fails, or abort
42    /// cleanup fails after another error.
43    pub async fn multipart_upload(
44        &self,
45        request: ManagedMultipartUploadRequest,
46    ) -> Result<CompleteMultipartUploadOutput, S3Error> {
47        let runtime = tokio::runtime::Handle::try_current().map_err(S3Error::transport)?;
48        let transfer_deadline = OperationDeadline::new(request.options().transfer_timeout());
49        let cleanup_timeout = request.options().cleanup_timeout();
50        let cancellation = Cancellation::new();
51        let worker_cancellation = cancellation.clone();
52        let client = self.clone();
53        let handle = runtime.spawn(async move {
54            run_multipart_upload(
55                client,
56                request,
57                worker_cancellation,
58                transfer_deadline,
59                cleanup_timeout,
60            )
61            .await
62        });
63        OwnedMultipartTask {
64            cancellation,
65            handle,
66        }
67        .await
68    }
69}
70
71struct OwnedMultipartTask {
72    cancellation: Cancellation,
73    handle: JoinHandle<Result<CompleteMultipartUploadOutput, S3Error>>,
74}
75
76impl Future for OwnedMultipartTask {
77    type Output = Result<CompleteMultipartUploadOutput, S3Error>;
78
79    fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
80        match Pin::new(&mut self.handle).poll(context) {
81            Poll::Ready(Ok(result)) => Poll::Ready(result),
82            Poll::Ready(Err(error)) => Poll::Ready(Err(S3Error::transport(error))),
83            Poll::Pending => Poll::Pending,
84        }
85    }
86}
87
88impl Drop for OwnedMultipartTask {
89    fn drop(&mut self) {
90        self.cancellation.cancel();
91    }
92}
93
94#[derive(Clone)]
95struct Cancellation {
96    sender: watch::Sender<bool>,
97}
98
99impl Cancellation {
100    fn new() -> Self {
101        let (sender, _) = watch::channel(false);
102        Self { sender }
103    }
104
105    fn cancel(&self) {
106        self.sender.send_replace(true);
107    }
108
109    fn is_cancelled(&self) -> bool {
110        *self.sender.borrow()
111    }
112
113    async fn cancelled(&self) {
114        let mut receiver = self.sender.subscribe();
115        while !*receiver.borrow() {
116            if receiver.changed().await.is_err() {
117                return;
118            }
119        }
120    }
121}
122
123async fn run_multipart_upload(
124    client: S3Client,
125    request: ManagedMultipartUploadRequest,
126    cancellation: Cancellation,
127    transfer_deadline: OperationDeadline,
128    cleanup_timeout: Duration,
129) -> Result<CompleteMultipartUploadOutput, S3Error> {
130    let options = request.options();
131    let phase_headers = request.headers;
132    let source = tokio::select! {
133        biased;
134        () = cancellation.cancelled() => return Err(cancelled()),
135        result = tokio::time::timeout_at(
136            transfer_deadline.instant(),
137            PreparedMultipartSource::prepare(request.source),
138        ) => result.map_err(|_| transfer_timeout())??,
139    };
140    let part_size = options.part_size();
141    let part_count = source.part_count(part_size)?;
142    if cancellation.is_cancelled() {
143        return Err(cancelled());
144    }
145
146    let mut create = CreateMultipartUploadRequest::new(request.key.clone());
147    create.content_type = request.content_type;
148    create.user_metadata = request.user_metadata;
149    create.checksum_algorithm = request.checksum_algorithm;
150    create.headers = phase_headers.create;
151    // Once creation starts, allow it to finish even after cancellation so a
152    // successful response cannot be discarded together with the upload ID
153    // needed for cleanup.
154    let created = client
155        .create_multipart_upload_with_deadline(create, &transfer_deadline)
156        .await?;
157    let upload_id = created.upload_id().clone();
158
159    let result = upload_parts(
160        PartUploadContext {
161            client: client.clone(),
162            key: request.key.clone(),
163            upload_id: upload_id.clone(),
164            source: Arc::new(source),
165            part_size,
166            checksum_algorithm: request.checksum_algorithm,
167            headers: phase_headers.upload_part,
168            deadline: transfer_deadline,
169        },
170        part_count,
171        options.concurrency(),
172        &cancellation,
173        cleanup_timeout,
174    )
175    .await;
176
177    let upload = match result {
178        Ok(upload) => upload,
179        Err(failure) => {
180            return fail_with_abort(
181                &client,
182                request.key,
183                upload_id,
184                phase_headers.abort,
185                failure,
186            )
187            .await;
188        }
189    };
190    let mut completion = match CompleteMultipartUploadRequest::new(
191        upload.key().clone(),
192        upload.upload_id().clone(),
193        upload.completed_parts().to_vec(),
194    ) {
195        Ok(completion) => completion,
196        Err(error) => {
197            return fail_with_abort(
198                &client,
199                request.key,
200                upload_id,
201                phase_headers.abort,
202                ManagedUploadFailure::new(S3Error::integrity(error.to_string()), cleanup_timeout),
203            )
204            .await;
205        }
206    };
207    completion.headers = phase_headers.complete;
208
209    let completion = tokio::select! {
210        biased;
211        result = client.complete_multipart_upload_with_deadline(
212            completion,
213            &transfer_deadline,
214        ) => result,
215        () = cancellation.cancelled() => Err(cancelled()),
216    };
217
218    match completion {
219        Ok(output) => Ok(output),
220        Err(primary) => {
221            fail_with_abort(
222                &client,
223                request.key,
224                upload_id,
225                phase_headers.abort,
226                ManagedUploadFailure::new(primary, cleanup_timeout),
227            )
228            .await
229        }
230    }
231}
232
233#[derive(Clone)]
234struct PartUploadContext {
235    client: S3Client,
236    key: ObjectKey,
237    upload_id: UploadId,
238    source: Arc<PreparedMultipartSource>,
239    part_size: u64,
240    checksum_algorithm: Option<crate::operation::ChecksumAlgorithm>,
241    headers: HeaderMap,
242    deadline: OperationDeadline,
243}
244
245async fn upload_parts(
246    context: PartUploadContext,
247    part_count: u16,
248    concurrency: usize,
249    cancellation: &Cancellation,
250    cleanup_timeout: Duration,
251) -> Result<MultipartUpload, ManagedUploadFailure> {
252    let mut next_part = 1_u16;
253    let mut pending = FuturesUnordered::new();
254    let mut upload = MultipartUpload::new(context.key.clone(), context.upload_id.clone());
255    let deadline_sleep = tokio::time::sleep_until(context.deadline.instant());
256    tokio::pin!(deadline_sleep);
257
258    loop {
259        while pending.len() < concurrency && next_part <= part_count {
260            pending.push(upload_one_part(context.clone(), next_part));
261            next_part += 1;
262        }
263        if pending.is_empty() {
264            return Ok(upload);
265        }
266        let completed = tokio::select! {
267            () = cancellation.cancelled() => Err(cancelled()),
268            () = &mut deadline_sleep => Err(transfer_timeout()),
269            result = pending.next() => result
270                .ok_or_else(|| S3Error::integrity("multipart scheduler lost an in-flight part"))
271                .and_then(|result| result),
272        };
273        let completed = match completed {
274            Ok(completed) => completed,
275            Err(primary) => {
276                return Err(quiesce_part_requests(primary, &mut pending, cleanup_timeout).await);
277            }
278        };
279        if let Err(error) = upload.record_part(completed) {
280            return Err(quiesce_part_requests(
281                S3Error::integrity(error.to_string()),
282                &mut pending,
283                cleanup_timeout,
284            )
285            .await);
286        }
287    }
288}
289
290struct ManagedUploadFailure {
291    primary: S3Error,
292    cleanup_deadline: OperationDeadline,
293    quiesce_failure: Option<S3Error>,
294}
295
296impl ManagedUploadFailure {
297    fn new(primary: S3Error, cleanup_timeout: Duration) -> Self {
298        Self {
299            primary,
300            cleanup_deadline: OperationDeadline::new(cleanup_timeout),
301            quiesce_failure: None,
302        }
303    }
304}
305
306async fn quiesce_part_requests<F>(
307    primary: S3Error,
308    pending: &mut FuturesUnordered<F>,
309    cleanup_timeout: Duration,
310) -> ManagedUploadFailure
311where
312    F: Future<Output = Result<CompletedPart, S3Error>>,
313{
314    let mut failure = ManagedUploadFailure::new(primary, cleanup_timeout);
315    let settle_until = tokio::time::Instant::now() + cleanup_timeout / 2;
316    while !pending.is_empty() {
317        match tokio::time::timeout_at(settle_until, pending.next()).await {
318            Ok(Some(_)) => {}
319            Ok(None) => break,
320            Err(_) => {
321                failure.quiesce_failure = Some(S3Error::timeout(
322                    crate::error::TimeoutPhase::Operation,
323                    "multipart cleanup could not quiesce in-flight part requests",
324                ));
325                break;
326            }
327        }
328    }
329    failure
330}
331
332async fn fail_with_abort(
333    client: &S3Client,
334    key: ObjectKey,
335    upload_id: UploadId,
336    headers: HeaderMap,
337    failure: ManagedUploadFailure,
338) -> Result<CompleteMultipartUploadOutput, S3Error> {
339    let mut request = AbortMultipartUploadRequest::new(key, upload_id);
340    request.headers = headers;
341    let abort = client
342        .abort_multipart_upload_with_deadline(request, &failure.cleanup_deadline)
343        .await
344        .map(|_| ());
345    let cleanup = match (abort, failure.quiesce_failure) {
346        (Ok(()), None) => Ok(()),
347        (Err(error), None) | (Ok(()), Some(error)) => Err(error),
348        (Err(abort), Some(quiesce)) => Err(abort.with_cleanup_failure(quiesce)),
349    };
350    Err(attach_cleanup_failure(failure.primary, cleanup))
351}
352
353async fn upload_one_part(
354    context: PartUploadContext,
355    part_number: u16,
356) -> Result<CompletedPart, S3Error> {
357    let body = context
358        .source
359        .read_part(context.part_size, part_number)
360        .await?;
361    let number = PartNumber::new(part_number)
362        .ok_or_else(|| S3Error::integrity("generated multipart part number is invalid"))?;
363    let checksum = match context.checksum_algorithm {
364        Some(algorithm) => tokio::time::timeout_at(
365            context.deadline.instant(),
366            crate::operation::Checksum::calculate_cooperatively(algorithm, &body),
367        )
368        .await
369        .map_err(|_| transfer_timeout())?
370        .map_err(|error| S3Error::unsupported(error.to_string()))?,
371        None => crate::operation::Checksum::default(),
372    };
373    let mut request = UploadPartRequest::new(
374        context.key,
375        context.upload_id,
376        number,
377        ByteStream::from_bytes(body),
378    )
379    .with_checksum(checksum);
380    request.headers = context.headers;
381    let output = context
382        .client
383        .upload_part_with_deadline(request, &context.deadline)
384        .await?;
385    CompletedPart::new(output.part_number.get(), output.e_tag)
386        .map(|part| part.with_checksum(output.checksum))
387        .map_err(|error| S3Error::invalid_response(error.to_string()))
388}
389
390fn attach_cleanup_failure(primary: S3Error, cleanup: Result<(), S3Error>) -> S3Error {
391    match cleanup {
392        Ok(()) => primary,
393        Err(cleanup) => primary.with_cleanup_failure(cleanup),
394    }
395}
396
397fn cancelled() -> S3Error {
398    S3Error::cancellation("multipart upload was cancelled")
399}
400
401fn transfer_timeout() -> S3Error {
402    S3Error::timeout(
403        crate::error::TimeoutPhase::Operation,
404        "managed multipart upload exceeded its transfer deadline",
405    )
406}
407
408enum PreparedMultipartSource {
409    Bytes(Bytes),
410    File(FileSnapshot),
411}
412
413impl PreparedMultipartSource {
414    async fn prepare(source: MultipartUploadSource) -> Result<Self, S3Error> {
415        match source {
416            MultipartUploadSource::Bytes(bytes) => Ok(Self::Bytes(bytes)),
417            MultipartUploadSource::File(path) => Self::snapshot(path).await,
418        }
419    }
420
421    async fn snapshot(path: PathBuf) -> Result<Self, S3Error> {
422        FileSnapshot::create(path, false).await.map(Self::File)
423    }
424
425    fn length(&self) -> Result<u64, S3Error> {
426        match self {
427            Self::Bytes(bytes) => u64::try_from(bytes.len())
428                .map_err(|_| S3Error::configuration("multipart byte length does not fit in u64")),
429            Self::File(snapshot) => Ok(snapshot.length()),
430        }
431    }
432
433    fn part_count(&self, part_size: u64) -> Result<u16, S3Error> {
434        let length = self.length()?;
435        if length == 0 {
436            return Err(S3Error::configuration(
437                "managed multipart uploads require a non-empty source",
438            ));
439        }
440        let count = length.div_ceil(part_size);
441        if count > u64::from(PartNumber::MAX) {
442            return Err(S3Error::configuration(
443                "multipart source exceeds the configured 10,000-part limit",
444            ));
445        }
446        u16::try_from(count)
447            .map_err(|_| S3Error::configuration("multipart part count does not fit in u16"))
448    }
449
450    async fn read_part(&self, part_size: u64, part_number: u16) -> Result<Bytes, S3Error> {
451        let offset = u64::from(part_number - 1)
452            .checked_mul(part_size)
453            .ok_or_else(|| S3Error::integrity("multipart part offset overflow"))?;
454        let remaining = self
455            .length()?
456            .checked_sub(offset)
457            .ok_or_else(|| S3Error::integrity("multipart part offset exceeds source length"))?;
458        let length = remaining.min(part_size);
459        let length_usize = usize::try_from(length)
460            .map_err(|_| S3Error::configuration("multipart part size does not fit in usize"))?;
461        match self {
462            Self::Bytes(bytes) => {
463                let start = usize::try_from(offset).map_err(|_| {
464                    S3Error::configuration("multipart byte offset does not fit in usize")
465                })?;
466                Ok(bytes.slice(start..start + length_usize))
467            }
468            Self::File(snapshot) => snapshot.read_range(offset, length_usize).await,
469        }
470    }
471}
472
473#[cfg(test)]
474mod tests {
475    use super::*;
476    use crate::error::ErrorCategory;
477
478    #[tokio::test]
479    async fn task_drop_signals_owned_cleanup_worker() {
480        let cancellation = Cancellation::new();
481        let worker = cancellation.clone();
482        let (finished, observed) = tokio::sync::oneshot::channel();
483        let handle = tokio::spawn(async move {
484            worker.cancelled().await;
485            let _ = finished.send(());
486            Err(cancelled())
487        });
488        drop(OwnedMultipartTask {
489            cancellation,
490            handle,
491        });
492        observed.await.unwrap();
493    }
494
495    #[test]
496    fn cleanup_failure_does_not_replace_primary_error() {
497        let primary = S3Error::integrity("primary");
498        let cleanup = S3Error::invalid_response("cleanup");
499        let combined = attach_cleanup_failure(primary, Err(cleanup));
500        assert_eq!(combined.category(), ErrorCategory::Integrity);
501        assert_eq!(combined.message(), "primary");
502        assert_eq!(
503            combined.cleanup_failure().map(S3Error::category),
504            Some(ErrorCategory::InvalidResponse)
505        );
506    }
507
508    #[tokio::test]
509    async fn byte_parts_are_exact_and_bounded() {
510        let source = PreparedMultipartSource::Bytes(Bytes::from_static(b"abcdefghij"));
511        assert_eq!(source.part_count(4).unwrap(), 3);
512        assert_eq!(source.read_part(4, 1).await.unwrap(), "abcd");
513        assert_eq!(source.read_part(4, 2).await.unwrap(), "efgh");
514        assert_eq!(source.read_part(4, 3).await.unwrap(), "ij");
515    }
516
517    #[tokio::test]
518    async fn file_source_is_snapshotted_before_parts_are_read() {
519        let directory = tempfile::tempdir().unwrap();
520        let path = directory.path().join("source");
521        tokio::fs::write(&path, b"original").await.unwrap();
522        let source = PreparedMultipartSource::snapshot(path.clone())
523            .await
524            .unwrap();
525        tokio::fs::write(path, b"changed!").await.unwrap();
526        assert_eq!(source.read_part(8, 1).await.unwrap(), "original");
527    }
528
529    #[tokio::test]
530    async fn cleanup_reports_part_requests_that_cannot_quiesce() {
531        let mut pending = FuturesUnordered::new();
532        pending.push(std::future::pending::<Result<CompletedPart, S3Error>>());
533
534        let failure = quiesce_part_requests(
535            S3Error::cancellation("primary"),
536            &mut pending,
537            Duration::from_millis(10),
538        )
539        .await;
540
541        assert!(failure.quiesce_failure.is_some());
542        assert_eq!(failure.primary.category(), ErrorCategory::Cancellation);
543    }
544
545    #[test]
546    fn part_count_rejects_empty_and_excessive_sources() {
547        let empty = PreparedMultipartSource::Bytes(Bytes::new());
548        assert!(empty.part_count(5).is_err());
549        let too_large = PreparedMultipartSource::Bytes(Bytes::from(vec![0_u8; 10_001]));
550        assert!(too_large.part_count(1).is_err());
551    }
552}