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