Skip to main content

mz_persist/
s3.rs

1// Copyright Materialize, Inc. and contributors. All rights reserved.
2//
3// Use of this software is governed by the Business Source License
4// included in the LICENSE file.
5//
6// As of the Change Date specified in that file, in accordance with
7// the Business Source License, use of this software will be governed
8// by the Apache License, Version 2.0.
9
10//! An S3 implementation of [Blob] storage.
11
12use std::cmp;
13use std::fmt::{Debug, Formatter};
14use std::ops::Range;
15use std::sync::Arc;
16use std::sync::atomic::{self, AtomicU64};
17use std::time::{Duration, Instant};
18
19use anyhow::{Context, anyhow};
20use async_trait::async_trait;
21use aws_config::sts::AssumeRoleProvider;
22use aws_config::timeout::TimeoutConfig;
23use aws_credential_types::Credentials;
24use aws_sdk_s3::Client as S3Client;
25use aws_sdk_s3::config::{AsyncSleep, Sleep};
26use aws_sdk_s3::error::{ProvideErrorMetadata, SdkError};
27use aws_sdk_s3::operation::get_object::GetObjectOutput;
28use aws_sdk_s3::primitives::ByteStream;
29use aws_sdk_s3::types::{CompletedMultipartUpload, CompletedPart};
30use aws_types::region::Region;
31use bytes::{Bytes, BytesMut};
32use futures_util::stream::FuturesOrdered;
33use futures_util::{FutureExt, StreamExt};
34use mz_ore::bytes::SegmentedBytes;
35use mz_ore::cast::CastFrom;
36use mz_ore::metrics::MetricsRegistry;
37use mz_ore::task::RuntimeExt;
38use tokio::runtime::Handle as AsyncHandle;
39use tracing::{Instrument, debug, debug_span, trace, trace_span};
40use uuid::Uuid;
41
42use crate::cfg::BlobKnobs;
43use crate::error::Error;
44use crate::location::{Blob, BlobMetadata, Determinate, ExternalError};
45use crate::metrics::S3BlobMetrics;
46
47/// Configuration for opening an [S3Blob].
48///
49/// NOTE: cloning shares the underlying `S3Client` and therefore its HTTP
50/// connection pool. Connection-pool isolation (as hedged gets require, see
51/// [crate::hedge]) needs a fresh [S3BlobConfig::new].
52#[derive(Clone, Debug)]
53pub struct S3BlobConfig {
54    metrics: S3BlobMetrics,
55    client: S3Client,
56    bucket: String,
57    prefix: String,
58}
59
60// There is no simple way to hook into the S3 client to capture when its various timeouts
61// are hit. Instead, we pass along marker values that inform our [MetricsSleep] impl which
62// type of timeout was requested so it can substitute in a dynamic value set by config
63// from the caller.
64const OPERATION_TIMEOUT_MARKER: Duration = Duration::new(111, 1111);
65const OPERATION_ATTEMPT_TIMEOUT_MARKER: Duration = Duration::new(222, 2222);
66const CONNECT_TIMEOUT_MARKER: Duration = Duration::new(333, 3333);
67const READ_TIMEOUT_MARKER: Duration = Duration::new(444, 4444);
68
69#[derive(Debug)]
70struct MetricsSleep {
71    knobs: Box<dyn BlobKnobs>,
72    metrics: S3BlobMetrics,
73}
74
75impl AsyncSleep for MetricsSleep {
76    fn sleep(&self, duration: Duration) -> Sleep {
77        let (duration, metric) = match duration {
78            OPERATION_TIMEOUT_MARKER => (
79                self.knobs.operation_timeout(),
80                Some(self.metrics.operation_timeouts.clone()),
81            ),
82            OPERATION_ATTEMPT_TIMEOUT_MARKER => (
83                self.knobs.operation_attempt_timeout(),
84                Some(self.metrics.operation_attempt_timeouts.clone()),
85            ),
86            CONNECT_TIMEOUT_MARKER => (
87                self.knobs.connect_timeout(),
88                Some(self.metrics.connect_timeouts.clone()),
89            ),
90            READ_TIMEOUT_MARKER => (
91                self.knobs.read_timeout(),
92                Some(self.metrics.read_timeouts.clone()),
93            ),
94            duration => (duration, None),
95        };
96
97        // the sleep future we return here will only be polled to
98        // completion if its corresponding http request to S3 times
99        // out, meaning we can chain incrementing the appropriate
100        // timeout counter to when it finishes
101        Sleep::new(tokio::time::sleep(duration).map(|x| {
102            if let Some(counter) = metric {
103                counter.inc();
104            }
105            x
106        }))
107    }
108}
109
110impl S3BlobConfig {
111    const EXTERNAL_TESTS_S3_BUCKET: &'static str = "MZ_PERSIST_EXTERNAL_STORAGE_TEST_S3_BUCKET";
112
113    /// Returns a new [S3BlobConfig] for use in production.
114    ///
115    /// Stores objects in the given bucket prepended with the (possibly empty)
116    /// prefix. S3 credentials and region must be available in the process or
117    /// environment.
118    pub async fn new(
119        bucket: String,
120        prefix: String,
121        role_arn: Option<String>,
122        endpoint: Option<String>,
123        region: Option<String>,
124        credentials: Option<(String, String)>,
125        knobs: Box<dyn BlobKnobs>,
126        metrics: S3BlobMetrics,
127    ) -> Result<Self, Error> {
128        let mut loader = mz_aws_util::defaults();
129
130        if let Some(region) = region {
131            loader = loader.region(Region::new(region));
132        };
133
134        if let Some(role_arn) = role_arn {
135            let assume_role_sdk_config = mz_aws_util::defaults().load().await;
136            let role_provider = AssumeRoleProvider::builder(role_arn)
137                .configure(&assume_role_sdk_config)
138                .session_name("persist")
139                .build()
140                .await;
141            loader = loader.credentials_provider(role_provider);
142        }
143
144        if let Some((access_key_id, secret_access_key)) = credentials {
145            loader = loader.credentials_provider(Credentials::from_keys(
146                access_key_id,
147                secret_access_key,
148                None,
149            ));
150        }
151
152        if let Some(endpoint) = endpoint {
153            loader = loader.endpoint_url(endpoint)
154        }
155
156        // NB: we must always use the custom sleep impl if we use the timeout marker values
157        loader = loader.sleep_impl(MetricsSleep {
158            knobs,
159            metrics: metrics.clone(),
160        });
161        loader = loader.timeout_config(
162            TimeoutConfig::builder()
163                // maximum time allowed for a top-level S3 API call (including internal retries)
164                .operation_timeout(OPERATION_TIMEOUT_MARKER)
165                // maximum time allowed for a single network call
166                .operation_attempt_timeout(OPERATION_ATTEMPT_TIMEOUT_MARKER)
167                // maximum time until a connection succeeds
168                .connect_timeout(CONNECT_TIMEOUT_MARKER)
169                // maximum time to read the first byte of a response
170                .read_timeout(READ_TIMEOUT_MARKER)
171                .build(),
172        );
173
174        let client = mz_aws_util::s3::new_client(&loader.load().await);
175        Ok(S3BlobConfig {
176            metrics,
177            client,
178            bucket,
179            prefix,
180        })
181    }
182
183    /// Returns a new [S3BlobConfig] for use in unit tests.
184    ///
185    /// By default, persist tests that use external storage (like s3) are
186    /// no-ops, so that `cargo test` does the right thing without any
187    /// configuration. To activate the tests, set the
188    /// `MZ_PERSIST_EXTERNAL_STORAGE_TEST_S3_BUCKET` environment variable and
189    /// ensure you have valid AWS credentials available in a location where the
190    /// AWS Rust SDK can discovery them.
191    ///
192    /// This intentionally uses the `MZ_PERSIST_EXTERNAL_STORAGE_TEST_S3_BUCKET`
193    /// env as the switch for test no-op-ness instead of the presence of a valid
194    /// AWS authentication configuration envs because a developers might have
195    /// valid credentials present and this isn't an explicit enough signal from
196    /// a developer running `cargo test` that it's okay to use these
197    /// credentials. It also intentionally does not use the local drop-in s3
198    /// replacement to keep persist unit tests light.
199    ///
200    /// On CI, these tests are enabled by adding the scratch-aws-access plugin
201    /// to the `cargo-test` step in `ci/test/pipeline.template.yml` and setting
202    /// `MZ_PERSIST_EXTERNAL_STORAGE_TEST_S3_BUCKET` in
203    /// `ci/test/cargo-test/mzcompose.py`.
204    ///
205    /// For a Materialize developer, to opt in to these tests locally for
206    /// development, follow the AWS access guide:
207    ///
208    /// ```text
209    /// https://github.com/MaterializeInc/i2/blob/main/doc/aws-access.md
210    /// ```
211    ///
212    /// then running `source src/persist/s3_test_env_mz.sh`. You will also have
213    /// to run `aws sso login` if you haven't recently.
214    ///
215    /// Non-Materialize developers will have to set up their own auto-deleting
216    /// bucket and export the same env vars that s3_test_env_mz.sh does.
217    ///
218    /// Only public for use in src/benches.
219    pub async fn new_for_test() -> Result<Option<Self>, Error> {
220        let bucket = match std::env::var(Self::EXTERNAL_TESTS_S3_BUCKET) {
221            Ok(bucket) => bucket,
222            Err(_) => {
223                if mz_ore::env::is_var_truthy("CI") {
224                    panic!("CI is supposed to run this test but something has gone wrong!");
225                }
226                return Ok(None);
227            }
228        };
229
230        struct TestBlobKnobs;
231        impl std::fmt::Debug for TestBlobKnobs {
232            fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
233                f.debug_struct("TestBlobKnobs").finish_non_exhaustive()
234            }
235        }
236        impl BlobKnobs for TestBlobKnobs {
237            fn operation_timeout(&self) -> Duration {
238                OPERATION_TIMEOUT_MARKER
239            }
240
241            fn operation_attempt_timeout(&self) -> Duration {
242                OPERATION_ATTEMPT_TIMEOUT_MARKER
243            }
244
245            fn connect_timeout(&self) -> Duration {
246                CONNECT_TIMEOUT_MARKER
247            }
248
249            fn read_timeout(&self) -> Duration {
250                READ_TIMEOUT_MARKER
251            }
252
253            fn is_cc_active(&self) -> bool {
254                false
255            }
256        }
257
258        // Give each test a unique prefix so they don't conflict. We don't have
259        // to worry about deleting any data that we create because the bucket is
260        // set to auto-delete after 1 day.
261        let prefix = Uuid::new_v4().to_string();
262        let role_arn = None;
263        let metrics = S3BlobMetrics::new(&MetricsRegistry::new());
264        let config = S3BlobConfig::new(
265            bucket,
266            prefix,
267            role_arn,
268            None,
269            None,
270            None,
271            Box::new(TestBlobKnobs),
272            metrics,
273        )
274        .await?;
275        Ok(Some(config))
276    }
277
278    /// Returns a clone of Self with a new v4 uuid prefix.
279    pub fn clone_with_new_uuid_prefix(&self) -> Self {
280        let mut ret = self.clone();
281        ret.prefix = Uuid::new_v4().to_string();
282        ret
283    }
284}
285
286/// Implementation of [Blob] backed by S3.
287#[derive(Debug)]
288pub struct S3Blob {
289    metrics: S3BlobMetrics,
290    client: S3Client,
291    bucket: String,
292    prefix: String,
293    // Maximum number of keys we get information about per list-objects request.
294    //
295    // Defaults to 1000 which is the current AWS max.
296    max_keys: i32,
297    multipart_config: MultipartConfig,
298}
299
300impl S3Blob {
301    /// Opens the given location for non-exclusive read-write access.
302    pub async fn open(config: S3BlobConfig) -> Result<Self, ExternalError> {
303        let ret = S3Blob {
304            metrics: config.metrics,
305            client: config.client,
306            bucket: config.bucket,
307            prefix: config.prefix,
308            max_keys: 1_000,
309            multipart_config: MultipartConfig::default(),
310        };
311        // Connect before returning success. We don't particularly care about
312        // what's stored in this blob (nothing writes to it, so presumably it's
313        // empty) just that we were able and allowed to fetch it.
314        let _ = ret.get("HEALTH_CHECK").await?;
315        Ok(ret)
316    }
317
318    fn get_path(&self, key: &str) -> String {
319        format!("{}/{}", self.prefix, key)
320    }
321}
322
323#[async_trait]
324impl Blob for S3Blob {
325    async fn get(&self, key: &str) -> Result<Option<SegmentedBytes>, ExternalError> {
326        let start_overall = Instant::now();
327        let path = self.get_path(key);
328
329        // S3 advises that it's fastest to download large objects along the part
330        // boundaries they were originally uploaded with [1].
331        //
332        // [1]: https://docs.aws.amazon.com/whitepapers/latest/s3-optimizing-performance-best-practices/use-byte-range-fetches.html
333        //
334        // One option is to run the same logic as multipart does and do the
335        // requests using the resulting byte ranges, but if we ever changed the
336        // multipart chunking logic, they wouldn't line up for old blobs written
337        // by a previous version.
338        //
339        // Another option is to store the part boundaries in the metadata we
340        // keep about the batch, but this would be large and wasteful.
341        //
342        // Luckily, s3 exposes a part_number param on GetObject requests that we
343        // can use. If an object was created with multipart, it allows
344        // requesting each part as they were originally uploaded by the part
345        // number index. With this, we can simply send off requests for part
346        // number 1..=num_parts and reassemble the results.
347        //
348        // We could roundtrip the number of parts through persist batch
349        // metadata, but with some cleverness, we can avoid even this. Turns
350        // out, if multipart upload wasn't used (it was just a normal PutObject
351        // request), s3 will still happily return it for a request specifying a
352        // part_number of 1. This lets us fire off a first request, which
353        // contains the metadata we need to determine how many additional parts
354        // we need, if any.
355        //
356        // So, the following call sends this first request. The SDK even returns
357        // the headers before the full data body has completed. This gives us
358        // the number of parts. We can then proceed to fetch the body of the
359        // first request concurrently with the rest of the parts of the object.
360        //
361        // Not every S3-compatible store reports the part count. Its response
362        // to the first request is then indistinguishable from a single-part
363        // object's, except that `Content-Range` still carries the object's
364        // total size. The remainder is fetched by byte range in that case, and
365        // the reassembled length is checked against the total either way, so a
366        // store that misbehaves fails the get instead of handing the decoder a
367        // truncated blob.
368
369        // For each header and body that we fetch, we track the fastest, and
370        // any large deviations from it.
371        let min_body_elapsed = Arc::new(MinElapsed::default());
372        let min_header_elapsed = Arc::new(MinElapsed::default());
373        self.metrics.get_part.inc();
374
375        // Fetch our first header, this tells us how many more are left.
376        let header_start = Instant::now();
377        let object = self
378            .client
379            .get_object()
380            .bucket(&self.bucket)
381            .key(&path)
382            .part_number(1)
383            .send()
384            .await;
385        let elapsed = header_start.elapsed();
386        min_header_elapsed.observe(elapsed, "s3 download first part header");
387
388        let first_part = match object {
389            Ok(object) => object,
390            Err(SdkError::ServiceError(err)) if err.err().is_no_such_key() => return Ok(None),
391            Err(err) => {
392                self.update_error_metrics("GetObject", &err);
393                Err(anyhow!(err).context("s3 get meta err"))?
394            }
395        };
396
397        // The object's total size comes from `Content-Range`, which s3 returns
398        // for any request that names a part. Its denominator is the length of
399        // the whole object, not of the part that was asked for, which is what
400        // makes it usable as the completeness check below. A store that put
401        // the part's length there instead would fail every multipart get.
402        let total_len = first_part
403            .content_range()
404            .and_then(parse_content_range_total);
405        let first_len = first_part
406            .content_length()
407            .and_then(|len| u64::try_from(len).ok());
408
409        let remaining: Vec<PartSelector> = match first_part.parts_count() {
410            // For a non-multipart upload, parts_count will be None, and the
411            // first request already returned the whole object. A store that
412            // does not report the count looks the same, except that the first
413            // part falls short of the total. Its remainder is fetched by byte
414            // range in chunks of the first part's size, which is the size the
415            // upload used for every part but the last.
416            None => match (first_len, total_len) {
417                (Some(first_len), Some(total_len)) if first_len < total_len => {
418                    remaining_ranges(first_len, total_len)
419                        .into_iter()
420                        .map(PartSelector::Range)
421                        .collect()
422                }
423                _ => Vec::new(),
424            },
425            Some(parts @ 1..) => (2..=parts).map(PartSelector::Number).collect(),
426            // A non-positive value is invalid.
427            Some(bad) => {
428                assert!(bad <= 0);
429                return Err(anyhow!("unexpected number of s3 object parts: {}", bad).into());
430            }
431        };
432        let num_requests = remaining.len() + 1;
433
434        trace!(
435            "s3 download first header took {:?} ({num_requests} requests)",
436            start_overall.elapsed(),
437        );
438
439        let mut body_futures = FuturesOrdered::new();
440        let mut requests = vec![PartRequest::First(first_part)];
441        requests.extend(remaining.into_iter().map(PartRequest::Fetch));
442
443        for request in requests {
444            // Clone a handle to our MinElapsed trackers so we can give one to
445            // each download task.
446            let min_header_elapsed = Arc::clone(&min_header_elapsed);
447            let min_body_elapsed = Arc::clone(&min_body_elapsed);
448            let get_invalid_resp = self.metrics.get_invalid_resp.clone();
449            let path = &path;
450            let request_future = async move {
451                let mut object = match request {
452                    // Fetched above, together with the headers that shaped
453                    // the remaining requests.
454                    PartRequest::First(object) => object,
455                    PartRequest::Fetch(selector) => {
456                        let header_start = Instant::now();
457                        let req = self.client.get_object().bucket(&self.bucket).key(path);
458                        let req = match selector {
459                            PartSelector::Number(part_num) => req.part_number(part_num),
460                            PartSelector::Range(range) => {
461                                req.range(format!("bytes={}-{}", range.start, range.end - 1))
462                            }
463                        };
464                        let object = req
465                            .send()
466                            .await
467                            .inspect_err(|err| self.update_error_metrics("GetObject", err))
468                            .context("s3 get meta err")?;
469                        min_header_elapsed
470                            .observe(header_start.elapsed(), "s3 download part header");
471                        object
472                    }
473                };
474                // Request the body.
475                let body_start = Instant::now();
476
477                // Coalesce all hyper chunks for this part into a single contiguous
478                // allocation. Pushing each SDK `Bytes` chunk separately into
479                // `SegmentedBytes` yields hundreds of segments per blob, which makes
480                // every parquet `ChunkReader::get_bytes` call O(N) and dominates CPU
481                // in `SegmentedBytes::advance`/`get_bytes` during decode. Copying
482                // also releases the hyper pool buffer so it doesn't stay pinned for
483                // the lifetime of the blob.
484                let mut buf = match object.content_length() {
485                    Some(len @ 1..) => BytesMut::with_capacity(usize::cast_from(
486                        u64::try_from(len).expect("positive integer"),
487                    )),
488                    Some(len @ ..=-1) => {
489                        tracing::trace!(?len, "found invalid content-length");
490                        get_invalid_resp.inc();
491                        BytesMut::new()
492                    }
493                    Some(0) | None => BytesMut::new(),
494                };
495
496                while let Some(data) = object.body.next().await {
497                    let data = data.context("s3 get body err")?;
498                    buf.extend_from_slice(&data);
499                }
500
501                let body_elapsed = body_start.elapsed();
502                min_body_elapsed.observe(body_elapsed, "s3 download part body");
503
504                let body_parts = if buf.is_empty() {
505                    Vec::new()
506                } else {
507                    vec![buf.freeze()]
508                };
509                Ok::<_, anyhow::Error>(body_parts)
510            };
511
512            body_futures.push_back(request_future);
513        }
514
515        // Await on all of our parts requests.
516        let mut segments = vec![];
517        while let Some(result) = body_futures.next().await {
518            // Download failure, we failed to fetch the body from S3.
519            let mut part_body = result
520                .inspect_err(|e| {
521                    self.metrics
522                        .error_counts
523                        .with_label_values(&["GetObjectStream", e.to_string().as_str()])
524                        .inc()
525                })
526                .context("s3 get body err")?;
527
528            // Collect all of our segments.
529            segments.append(&mut part_body);
530        }
531
532        // A store that reports neither the part count nor a range it honors
533        // hands back less than the object. Fail the get here rather than let
534        // the decoder find out. This runs for every store, not just the ones
535        // needing the byte-range fallback, so any short read becomes a
536        // retryable error instead of a corrupt blob.
537        if let Some(total_len) = total_len {
538            let fetched_len: u64 = segments.iter().map(|s| u64::cast_from(s.len())).sum();
539            if fetched_len != total_len {
540                self.metrics.get_invalid_resp.inc();
541                return Err(anyhow!(
542                    "s3 GetObject {path} returned {fetched_len} bytes of a {total_len} byte object"
543                )
544                .into());
545            }
546        }
547
548        debug!(
549            "s3 GetObject took {:?} ({} requests)",
550            start_overall.elapsed(),
551            num_requests
552        );
553        Ok(Some(SegmentedBytes::from(segments)))
554    }
555
556    async fn list_keys_and_metadata(
557        &self,
558        key_prefix: &str,
559        f: &mut (dyn FnMut(BlobMetadata) + Send + Sync),
560    ) -> Result<(), ExternalError> {
561        let mut continuation_token = None;
562        // we only want to return keys that match the specified blob key prefix
563        let blob_key_prefix = self.get_path(key_prefix);
564        // but we want to exclude the shared root prefix from our returned keys,
565        // so only the blob key itself is passed in to `f`
566        let strippable_root_prefix = format!("{}/", self.prefix);
567
568        loop {
569            self.metrics.list_objects.inc();
570            let resp = self
571                .client
572                .list_objects_v2()
573                .bucket(&self.bucket)
574                .prefix(&blob_key_prefix)
575                .max_keys(self.max_keys)
576                .set_continuation_token(continuation_token)
577                .send()
578                .await
579                .inspect_err(|err| self.update_error_metrics("ListObjectsV2", err))
580                .context("list bucket error")?;
581            if let Some(contents) = resp.contents {
582                for object in contents.iter() {
583                    if let Some(key) = object.key.as_ref() {
584                        if let Some(key) = key.strip_prefix(&strippable_root_prefix) {
585                            let size_in_bytes = match object.size {
586                                None => {
587                                    return Err(ExternalError::from(anyhow!(
588                                        "object missing size: {key}"
589                                    )));
590                                }
591                                Some(size) => size
592                                    .try_into()
593                                    .expect("file in S3 cannot have negative size"),
594                            };
595                            f(BlobMetadata { key, size_in_bytes });
596                        } else {
597                            return Err(ExternalError::from(anyhow!(
598                                "found key with invalid prefix: {}",
599                                key
600                            )));
601                        }
602                    }
603                }
604            }
605
606            if resp.next_continuation_token.is_some() {
607                continuation_token = resp.next_continuation_token;
608            } else {
609                break;
610            }
611        }
612
613        Ok(())
614    }
615
616    async fn set(&self, key: &str, value: Bytes) -> Result<(), ExternalError> {
617        let value_len = value.len();
618        if self
619            .multipart_config
620            .should_multipart(value_len)
621            .map_err(anyhow::Error::msg)?
622        {
623            self.set_multi_part(key, value)
624                .instrument(debug_span!("s3set_multi", payload_len = value_len))
625                .await
626        } else {
627            self.set_single_part(key, value).await
628        }
629    }
630
631    async fn delete(&self, key: &str) -> Result<Option<usize>, ExternalError> {
632        // There is a race condition here where, if two delete calls for the
633        // same key occur simultaneously, both might think they did the actual
634        // deletion. This return value is only used for metrics, so it's
635        // unfortunate, but fine.
636        let path = self.get_path(key);
637        self.metrics.delete_head.inc();
638        let head_res = self
639            .client
640            .head_object()
641            .bucket(&self.bucket)
642            .key(&path)
643            .send()
644            .await;
645        let size_bytes = match head_res {
646            Ok(x) => match x.content_length {
647                None => {
648                    return Err(ExternalError::from(anyhow!(
649                        "s3 delete content length was none"
650                    )));
651                }
652                Some(content_length) => {
653                    u64::try_from(content_length).expect("file in S3 cannot have negative size")
654                }
655            },
656            Err(SdkError::ServiceError(err)) if err.err().is_not_found() => return Ok(None),
657            Err(err) => {
658                self.update_error_metrics("HeadObject", &err);
659                return Err(ExternalError::from(
660                    anyhow!(err).context("s3 delete head err"),
661                ));
662            }
663        };
664        self.metrics.delete_object.inc();
665        let _ = self
666            .client
667            .delete_object()
668            .bucket(&self.bucket)
669            .key(&path)
670            .send()
671            .await
672            .inspect_err(|err| self.update_error_metrics("DeleteObject", err))
673            .context("s3 delete object err")?;
674        Ok(Some(usize::cast_from(size_bytes)))
675    }
676
677    async fn restore(&self, key: &str) -> Result<(), ExternalError> {
678        let path = self.get_path(key);
679        // Fetch the latest version of the object. If it's a normal version, return true;
680        // if it's a delete marker, delete it and loop; if there is no such version,
681        // return false.
682        // TODO: limit the number of delete markers we'll peel back?
683        loop {
684            // S3 only lets us fetch the versions of an object with a list requests.
685            // Seems a bit wasteful to just fetch one at a time, but otherwise we can only
686            // guess the order of versions via the timestamp, and that feels brittle.
687            let list_res = self
688                .client
689                .list_object_versions()
690                .bucket(&self.bucket)
691                .prefix(&path)
692                .max_keys(1)
693                .send()
694                .await
695                .inspect_err(|err| self.update_error_metrics("ListObjectVersions", err))
696                .context("listing object versions during restore")?;
697
698            let current_delete = list_res
699                .delete_markers()
700                .into_iter()
701                .filter(|d| {
702                    // We need to check that any versions we're looking at have the right key,
703                    // not just a key with our key as a prefix.
704                    d.key() == Some(path.as_str())
705                })
706                .find(|d| d.is_latest().unwrap_or(false))
707                .and_then(|d| d.version_id());
708
709            if let Some(version) = current_delete {
710                let deleted = self
711                    .client
712                    .delete_object()
713                    .bucket(&self.bucket)
714                    .key(&path)
715                    .version_id(version)
716                    .send()
717                    .await
718                    .inspect_err(|err| self.update_error_metrics("DeleteObject", err))
719                    .context("deleting a delete marker")?;
720                assert!(
721                    deleted.delete_marker().unwrap_or(false),
722                    "deleting a delete marker"
723                );
724            } else {
725                let has_current_version = list_res
726                    .versions()
727                    .into_iter()
728                    .filter(|d| d.key() == Some(path.as_str()))
729                    .any(|v| v.is_latest().unwrap_or(false));
730
731                if !has_current_version {
732                    return Err(Determinate::new(anyhow!(
733                        "unable to restore {key} in s3: no valid version exists"
734                    ))
735                    .into());
736                }
737                return Ok(());
738            }
739        }
740    }
741}
742
743impl S3Blob {
744    async fn set_single_part(&self, key: &str, value: Bytes) -> Result<(), ExternalError> {
745        let start_overall = Instant::now();
746        let path = self.get_path(key);
747
748        let value_len = value.len();
749        let part_span = trace_span!("s3set_single", payload_len = value_len);
750        self.metrics.set_single.inc();
751        self.client
752            .put_object()
753            .bucket(&self.bucket)
754            .key(path)
755            .body(ByteStream::from(value))
756            .send()
757            .instrument(part_span)
758            .await
759            .inspect_err(|err| self.update_error_metrics("PutObject", err))
760            .context("set single part")?;
761        debug!(
762            "s3 PutObject single done {}b / {:?}",
763            value_len,
764            start_overall.elapsed()
765        );
766        Ok(())
767    }
768
769    // TODO(benesch): remove this once this function no longer makes use of
770    // potentially dangerous `as` conversions.
771    #[allow(clippy::as_conversions)]
772    async fn set_multi_part(&self, key: &str, value: Bytes) -> Result<(), ExternalError> {
773        let start_overall = Instant::now();
774        let path = self.get_path(key);
775
776        // Start the multi part request and get an upload id.
777        trace!("s3 PutObject multi start {}b", value.len());
778        self.metrics.set_multi_create.inc();
779        let upload_res = self
780            .client
781            .create_multipart_upload()
782            .bucket(&self.bucket)
783            .key(&path)
784            .customize()
785            .mutate_request(|req| {
786                // By default the Rust AWS SDK does not set the Content-Length
787                // header on POST calls with empty bodies. This is fine for S3,
788                // but when running against GCS's S3 interop mode these calls
789                // will be rejected unless we set this header manually.
790                req.headers_mut().insert("Content-Length", "0");
791            })
792            .send()
793            .instrument(debug_span!("s3set_multi_start"))
794            .await
795            .inspect_err(|err| self.update_error_metrics("CreateMultipartUpload", err))
796            .context("create_multipart_upload err")?;
797        let upload_id = upload_res
798            .upload_id()
799            .ok_or_else(|| anyhow!("create_multipart_upload response missing upload_id"))?;
800        trace!(
801            "s3 create_multipart_upload took {:?}",
802            start_overall.elapsed()
803        );
804
805        let async_runtime = AsyncHandle::try_current().map_err(anyhow::Error::new)?;
806
807        // Fire off all the individual parts.
808        //
809        // TODO: The aws cli throttles how many of these are outstanding at any
810        // given point. We'll likely want to do the same at some point.
811        let start_parts = Instant::now();
812        let mut part_futs = Vec::new();
813        for (part_num, part_range) in self.multipart_config.part_iter(value.len()) {
814            // NB: Without this spawn, these will execute serially. This is rust
815            // async 101 stuff, but there isn't much async in the persist
816            // codebase (yet?) so I thought it worth calling out.
817            let part_span = debug_span!("s3set_multi_part", payload_len = part_range.len());
818            let part_fut = async_runtime.spawn_named(
819                // TODO: Add the key and part number once this can be annotated
820                // with metadata.
821                || "persist_s3blob_put_part",
822                {
823                    self.metrics.set_multi_part.inc();
824                    self.client
825                        .upload_part()
826                        .bucket(&self.bucket)
827                        .key(&path)
828                        .upload_id(upload_id)
829                        .part_number(part_num as i32)
830                        .body(ByteStream::from(value.slice(part_range)))
831                        .send()
832                        .instrument(part_span)
833                        .map(move |res| (start_parts.elapsed(), res))
834                },
835            );
836            part_futs.push((part_num, part_fut));
837        }
838        let parts_len = part_futs.len();
839
840        // Wait on all the parts to finish. This is done in part order, no need
841        // for joining them in the order they finish.
842        //
843        // TODO: Consider using something like futures::future::join_all() for
844        // this. That would cancel outstanding requests for us if any of them
845        // fails. However, it might not play well with using retries for tail
846        // latencies. Investigate.
847        let min_part_elapsed = MinElapsed::default();
848        let mut parts = Vec::with_capacity(parts_len);
849        for (part_num, part_fut) in part_futs.into_iter() {
850            let (this_part_elapsed, part_res) = part_fut.await;
851            let part_res = part_res
852                .inspect_err(|err| self.update_error_metrics("UploadPart", err))
853                .context("s3 upload_part err")?;
854            let part_e_tag = part_res.e_tag().ok_or_else(|| {
855                self.metrics
856                    .error_counts
857                    .with_label_values(&["UploadPart", "MissingEtag"])
858                    .inc();
859                anyhow!("s3 upload part missing e_tag")
860            })?;
861            parts.push(
862                CompletedPart::builder()
863                    .e_tag(part_e_tag)
864                    .part_number(part_num as i32)
865                    .build(),
866            );
867            min_part_elapsed.observe(this_part_elapsed, "s3 upload_part took");
868        }
869        trace!(
870            "s3 upload_parts overall took {:?} ({} parts)",
871            start_parts.elapsed(),
872            parts_len
873        );
874
875        // Complete the upload.
876        //
877        // Currently, we early return if any of the individual parts fail. This
878        // permanently orphans any parts that succeeded. One fix is to call
879        // abort_multipart_upload, which deletes them. However, there's also an
880        // option for an s3 bucket to auto-delete parts that haven't been
881        // completed or aborted after a given amount of time. This latter is
882        // simpler and also resilient to ill-timed mz restarts, so we use it for
883        // now. We could likely add the accounting necessary to make
884        // abort_multipart_upload work, but it would be complex and affect perf.
885        // Let's see how far we can get without it.
886        let start_complete = Instant::now();
887        self.metrics.set_multi_complete.inc();
888        self.client
889            .complete_multipart_upload()
890            .bucket(&self.bucket)
891            .key(&path)
892            .upload_id(upload_id)
893            .multipart_upload(
894                CompletedMultipartUpload::builder()
895                    .set_parts(Some(parts))
896                    .build(),
897            )
898            .send()
899            .instrument(debug_span!("s3set_multi_complete", num_parts = parts_len))
900            .await
901            .inspect_err(|err| self.update_error_metrics("CompleteMultipartUpload", err))
902            .context("complete_multipart_upload err")?;
903        trace!(
904            "s3 complete_multipart_upload took {:?}",
905            start_complete.elapsed()
906        );
907
908        debug!(
909            "s3 PutObject multi done {}b / {:?} ({} parts)",
910            value.len(),
911            start_overall.elapsed(),
912            parts_len
913        );
914        Ok(())
915    }
916
917    fn update_error_metrics<E, R>(&self, op: &str, err: &SdkError<E, R>)
918    where
919        E: ProvideErrorMetadata,
920    {
921        let code = match err {
922            SdkError::ServiceError(e) => match e.err().code() {
923                Some(code) => code,
924                None => "UnknownServiceError",
925            },
926            SdkError::DispatchFailure(e) => {
927                if let Some(other_error) = e.as_other() {
928                    match other_error {
929                        aws_config::retry::ErrorKind::TransientError => "TransientError",
930                        aws_config::retry::ErrorKind::ThrottlingError => "ThrottlingError",
931                        aws_config::retry::ErrorKind::ServerError => "ServerError",
932                        aws_config::retry::ErrorKind::ClientError => "ClientError",
933                        _ => "UnknownDispatchFailure",
934                    }
935                } else if e.is_timeout() {
936                    "TimeoutError"
937                } else if e.is_io() {
938                    "IOError"
939                } else if e.is_user() {
940                    "UserError"
941                } else {
942                    "UnknownDispathFailure"
943                }
944            }
945            SdkError::ResponseError(_) => "ResponseError",
946            SdkError::ConstructionFailure(_) => "ConstructionFailure",
947            // There is some overlap with MetricsSleep. MetricsSleep is more granular
948            // but does not contain the operation.
949            SdkError::TimeoutError(_) => "TimeoutError",
950            // an error was added at some point in the future
951            _ => "UnknownSdkError",
952        };
953        self.metrics
954            .error_counts
955            .with_label_values(&[op, code])
956            .inc();
957    }
958}
959
960/// One request of a multi-request `get`.
961enum PartRequest {
962    /// The first part, already fetched to learn the object's shape.
963    First(GetObjectOutput),
964    /// A piece of the object still to fetch.
965    Fetch(PartSelector),
966}
967
968/// How a `GetObject` request names the piece of the object it wants.
969enum PartSelector {
970    /// A part by the number it was uploaded with.
971    Number(i32),
972    /// A byte range, for stores that do not report the part count.
973    Range(Range<u64>),
974}
975
976/// The total size in a `Content-Range` header (`bytes 0-8388607/9711660`), if
977/// the header states one.
978fn parse_content_range_total(content_range: &str) -> Option<u64> {
979    content_range.rsplit_once('/')?.1.trim().parse().ok()
980}
981
982/// Byte ranges covering `[first_len, total_len)` in chunks of `first_len`.
983///
984/// A store reporting a zero-length first part leaves no chunk size to reuse,
985/// so the whole object is fetched in a single range instead.
986fn remaining_ranges(first_len: u64, total_len: u64) -> Vec<Range<u64>> {
987    let chunk_len = if first_len == 0 { total_len } else { first_len };
988    let mut ranges = Vec::new();
989    let mut start = first_len;
990    while start < total_len {
991        let end = cmp::min(start + chunk_len, total_len);
992        ranges.push(start..end);
993        start = end;
994    }
995    ranges
996}
997
998#[derive(Clone, Debug)]
999struct MultipartConfig {
1000    multipart_threshold: usize,
1001    multipart_chunk_size: usize,
1002}
1003
1004impl Default for MultipartConfig {
1005    fn default() -> Self {
1006        Self {
1007            multipart_threshold: Self::DEFAULT_MULTIPART_THRESHOLD,
1008            multipart_chunk_size: Self::DEFAULT_MULTIPART_CHUNK_SIZE,
1009        }
1010    }
1011}
1012
1013const MB: usize = 1024 * 1024;
1014const TB: usize = 1024 * 1024 * MB;
1015
1016impl MultipartConfig {
1017    /// The minimum object size for which we start using multipart upload.
1018    ///
1019    /// From the official `aws cli` tool implementation:
1020    ///
1021    /// <https://github.com/aws/aws-cli/blob/2.4.14/awscli/customizations/s3/transferconfig.py#L18-L29>
1022    const DEFAULT_MULTIPART_THRESHOLD: usize = 8 * MB;
1023    /// The size of each part (except the last) in a multipart upload.
1024    ///
1025    /// From the official `aws cli` tool implementation:
1026    ///
1027    /// <https://github.com/aws/aws-cli/blob/2.4.14/awscli/customizations/s3/transferconfig.py#L18-L29>
1028    const DEFAULT_MULTIPART_CHUNK_SIZE: usize = 8 * MB;
1029
1030    /// The largest size object creatable in S3.
1031    ///
1032    /// From <https://docs.aws.amazon.com/AmazonS3/latest/userguide/qfacts.html>
1033    const MAX_SINGLE_UPLOAD_SIZE: usize = 5 * TB;
1034    /// The minimum size of a part in a multipart upload.
1035    ///
1036    /// This minimum doesn't apply to the last chunk, which can be any size.
1037    ///
1038    /// From <https://docs.aws.amazon.com/AmazonS3/latest/userguide/qfacts.html>
1039    const MIN_UPLOAD_CHUNK_SIZE: usize = 5 * MB;
1040    /// The smallest allowable part number (inclusive).
1041    ///
1042    /// From <https://docs.aws.amazon.com/AmazonS3/latest/userguide/qfacts.html>
1043    const MIN_PART_NUM: u32 = 1;
1044    /// The largest allowable part number (inclusive).
1045    ///
1046    /// From <https://docs.aws.amazon.com/AmazonS3/latest/userguide/qfacts.html>
1047    const MAX_PART_NUM: u32 = 10_000;
1048
1049    fn should_multipart(&self, blob_len: usize) -> Result<bool, String> {
1050        if blob_len > Self::MAX_SINGLE_UPLOAD_SIZE {
1051            return Err(format!(
1052                "S3 does not support blobs larger than {} bytes got: {}",
1053                Self::MAX_SINGLE_UPLOAD_SIZE,
1054                blob_len
1055            ));
1056        }
1057        Ok(blob_len > self.multipart_threshold)
1058    }
1059
1060    fn part_iter(&self, blob_len: usize) -> MultipartChunkIter {
1061        mz_ore::soft_assert_no_log!(
1062            self.multipart_chunk_size >= MultipartConfig::MIN_UPLOAD_CHUNK_SIZE
1063        );
1064        MultipartChunkIter::new(self.multipart_chunk_size, blob_len)
1065    }
1066}
1067
1068#[derive(Clone, Debug)]
1069struct MultipartChunkIter {
1070    total_len: usize,
1071    part_size: usize,
1072    part_idx: u32,
1073}
1074
1075impl MultipartChunkIter {
1076    fn new(default_part_size: usize, blob_len: usize) -> Self {
1077        let max_parts: usize = usize::cast_from(MultipartConfig::MAX_PART_NUM);
1078
1079        // Compute the minimum part size we can use without going over the max
1080        // number of parts that S3 allows: `ceil(blob_len / max_parts)`.This
1081        // will end up getting thrown away by the `cmp::max` for anything
1082        // smaller than `max_parts * default_part_size = 80GiB`.
1083        let min_part_size = (blob_len + max_parts - 1) / max_parts;
1084        let part_size = cmp::max(min_part_size, default_part_size);
1085
1086        // Part nums are 1-indexed in S3. Convert back to 0-indexed to make the
1087        // range math easier to follow.
1088        let part_idx = MultipartConfig::MIN_PART_NUM - 1;
1089        MultipartChunkIter {
1090            total_len: blob_len,
1091            part_size,
1092            part_idx,
1093        }
1094    }
1095}
1096
1097impl Iterator for MultipartChunkIter {
1098    type Item = (u32, Range<usize>);
1099
1100    fn next(&mut self) -> Option<Self::Item> {
1101        let part_idx = self.part_idx;
1102        self.part_idx += 1;
1103
1104        let start = usize::cast_from(part_idx) * self.part_size;
1105        if start >= self.total_len {
1106            return None;
1107        }
1108        let end = cmp::min(start + self.part_size, self.total_len);
1109        let part_num = part_idx + 1;
1110        Some((part_num, start..end))
1111    }
1112}
1113
1114/// A helper for tracking the minimum of a set of Durations.
1115#[derive(Debug)]
1116struct MinElapsed {
1117    min: AtomicU64,
1118    alert_factor: u64,
1119}
1120
1121impl Default for MinElapsed {
1122    fn default() -> Self {
1123        MinElapsed {
1124            min: AtomicU64::new(u64::MAX),
1125            alert_factor: 8,
1126        }
1127    }
1128}
1129
1130impl MinElapsed {
1131    fn observe(&self, x: Duration, msg: &'static str) {
1132        let nanos = x.as_nanos();
1133        let nanos = u64::try_from(nanos).unwrap_or(u64::MAX);
1134
1135        // Possibly set a new minimum.
1136        let prev_min = self.min.fetch_min(nanos, atomic::Ordering::SeqCst);
1137
1138        // Trace if our provided duration was much larger than our minimum.
1139        let new_min = std::cmp::min(prev_min, nanos);
1140        if nanos > new_min.saturating_mul(self.alert_factor) {
1141            let min_duration = Duration::from_nanos(new_min);
1142            let factor = self.alert_factor;
1143            debug!("{msg} took {x:?} more than {factor}x the min {min_duration:?}");
1144        } else {
1145            trace!("{msg} took {x:?}");
1146        }
1147    }
1148}
1149
1150// Make sure the "vendored" feature of the openssl_sys crate makes it into the
1151// transitive dep graph of persist, so that we don't attempt to link against the
1152// system OpenSSL library. Fake a usage of the crate here so that a good
1153// samaritan doesn't remove our unused dep.
1154#[allow(dead_code)]
1155fn openssl_sys_hack() {
1156    openssl_sys::init();
1157}
1158
1159#[cfg(test)]
1160mod tests {
1161    use tracing::info;
1162
1163    use crate::location::tests::blob_impl_test;
1164
1165    use super::*;
1166
1167    #[mz_ore::test]
1168    fn content_range_total() {
1169        assert_eq!(
1170            parse_content_range_total("bytes 0-8388607/9711660"),
1171            Some(9711660)
1172        );
1173        assert_eq!(parse_content_range_total("bytes */42"), Some(42));
1174        assert_eq!(parse_content_range_total("bytes 0-1/*"), None);
1175        assert_eq!(parse_content_range_total("garbage"), None);
1176    }
1177
1178    #[mz_ore::test]
1179    fn ranges_for_unreported_parts() {
1180        let mib = 1024 * 1024;
1181        assert_eq!(remaining_ranges(8 * mib, 9711660), vec![8 * mib..9711660]);
1182        assert_eq!(
1183            remaining_ranges(8 * mib, 3 * 8 * mib + 5),
1184            vec![
1185                8 * mib..16 * mib,
1186                16 * mib..24 * mib,
1187                24 * mib..24 * mib + 5
1188            ]
1189        );
1190        assert_eq!(remaining_ranges(8 * mib, 8 * mib), Vec::<Range<u64>>::new());
1191        assert_eq!(remaining_ranges(0, 10), vec![0..10]);
1192    }
1193
1194    #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1195    #[cfg_attr(coverage, ignore)] // https://github.com/MaterializeInc/database-issues/issues/5586
1196    #[cfg_attr(miri, ignore)] // error: unsupported operation: can't call foreign function `TLS_method` on OS `linux`
1197    #[ignore] // TODO: Reenable against minio so it can run locally
1198    async fn s3_blob() -> Result<(), ExternalError> {
1199        let config = match S3BlobConfig::new_for_test().await? {
1200            Some(client) => client,
1201            None => {
1202                info!(
1203                    "{} env not set: skipping test that uses external service",
1204                    S3BlobConfig::EXTERNAL_TESTS_S3_BUCKET
1205                );
1206                return Ok(());
1207            }
1208        };
1209        let config_multipart = config.clone_with_new_uuid_prefix();
1210
1211        blob_impl_test(move |path| {
1212            let path = path.to_owned();
1213            let config = config.clone();
1214            async move {
1215                let config = S3BlobConfig {
1216                    metrics: config.metrics.clone(),
1217                    client: config.client.clone(),
1218                    bucket: config.bucket.clone(),
1219                    prefix: format!("{}/s3_blob_impl_test/{}", config.prefix, path),
1220                };
1221                let mut blob = S3Blob::open(config).await?;
1222                blob.max_keys = 2;
1223                Ok(blob)
1224            }
1225        })
1226        .await?;
1227
1228        // Also specifically test multipart. S3 requires all parts but the last
1229        // to be at least 5MB, which we don't want to do from a test, so this
1230        // uses the multipart code path but only writes a single part.
1231        {
1232            let blob = S3Blob::open(config_multipart).await?;
1233            blob.set_multi_part("multipart", "foobar".into()).await?;
1234            assert_eq!(
1235                blob.get("multipart").await?,
1236                Some(b"foobar".to_vec().into())
1237            );
1238        }
1239
1240        Ok(())
1241    }
1242
1243    /// Runs the conformance suite through [crate::hedge::HedgedBlob] with two
1244    /// genuinely independent S3 clients (separate connection pools) pointed
1245    /// at the same bucket/prefix, a hedge racing on every get. Ignored by
1246    /// default like `s3_blob` above. When run against the external test
1247    /// bucket, it is the one exercise of the real pool-isolation path.
1248    #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1249    #[cfg_attr(coverage, ignore)] // https://github.com/MaterializeInc/database-issues/issues/5586
1250    #[cfg_attr(miri, ignore)] // error: unsupported operation: can't call foreign function `TLS_method` on OS `linux`
1251    #[ignore] // TODO: Reenable against minio so it can run locally
1252    async fn s3_blob_hedged() -> Result<(), ExternalError> {
1253        use crate::hedge::{
1254            BLOB_HEDGED_GET_BUDGET_RATIO, BLOB_HEDGED_GET_DELAY, BLOB_HEDGED_GET_ENABLED,
1255            HedgeSibling, HedgedBlob,
1256        };
1257        use crate::metrics::BlobHedgeMetrics;
1258        use mz_dyncfg::{ConfigSet, ConfigUpdates};
1259
1260        let config = match S3BlobConfig::new_for_test().await? {
1261            Some(client) => client,
1262            None => return Ok(()),
1263        };
1264        // A second client with its own connection pool. Its generated prefix
1265        // is discarded below: both sides must point at the same store.
1266        let sibling = match S3BlobConfig::new_for_test().await? {
1267            Some(client) => client,
1268            None => return Ok(()),
1269        };
1270
1271        let cfg = crate::cfg::all_dyn_configs(ConfigSet::default());
1272        let mut updates = ConfigUpdates::default();
1273        updates.add(&BLOB_HEDGED_GET_ENABLED, true);
1274        updates.add(&BLOB_HEDGED_GET_DELAY, Duration::ZERO);
1275        updates.add(&BLOB_HEDGED_GET_BUDGET_RATIO, 1.0);
1276        updates.apply(&cfg);
1277        let cfg = Arc::new(cfg);
1278
1279        blob_impl_test(move |path| {
1280            let path = path.to_owned();
1281            let config = config.clone();
1282            let sibling = sibling.clone();
1283            let cfg = Arc::clone(&cfg);
1284            async move {
1285                let prefix = format!("{}/s3_blob_hedged_test/{}", config.prefix, path);
1286                let primary_config = S3BlobConfig {
1287                    metrics: config.metrics.clone(),
1288                    client: config.client.clone(),
1289                    bucket: config.bucket.clone(),
1290                    prefix: prefix.clone(),
1291                };
1292                let hedge_config = S3BlobConfig {
1293                    metrics: sibling.metrics.clone(),
1294                    client: sibling.client.clone(),
1295                    bucket: config.bucket.clone(),
1296                    prefix,
1297                };
1298                let primary: Arc<dyn Blob> = Arc::new(S3Blob::open(primary_config).await?);
1299                let hedge: Arc<dyn Blob> = Arc::new(S3Blob::open(hedge_config).await?);
1300                Ok(HedgedBlob::new(
1301                    primary,
1302                    HedgeSibling::Isolated(hedge),
1303                    cfg,
1304                    BlobHedgeMetrics::new(&MetricsRegistry::new()),
1305                ))
1306            }
1307        })
1308        .await?;
1309
1310        Ok(())
1311    }
1312
1313    #[mz_ore::test]
1314    fn should_multipart() {
1315        let config = MultipartConfig::default();
1316        assert_eq!(config.should_multipart(0), Ok(false));
1317        assert_eq!(config.should_multipart(1), Ok(false));
1318        assert_eq!(
1319            config.should_multipart(MultipartConfig::DEFAULT_MULTIPART_THRESHOLD),
1320            Ok(false)
1321        );
1322        assert_eq!(
1323            config.should_multipart(MultipartConfig::DEFAULT_MULTIPART_THRESHOLD + 1),
1324            Ok(true)
1325        );
1326        assert_eq!(
1327            config.should_multipart(MultipartConfig::DEFAULT_MULTIPART_THRESHOLD * 2),
1328            Ok(true)
1329        );
1330        assert_eq!(
1331            config.should_multipart(MultipartConfig::MAX_SINGLE_UPLOAD_SIZE),
1332            Ok(true)
1333        );
1334        assert_eq!(
1335            config.should_multipart(MultipartConfig::MAX_SINGLE_UPLOAD_SIZE + 1),
1336            Err(
1337                "S3 does not support blobs larger than 5497558138880 bytes got: 5497558138881"
1338                    .into()
1339            )
1340        );
1341    }
1342
1343    #[mz_ore::test]
1344    fn multipart_iter() {
1345        let iter = MultipartChunkIter::new(10, 0);
1346        assert_eq!(iter.collect::<Vec<_>>(), vec![]);
1347
1348        let iter = MultipartChunkIter::new(10, 9);
1349        assert_eq!(iter.collect::<Vec<_>>(), vec![(1, 0..9)]);
1350
1351        let iter = MultipartChunkIter::new(10, 10);
1352        assert_eq!(iter.collect::<Vec<_>>(), vec![(1, 0..10)]);
1353
1354        let iter = MultipartChunkIter::new(10, 11);
1355        assert_eq!(iter.collect::<Vec<_>>(), vec![(1, 0..10), (2, 10..11)]);
1356
1357        let iter = MultipartChunkIter::new(10, 19);
1358        assert_eq!(iter.collect::<Vec<_>>(), vec![(1, 0..10), (2, 10..19)]);
1359
1360        let iter = MultipartChunkIter::new(10, 20);
1361        assert_eq!(iter.collect::<Vec<_>>(), vec![(1, 0..10), (2, 10..20)]);
1362
1363        let iter = MultipartChunkIter::new(10, 21);
1364        assert_eq!(
1365            iter.collect::<Vec<_>>(),
1366            vec![(1, 0..10), (2, 10..20), (3, 20..21)]
1367        );
1368    }
1369
1370    /// None of the SDK timeouts covers a response body, so a body that stops
1371    /// arriving is bounded only by the SDK's stalled-stream protection (on by
1372    /// default for downloads). The hedged-gets design relies on it, so these
1373    /// tests pin it against a fake S3 endpoint.
1374    mod stalled_body {
1375        use std::net::SocketAddr;
1376
1377        use mz_dyncfg::{ConfigSet, ConfigUpdates};
1378        use mz_ore::task::{AbortOnDropHandle, JoinSetExt};
1379        use tokio::io::{AsyncReadExt, AsyncWriteExt};
1380        use tokio::net::{TcpListener, TcpStream};
1381        use tokio::task::JoinSet;
1382
1383        use crate::hedge::{BLOB_HEDGED_GET_ENABLED, HedgeSibling, HedgedBlob};
1384        use crate::metrics::BlobHedgeMetrics;
1385
1386        use super::*;
1387
1388        const BODY_LEN: usize = 1024 * 1024;
1389
1390        /// The production defaults of the `persist_blob_*_timeout` configs.
1391        #[derive(Debug)]
1392        struct ProdKnobs;
1393
1394        impl BlobKnobs for ProdKnobs {
1395            fn operation_timeout(&self) -> Duration {
1396                Duration::from_secs(180)
1397            }
1398            fn operation_attempt_timeout(&self) -> Duration {
1399                Duration::from_secs(90)
1400            }
1401            fn connect_timeout(&self) -> Duration {
1402                Duration::from_secs(7)
1403            }
1404            fn read_timeout(&self) -> Duration {
1405                Duration::from_secs(10)
1406            }
1407            fn is_cc_active(&self) -> bool {
1408                false
1409            }
1410        }
1411
1412        /// A fake S3 endpoint on 127.0.0.1. A GET of the key `target` gets a
1413        /// response head promising `BODY_LEN` bytes, then the body, or, if
1414        /// `stall` is set, no body byte at all while the connection stays
1415        /// open. Every other key gets a 404 NoSuchKey, which is what the health
1416        /// check in `S3Blob::open` expects.
1417        struct FakeS3 {
1418            addr: SocketAddr,
1419            _accept: AbortOnDropHandle<()>,
1420        }
1421
1422        impl FakeS3 {
1423            async fn start(stall: bool) -> FakeS3 {
1424                let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
1425                let addr = listener.local_addr().expect("local addr");
1426                let accept = mz_ore::task::spawn(|| "fake_s3_accept", async move {
1427                    // Dropping this task aborts every connection handler.
1428                    let mut conns = JoinSet::new();
1429                    loop {
1430                        let (stream, _) = listener.accept().await.expect("accept");
1431                        conns.spawn_named(|| "fake_s3_conn", serve(stream, stall));
1432                    }
1433                })
1434                .abort_on_drop();
1435                FakeS3 {
1436                    addr,
1437                    _accept: accept,
1438                }
1439            }
1440
1441            async fn open(&self) -> S3Blob {
1442                let config = S3BlobConfig::new(
1443                    "bucket".into(),
1444                    "prefix".into(),
1445                    None,
1446                    Some(format!("http://{}", self.addr)),
1447                    Some("us-east-1".into()),
1448                    Some(("user".into(), "pass".into())),
1449                    Box::new(ProdKnobs),
1450                    S3BlobMetrics::new(&MetricsRegistry::new()),
1451                )
1452                .await
1453                .expect("config");
1454                S3Blob::open(config).await.expect("open")
1455            }
1456        }
1457
1458        async fn serve(mut stream: TcpStream, stall: bool) {
1459            let mut buf = Vec::new();
1460            loop {
1461                let head_end = loop {
1462                    if let Some(pos) = buf.windows(4).position(|w| w == b"\r\n\r\n") {
1463                        break pos;
1464                    }
1465                    let mut chunk = [0u8; 4096];
1466                    match stream.read(&mut chunk).await {
1467                        Ok(0) | Err(_) => return,
1468                        Ok(n) => buf.extend_from_slice(&chunk[..n]),
1469                    }
1470                };
1471                let head = String::from_utf8_lossy(&buf[..head_end]).into_owned();
1472                buf.drain(..head_end + 4);
1473                let path = head.split_whitespace().nth(1).unwrap_or_default();
1474                if !path.contains("/target") {
1475                    let body = "<?xml version=\"1.0\" encoding=\"UTF-8\"?>\n\
1476                                <Error><Code>NoSuchKey</Code></Error>";
1477                    let resp = format!(
1478                        "HTTP/1.1 404 Not Found\r\n\
1479                         Content-Type: application/xml\r\n\
1480                         Content-Length: {}\r\n\r\n{body}",
1481                        body.len()
1482                    );
1483                    if stream.write_all(resp.as_bytes()).await.is_err() {
1484                        return;
1485                    }
1486                    continue;
1487                }
1488                let mut resp = format!(
1489                    "HTTP/1.1 206 Partial Content\r\n\
1490                     ETag: \"etag\"\r\n\
1491                     Content-Range: bytes 0-{last}/{BODY_LEN}\r\n\
1492                     Content-Type: application/octet-stream\r\n\
1493                     Content-Length: {BODY_LEN}\r\n\r\n",
1494                    last = BODY_LEN - 1,
1495                )
1496                .into_bytes();
1497                if !stall {
1498                    resp.resize(resp.len() + BODY_LEN, b'x');
1499                }
1500                if stream.write_all(&resp).await.is_err() {
1501                    return;
1502                }
1503                if stall {
1504                    // Hold the connection open until the client closes it.
1505                    let mut chunk = [0u8; 4096];
1506                    while let Ok(n) = stream.read(&mut chunk).await {
1507                        if n == 0 {
1508                            return;
1509                        }
1510                    }
1511                    return;
1512                }
1513            }
1514        }
1515
1516        #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1517        #[cfg_attr(miri, ignore)] // error: unsupported operation: can't call foreign function `TLS_method` on OS `linux`
1518        async fn stalled_body_fails_the_get() {
1519            let server = FakeS3::start(true).await;
1520            let blob = server.open().await;
1521            // The protection fires after about 6 s without a byte. The
1522            // generous timeout only catches a get that hangs.
1523            let res = tokio::time::timeout(Duration::from_secs(60), blob.get("target"))
1524                .await
1525                .expect("a stalled body must fail the get, not hang it");
1526            assert!(res.is_err(), "unexpected success");
1527        }
1528
1529        #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1530        #[cfg_attr(miri, ignore)] // error: unsupported operation: can't call foreign function `TLS_method` on OS `linux`
1531        async fn hedge_rescues_stalled_body() {
1532            let primary_server = FakeS3::start(true).await;
1533            let sibling_server = FakeS3::start(false).await;
1534            let primary: Arc<dyn Blob> = Arc::new(primary_server.open().await);
1535            let sibling: Arc<dyn Blob> = Arc::new(sibling_server.open().await);
1536            let cfg = crate::cfg::all_dyn_configs(ConfigSet::default());
1537            let mut updates = ConfigUpdates::default();
1538            updates.add(&BLOB_HEDGED_GET_ENABLED, true);
1539            updates.apply(&cfg);
1540            let metrics = BlobHedgeMetrics::new(&MetricsRegistry::new());
1541            let blob = HedgedBlob::new(
1542                primary,
1543                HedgeSibling::Isolated(sibling),
1544                Arc::new(cfg),
1545                metrics.clone(),
1546            );
1547            let res = tokio::time::timeout(Duration::from_secs(60), blob.get("target"))
1548                .await
1549                .expect("the hedge must rescue the get");
1550            let bytes = res.expect("get").expect("blob exists");
1551            assert_eq!(bytes.len(), BODY_LEN);
1552            assert_eq!(metrics.won.get(), 1);
1553        }
1554    }
1555}