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