1use anyhow::anyhow;
13use async_trait::async_trait;
14use azure_core::credentials::{AccessToken, Secret, TokenCredential, TokenRequestOptions};
15use azure_core::error::ErrorKind;
16use azure_core::http::headers::{HeaderName, Headers};
17use azure_core::http::{
18 AsyncResponseBody, ClientMethodOptions, ClientOptions, Etag, ExponentialRetryOptions,
19 RequestContent, RetryOptions, StatusCode, Transport,
20};
21use azure_identity::{
22 AzureCliCredential, ClientAssertion, ClientAssertionCredential, ClientSecretCredential,
23 ManagedIdentityCredential,
24};
25use azure_storage_blob::models::{
26 BlobClientDownloadOptions, BlobClientDownloadResult, BlobClientGetPropertiesResultHeaders,
27 BlobClientUploadOptions, BlobContainerClientListBlobsOptions, HttpRange,
28};
29use azure_storage_blob::{BlobContainerClient, BlobContainerClientOptions};
30use bytes::Bytes;
31use futures_util::future::BoxFuture;
32use futures_util::{FutureExt, StreamExt};
33use std::collections::BTreeMap;
34use std::fmt::{Debug, Formatter};
35use std::num::NonZero;
36use std::path::{Path, PathBuf};
37use std::sync::Arc;
38use std::time::Duration;
39use time::OffsetDateTime;
40use tokio::sync::RwLock;
41use tracing::{info, warn};
42use url::Url;
43use uuid::Uuid;
44
45use mz_ore::bytes::SegmentedBytes;
46use mz_ore::cast::CastFrom;
47use mz_ore::metrics::MetricsRegistry;
48use mz_ore::task::AbortOnDropHandle;
49
50use crate::cfg::BlobKnobs;
51use crate::error::Error;
52use crate::location::{Blob, BlobMetadata, Determinate, ExternalError};
53use crate::metrics::S3BlobMetrics;
54
55const AZURE_TENANT_ID: &str = "AZURE_TENANT_ID";
59const AZURE_CLIENT_ID: &str = "AZURE_CLIENT_ID";
60const AZURE_CLIENT_SECRET: &str = "AZURE_CLIENT_SECRET";
61const AZURE_FEDERATED_TOKEN: &str = "AZURE_FEDERATED_TOKEN";
62const AZURE_FEDERATED_TOKEN_FILE: &str = "AZURE_FEDERATED_TOKEN_FILE";
63
64const EMULATOR_ACCOUNT: &str = "devstoreaccount1";
66
67const EMULATOR_ACCOUNT_KEY: &str =
73 "Eby8vdM02xNOcqFlqUwJPLlmEtlCDXJ1OUzFT50uSRZ6IFsuFq2UVErCz4I6tq/K1SZFPTOtr/KBHBeksoGMGw==";
74
75const TOKEN_REFRESH_BUFFER: Duration = Duration::from_secs(5 * 60);
79
80const TOKEN_REFRESH_RETRY_INTERVAL: Duration = Duration::from_secs(10);
85
86const MANAGED_IDENTITY_TIMEOUT: Duration = Duration::from_secs(1);
91
92const GET_PARTITION_SIZE: NonZero<usize> = NonZero::new(1 << 20).unwrap();
96
97const GET_CONCURRENCY: usize = 8;
99
100type ExchangeFn = Arc<
103 dyn Fn(String, Vec<String>) -> BoxFuture<'static, azure_core::Result<AccessToken>>
104 + Send
105 + Sync,
106>;
107
108type TokenSlot = Arc<std::sync::RwLock<AccessToken>>;
110
111#[derive(Debug)]
113struct FixedAssertion(Secret);
114
115#[async_trait]
116impl ClientAssertion for FixedAssertion {
117 async fn secret(
118 &self,
119 _options: Option<ClientMethodOptions<'_>>,
120 ) -> azure_core::Result<String> {
121 Ok(self.0.secret().to_string())
122 }
123}
124
125struct RefreshingWorkloadIdentityCredential {
136 federated_token_file: PathBuf,
137 exchange: ExchangeFn,
138 cache: RwLock<BTreeMap<Vec<String>, (TokenSlot, AbortOnDropHandle<()>)>>,
142 refresh_buffer: Duration,
143 retry_interval: Duration,
144}
145
146impl Debug for RefreshingWorkloadIdentityCredential {
147 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
148 f.debug_struct("RefreshingWorkloadIdentityCredential")
149 .field("federated_token_file", &self.federated_token_file)
150 .finish_non_exhaustive()
151 }
152}
153
154impl RefreshingWorkloadIdentityCredential {
155 fn from_env() -> Option<azure_core::Result<Self>> {
159 if std::env::var(AZURE_FEDERATED_TOKEN).is_ok() {
163 return None;
164 }
165 let (Ok(tenant_id), Ok(client_id), Ok(token_file)) = (
166 std::env::var(AZURE_TENANT_ID),
167 std::env::var(AZURE_CLIENT_ID),
168 std::env::var(AZURE_FEDERATED_TOKEN_FILE),
169 ) else {
170 return None;
171 };
172 Some(Self::new(tenant_id, client_id, PathBuf::from(token_file)))
173 }
174
175 fn new(
176 tenant_id: String,
177 client_id: String,
178 federated_token_file: PathBuf,
179 ) -> azure_core::Result<Self> {
180 ClientAssertionCredential::new(
184 tenant_id.clone(),
185 client_id.clone(),
186 FixedAssertion(Secret::new("")),
187 None,
188 )?;
189 let exchange: ExchangeFn = Arc::new(move |assertion, scopes| {
190 let tenant_id = tenant_id.clone();
191 let client_id = client_id.clone();
192 async move {
193 let credential = ClientAssertionCredential::new(
197 tenant_id,
198 client_id,
199 FixedAssertion(Secret::new(assertion)),
200 None,
201 )?;
202 let scopes: Vec<&str> = scopes.iter().map(String::as_str).collect();
203 credential.get_token(&scopes, None).await
204 }
205 .boxed()
206 });
207 Ok(Self::with_exchange(
208 federated_token_file,
209 exchange,
210 TOKEN_REFRESH_BUFFER,
211 TOKEN_REFRESH_RETRY_INTERVAL,
212 ))
213 }
214
215 fn with_exchange(
216 federated_token_file: PathBuf,
217 exchange: ExchangeFn,
218 refresh_buffer: Duration,
219 retry_interval: Duration,
220 ) -> Self {
221 Self {
222 federated_token_file,
223 exchange,
224 cache: RwLock::new(BTreeMap::new()),
225 refresh_buffer,
226 retry_interval,
227 }
228 }
229
230 #[cfg(test)]
233 async fn clear_cache(&self) {
234 self.cache.write().await.clear();
236 }
237}
238
239async fn fetch_token(
242 federated_token_file: &Path,
243 exchange: &ExchangeFn,
244 scopes: Vec<String>,
245) -> azure_core::Result<AccessToken> {
246 let assertion = tokio::fs::read_to_string(federated_token_file)
247 .await
248 .map_err(|err| {
249 azure_core::Error::with_error(
250 ErrorKind::Credential,
251 err,
252 format!(
253 "failed to read federated token from file {}",
254 federated_token_file.display()
255 ),
256 )
257 })?;
258 (exchange)(assertion.trim().to_string(), scopes).await
262}
263
264async fn refresh_task(
268 federated_token_file: PathBuf,
269 exchange: ExchangeFn,
270 slot: TokenSlot,
271 scopes: Vec<String>,
272 refresh_buffer: Duration,
273 retry_interval: Duration,
274) {
275 loop {
276 let refresh_at = slot.read().expect("lock poisoned").expires_on - refresh_buffer;
277 let wait = refresh_at - OffsetDateTime::now_utc();
278 let wait = if wait.is_positive() {
279 wait.unsigned_abs()
280 } else {
281 Duration::ZERO
282 };
283 tokio::time::sleep(wait.max(retry_interval)).await;
284 match fetch_token(&federated_token_file, &exchange, scopes.clone()).await {
285 Ok(token) => *slot.write().expect("lock poisoned") = token,
286 Err(err) => {
287 warn!("failed to refresh Azure workload identity token, will retry: {err}")
288 }
289 }
290 }
291}
292
293#[async_trait]
294impl TokenCredential for RefreshingWorkloadIdentityCredential {
295 async fn get_token(
296 &self,
297 scopes: &[&str],
298 _options: Option<TokenRequestOptions<'_>>,
299 ) -> azure_core::Result<AccessToken> {
300 let scopes_key: Vec<String> = scopes.iter().map(ToString::to_string).collect();
301
302 {
303 let cache = self.cache.read().await;
304 if let Some((slot, _refresh)) = cache.get(&scopes_key) {
305 return Ok(slot.read().expect("lock poisoned").clone());
306 }
307 }
308
309 let mut cache = self.cache.write().await;
310 if let Some((slot, _refresh)) = cache.get(&scopes_key) {
311 return Ok(slot.read().expect("lock poisoned").clone());
312 }
313
314 let token = fetch_token(
318 &self.federated_token_file,
319 &self.exchange,
320 scopes_key.clone(),
321 )
322 .await?;
323 let slot = Arc::new(std::sync::RwLock::new(token.clone()));
324 let refresh = mz_ore::task::spawn(
325 || "azure-workload-identity-token-refresh",
326 refresh_task(
327 self.federated_token_file.clone(),
328 Arc::clone(&self.exchange),
329 Arc::clone(&slot),
330 scopes_key.clone(),
331 self.refresh_buffer,
332 self.retry_interval,
333 ),
334 )
335 .abort_on_drop();
336 cache.insert(scopes_key, (slot, refresh));
337 Ok(token)
338 }
339}
340
341#[derive(Debug)]
344struct ChainedCredential {
345 sources: Vec<Arc<dyn TokenCredential>>,
346}
347
348impl ChainedCredential {
349 fn default_chain() -> azure_core::Result<Self> {
355 let mut sources: Vec<Arc<dyn TokenCredential>> = Vec::new();
356 let env = |key| std::env::var(key).ok();
357 if let (Some(tenant_id), Some(client_id)) = (env(AZURE_TENANT_ID), env(AZURE_CLIENT_ID)) {
358 if let Some(token) = env(AZURE_FEDERATED_TOKEN) {
359 sources.push(ClientAssertionCredential::new(
360 tenant_id,
361 client_id,
362 FixedAssertion(Secret::new(token)),
363 None,
364 )?);
365 } else if let Some(secret) = env(AZURE_CLIENT_SECRET) {
366 sources.push(ClientSecretCredential::new(
367 &tenant_id,
368 client_id,
369 Secret::new(secret),
370 None,
371 )?);
372 }
373 }
374 match ManagedIdentityCredential::new(None) {
375 Ok(credential) => sources.push(Arc::new(TimeoutCredential {
376 inner: credential,
377 timeout: MANAGED_IDENTITY_TIMEOUT,
378 })),
379 Err(err) => info!("azure: managed identity credential unavailable: {err}"),
380 }
381 sources.push(AzureCliCredential::new(None)?);
382 Ok(Self { sources })
383 }
384}
385
386#[async_trait]
387impl TokenCredential for ChainedCredential {
388 async fn get_token(
389 &self,
390 scopes: &[&str],
391 options: Option<TokenRequestOptions<'_>>,
392 ) -> azure_core::Result<AccessToken> {
393 let mut errors = Vec::new();
394 for source in &self.sources {
395 match source.get_token(scopes, options.clone()).await {
396 Ok(token) => return Ok(token),
397 Err(err) => errors.push(err.to_string()),
398 }
399 }
400 Err(azure_core::Error::with_message(
401 ErrorKind::Credential,
402 format!(
403 "Multiple errors were encountered while attempting to authenticate:\n{}",
404 errors.join("\n")
405 ),
406 ))
407 }
408}
409
410#[derive(Debug)]
413struct TimeoutCredential {
414 inner: Arc<dyn TokenCredential>,
415 timeout: Duration,
416}
417
418#[async_trait]
419impl TokenCredential for TimeoutCredential {
420 async fn get_token(
421 &self,
422 scopes: &[&str],
423 options: Option<TokenRequestOptions<'_>>,
424 ) -> azure_core::Result<AccessToken> {
425 tokio::time::timeout(self.timeout, self.inner.get_token(scopes, options))
426 .await
427 .map_err(|_| {
428 azure_core::Error::with_message(
429 ErrorKind::Credential,
430 format!("token request timed out after {:?}", self.timeout),
431 )
432 })?
433 }
434}
435
436fn token_credential() -> Arc<dyn TokenCredential> {
445 match RefreshingWorkloadIdentityCredential::from_env() {
446 Some(credential) => {
447 info!("azure: using refreshing workload identity credentials");
448 Arc::new(credential.expect("Azure workload identity credentials"))
449 }
450 None => Arc::new(ChainedCredential::default_chain().expect("Azure default credentials")),
451 }
452}
453
454fn emulator_sas_token() -> azure_core::Result<String> {
461 const PERMISSIONS: &str = "rwdlac";
462 const SERVICES: &str = "b";
463 const RESOURCE_TYPES: &str = "sco";
464 const EXPIRY: &str = "2099-12-31T23:59:59Z";
467 const PROTOCOLS: &str = "https,http";
468 const VERSION: &str = "2022-11-02";
469 let string_to_sign = format!(
473 "{EMULATOR_ACCOUNT}\n{PERMISSIONS}\n{SERVICES}\n{RESOURCE_TYPES}\n\n{EXPIRY}\n\n\
474 {PROTOCOLS}\n{VERSION}\n\n"
475 );
476 let signature =
477 azure_core::hmac::hmac_sha256(&string_to_sign, &Secret::new(EMULATOR_ACCOUNT_KEY))?;
478 Ok(url::form_urlencoded::Serializer::new(String::new())
479 .append_pair("sv", VERSION)
480 .append_pair("ss", SERVICES)
481 .append_pair("srt", RESOURCE_TYPES)
482 .append_pair("sp", PERMISSIONS)
483 .append_pair("se", EXPIRY)
484 .append_pair("spr", PROTOCOLS)
485 .append_pair("sig", &signature)
486 .finish())
487}
488
489#[derive(Clone)]
495pub struct AzureBlobConfig {
496 metrics: S3BlobMetrics,
497 client: Arc<BlobContainerClient>,
498 prefix: String,
499 emulator: bool,
500}
501
502impl Debug for AzureBlobConfig {
503 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
504 let AzureBlobConfig {
505 metrics,
506 client,
507 prefix,
508 emulator,
509 } = self;
510 f.debug_struct("AzureBlobConfig")
511 .field("metrics", metrics)
512 .field("container_url", &redacted_url(client.url()))
513 .field("prefix", prefix)
514 .field("emulator", emulator)
515 .finish()
516 }
517}
518
519fn redacted_url(url: &Url) -> Url {
521 let mut url = url.clone();
522 url.set_query(None);
523 url
524}
525
526impl AzureBlobConfig {
527 const EXTERNAL_TESTS_AZURE_CONTAINER: &'static str =
528 "MZ_PERSIST_EXTERNAL_STORAGE_TEST_AZURE_CONTAINER";
529
530 pub fn new(
535 account: String,
536 container: String,
537 prefix: String,
538 metrics: S3BlobMetrics,
539 url: Url,
540 knobs: Box<dyn BlobKnobs>,
541 ) -> Result<Self, Error> {
542 let http_client = reqwest::ClientBuilder::new()
543 .timeout(knobs.operation_attempt_timeout())
544 .read_timeout(knobs.read_timeout())
545 .connect_timeout(knobs.connect_timeout())
546 .no_gzip()
548 .no_brotli()
549 .no_zstd()
550 .no_deflate()
551 .build()
552 .expect("valid config for azure HTTP client");
553 let operation_timeout = time::Duration::try_from(knobs.operation_timeout())
554 .expect("operation timeout fits a time::Duration");
555 let client_options = ClientOptions {
556 transport: Some(Transport::new(Arc::new(http_client))),
557 retry: RetryOptions::exponential(ExponentialRetryOptions {
558 max_total_elapsed: operation_timeout,
559 ..Default::default()
560 }),
561 ..Default::default()
562 };
563
564 let emulator = account == EMULATOR_ACCOUNT;
565 let (mut container_url, credential) = if emulator {
566 info!("Connecting to Azure emulator");
567 let mut container_url = Url::parse(&format!(
568 "http://{}:{}/",
569 url.domain().expect("domain for Azure emulator"),
570 url.port().expect("port for Azure emulator"),
571 ))
572 .expect("valid Azure emulator URL");
573 container_url
574 .path_segments_mut()
575 .expect("base URL")
576 .push(EMULATOR_ACCOUNT);
577 let sas_token = emulator_sas_token().expect("valid Azure emulator SAS token");
578 container_url.set_query(Some(&sas_token));
579 (container_url, None)
580 } else {
581 let mut container_url =
582 Url::parse(&format!("https://{account}.blob.core.windows.net/"))
583 .map_err(|err| format!("invalid Azure storage account {account}: {err}"))?;
584 let credential = match url.query() {
587 Some(query) => {
588 container_url.set_query(Some(query));
589 None
590 }
591 None => Some(token_credential()),
592 };
593 (container_url, credential)
594 };
595 container_url
596 .path_segments_mut()
597 .expect("base URL")
598 .push(&container);
599
600 let client = BlobContainerClient::new(
601 container_url,
602 credential,
603 Some(BlobContainerClientOptions {
604 client_options,
605 ..Default::default()
606 }),
607 )
608 .map_err(|err| format!("azure blob client: {err}"))?;
609
610 Ok(AzureBlobConfig {
616 metrics,
617 client: Arc::new(client),
618 prefix,
619 emulator,
620 })
621 }
622
623 pub fn new_for_test() -> Result<Option<Self>, Error> {
625 struct TestBlobKnobs;
626 impl Debug for TestBlobKnobs {
627 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
628 f.debug_struct("TestBlobKnobs").finish_non_exhaustive()
629 }
630 }
631 impl BlobKnobs for TestBlobKnobs {
632 fn operation_timeout(&self) -> Duration {
633 Duration::from_secs(30)
634 }
635
636 fn operation_attempt_timeout(&self) -> Duration {
637 Duration::from_secs(10)
638 }
639
640 fn connect_timeout(&self) -> Duration {
641 Duration::from_secs(5)
642 }
643
644 fn read_timeout(&self) -> Duration {
645 Duration::from_secs(5)
646 }
647
648 fn is_cc_active(&self) -> bool {
649 false
650 }
651 }
652
653 let container_name = match std::env::var(Self::EXTERNAL_TESTS_AZURE_CONTAINER) {
654 Ok(container) => container,
655 Err(_) => {
656 assert!(
657 !mz_ore::env::is_var_truthy("CI"),
658 "CI is supposed to run this test but something has gone wrong!"
659 );
660 return Ok(None);
661 }
662 };
663
664 let prefix = Uuid::new_v4().to_string();
665 let metrics = S3BlobMetrics::new(&MetricsRegistry::new());
666
667 let config = AzureBlobConfig::new(
668 EMULATOR_ACCOUNT.to_string(),
669 container_name.clone(),
670 prefix,
671 metrics,
672 Url::parse(&format!("http://localhost:40111/{}", container_name)).expect("valid url"),
673 Box::new(TestBlobKnobs),
674 )?;
675
676 Ok(Some(config))
677 }
678}
679
680pub struct AzureBlob {
682 metrics: S3BlobMetrics,
683 client: Arc<BlobContainerClient>,
684 prefix: String,
685}
686
687impl Debug for AzureBlob {
688 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
689 let AzureBlob {
690 metrics,
691 client,
692 prefix,
693 } = self;
694 f.debug_struct("AzureBlob")
695 .field("metrics", metrics)
696 .field("container_url", &redacted_url(client.url()))
697 .field("prefix", prefix)
698 .finish()
699 }
700}
701
702impl AzureBlob {
703 pub async fn open(config: AzureBlobConfig) -> Result<Self, ExternalError> {
705 if config.emulator {
706 if let Err(error) = config.client.create(None).await {
710 info!(
711 ?error,
712 "failed to create emulator container; this is expected on repeat runs"
713 );
714 }
715 }
716
717 let ret = AzureBlob {
718 metrics: config.metrics,
719 client: config.client,
720 prefix: config.prefix,
721 };
722
723 Ok(ret)
724 }
725
726 fn get_path(&self, key: &str) -> String {
727 format!("{}/{}", self.prefix, key)
728 }
729
730 fn blob_client(&self, path: &str) -> azure_storage_blob::BlobClient {
732 self.client.blob_client(path.trim_matches('/'))
736 }
737}
738
739fn download_total_len(headers: &Headers) -> Option<u64> {
745 match headers.get_optional_str(&HeaderName::from_static("content-range")) {
746 Some(range) => range.rsplit_once('/')?.1.parse().ok(),
747 None => headers
748 .get_optional_str(&HeaderName::from_static("content-length"))?
749 .parse()
750 .ok(),
751 }
752}
753
754async fn download_range(
759 blob: &azure_storage_blob::BlobClient,
760 offset: u64,
761 etag: Option<Etag>,
762) -> Result<Option<BlobClientDownloadResult>, ExternalError> {
763 let options = BlobClientDownloadOptions {
766 range: Some(HttpRange::new(
767 offset,
768 u64::cast_from(GET_PARTITION_SIZE.get()),
769 )),
770 partition_size: Some(GET_PARTITION_SIZE),
771 if_match: etag,
772 ..Default::default()
773 };
774 match blob.download(Some(options)).await {
775 Ok(response) => Ok(Some(response)),
776 Err(e) if e.http_status() == Some(StatusCode::NotFound) => Ok(None),
777 Err(e) => Err(ExternalError::from(e.with_context("azure blob get error"))),
778 }
779}
780
781async fn read_body(mut body: AsyncResponseBody) -> Result<Vec<Bytes>, ExternalError> {
783 let mut parts = Vec::new();
784 while let Some(part) = body.next().await {
785 parts.push(
786 part.map_err(|e| ExternalError::from(e.with_context("azure blob get body error")))?,
787 );
788 }
789 Ok(parts)
790}
791
792#[async_trait]
793impl Blob for AzureBlob {
794 async fn get(&self, key: &str) -> Result<Option<SegmentedBytes>, ExternalError> {
795 let path = self.get_path(key);
796 let blob = self.blob_client(&path);
797
798 let Some(first) = download_range(&blob, 0, None).await? else {
804 return Ok(None);
805 };
806 let content_length = download_total_len(&first.headers)
809 .ok_or_else(|| anyhow!("azure blob get error: no length for blob {path}"))?;
810 let etag = first.properties.etag.clone();
812
813 let mut segments = SegmentedBytes::new();
814 let mut total_len: u64 = 0;
815 for part in read_body(first.body).await? {
816 total_len += u64::cast_from(part.len());
817 segments.push(part);
818 }
819 let offsets = (total_len..content_length).step_by(GET_PARTITION_SIZE.get());
823 let mut rest = futures_util::stream::iter(offsets)
824 .map(|offset| {
825 let (blob, etag) = (&blob, etag.clone());
826 async move {
827 match download_range(blob, offset, etag).await? {
828 Some(response) => read_body(response.body).await.map(Some),
829 None => Ok(None),
830 }
831 }
832 })
833 .buffered(GET_CONCURRENCY);
834 while let Some(parts) = rest.next().await {
835 let Some(parts) = parts? else {
836 return Ok(None);
837 };
838 for part in parts {
839 total_len += u64::cast_from(part.len());
840 segments.push(part);
841 }
842 }
843
844 if content_length != total_len {
847 self.metrics.get_invalid_resp.inc();
848 }
849
850 Ok(Some(segments))
851 }
852
853 async fn list_keys_and_metadata(
854 &self,
855 key_prefix: &str,
856 f: &mut (dyn FnMut(BlobMetadata) + Send + Sync),
857 ) -> Result<(), ExternalError> {
858 let blob_key_prefix = self.get_path(key_prefix);
859 let strippable_root_prefix = format!("{}/", self.prefix);
860
861 let mut stream = self
862 .client
863 .list_blobs(Some(BlobContainerClientListBlobsOptions {
864 prefix: Some(blob_key_prefix),
865 ..Default::default()
866 }))
867 .map_err(|e| ExternalError::from(e.with_context("azure blob list error")))?;
868
869 while let Some(blob) = stream.next().await {
870 let blob =
871 blob.map_err(|e| ExternalError::from(e.with_context("azure blob list error")))?;
872 let Some(name) = blob.name else {
873 continue;
874 };
875
876 if let Some(key) = name.strip_prefix(&strippable_root_prefix) {
877 let size_in_bytes = blob
878 .properties
879 .and_then(|properties| properties.content_length)
880 .ok_or_else(|| anyhow!("azure blob list error: no size for blob {name}"))?;
881 f(BlobMetadata { key, size_in_bytes });
882 }
883 }
884
885 Ok(())
886 }
887
888 async fn set(&self, key: &str, value: Bytes) -> Result<(), ExternalError> {
889 let path = self.get_path(key);
890 let blob = self.blob_client(&path);
891
892 let options = BlobClientUploadOptions {
895 partition_size: Some(NonZero::<u64>::MAX),
896 ..Default::default()
897 };
898 let content: RequestContent<Bytes, _> = value.into();
899 blob.upload(content, Some(options))
900 .await
901 .map_err(|e| ExternalError::from(e.with_context("azure blob put error")))?;
902
903 Ok(())
904 }
905
906 async fn delete(&self, key: &str) -> Result<Option<usize>, ExternalError> {
907 let path = self.get_path(key);
908 let blob = self.blob_client(&path);
909
910 match blob.get_properties(None).await {
911 Ok(props) => {
912 let size = props
913 .content_length()
914 .map_err(|e| ExternalError::from(e.with_context("azure blob error")))?
915 .ok_or_else(|| anyhow!("azure blob error: no size for blob {path}"))?;
916 blob.delete(None)
917 .await
918 .map_err(|e| ExternalError::from(e.with_context("azure blob delete error")))?;
919 Ok(Some(usize::cast_from(size)))
920 }
921 Err(e) if e.http_status() == Some(StatusCode::NotFound) => Ok(None),
922 Err(e) => Err(ExternalError::from(e.with_context("azure blob error"))),
923 }
924 }
925
926 async fn restore(&self, key: &str) -> Result<(), ExternalError> {
927 let path = self.get_path(key);
928 let blob = self.blob_client(&path);
929
930 match blob.get_properties(None).await {
931 Ok(_) => Ok(()),
932 Err(e) if e.http_status() == Some(StatusCode::NotFound) => Err(Determinate::new(
933 anyhow!("azure blob error: unable to restore non-existent key {key}"),
934 )
935 .into()),
936 Err(e) => Err(ExternalError::from(e.with_context("azure blob error"))),
937 }
938 }
939}
940
941#[cfg(test)]
942mod tests {
943 use std::sync::Mutex;
944 use std::sync::atomic::{AtomicUsize, Ordering};
945
946 use azure_core::http::{AsyncRawResponse, HttpClient, Request};
947 use tracing::info;
948
949 use crate::location::tests::blob_impl_test;
950
951 use super::*;
952
953 struct DropCounter(Arc<AtomicUsize>);
955
956 impl Drop for DropCounter {
957 fn drop(&mut self) {
958 self.0.fetch_add(1, Ordering::SeqCst);
959 }
960 }
961
962 #[derive(Debug)]
966 struct HangingRangesClient {
967 started: Arc<AtomicUsize>,
969 dropped: Arc<AtomicUsize>,
971 }
972
973 #[async_trait]
974 impl HttpClient for HangingRangesClient {
975 async fn execute_request(&self, request: &Request) -> azure_core::Result<AsyncRawResponse> {
976 let partition = GET_PARTITION_SIZE.get();
977 let range = request
978 .headers()
979 .get_optional_str(&HeaderName::from_static("range"))
980 .unwrap_or_default();
981 if range.starts_with("bytes=0-") {
982 let mut headers = Headers::new();
983 headers.insert(
984 "content-range",
985 format!("bytes 0-{}/{}", partition - 1, 3 * partition),
986 );
987 headers.insert("content-length", partition.to_string());
988 headers.insert("etag", "\"v1\"");
989 return Ok(AsyncRawResponse::from_bytes(
990 StatusCode::PartialContent,
991 headers,
992 vec![0u8; partition],
993 ));
994 }
995 let _dropped = DropCounter(Arc::clone(&self.dropped));
996 self.started.fetch_add(1, Ordering::SeqCst);
997 std::future::pending().await
998 }
999 }
1000
1001 fn mock_blob(http_client: impl HttpClient + 'static) -> AzureBlob {
1003 let client = BlobContainerClient::new(
1004 Url::parse("https://account.blob.core.windows.net/container").expect("valid url"),
1005 None,
1006 Some(BlobContainerClientOptions {
1007 client_options: ClientOptions {
1008 transport: Some(Transport::new(Arc::new(http_client))),
1009 ..Default::default()
1010 },
1011 ..Default::default()
1012 }),
1013 )
1014 .expect("valid client");
1015 AzureBlob {
1016 metrics: S3BlobMetrics::new(&MetricsRegistry::new()),
1017 client: Arc::new(client),
1018 prefix: "prefix".to_string(),
1019 }
1020 }
1021
1022 #[derive(Debug)]
1025 struct FullBodyClient {
1026 with_length: bool,
1028 requests: Arc<AtomicUsize>,
1030 }
1031
1032 #[async_trait]
1033 impl HttpClient for FullBodyClient {
1034 async fn execute_request(
1035 &self,
1036 _request: &Request,
1037 ) -> azure_core::Result<AsyncRawResponse> {
1038 self.requests.fetch_add(1, Ordering::SeqCst);
1039 let len = 3 * GET_PARTITION_SIZE.get();
1040 let mut headers = Headers::new();
1041 if self.with_length {
1042 headers.insert("content-length", len.to_string());
1043 }
1044 headers.insert("etag", "\"v1\"");
1045 Ok(AsyncRawResponse::from_bytes(
1046 StatusCode::Ok,
1047 headers,
1048 vec![7u8; len],
1049 ))
1050 }
1051 }
1052
1053 #[mz_ore::test(tokio::test)]
1054 #[cfg_attr(miri, ignore)] async fn azure_blob_get_full_body_response() {
1056 let requests = Arc::new(AtomicUsize::new(0));
1057 let blob = mock_blob(FullBodyClient {
1058 with_length: true,
1059 requests: Arc::clone(&requests),
1060 });
1061 let value = blob
1062 .get("key")
1063 .await
1064 .expect("get")
1065 .expect("blob exists")
1066 .into_contiguous();
1067 assert_eq!(value.len(), 3 * GET_PARTITION_SIZE.get());
1069 assert!(value.iter().all(|b| *b == 7), "unexpected blob contents");
1070 assert_eq!(requests.load(Ordering::SeqCst), 1);
1071 assert_eq!(blob.metrics.get_invalid_resp.get(), 0);
1072 }
1073
1074 #[mz_ore::test(tokio::test)]
1075 #[cfg_attr(miri, ignore)] async fn azure_blob_get_without_length_fails() {
1077 let blob = mock_blob(FullBodyClient {
1078 with_length: false,
1079 requests: Arc::new(AtomicUsize::new(0)),
1080 });
1081 let Err(err) = blob.get("key").await else {
1082 panic!("get without length succeeded");
1083 };
1084 assert!(err.to_string().contains("no length"), "{err}");
1085 }
1086
1087 #[mz_ore::test(tokio::test)]
1090 #[cfg_attr(miri, ignore)] async fn azure_blob_get_drop_cancels_requests() {
1092 let started = Arc::new(AtomicUsize::new(0));
1093 let dropped = Arc::new(AtomicUsize::new(0));
1094 let http_client = HangingRangesClient {
1095 started: Arc::clone(&started),
1096 dropped: Arc::clone(&dropped),
1097 };
1098 let blob = mock_blob(http_client);
1099
1100 tokio::select! {
1103 _ = blob.get("key") => panic!("get of a hanging blob completed"),
1104 () = async {
1105 while started.load(Ordering::SeqCst) < 2 {
1106 tokio::task::yield_now().await;
1107 }
1108 } => {}
1109 }
1110 assert_eq!(dropped.load(Ordering::SeqCst), 2);
1111 }
1112
1113 struct MockExchange {
1116 assertions: Vec<String>,
1118 fail: bool,
1120 }
1121
1122 fn mock_exchange(state: &Arc<Mutex<MockExchange>>) -> ExchangeFn {
1123 let state = Arc::clone(state);
1124 Arc::new(move |assertion, _scopes| {
1125 let state = Arc::clone(&state);
1126 async move {
1127 let mut state = state.lock().unwrap();
1128 state.assertions.push(assertion);
1129 if state.fail {
1130 return Err(azure_core::Error::with_message(
1131 ErrorKind::Credential,
1132 "mock exchange failure",
1133 ));
1134 }
1135 Ok(AccessToken::new(
1136 Secret::new(format!("aad-{}", state.assertions.len())),
1137 OffsetDateTime::now_utc() + Duration::from_secs(3600),
1138 ))
1139 }
1140 .boxed()
1141 })
1142 }
1143
1144 #[mz_ore::test(tokio::test)]
1148 async fn refreshing_workload_identity_credential() {
1149 let token_file = tempfile::NamedTempFile::new().expect("create temp token file");
1150 std::fs::write(token_file.path(), "token-a\n").expect("write token file");
1151
1152 let state = Arc::new(Mutex::new(MockExchange {
1153 assertions: Vec::new(),
1154 fail: false,
1155 }));
1156 let credential = RefreshingWorkloadIdentityCredential::with_exchange(
1157 token_file.path().to_path_buf(),
1158 mock_exchange(&state),
1159 TOKEN_REFRESH_BUFFER,
1160 TOKEN_REFRESH_RETRY_INTERVAL,
1161 );
1162 let scopes = &["https://storage.azure.com/"];
1163
1164 let token = credential.get_token(scopes, None).await.expect("token");
1165 assert_eq!(token.token.secret(), "aad-1");
1166 let token = credential.get_token(scopes, None).await.expect("token");
1167 assert_eq!(token.token.secret(), "aad-1");
1168 assert_eq!(state.lock().unwrap().assertions, vec!["token-a"]);
1169
1170 std::fs::write(token_file.path(), "token-b").expect("write token file");
1173 credential.clear_cache().await;
1174 state.lock().unwrap().fail = true;
1175 assert!(credential.get_token(scopes, None).await.is_err());
1176 state.lock().unwrap().fail = false;
1177 let token = credential.get_token(scopes, None).await.expect("token");
1178 assert_eq!(token.token.secret(), "aad-3");
1179 assert_eq!(
1180 state.lock().unwrap().assertions,
1181 vec!["token-a", "token-b", "token-b"]
1182 );
1183 }
1184
1185 #[mz_ore::test(tokio::test)]
1188 async fn workload_identity_credential_background_refresh() {
1189 let token_file = tempfile::NamedTempFile::new().expect("create temp token file");
1190 std::fs::write(token_file.path(), "token-a").expect("write token file");
1191
1192 let state = Arc::new(Mutex::new(MockExchange {
1193 assertions: Vec::new(),
1194 fail: false,
1195 }));
1196 let credential = RefreshingWorkloadIdentityCredential::with_exchange(
1200 token_file.path().to_path_buf(),
1201 mock_exchange(&state),
1202 Duration::from_secs(7200),
1203 Duration::from_millis(10),
1204 );
1205 let scopes = &["https://storage.azure.com/"];
1206
1207 let token = credential.get_token(scopes, None).await.expect("token");
1208 assert_eq!(token.token.secret(), "aad-1");
1209
1210 std::fs::write(token_file.path(), "token-b").expect("write token file");
1213 tokio::time::timeout(Duration::from_secs(30), async {
1214 loop {
1215 let token = credential.get_token(scopes, None).await.expect("token");
1216 if token.token.secret() != "aad-1" {
1217 break;
1218 }
1219 tokio::time::sleep(Duration::from_millis(10)).await;
1220 }
1221 })
1222 .await
1223 .expect("token refreshed within timeout");
1224 assert_eq!(
1225 state.lock().unwrap().assertions.last().map(String::as_str),
1226 Some("token-b")
1227 );
1228
1229 state.lock().unwrap().fail = true;
1231 let held = credential.get_token(scopes, None).await.expect("token");
1232 let calls_when_failing = state.lock().unwrap().assertions.len();
1233 tokio::time::timeout(Duration::from_secs(30), async {
1234 while state.lock().unwrap().assertions.len() <= calls_when_failing + 2 {
1235 tokio::time::sleep(Duration::from_millis(10)).await;
1236 }
1237 })
1238 .await
1239 .expect("retries within timeout");
1240 let token = credential.get_token(scopes, None).await.expect("token");
1241 assert_eq!(token.token.secret(), held.token.secret());
1242 }
1243
1244 #[mz_ore::test]
1245 fn download_total_len_from_headers() {
1246 let mut ranged = Headers::new();
1247 ranged.insert("content-range", "bytes 0-1023/4096");
1248 ranged.insert("content-length", "1024");
1249 assert_eq!(download_total_len(&ranged), Some(4096));
1250
1251 let mut full = Headers::new();
1252 full.insert("content-length", "17");
1253 assert_eq!(download_total_len(&full), Some(17));
1254
1255 assert_eq!(download_total_len(&Headers::new()), None);
1256 }
1257
1258 #[cfg_attr(miri, ignore)] #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1260 async fn azure_blob() -> Result<(), ExternalError> {
1261 let config = match AzureBlobConfig::new_for_test()? {
1262 Some(client) => client,
1263 None => {
1264 info!(
1265 "{} env not set: skipping test that uses external service",
1266 AzureBlobConfig::EXTERNAL_TESTS_AZURE_CONTAINER
1267 );
1268 return Ok(());
1269 }
1270 };
1271
1272 blob_impl_test(move |_path| {
1273 let config = config.clone();
1274 async move {
1275 let config = AzureBlobConfig {
1276 metrics: config.metrics.clone(),
1277 client: Arc::clone(&config.client),
1278 prefix: config.prefix.clone(),
1279 emulator: config.emulator,
1280 };
1281 AzureBlob::open(config).await
1282 }
1283 })
1284 .await
1285 }
1286
1287 #[cfg_attr(miri, ignore)] #[mz_ore::test(tokio::test(flavor = "multi_thread"))]
1289 async fn azure_blob_multi_range_get() -> Result<(), ExternalError> {
1290 let Some(config) = AzureBlobConfig::new_for_test()? else {
1291 info!(
1292 "{} env not set: skipping test that uses external service",
1293 AzureBlobConfig::EXTERNAL_TESTS_AZURE_CONTAINER
1294 );
1295 return Ok(());
1296 };
1297 let blob = AzureBlob::open(config).await?;
1298
1299 let len = 2 * GET_PARTITION_SIZE.get() + 12345;
1301 let value: Vec<u8> = (0..251u8).cycle().take(len).collect();
1302 blob.set("large", Bytes::from(value.clone())).await?;
1303
1304 let got = blob.get("large").await?.expect("blob exists");
1305 assert_eq!(got.into_contiguous(), value);
1306 assert_eq!(blob.metrics.get_invalid_resp.get(), 0);
1307
1308 let mut listed = Vec::new();
1309 blob.list_keys_and_metadata("", &mut |m| {
1310 listed.push((m.key.to_string(), m.size_in_bytes))
1311 })
1312 .await?;
1313 assert_eq!(listed, vec![("large".to_string(), u64::cast_from(len))]);
1314
1315 assert_eq!(blob.delete("large").await?, Some(len));
1316 assert_eq!(blob.get("large").await?, None);
1317 Ok(())
1318 }
1319}