1use 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 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 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}