1use 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#[derive(Clone, Debug)]
53pub struct S3BlobConfig {
54 metrics: S3BlobMetrics,
55 client: S3Client,
56 bucket: String,
57 prefix: String,
58}
59
60const 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 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 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 loader = loader.sleep_impl(MetricsSleep {
158 knobs,
159 metrics: metrics.clone(),
160 });
161 loader = loader.timeout_config(
162 TimeoutConfig::builder()
163 .operation_timeout(OPERATION_TIMEOUT_MARKER)
165 .operation_attempt_timeout(OPERATION_ATTEMPT_TIMEOUT_MARKER)
167 .connect_timeout(CONNECT_TIMEOUT_MARKER)
169 .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 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 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 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#[derive(Debug)]
288pub struct S3Blob {
289 metrics: S3BlobMetrics,
290 client: S3Client,
291 bucket: String,
292 prefix: String,
293 max_keys: i32,
297 multipart_config: MultipartConfig,
298}
299
300impl S3Blob {
301 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 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 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 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 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 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 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 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 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 let body_start = Instant::now();
476
477 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 let mut segments = vec![];
517 while let Some(result) = body_futures.next().await {
518 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 segments.append(&mut part_body);
530 }
531
532 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 let blob_key_prefix = self.get_path(key_prefix);
564 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 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 loop {
684 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 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 #[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 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 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 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 let part_span = debug_span!("s3set_multi_part", payload_len = part_range.len());
818 let part_fut = async_runtime.spawn_named(
819 || "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 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 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 SdkError::TimeoutError(_) => "TimeoutError",
950 _ => "UnknownSdkError",
952 };
953 self.metrics
954 .error_counts
955 .with_label_values(&[op, code])
956 .inc();
957 }
958}
959
960enum PartRequest {
962 First(GetObjectOutput),
964 Fetch(PartSelector),
966}
967
968enum PartSelector {
970 Number(i32),
972 Range(Range<u64>),
974}
975
976fn parse_content_range_total(content_range: &str) -> Option<u64> {
979 content_range.rsplit_once('/')?.1.trim().parse().ok()
980}
981
982fn 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 const DEFAULT_MULTIPART_THRESHOLD: usize = 8 * MB;
1023 const DEFAULT_MULTIPART_CHUNK_SIZE: usize = 8 * MB;
1029
1030 const MAX_SINGLE_UPLOAD_SIZE: usize = 5 * TB;
1034 const MIN_UPLOAD_CHUNK_SIZE: usize = 5 * MB;
1040 const MIN_PART_NUM: u32 = 1;
1044 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 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 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#[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 let prev_min = self.min.fetch_min(nanos, atomic::Ordering::SeqCst);
1137
1138 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#[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)] #[cfg_attr(miri, ignore)] #[ignore] 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 {
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 #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1249 #[cfg_attr(coverage, ignore)] #[cfg_attr(miri, ignore)] #[ignore] 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 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 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 #[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 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 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 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)] async fn stalled_body_fails_the_get() {
1519 let server = FakeS3::start(true).await;
1520 let blob = server.open().await;
1521 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)] 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}