From 992bc103c8897689bd9e169a4162cd1a14f20ce9 Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Thu, 30 Jul 2026 22:25:12 +0100 Subject: [PATCH 1/7] feat(rest): support refreshing vended storage credentials --- Cargo.lock | 2 + crates/catalog/rest/Cargo.toml | 1 + crates/catalog/rest/public-api.txt | 15 + crates/catalog/rest/src/catalog.rs | 53 +- crates/catalog/rest/src/client.rs | 59 +- crates/catalog/rest/src/credential.rs | 1193 +++++++++++++++++++ crates/catalog/rest/src/lib.rs | 1 + crates/catalog/rest/src/types.rs | 38 +- crates/iceberg/public-api.txt | 43 + crates/iceberg/src/io/file_io.rs | 27 +- crates/iceberg/src/io/storage/config/gcs.rs | 6 + crates/iceberg/src/io/storage/config/mod.rs | 12 +- crates/iceberg/src/io/storage/config/s3.rs | 6 + crates/iceberg/src/io/storage/mod.rs | 102 +- crates/storage/opendal/Cargo.toml | 3 +- crates/storage/opendal/public-api.txt | 4 + crates/storage/opendal/src/gcs.rs | 119 +- crates/storage/opendal/src/lib.rs | 163 ++- crates/storage/opendal/src/resolving.rs | 38 +- crates/storage/opendal/src/s3.rs | 90 +- crates/storage/opendal/src/utils.rs | 22 + 21 files changed, 1934 insertions(+), 63 deletions(-) create mode 100644 crates/catalog/rest/src/credential.rs diff --git a/Cargo.lock b/Cargo.lock index b2269dc51f..8a359a3404 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3889,6 +3889,7 @@ dependencies = [ "iceberg_test_utils", "itertools 0.13.0", "mockito", + "rand 0.9.5", "reqwest 0.12.28", "serde", "serde_derive", @@ -4035,6 +4036,7 @@ dependencies = [ "opendal", "reqsign-aws-v4", "reqsign-core", + "reqsign-google", "reqwest 0.12.28", "serde", "tokio", diff --git a/crates/catalog/rest/Cargo.toml b/crates/catalog/rest/Cargo.toml index 247709efd4..d9526aa6bf 100644 --- a/crates/catalog/rest/Cargo.toml +++ b/crates/catalog/rest/Cargo.toml @@ -35,6 +35,7 @@ chrono = { workspace = true } http = { workspace = true } iceberg = { workspace = true } itertools = { workspace = true } +rand = { workspace = true } reqwest = { workspace = true } serde = { workspace = true } serde_derive = { workspace = true } diff --git a/crates/catalog/rest/public-api.txt b/crates/catalog/rest/public-api.txt index 776b11c40a..c3e226fc1c 100644 --- a/crates/catalog/rest/public-api.txt +++ b/crates/catalog/rest/public-api.txt @@ -139,6 +139,20 @@ impl serde_core::ser::Serialize for iceberg_catalog_rest::ListTablesResponse pub fn iceberg_catalog_rest::ListTablesResponse::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_catalog_rest::ListTablesResponse pub fn iceberg_catalog_rest::ListTablesResponse::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg_catalog_rest::LoadCredentialsResponse +pub iceberg_catalog_rest::LoadCredentialsResponse::storage_credentials: alloc::vec::Vec +impl core::clone::Clone for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::clone(&self) -> iceberg_catalog_rest::LoadCredentialsResponse +impl core::cmp::Eq for iceberg_catalog_rest::LoadCredentialsResponse +impl core::cmp::PartialEq for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::eq(&self, other: &iceberg_catalog_rest::LoadCredentialsResponse) -> bool +impl core::fmt::Debug for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result +impl core::marker::StructuralPartialEq for iceberg_catalog_rest::LoadCredentialsResponse +impl serde_core::ser::Serialize for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer +impl<'de> serde_core::de::Deserialize<'de> for iceberg_catalog_rest::LoadCredentialsResponse +pub fn iceberg_catalog_rest::LoadCredentialsResponse::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg_catalog_rest::LoadTableResult pub iceberg_catalog_rest::LoadTableResult::config: std::collections::hash::map::HashMap pub iceberg_catalog_rest::LoadTableResult::metadata: iceberg::spec::table_metadata::TableMetadata @@ -288,5 +302,6 @@ pub fn iceberg_catalog_rest::UpdateNamespacePropertiesResponse::serialize<__S>(& impl<'de> serde_core::de::Deserialize<'de> for iceberg_catalog_rest::UpdateNamespacePropertiesResponse pub fn iceberg_catalog_rest::UpdateNamespacePropertiesResponse::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub const iceberg_catalog_rest::REST_CATALOG_PROP_DISABLE_HEADER_REDACTION: &str +pub const iceberg_catalog_rest::REST_CATALOG_PROP_SCAN_PLAN_ID: &str pub const iceberg_catalog_rest::REST_CATALOG_PROP_URI: &str pub const iceberg_catalog_rest::REST_CATALOG_PROP_WAREHOUSE: &str diff --git a/crates/catalog/rest/src/catalog.rs b/crates/catalog/rest/src/catalog.rs index 8642b32d22..1a400fd7da 100644 --- a/crates/catalog/rest/src/catalog.rs +++ b/crates/catalog/rest/src/catalog.rs @@ -54,6 +54,8 @@ pub const REST_CATALOG_PROP_URI: &str = "uri"; pub const REST_CATALOG_PROP_WAREHOUSE: &str = "warehouse"; /// Disable header redaction in error logs (defaults to false for security) pub const REST_CATALOG_PROP_DISABLE_HEADER_REDACTION: &str = "disable-header-redaction"; +/// Identifier for a server-side scan plan associated with credential requests. +pub const REST_CATALOG_PROP_SCAN_PLAN_ID: &str = "rest.scan.plan-id"; const ICEBERG_REST_SPEC_VERSION: &str = "0.14.1"; const CARGO_PKG_VERSION: &str = env!("CARGO_PKG_VERSION"); @@ -164,7 +166,7 @@ impl RestCatalogBuilder { } /// Rest catalog configuration. -#[derive(Clone, Debug, TypedBuilder)] +#[derive(Clone, TypedBuilder)] pub(crate) struct RestCatalogConfig { #[builder(default, setter(strip_option))] name: Option, @@ -181,6 +183,19 @@ pub(crate) struct RestCatalogConfig { client: Option, } +impl std::fmt::Debug for RestCatalogConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Catalog and table properties may contain OAuth credentials, delegated + // storage credentials, or arbitrary secret-bearing `header.*` values. + f.debug_struct("RestCatalogConfig") + .field("name", &self.name) + .field("uri", &self.uri) + .field("warehouse", &self.warehouse) + .field("property_keys", &self.props.keys().collect::>()) + .finish_non_exhaustive() + } +} + impl RestCatalogConfig { fn url_prefixed(&self, parts: &[&str]) -> String { [&self.uri, PATH_V1] @@ -358,7 +373,7 @@ impl RestCatalogConfig { #[derive(Debug)] struct RestContext { - client: HttpClient, + client: Arc, /// Runtime config is fetched from rest server and stored here. /// /// It's could be different from the user config. @@ -447,7 +462,7 @@ impl RestCatalog { Ok(RestContext { config, - client, + client: Arc::new(client), endpoints, }) }) @@ -510,17 +525,17 @@ impl RestCatalog { metadata_location: Option<&str>, extra_config: Option>, ) -> Result { - let mut props = self.context().await?.config.props.clone(); + let context = self.context().await?; + let mut props = context.config.props.clone(); if let Some(config) = extra_config { props.extend(config); } // If the warehouse is a logical identifier instead of a URL we don't want // to raise an exception - let warehouse_path = match self.context().await?.config.warehouse.as_deref() { + let warehouse_path = match context.config.warehouse.as_deref() { Some(url) if Url::parse(url).is_ok() => Some(url), - Some(_) => None, - None => None, + _ => None, }; if metadata_location.or(warehouse_path).is_none() { @@ -541,9 +556,20 @@ impl RestCatalog { ) })?; - let file_io = FileIOBuilder::new(factory).with_props(props).build(); + // If the catalog vends refreshable credentials for this table's storage, + // attach a provider so the backend re-fetches them before they expire. + let credential_provider = crate::credential::build_vended_credential_provider( + context.client.clone(), + &context.config.uri, + &props, + )?; + + let mut builder = FileIOBuilder::new(factory).with_props(props); + if let Some(provider) = credential_provider { + builder = builder.with_credential_provider(provider); + } - Ok(file_io) + Ok(builder.build()) } /// Invalidate the current token without generating a new one. On the next request, the client @@ -1048,7 +1074,14 @@ impl Catalog for RestCatalog { "Metadata location missing in `register_table` response!", ))?; - let file_io = self.load_file_io(Some(metadata_location), None).await?; + let config = response + .config + .into_iter() + .chain(self.user_config.props.clone()) + .collect(); + let file_io = self + .load_file_io(Some(metadata_location), Some(config)) + .await?; let mut table_builder = Table::builder() .identifier(table_ident.clone()) diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index 07dc0620da..cc4b1c2f97 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -49,9 +49,13 @@ pub(crate) struct HttpClient { impl Debug for HttpClient { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + // Omit the reqwest client: injected clients may carry secret default + // headers. Explicit headers use the same redaction policy as errors. f.debug_struct("HttpClient") - .field("client", &self.client) - .field("extra_headers", &self.extra_headers) + .field( + "extra_headers", + &format_headers_redacted(&self.extra_headers, self.disable_header_redaction), + ) .finish_non_exhaustive() } } @@ -71,6 +75,24 @@ impl HttpClient { }) } + /// Create a client for table-scoped resources while reusing this client's + /// underlying connection pool. + /// + /// A load-table response may supply a table token or `header.*` values that + /// must be used for subsequent table requests such as credential refresh. + pub(crate) fn for_table( + &self, + catalog_uri: &str, + props: HashMap, + ) -> Result { + let cfg = RestCatalogConfig::builder() + .uri(catalog_uri.to_string()) + .props(props) + .client(Some(self.client.clone())) + .build(); + Self::new(&cfg) + } + /// Update the http client with new configuration. /// /// If cfg carries new value, we will use cfg instead. @@ -156,7 +178,6 @@ impl HttpClient { ) .with_context("operation", "auth") .with_context("url", auth_url.to_string()) - .with_context("json", String::from_utf8_lossy(&text)) .with_source(e) })?) } else { @@ -170,7 +191,6 @@ impl HttpClient { .with_context("code", code.to_string()) .with_context("operation", "auth") .with_context("url", auth_url.to_string()) - .with_context("json", String::from_utf8_lossy(&text)) .with_source(e) })?; Err(Error::from(e)) @@ -278,29 +298,30 @@ pub(crate) async fn deserialize_catalog_response( let bytes = response.bytes().await?; serde_json::from_slice::(&bytes).map_err(|e| { + // Successful REST responses can contain OAuth tokens and delegated + // storage credentials. Never copy an unparseable response into an error. Error::new( ErrorKind::Unexpected, "Failed to parse response from rest catalog server", ) - .with_context("json", String::from_utf8_lossy(&bytes)) .with_source(e) }) } -/// Headers that contain sensitive information and should be excluded from logs. -const SENSITIVE_HEADERS: &[&str] = &[ - "authorization", - "proxy-authorization", - "set-cookie", - "cookie", - "x-api-key", - "x-auth-token", -]; - -/// Returns true if the header name is considered sensitive. +/// Returns true if the header may carry a secret. fn is_sensitive_header(name: &str) -> bool { let name_lower = name.to_lowercase(); - SENSITIVE_HEADERS.iter().any(|h| name_lower == *h) + [ + "auth", + "token", + "secret", + "key", + "password", + "cookie", + "credential", + ] + .iter() + .any(|pattern| name_lower.contains(pattern)) } /// Redacts sensitive headers and returns a debug-formatted string. @@ -386,6 +407,8 @@ mod tests { fn test_format_headers_redacted_filters_sensitive() { let mut headers = HeaderMap::new(); headers.insert("authorization", "Bearer secret-token".parse().unwrap()); + headers.insert("x-client-secret", "private-value".parse().unwrap()); + headers.insert("x-client-credential", "credential-value".parse().unwrap()); headers.insert("content-type", "application/json".parse().unwrap()); let result = format_headers_redacted(&headers, false); @@ -395,6 +418,8 @@ mod tests { assert!(result.contains("[REDACTED]")); // Sensitive value should NOT be present assert!(!result.contains("secret-token")); + assert!(!result.contains("private-value")); + assert!(!result.contains("credential-value")); // Non-sensitive header should be present with actual value assert!(result.contains("content-type")); assert!(result.contains("application/json")); diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs new file mode 100644 index 0000000000..14363b3a71 --- /dev/null +++ b/crates/catalog/rest/src/credential.rs @@ -0,0 +1,1193 @@ +// Licensed to the Apache Software Foundation (ASF) under one +// or more contributor license agreements. See the NOTICE file +// distributed with this work for additional information +// regarding copyright ownership. The ASF licenses this file +// to you under the Apache License, Version 2.0 (the +// "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +//! Refresh of vended storage credentials against a REST catalog. +//! +//! A REST catalog can vend short-lived storage credentials whose lifetime the +//! client does not control. [`RestVendedCredentialProvider`] implements the +//! core [`StorageCredentialProvider`] trait so storage backends re-fetch those +//! credentials from the catalog's table credentials endpoint before they +//! expire, keeping long-running jobs authenticated instead of failing with a +//! `403` once the initial token's TTL elapses. +//! +//! Unlike the Java client, which has one provider per cloud SDK, this is a +//! single backend-agnostic provider with an independent endpoint and cache for +//! each configured cloud. The path being accessed selects the cloud cache, and +//! the returned [`StorageCredential`] enum lets the storage adapter enforce the +//! expected backend-specific type. This preserves Java's per-cloud refresh +//! policy while supporting mixed-cloud tables through a resolving FileIO. +//! +//! # Adding a cloud +//! +//! The refresh policy for each cloud lives in one [`CloudRefresh`] constant. To +//! add a backend, first add its credential type to Iceberg's storage API and +//! teach the storage adapter to consume it. Then write its `parse_*` function, +//! add a `CloudRefresh` constant, and list it in [`CloudRefresh::SUPPORTED`]. + +use std::collections::HashMap; +use std::sync::Arc; +use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; + +use async_trait::async_trait; +use iceberg::io::{ + AWS_REFRESH_CREDENTIALS_ENABLED, AWS_REFRESH_CREDENTIALS_ENDPOINT, + GCS_REFRESH_CREDENTIALS_ENABLED, GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, + GCS_TOKEN_EXPIRES_AT, GcsCredential, S3_ACCESS_KEY_ID, S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, + S3_SESSION_TOKEN_EXPIRES_AT_MS, S3Credential, StorageCredential, StorageCredentialKind, + StorageCredentialProvider, +}; +use iceberg::{Error, ErrorKind, Result}; +use rand::Rng; +use reqwest::{Method, StatusCode, Url}; +use tokio::sync::Mutex; + +use crate::REST_CATALOG_PROP_SCAN_PLAN_ID; +use crate::client::{HttpClient, deserialize_unexpected_catalog_error}; +use crate::types::LoadCredentialsResponse; + +/// Cloud-specific details regarding vended-credential refresh. +/// +/// It contains the location schemes it backs, the property keys it is configured +/// with, and how to parse its credential. The generic provider stays free of +/// any per-cloud knowledge. +struct CloudRefresh { + /// Location URL schemes this backend serves. + schemes: &'static [&'static str], + /// Table property naming the refresh endpoint (absolute or catalog-relative). + endpoint_key: &'static str, + /// Table property to opt out; refresh is enabled unless this is `"false"`. + enabled_key: &'static str, + /// Whether to jitter successful prefetch times like AWS `CachedSupplier`. + jitter_prefetch: bool, + /// Parse a complete credential from catalog-supplied properties. + parse_credential: fn(&HashMap) -> Result, +} + +impl CloudRefresh { + /// S3 / AWS + const AWS: Self = Self { + schemes: &["s3", "s3a", "s3n"], + endpoint_key: AWS_REFRESH_CREDENTIALS_ENDPOINT, + enabled_key: AWS_REFRESH_CREDENTIALS_ENABLED, + jitter_prefetch: true, + parse_credential: parse_s3_credential, + }; + /// Google Cloud Storage + const GCP: Self = Self { + schemes: &["gs", "gcs"], + endpoint_key: GCS_REFRESH_CREDENTIALS_ENDPOINT, + enabled_key: GCS_REFRESH_CREDENTIALS_ENABLED, + jitter_prefetch: false, + parse_credential: parse_gcs_credential, + }; + // TODO: Azure (ADLS) is not yet supported: opendal 0.57's Azdls builder exposes no + // custom credential-provider hook, and reqsign's SAS-token credential has no + // expiry, so reqsign-based refresh isn't possible. + + /// Backends with refresh support + const SUPPORTED: &[Self] = &[Self::AWS, Self::GCP]; + + /// The backend that serves `location`, by its URL scheme, or `None` if no + /// supported backend matches (in which case static credentials are used + /// as-is, as before). + fn for_location(location: &str) -> Option<&'static Self> { + let scheme = scheme_of(location)?; + Self::SUPPORTED + .iter() + .find(|cloud| cloud.schemes.contains(&scheme.as_str())) + } + + fn matches_credential_prefix(&self, prefix: &str) -> bool { + scheme_of(prefix).is_some_and(|scheme| self.schemes.contains(&scheme.as_str())) + } +} + +/// Re-fetch a credential once it is within this window of expiry, so a fresh +/// token is in hand before the object store would reject the old one. +const REFRESH_BUFFER: Duration = Duration::from_mins(5); + +/// AWS keeps at least one minute between its jittered prefetch time and expiry. +const MIN_REFRESH_BUFFER: Duration = Duration::from_mins(1); + +/// Initial ceiling for failure backoff. Equal jitter chooses from half this +/// value through the full value. +const INITIAL_FAILURE_BACKOFF: Duration = Duration::from_secs(1); + +/// Maximum failure backoff while a cached credential remains usable. +const MAX_FAILURE_BACKOFF: Duration = Duration::from_secs(30); + +/// A vended credential paired with the storage-location prefix it applies to. +#[derive(Clone)] +struct CachedEntry { + prefix: String, + credential: StorageCredential, + /// When this entry becomes eligible for prefetch. `None` means it does not + /// expire and therefore never needs proactive refresh. + refresh_at: Option, +} + +impl CachedEntry { + fn new(prefix: String, credential: StorageCredential, jitter_prefetch: bool) -> Self { + let refresh_at = credential + .expires_at + .map(|expires_at| prefetch_time(expires_at, jitter_prefetch)); + Self { + prefix, + credential, + refresh_at, + } + } + + /// Seed entries that are already inside the nominal five-minute window are + /// immediately due. Otherwise AWS applies the same jitter as it does to a + /// freshly fetched value. + fn seed(prefix: String, credential: StorageCredential, jitter_prefetch: bool) -> Self { + let due = credential.expires_at.is_some_and(|expires_at| { + SystemTime::now() + .checked_add(REFRESH_BUFFER) + .is_none_or(|refresh_boundary| refresh_boundary >= expires_at) + }); + let mut entry = Self::new(prefix, credential, jitter_prefetch); + if due { + entry.refresh_at = Some(UNIX_EPOCH); + } + entry + } + + fn is_fresh(&self, now: SystemTime) -> bool { + self.refresh_at.is_none_or(|refresh_at| now < refresh_at) + } + + fn is_unexpired(&self, now: SystemTime) -> bool { + self.credential + .expires_at + .is_none_or(|expires_at| now < expires_at) + } +} + +/// Cached credentials plus failure-backoff state. +struct CacheState { + entries: Vec, + consecutive_failures: u32, + retry_not_before: Option, +} + +struct ConfiguredCloud { + cloud: &'static CloudRefresh, + endpoint: String, + cache: Mutex, + /// Only one caller fetches at a time. The cache lock is deliberately + /// separate so other callers can keep using an unexpired credential while + /// the refresh is in flight. + refresh: Mutex<()>, +} + +/// Fetches and refreshes vended credentials from a REST catalog's table +/// credentials endpoint. +/// +/// Each cloud cache is seeded with the credential from the initial table +/// properties (when complete) and re-fetched from its endpoint as it +/// nears expiry. +pub(crate) struct RestVendedCredentialProvider { + client: Arc, + /// Optional scan-plan identifier. + plan_id: Option, + /// Independently configured endpoint and cache for each backing cloud. + clouds: Vec, +} + +impl RestVendedCredentialProvider { + fn new(client: Arc, plan_id: Option, clouds: Vec) -> Self { + Self { + client, + plan_id, + clouds, + } + } + + /// Fetch fresh credentials from the catalog's credentials endpoint. + async fn fetch(&self, configured: &ConfiguredCloud) -> Result> { + let mut request = self.client.request(Method::GET, &configured.endpoint); + if let Some(plan_id) = &self.plan_id { + request = request.query(&[("planId", plan_id)]); + } + let request = request.build()?; + let response = self.client.query_catalog(request).await?; + + match response.status() { + StatusCode::OK => { + let parsed: LoadCredentialsResponse = response.json().await?; + parsed + .storage_credentials + .into_iter() + .filter(|sc| configured.cloud.matches_credential_prefix(&sc.prefix)) + .map(|sc| { + (configured.cloud.parse_credential)(&sc.config).map(|credential| { + CachedEntry::new( + sc.prefix, + credential, + configured.cloud.jitter_prefetch, + ) + }) + }) + .collect() + } + _ => Err(deserialize_unexpected_catalog_error( + response, + self.client.disable_header_redaction(), + ) + .await), + } + } + + async fn refresh_credential( + &self, + configured: &ConfiguredCloud, + path: &str, + fallback: Option, + ) -> Result { + let refreshed = self.fetch(configured).await.and_then(|entries| { + let credential = longest_prefix_match(&entries, path) + .filter(|entry| entry.is_unexpired(SystemTime::now())) + .map(|entry| entry.credential.clone()) + .ok_or_else(|| { + Error::new( + ErrorKind::Unexpected, + format!("no unexpired vended credential matches storage location: {path}"), + ) + })?; + Ok((entries, credential)) + }); + + match refreshed { + Ok((entries, credential)) => { + let mut cache = configured.cache.lock().await; + cache.entries = entries; + cache.consecutive_failures = 0; + cache.retry_not_before = None; + Ok(credential) + } + Err(fetch_error) => { + let mut cache = configured.cache.lock().await; + cache.consecutive_failures = cache.consecutive_failures.saturating_add(1); + cache.retry_not_before = + Instant::now().checked_add(failure_backoff(cache.consecutive_failures)); + + // Graceful degradation: while the cached credential remains + // usable, serve it and retry after jittered backoff. Expired + // credentials are never served. + fallback + .filter(|entry| entry.is_unexpired(SystemTime::now())) + .map(|entry| entry.credential) + .ok_or(fetch_error) + } + } + } +} + +impl std::fmt::Debug for RestVendedCredentialProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RestVendedCredentialProvider") + .field("configured_clouds", &self.clouds.len()) + .finish_non_exhaustive() + } +} + +enum CacheDecision { + Use(StorageCredential), + Refresh(Option), +} + +async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> CacheDecision { + let cache = configured.cache.lock().await; + let current = longest_prefix_match(&cache.entries, path).cloned(); + let now = SystemTime::now(); + + if let Some(entry) = current.as_ref().filter(|entry| entry.is_fresh(now)) { + return CacheDecision::Use(entry.credential.clone()); + } + + if cache + .retry_not_before + .is_some_and(|retry_at| Instant::now() < retry_at) + && let Some(entry) = current.as_ref().filter(|entry| entry.is_unexpired(now)) + { + return CacheDecision::Use(entry.credential.clone()); + } + + CacheDecision::Refresh(current) +} + +#[async_trait] +impl StorageCredentialProvider for RestVendedCredentialProvider { + fn supports_path(&self, path: &str) -> bool { + CloudRefresh::for_location(path).is_some_and(|cloud| { + self.clouds + .iter() + .any(|configured| configured.cloud.endpoint_key == cloud.endpoint_key) + }) + } + + async fn load_credential(&self, path: &str) -> Result { + let cloud = CloudRefresh::for_location(path).ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!("no credential refresh implementation for storage location: {path}"), + ) + })?; + let configured = self + .clouds + .iter() + .find(|configured| configured.cloud.endpoint_key == cloud.endpoint_key) + .ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!("credential refresh is not configured for storage location: {path}"), + ) + })?; + + let current = match cache_decision(configured, path).await { + CacheDecision::Use(credential) => return Ok(credential), + CacheDecision::Refresh(current) => current, + }; + + // One caller refreshes, while concurrent callers immediately keep using the + // unexpired cached credential. With no usable credential, callers wait + // for the in-flight refresh instead. + let usable = current + .as_ref() + .filter(|entry| entry.is_unexpired(SystemTime::now())); + let _refresh_guard = if let Some(entry) = usable { + match configured.refresh.try_lock() { + Ok(guard) => guard, + Err(_) => return Ok(entry.credential.clone()), + } + } else { + configured.refresh.lock().await + }; + + // Another caller may have completed a refresh between our cache check + // and acquiring the single-flight guard. + let current = match cache_decision(configured, path).await { + CacheDecision::Use(credential) => return Ok(credential), + CacheDecision::Refresh(current) => current, + }; + + self.refresh_credential(configured, path, current).await + } +} + +/// Select the credential whose prefix is the longest match for `path`. +fn longest_prefix_match<'a>(entries: &'a [CachedEntry], path: &str) -> Option<&'a CachedEntry> { + entries + .iter() + .filter(|entry| path.starts_with(&entry.prefix)) + .max_by_key(|entry| entry.prefix.len()) +} + +/// Compute a successful credential's prefetch time. +fn prefetch_time(expires_at: SystemTime, jitter: bool) -> SystemTime { + let base = expires_at.checked_sub(REFRESH_BUFFER).unwrap_or(UNIX_EPOCH); + if !jitter { + return base; + } + + let jitter_window = REFRESH_BUFFER.saturating_sub(MIN_REFRESH_BUFFER); + let jitter_millis = rand::rng().random_range(0..jitter_window.as_millis() as u64); + base.checked_add(Duration::from_millis(jitter_millis)) + .unwrap_or(base) +} + +/// Equal-jitter exponential backoff. The random lower half avoids both hot +/// retry loops and synchronized retries across clients. +fn failure_backoff(consecutive_failures: u32) -> Duration { + let exponent = consecutive_failures.saturating_sub(1).min(5); + let ceiling = INITIAL_FAILURE_BACKOFF + .checked_mul(1 << exponent) + .unwrap_or(MAX_FAILURE_BACKOFF) + .min(MAX_FAILURE_BACKOFF); + let ceiling_millis = ceiling.as_millis() as u64; + let floor_millis = ceiling_millis / 2; + Duration::from_millis(rand::rng().random_range(floor_millis..=ceiling_millis)) +} + +/// Build a credential provider from a table's properties, +/// or `None` when no supported cloud advertises an enabled refresh endpoint. +/// +/// `base_uri` is the catalog URI, used to resolve a relative endpoint. +pub(crate) fn build_vended_credential_provider( + client: Arc, + base_uri: &str, + props: &HashMap, +) -> Result>> { + let clouds = CloudRefresh::SUPPORTED + .iter() + .filter_map(|cloud| { + // Refresh is enabled by default and invalid booleans disable refresh. + let enabled = props + .get(cloud.enabled_key) + .is_none_or(|value| value.parse().unwrap_or(false)); + if !enabled { + return None; + } + + let endpoint = resolve_endpoint( + base_uri, + props + .get(cloud.endpoint_key) + .filter(|value| !value.is_empty())?, + ); + let entries = (cloud.parse_credential)(props) + .ok() + .map(|credential| { + vec![CachedEntry::seed( + // Flat table properties carry no prefix. This cache is + // already cloud-specific, so an empty prefix is safe. + String::new(), + credential, + cloud.jitter_prefetch, + )] + }) + .unwrap_or_default(); + + Some(ConfiguredCloud { + cloud, + endpoint, + cache: Mutex::new(CacheState { + entries, + consecutive_failures: 0, + retry_not_before: None, + }), + refresh: Mutex::new(()), + }) + }) + .collect::>(); + + if clouds.is_empty() { + return Ok(None); + } + + let table_client = Arc::new(client.for_table(base_uri, props.clone())?); + let plan_id = props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(); + Ok(Some(Arc::new(RestVendedCredentialProvider::new( + table_client, + plan_id, + clouds, + )))) +} + +/// Resolve a possibly-relative refresh endpoint against the catalog base URI. +fn resolve_endpoint(base_uri: &str, endpoint: &str) -> String { + if endpoint.starts_with("http://") || endpoint.starts_with("https://") { + return endpoint.to_string(); + } + + let base = base_uri.trim_end_matches('/'); + let separator = if endpoint.starts_with('/') { "" } else { "/" }; + format!("{base}{separator}{endpoint}") +} + +/// The URL scheme of `location`, lowercased (e.g. `"s3"` for `s3://bucket/k`). +fn scheme_of(location: &str) -> Option { + Url::parse(location) + .ok() + .map(|url| url.scheme().to_string()) +} + +/// Parse a complete S3 credential returned by the credentials endpoint. +fn parse_s3_credential(config: &HashMap) -> Result { + let access_key_id = required_nonempty(config, S3_ACCESS_KEY_ID)?; + let secret_access_key = required_nonempty(config, S3_SECRET_ACCESS_KEY)?; + let session_token = required_nonempty(config, S3_SESSION_TOKEN)?; + let expires_at = required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?; + Ok(StorageCredential { + kind: StorageCredentialKind::S3(S3Credential { + access_key_id, + secret_access_key, + session_token: Some(session_token), + }), + expires_at: Some(expires_at), + }) +} + +/// Parse a complete GCS credential returned by the credentials endpoint. +fn parse_gcs_credential(config: &HashMap) -> Result { + let token = required_nonempty(config, GCS_TOKEN)?; + let expires_at = required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?; + Ok(StorageCredential { + kind: StorageCredentialKind::Gcs(GcsCredential { token }), + expires_at: Some(expires_at), + }) +} + +fn required_nonempty(config: &HashMap, key: &str) -> Result { + config + .get(key) + .filter(|value| !value.is_empty()) + .cloned() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid vended credential: {key} is missing or empty"), + ) + }) +} + +fn required_epoch_millis(config: &HashMap, key: &str) -> Result { + let value = required_nonempty(config, key)?; + parse_epoch_millis(&value).ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid vended credential: {key} is not a valid epoch-millisecond timestamp"), + ) + }) +} + +/// Parse an epoch-millisecond timestamp into a [`SystemTime`]. +fn parse_epoch_millis(millis: &str) -> Option { + millis + .parse() + .ok() + .and_then(|millis| UNIX_EPOCH.checked_add(Duration::from_millis(millis))) +} + +#[cfg(test)] +mod tests { + use std::sync::{Barrier, mpsc}; + + use mockito::{Matcher, Server}; + + use super::*; + use crate::RestCatalogConfig; + + fn epoch_millis(time: SystemTime) -> String { + time.duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() + .to_string() + } + + fn s3_cred(access_key_id: &str, expires_at: Option) -> StorageCredential { + StorageCredential { + kind: StorageCredentialKind::S3(S3Credential { + access_key_id: access_key_id.to_string(), + secret_access_key: "secret".to_string(), + session_token: None, + }), + expires_at, + } + } + + fn s3_access_key_id(credential: &StorageCredential) -> &str { + match &credential.kind { + StorageCredentialKind::S3(s3) => &s3.access_key_id, + other => panic!("expected S3 credential, got {other:?}"), + } + } + + fn cached_s3(prefix: &str, access_key_id: &str, expires_at: Option) -> CachedEntry { + CachedEntry::new( + prefix.to_string(), + s3_cred(access_key_id, expires_at), + false, + ) + } + + #[test] + fn resolve_endpoint_matches_java_semantics() { + assert_eq!( + resolve_endpoint("https://catalog/", "https://other/creds"), + "https://other/creds" + ); + assert_eq!( + resolve_endpoint("https://catalog", "http://other/creds"), + "http://other/creds" + ); + assert_eq!( + resolve_endpoint("https://catalog", "v1/creds"), + "https://catalog/v1/creds" + ); + assert_eq!( + resolve_endpoint("https://catalog/", "v1/creds"), + "https://catalog/v1/creds" + ); + // All trailing slashes stripped from the base (Java stripTrailingSlash). + assert_eq!( + resolve_endpoint("https://catalog///", "/v1/creds"), + "https://catalog/v1/creds" + ); + // Existing leading slashes on the endpoint are preserved, not collapsed + // (matches Java's resolveEndpoint, which only prepends when absent). + assert_eq!( + resolve_endpoint("https://catalog/", "//v1/creds"), + "https://catalog//v1/creds" + ); + } + + #[test] + fn cloud_selection_uses_url_scheme() { + assert!(CloudRefresh::for_location("s3://b/k").is_some()); + assert!(CloudRefresh::for_location("S3://b/k").is_some()); + assert!(CloudRefresh::for_location("s3a://b/k").is_some()); + assert!(CloudRefresh::for_location("s3n://b/k").is_some()); + assert!(CloudRefresh::for_location("gs://b/k").is_some()); + assert!(CloudRefresh::for_location("gcs://b/k").is_some()); + assert!(CloudRefresh::for_location("not a url").is_none()); + // Azure not supported yet -> no provider, static creds used as-is. + assert!(CloudRefresh::for_location("abfss://fs@acct.dfs.core.windows.net/k").is_none()); + assert!(CloudRefresh::AWS.matches_credential_prefix("s3://bucket/path")); + assert!(!CloudRefresh::AWS.matches_credential_prefix("s3evil://bucket/path")); + } + + #[test] + fn parses_s3_credential() { + let mut config = HashMap::new(); + assert!(parse_s3_credential(&config).is_err()); + config.insert(S3_ACCESS_KEY_ID.to_string(), "AK".to_string()); + assert!(parse_s3_credential(&config).is_err()); + config.insert(S3_SECRET_ACCESS_KEY.to_string(), "SK".to_string()); + assert!(parse_s3_credential(&config).is_err()); + + config.insert(S3_SESSION_TOKEN.to_string(), "TOK".to_string()); + config.insert( + S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), + "1500".to_string(), + ); + let credential = parse_s3_credential(&config).unwrap(); + assert_eq!( + credential.expires_at, + Some(UNIX_EPOCH + Duration::from_millis(1500)) + ); + match credential.kind { + StorageCredentialKind::S3(s3) => assert_eq!(s3.session_token.as_deref(), Some("TOK")), + other => panic!("expected S3, got {other:?}"), + } + } + + #[test] + fn parse_gcs_requires_token_and_expiry() { + let mut config = HashMap::new(); + assert!(parse_gcs_credential(&config).is_err()); + config.insert(GCS_TOKEN.to_string(), "ya29.token".to_string()); + assert!(parse_gcs_credential(&config).is_err()); + + config.insert(GCS_TOKEN_EXPIRES_AT.to_string(), "2000".to_string()); + let credential = parse_gcs_credential(&config).unwrap(); + match &credential.kind { + StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token, "ya29.token"), + other => panic!("expected GCS, got {other:?}"), + } + assert_eq!( + credential.expires_at, + Some(UNIX_EPOCH + Duration::from_millis(2000)) + ); + } + + #[test] + fn cached_entry_freshness() { + let no_expiry = cached_s3("", "a", None); + let far = cached_s3( + "", + "a", + Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)), + ); + // Within the buffer but not yet expired: stale for a fast-path read, but + // still usable for graceful degradation. + let soon = cached_s3("", "a", Some(SystemTime::now() + Duration::from_secs(60))); + let past = cached_s3("", "a", Some(SystemTime::now() - Duration::from_secs(60))); + + let now = SystemTime::now(); + assert!(no_expiry.is_fresh(now)); + assert!(no_expiry.is_unexpired(now)); + assert!(far.is_fresh(now)); + assert!(!soon.is_fresh(now)); + assert!(soon.is_unexpired(now)); + assert!(!past.is_unexpired(now)); + } + + #[test] + fn prefetch_time_matches_cloud_policy() { + let expires_at = SystemTime::now() + Duration::from_secs(3600); + assert_eq!( + prefetch_time(expires_at, false), + expires_at - REFRESH_BUFFER + ); + + for _ in 0..100 { + let refresh_at = prefetch_time(expires_at, true); + assert!(refresh_at >= expires_at - REFRESH_BUFFER); + assert!(refresh_at < expires_at - MIN_REFRESH_BUFFER); + } + } + + #[test] + fn seed_inside_nominal_window_is_immediately_due() { + let credential = s3_cred( + "a", + Some(SystemTime::now() + REFRESH_BUFFER - Duration::from_secs(1)), + ); + let entry = CachedEntry::seed(String::new(), credential, true); + let now = SystemTime::now(); + assert!(!entry.is_fresh(now)); + assert!(entry.is_unexpired(now)); + } + + #[test] + fn failure_backoff_is_jittered_and_capped() { + for failures in 1_u32..=40 { + let ceiling = INITIAL_FAILURE_BACKOFF + .checked_mul(1_u32 << failures.saturating_sub(1).min(5)) + .unwrap_or(MAX_FAILURE_BACKOFF) + .min(MAX_FAILURE_BACKOFF); + for _ in 0..100 { + let backoff = failure_backoff(failures); + assert!(backoff >= ceiling / 2); + assert!(backoff <= ceiling); + } + } + } + + #[test] + fn longest_prefix_match_ignores_freshness() { + let far = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(3600)); + let entries = vec![ + cached_s3("s3://bucket", "wide", far), + cached_s3("s3://bucket/warehouse/db", "narrow", far), + ]; + let got = longest_prefix_match(&entries, "s3://bucket/warehouse/db/t/f").unwrap(); + assert_eq!(s3_access_key_id(&got.credential), "narrow"); + assert!(longest_prefix_match(&entries, "s3://other/x").is_none()); + + let fresh = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)); + let stale = Some(SystemTime::now() - Duration::from_secs(60)); + let entries = vec![ + cached_s3("s3://bucket", "wide", fresh), + cached_s3("s3://bucket/table", "narrow-stale", stale), + ]; + let selected = longest_prefix_match(&entries, "s3://bucket/table/f").unwrap(); + assert_eq!(s3_access_key_id(&selected.credential), "narrow-stale"); + assert!(!selected.is_fresh(SystemTime::now())); + } + + #[test] + fn provider_support_tracks_configured_cloud_endpoints() { + let config = RestCatalogConfig::builder() + .uri("http://cat".to_string()) + .build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + let props = HashMap::from([( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/creds".to_string(), + )]); + + // The provider is configured independently of the table metadata scheme, + // but advertises support only for clouds whose endpoint is present. + let provider = build_vended_credential_provider(client.clone(), "http://cat", &props) + .unwrap() + .unwrap(); + assert!(provider.supports_path("s3://b/k")); + assert!(!provider.supports_path("abfss://fs@acct.dfs.core.windows.net/k")); + // No endpoint advertised. + assert!( + build_vended_credential_provider(client.clone(), "http://cat", &HashMap::new()) + .unwrap() + .is_none() + ); + // Explicitly disabled. + let disabled = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/creds".to_string(), + ), + ( + CloudRefresh::AWS.enabled_key.to_string(), + "false".to_string(), + ), + ]); + assert!( + build_vended_credential_provider(client, "http://cat", &disabled) + .unwrap() + .is_none() + ); + // Java's `Strings.isNullOrEmpty` check treats an empty endpoint as absent. + let empty = HashMap::from([(CloudRefresh::AWS.endpoint_key.to_string(), String::new())]); + let config = RestCatalogConfig::builder() + .uri("http://cat".to_string()) + .build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + assert!( + build_vended_credential_provider(client, "http://cat", &empty) + .unwrap() + .is_none() + ); + } + + #[tokio::test] + async fn fetches_and_selects_vended_credential() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_query(Matcher::UrlEncoded( + "planId".to_string(), + "scan-plan-1".to_string(), + )) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + // No static creds -> no seed -> first load fetches from the endpoint. + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + REST_CATALOG_PROP_SCAN_PLAN_ID.to_string(), + "scan-plan-1".to_string(), + ), + ]); + + let provider = build_vended_credential_provider(client, &server.url(), &props) + .expect("provider construction should succeed") + .expect("provider should be built"); + let credential = provider + .load_credential("s3://bucket/warehouse/f") + .await + .unwrap(); + + match credential.kind { + StorageCredentialKind::S3(s3) => { + assert_eq!(s3.access_key_id, "AK"); + assert_eq!(s3.session_token.as_deref(), Some("TOK")); + } + other => panic!("expected S3, got {other:?}"), + } + mock.assert_async().await; + } + + #[tokio::test] + async fn fresh_seed_is_served_without_fetching() { + let mut server = Server::new_async().await; + // Any call to the server is a failure: the fresh seed must be reused. + let mock = server + .mock("GET", Matcher::Any) + .expect(0) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + (S3_ACCESS_KEY_ID.to_string(), "SEED_AK".to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), + (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), + (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), expires), + ]); + + let provider = build_vended_credential_provider(client, &server.url(), &props) + .expect("provider construction should succeed") + .expect("provider should be built"); + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + + assert_eq!(s3_access_key_id(&credential), "SEED_AK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn failed_refresh_is_backed_off_while_credential_is_unexpired() { + let mut server = Server::new_async().await; + // A refresh is due (the seed is within the buffer) but the catalog errors. + let mock = server + .mock("GET", Matcher::Any) + .expect(1) + .with_status(500) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + // Seeded credential is within the refresh buffer but not yet expired. + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(60)); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + (S3_ACCESS_KEY_ID.to_string(), "SEED_AK".to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), + (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), + (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), expires), + ]); + + let provider = build_vended_credential_provider(client, &server.url(), &props) + .expect("provider construction should succeed") + .expect("provider should be built"); + // The first refresh fails, but the still-valid seed is served. Immediate + // follow-up operations stay inside the first jittered backoff window. + for _ in 0..3 { + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(s3_access_key_id(&credential), "SEED_AK"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn concurrent_prefetch_serves_unexpired_credential_without_waiting() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"NEW_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + ); + let (started_tx, started_rx) = mpsc::channel(); + let release = Arc::new(Barrier::new(2)); + let callback_release = Arc::clone(&release); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_chunked_body(move |writer| { + started_tx.send(()).unwrap(); + callback_release.wait(); + writer.write_all(body.as_bytes()) + }) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + let seed_expires = epoch_millis(SystemTime::now() + Duration::from_secs(60)); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + (S3_ACCESS_KEY_ID.to_string(), "SEED_AK".to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), + (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), + (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), seed_expires), + ]); + let provider = build_vended_credential_provider(client, &server.url(), &props) + .unwrap() + .unwrap(); + + let first_provider = Arc::clone(&provider); + let first = + tokio::spawn(async move { first_provider.load_credential("s3://bucket/x/f").await }); + tokio::task::spawn_blocking(move || { + started_rx.recv_timeout(Duration::from_secs(5)).unwrap() + }) + .await + .unwrap(); + + let second_provider = Arc::clone(&provider); + let second = + tokio::spawn(async move { second_provider.load_credential("s3://bucket/x/f").await }); + for _ in 0..100 { + if second.is_finished() { + break; + } + tokio::task::yield_now().await; + } + let completed_without_waiting = second.is_finished(); + release.wait(); + + assert!( + completed_without_waiting, + "a concurrent prefetch waited instead of using the unexpired credential" + ); + assert_eq!(s3_access_key_id(&second.await.unwrap().unwrap()), "SEED_AK"); + assert_eq!(s3_access_key_id(&first.await.unwrap().unwrap()), "NEW_AK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn refresh_uses_table_scoped_token() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer table-token") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ("token".to_string(), "table-token".to_string()), + ]); + let provider = build_vended_credential_provider(client, &server.url(), &props) + .unwrap() + .unwrap(); + + provider.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn one_provider_refreshes_multiple_clouds() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let aws_body = format!( + r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + ); + let gcp_body = format!( + r#"{{"storage-credentials":[{{"prefix":"gs://bucket","config":{{"gcs.oauth2.token":"GCS","gcs.oauth2.token-expires-at":"{expires}"}}}}]}}"# + ); + let aws_mock = server + .mock("GET", "/v1/aws-credentials") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(aws_body) + .create_async() + .await; + let gcp_mock = server + .mock("GET", "/v1/gcp-credentials") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(gcp_body) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/aws-credentials".to_string(), + ), + ( + CloudRefresh::GCP.endpoint_key.to_string(), + "/v1/gcp-credentials".to_string(), + ), + ]); + let provider = build_vended_credential_provider(client, &server.url(), &props) + .unwrap() + .unwrap(); + + assert!(provider.supports_path("s3://bucket/x")); + assert!(provider.supports_path("gs://bucket/x")); + assert_eq!( + s3_access_key_id(&provider.load_credential("s3://bucket/x").await.unwrap()), + "AK" + ); + match provider + .load_credential("gs://bucket/x") + .await + .unwrap() + .kind + { + StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token, "GCS"), + other => panic!("expected GCS credential, got {other:?}"), + } + aws_mock.assert_async().await; + gcp_mock.assert_async().await; + } + + #[tokio::test] + async fn refresh_failure_never_serves_expired_credential() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", Matcher::Any) + .with_status(500) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + let expired = epoch_millis(SystemTime::now() - Duration::from_secs(60)); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + (S3_ACCESS_KEY_ID.to_string(), "EXPIRED_AK".to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "EXPIRED_SK".to_string()), + (S3_SESSION_TOKEN.to_string(), "EXPIRED_TOK".to_string()), + (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), expired), + ]); + + let provider = build_vended_credential_provider(client, &server.url(), &props) + .expect("provider construction should succeed") + .expect("provider should be built"); + + assert!(provider.load_credential("s3://bucket/x/f").await.is_err()); + mock.assert_async().await; + } + + #[tokio::test] + async fn sub_buffer_ttl_is_refetched_on_each_sequential_operation() { + let mut server = Server::new_async().await; + // The vended TTL (60s) is shorter than REFRESH_BUFFER, so every operation + // is eligible for refresh. This matches AWS CachedSupplier: successful + // results do not receive failure backoff. + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(60)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(3) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let config = RestCatalogConfig::builder().uri(server.url()).build(); + let client = Arc::new(HttpClient::new(&config).unwrap()); + // No static creds -> no seed -> each sequential load fetches because the + // returned credential is already inside its prefetch window. + let props = HashMap::from([( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + )]); + + let provider = build_vended_credential_provider(client, &server.url(), &props) + .expect("provider construction should succeed") + .expect("provider should be built"); + + for _ in 0..3 { + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(s3_access_key_id(&credential), "AK"); + } + mock.assert_async().await; + } +} diff --git a/crates/catalog/rest/src/lib.rs b/crates/catalog/rest/src/lib.rs index 383728401f..a3f13742b8 100644 --- a/crates/catalog/rest/src/lib.rs +++ b/crates/catalog/rest/src/lib.rs @@ -53,6 +53,7 @@ mod catalog; mod client; +mod credential; mod endpoint; mod types; diff --git a/crates/catalog/rest/src/types.rs b/crates/catalog/rest/src/types.rs index 390521229e..8265ceb40b 100644 --- a/crates/catalog/rest/src/types.rs +++ b/crates/catalog/rest/src/types.rs @@ -101,7 +101,7 @@ impl From for Error { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Serialize, Deserialize)] pub(super) struct TokenResponse { pub(super) access_token: String, pub(super) token_type: String, @@ -205,7 +205,7 @@ pub struct RenameTableRequest { pub destination: TableIdent, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[derive(Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "kebab-case")] /// Result returned when a table is successfully loaded or created. /// @@ -231,7 +231,18 @@ pub struct LoadTableResult { pub storage_credentials: Option>, } -#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +impl std::fmt::Debug for LoadTableResult { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LoadTableResult") + .field("metadata_location", &self.metadata_location) + .field("metadata", &self.metadata) + .field("config_keys", &self.config.keys().collect::>()) + .field("storage_credentials", &self.storage_credentials) + .finish_non_exhaustive() + } +} + +#[derive(Clone, Serialize, Deserialize, PartialEq, Eq)] /// Storage credential for a specific location prefix. /// /// Indicates a storage location prefix where the credential is relevant. Clients should @@ -244,6 +255,27 @@ pub struct StorageCredential { pub config: HashMap, } +impl std::fmt::Debug for StorageCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("StorageCredential") + .field("prefix", &self.prefix) + .field("config_keys", &self.config.keys().collect::>()) + .finish_non_exhaustive() + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] +#[serde(rename_all = "kebab-case")] +/// Response from the table credentials endpoint +/// (`GET /v1/{prefix}/namespaces/{namespace}/tables/{table}/credentials`). +/// +/// Returns freshly vended storage credentials so clients can refresh temporary +/// credentials before they expire. +pub struct LoadCredentialsResponse { + /// Storage credentials, one entry per location prefix. + pub storage_credentials: Vec, +} + #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)] #[serde(rename_all = "kebab-case")] /// Request to create a new table in a namespace. diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index 610675c83d..2771ffd998 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -689,6 +689,13 @@ pub fn iceberg::inspect::SnapshotsTable<'a>::new(table: &'a iceberg::table::Tabl pub async fn iceberg::inspect::SnapshotsTable<'a>::scan(&self) -> iceberg::Result pub fn iceberg::inspect::SnapshotsTable<'a>::schema(&self) -> iceberg::spec::Schema pub mod iceberg::io +pub enum iceberg::io::StorageCredentialKind +pub iceberg::io::StorageCredentialKind::Gcs(iceberg::io::GcsCredential) +pub iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential) +impl core::clone::Clone for iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredentialKind::clone(&self) -> iceberg::io::StorageCredentialKind +impl core::fmt::Debug for iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredentialKind::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::AzdlsConfig pub iceberg::io::AzdlsConfig::account_key: core::option::Option pub iceberg::io::AzdlsConfig::account_name: core::option::Option @@ -739,6 +746,7 @@ impl iceberg::io::FileIOBuilder pub fn iceberg::io::FileIOBuilder::build(self) -> iceberg::io::FileIO pub fn iceberg::io::FileIOBuilder::config(&self) -> &iceberg::io::StorageConfig pub fn iceberg::io::FileIOBuilder::new(factory: alloc::sync::Arc) -> Self +pub fn iceberg::io::FileIOBuilder::with_credential_provider(self, provider: alloc::sync::Arc) -> Self pub fn iceberg::io::FileIOBuilder::with_prop(self, key: impl alloc::string::ToString, value: impl alloc::string::ToString) -> Self pub fn iceberg::io::FileIOBuilder::with_props(self, args: impl core::iter::traits::collect::IntoIterator) -> Self impl core::clone::Clone for iceberg::io::FileIOBuilder @@ -775,6 +783,12 @@ impl serde_core::ser::Serialize for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::GcsCredential +pub iceberg::io::GcsCredential::token: alloc::string::String +impl core::clone::Clone for iceberg::io::GcsCredential +pub fn iceberg::io::GcsCredential::clone(&self) -> iceberg::io::GcsCredential +impl core::fmt::Debug for iceberg::io::GcsCredential +pub fn iceberg::io::GcsCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::HfConfig pub iceberg::io::HfConfig::endpoint: core::option::Option pub iceberg::io::HfConfig::revision: core::option::Option @@ -842,6 +856,7 @@ impl core::fmt::Debug for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::StorageFactory for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl serde_core::ser::Serialize for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::LocalFsStorageFactory @@ -880,6 +895,7 @@ impl core::fmt::Debug for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl serde_core::ser::Serialize for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::MemoryStorageFactory @@ -955,6 +971,14 @@ impl serde_core::ser::Serialize for iceberg::io::S3Config pub fn iceberg::io::S3Config::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::S3Config pub fn iceberg::io::S3Config::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::S3Credential +pub iceberg::io::S3Credential::access_key_id: alloc::string::String +pub iceberg::io::S3Credential::secret_access_key: alloc::string::String +pub iceberg::io::S3Credential::session_token: core::option::Option +impl core::clone::Clone for iceberg::io::S3Credential +pub fn iceberg::io::S3Credential::clone(&self) -> iceberg::io::S3Credential +impl core::fmt::Debug for iceberg::io::S3Credential +pub fn iceberg::io::S3Credential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::StorageConfig impl iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::from_props(props: std::collections::hash::map::HashMap) -> Self @@ -992,6 +1016,13 @@ impl serde_core::ser::Serialize for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::StorageCredential +pub iceberg::io::StorageCredential::expires_at: core::option::Option +pub iceberg::io::StorageCredential::kind: iceberg::io::StorageCredentialKind +impl core::clone::Clone for iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::clone(&self) -> iceberg::io::StorageCredential +impl core::fmt::Debug for iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub const iceberg::io::ADLS_ACCOUNT_KEY: &str pub const iceberg::io::ADLS_ACCOUNT_NAME: &str pub const iceberg::io::ADLS_AUTHORITY_HOST: &str @@ -1000,6 +1031,8 @@ pub const iceberg::io::ADLS_CLIENT_SECRET: &str pub const iceberg::io::ADLS_CONNECTION_STRING: &str pub const iceberg::io::ADLS_SAS_TOKEN: &str pub const iceberg::io::ADLS_TENANT_ID: &str +pub const iceberg::io::AWS_REFRESH_CREDENTIALS_ENABLED: &str +pub const iceberg::io::AWS_REFRESH_CREDENTIALS_ENDPOINT: &str pub const iceberg::io::CLIENT_REGION: &str pub const iceberg::io::GCS_ALLOW_ANONYMOUS: &str pub const iceberg::io::GCS_CREDENTIALS_JSON: &str @@ -1007,8 +1040,11 @@ pub const iceberg::io::GCS_DISABLE_CONFIG_LOAD: &str pub const iceberg::io::GCS_DISABLE_VM_METADATA: &str pub const iceberg::io::GCS_NO_AUTH: &str pub const iceberg::io::GCS_PROJECT_ID: &str +pub const iceberg::io::GCS_REFRESH_CREDENTIALS_ENABLED: &str +pub const iceberg::io::GCS_REFRESH_CREDENTIALS_ENDPOINT: &str pub const iceberg::io::GCS_SERVICE_PATH: &str pub const iceberg::io::GCS_TOKEN: &str +pub const iceberg::io::GCS_TOKEN_EXPIRES_AT: &str pub const iceberg::io::GCS_USER_PROJECT: &str pub const iceberg::io::HF_ENDPOINT: &str pub const iceberg::io::HF_REVISION: &str @@ -1028,6 +1064,7 @@ pub const iceberg::io::S3_PATH_STYLE_ACCESS: &str pub const iceberg::io::S3_REGION: &str pub const iceberg::io::S3_SECRET_ACCESS_KEY: &str pub const iceberg::io::S3_SESSION_TOKEN: &str +pub const iceberg::io::S3_SESSION_TOKEN_EXPIRES_AT_MS: &str pub const iceberg::io::S3_SSE_KEY: &str pub const iceberg::io::S3_SSE_MD5: &str pub const iceberg::io::S3_SSE_TYPE: &str @@ -1079,12 +1116,18 @@ pub fn iceberg::io::MemoryStorage::read<'life0, 'life1, 'async_trait>(&'life0 se pub fn iceberg::io::MemoryStorage::reader<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub fn iceberg::io::MemoryStorage::write<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str, bs: bytes::bytes::Bytes) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub fn iceberg::io::MemoryStorage::writer<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait +pub trait iceberg::io::StorageCredentialProvider: core::fmt::Debug + core::marker::Send + core::marker::Sync +pub fn iceberg::io::StorageCredentialProvider::load_credential<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait +pub fn iceberg::io::StorageCredentialProvider::supports_path(&self, _path: &str) -> bool pub trait iceberg::io::StorageFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize pub fn iceberg::io::StorageFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::StorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl iceberg::io::StorageFactory for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> +pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> pub mod iceberg::memory pub struct iceberg::memory::MemoryCatalog impl core::fmt::Debug for iceberg::memory::MemoryCatalog diff --git a/crates/iceberg/src/io/file_io.rs b/crates/iceberg/src/io/file_io.rs index cd0a4434c4..90d9e48785 100644 --- a/crates/iceberg/src/io/file_io.rs +++ b/crates/iceberg/src/io/file_io.rs @@ -22,7 +22,8 @@ use bytes::Bytes; use futures::{Stream, StreamExt}; use super::storage::{ - LocalFsStorageFactory, MemoryStorageFactory, Storage, StorageConfig, StorageFactory, + LocalFsStorageFactory, MemoryStorageFactory, Storage, StorageConfig, StorageCredentialProvider, + StorageFactory, }; use crate::Result; @@ -65,6 +66,8 @@ pub struct FileIO { config: StorageConfig, /// Factory for creating storage instances factory: Arc, + /// Optional provider of refreshable, backend-specific credentials + credential_provider: Option>, /// Cached storage instance (lazily initialized) storage: Arc>>, } @@ -77,6 +80,7 @@ impl FileIO { Self { config: StorageConfig::new(), factory: Arc::new(MemoryStorageFactory), + credential_provider: None, storage: Arc::new(OnceLock::new()), } } @@ -88,6 +92,7 @@ impl FileIO { Self { config: StorageConfig::new(), factory: Arc::new(LocalFsStorageFactory), + credential_provider: None, storage: Arc::new(OnceLock::new()), } } @@ -107,8 +112,11 @@ impl FileIO { return Ok(storage.clone()); } - // Build the storage - let storage = self.factory.build(&self.config)?; + // Build the storage, passing any credential provider so backends that + // support refreshable credentials can wire it into their operators. + let storage = self + .factory + .build_with_credentials(&self.config, self.credential_provider.clone())?; // Try to set it (another thread might have set it first) let _ = self.storage.set(storage.clone()); @@ -191,6 +199,8 @@ pub struct FileIOBuilder { factory: Arc, /// Storage configuration config: StorageConfig, + /// Optional provider of refreshable, backend-specific credentials + credential_provider: Option>, } impl FileIOBuilder { @@ -199,6 +209,7 @@ impl FileIOBuilder { Self { factory, config: StorageConfig::new(), + credential_provider: None, } } @@ -224,11 +235,21 @@ impl FileIOBuilder { &self.config } + /// Attach a provider of refreshable, backend-specific credentials. + pub fn with_credential_provider( + mut self, + provider: Arc, + ) -> Self { + self.credential_provider = Some(provider); + self + } + /// Builds [`FileIO`]. pub fn build(self) -> FileIO { FileIO { config: self.config, factory: self.factory, + credential_provider: self.credential_provider, storage: Arc::new(OnceLock::new()), } } diff --git a/crates/iceberg/src/io/storage/config/gcs.rs b/crates/iceberg/src/io/storage/config/gcs.rs index 5b11567f4d..1392e35fd1 100644 --- a/crates/iceberg/src/io/storage/config/gcs.rs +++ b/crates/iceberg/src/io/storage/config/gcs.rs @@ -41,6 +41,12 @@ pub const GCS_NO_AUTH: &str = "gcs.no-auth"; pub const GCS_CREDENTIALS_JSON: &str = "gcs.credentials-json"; /// Google Cloud Storage token. pub const GCS_TOKEN: &str = "gcs.oauth2.token"; +/// Epoch-millisecond timestamp at which the vended GCS OAuth2 token expires. +pub const GCS_TOKEN_EXPIRES_AT: &str = "gcs.oauth2.token-expires-at"; +/// Endpoint used to fetch and refresh vended GCS OAuth2 credentials. +pub const GCS_REFRESH_CREDENTIALS_ENDPOINT: &str = "gcs.oauth2.refresh-credentials-endpoint"; +/// Whether vended GCS OAuth2 credentials should be refreshed. Defaults to `true`. +pub const GCS_REFRESH_CREDENTIALS_ENABLED: &str = "gcs.oauth2.refresh-credentials-enabled"; /// Option to skip signing requests (e.g. for public buckets/folders). pub const GCS_ALLOW_ANONYMOUS: &str = "gcs.allow-anonymous"; /// Option to skip loading the credential from GCE metadata server. diff --git a/crates/iceberg/src/io/storage/config/mod.rs b/crates/iceberg/src/io/storage/config/mod.rs index 2350aab6dd..a79f871e8d 100644 --- a/crates/iceberg/src/io/storage/config/mod.rs +++ b/crates/iceberg/src/io/storage/config/mod.rs @@ -51,12 +51,22 @@ use serde::{Deserialize, Serialize}; /// which storage backend to use. The storage type is determined by the /// explicit factory selection. /// ``` -#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)] +#[derive(Clone, PartialEq, Eq, Serialize, Deserialize, Default)] pub struct StorageConfig { /// Configuration properties for the storage backend props: HashMap, } +impl std::fmt::Debug for StorageConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // Property values may hold vended credentials (e.g. `s3.secret-access-key`, + // `s3.session-token`) + f.debug_struct("StorageConfig") + .field("keys", &self.props.keys().collect::>()) + .finish_non_exhaustive() + } +} + impl StorageConfig { /// Create a new empty StorageConfig. pub fn new() -> Self { diff --git a/crates/iceberg/src/io/storage/config/s3.rs b/crates/iceberg/src/io/storage/config/s3.rs index 664e8637b8..eb8c9f54c5 100644 --- a/crates/iceberg/src/io/storage/config/s3.rs +++ b/crates/iceberg/src/io/storage/config/s3.rs @@ -35,10 +35,16 @@ pub const S3_ACCESS_KEY_ID: &str = "s3.access-key-id"; pub const S3_SECRET_ACCESS_KEY: &str = "s3.secret-access-key"; /// S3 session token (required when using temporary credentials). pub const S3_SESSION_TOKEN: &str = "s3.session-token"; +/// Epoch-millisecond timestamp at which the vended S3 session token expires. +pub const S3_SESSION_TOKEN_EXPIRES_AT_MS: &str = "s3.session-token-expires-at-ms"; /// S3 region. pub const S3_REGION: &str = "s3.region"; /// Region to use for the S3 client (takes precedence over [`S3_REGION`]). pub const CLIENT_REGION: &str = "client.region"; +/// Endpoint used to fetch and refresh vended AWS credentials. +pub const AWS_REFRESH_CREDENTIALS_ENDPOINT: &str = "client.refresh-credentials-endpoint"; +/// Whether vended AWS credentials should be refreshed. Defaults to `true`. +pub const AWS_REFRESH_CREDENTIALS_ENABLED: &str = "client.refresh-credentials-enabled"; /// S3 Path Style Access. pub const S3_PATH_STYLE_ACCESS: &str = "s3.path-style-access"; /// S3 Server Side Encryption Type. diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index 5276c7771f..b10bd380c0 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -23,6 +23,7 @@ mod memory; use std::fmt::Debug; use std::sync::Arc; +use std::time::SystemTime; use async_trait::async_trait; use bytes::Bytes; @@ -32,7 +33,7 @@ pub use local_fs::{LocalFsStorage, LocalFsStorageFactory}; pub use memory::{MemoryStorage, MemoryStorageFactory}; use super::{FileMetadata, FileRead, FileWrite, InputFile, OutputFile}; -use crate::Result; +use crate::{Error, ErrorKind, Result}; /// Trait for storage operations in Iceberg. /// @@ -139,4 +140,103 @@ pub trait StorageFactory: Debug + Send + Sync { /// A `Result` containing an `Arc` on success, or an error /// if the storage could not be created. fn build(&self, config: &StorageConfig) -> Result>; + + /// Build a new Storage instance, optionally supplying a credential provider + /// that the backend can call to obtain and refresh short-lived credentials. + fn build_with_credentials( + &self, + config: &StorageConfig, + credential_provider: Option>, + ) -> Result> { + if credential_provider.is_some() { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + "Storage factory does not support refreshable credential providers", + )); + } + + self.build(config) + } +} + +/// Supplies fresh, backend-specific storage credentials on demand. +/// +/// A catalog that vends temporary credentials implements this trait so that +/// storage backends can re-fetch credentials as they approach expiry instead +/// of failing once the initial token's TTL runs out. +/// +/// # Caching +/// +/// [`load_credential`](Self::load_credential) may be called very frequently — +/// the S3 backend, for example, rebuilds its operator (and therefore its +/// signer) on every file operation. Implementations must cache internally and +/// only re-fetch when the current credential is at or near expiry; otherwise +/// every object-store request would trigger a call back to the catalog. +#[async_trait] +pub trait StorageCredentialProvider: Debug + Send + Sync { + /// Return whether this provider has refresh configuration for `path`. + /// + /// Backends use this before replacing their normal credential chain. The + /// default is `true` for single-backend providers; multi-backend providers + /// should return `false` for schemes they do not configure. + fn supports_path(&self, _path: &str) -> bool { + true + } + + /// Load a fresh credential for the storage location identified by `path`. + /// + /// `path` is the absolute location being accessed (e.g. + /// `s3://bucket/warehouse/db/table/...`). Providers that vend distinct + /// credentials per location prefix use it to select the most specific + /// match. + async fn load_credential(&self, path: &str) -> Result; +} + +/// A vended storage credential together with when it expires. +#[derive(Clone, Debug)] +pub struct StorageCredential { + /// The backend-specific credential material. + pub kind: StorageCredentialKind, + /// When the credential expires, if known. `None` means non-expiring and + /// backends treat such a credential as always valid and never refresh it. + pub expires_at: Option, +} + +/// Backend-specific credential material. +#[derive(Clone, Debug)] +pub enum StorageCredentialKind { + /// Amazon S3 credentials. + S3(S3Credential), + /// Google Cloud Storage credentials. + Gcs(GcsCredential), +} + +/// Temporary Amazon S3 credentials. +#[derive(Clone)] +pub struct S3Credential { + /// AWS access key ID. + pub access_key_id: String, + /// AWS secret access key. + pub secret_access_key: String, + /// AWS session token, set for temporary (STS/vended) credentials. + pub session_token: Option, +} + +impl Debug for S3Credential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("S3Credential").finish_non_exhaustive() + } +} + +/// Temporary Google Cloud Storage credentials (an OAuth2 access token). +#[derive(Clone)] +pub struct GcsCredential { + /// OAuth2 bearer token used to access GCS. + pub token: String, +} + +impl Debug for GcsCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("GcsCredential").finish_non_exhaustive() + } } diff --git a/crates/storage/opendal/Cargo.toml b/crates/storage/opendal/Cargo.toml index e43e7845b3..4cbf2892d9 100644 --- a/crates/storage/opendal/Cargo.toml +++ b/crates/storage/opendal/Cargo.toml @@ -41,7 +41,7 @@ opendal-all = [ opendal-azdls = ["opendal/services-azdls"] opendal-fs = ["opendal/services-fs"] -opendal-gcs = ["opendal/services-gcs"] +opendal-gcs = ["opendal/services-gcs", "reqsign-google", "reqsign-core"] opendal-hf = ["opendal/services-hf"] opendal-memory = ["opendal/services-memory"] opendal-oss = ["opendal/services-oss"] @@ -57,6 +57,7 @@ iceberg = { workspace = true } opendal = { workspace = true } reqsign-aws-v4 = { version = "3.0.0", optional = true } reqsign-core = { version = "3.0.0", optional = true } +reqsign-google = { version = "3.0.0", optional = true } serde = { workspace = true } typetag = { workspace = true } url = { workspace = true } diff --git a/crates/storage/opendal/public-api.txt b/crates/storage/opendal/public-api.txt index fc1ed7cf78..f420a891f1 100644 --- a/crates/storage/opendal/public-api.txt +++ b/crates/storage/opendal/public-api.txt @@ -6,6 +6,7 @@ pub iceberg_storage_opendal::OpenDalStorage::Azdls pub iceberg_storage_opendal::OpenDalStorage::Azdls::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Gcs pub iceberg_storage_opendal::OpenDalStorage::Gcs::config: alloc::sync::Arc +pub iceberg_storage_opendal::OpenDalStorage::Gcs::credential_provider: core::option::Option> pub iceberg_storage_opendal::OpenDalStorage::Hf pub iceberg_storage_opendal::OpenDalStorage::Hf::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::LocalFs @@ -14,6 +15,7 @@ pub iceberg_storage_opendal::OpenDalStorage::Oss pub iceberg_storage_opendal::OpenDalStorage::Oss::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::S3 pub iceberg_storage_opendal::OpenDalStorage::S3::config: alloc::sync::Arc +pub iceberg_storage_opendal::OpenDalStorage::S3::credential_provider: core::option::Option> pub iceberg_storage_opendal::OpenDalStorage::S3::customized_credential_load: core::option::Option impl core::clone::Clone for iceberg_storage_opendal::OpenDalStorage pub fn iceberg_storage_opendal::OpenDalStorage::clone(&self) -> iceberg_storage_opendal::OpenDalStorage @@ -50,6 +52,7 @@ impl core::fmt::Debug for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::storage::StorageFactory for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::build(&self, config: &iceberg::io::storage::config::StorageConfig) -> iceberg::error::Result> +pub fn iceberg_storage_opendal::OpenDalStorageFactory::build_with_credentials(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalStorageFactory @@ -92,6 +95,7 @@ impl core::fmt::Debug for iceberg_storage_opendal::OpenDalResolvingStorageFactor pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::storage::StorageFactory for iceberg_storage_opendal::OpenDalResolvingStorageFactory pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build(&self, config: &iceberg::io::storage::config::StorageConfig) -> iceberg::error::Result> +pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build_with_credentials(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalResolvingStorageFactory pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalResolvingStorageFactory diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index 47ca52ccc6..3bbaba2c5b 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -17,17 +17,20 @@ //! Google Cloud Storage properties use std::collections::HashMap; +use std::sync::Arc; use iceberg::io::{ GCS_ALLOW_ANONYMOUS, GCS_CREDENTIALS_JSON, GCS_DISABLE_CONFIG_LOAD, GCS_DISABLE_VM_METADATA, - GCS_NO_AUTH, GCS_SERVICE_PATH, GCS_TOKEN, + GCS_NO_AUTH, GCS_SERVICE_PATH, GCS_TOKEN, StorageCredentialKind, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; -use opendal::Operator; use opendal::services::GcsConfig; +use opendal::{Configurator, Operator}; +use reqsign_core::{Context, Error as ReqsignError, ProvideCredential, Result as ReqsignResult}; +use reqsign_google::{Credential as GoogleCredential, Token as GoogleToken}; use url::Url; -use crate::utils::{from_opendal_error, is_truthy}; +use crate::utils::{from_opendal_error, is_truthy, system_time_to_timestamp}; /// Parse iceberg properties to [`GcsConfig`]. pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result { @@ -45,7 +48,9 @@ pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result) -> Result Result { +pub(crate) fn gcs_config_build( + cfg: &GcsConfig, + credential_provider: &Option>, + path: &str, +) -> Result { let url = Url::parse(path)?; + if !matches!(url.scheme(), "gs" | "gcs") { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("Invalid gcs url: {path}, expected gs:// or gcs://"), + )); + } let bucket = url.host_str().ok_or_else(|| { Error::new( ErrorKind::DataInvalid, @@ -82,7 +97,95 @@ pub(crate) fn gcs_config_build(cfg: &GcsConfig, path: &str) -> Result let mut cfg = cfg.clone(); cfg.bucket = bucket.to_string(); - Ok(Operator::from_config(cfg) - .map_err(from_opendal_error)? - .finish()) + + // When a catalog-supplied provider is present, make it the sole credential + // source. `reqsign_google` only prepends a custom provider (unlike S3, which + // replaces the chain) and its chain continues to the next provider on error, + // so without this a failed refresh would silently fall back to the stale seed + // token or ambient GCP credentials. + let credential_provider = credential_provider + .as_ref() + .filter(|provider| provider.supports_path(path)); + if credential_provider.is_some() { + if cfg.skip_signature { + return Err(Error::new( + ErrorKind::DataInvalid, + "Invalid GCS auth settings: anonymous access cannot be combined with refreshable credentials", + )); + } + cfg.token = None; + cfg.credential = None; + cfg.disable_vm_metadata = true; + cfg.disable_config_load = true; + } + + let mut builder = cfg.into_builder(); + + // A catalog-supplied provider re-fetches the vended OAuth2 token as it nears expiry + if let Some(provider) = credential_provider { + builder = builder.credential_provider(VendedGcsCredentialProvider::new( + Arc::clone(provider), + path.to_string(), + )); + } + + Ok(Operator::new(builder).map_err(from_opendal_error)?.finish()) +} + +/// Adapts a generic [`StorageCredentialProvider`] into a `reqsign` +/// [`ProvideCredential`], so the GCS signer can obtain and refresh vended OAuth2 +/// tokens. +struct VendedGcsCredentialProvider { + provider: Arc, + /// Absolute path this operator serves: handed back to the provider so it can + /// select the vended credential whose prefix best matches the location. + path: String, +} + +impl std::fmt::Debug for VendedGcsCredentialProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("VendedGcsCredentialProvider") + .field("path", &self.path) + .finish_non_exhaustive() + } +} + +impl VendedGcsCredentialProvider { + fn new(provider: Arc, path: String) -> Self { + Self { provider, path } + } +} + +impl ProvideCredential for VendedGcsCredentialProvider { + type Credential = GoogleCredential; + + async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { + let credential = self + .provider + .load_credential(&self.path) + .await + .map_err(|e| { + ReqsignError::unexpected(format!( + "failed to load vended GCS credential for {}", + self.path + )) + .with_source(e) + })?; + + let expires_at = credential + .expires_at + .map(system_time_to_timestamp) + .transpose()?; + match credential.kind { + StorageCredentialKind::Gcs(gcs) => { + Ok(Some(GoogleCredential::with_token(GoogleToken { + access_token: gcs.token, + expires_at, + }))) + } + _ => Err(ReqsignError::unexpected( + "GCS storage received a non-GCS credential from the provider", + )), + } + } } diff --git a/crates/storage/opendal/src/lib.rs b/crates/storage/opendal/src/lib.rs index 768d3ff77f..0f301dfc91 100644 --- a/crates/storage/opendal/src/lib.rs +++ b/crates/storage/opendal/src/lib.rs @@ -35,7 +35,7 @@ use futures::StreamExt; use futures::stream::BoxStream; use iceberg::io::{ FileMetadata, FileRead, FileWrite, InputFile, OutputFile, Storage, StorageConfig, - StorageFactory, + StorageCredentialProvider, StorageFactory, }; use iceberg::{Error, ErrorKind, Result}; use opendal::Operator; @@ -135,8 +135,16 @@ pub enum OpenDalStorageFactory { #[typetag::serde(name = "OpenDalStorageFactory")] impl StorageFactory for OpenDalStorageFactory { - #[allow(unused_variables)] fn build(&self, config: &StorageConfig) -> Result> { + self.build_with_credentials(config, None) + } + + #[allow(unused_variables)] + fn build_with_credentials( + &self, + config: &StorageConfig, + credential_provider: Option>, + ) -> Result> { match self { #[cfg(feature = "opendal-memory")] OpenDalStorageFactory::Memory => { @@ -150,10 +158,12 @@ impl StorageFactory for OpenDalStorageFactory { } => Ok(Arc::new(OpenDalStorage::S3 { config: s3_config_parse(config.props().clone())?.into(), customized_credential_load: customized_credential_load.clone(), + credential_provider, })), #[cfg(feature = "opendal-gcs")] OpenDalStorageFactory::Gcs => Ok(Arc::new(OpenDalStorage::Gcs { config: gcs_config_parse(config.props().clone())?.into(), + credential_provider, })), #[cfg(feature = "opendal-oss")] OpenDalStorageFactory::Oss => Ok(Arc::new(OpenDalStorage::Oss { @@ -210,12 +220,18 @@ pub enum OpenDalStorage { /// Custom AWS credential loader. #[serde(skip)] customized_credential_load: Option, + /// Provider of refreshable vended credentials, supplied by the catalog. + #[serde(skip)] + credential_provider: Option>, }, /// GCS storage variant. #[cfg(feature = "opendal-gcs")] Gcs { /// GCS configuration. config: Arc, + /// Provider of refreshable vended credentials, supplied by the catalog. + #[serde(skip)] + credential_provider: Option>, }, /// OSS storage variant. #[cfg(feature = "opendal-oss")] @@ -287,8 +303,14 @@ impl OpenDalStorage { OpenDalStorage::S3 { config, customized_credential_load, + credential_provider, } => { - let op = s3_config_build(config, customized_credential_load, path)?; + let op = s3_config_build( + config, + customized_credential_load, + credential_provider, + path, + )?; let op_info = op.info(); // Use the URL scheme in the path for prefix matching. This enables @@ -310,9 +332,18 @@ impl OpenDalStorage { } } #[cfg(feature = "opendal-gcs")] - OpenDalStorage::Gcs { config } => { - let operator = gcs_config_build(config, path)?; - let prefix = format!("gs://{}/", operator.info().name()); + OpenDalStorage::Gcs { + config, + credential_provider, + } => { + let operator = gcs_config_build(config, credential_provider, path)?; + let url = url::Url::parse(path).map_err(|e| { + Error::new( + ErrorKind::DataInvalid, + format!("Invalid gcs url: {path}: {e}"), + ) + })?; + let prefix = format!("{}://{}/", url.scheme(), operator.info().name()); if path.starts_with(&prefix) { (operator, &path[prefix.len()..]) } else { @@ -384,6 +415,27 @@ impl OpenDalStorage { } } + /// Returns whether `path` uses a path-scoped dynamic credential provider. + /// + /// Such paths cannot share an OpenDAL deleter unless the provider can expose + /// the credential scope that applies to each path. Process them independently + /// so bulk deletion remains bounded without crossing credential boundaries. + fn uses_dynamic_credentials(&self, path: &str) -> bool { + match self { + #[cfg(feature = "opendal-s3")] + OpenDalStorage::S3 { + credential_provider: Some(provider), + .. + } => provider.supports_path(path), + #[cfg(feature = "opendal-gcs")] + OpenDalStorage::Gcs { + credential_provider: Some(provider), + .. + } => provider.supports_path(path), + _ => false, + } + } + /// Extracts the relative path from an absolute path without building an operator. /// /// This is a lightweight alternative to [`create_operator`](Self::create_operator) for cases @@ -418,13 +470,19 @@ impl OpenDalStorage { #[cfg(feature = "opendal-gcs")] OpenDalStorage::Gcs { .. } => { let url = url::Url::parse(path)?; + if !matches!(url.scheme(), "gs" | "gcs") { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("Invalid gcs url: {path}, expected gs:// or gcs://"), + )); + } let bucket = url.host_str().ok_or_else(|| { Error::new( ErrorKind::DataInvalid, format!("Invalid gcs url: {path}, missing bucket"), ) })?; - let prefix = format!("gs://{}/", bucket); + let prefix = format!("{}://{}/", url.scheme(), bucket); if path.starts_with(&prefix) { Ok(&path[prefix.len()..]) } else { @@ -553,6 +611,12 @@ impl Storage for OpenDalStorage { let mut deleters: HashMap = HashMap::new(); while let Some(path) = paths.next().await { + if self.uses_dynamic_credentials(&path) { + let (op, relative_path) = self.create_operator(&path)?; + op.delete(relative_path).await.map_err(from_opendal_error)?; + continue; + } + let bucket = self.batch_key_for_path(&path); let (relative_path, deleter) = match deleters.entry(bucket) { @@ -631,6 +695,18 @@ impl FileWrite for OpenDalWriter { mod tests { use super::*; + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[derive(Debug)] + struct AlwaysSupportedCredentialProvider; + + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[async_trait] + impl StorageCredentialProvider for AlwaysSupportedCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + unreachable!("batch key selection must not load a credential") + } + } + #[cfg(feature = "opendal-memory")] #[test] fn test_default_memory_operator() { @@ -677,6 +753,7 @@ mod tests { let storage = OpenDalStorage::S3 { config: Arc::new(S3Config::default()), customized_credential_load: None, + credential_provider: None, }; // All S3-family schemes are accepted by the same storage instance. @@ -692,19 +769,73 @@ mod tests { } } + #[cfg(feature = "opendal-s3")] + #[test] + fn test_dynamic_credentials_use_isolated_deletes() { + let storage = OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: None, + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + let first = "s3://bucket/table-a/file.parquet"; + let second = "s3://bucket/table-b/file.parquet"; + + assert!(storage.uses_dynamic_credentials(first)); + assert!(storage.uses_dynamic_credentials(second)); + assert_eq!(storage.batch_key_for_path(first), "bucket"); + assert_eq!(storage.batch_key_for_path(second), "bucket"); + } + + #[cfg(feature = "opendal-s3")] + #[test] + fn test_s3_rejects_anonymous_dynamic_credentials() { + let mut config = S3Config::default(); + config.skip_signature = true; + let storage = OpenDalStorage::S3 { + config: Arc::new(config), + customized_credential_load: None, + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + + let error = storage + .create_operator(&"s3://bucket/file.parquet") + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::DataInvalid); + } + + #[cfg(feature = "opendal-gcs")] + #[test] + fn test_gcs_rejects_anonymous_dynamic_credentials() { + let mut config = GcsConfig::default(); + config.skip_signature = true; + let storage = OpenDalStorage::Gcs { + config: Arc::new(config), + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + + let error = storage + .create_operator(&"gs://bucket/file.parquet") + .unwrap_err(); + assert_eq!(error.kind(), ErrorKind::DataInvalid); + } + #[cfg(feature = "opendal-gcs")] #[test] fn test_relativize_path_gcs() { let storage = OpenDalStorage::Gcs { config: Arc::new(GcsConfig::default()), + credential_provider: None, }; - assert_eq!( - storage - .relativize_path("gs://my-bucket/path/to/file.parquet") - .unwrap(), - "path/to/file.parquet" - ); + for scheme in ["gs", "gcs"] { + let path = format!("{scheme}://my-bucket/path/to/file.parquet"); + assert_eq!( + storage.relativize_path(&path).unwrap(), + "path/to/file.parquet" + ); + let (_, relative) = storage.create_operator(&path).unwrap(); + assert_eq!(relative, "path/to/file.parquet"); + } } #[cfg(feature = "opendal-gcs")] @@ -712,6 +843,7 @@ mod tests { fn test_relativize_path_gcs_invalid_scheme() { let storage = OpenDalStorage::Gcs { config: Arc::new(GcsConfig::default()), + credential_provider: None, }; assert!( @@ -719,6 +851,11 @@ mod tests { .relativize_path("s3://my-bucket/path/to/file.parquet") .is_err() ); + assert!( + storage + .create_operator(&"s3://my-bucket/path/to/file.parquet") + .is_err() + ); } #[cfg(feature = "opendal-oss")] diff --git a/crates/storage/opendal/src/resolving.rs b/crates/storage/opendal/src/resolving.rs index 86993220a8..0fad716e52 100644 --- a/crates/storage/opendal/src/resolving.rs +++ b/crates/storage/opendal/src/resolving.rs @@ -27,7 +27,7 @@ use futures::StreamExt; use futures::stream::BoxStream; use iceberg::io::{ FileMetadata, FileRead, FileWrite, InputFile, OutputFile, Storage, StorageConfig, - StorageFactory, + StorageCredentialProvider, StorageFactory, }; use iceberg::{Error, ErrorKind, Result}; use serde::{Deserialize, Serialize}; @@ -85,6 +85,9 @@ fn build_storage_for_scheme( scheme: &'static str, props: &HashMap, #[cfg(feature = "opendal-s3")] customized_credential_load: &Option, + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] credential_provider: &Option< + Arc, + >, ) -> Result { match scheme { #[cfg(feature = "opendal-s3")] @@ -93,6 +96,7 @@ fn build_storage_for_scheme( Ok(OpenDalStorage::S3 { config: Arc::new(config), customized_credential_load: customized_credential_load.clone(), + credential_provider: credential_provider.clone(), }) } #[cfg(feature = "opendal-gcs")] @@ -100,6 +104,7 @@ fn build_storage_for_scheme( let config = crate::gcs::gcs_config_parse(props.clone())?; Ok(OpenDalStorage::Gcs { config: Arc::new(config), + credential_provider: credential_provider.clone(), }) } #[cfg(feature = "opendal-oss")] @@ -186,11 +191,22 @@ impl OpenDalResolvingStorageFactory { #[typetag::serde] impl StorageFactory for OpenDalResolvingStorageFactory { fn build(&self, config: &StorageConfig) -> Result> { + self.build_with_credentials(config, None) + } + + #[allow(unused_variables)] + fn build_with_credentials( + &self, + config: &StorageConfig, + credential_provider: Option>, + ) -> Result> { Ok(Arc::new(OpenDalResolvingStorage { props: config.props().clone(), storages: RwLock::new(HashMap::new()), #[cfg(feature = "opendal-s3")] customized_credential_load: self.customized_credential_load.clone(), + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + credential_provider, })) } } @@ -201,7 +217,7 @@ impl StorageFactory for OpenDalResolvingStorageFactory { /// Sub-storages are lazily created on first use for each scheme and cached /// for subsequent operations. Scheme aliases like `s3`/`s3a`/`s3n` map to /// the same canonical scheme, so they share a storage instance. -#[derive(Debug, Serialize, Deserialize)] +#[derive(Serialize, Deserialize)] pub struct OpenDalResolvingStorage { /// Configuration properties shared across all backends. props: HashMap, @@ -212,6 +228,20 @@ pub struct OpenDalResolvingStorage { #[cfg(feature = "opendal-s3")] #[serde(skip)] customized_credential_load: Option, + /// Provider of refreshable vended credentials. + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[serde(skip)] + credential_provider: Option>, +} + +impl std::fmt::Debug for OpenDalResolvingStorage { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + // `props` can contain storage secrets, and a custom credential provider + // may carry secret state in its own Debug implementation + f.debug_struct("OpenDalResolvingStorage") + .field("property_keys", &self.props.keys().collect::>()) + .finish_non_exhaustive() + } } impl OpenDalResolvingStorage { @@ -247,6 +277,8 @@ impl OpenDalResolvingStorage { &self.props, #[cfg(feature = "opendal-s3")] &self.customized_credential_load, + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + &self.credential_provider, )?; let storage = Arc::new(storage); cache.insert(scheme, storage.clone()); @@ -334,6 +366,8 @@ mod tests { storages: RwLock::new(HashMap::new()), #[cfg(feature = "opendal-s3")] customized_credential_load: None, + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + credential_provider: None, } } diff --git a/crates/storage/opendal/src/s3.rs b/crates/storage/opendal/src/s3.rs index 4b3893b39c..8e57cdfdb6 100644 --- a/crates/storage/opendal/src/s3.rs +++ b/crates/storage/opendal/src/s3.rs @@ -22,7 +22,8 @@ use iceberg::io::{ CLIENT_REGION, S3_ACCESS_KEY_ID, S3_ALLOW_ANONYMOUS, S3_ASSUME_ROLE_ARN, S3_ASSUME_ROLE_EXTERNAL_ID, S3_ASSUME_ROLE_SESSION_NAME, S3_DISABLE_CONFIG_LOAD, S3_DISABLE_EC2_METADATA, S3_ENDPOINT, S3_PATH_STYLE_ACCESS, S3_REGION, S3_SECRET_ACCESS_KEY, - S3_SESSION_TOKEN, S3_SSE_KEY, S3_SSE_MD5, S3_SSE_TYPE, + S3_SESSION_TOKEN, S3_SSE_KEY, S3_SSE_MD5, S3_SSE_TYPE, StorageCredentialKind, + StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; use opendal::services::S3Config; @@ -31,10 +32,13 @@ use opendal::{Configurator, Operator}; pub use reqsign_aws_v4::Credential as AwsCredential; /// Trait for types that can asynchronously supply [`AwsCredential`] to a [`CustomAwsCredentialLoader`]. pub use reqsign_core::ProvideCredential; -use reqsign_core::{ProvideCredentialChain, ProvideCredentialDyn}; +use reqsign_core::{ + Context, Error as ReqsignError, ProvideCredentialChain, ProvideCredentialDyn, + Result as ReqsignResult, +}; use url::Url; -use crate::utils::{from_opendal_error, is_truthy}; +use crate::utils::{from_opendal_error, is_truthy, system_time_to_timestamp}; /// Parse iceberg props to s3 config. pub(crate) fn s3_config_parse(mut m: HashMap) -> Result { @@ -129,6 +133,7 @@ pub(crate) fn s3_config_parse(mut m: HashMap) -> Result, + credential_provider: &Option>, path: &str, ) -> Result { let url = Url::parse(path)?; @@ -139,6 +144,19 @@ pub(crate) fn s3_config_build( ) })?; + // Preserve the existing custom-loader precedence: an explicitly configured loader + // is the sole source, otherwise install the catalog provider as a replacement chain so + // refresh failures cannot fall through to broader ambient AWS credentials. + let credential_provider = credential_provider + .as_ref() + .filter(|provider| provider.supports_path(path)); + if customized_credential_load.is_none() && credential_provider.is_some() && cfg.skip_signature { + return Err(Error::new( + ErrorKind::DataInvalid, + "Invalid S3 auth settings: anonymous access cannot be combined with refreshable credentials", + )); + } + let mut builder = cfg .clone() .into_builder() @@ -148,11 +166,75 @@ pub(crate) fn s3_config_build( if let Some(loader) = customized_credential_load { let chain = ProvideCredentialChain::new().push(Arc::clone(&loader.0)); builder = builder.credential_provider_chain(chain); + } else if let Some(provider) = credential_provider { + let chain = ProvideCredentialChain::new().push(VendedS3CredentialProvider::new( + Arc::clone(provider), + path.to_string(), + )); + builder = builder.credential_provider_chain(chain); } Ok(Operator::new(builder).map_err(from_opendal_error)?.finish()) } +/// Adapts a generic [`StorageCredentialProvider`] into a reqsign +/// [`ProvideCredential`], so the S3 signer can obtain and refresh vended +/// credentials. +struct VendedS3CredentialProvider { + provider: Arc, + /// Absolute path this operator serves; handed back to the provider so it can + /// select the vended credential whose prefix best matches the location. + path: String, +} + +impl std::fmt::Debug for VendedS3CredentialProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("VendedS3CredentialProvider") + .field("path", &self.path) + .finish_non_exhaustive() + } +} + +impl VendedS3CredentialProvider { + fn new(provider: Arc, path: String) -> Self { + Self { provider, path } + } +} + +impl ProvideCredential for VendedS3CredentialProvider { + type Credential = AwsCredential; + + async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { + let credential = self + .provider + .load_credential(&self.path) + .await + .map_err(|e| { + ReqsignError::unexpected(format!( + "failed to load vended S3 credential for {}", + self.path + )) + .with_source(e) + })?; + + let expires_in = credential + .expires_at + .map(system_time_to_timestamp) + .transpose()?; + match credential.kind { + StorageCredentialKind::S3(s3) => Ok(Some(AwsCredential { + access_key_id: s3.access_key_id, + secret_access_key: s3.secret_access_key, + session_token: s3.session_token, + expires_in, + })), + _ => Err(ReqsignError::unexpected( + "S3 storage received a non-S3 credential from the provider", + )), + } + } +} + /// Custom AWS credential loader. /// /// Wraps any [`ProvideCredential`] implementation for use with the S3 storage backend. @@ -176,7 +258,7 @@ impl std::fmt::Debug for CustomAwsCredentialLoader { impl CustomAwsCredentialLoader { /// Create a new custom AWS credential loader from any [`ProvideCredential`] implementation. pub fn new(provider: impl ProvideCredential + 'static) -> Self { - Self(Arc::new(provider) as Arc>) + Self(Arc::new(provider)) } } diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index 56f8c18059..3afd319f22 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -27,3 +27,25 @@ pub(crate) fn from_opendal_error(e: opendal::Error) -> iceberg::Error { ) .with_source(e) } + +/// Convert a [`SystemTime`](std::time::SystemTime) credential expiry into the +/// `reqsign` [`Timestamp`](reqsign_core::time::Timestamp) used on backend +/// credential types (e.g. `AwsCredential::expires_in`, `google::Token::expires_at`). +#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +pub(crate) fn system_time_to_timestamp( + time: std::time::SystemTime, +) -> reqsign_core::Result { + let millis = time + .duration_since(std::time::UNIX_EPOCH) + .map_err(|e| { + reqsign_core::Error::unexpected(format!( + "credential expiry precedes the UNIX epoch: {e}" + )) + })? + .as_millis(); + let millis = i64::try_from(millis).map_err(|_| { + reqsign_core::Error::unexpected("credential expiry overflows i64 milliseconds") + })?; + reqsign_core::time::Timestamp::from_millisecond(millis) + .map_err(|e| reqsign_core::Error::unexpected(format!("invalid credential expiry: {e}"))) +} From 26d7fcc6d299a8e814b6cceca480409b2c2cc3c9 Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Thu, 30 Jul 2026 23:08:08 +0100 Subject: [PATCH 2/7] chore: fix typo --- crates/catalog/rest/src/client.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index cc4b1c2f97..8ac2bbfb45 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -299,7 +299,7 @@ pub(crate) async fn deserialize_catalog_response( serde_json::from_slice::(&bytes).map_err(|e| { // Successful REST responses can contain OAuth tokens and delegated - // storage credentials. Never copy an unparseable response into an error. + // storage credentials. Never copy an unparsable response into an error. Error::new( ErrorKind::Unexpected, "Failed to parse response from rest catalog server", From 8af37a53d8a5ec05bd523a8391ebd9c774bbafae Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Sun, 2 Aug 2026 17:36:00 +0100 Subject: [PATCH 3/7] refactor(rest): preserve vended credential scopes --- crates/catalog/rest/src/credential.rs | 142 ++++++++++++++------------ crates/iceberg/public-api.txt | 1 + crates/iceberg/src/io/storage/mod.rs | 8 +- crates/storage/opendal/src/gcs.rs | 6 +- crates/storage/opendal/src/s3.rs | 6 +- crates/storage/opendal/src/utils.rs | 18 ++++ 6 files changed, 111 insertions(+), 70 deletions(-) diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs index 14363b3a71..457e90299f 100644 --- a/crates/catalog/rest/src/credential.rs +++ b/crates/catalog/rest/src/credential.rs @@ -74,7 +74,8 @@ struct CloudRefresh { /// Whether to jitter successful prefetch times like AWS `CachedSupplier`. jitter_prefetch: bool, /// Parse a complete credential from catalog-supplied properties. - parse_credential: fn(&HashMap) -> Result, + parse_credential: + fn(config: &HashMap, prefix: Option) -> Result, } impl CloudRefresh { @@ -105,14 +106,13 @@ impl CloudRefresh { /// supported backend matches (in which case static credentials are used /// as-is, as before). fn for_location(location: &str) -> Option<&'static Self> { - let scheme = scheme_of(location)?; Self::SUPPORTED .iter() - .find(|cloud| cloud.schemes.contains(&scheme.as_str())) + .find(|cloud| cloud.matches_location(location)) } - fn matches_credential_prefix(&self, prefix: &str) -> bool { - scheme_of(prefix).is_some_and(|scheme| self.schemes.contains(&scheme.as_str())) + fn matches_location(&self, location: &str) -> bool { + scheme_of(location).is_some_and(|scheme| self.schemes.contains(&scheme.as_str())) } } @@ -130,10 +130,9 @@ const INITIAL_FAILURE_BACKOFF: Duration = Duration::from_secs(1); /// Maximum failure backoff while a cached credential remains usable. const MAX_FAILURE_BACKOFF: Duration = Duration::from_secs(30); -/// A vended credential paired with the storage-location prefix it applies to. +/// A cached vended credential and its refresh schedule. #[derive(Clone)] struct CachedEntry { - prefix: String, credential: StorageCredential, /// When this entry becomes eligible for prefetch. `None` means it does not /// expire and therefore never needs proactive refresh. @@ -141,12 +140,11 @@ struct CachedEntry { } impl CachedEntry { - fn new(prefix: String, credential: StorageCredential, jitter_prefetch: bool) -> Self { + fn new(credential: StorageCredential, jitter_prefetch: bool) -> Self { let refresh_at = credential .expires_at .map(|expires_at| prefetch_time(expires_at, jitter_prefetch)); Self { - prefix, credential, refresh_at, } @@ -155,13 +153,13 @@ impl CachedEntry { /// Seed entries that are already inside the nominal five-minute window are /// immediately due. Otherwise AWS applies the same jitter as it does to a /// freshly fetched value. - fn seed(prefix: String, credential: StorageCredential, jitter_prefetch: bool) -> Self { + fn seed(credential: StorageCredential, jitter_prefetch: bool) -> Self { let due = credential.expires_at.is_some_and(|expires_at| { SystemTime::now() .checked_add(REFRESH_BUFFER) .is_none_or(|refresh_boundary| refresh_boundary >= expires_at) }); - let mut entry = Self::new(prefix, credential, jitter_prefetch); + let mut entry = Self::new(credential, jitter_prefetch); if due { entry.refresh_at = Some(UNIX_EPOCH); } @@ -219,6 +217,12 @@ impl RestVendedCredentialProvider { } } + fn configured_cloud_for_location(&self, location: &str) -> Option<&ConfiguredCloud> { + self.clouds + .iter() + .find(|configured| configured.cloud.matches_location(location)) + } + /// Fetch fresh credentials from the catalog's credentials endpoint. async fn fetch(&self, configured: &ConfiguredCloud) -> Result> { let mut request = self.client.request(Method::GET, &configured.endpoint); @@ -234,15 +238,13 @@ impl RestVendedCredentialProvider { parsed .storage_credentials .into_iter() - .filter(|sc| configured.cloud.matches_credential_prefix(&sc.prefix)) + .filter(|sc| configured.cloud.matches_location(&sc.prefix)) .map(|sc| { - (configured.cloud.parse_credential)(&sc.config).map(|credential| { - CachedEntry::new( - sc.prefix, - credential, - configured.cloud.jitter_prefetch, - ) - }) + (configured.cloud.parse_credential)(&sc.config, Some(sc.prefix)).map( + |credential| { + CachedEntry::new(credential, configured.cloud.jitter_prefetch) + }, + ) }) .collect() } @@ -335,30 +337,22 @@ async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> CacheDecisi #[async_trait] impl StorageCredentialProvider for RestVendedCredentialProvider { fn supports_path(&self, path: &str) -> bool { - CloudRefresh::for_location(path).is_some_and(|cloud| { - self.clouds - .iter() - .any(|configured| configured.cloud.endpoint_key == cloud.endpoint_key) - }) + self.configured_cloud_for_location(path).is_some() } async fn load_credential(&self, path: &str) -> Result { - let cloud = CloudRefresh::for_location(path).ok_or_else(|| { - Error::new( + if CloudRefresh::for_location(path).is_none() { + return Err(Error::new( ErrorKind::FeatureUnsupported, format!("no credential refresh implementation for storage location: {path}"), + )); + } + let configured = self.configured_cloud_for_location(path).ok_or_else(|| { + Error::new( + ErrorKind::FeatureUnsupported, + format!("credential refresh is not configured for storage location: {path}"), ) })?; - let configured = self - .clouds - .iter() - .find(|configured| configured.cloud.endpoint_key == cloud.endpoint_key) - .ok_or_else(|| { - Error::new( - ErrorKind::FeatureUnsupported, - format!("credential refresh is not configured for storage location: {path}"), - ) - })?; let current = match cache_decision(configured, path).await { CacheDecision::Use(credential) => return Ok(credential), @@ -395,8 +389,14 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { fn longest_prefix_match<'a>(entries: &'a [CachedEntry], path: &str) -> Option<&'a CachedEntry> { entries .iter() - .filter(|entry| path.starts_with(&entry.prefix)) - .max_by_key(|entry| entry.prefix.len()) + .filter(|entry| { + entry + .credential + .prefix + .as_deref() + .is_none_or(|prefix| path.starts_with(prefix)) + }) + .max_by_key(|entry| entry.credential.prefix.as_deref().map_or(0, str::len)) } /// Compute a successful credential's prefetch time. @@ -451,17 +451,9 @@ pub(crate) fn build_vended_credential_provider( .get(cloud.endpoint_key) .filter(|value| !value.is_empty())?, ); - let entries = (cloud.parse_credential)(props) + let entries = (cloud.parse_credential)(props, None) .ok() - .map(|credential| { - vec![CachedEntry::seed( - // Flat table properties carry no prefix. This cache is - // already cloud-specific, so an empty prefix is safe. - String::new(), - credential, - cloud.jitter_prefetch, - )] - }) + .map(|credential| vec![CachedEntry::seed(credential, cloud.jitter_prefetch)]) .unwrap_or_default(); Some(ConfiguredCloud { @@ -509,12 +501,16 @@ fn scheme_of(location: &str) -> Option { } /// Parse a complete S3 credential returned by the credentials endpoint. -fn parse_s3_credential(config: &HashMap) -> Result { +fn parse_s3_credential( + config: &HashMap, + prefix: Option, +) -> Result { let access_key_id = required_nonempty(config, S3_ACCESS_KEY_ID)?; let secret_access_key = required_nonempty(config, S3_SECRET_ACCESS_KEY)?; let session_token = required_nonempty(config, S3_SESSION_TOKEN)?; let expires_at = required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?; Ok(StorageCredential { + prefix, kind: StorageCredentialKind::S3(S3Credential { access_key_id, secret_access_key, @@ -525,10 +521,14 @@ fn parse_s3_credential(config: &HashMap) -> Result) -> Result { +fn parse_gcs_credential( + config: &HashMap, + prefix: Option, +) -> Result { let token = required_nonempty(config, GCS_TOKEN)?; let expires_at = required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?; Ok(StorageCredential { + prefix, kind: StorageCredentialKind::Gcs(GcsCredential { token }), expires_at: Some(expires_at), }) @@ -581,8 +581,13 @@ mod tests { .to_string() } - fn s3_cred(access_key_id: &str, expires_at: Option) -> StorageCredential { + fn s3_cred( + prefix: Option<&str>, + access_key_id: &str, + expires_at: Option, + ) -> StorageCredential { StorageCredential { + prefix: prefix.map(str::to_string), kind: StorageCredentialKind::S3(S3Credential { access_key_id: access_key_id.to_string(), secret_access_key: "secret".to_string(), @@ -600,11 +605,7 @@ mod tests { } fn cached_s3(prefix: &str, access_key_id: &str, expires_at: Option) -> CachedEntry { - CachedEntry::new( - prefix.to_string(), - s3_cred(access_key_id, expires_at), - false, - ) + CachedEntry::new(s3_cred(Some(prefix), access_key_id, expires_at), false) } #[test] @@ -649,25 +650,26 @@ mod tests { assert!(CloudRefresh::for_location("not a url").is_none()); // Azure not supported yet -> no provider, static creds used as-is. assert!(CloudRefresh::for_location("abfss://fs@acct.dfs.core.windows.net/k").is_none()); - assert!(CloudRefresh::AWS.matches_credential_prefix("s3://bucket/path")); - assert!(!CloudRefresh::AWS.matches_credential_prefix("s3evil://bucket/path")); + assert!(CloudRefresh::AWS.matches_location("s3://bucket/path")); + assert!(!CloudRefresh::AWS.matches_location("s3evil://bucket/path")); } #[test] fn parses_s3_credential() { let mut config = HashMap::new(); - assert!(parse_s3_credential(&config).is_err()); + assert!(parse_s3_credential(&config, None).is_err()); config.insert(S3_ACCESS_KEY_ID.to_string(), "AK".to_string()); - assert!(parse_s3_credential(&config).is_err()); + assert!(parse_s3_credential(&config, None).is_err()); config.insert(S3_SECRET_ACCESS_KEY.to_string(), "SK".to_string()); - assert!(parse_s3_credential(&config).is_err()); + assert!(parse_s3_credential(&config, None).is_err()); config.insert(S3_SESSION_TOKEN.to_string(), "TOK".to_string()); config.insert( S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), "1500".to_string(), ); - let credential = parse_s3_credential(&config).unwrap(); + let credential = parse_s3_credential(&config, Some("s3://bucket".to_string())).unwrap(); + assert_eq!(credential.prefix.as_deref(), Some("s3://bucket")); assert_eq!( credential.expires_at, Some(UNIX_EPOCH + Duration::from_millis(1500)) @@ -681,12 +683,13 @@ mod tests { #[test] fn parse_gcs_requires_token_and_expiry() { let mut config = HashMap::new(); - assert!(parse_gcs_credential(&config).is_err()); + assert!(parse_gcs_credential(&config, None).is_err()); config.insert(GCS_TOKEN.to_string(), "ya29.token".to_string()); - assert!(parse_gcs_credential(&config).is_err()); + assert!(parse_gcs_credential(&config, None).is_err()); config.insert(GCS_TOKEN_EXPIRES_AT.to_string(), "2000".to_string()); - let credential = parse_gcs_credential(&config).unwrap(); + let credential = parse_gcs_credential(&config, Some("gs://bucket".to_string())).unwrap(); + assert_eq!(credential.prefix.as_deref(), Some("gs://bucket")); match &credential.kind { StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token, "ya29.token"), other => panic!("expected GCS, got {other:?}"), @@ -737,10 +740,11 @@ mod tests { #[test] fn seed_inside_nominal_window_is_immediately_due() { let credential = s3_cred( + None, "a", Some(SystemTime::now() + REFRESH_BUFFER - Duration::from_secs(1)), ); - let entry = CachedEntry::seed(String::new(), credential, true); + let entry = CachedEntry::seed(credential, true); let now = SystemTime::now(); assert!(!entry.is_fresh(now)); assert!(entry.is_unexpired(now)); @@ -770,6 +774,10 @@ mod tests { ]; let got = longest_prefix_match(&entries, "s3://bucket/warehouse/db/t/f").unwrap(); assert_eq!(s3_access_key_id(&got.credential), "narrow"); + assert_eq!( + got.credential.prefix.as_deref(), + Some("s3://bucket/warehouse/db") + ); assert!(longest_prefix_match(&entries, "s3://other/x").is_none()); let fresh = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)); @@ -876,6 +884,7 @@ mod tests { .load_credential("s3://bucket/warehouse/f") .await .unwrap(); + assert_eq!(credential.prefix.as_deref(), Some("s3://bucket")); match credential.kind { StorageCredentialKind::S3(s3) => { @@ -916,6 +925,7 @@ mod tests { .expect("provider should be built"); let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(credential.prefix, None); assert_eq!(s3_access_key_id(&credential), "SEED_AK"); mock.assert_async().await; } diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index 2771ffd998..c8caceacf1 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -1019,6 +1019,7 @@ pub fn iceberg::io::StorageConfig::deserialize<__D>(__deserializer: __D) -> core pub struct iceberg::io::StorageCredential pub iceberg::io::StorageCredential::expires_at: core::option::Option pub iceberg::io::StorageCredential::kind: iceberg::io::StorageCredentialKind +pub iceberg::io::StorageCredential::prefix: core::option::Option impl core::clone::Clone for iceberg::io::StorageCredential pub fn iceberg::io::StorageCredential::clone(&self) -> iceberg::io::StorageCredential impl core::fmt::Debug for iceberg::io::StorageCredential diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index b10bd380c0..dd812a2d45 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -188,13 +188,17 @@ pub trait StorageCredentialProvider: Debug + Send + Sync { /// `path` is the absolute location being accessed (e.g. /// `s3://bucket/warehouse/db/table/...`). Providers that vend distinct /// credentials per location prefix use it to select the most specific - /// match. + /// match. When the selected credential has a declared + /// [`StorageCredential::prefix`], it must cover `path`. async fn load_credential(&self, path: &str) -> Result; } -/// A vended storage credential together with when it expires. +/// A vended storage credential together with its scope and expiry. #[derive(Clone, Debug)] pub struct StorageCredential { + /// Storage-location prefix this credential is scoped to. `None` represents a + /// credential without a declared scope, sourced from flat storage properties. + pub prefix: Option, /// The backend-specific credential material. pub kind: StorageCredentialKind, /// When the credential expires, if known. `None` means non-expiring and diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index 3bbaba2c5b..311bd1ac88 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -30,7 +30,9 @@ use reqsign_core::{Context, Error as ReqsignError, ProvideCredential, Result as use reqsign_google::{Credential as GoogleCredential, Token as GoogleToken}; use url::Url; -use crate::utils::{from_opendal_error, is_truthy, system_time_to_timestamp}; +use crate::utils::{ + from_opendal_error, is_truthy, system_time_to_timestamp, validate_credential_prefix, +}; /// Parse iceberg properties to [`GcsConfig`]. pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result { @@ -172,6 +174,8 @@ impl ProvideCredential for VendedGcsCredentialProvider { .with_source(e) })?; + validate_credential_prefix(&self.path, credential.prefix.as_deref())?; + let expires_at = credential .expires_at .map(system_time_to_timestamp) diff --git a/crates/storage/opendal/src/s3.rs b/crates/storage/opendal/src/s3.rs index 8e57cdfdb6..e6b825347d 100644 --- a/crates/storage/opendal/src/s3.rs +++ b/crates/storage/opendal/src/s3.rs @@ -38,7 +38,9 @@ use reqsign_core::{ }; use url::Url; -use crate::utils::{from_opendal_error, is_truthy, system_time_to_timestamp}; +use crate::utils::{ + from_opendal_error, is_truthy, system_time_to_timestamp, validate_credential_prefix, +}; /// Parse iceberg props to s3 config. pub(crate) fn s3_config_parse(mut m: HashMap) -> Result { @@ -217,6 +219,8 @@ impl ProvideCredential for VendedS3CredentialProvider { .with_source(e) })?; + validate_credential_prefix(&self.path, credential.prefix.as_deref())?; + let expires_in = credential .expires_at .map(system_time_to_timestamp) diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index 3afd319f22..c168093fc2 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -49,3 +49,21 @@ pub(crate) fn system_time_to_timestamp( reqsign_core::time::Timestamp::from_millisecond(millis) .map_err(|e| reqsign_core::Error::unexpected(format!("invalid credential expiry: {e}"))) } + +/// Validate that a provider's declared credential prefix covers the path for +/// which the backend requested the credential. +#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +pub(crate) fn validate_credential_prefix( + path: &str, + prefix: Option<&str>, +) -> reqsign_core::Result<()> { + match prefix { + Some("") => Err(reqsign_core::Error::unexpected( + "vended credential has an empty storage prefix", + )), + Some(prefix) if !path.starts_with(prefix) => Err(reqsign_core::Error::unexpected(format!( + "vended credential prefix {prefix:?} does not cover storage location {path:?}" + ))), + _ => Ok(()), + } +} From f7b065a3fc8a3f4f93320eeafba340ed88cd4e3f Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Fri, 7 Aug 2026 10:38:27 +0100 Subject: [PATCH 4/7] fix(rest): harden vended credential refresh --- crates/catalog/rest/src/catalog.rs | 34 +- crates/catalog/rest/src/client.rs | 285 +++++--- crates/catalog/rest/src/credential.rs | 879 +++++++++++++++++------- crates/iceberg/public-api.txt | 25 +- crates/iceberg/src/io/storage/mod.rs | 116 +++- crates/storage/opendal/src/gcs.rs | 120 +++- crates/storage/opendal/src/lib.rs | 207 +++++- crates/storage/opendal/src/resolving.rs | 77 +++ crates/storage/opendal/src/s3.rs | 132 +++- crates/storage/opendal/src/utils.rs | 18 +- 10 files changed, 1474 insertions(+), 419 deletions(-) diff --git a/crates/catalog/rest/src/catalog.rs b/crates/catalog/rest/src/catalog.rs index 1a400fd7da..f86e29c0d9 100644 --- a/crates/catalog/rest/src/catalog.rs +++ b/crates/catalog/rest/src/catalog.rs @@ -41,6 +41,7 @@ use typed_builder::TypedBuilder; use crate::client::{ HttpClient, deserialize_catalog_response, deserialize_unexpected_catalog_error, }; +use crate::credential::build_vended_credential_provider; use crate::endpoint::{Endpoint, V1_NAMESPACE_EXISTS, V1_TABLE_EXISTS}; use crate::types::{ CatalogConfig, CommitTableRequest, CommitTableResponse, CreateNamespaceRequest, @@ -315,7 +316,7 @@ impl RestCatalogConfig { HeaderValue::from_str(value).map_err(|e| { Error::new( ErrorKind::DataInvalid, - format!("Invalid header value: {value}"), + format!("Invalid value for header: {key}"), ) .with_source(e) })?, @@ -524,6 +525,7 @@ impl RestCatalog { &self, metadata_location: Option<&str>, extra_config: Option>, + table_auth_config: Option>, ) -> Result { let context = self.context().await?; let mut props = context.config.props.clone(); @@ -558,10 +560,11 @@ impl RestCatalog { // If the catalog vends refreshable credentials for this table's storage, // attach a provider so the backend re-fetches them before they expire. - let credential_provider = crate::credential::build_vended_credential_provider( + let credential_provider = build_vended_credential_provider( context.client.clone(), &context.config.uri, &props, + table_auth_config.as_ref(), )?; let mut builder = FileIOBuilder::new(factory).with_props(props); @@ -872,14 +875,15 @@ impl Catalog for RestCatalog { "Metadata location missing in `create_table` response!", ))?; - let config = response - .config + let table_config = response.config; + let config = table_config + .clone() .into_iter() .chain(self.user_config.props.clone()) .collect(); let file_io = self - .load_file_io(Some(metadata_location), Some(config)) + .load_file_io(Some(metadata_location), Some(config), Some(table_config)) .await?; let mut table_builder = Table::builder() @@ -932,14 +936,19 @@ impl Catalog for RestCatalog { } }; - let config = response - .config + let table_config = response.config; + let config = table_config + .clone() .into_iter() .chain(self.user_config.props.clone()) .collect(); let file_io = self - .load_file_io(response.metadata_location.as_deref(), Some(config)) + .load_file_io( + response.metadata_location.as_deref(), + Some(config), + Some(table_config), + ) .await?; let mut table_builder = Table::builder() @@ -1074,13 +1083,14 @@ impl Catalog for RestCatalog { "Metadata location missing in `register_table` response!", ))?; - let config = response - .config + let table_config = response.config; + let config = table_config + .clone() .into_iter() .chain(self.user_config.props.clone()) .collect(); let file_io = self - .load_file_io(Some(metadata_location), Some(config)) + .load_file_io(Some(metadata_location), Some(config), Some(table_config)) .await?; let mut table_builder = Table::builder() @@ -1156,7 +1166,7 @@ impl Catalog for RestCatalog { }; let file_io = self - .load_file_io(Some(&response.metadata_location), None) + .load_file_io(Some(&response.metadata_location), None, None) .await?; let mut table_builder = Table::builder() diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index 8ac2bbfb45..bc9157212b 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -17,6 +17,7 @@ use std::collections::HashMap; use std::fmt::{Debug, Formatter}; +use std::sync::Arc; use http::StatusCode; use iceberg::{Error, ErrorKind, Result}; @@ -34,7 +35,7 @@ pub(crate) struct HttpClient { /// The token to be used for authentication. /// /// It's possible to fetch the token from the server while needed. - token: Mutex>, + token: Arc>>, /// The token endpoint to be used for authentication. token_endpoint: String, /// The credential to be used for authentication. @@ -66,7 +67,7 @@ impl HttpClient { let extra_headers = cfg.extra_headers()?; Ok(HttpClient { client: cfg.client().unwrap_or_default(), - token: Mutex::new(cfg.token()), + token: Arc::new(Mutex::new(cfg.token())), token_endpoint: cfg.get_token_endpoint(), credential: cfg.credential(), extra_headers, @@ -83,14 +84,70 @@ impl HttpClient { pub(crate) fn for_table( &self, catalog_uri: &str, - props: HashMap, + props: &HashMap, ) -> Result { let cfg = RestCatalogConfig::builder() .uri(catalog_uri.to_string()) - .props(props) + .props(props.clone()) .client(Some(self.client.clone())) .build(); - Self::new(&cfg) + let table_token = cfg.token(); + let configured_credential = cfg.credential(); + let has_table_token = props.contains_key("token"); + let has_table_credential = props.contains_key("credential"); + let has_table_auth_header = props.keys().any(|key| { + key.strip_prefix("header.") + .is_some_and(|name| name.eq_ignore_ascii_case("authorization")) + }); + let has_oauth_params = ["scope", "audience", "resource"] + .iter() + .any(|key| props.contains_key(*key)); + let has_oauth_config = props.contains_key("oauth2-server-uri") || has_oauth_params; + let has_oauth_override = + has_table_credential || (has_oauth_config && self.credential.is_some()); + let has_table_auth = has_table_token || has_table_auth_header || has_oauth_override; + let mut extra_headers = self.extra_headers.clone(); + if has_table_auth && !has_table_auth_header { + // Authentication is applied after request headers are initially built, but + // `execute` reapplies `extra_headers`. Do not let an inherited catalog + // Authorization header overwrite table-scoped authentication at that point. + extra_headers.remove(http::header::AUTHORIZATION); + } + extra_headers.extend(cfg.extra_headers()?); + + Ok(Self { + client: self.client.clone(), + token: if has_table_auth { + Arc::new(Mutex::new(table_token)) + } else { + Arc::clone(&self.token) + }, + token_endpoint: if has_oauth_override && props.contains_key("oauth2-server-uri") { + cfg.get_token_endpoint() + } else { + self.token_endpoint.clone() + }, + credential: if has_table_token || has_table_auth_header { + None + } else if has_table_credential { + configured_credential + } else { + self.credential.clone() + }, + extra_headers, + extra_oauth_params: if has_oauth_override && has_oauth_params { + cfg.extra_oauth_params() + } else { + self.extra_oauth_params.clone() + }, + disable_header_redaction: if props + .contains_key(crate::REST_CATALOG_PROP_DISABLE_HEADER_REDACTION) + { + cfg.disable_header_redaction() + } else { + self.disable_header_redaction + }, + }) } /// Update the http client with new configuration. @@ -104,7 +161,10 @@ impl HttpClient { .unwrap_or(self.extra_headers); Ok(HttpClient { client: cfg.client().unwrap_or(self.client), - token: Mutex::new(cfg.token().or_else(|| self.token.into_inner())), + token: match cfg.token() { + Some(token) => Arc::new(Mutex::new(Some(token))), + None => self.token, + }, token_endpoint: if !cfg.get_token_endpoint().is_empty() { cfg.get_token_endpoint() } else { @@ -308,26 +368,10 @@ pub(crate) async fn deserialize_catalog_response( }) } -/// Returns true if the header may carry a secret. -fn is_sensitive_header(name: &str) -> bool { - let name_lower = name.to_lowercase(); - [ - "auth", - "token", - "secret", - "key", - "password", - "cookie", - "credential", - ] - .iter() - .any(|pattern| name_lower.contains(pattern)) -} - -/// Redacts sensitive headers and returns a debug-formatted string. +/// Redacts header values and returns a debug-formatted string. /// /// If `disable_redaction` is true, returns all headers without redaction. -/// Otherwise, replaces sensitive header values with "[REDACTED]". +/// Otherwise, redacts every header value. fn format_headers_redacted(headers: &HeaderMap, disable_redaction: bool) -> String { if disable_redaction { // Return all headers as-is without redaction @@ -338,16 +382,10 @@ fn format_headers_redacted(headers: &HeaderMap, disable_redaction: bool) -> Stri return format!("{all:?}"); } - // Redact sensitive headers by replacing their values with "[REDACTED]" + // Retain names for diagnostics but redact every value let redacted: HashMap<&str, &str> = headers .iter() - .filter_map(|(name, value)| { - if is_sensitive_header(name.as_str()) { - Some((name.as_str(), "[REDACTED]")) - } else { - value.to_str().ok().map(|v| (name.as_str(), v)) - } - }) + .map(|(name, _)| (name.as_str(), "[REDACTED]")) .collect(); format!("{redacted:?}") } @@ -390,96 +428,21 @@ mod tests { } #[test] - fn test_format_headers_redacted_non_sensitive() { + fn test_format_headers_redacts_all_values() { let mut headers = HeaderMap::new(); + headers.insert("authorization", "Bearer secret-token".parse().unwrap()); headers.insert("content-type", "application/json".parse().unwrap()); headers.insert("x-request-id", "abc123".parse().unwrap()); let result = format_headers_redacted(&headers, false); + assert!(result.contains("authorization")); assert!(result.contains("content-type")); - assert!(result.contains("application/json")); assert!(result.contains("x-request-id")); - assert!(result.contains("abc123")); - } - - #[test] - fn test_format_headers_redacted_filters_sensitive() { - let mut headers = HeaderMap::new(); - headers.insert("authorization", "Bearer secret-token".parse().unwrap()); - headers.insert("x-client-secret", "private-value".parse().unwrap()); - headers.insert("x-client-credential", "credential-value".parse().unwrap()); - headers.insert("content-type", "application/json".parse().unwrap()); - - let result = format_headers_redacted(&headers, false); - - // Sensitive header should be present but with redacted value - assert!(result.contains("authorization")); assert!(result.contains("[REDACTED]")); - // Sensitive value should NOT be present assert!(!result.contains("secret-token")); - assert!(!result.contains("private-value")); - assert!(!result.contains("credential-value")); - // Non-sensitive header should be present with actual value - assert!(result.contains("content-type")); - assert!(result.contains("application/json")); - } - - #[test] - fn test_format_headers_redacted_filters_set_cookie() { - let mut headers = HeaderMap::new(); - headers.insert( - "set-cookie", - "CF_Authorization=sensitive-session-token; Path=/; Secure;" - .parse() - .unwrap(), - ); - headers.insert("server", "cloudflare".parse().unwrap()); - - let result = format_headers_redacted(&headers, false); - - // Sensitive header should be present but with redacted value - assert!(result.contains("set-cookie")); - assert!(result.contains("[REDACTED]")); - // Sensitive value should NOT be present - assert!(!result.contains("sensitive-session-token")); - // Non-sensitive header should be present with actual value - assert!(result.contains("server")); - assert!(result.contains("cloudflare")); - } - - #[test] - fn test_format_headers_redacted_filters_all_sensitive() { - let mut headers = HeaderMap::new(); - headers.insert("authorization", "Bearer token".parse().unwrap()); - headers.insert("proxy-authorization", "Basic creds".parse().unwrap()); - headers.insert("set-cookie", "session=abc".parse().unwrap()); - headers.insert("cookie", "session=abc".parse().unwrap()); - headers.insert("x-api-key", "api-key-123".parse().unwrap()); - headers.insert("x-auth-token", "auth-token-456".parse().unwrap()); - headers.insert("x-request-id", "req-123".parse().unwrap()); - - let result = format_headers_redacted(&headers, false); - - // All sensitive headers should be present but with redacted values - assert!(result.contains("authorization")); - assert!(result.contains("proxy-authorization")); - assert!(result.contains("set-cookie")); - assert!(result.contains("cookie")); - assert!(result.contains("x-api-key")); - assert!(result.contains("x-auth-token")); - assert!(result.contains("[REDACTED]")); - - // Ensure no sensitive values leaked - assert!(!result.contains("Bearer token")); - assert!(!result.contains("Basic creds")); - assert!(!result.contains("session=abc")); - assert!(!result.contains("api-key-123")); - assert!(!result.contains("auth-token-456")); - - // Non-sensitive header should be present with actual value - assert!(result.contains("x-request-id")); - assert!(result.contains("req-123")); + assert!(!result.contains("application/json")); + assert!(!result.contains("abc123")); } #[test] @@ -501,4 +464,102 @@ mod tests { // [REDACTED] should NOT be present when redaction is disabled assert!(!result.contains("[REDACTED]")); } + + #[tokio::test] + async fn table_client_reuses_parent_token_unless_overridden() { + let inherited_props = HashMap::from([( + "credential".to_string(), + "client-id:client-secret".to_string(), + )]); + let config = RestCatalogConfig::builder() + .uri("https://catalog.example".to_string()) + .props(inherited_props.clone()) + .build(); + let client = HttpClient::new(&config).unwrap(); + *client.token.lock().await = Some("catalog-token".to_string()); + + let inherited = client + .for_table("https://catalog.example", &HashMap::new()) + .unwrap(); + assert!(Arc::ptr_eq(&client.token, &inherited.token)); + assert_eq!( + inherited.token.lock().await.as_deref(), + Some("catalog-token") + ); + + let overridden = client + .for_table( + "https://catalog.example", + &HashMap::from([("token".to_string(), "table-token".to_string())]), + ) + .unwrap(); + assert!(!Arc::ptr_eq(&client.token, &overridden.token)); + assert_eq!( + overridden.token.lock().await.as_deref(), + Some("table-token") + ); + } + + #[tokio::test] + async fn table_auth_removes_inherited_authorization_header() { + let config = RestCatalogConfig::builder() + .uri("https://catalog.example".to_string()) + .props(HashMap::from([ + ("token".to_string(), "catalog-token".to_string()), + ( + "header.Authorization".to_string(), + "Bearer catalog-header-token".to_string(), + ), + ])) + .build(); + let client = HttpClient::new(&config).unwrap(); + + let table_client = client + .for_table( + "https://catalog.example", + &HashMap::from([("token".to_string(), "table-token".to_string())]), + ) + .unwrap(); + + assert!( + !table_client + .extra_headers + .contains_key(http::header::AUTHORIZATION) + ); + assert_eq!( + table_client.token.lock().await.as_deref(), + Some("table-token") + ); + } + + #[tokio::test] + async fn table_oauth_options_without_credential_reuse_parent_token() { + let config = RestCatalogConfig::builder() + .uri("https://catalog.example".to_string()) + .props(HashMap::from([( + "token".to_string(), + "catalog-token".to_string(), + )])) + .build(); + let client = HttpClient::new(&config).unwrap(); + let table_props = HashMap::from([ + ("scope".to_string(), "table-scope".to_string()), + ( + "oauth2-server-uri".to_string(), + "https://table-auth.example/token".to_string(), + ), + ]); + + let table_client = client + .for_table("https://catalog.example", &table_props) + .unwrap(); + + assert!(Arc::ptr_eq(&client.token, &table_client.token)); + assert_eq!( + table_client.token.lock().await.as_deref(), + Some("catalog-token") + ); + assert_eq!(table_client.token_endpoint, client.token_endpoint); + assert_eq!(table_client.extra_oauth_params, client.extra_oauth_params); + } } diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs index 457e90299f..77eed648fc 100644 --- a/crates/catalog/rest/src/credential.rs +++ b/crates/catalog/rest/src/credential.rs @@ -28,8 +28,11 @@ //! single backend-agnostic provider with an independent endpoint and cache for //! each configured cloud. The path being accessed selects the cloud cache, and //! the returned [`StorageCredential`] enum lets the storage adapter enforce the -//! expected backend-specific type. This preserves Java's per-cloud refresh -//! policy while supporting mixed-cloud tables through a resolving FileIO. +//! expected backend-specific type. This preserves Java's per-cloud credential +//! selection and prefetch policies while supporting mixed-cloud tables through +//! a resolving FileIO. Unlike Java's scheduled refresh, which permanently stops +//! after a failed fetch, transient failures are retried here with jittered +//! exponential backoff while an unexpired credential remains available. //! //! # Adding a cloud //! @@ -38,7 +41,7 @@ //! teach the storage adapter to consume it. Then write its `parse_*` function, //! add a `CloudRefresh` constant, and list it in [`CloudRefresh::SUPPORTED`]. -use std::collections::HashMap; +use std::collections::{HashMap, HashSet}; use std::sync::Arc; use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; @@ -69,7 +72,7 @@ struct CloudRefresh { schemes: &'static [&'static str], /// Table property naming the refresh endpoint (absolute or catalog-relative). endpoint_key: &'static str, - /// Table property to opt out; refresh is enabled unless this is `"false"`. + /// Table property controlling refresh; only missing or case-insensitive `"true"` enables it. enabled_key: &'static str, /// Whether to jitter successful prefetch times like AWS `CachedSupplier`. jitter_prefetch: bool, @@ -95,7 +98,7 @@ impl CloudRefresh { jitter_prefetch: false, parse_credential: parse_gcs_credential, }; - // TODO: Azure (ADLS) is not yet supported: opendal 0.57's Azdls builder exposes no + // TODO: Azure (ADLS) is not yet supported: opendal's Azdls builder exposes no // custom credential-provider hook, and reqsign's SAS-token credential has no // expiry, so reqsign-based refresh isn't possible. @@ -112,7 +115,10 @@ impl CloudRefresh { } fn matches_location(&self, location: &str) -> bool { - scheme_of(location).is_some_and(|scheme| self.schemes.contains(&scheme.as_str())) + self.schemes + .iter() + .any(|scheme| location.eq_ignore_ascii_case(scheme)) + || scheme_of(location).is_some_and(|scheme| self.schemes.contains(&scheme.as_str())) } } @@ -127,7 +133,7 @@ const MIN_REFRESH_BUFFER: Duration = Duration::from_mins(1); /// value through the full value. const INITIAL_FAILURE_BACKOFF: Duration = Duration::from_secs(1); -/// Maximum failure backoff while a cached credential remains usable. +/// Maximum delay between failed refresh attempts. const MAX_FAILURE_BACKOFF: Duration = Duration::from_secs(30); /// A cached vended credential and its refresh schedule. @@ -142,7 +148,7 @@ struct CachedEntry { impl CachedEntry { fn new(credential: StorageCredential, jitter_prefetch: bool) -> Self { let refresh_at = credential - .expires_at + .expires_at() .map(|expires_at| prefetch_time(expires_at, jitter_prefetch)); Self { credential, @@ -154,7 +160,7 @@ impl CachedEntry { /// immediately due. Otherwise AWS applies the same jitter as it does to a /// freshly fetched value. fn seed(credential: StorageCredential, jitter_prefetch: bool) -> Self { - let due = credential.expires_at.is_some_and(|expires_at| { + let due = credential.expires_at().is_some_and(|expires_at| { SystemTime::now() .checked_add(REFRESH_BUFFER) .is_none_or(|refresh_boundary| refresh_boundary >= expires_at) @@ -172,11 +178,22 @@ impl CachedEntry { fn is_unexpired(&self, now: SystemTime) -> bool { self.credential - .expires_at + .expires_at() .is_none_or(|expires_at| now < expires_at) } } +/// Parsed entries from one successful credentials response. +struct ParsedCredentials { + entries: Vec, + errors: Vec, +} + +struct CredentialError { + prefix: String, + error: Error, +} + /// Cached credentials plus failure-backoff state. struct CacheState { entries: Vec, @@ -184,6 +201,31 @@ struct CacheState { retry_not_before: Option, } +impl CacheState { + fn record_success(&mut self) { + self.consecutive_failures = 0; + self.retry_not_before = None; + } + + fn record_failure(&mut self) { + self.consecutive_failures = self.consecutive_failures.saturating_add(1); + self.retry_not_before = + Instant::now().checked_add(failure_backoff(self.consecutive_failures)); + } + + fn insert_fallback_if_missing(&mut self, fallback: CachedEntry, now: SystemTime) -> bool { + let has_unexpired_entry = self.entries.iter().any(|entry| { + entry.is_unexpired(now) && entry.credential.prefix() == fallback.credential.prefix() + }); + if has_unexpired_entry { + false + } else { + self.entries.push(fallback); + true + } + } +} + struct ConfiguredCloud { cloud: &'static CloudRefresh, endpoint: String, @@ -224,7 +266,7 @@ impl RestVendedCredentialProvider { } /// Fetch fresh credentials from the catalog's credentials endpoint. - async fn fetch(&self, configured: &ConfiguredCloud) -> Result> { + async fn fetch(&self, configured: &ConfiguredCloud) -> Result { let mut request = self.client.request(Method::GET, &configured.endpoint); if let Some(plan_id) = &self.plan_id { request = request.query(&[("planId", plan_id)]); @@ -235,18 +277,39 @@ impl RestVendedCredentialProvider { match response.status() { StatusCode::OK => { let parsed: LoadCredentialsResponse = response.json().await?; - parsed + let mut entries = Vec::new(); + let mut errors = Vec::new(); + let matching_credentials = parsed .storage_credentials .into_iter() - .filter(|sc| configured.cloud.matches_location(&sc.prefix)) - .map(|sc| { - (configured.cloud.parse_credential)(&sc.config, Some(sc.prefix)).map( - |credential| { - CachedEntry::new(credential, configured.cloud.jitter_prefetch) - }, - ) - }) - .collect() + .filter(|credential| configured.cloud.matches_location(&credential.prefix)); + let parse_credential = configured.cloud.parse_credential; + let now = SystemTime::now(); + + for storage_credential in matching_credentials { + let prefix = storage_credential.prefix; + let parsed_credential = + parse_credential(&storage_credential.config, Some(prefix.clone())); + match parsed_credential { + Ok(credential) => { + let entry = + CachedEntry::new(credential, configured.cloud.jitter_prefetch); + if entry.is_unexpired(now) { + entries.push(entry); + } else { + errors.push(CredentialError { + prefix, + error: Error::new( + ErrorKind::DataInvalid, + "invalid vended credential: credential is already expired", + ), + }); + } + } + Err(error) => errors.push(CredentialError { prefix, error }), + } + } + Ok(ParsedCredentials { entries, errors }) } _ => Err(deserialize_unexpected_catalog_error( response, @@ -262,32 +325,101 @@ impl RestVendedCredentialProvider { path: &str, fallback: Option, ) -> Result { - let refreshed = self.fetch(configured).await.and_then(|entries| { - let credential = longest_prefix_match(&entries, path) - .filter(|entry| entry.is_unexpired(SystemTime::now())) - .map(|entry| entry.credential.clone()) - .ok_or_else(|| { - Error::new( - ErrorKind::Unexpected, - format!("no unexpired vended credential matches storage location: {path}"), - ) - })?; - Ok((entries, credential)) - }); + match self.fetch(configured).await { + Ok(ParsedCredentials { entries, errors }) => { + let invalid_prefixes = errors + .iter() + .map(|error| error.prefix.clone()) + .collect::>(); + let credential_error = errors + .into_iter() + .filter(|error| path.starts_with(&error.prefix)) + .max_by_key(|error| error.prefix.len()); + let failure = credential_error + .map(|error| error.error) + .unwrap_or_else(|| { + Error::new( + ErrorKind::Unexpected, + format!( + "no unexpired vended credential matches storage location: {path}" + ), + ) + }); - match refreshed { - Ok((entries, credential)) => { let mut cache = configured.cache.lock().await; + if entries.is_empty() { + cache.record_failure(); + return fallback + .filter(|entry| entry.is_unexpired(SystemTime::now())) + .map(|entry| entry.credential) + .ok_or(failure); + } + + let now = SystemTime::now(); + let invalid_fallbacks = cache + .entries + .iter() + .filter(|entry| entry.is_unexpired(now)) + .filter(|entry| { + entry + .credential + .prefix() + .is_some_and(|prefix| invalid_prefixes.contains(prefix)) + }) + .cloned() + .collect::>(); cache.entries = entries; - cache.consecutive_failures = 0; - cache.retry_not_before = None; - Ok(credential) + + // An invalid replacement for one prefix must not evict that + // prefix's still-valid cached credential. Iterate in reverse + // because equal-length prefix selection uses the last cached entry. + // This preserves the same credential if a bad response follows a + // response with duplicate prefixes. + let mut restored_fallback_prefixes = HashSet::new(); + for fallback in invalid_fallbacks.into_iter().rev() { + let prefix = fallback.credential.prefix().map(str::to_owned); + if cache.insert_fallback_if_missing(fallback, now) + && let Some(prefix) = prefix + { + restored_fallback_prefixes.insert(prefix); + } + } + + let selected = longest_prefix_match(&cache.entries, path) + .filter(|entry| entry.is_unexpired(now)) + .map(|entry| { + let is_restored_fallback = entry + .credential + .prefix() + .is_some_and(|prefix| restored_fallback_prefixes.contains(prefix)); + (entry.credential.clone(), is_restored_fallback) + }); + + if let Some((credential, is_restored_fallback)) = selected { + if is_restored_fallback { + cache.record_failure(); + } else { + cache.record_success(); + } + return Ok(credential); + } + + // Cache valid credentials for other prefixes while retaining this + // path's still-usable credential when its replacement was absent or + // invalid. Backoff ensures the failed path is retried promptly. + if let Some(fallback) = fallback.filter(|entry| entry.is_unexpired(now)) { + let credential = fallback.credential.clone(); + cache.insert_fallback_if_missing(fallback, now); + cache.record_failure(); + return Ok(credential); + } + + cache.record_failure(); + Err(failure) } Err(fetch_error) => { let mut cache = configured.cache.lock().await; - cache.consecutive_failures = cache.consecutive_failures.saturating_add(1); - cache.retry_not_before = - Instant::now().checked_add(failure_backoff(cache.consecutive_failures)); + cache.record_failure(); // Graceful degradation: while the cached credential remains // usable, serve it and retry after jittered backoff. Expired @@ -312,6 +444,14 @@ impl std::fmt::Debug for RestVendedCredentialProvider { enum CacheDecision { Use(StorageCredential), Refresh(Option), + Backoff, +} + +fn refresh_backoff_error(path: &str) -> Error { + Error::new( + ErrorKind::Unexpected, + format!("vended credential refresh is temporarily backed off for storage location: {path}"), + ) } async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> CacheDecision { @@ -326,9 +466,12 @@ async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> CacheDecisi if cache .retry_not_before .is_some_and(|retry_at| Instant::now() < retry_at) - && let Some(entry) = current.as_ref().filter(|entry| entry.is_unexpired(now)) { - return CacheDecision::Use(entry.credential.clone()); + return current + .as_ref() + .filter(|entry| entry.is_unexpired(now)) + .map(|entry| CacheDecision::Use(entry.credential.clone())) + .unwrap_or(CacheDecision::Backoff); } CacheDecision::Refresh(current) @@ -357,6 +500,7 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { let current = match cache_decision(configured, path).await { CacheDecision::Use(credential) => return Ok(credential), CacheDecision::Refresh(current) => current, + CacheDecision::Backoff => return Err(refresh_backoff_error(path)), }; // One caller refreshes, while concurrent callers immediately keep using the @@ -379,6 +523,7 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { let current = match cache_decision(configured, path).await { CacheDecision::Use(credential) => return Ok(credential), CacheDecision::Refresh(current) => current, + CacheDecision::Backoff => return Err(refresh_backoff_error(path)), }; self.refresh_credential(configured, path, current).await @@ -392,11 +537,10 @@ fn longest_prefix_match<'a>(entries: &'a [CachedEntry], path: &str) -> Option<&' .filter(|entry| { entry .credential - .prefix - .as_deref() + .prefix() .is_none_or(|prefix| path.starts_with(prefix)) }) - .max_by_key(|entry| entry.credential.prefix.as_deref().map_or(0, str::len)) + .max_by_key(|entry| entry.credential.prefix().map_or(0, str::len)) } /// Compute a successful credential's prefetch time. @@ -429,18 +573,20 @@ fn failure_backoff(consecutive_failures: u32) -> Duration { /// or `None` when no supported cloud advertises an enabled refresh endpoint. /// /// `base_uri` is the catalog URI, used to resolve a relative endpoint. +/// `table_auth_props` is the unmerged config returned by the table endpoint: +/// keeping it separate prevents local FileIO overrides from masking table auth. pub(crate) fn build_vended_credential_provider( client: Arc, base_uri: &str, props: &HashMap, + table_auth_props: Option<&HashMap>, ) -> Result>> { let clouds = CloudRefresh::SUPPORTED .iter() .filter_map(|cloud| { - // Refresh is enabled by default and invalid booleans disable refresh. let enabled = props .get(cloud.enabled_key) - .is_none_or(|value| value.parse().unwrap_or(false)); + .is_none_or(|value| value.eq_ignore_ascii_case("true")); if !enabled { return None; } @@ -473,7 +619,9 @@ pub(crate) fn build_vended_credential_provider( return Ok(None); } - let table_client = Arc::new(client.for_table(base_uri, props.clone())?); + let empty_auth_props = HashMap::new(); + let table_client = + Arc::new(client.for_table(base_uri, table_auth_props.unwrap_or(&empty_auth_props))?); let plan_id = props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(); Ok(Some(Arc::new(RestVendedCredentialProvider::new( table_client, @@ -509,14 +657,15 @@ fn parse_s3_credential( let secret_access_key = required_nonempty(config, S3_SECRET_ACCESS_KEY)?; let session_token = required_nonempty(config, S3_SESSION_TOKEN)?; let expires_at = required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?; - Ok(StorageCredential { - prefix, - kind: StorageCredentialKind::S3(S3Credential { - access_key_id, - secret_access_key, - session_token: Some(session_token), - }), - expires_at: Some(expires_at), + let credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( + access_key_id, + secret_access_key, + Some(session_token), + ))) + .with_expiration(expires_at); + Ok(match prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, }) } @@ -527,10 +676,11 @@ fn parse_gcs_credential( ) -> Result { let token = required_nonempty(config, GCS_TOKEN)?; let expires_at = required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?; - Ok(StorageCredential { - prefix, - kind: StorageCredentialKind::Gcs(GcsCredential { token }), - expires_at: Some(expires_at), + let credential = StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new(token))) + .with_expiration(expires_at); + Ok(match prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, }) } @@ -586,20 +736,23 @@ mod tests { access_key_id: &str, expires_at: Option, ) -> StorageCredential { - StorageCredential { - prefix: prefix.map(str::to_string), - kind: StorageCredentialKind::S3(S3Credential { - access_key_id: access_key_id.to_string(), - secret_access_key: "secret".to_string(), - session_token: None, - }), - expires_at, + let mut credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( + access_key_id, + "secret", + None, + ))); + if let Some(prefix) = prefix { + credential = credential.with_prefix(prefix); } + if let Some(expires_at) = expires_at { + credential = credential.with_expiration(expires_at); + } + credential } fn s3_access_key_id(credential: &StorageCredential) -> &str { - match &credential.kind { - StorageCredentialKind::S3(s3) => &s3.access_key_id, + match credential.kind() { + StorageCredentialKind::S3(s3) => s3.access_key_id(), other => panic!("expected S3 credential, got {other:?}"), } } @@ -608,6 +761,53 @@ mod tests { CachedEntry::new(s3_cred(Some(prefix), access_key_id, expires_at), false) } + fn test_client(base_uri: &str) -> Arc { + let config = RestCatalogConfig::builder() + .uri(base_uri.to_string()) + .build(); + Arc::new(HttpClient::new(&config).unwrap()) + } + + fn aws_refresh_props(endpoint: &str) -> HashMap { + HashMap::from([( + CloudRefresh::AWS.endpoint_key.to_string(), + endpoint.to_string(), + )]) + } + + fn with_s3_seed( + mut props: HashMap, + access_key_id: &str, + expires_at: SystemTime, + ) -> HashMap { + props.extend([ + (S3_ACCESS_KEY_ID.to_string(), access_key_id.to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), + (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), + ( + S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), + epoch_millis(expires_at), + ), + ]); + props + } + + fn test_provider( + base_uri: &str, + props: &HashMap, + ) -> Arc { + build_vended_credential_provider(test_client(base_uri), base_uri, props, None) + .expect("provider construction should succeed") + .expect("provider should be built") + } + + fn s3_response(prefix: &str, access_key_id: &str, expires_at: SystemTime) -> String { + let expires_at = epoch_millis(expires_at); + format!( + r#"{{"storage-credentials":[{{"prefix":"{prefix}","config":{{"s3.access-key-id":"{access_key_id}","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires_at}"}}}}]}}"# + ) + } + #[test] fn resolve_endpoint_matches_java_semantics() { assert_eq!( @@ -640,7 +840,11 @@ mod tests { } #[test] - fn cloud_selection_uses_url_scheme() { + fn cloud_selection_accepts_root_prefixes_and_url_schemes() { + // Java uses these root prefixes for its fallback clients and accepts + // credentials scoped directly to them. + assert!(CloudRefresh::for_location("s3").is_some()); + assert!(CloudRefresh::for_location("gs").is_some()); assert!(CloudRefresh::for_location("s3://b/k").is_some()); assert!(CloudRefresh::for_location("S3://b/k").is_some()); assert!(CloudRefresh::for_location("s3a://b/k").is_some()); @@ -669,13 +873,13 @@ mod tests { "1500".to_string(), ); let credential = parse_s3_credential(&config, Some("s3://bucket".to_string())).unwrap(); - assert_eq!(credential.prefix.as_deref(), Some("s3://bucket")); + assert_eq!(credential.prefix(), Some("s3://bucket")); assert_eq!( - credential.expires_at, + credential.expires_at(), Some(UNIX_EPOCH + Duration::from_millis(1500)) ); - match credential.kind { - StorageCredentialKind::S3(s3) => assert_eq!(s3.session_token.as_deref(), Some("TOK")), + match credential.kind() { + StorageCredentialKind::S3(s3) => assert_eq!(s3.session_token(), Some("TOK")), other => panic!("expected S3, got {other:?}"), } } @@ -689,13 +893,13 @@ mod tests { config.insert(GCS_TOKEN_EXPIRES_AT.to_string(), "2000".to_string()); let credential = parse_gcs_credential(&config, Some("gs://bucket".to_string())).unwrap(); - assert_eq!(credential.prefix.as_deref(), Some("gs://bucket")); - match &credential.kind { - StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token, "ya29.token"), + assert_eq!(credential.prefix(), Some("gs://bucket")); + match credential.kind() { + StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token(), "ya29.token"), other => panic!("expected GCS, got {other:?}"), } assert_eq!( - credential.expires_at, + credential.expires_at(), Some(UNIX_EPOCH + Duration::from_millis(2000)) ); } @@ -730,7 +934,7 @@ mod tests { expires_at - REFRESH_BUFFER ); - for _ in 0..100 { + for _ in 0..16 { let refresh_at = prefetch_time(expires_at, true); assert!(refresh_at >= expires_at - REFRESH_BUFFER); assert!(refresh_at < expires_at - MIN_REFRESH_BUFFER); @@ -752,12 +956,14 @@ mod tests { #[test] fn failure_backoff_is_jittered_and_capped() { - for failures in 1_u32..=40 { - let ceiling = INITIAL_FAILURE_BACKOFF - .checked_mul(1_u32 << failures.saturating_sub(1).min(5)) - .unwrap_or(MAX_FAILURE_BACKOFF) - .min(MAX_FAILURE_BACKOFF); - for _ in 0..100 { + for (failures, ceiling) in [ + (1, Duration::from_secs(1)), + (2, Duration::from_secs(2)), + (5, Duration::from_secs(16)), + (6, MAX_FAILURE_BACKOFF), + (u32::MAX, MAX_FAILURE_BACKOFF), + ] { + for _ in 0..16 { let backoff = failure_backoff(failures); assert!(backoff >= ceiling / 2); assert!(backoff <= ceiling); @@ -774,10 +980,7 @@ mod tests { ]; let got = longest_prefix_match(&entries, "s3://bucket/warehouse/db/t/f").unwrap(); assert_eq!(s3_access_key_id(&got.credential), "narrow"); - assert_eq!( - got.credential.prefix.as_deref(), - Some("s3://bucket/warehouse/db") - ); + assert_eq!(got.credential.prefix(), Some("s3://bucket/warehouse/db")); assert!(longest_prefix_match(&entries, "s3://other/x").is_none()); let fresh = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)); @@ -793,25 +996,34 @@ mod tests { #[test] fn provider_support_tracks_configured_cloud_endpoints() { - let config = RestCatalogConfig::builder() - .uri("http://cat".to_string()) - .build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); - let props = HashMap::from([( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/creds".to_string(), - )]); + let client = test_client("http://cat"); + let props = aws_refresh_props("/v1/creds"); // The provider is configured independently of the table metadata scheme, // but advertises support only for clouds whose endpoint is present. - let provider = build_vended_credential_provider(client.clone(), "http://cat", &props) + let provider = build_vended_credential_provider(client.clone(), "http://cat", &props, None) .unwrap() .unwrap(); assert!(provider.supports_path("s3://b/k")); assert!(!provider.supports_path("abfss://fs@acct.dfs.core.windows.net/k")); + let enabled = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/creds".to_string(), + ), + ( + CloudRefresh::AWS.enabled_key.to_string(), + "True".to_string(), + ), + ]); + assert!( + build_vended_credential_provider(client.clone(), "http://cat", &enabled, None) + .unwrap() + .is_some() + ); // No endpoint advertised. assert!( - build_vended_credential_provider(client.clone(), "http://cat", &HashMap::new()) + build_vended_credential_provider(client.clone(), "http://cat", &HashMap::new(), None) .unwrap() .is_none() ); @@ -827,29 +1039,27 @@ mod tests { ), ]); assert!( - build_vended_credential_provider(client, "http://cat", &disabled) + build_vended_credential_provider(client, "http://cat", &disabled, None) .unwrap() .is_none() ); // Java's `Strings.isNullOrEmpty` check treats an empty endpoint as absent. let empty = HashMap::from([(CloudRefresh::AWS.endpoint_key.to_string(), String::new())]); - let config = RestCatalogConfig::builder() - .uri("http://cat".to_string()) - .build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); + let client = test_client("http://cat"); assert!( - build_vended_credential_provider(client, "http://cat", &empty) + build_vended_credential_provider(client, "http://cat", &empty, None) .unwrap() .is_none() ); } #[tokio::test] - async fn fetches_and_selects_vended_credential() { + async fn refresh_includes_scan_plan_id() { let mut server = Server::new_async().await; - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); - let body = format!( - r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), ); let mock = server .mock("GET", "/v1/credentials") @@ -863,39 +1073,122 @@ mod tests { .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); // No static creds -> no seed -> first load fetches from the endpoint. - let props = HashMap::from([ - ( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/credentials".to_string(), - ), - ( - REST_CATALOG_PROP_SCAN_PLAN_ID.to_string(), - "scan-plan-1".to_string(), - ), - ]); + let mut props = aws_refresh_props("/v1/credentials"); + props.insert( + REST_CATALOG_PROP_SCAN_PLAN_ID.to_string(), + "scan-plan-1".to_string(), + ); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .expect("provider construction should succeed") - .expect("provider should be built"); + let provider = test_provider(&server.url(), &props); let credential = provider .load_credential("s3://bucket/warehouse/f") .await .unwrap(); - assert_eq!(credential.prefix.as_deref(), Some("s3://bucket")); + assert_eq!(credential.prefix(), Some("s3://bucket")); - match credential.kind { + match credential.kind() { StorageCredentialKind::S3(s3) => { - assert_eq!(s3.access_key_id, "AK"); - assert_eq!(s3.session_token.as_deref(), Some("TOK")); + assert_eq!(s3.access_key_id(), "AK"); + assert_eq!(s3.session_token(), Some("TOK")); } other => panic!("expected S3, got {other:?}"), } mock.assert_async().await; } + #[tokio::test] + async fn successful_refresh_caches_entries_before_selecting_path() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket/table-a", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let provider = test_provider(&server.url(), &props); + + assert!( + provider + .load_credential("s3://bucket/table-b/file") + .await + .is_err() + ); + // A successful response with no matching credential is negatively + // cached instead of immediately hitting the endpoint again. + assert!( + provider + .load_credential("s3://bucket/table-b/file") + .await + .is_err() + ); + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/table-a/file") + .await + .unwrap() + ), + "AK" + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn malformed_entry_does_not_discard_valid_entries() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/invalid","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket/valid","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let provider = test_provider(&server.url(), &props); + + assert!( + provider + .load_credential("s3://bucket/invalid/file") + .await + .is_err() + ); + assert!( + provider + .load_credential("s3://bucket/invalid/file") + .await + .is_err() + ); + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/valid/file") + .await + .unwrap() + ), + "AK" + ); + mock.assert_async().await; + } + #[tokio::test] async fn fresh_seed_is_served_without_fetching() { let mut server = Server::new_async().await; @@ -906,26 +1199,16 @@ mod tests { .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); - let props = HashMap::from([ - ( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/credentials".to_string(), - ), - (S3_ACCESS_KEY_ID.to_string(), "SEED_AK".to_string()), - (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), - (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), - (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), expires), - ]); + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(3600), + ); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .expect("provider construction should succeed") - .expect("provider should be built"); + let provider = test_provider(&server.url(), &props); let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); - assert_eq!(credential.prefix, None); + assert_eq!(credential.prefix(), None); assert_eq!(s3_access_key_id(&credential), "SEED_AK"); mock.assert_async().await; } @@ -941,27 +1224,17 @@ mod tests { .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); // Seeded credential is within the refresh buffer but not yet expired. - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(60)); - let props = HashMap::from([ - ( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/credentials".to_string(), - ), - (S3_ACCESS_KEY_ID.to_string(), "SEED_AK".to_string()), - (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), - (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), - (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), expires), - ]); + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(60), + ); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .expect("provider construction should succeed") - .expect("provider should be built"); + let provider = test_provider(&server.url(), &props); // The first refresh fails, but the still-valid seed is served. Immediate // follow-up operations stay inside the first jittered backoff window. - for _ in 0..3 { + for _ in 0..2 { let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); assert_eq!(s3_access_key_id(&credential), "SEED_AK"); } @@ -969,11 +1242,172 @@ mod tests { } #[tokio::test] - async fn concurrent_prefetch_serves_unexpired_credential_without_waiting() { + async fn empty_refresh_preserves_unexpired_seed() { let mut server = Server::new_async().await; - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(r#"{"storage-credentials":[]}"#) + .create_async() + .await; + + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(60), + ); + let provider = test_provider(&server.url(), &props); + + for _ in 0..2 { + let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); + assert_eq!(s3_access_key_id(&credential), "SEED_AK"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn malformed_specific_refresh_prefers_fallback_over_valid_broader_entry() { + let mut server = Server::new_async().await; + let refreshed_expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/requested","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket","config":{{"s3.access-key-id":"NEW_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{refreshed_expires}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = RestVendedCredentialProvider::new(test_client(&server.url()), None, vec![ + ConfiguredCloud { + cloud: &CloudRefresh::AWS, + endpoint: format!("{}/v1/credentials", server.url()), + cache: Mutex::new(CacheState { + entries: vec![cached_s3( + "s3://bucket/requested", + "SEED_AK", + Some(SystemTime::now() + Duration::from_secs(60)), + )], + consecutive_failures: 0, + retry_not_before: None, + }), + refresh: Mutex::new(()), + }, + ]); + + let fallback = provider + .load_credential("s3://bucket/requested/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "SEED_AK"); + + // The malformed specific replacement is a failed refresh for this path. + // Its still-valid fallback wins over the broader fetched credential and + // is backed off, so an immediate retry does not fetch again. + let fallback = provider + .load_credential("s3://bucket/requested/other-file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "SEED_AK"); + + let refreshed = provider + .load_credential("s3://bucket/other/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&refreshed), "NEW_AK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn invalid_refresh_preserves_other_prefix_fallbacks() { + let mut server = Server::new_async().await; + let refreshed_expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let expired = epoch_millis(SystemTime::now() - Duration::from_secs(60)); let body = format!( - r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"NEW_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/table-a","config":{{"s3.access-key-id":"NEW_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{refreshed_expires}"}}}}, + {{"prefix":"s3://bucket/table-b","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket/table-c","config":{{"s3.access-key-id":"EXPIRED_CK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expired}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = RestVendedCredentialProvider::new(test_client(&server.url()), None, vec![ + ConfiguredCloud { + cloud: &CloudRefresh::AWS, + endpoint: format!("{}/v1/credentials", server.url()), + cache: Mutex::new(CacheState { + entries: vec![ + cached_s3( + "s3://bucket/table-a", + "OLD_AK", + Some(SystemTime::now() + Duration::from_secs(60)), + ), + cached_s3( + "s3://bucket/table-b", + "OLDER_BK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + cached_s3( + "s3://bucket/table-b", + "LATEST_BK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + cached_s3( + "s3://bucket/table-c", + "VALID_CK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + ], + consecutive_failures: 0, + retry_not_before: None, + }), + refresh: Mutex::new(()), + }, + ]); + + let refreshed = provider + .load_credential("s3://bucket/table-a/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&refreshed), "NEW_AK"); + + let fallback = provider + .load_credential("s3://bucket/table-b/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "LATEST_BK"); + + let fallback = provider + .load_credential("s3://bucket/table-c/file") + .await + .unwrap(); + assert_eq!(s3_access_key_id(&fallback), "VALID_CK"); + mock.assert_async().await; + } + + #[tokio::test] + async fn concurrent_prefetch_serves_unexpired_credential_without_waiting() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "NEW_AK", + SystemTime::now() + Duration::from_secs(3600), ); let (started_tx, started_rx) = mpsc::channel(); let release = Arc::new(Barrier::new(2)); @@ -991,22 +1425,12 @@ mod tests { .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); - let seed_expires = epoch_millis(SystemTime::now() + Duration::from_secs(60)); - let props = HashMap::from([ - ( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/credentials".to_string(), - ), - (S3_ACCESS_KEY_ID.to_string(), "SEED_AK".to_string()), - (S3_SECRET_ACCESS_KEY.to_string(), "SEED_SK".to_string()), - (S3_SESSION_TOKEN.to_string(), "SEED_TOK".to_string()), - (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), seed_expires), - ]); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .unwrap() - .unwrap(); + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "SEED_AK", + SystemTime::now() + Duration::from_secs(60), + ); + let provider = test_provider(&server.url(), &props); let first_provider = Arc::clone(&provider); let first = @@ -1041,9 +1465,10 @@ mod tests { #[tokio::test] async fn refresh_uses_table_scoped_token() { let mut server = Server::new_async().await; - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); - let body = format!( - r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), ); let mock = server .mock("GET", "/v1/credentials") @@ -1054,18 +1479,28 @@ mod tests { .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); + let config = RestCatalogConfig::builder() + .uri(server.url()) + .props(HashMap::from([( + "token".to_string(), + "catalog-token".to_string(), + )])) + .build(); let client = Arc::new(HttpClient::new(&config).unwrap()); let props = HashMap::from([ ( CloudRefresh::AWS.endpoint_key.to_string(), "/v1/credentials".to_string(), ), - ("token".to_string(), "table-token".to_string()), + // Local properties win in the FileIO configuration merge, but must + // not mask auth returned by the table endpoint. + ("token".to_string(), "user-token".to_string()), ]); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .unwrap() - .unwrap(); + let table_auth = HashMap::from([("token".to_string(), "table-token".to_string())]); + let provider = + build_vended_credential_provider(client, &server.url(), &props, Some(&table_auth)) + .unwrap() + .unwrap(); provider.load_credential("s3://bucket/x/f").await.unwrap(); mock.assert_async().await; @@ -1075,8 +1510,10 @@ mod tests { async fn one_provider_refreshes_multiple_clouds() { let mut server = Server::new_async().await; let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); - let aws_body = format!( - r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + let aws_body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), ); let gcp_body = format!( r#"{{"storage-credentials":[{{"prefix":"gs://bucket","config":{{"gcs.oauth2.token":"GCS","gcs.oauth2.token-expires-at":"{expires}"}}}}]}}"# @@ -1096,8 +1533,6 @@ mod tests { .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); let props = HashMap::from([ ( CloudRefresh::AWS.endpoint_key.to_string(), @@ -1108,9 +1543,7 @@ mod tests { "/v1/gcp-credentials".to_string(), ), ]); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .unwrap() - .unwrap(); + let provider = test_provider(&server.url(), &props); assert!(provider.supports_path("s3://bucket/x")); assert!(provider.supports_path("gs://bucket/x")); @@ -1122,9 +1555,9 @@ mod tests { .load_credential("gs://bucket/x") .await .unwrap() - .kind + .kind() { - StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token, "GCS"), + StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token(), "GCS"), other => panic!("expected GCS credential, got {other:?}"), } aws_mock.assert_async().await; @@ -1136,28 +1569,20 @@ mod tests { let mut server = Server::new_async().await; let mock = server .mock("GET", Matcher::Any) + .expect(1) .with_status(500) .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); - let expired = epoch_millis(SystemTime::now() - Duration::from_secs(60)); - let props = HashMap::from([ - ( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/credentials".to_string(), - ), - (S3_ACCESS_KEY_ID.to_string(), "EXPIRED_AK".to_string()), - (S3_SECRET_ACCESS_KEY.to_string(), "EXPIRED_SK".to_string()), - (S3_SESSION_TOKEN.to_string(), "EXPIRED_TOK".to_string()), - (S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), expired), - ]); + let props = with_s3_seed( + aws_refresh_props("/v1/credentials"), + "EXPIRED_AK", + SystemTime::now() - Duration::from_secs(60), + ); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .expect("provider construction should succeed") - .expect("provider should be built"); + let provider = test_provider(&server.url(), &props); + assert!(provider.load_credential("s3://bucket/x/f").await.is_err()); assert!(provider.load_credential("s3://bucket/x/f").await.is_err()); mock.assert_async().await; } @@ -1168,33 +1593,27 @@ mod tests { // The vended TTL (60s) is shorter than REFRESH_BUFFER, so every operation // is eligible for refresh. This matches AWS CachedSupplier: successful // results do not receive failure backoff. - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(60)); - let body = format!( - r#"{{"storage-credentials":[{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}]}}"# + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(60), ); let mock = server .mock("GET", "/v1/credentials") - .expect(3) + .expect(2) .with_status(200) .with_header("content-type", "application/json") .with_body(body) .create_async() .await; - let config = RestCatalogConfig::builder().uri(server.url()).build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); // No static creds -> no seed -> each sequential load fetches because the // returned credential is already inside its prefetch window. - let props = HashMap::from([( - CloudRefresh::AWS.endpoint_key.to_string(), - "/v1/credentials".to_string(), - )]); + let props = aws_refresh_props("/v1/credentials"); - let provider = build_vended_credential_provider(client, &server.url(), &props) - .expect("provider construction should succeed") - .expect("provider should be built"); + let provider = test_provider(&server.url(), &props); - for _ in 0..3 { + for _ in 0..2 { let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); assert_eq!(s3_access_key_id(&credential), "AK"); } diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index 93eaa414cf..698dcb1a68 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -784,7 +784,10 @@ pub fn iceberg::io::GcsConfig::serialize<__S>(&self, __serializer: __S) -> core: impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg::io::GcsCredential -pub iceberg::io::GcsCredential::token: alloc::string::String +impl iceberg::io::GcsCredential +pub fn iceberg::io::GcsCredential::into_token(self) -> alloc::string::String +pub fn iceberg::io::GcsCredential::new(token: impl core::convert::Into) -> Self +pub fn iceberg::io::GcsCredential::token(&self) -> &str impl core::clone::Clone for iceberg::io::GcsCredential pub fn iceberg::io::GcsCredential::clone(&self) -> iceberg::io::GcsCredential impl core::fmt::Debug for iceberg::io::GcsCredential @@ -972,9 +975,12 @@ pub fn iceberg::io::S3Config::serialize<__S>(&self, __serializer: __S) -> core:: impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::S3Config pub fn iceberg::io::S3Config::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg::io::S3Credential -pub iceberg::io::S3Credential::access_key_id: alloc::string::String -pub iceberg::io::S3Credential::secret_access_key: alloc::string::String -pub iceberg::io::S3Credential::session_token: core::option::Option +impl iceberg::io::S3Credential +pub fn iceberg::io::S3Credential::access_key_id(&self) -> &str +pub fn iceberg::io::S3Credential::into_parts(self) -> (alloc::string::String, alloc::string::String, core::option::Option) +pub fn iceberg::io::S3Credential::new(access_key_id: impl core::convert::Into, secret_access_key: impl core::convert::Into, session_token: core::option::Option) -> Self +pub fn iceberg::io::S3Credential::secret_access_key(&self) -> &str +pub fn iceberg::io::S3Credential::session_token(&self) -> core::option::Option<&str> impl core::clone::Clone for iceberg::io::S3Credential pub fn iceberg::io::S3Credential::clone(&self) -> iceberg::io::S3Credential impl core::fmt::Debug for iceberg::io::S3Credential @@ -1017,9 +1023,14 @@ pub fn iceberg::io::StorageConfig::serialize<__S>(&self, __serializer: __S) -> c impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg::io::StorageCredential -pub iceberg::io::StorageCredential::expires_at: core::option::Option -pub iceberg::io::StorageCredential::kind: iceberg::io::StorageCredentialKind -pub iceberg::io::StorageCredential::prefix: core::option::Option +impl iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::expires_at(&self) -> core::option::Option +pub fn iceberg::io::StorageCredential::into_kind(self) -> iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredential::kind(&self) -> &iceberg::io::StorageCredentialKind +pub fn iceberg::io::StorageCredential::new(kind: iceberg::io::StorageCredentialKind) -> Self +pub fn iceberg::io::StorageCredential::prefix(&self) -> core::option::Option<&str> +pub fn iceberg::io::StorageCredential::with_expiration(self, expires_at: std::time::SystemTime) -> Self +pub fn iceberg::io::StorageCredential::with_prefix(self, prefix: impl core::convert::Into) -> Self impl core::clone::Clone for iceberg::io::StorageCredential pub fn iceberg::io::StorageCredential::clone(&self) -> iceberg::io::StorageCredential impl core::fmt::Debug for iceberg::io::StorageCredential diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index dd812a2d45..8ce1780c44 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -198,12 +198,55 @@ pub trait StorageCredentialProvider: Debug + Send + Sync { pub struct StorageCredential { /// Storage-location prefix this credential is scoped to. `None` represents a /// credential without a declared scope, sourced from flat storage properties. - pub prefix: Option, + prefix: Option, /// The backend-specific credential material. - pub kind: StorageCredentialKind, + kind: StorageCredentialKind, /// When the credential expires, if known. `None` means non-expiring and /// backends treat such a credential as always valid and never refresh it. - pub expires_at: Option, + expires_at: Option, +} + +impl StorageCredential { + /// Create a storage credential with no declared scope or expiration. + pub fn new(kind: StorageCredentialKind) -> Self { + Self { + prefix: None, + kind, + expires_at: None, + } + } + + /// Set the storage-location prefix this credential is scoped to. + pub fn with_prefix(mut self, prefix: impl Into) -> Self { + self.prefix = Some(prefix.into()); + self + } + + /// Set when this credential expires. + pub fn with_expiration(mut self, expires_at: SystemTime) -> Self { + self.expires_at = Some(expires_at); + self + } + + /// Return the storage-location prefix this credential is scoped to. + pub fn prefix(&self) -> Option<&str> { + self.prefix.as_deref() + } + + /// Return the backend-specific credential material. + pub fn kind(&self) -> &StorageCredentialKind { + &self.kind + } + + /// Consume this credential and return its backend-specific material. + pub fn into_kind(self) -> StorageCredentialKind { + self.kind + } + + /// Return when this credential expires. + pub fn expires_at(&self) -> Option { + self.expires_at + } } /// Backend-specific credential material. @@ -219,11 +262,51 @@ pub enum StorageCredentialKind { #[derive(Clone)] pub struct S3Credential { /// AWS access key ID. - pub access_key_id: String, + access_key_id: String, /// AWS secret access key. - pub secret_access_key: String, + secret_access_key: String, /// AWS session token, set for temporary (STS/vended) credentials. - pub session_token: Option, + session_token: Option, +} + +impl S3Credential { + /// Create temporary Amazon S3 credentials. + pub fn new( + access_key_id: impl Into, + secret_access_key: impl Into, + session_token: Option, + ) -> Self { + Self { + access_key_id: access_key_id.into(), + secret_access_key: secret_access_key.into(), + session_token, + } + } + + /// Return the AWS access key ID. + pub fn access_key_id(&self) -> &str { + &self.access_key_id + } + + /// Return the AWS secret access key. + pub fn secret_access_key(&self) -> &str { + &self.secret_access_key + } + + /// Return the AWS session token, if present. + pub fn session_token(&self) -> Option<&str> { + self.session_token.as_deref() + } + + /// Consume these credentials and return their component values. + pub fn into_parts(self) -> (String, String, Option) { + let Self { + access_key_id, + secret_access_key, + session_token, + } = self; + (access_key_id, secret_access_key, session_token) + } } impl Debug for S3Credential { @@ -236,7 +319,26 @@ impl Debug for S3Credential { #[derive(Clone)] pub struct GcsCredential { /// OAuth2 bearer token used to access GCS. - pub token: String, + token: String, +} + +impl GcsCredential { + /// Create a Google Cloud Storage credential. + pub fn new(token: impl Into) -> Self { + Self { + token: token.into(), + } + } + + /// Return the OAuth2 bearer token used to access GCS. + pub fn token(&self) -> &str { + &self.token + } + + /// Consume this credential and return its OAuth2 bearer token. + pub fn into_token(self) -> String { + self.token + } } impl Debug for GcsCredential { diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index 311bd1ac88..20cf1919a1 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -30,6 +30,7 @@ use reqsign_core::{Context, Error as ReqsignError, ProvideCredential, Result as use reqsign_google::{Credential as GoogleCredential, Token as GoogleToken}; use url::Url; +use crate::DynamicCredentialScope; use crate::utils::{ from_opendal_error, is_truthy, system_time_to_timestamp, validate_credential_prefix, }; @@ -82,6 +83,7 @@ pub(crate) fn gcs_config_build( cfg: &GcsConfig, credential_provider: &Option>, path: &str, + credential_scope: Option<&DynamicCredentialScope>, ) -> Result { let url = Url::parse(path)?; if !matches!(url.scheme(), "gs" | "gcs") { @@ -128,6 +130,7 @@ pub(crate) fn gcs_config_build( builder = builder.credential_provider(VendedGcsCredentialProvider::new( Arc::clone(provider), path.to_string(), + credential_scope.cloned(), )); } @@ -142,19 +145,31 @@ struct VendedGcsCredentialProvider { /// Absolute path this operator serves: handed back to the provider so it can /// select the vended credential whose prefix best matches the location. path: String, + /// Exact scope selected when a bulk-delete operator was created. Ordinary + /// operators are unbound so they can follow the provider's best match. + credential_scope: Option, } impl std::fmt::Debug for VendedGcsCredentialProvider { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("VendedGcsCredentialProvider") .field("path", &self.path) + .field("credential_scope", &self.credential_scope) .finish_non_exhaustive() } } impl VendedGcsCredentialProvider { - fn new(provider: Arc, path: String) -> Self { - Self { provider, path } + fn new( + provider: Arc, + path: String, + credential_scope: Option, + ) -> Self { + Self { + provider, + path, + credential_scope, + } } } @@ -174,16 +189,20 @@ impl ProvideCredential for VendedGcsCredentialProvider { .with_source(e) })?; - validate_credential_prefix(&self.path, credential.prefix.as_deref())?; + validate_credential_prefix( + &self.path, + credential.prefix(), + self.credential_scope.as_ref(), + )?; let expires_at = credential - .expires_at + .expires_at() .map(system_time_to_timestamp) .transpose()?; - match credential.kind { + match credential.into_kind() { StorageCredentialKind::Gcs(gcs) => { Ok(Some(GoogleCredential::with_token(GoogleToken { - access_token: gcs.token, + access_token: gcs.into_token(), expires_at, }))) } @@ -193,3 +212,92 @@ impl ProvideCredential for VendedGcsCredentialProvider { } } } + +#[cfg(test)] +mod tests { + use async_trait::async_trait; + use iceberg::io::{GcsCredential, StorageCredential}; + + use super::*; + + #[derive(Debug)] + struct FixedCredentialProvider(StorageCredential); + + #[async_trait] + impl StorageCredentialProvider for FixedCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(self.0.clone()) + } + } + + fn credential(prefix: Option<&str>) -> StorageCredential { + let credential = StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new( + "access-token", + ))); + match prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, + } + } + + fn provider( + path: &str, + returned_prefix: Option<&str>, + credential_scope: Option, + ) -> VendedGcsCredentialProvider { + VendedGcsCredentialProvider::new( + Arc::new(FixedCredentialProvider(credential(returned_prefix))), + path.to_string(), + credential_scope, + ) + } + + #[tokio::test] + async fn vended_provider_enforces_bound_credential_scope() { + let path = "gs://bucket/table/file.parquet"; + let table_prefix = "gs://bucket/table"; + + assert!( + provider( + path, + Some(table_prefix), + Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), + ) + .provide_credential(&Context::default()) + .await + .is_ok() + ); + assert!( + provider( + path, + Some("gs://bucket"), + Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), + ) + .provide_credential(&Context::default()) + .await + .is_err() + ); + assert!( + provider(path, Some("gs://bucket"), None) + .provide_credential(&Context::default()) + .await + .is_ok() + ); + assert!( + provider(path, None, Some(DynamicCredentialScope::Unscoped)) + .provide_credential(&Context::default()) + .await + .is_ok() + ); + assert!( + provider( + path, + Some(table_prefix), + Some(DynamicCredentialScope::Unscoped), + ) + .provide_credential(&Context::default()) + .await + .is_err() + ); + } +} diff --git a/crates/storage/opendal/src/lib.rs b/crates/storage/opendal/src/lib.rs index 0f301dfc91..e58b94575e 100644 --- a/crates/storage/opendal/src/lib.rs +++ b/crates/storage/opendal/src/lib.rs @@ -145,6 +145,21 @@ impl StorageFactory for OpenDalStorageFactory { config: &StorageConfig, credential_provider: Option>, ) -> Result> { + #[allow(unreachable_patterns)] + let supports_credential_provider = match self { + #[cfg(feature = "opendal-s3")] + OpenDalStorageFactory::S3 { .. } => true, + #[cfg(feature = "opendal-gcs")] + OpenDalStorageFactory::Gcs => true, + _ => false, + }; + if credential_provider.is_some() && !supports_credential_provider { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + "OpenDAL storage factory does not support refreshable credentials for this backend", + )); + } + match self { #[cfg(feature = "opendal-memory")] OpenDalStorageFactory::Memory => { @@ -262,6 +277,33 @@ pub enum OpenDalStorage { }, } +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +enum DeleteCredentialScope { + Static, + Dynamic(DynamicCredentialScope), +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +enum DynamicCredentialScope { + Unscoped, + Prefix(String), +} + +impl DynamicCredentialScope { + fn from_prefix(prefix: Option<&str>) -> Self { + match prefix { + Some(prefix) => Self::Prefix(prefix.to_string()), + None => Self::Unscoped, + } + } +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct DeleteBatchKey { + storage: String, + credential_scope: DeleteCredentialScope, +} + impl OpenDalStorage { /// Creates operator from path. /// @@ -279,6 +321,17 @@ impl OpenDalStorage { pub(crate) fn create_operator<'a>( &self, path: &'a impl AsRef, + ) -> Result<(Operator, &'a str)> { + self.create_operator_with_scope(path, None) + } + + /// Creates an operator, optionally binding dynamic credentials to the exact + /// scope used to group a bulk-delete batch. + #[allow(unreachable_code, unused_variables)] + fn create_operator_with_scope<'a>( + &self, + path: &'a impl AsRef, + credential_scope: Option<&DynamicCredentialScope>, ) -> Result<(Operator, &'a str)> { let path = path.as_ref(); let (operator, relative_path): (Operator, &str) = match self { @@ -310,6 +363,7 @@ impl OpenDalStorage { customized_credential_load, credential_provider, path, + credential_scope, )?; let op_info = op.info(); @@ -336,7 +390,8 @@ impl OpenDalStorage { config, credential_provider, } => { - let operator = gcs_config_build(config, credential_provider, path)?; + let operator = + gcs_config_build(config, credential_provider, path, credential_scope)?; let url = url::Url::parse(path).map_err(|e| { Error::new( ErrorKind::DataInvalid, @@ -415,25 +470,58 @@ impl OpenDalStorage { } } - /// Returns whether `path` uses a path-scoped dynamic credential provider. - /// - /// Such paths cannot share an OpenDAL deleter unless the provider can expose - /// the credential scope that applies to each path. Process them independently - /// so bulk deletion remains bounded without crossing credential boundaries. - fn uses_dynamic_credentials(&self, path: &str) -> bool { - match self { + /// Return the dynamic credential provider that serves `path`, if any. + fn credential_provider_for_path( + &self, + path: &str, + ) -> Option<&Arc> { + let provider: &Arc = (match self { #[cfg(feature = "opendal-s3")] OpenDalStorage::S3 { + customized_credential_load: None, credential_provider: Some(provider), .. - } => provider.supports_path(path), + } => Some(provider), #[cfg(feature = "opendal-gcs")] OpenDalStorage::Gcs { credential_provider: Some(provider), .. - } => provider.supports_path(path), - _ => false, - } + } => Some(provider), + _ => None, + })?; + provider.supports_path(path).then_some(provider) + } + + /// Returns a key that keeps bulk deletes within one operator and credential + /// scope. Loading the credential is normally a cache hit and avoids rebuilding + /// an operator for every path while preventing a batch from crossing prefixes. + async fn delete_batch_key_for_path(&self, path: &str) -> Result { + let credential_scope = match self.credential_provider_for_path(path) { + Some(provider) => { + let credential = provider.load_credential(path).await?; + if credential + .prefix() + .is_some_and(|prefix| prefix.is_empty() || !path.starts_with(prefix)) + { + return Err(Error::new( + ErrorKind::DataInvalid, + format!( + "vended credential prefix {:?} does not cover storage location {path:?}", + credential.prefix() + ), + )); + } + DeleteCredentialScope::Dynamic(DynamicCredentialScope::from_prefix( + credential.prefix(), + )) + } + None => DeleteCredentialScope::Static, + }; + + Ok(DeleteBatchKey { + storage: self.batch_key_for_path(path), + credential_scope, + }) } /// Extracts the relative path from an absolute path without building an operator. @@ -608,23 +696,21 @@ impl Storage for OpenDalStorage { } async fn delete_stream(&self, mut paths: BoxStream<'static, String>) -> Result<()> { - let mut deleters: HashMap = HashMap::new(); + let mut deleters: HashMap = HashMap::new(); while let Some(path) = paths.next().await { - if self.uses_dynamic_credentials(&path) { - let (op, relative_path) = self.create_operator(&path)?; - op.delete(relative_path).await.map_err(from_opendal_error)?; - continue; - } + let batch_key = self.delete_batch_key_for_path(&path).await?; - let bucket = self.batch_key_for_path(&path); - - let (relative_path, deleter) = match deleters.entry(bucket) { + let (relative_path, deleter) = match deleters.entry(batch_key) { Entry::Occupied(entry) => { (self.relativize_path(&path)?.to_string(), entry.into_mut()) } Entry::Vacant(entry) => { - let (op, rel) = self.create_operator(&path)?; + let credential_scope = match &entry.key().credential_scope { + DeleteCredentialScope::Static => None, + DeleteCredentialScope::Dynamic(scope) => Some(scope), + }; + let (op, rel) = self.create_operator_with_scope(&path, credential_scope)?; let rel = rel.to_string(); let deleter = op.deleter().await.map_err(from_opendal_error)?; (rel, entry.insert(deleter)) @@ -702,8 +788,18 @@ mod tests { #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] #[async_trait] impl StorageCredentialProvider for AlwaysSupportedCredentialProvider { - async fn load_credential(&self, _path: &str) -> Result { - unreachable!("batch key selection must not load a credential") + async fn load_credential(&self, path: &str) -> Result { + let prefix = if path.contains("/table-a/") { + "s3://bucket/table-a" + } else { + "s3://bucket/table-b" + }; + Ok( + iceberg::io::StorageCredential::new(iceberg::io::StorageCredentialKind::S3( + iceberg::io::S3Credential::new("access-key", "secret-key", None), + )) + .with_prefix(prefix), + ) } } @@ -714,6 +810,22 @@ mod tests { assert_eq!(op.info().scheme().to_string(), "memory"); } + #[cfg(all( + feature = "opendal-memory", + any(feature = "opendal-s3", feature = "opendal-gcs") + ))] + #[test] + fn test_factory_rejects_credentials_for_unsupported_backend() { + let error = OpenDalStorageFactory::Memory + .build_with_credentials( + &StorageConfig::new(), + Some(Arc::new(AlwaysSupportedCredentialProvider)), + ) + .expect_err("memory must reject a credential provider"); + + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + } + #[cfg(feature = "opendal-memory")] #[test] fn test_relativize_path_memory() { @@ -770,20 +882,53 @@ mod tests { } #[cfg(feature = "opendal-s3")] - #[test] - fn test_dynamic_credentials_use_isolated_deletes() { + #[tokio::test] + async fn test_dynamic_credentials_batch_by_prefix() { let storage = OpenDalStorage::S3 { config: Arc::new(S3Config::default()), customized_credential_load: None, credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), }; let first = "s3://bucket/table-a/file.parquet"; - let second = "s3://bucket/table-b/file.parquet"; + let same_scope = "s3://bucket/table-a/other.parquet"; + let other_scope = "s3://bucket/table-b/file.parquet"; - assert!(storage.uses_dynamic_credentials(first)); - assert!(storage.uses_dynamic_credentials(second)); - assert_eq!(storage.batch_key_for_path(first), "bucket"); - assert_eq!(storage.batch_key_for_path(second), "bucket"); + let first_key = storage.delete_batch_key_for_path(first).await.unwrap(); + assert_eq!( + first_key.credential_scope, + DeleteCredentialScope::Dynamic(DynamicCredentialScope::Prefix( + "s3://bucket/table-a".to_string() + )) + ); + assert_eq!( + first_key, + storage.delete_batch_key_for_path(same_scope).await.unwrap() + ); + assert_ne!( + first_key, + storage + .delete_batch_key_for_path(other_scope) + .await + .unwrap() + ); + } + + #[cfg(feature = "opendal-s3")] + #[tokio::test] + async fn test_custom_s3_credential_loader_ignores_dynamic_provider_for_batching() { + let storage = OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: Some(CustomAwsCredentialLoader::new( + reqsign_aws_v4::StaticCredentialProvider::new("access-key", "secret-key"), + )), + credential_provider: Some(Arc::new(AlwaysSupportedCredentialProvider)), + }; + + let key = storage + .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") + .await + .unwrap(); + assert_eq!(key.credential_scope, DeleteCredentialScope::Static); } #[cfg(feature = "opendal-s3")] diff --git a/crates/storage/opendal/src/resolving.rs b/crates/storage/opendal/src/resolving.rs index 0fad716e52..05194bb292 100644 --- a/crates/storage/opendal/src/resolving.rs +++ b/crates/storage/opendal/src/resolving.rs @@ -80,6 +80,17 @@ fn extract_scheme(path: &str) -> Result<&'static str> { parse_scheme(url.scheme()) } +#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +fn supports_dynamic_credentials(scheme: &str) -> bool { + match scheme { + #[cfg(feature = "opendal-s3")] + "s3" => true, + #[cfg(feature = "opendal-gcs")] + "gcs" => true, + _ => false, + } +} + /// Build an [`OpenDalStorage`] variant for the given scheme and config properties. fn build_storage_for_scheme( scheme: &'static str, @@ -200,6 +211,14 @@ impl StorageFactory for OpenDalResolvingStorageFactory { config: &StorageConfig, credential_provider: Option>, ) -> Result> { + #[cfg(not(any(feature = "opendal-s3", feature = "opendal-gcs")))] + if credential_provider.is_some() { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + "OpenDAL resolving storage does not support refreshable credentials because no compatible backend is enabled", + )); + } + Ok(Arc::new(OpenDalResolvingStorage { props: config.props().clone(), storages: RwLock::new(HashMap::new()), @@ -250,6 +269,21 @@ impl OpenDalResolvingStorage { fn resolve(&self, path: &str) -> Result> { let scheme = extract_scheme(path)?; + #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + if self + .credential_provider + .as_ref() + .is_some_and(|provider| provider.supports_path(path)) + && !supports_dynamic_credentials(scheme) + { + return Err(Error::new( + ErrorKind::FeatureUnsupported, + format!( + "OpenDAL resolving storage does not support refreshable credentials for scheme: {scheme}" + ), + )); + } + // Fast path: check read lock first. { let cache = self @@ -358,8 +392,36 @@ impl Storage for OpenDalResolvingStorage { mod tests { use super::*; + #[derive(Debug)] + struct AllPathsCredentialProvider; + + #[async_trait] + impl StorageCredentialProvider for AllPathsCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + unreachable!("unsupported backends must reject the provider before loading") + } + } + + #[cfg(not(any(feature = "opendal-s3", feature = "opendal-gcs")))] + #[test] + fn test_factory_rejects_credentials_without_compatible_backend() { + let error = OpenDalResolvingStorageFactory::new() + .build_with_credentials( + &StorageConfig::new(), + Some(Arc::new(AllPathsCredentialProvider)), + ) + .expect_err("a provider must not be silently discarded"); + + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + } + /// Builds a resolving storage with empty props, suitable for `resolve()` /// calls that don't actually hit any backend. + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] fn empty_resolving_storage() -> OpenDalResolvingStorage { OpenDalResolvingStorage { props: HashMap::new(), @@ -387,6 +449,21 @@ mod tests { assert!(Arc::ptr_eq(&a, &c), "s3 and s3n should share one instance"); } + #[cfg(all( + feature = "opendal-memory", + any(feature = "opendal-s3", feature = "opendal-gcs") + ))] + #[test] + fn test_resolver_rejects_credentials_for_unsupported_backend() { + let mut storage = empty_resolving_storage(); + storage.credential_provider = Some(Arc::new(AllPathsCredentialProvider)); + + let error = storage + .resolve("memory:/key") + .expect_err("memory must reject a credential provider"); + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + } + #[cfg(feature = "opendal-azdls")] #[test] fn test_resolve_azdls_aliases_share_instance() { diff --git a/crates/storage/opendal/src/s3.rs b/crates/storage/opendal/src/s3.rs index e6b825347d..9d0282e007 100644 --- a/crates/storage/opendal/src/s3.rs +++ b/crates/storage/opendal/src/s3.rs @@ -38,6 +38,7 @@ use reqsign_core::{ }; use url::Url; +use crate::DynamicCredentialScope; use crate::utils::{ from_opendal_error, is_truthy, system_time_to_timestamp, validate_credential_prefix, }; @@ -137,6 +138,7 @@ pub(crate) fn s3_config_build( customized_credential_load: &Option, credential_provider: &Option>, path: &str, + credential_scope: Option<&DynamicCredentialScope>, ) -> Result { let url = Url::parse(path)?; let bucket = url.host_str().ok_or_else(|| { @@ -172,6 +174,7 @@ pub(crate) fn s3_config_build( let chain = ProvideCredentialChain::new().push(VendedS3CredentialProvider::new( Arc::clone(provider), path.to_string(), + credential_scope.cloned(), )); builder = builder.credential_provider_chain(chain); } @@ -187,19 +190,31 @@ struct VendedS3CredentialProvider { /// Absolute path this operator serves; handed back to the provider so it can /// select the vended credential whose prefix best matches the location. path: String, + /// Exact scope selected when a bulk-delete operator was created. Ordinary + /// operators are unbound so they can follow the provider's best match. + credential_scope: Option, } impl std::fmt::Debug for VendedS3CredentialProvider { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("VendedS3CredentialProvider") .field("path", &self.path) + .field("credential_scope", &self.credential_scope) .finish_non_exhaustive() } } impl VendedS3CredentialProvider { - fn new(provider: Arc, path: String) -> Self { - Self { provider, path } + fn new( + provider: Arc, + path: String, + credential_scope: Option, + ) -> Self { + Self { + provider, + path, + credential_scope, + } } } @@ -219,19 +234,26 @@ impl ProvideCredential for VendedS3CredentialProvider { .with_source(e) })?; - validate_credential_prefix(&self.path, credential.prefix.as_deref())?; + validate_credential_prefix( + &self.path, + credential.prefix(), + self.credential_scope.as_ref(), + )?; let expires_in = credential - .expires_at + .expires_at() .map(system_time_to_timestamp) .transpose()?; - match credential.kind { - StorageCredentialKind::S3(s3) => Ok(Some(AwsCredential { - access_key_id: s3.access_key_id, - secret_access_key: s3.secret_access_key, - session_token: s3.session_token, - expires_in, - })), + match credential.into_kind() { + StorageCredentialKind::S3(s3) => { + let (access_key_id, secret_access_key, session_token) = s3.into_parts(); + Ok(Some(AwsCredential { + access_key_id, + secret_access_key, + session_token, + expires_in, + })) + } _ => Err(ReqsignError::unexpected( "S3 storage received a non-S3 credential from the provider", )), @@ -270,9 +292,44 @@ impl CustomAwsCredentialLoader { mod tests { use std::collections::HashMap; - use iceberg::io::S3_PATH_STYLE_ACCESS; + use async_trait::async_trait; + use iceberg::io::{S3_PATH_STYLE_ACCESS, S3Credential, StorageCredential}; + + use super::*; + + #[derive(Debug)] + struct FixedCredentialProvider(StorageCredential); + + #[async_trait] + impl StorageCredentialProvider for FixedCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(self.0.clone()) + } + } - use super::s3_config_parse; + fn credential(prefix: Option<&str>) -> StorageCredential { + let credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( + "access-key", + "secret-key", + Some("session-token".to_string()), + ))); + match prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, + } + } + + fn provider( + path: &str, + returned_prefix: Option<&str>, + credential_scope: Option, + ) -> VendedS3CredentialProvider { + VendedS3CredentialProvider::new( + Arc::new(FixedCredentialProvider(credential(returned_prefix))), + path.to_string(), + credential_scope, + ) + } fn parse_with(prop: Option<&str>) -> bool { let mut props = HashMap::new(); @@ -289,4 +346,53 @@ mod tests { assert!(parse_with(Some("false"))); assert!(!parse_with(Some("true"))); } + + #[tokio::test] + async fn vended_provider_enforces_bound_credential_scope() { + let path = "s3://bucket/table/file.parquet"; + let table_prefix = "s3://bucket/table"; + + assert!( + provider( + path, + Some(table_prefix), + Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), + ) + .provide_credential(&Context::default()) + .await + .is_ok() + ); + assert!( + provider( + path, + Some("s3://bucket"), + Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), + ) + .provide_credential(&Context::default()) + .await + .is_err() + ); + assert!( + provider(path, Some("s3://bucket"), None) + .provide_credential(&Context::default()) + .await + .is_ok() + ); + assert!( + provider(path, None, Some(DynamicCredentialScope::Unscoped)) + .provide_credential(&Context::default()) + .await + .is_ok() + ); + assert!( + provider( + path, + Some(table_prefix), + Some(DynamicCredentialScope::Unscoped), + ) + .provide_credential(&Context::default()) + .await + .is_err() + ); + } } diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index c168093fc2..ddb4f08b2d 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -15,6 +15,9 @@ // specific language governing permissions and limitations // under the License. +#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +use crate::DynamicCredentialScope; + pub(crate) fn is_truthy(value: &str) -> bool { ["true", "t", "1", "on"].contains(&value.to_lowercase().as_str()) } @@ -56,6 +59,7 @@ pub(crate) fn system_time_to_timestamp( pub(crate) fn validate_credential_prefix( path: &str, prefix: Option<&str>, + credential_scope: Option<&DynamicCredentialScope>, ) -> reqsign_core::Result<()> { match prefix { Some("") => Err(reqsign_core::Error::unexpected( @@ -64,6 +68,18 @@ pub(crate) fn validate_credential_prefix( Some(prefix) if !path.starts_with(prefix) => Err(reqsign_core::Error::unexpected(format!( "vended credential prefix {prefix:?} does not cover storage location {path:?}" ))), - _ => Ok(()), + _ => match credential_scope { + Some(DynamicCredentialScope::Unscoped) if prefix.is_some() => { + Err(reqsign_core::Error::unexpected( + "vended credential scope changed while a bulk-delete operator was active", + )) + } + Some(DynamicCredentialScope::Prefix(expected)) if prefix != Some(expected.as_str()) => { + Err(reqsign_core::Error::unexpected( + "vended credential scope changed while a bulk-delete operator was active", + )) + } + _ => Ok(()), + }, } } From 3dcec6eb50bf306b32e610ef8c51a74e87b66283 Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Mon, 7 Sep 2026 11:42:43 +0100 Subject: [PATCH 5/7] feat(rest): support refreshing vended ADLS credentials --- Cargo.lock | 222 ++++--- Cargo.toml | 2 +- crates/catalog/rest/src/client.rs | 45 +- crates/catalog/rest/src/credential.rs | 561 ++++++++++++++++-- crates/iceberg/public-api.txt | 14 + crates/iceberg/src/io/storage/config/azdls.rs | 8 + crates/iceberg/src/io/storage/mod.rs | 34 ++ crates/storage/opendal/Cargo.toml | 7 +- crates/storage/opendal/public-api.txt | 13 + crates/storage/opendal/src/azdls.rs | 288 ++++++++- crates/storage/opendal/src/gcs.rs | 11 +- crates/storage/opendal/src/lib.rs | 54 +- crates/storage/opendal/src/resolving.rs | 94 ++- crates/storage/opendal/src/utils.rs | 19 +- 14 files changed, 1161 insertions(+), 211 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index dc3b3cc440..fbfe9c2855 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -561,6 +561,16 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "asyncband" +version = "0.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94a214ba60d6231afd0e805e3c27c45a1626d9debaa5a5061c45a1ea1b2f1ed0" +dependencies = [ + "hashbrown 0.17.1", + "slab", +] + [[package]] name = "atoi" version = "2.0.0" @@ -1769,16 +1779,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "ctor" -version = "0.6.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "424e0138278faeb2b401f174ad17e715c829512d74f3d1e81eb43365c2e0590e" -dependencies = [ - "ctor-proc-macro", - "dtor", -] - [[package]] name = "ctor" version = "1.0.8" @@ -1789,12 +1789,6 @@ dependencies = [ "linktime-proc-macro", ] -[[package]] -name = "ctor-proc-macro" -version = "0.0.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52560adf09603e58c9a7ee1fe1dcb95a16927b17c127f0ac02d6e768a0e25bc1" - [[package]] name = "ctr" version = "0.9.2" @@ -2841,21 +2835,6 @@ version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1435fa1053d8b2fbbe9be7e97eca7f33d37b28409959813daefc1446a14247f1" -[[package]] -name = "dtor" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "404d02eeb088a82cfd873006cb713fe411306c7d182c344905e101fb1167d301" -dependencies = [ - "dtor-proc-macro", -] - -[[package]] -name = "dtor-proc-macro" -version = "0.0.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f678cf4a922c215c63e0de95eb1ff08a958a81d47e485cf9da1e27bf6305cfa5" - [[package]] name = "dunce" version = "1.0.5" @@ -2968,7 +2947,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -3456,18 +3435,21 @@ checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" [[package]] name = "hf-xet" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "430b33fa84f92796d4d263070b6c0d3ca219df7b9a0e1853ee431029b1612bcd" +checksum = "c237ef4fb0ce1962a5117f8bd8c74454b41629826a9df17d14a1840ca18f0754" dependencies = [ + "anyhow", "async-trait", "bytes", "http 1.5.0", "more-asserts", "serde", + "serde_json", "thiserror 2.0.18", "tokio", "tokio-util", + "tokio_with_wasm", "tracing", "uuid", "xet-client", @@ -3998,6 +3980,7 @@ dependencies = [ "iceberg_test_utils", "opendal", "reqsign-aws-v4", + "reqsign-azure-storage", "reqsign-core", "reqsign-google", "reqwest 0.12.28", @@ -5110,11 +5093,11 @@ checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" [[package]] name = "opendal" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4f20562cc7447fcc915fc5c23df305a412ea80a733c9f2fd9e2d267e2815be6d" +checksum = "379a66c8c20c51a2321e444338ebc16a6cc0c73e39e43e9a3eea376bd71fef49" dependencies = [ - "ctor 1.0.8", + "ctor", "opendal-core", "opendal-http-transport-reqwest", "opendal-layer-concurrent-limit", @@ -5131,11 +5114,12 @@ dependencies = [ [[package]] name = "opendal-core" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ec75551ff4cf3e57da98979f6a937aaa9ddb3915bf68cc17d03df733be6646ed" +checksum = "45165de3ea19ec39e5da504fb159cc47e7c37a58a93ee49ac119f8e65c7ce271" dependencies = [ "anyhow", + "asyncband", "base64 0.23.0", "bytes", "futures", @@ -5143,7 +5127,6 @@ dependencies = [ "jiff", "log", "md-5 0.11.0", - "mea", "percent-encoding", "quick-xml 0.41.0", "reqsign-core", @@ -5157,9 +5140,9 @@ dependencies = [ [[package]] name = "opendal-http-transport-reqwest" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ad4d4f19c3ce01126a30611f8e544eaa217104a278c889ac17c9374fe4f9e4ef" +checksum = "740229ef5aa63335415ff9410c26301358faeef645945e92755916b39cdccd34" dependencies = [ "bytes", "futures", @@ -5171,21 +5154,21 @@ dependencies = [ [[package]] name = "opendal-layer-concurrent-limit" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "249ac5b0aa5a7a6c3737342d10456067937f9c9a6f3f02544271f7908ab91081" +checksum = "859214a85b9e3f75b2625b82401bde8f5c37223c5b1e208ad76551622ea84bca" dependencies = [ + "asyncband", "futures", "http 1.5.0", - "mea", "opendal-core", ] [[package]] name = "opendal-layer-logging" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5c75411ab00f77851ff086b686c1e9ca8175ac18c15afa2cb75b9036436cb06c" +checksum = "679106b611dbfb759b165ab27452c3a5b0f9c2b56df0f462d7ff5dfdc74e47dd" dependencies = [ "log", "opendal-core", @@ -5193,9 +5176,9 @@ dependencies = [ [[package]] name = "opendal-layer-retry" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "80b7738bd5f233ad8da39af9b9316b9b7a4eaddd91e8e32a1e19b7030688121d" +checksum = "3bc8bd833bd44e70bd91ed396e7d350203342766274ac8dbd4c50755cb8cf907" dependencies = [ "backon", "log", @@ -5204,9 +5187,9 @@ dependencies = [ [[package]] name = "opendal-layer-timeout" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a704141924500f3803c05ed871b53305d2a2f11cb5ef20160c3ee688a1857f66" +checksum = "90b4a2388510aed112ec81186dfde3ed72c637a7a33c323caca8a9f4a44829df" dependencies = [ "opendal-core", "tokio", @@ -5214,15 +5197,15 @@ dependencies = [ [[package]] name = "opendal-service-azdls" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2e3c406729935fe214ce574d68681a1ff7e0b322548f14094912bdbfe50e5c53" +checksum = "c7be53c9f526ce700562bcbc4b61279915ca4c5c0ea051cdd4e484435f8d8277" dependencies = [ + "asyncband", "base64 0.23.0", "bytes", "http 1.5.0", "log", - "mea", "opendal-core", "opendal-service-azure-common", "quick-xml 0.41.0", @@ -5235,9 +5218,9 @@ dependencies = [ [[package]] name = "opendal-service-azure-common" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7348c88edf15af435b7be930077746b569fac5e738c1bf6a363b675e7317c9df" +checksum = "e9c9bbb2670f63912a267ff9a0c46233165f7f84fb77f07e28eb1c8be51c633b" dependencies = [ "http 1.5.0", "opendal-core", @@ -5245,9 +5228,9 @@ dependencies = [ [[package]] name = "opendal-service-fs" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "826c4e17a30643b888fe983897f9a4b23b07066e1d069727a923cc8fb419a702" +checksum = "05332f2bfd6c38ebc8155b9362f1d3b31a27ccc36709a1a776c050c376b9dddf" dependencies = [ "bytes", "log", @@ -5259,9 +5242,9 @@ dependencies = [ [[package]] name = "opendal-service-gcs" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "007f3fba63c21e516c956b891e96ff9892d8175662bfb781cdada9d3766a11e6" +checksum = "b001c0eb6b8c40b8397bdd837001982fcdae163ae493b40327dd624b1ab1b307" dependencies = [ "async-trait", "bytes", @@ -5276,14 +5259,16 @@ dependencies = [ "serde", "serde_json", "tokio", + "uuid", ] [[package]] name = "opendal-service-hf" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b41fd41eb7ed03c5e66cefda61e8e117808ffd2908f2916737cb020a6beb02c7" +checksum = "f6e3d36689653ad52fce8f16c2c9181e930920446175812a452ebe5ee5ce2d64" dependencies = [ + "asyncband", "bytes", "hf-xet", "http 1.5.0", @@ -5296,9 +5281,9 @@ dependencies = [ [[package]] name = "opendal-service-oss" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cd528ec2d49c5ca69e674ffed7b3e0686fb9cfcfea0596870de381467fda4f1b" +checksum = "3e09763037e1dd404f8b3e2b71b0efef891365deaabfa7a8cc8944c14cea2bc1" dependencies = [ "bytes", "http 1.5.0", @@ -5313,9 +5298,9 @@ dependencies = [ [[package]] name = "opendal-service-s3" -version = "0.58.1" +version = "0.59.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "58e80cdf192d7eff05feed747894d64f81905ac4eaf132edf7ea270abdd2d663" +checksum = "067f5bdc362a85ed02926f5184e4a9ebb830a3c3ed8bfab0692b444c69d5210a" dependencies = [ "base64 0.23.0", "bytes", @@ -5477,11 +5462,11 @@ dependencies = [ [[package]] name = "pem" -version = "3.0.6" +version = "4.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d30c53c26bc5b31a98cd02d20f25a7c8567146caf63ed593a9d87b2775291be" +checksum = "d354a98a3d1251555de99e8fdd8afda05573c31b82f59063a7b0a29b5527f120" dependencies = [ - "base64 0.22.1", + "base64 0.23.0", "serde_core", ] @@ -5925,7 +5910,7 @@ dependencies = [ "once_cell", "socket2 0.5.10", "tracing", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -6199,9 +6184,9 @@ dependencies = [ [[package]] name = "reqsign-aws-core" -version = "3.1.0" +version = "3.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4d63b56638bb3cc7bd376a7cdce1ba3089777a08f47e4097888f2d784cc3f46c" +checksum = "bac4749b7dfa7bfaccd01eb03e9dc795ed37e3f20d6f0f38e2c67ee85ad6bc86" dependencies = [ "bytes", "form_urlencoded", @@ -6220,9 +6205,9 @@ dependencies = [ [[package]] name = "reqsign-aws-v4" -version = "3.2.0" +version = "3.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4a0c499f4ed12d04c3d4c78fe4cb01aee22c9dae22848c14db2c6313d9df9f43" +checksum = "ff250f0fd0b913fbd565e405acc553da0f13bde30bfb5403178c9d0313cdc15f" dependencies = [ "bytes", "http 1.5.0", @@ -6235,9 +6220,9 @@ dependencies = [ [[package]] name = "reqsign-azure-storage" -version = "3.1.2" +version = "3.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2824e7da3c2cc42ac3406c674eb57c89127fdcd97f3a73c608cfc680505ea134" +checksum = "c3f1fbc9add082f54e51e3bd9730f18d850a1f09d246baff5cd612c74ae4290e" dependencies = [ "anyhow", "base64 0.23.0", @@ -6256,9 +6241,9 @@ dependencies = [ [[package]] name = "reqsign-core" -version = "3.3.0" +version = "3.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f4ac1510872d9481205975d264deb39c109797e5068cc882ed9064270eaae5fa" +checksum = "ff052daffb0599681c50f85c59e7236438976efe991ab864edd9f3b235501a0f" dependencies = [ "anyhow", "base64 0.23.0", @@ -6269,6 +6254,7 @@ dependencies = [ "http 1.5.0", "jiff", "log", + "mea", "percent-encoding", "rsa", "serde", @@ -6280,9 +6266,9 @@ dependencies = [ [[package]] name = "reqsign-file-read-tokio" -version = "3.0.4" +version = "3.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "663d9d55abd0df0830ef0ae43708297cc1371cf4e8ca91f3ac813c309cca8c98" +checksum = "b3235df90a6bca681aa47dd86f2393d122a6d77042aa8a7c81e218cd45c5bfc0" dependencies = [ "anyhow", "reqsign-core", @@ -6544,7 +6530,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -6602,7 +6588,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -6948,7 +6934,6 @@ dependencies = [ "cfg-if 1.0.4", "cpufeatures 0.2.17", "digest 0.10.7", - "sha2-asm", ] [[package]] @@ -6962,15 +6947,6 @@ dependencies = [ "digest 0.11.3", ] -[[package]] -name = "sha2-asm" -version = "0.6.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b845214d6175804686b2bd482bcffe96651bb2d1200742b712003504a2dac1ab" -dependencies = [ - "cc", -] - [[package]] name = "sharded-slab" version = "0.1.7" @@ -7601,7 +7577,7 @@ dependencies = [ "getrandom 0.3.4", "once_cell", "rustix", - "windows-sys 0.61.2", + "windows-sys 0.52.0", ] [[package]] @@ -7806,6 +7782,30 @@ dependencies = [ "tokio", ] +[[package]] +name = "tokio_with_wasm" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "34e40fbbbd95441133fe9483f522db15dbfd26dc636164ebd8f2dd28759a6aa6" +dependencies = [ + "js-sys", + "tokio", + "tokio_with_wasm_proc", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "tokio_with_wasm_proc" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d01145a2c788d6aae4cd653afec1e8332534d7d783d01897cefcafe4428de992" +dependencies = [ + "quote", + "syn 2.0.119", +] + [[package]] name = "toml" version = "0.8.23" @@ -8863,20 +8863,18 @@ dependencies = [ [[package]] name = "xet-client" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3e1e496dcbe6a09017acdfaf48e1a646735e7ff5b2a49e2c7e081cca77a59bc8" +checksum = "c3b8da8cc70aa2e3c500c0400e012df82c656ab9fca47f9f939fffc5afd89aca" dependencies = [ "anyhow", "async-trait", "base64 0.22.1", "bytes", - "clap", "crc32fast", "futures", "http 1.5.0", "hyper", - "lazy_static", "more-asserts", "rand 0.10.2", "redb", @@ -8890,8 +8888,8 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-retry", + "tokio_with_wasm", "tracing", - "tracing-subscriber", "url", "urlencoding", "web-time", @@ -8901,24 +8899,21 @@ dependencies = [ [[package]] name = "xet-core-structures" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cb838aa8eb67d730af301584cf003caad407487606058292a6750711b603fbee" +checksum = "73503c223783dccc864abde22115e09d12f190448a0baf58ab2c54bc709e2f99" dependencies = [ "async-trait", "base64 0.22.1", "blake3", "bytemuck", "bytes", - "clap", "countio", - "csv", "futures", "futures-util", "getrandom 0.4.3", "heapify", "itertools 0.14.0", - "lazy_static", "lz4_flex 0.13.1", "more-asserts", "rand 0.10.2", @@ -8926,7 +8921,6 @@ dependencies = [ "safe-transmute", "serde", "static_assertions", - "tempfile", "thiserror 2.0.18", "tokio", "tokio-util", @@ -8938,32 +8932,31 @@ dependencies = [ [[package]] name = "xet-data" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "67fd409bef621411a9d9013798540bb8036cb2678f03ab39af89a5e88034ed8c" +checksum = "c89052ec5dec2187cad30b86af92cc24fd61c4a57a795f1ff7ff5f38d49184eb" dependencies = [ "anyhow", "async-trait", "bytes", "chrono", - "clap", "gearhash", "http 1.5.0", "itertools 0.14.0", - "lazy_static", "more-asserts", "rand 0.10.2", "serde", "serde_json", - "sha2 0.10.9", + "sha2 0.11.0", "tempfile", "thiserror 2.0.18", "tokio", "tokio-util", + "tokio_with_wasm", "tracing", "url", "uuid", - "walkdir", + "web-time", "xet-client", "xet-core-structures", "xet-runtime", @@ -8971,9 +8964,9 @@ dependencies = [ [[package]] name = "xet-runtime" -version = "1.5.2" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "15d8f121c33866f7648b737abe70d0e2dd9c0af4ffdd7219207531d0283aa63d" +checksum = "af5c60d5eed38ab4c576f4421bae835e7bd07631fb381705605529d2015c106b" dependencies = [ "anyhow", "async-trait", @@ -8981,13 +8974,12 @@ dependencies = [ "chrono", "colored", "const-str", - "ctor 0.6.3", + "ctor", "dirs", "futures", "git-version", "humantime", "konst", - "lazy_static", "libc", "more-asserts", "oneshot", @@ -9001,9 +8993,11 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-util", + "tokio_with_wasm", "tracing", "tracing-appender", "tracing-subscriber", + "web-time", "whoami 2.1.2", "winapi", ] diff --git a/Cargo.toml b/Cargo.toml index 8c220f489c..3226546c9a 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -123,7 +123,7 @@ mockito = "1" motore-macros = "0.4.3" murmur3 = "0.5.2" once_cell = "1.20" -opendal = "0.58" +opendal = "0.59" ordered-float = "4" parquet = "59.2" pilota = "0.11.10" diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index 03dba8a020..2d90ead06f 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -145,9 +145,10 @@ impl HttpClient { /// Derives a client for table-scoped requests. /// - /// The connection pool and catalog headers are inherited, explicit table - /// headers override them, and `auth_session` replaces the catalog session - /// only when the auth manager selected a table-specific child session. + /// The connection pool and catalog headers are inherited, headers in the + /// effective table properties override them, and `auth_session` replaces + /// the catalog session only when the auth manager selected a table-specific + /// child session. pub(crate) fn for_table( &self, props: &HashMap, @@ -295,6 +296,25 @@ pub(crate) fn format_headers_redacted(headers: &HeaderMap, disable_redaction: bo pub(crate) fn deserialize_unexpected_catalog_error( response: HttpResponse, disable_header_redaction: bool, +) -> Error { + unexpected_catalog_error(response, disable_header_redaction, true) +} + +/// Builds an unexpected catalog error without retaining the response body. +/// +/// Credential endpoints use this because even an unsuccessful response may +/// contain credential material that must not be surfaced through an error. +pub(crate) fn unexpected_catalog_error_without_body( + response: HttpResponse, + disable_header_redaction: bool, +) -> Error { + unexpected_catalog_error(response, disable_header_redaction, false) +} + +fn unexpected_catalog_error( + response: HttpResponse, + disable_header_redaction: bool, + include_body: bool, ) -> Error { let err = Error::new( ErrorKind::Unexpected, @@ -307,7 +327,7 @@ pub(crate) fn deserialize_unexpected_catalog_error( ); let bytes = response.body(); - if bytes.is_empty() { + if !include_body || bytes.is_empty() { return err; } err.with_context("json", String::from_utf8_lossy(bytes)) @@ -423,6 +443,23 @@ mod tests { assert!(!err.contains("leaked"), "{err}"); } + #[test] + fn test_unexpected_error_can_omit_a_secret_body() { + let response = HttpResponse::new( + http::StatusCode::INTERNAL_SERVER_ERROR, + HeaderMap::new(), + br#"{"secret": "do-not-expose"}"#.to_vec(), + ); + + let err = format!( + "{:?}", + unexpected_catalog_error_without_body(response, false) + ); + + assert!(err.contains("500"), "{err}"); + assert!(!err.contains("do-not-expose"), "{err}"); + } + #[tokio::test] async fn test_post_form_carries_the_session_until_it_is_removed() { // Every request a client sends carries its session; a caller that diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs index e9b325b170..24e2ed7002 100644 --- a/crates/catalog/rest/src/credential.rs +++ b/crates/catalog/rest/src/credential.rs @@ -47,11 +47,12 @@ use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH}; use async_trait::async_trait; use iceberg::io::{ - AWS_REFRESH_CREDENTIALS_ENABLED, AWS_REFRESH_CREDENTIALS_ENDPOINT, - GCS_REFRESH_CREDENTIALS_ENABLED, GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, - GCS_TOKEN_EXPIRES_AT, GcsCredential, S3_ACCESS_KEY_ID, S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, - S3_SESSION_TOKEN_EXPIRES_AT_MS, S3Credential, StorageCredential, StorageCredentialKind, - StorageCredentialProvider, + ADLS_REFRESH_CREDENTIALS_ENABLED, ADLS_REFRESH_CREDENTIALS_ENDPOINT, + ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX, ADLS_SAS_TOKEN_PREFIX, AWS_REFRESH_CREDENTIALS_ENABLED, + AWS_REFRESH_CREDENTIALS_ENDPOINT, AzdlsCredential, GCS_REFRESH_CREDENTIALS_ENABLED, + GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, GcsCredential, + S3_ACCESS_KEY_ID, S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, S3_SESSION_TOKEN_EXPIRES_AT_MS, + S3Credential, StorageCredential, StorageCredentialKind, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result, TableIdent}; use rand::Rng; @@ -60,10 +61,33 @@ use tokio::sync::Mutex; use crate::REST_CATALOG_PROP_SCAN_PLAN_ID; use crate::auth::AuthManager; -use crate::client::{HttpClient, deserialize_unexpected_catalog_error}; +use crate::client::{HttpClient, unexpected_catalog_error_without_body}; use crate::request::HttpRequest; use crate::types::LoadCredentialsResponse; +type CredentialParser = + fn(config: &HashMap, prefix: Option) -> Result; +type KeyedSeedCredentialParser = + fn(config: &HashMap) -> HashMap; +type KeyedSeedPathResolver = fn(path: &str) -> Result; + +enum SeedStrategy { + /// One credential stored in the backend's flat properties. + Flat, + /// Credentials selected by a backend-specific key derived from each path. + Keyed(KeyedSeedStrategy), +} + +struct KeyedSeedStrategy { + parse_credentials: KeyedSeedCredentialParser, + resolve_path: KeyedSeedPathResolver, +} + +struct KeyedSeedPath { + key: String, + scope: String, +} + /// Cloud-specific details regarding vended-credential refresh. /// /// It contains the location schemes it backs, the property keys it is configured @@ -79,8 +103,9 @@ struct CloudRefresh { /// Whether to jitter successful prefetch times like AWS `CachedSupplier`. jitter_prefetch: bool, /// Parse a complete credential from catalog-supplied properties. - parse_credential: - fn(config: &HashMap, prefix: Option) -> Result, + parse_credential: CredentialParser, + /// How initial credentials in the table properties are selected. + seed_strategy: SeedStrategy, } impl CloudRefresh { @@ -91,6 +116,7 @@ impl CloudRefresh { enabled_key: AWS_REFRESH_CREDENTIALS_ENABLED, jitter_prefetch: true, parse_credential: parse_s3_credential, + seed_strategy: SeedStrategy::Flat, }; /// Google Cloud Storage const GCP: Self = Self { @@ -99,13 +125,23 @@ impl CloudRefresh { enabled_key: GCS_REFRESH_CREDENTIALS_ENABLED, jitter_prefetch: false, parse_credential: parse_gcs_credential, + seed_strategy: SeedStrategy::Flat, + }; + /// Azure Data Lake Storage + const AZURE: Self = Self { + schemes: &["abfs", "abfss", "wasb", "wasbs"], + endpoint_key: ADLS_REFRESH_CREDENTIALS_ENDPOINT, + enabled_key: ADLS_REFRESH_CREDENTIALS_ENABLED, + jitter_prefetch: false, + parse_credential: parse_azdls_credential, + seed_strategy: SeedStrategy::Keyed(KeyedSeedStrategy { + parse_credentials: parse_azdls_account_seeds, + resolve_path: resolve_azdls_seed_path, + }), }; - // TODO: Azure (ADLS) is not yet supported: opendal's Azdls builder exposes no - // custom credential-provider hook, and reqsign's SAS-token credential has no - // expiry, so reqsign-based refresh isn't possible. /// Backends with refresh support - const SUPPORTED: &[Self] = &[Self::AWS, Self::GCP]; + const SUPPORTED: &[Self] = &[Self::AWS, Self::GCP, Self::AZURE]; /// The backend that serves `location`, by its URL scheme, or `None` if no /// supported backend matches (in which case static credentials are used @@ -232,12 +268,31 @@ struct ConfiguredCloud { cloud: &'static CloudRefresh, endpoint: String, cache: Mutex, + /// Initial property credentials whose scope cannot be represented by a + /// single leading URI prefix, keyed according to the cloud strategy. + keyed_seeds: HashMap, /// Only one caller fetches at a time. The cache lock is deliberately /// separate so other callers can keep using an unexpired credential while /// the refresh is in flight. refresh: Mutex<()>, } +impl ConfiguredCloud { + fn keyed_seed_for_path(&self, path: &str) -> Result> { + let SeedStrategy::Keyed(strategy) = &self.cloud.seed_strategy else { + return Ok(None); + }; + let resolved = (strategy.resolve_path)(path)?; + let Some(seed) = self.keyed_seeds.get(&resolved.key) else { + return Ok(None); + }; + Ok(Some(CachedEntry { + credential: seed.credential.clone().with_prefix(resolved.scope), + refresh_at: seed.refresh_at, + })) + } +} + /// Fetches and refreshes vended credentials from a REST catalog's table /// credentials endpoint. /// @@ -322,7 +377,7 @@ impl RestVendedCredentialProvider { } Ok(ParsedCredentials { entries, errors }) } - _ => Err(deserialize_unexpected_catalog_error( + _ => Err(unexpected_catalog_error_without_body( response, self.client.disable_header_redaction(), )), @@ -333,7 +388,7 @@ impl RestVendedCredentialProvider { &self, configured: &ConfiguredCloud, path: &str, - fallback: Option, + fallback: Option, ) -> Result { match self.fetch(configured).await { Ok(ParsedCredentials { entries, errors }) => { @@ -360,8 +415,8 @@ impl RestVendedCredentialProvider { if entries.is_empty() { cache.record_failure(); return fallback - .filter(|entry| entry.is_unexpired(SystemTime::now())) - .map(|entry| entry.credential) + .filter(|fallback| fallback.entry.is_unexpired(SystemTime::now())) + .map(|fallback| fallback.entry.credential) .ok_or(failure); } @@ -414,12 +469,15 @@ impl RestVendedCredentialProvider { return Ok(credential); } - // Cache valid credentials for other prefixes while retaining this - // path's still-usable credential when its replacement was absent or - // invalid. Backoff ensures the failed path is retried promptly. - if let Some(fallback) = fallback.filter(|entry| entry.is_unexpired(now)) { - let credential = fallback.credential.clone(); - cache.insert_fallback_if_missing(fallback, now); + // Cache valid credentials for other prefixes. If none covers this + // path, retain its still-usable credential when its replacement was + // absent or invalid. Backoff ensures the failed path is retried promptly. + if let Some(fallback) = fallback.filter(|fallback| fallback.entry.is_unexpired(now)) + { + let credential = fallback.entry.credential.clone(); + if fallback.cacheable { + cache.insert_fallback_if_missing(fallback.entry, now); + } cache.record_failure(); return Ok(credential); } @@ -435,8 +493,8 @@ impl RestVendedCredentialProvider { // usable, serve it and retry after jittered backoff. Expired // credentials are never served. fallback - .filter(|entry| entry.is_unexpired(SystemTime::now())) - .map(|entry| entry.credential) + .filter(|fallback| fallback.entry.is_unexpired(SystemTime::now())) + .map(|fallback| fallback.entry.credential) .ok_or(fetch_error) } } @@ -453,10 +511,34 @@ impl std::fmt::Debug for RestVendedCredentialProvider { enum CacheDecision { Use(StorageCredential), - Refresh(Option), + Refresh(Option), Backoff, } +#[derive(Clone)] +struct CredentialFallback { + entry: CachedEntry, + /// Whether this entry belongs in the ordinary prefix cache after a failed + /// replacement. Keyed seeds remain in their separately scoped map. + cacheable: bool, +} + +impl CredentialFallback { + fn cached(entry: CachedEntry) -> Self { + Self { + entry, + cacheable: true, + } + } + + fn keyed_seed(entry: CachedEntry) -> Self { + Self { + entry, + cacheable: false, + } + } +} + fn refresh_backoff_error(path: &str) -> Error { Error::new( ErrorKind::Unexpected, @@ -464,27 +546,40 @@ fn refresh_backoff_error(path: &str) -> Error { ) } -async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> CacheDecision { +async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> Result { + let keyed_seed = configured + .keyed_seed_for_path(path)? + .map(CredentialFallback::keyed_seed); let cache = configured.cache.lock().await; - let current = longest_prefix_match(&cache.entries, path).cloned(); let now = SystemTime::now(); + let cached = longest_prefix_match(&cache.entries, path) + .cloned() + .map(CredentialFallback::cached); + let current = match cached { + Some(cached) if cached.entry.is_unexpired(now) => Some(cached), + Some(expired) => keyed_seed.or(Some(expired)), + None => keyed_seed, + }; - if let Some(entry) = current.as_ref().filter(|entry| entry.is_fresh(now)) { - return CacheDecision::Use(entry.credential.clone()); + if let Some(fallback) = current + .as_ref() + .filter(|fallback| fallback.entry.is_fresh(now)) + { + return Ok(CacheDecision::Use(fallback.entry.credential.clone())); } if cache .retry_not_before .is_some_and(|retry_at| Instant::now() < retry_at) { - return current + return Ok(current .as_ref() - .filter(|entry| entry.is_unexpired(now)) - .map(|entry| CacheDecision::Use(entry.credential.clone())) - .unwrap_or(CacheDecision::Backoff); + .filter(|fallback| fallback.entry.is_unexpired(now)) + .map(|fallback| CacheDecision::Use(fallback.entry.credential.clone())) + .unwrap_or(CacheDecision::Backoff)); } - CacheDecision::Refresh(current) + Ok(CacheDecision::Refresh(current)) } #[async_trait] @@ -507,7 +602,7 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { ) })?; - let current = match cache_decision(configured, path).await { + let current = match cache_decision(configured, path).await? { CacheDecision::Use(credential) => return Ok(credential), CacheDecision::Refresh(current) => current, CacheDecision::Backoff => return Err(refresh_backoff_error(path)), @@ -518,11 +613,11 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { // for the in-flight refresh instead. let usable = current .as_ref() - .filter(|entry| entry.is_unexpired(SystemTime::now())); - let _refresh_guard = if let Some(entry) = usable { + .filter(|fallback| fallback.entry.is_unexpired(SystemTime::now())); + let _refresh_guard = if let Some(fallback) = usable { match configured.refresh.try_lock() { Ok(guard) => guard, - Err(_) => return Ok(entry.credential.clone()), + Err(_) => return Ok(fallback.entry.credential.clone()), } } else { configured.refresh.lock().await @@ -530,7 +625,7 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { // Another caller may have completed a refresh between our cache check // and acquiring the single-flight guard. - let current = match cache_decision(configured, path).await { + let current = match cache_decision(configured, path).await? { CacheDecision::Use(credential) => return Ok(credential), CacheDecision::Refresh(current) => current, CacheDecision::Backoff => return Err(refresh_backoff_error(path)), @@ -583,6 +678,8 @@ fn failure_backoff(consecutive_failures: u32) -> Duration { /// or `None` when no supported cloud advertises an enabled refresh endpoint. /// /// `base_uri` is the catalog URI, used to resolve a relative endpoint. +/// `props` contains the effective FileIO properties after applying local +/// overrides, and therefore supplies headers and client policy. /// `table_auth_props` is the unmerged config returned by the table endpoint: /// keeping it separate prevents local FileIO overrides from masking table auth. /// `auth_manager` derives an isolated child session when those properties @@ -611,10 +708,26 @@ pub(crate) async fn build_vended_credential_provider( .get(cloud.endpoint_key) .filter(|value| !value.is_empty())?, ); - let entries = (cloud.parse_credential)(props, None) - .ok() - .map(|credential| vec![CachedEntry::seed(credential, cloud.jitter_prefetch)]) - .unwrap_or_default(); + let (entries, keyed_seeds) = match &cloud.seed_strategy { + SeedStrategy::Flat => { + let entries = (cloud.parse_credential)(props, None) + .ok() + .map(|credential| { + vec![CachedEntry::seed(credential, cloud.jitter_prefetch)] + }) + .unwrap_or_default(); + (entries, HashMap::new()) + } + SeedStrategy::Keyed(strategy) => { + let keyed_seeds = (strategy.parse_credentials)(props) + .into_iter() + .map(|(key, credential)| { + (key, CachedEntry::seed(credential, cloud.jitter_prefetch)) + }) + .collect(); + (Vec::new(), keyed_seeds) + } + }; Some(ConfiguredCloud { cloud, @@ -624,6 +737,7 @@ pub(crate) async fn build_vended_credential_provider( consecutive_failures: 0, retry_not_before: None, }), + keyed_seeds, refresh: Mutex::new(()), }) }) @@ -643,7 +757,7 @@ pub(crate) async fn build_vended_credential_provider( client.auth_session(), ) .await?; - let table_client = Arc::new(client.for_table(table_auth_props, table_session)?); + let table_client = Arc::new(client.for_table(props, table_session)?); let plan_id = props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(); Ok(Some(Arc::new(RestVendedCredentialProvider::new( table_client, @@ -706,6 +820,103 @@ fn parse_gcs_credential( }) } +/// Parse an account-specific ADLS SAS credential returned by the catalog. +fn parse_azdls_credential( + config: &HashMap, + prefix: Option, +) -> Result { + let prefix = prefix.ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "invalid vended ADLS credential: storage prefix is missing", + ) + })?; + let account = azdls_account_name(&prefix)?; + let sas_token = required_nonempty(config, &format!("{ADLS_SAS_TOKEN_PREFIX}{account}"))?; + let expires_at = required_epoch_millis( + config, + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{account}"), + )?; + + Ok( + StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new( + sas_token, + ))) + .with_prefix(prefix) + .with_expiration(expires_at), + ) +} + +/// Parse every complete account-qualified ADLS credential from the initial +/// table properties. Unlike a URI prefix, an Azure account occurs after the +/// filesystem in a location, so these seeds are selected by account name. +fn parse_azdls_account_seeds( + config: &HashMap, +) -> HashMap { + config + .iter() + .filter_map(|(key, value)| { + let account = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; + if account.is_empty() || value.is_empty() { + return None; + } + let expires_at = required_epoch_millis( + config, + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{account}"), + ) + .ok()?; + Some(( + account.to_string(), + StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new(value))) + .with_expiration(expires_at), + )) + }) + .collect() +} + +fn resolve_azdls_seed_path(location: &str) -> Result { + let (mut url, account) = parse_azdls_location(location)?; + url.set_path("/"); + url.set_query(None); + url.set_fragment(None); + Ok(KeyedSeedPath { + key: account, + scope: url.to_string(), + }) +} + +fn azdls_account_name(location: &str) -> Result { + parse_azdls_location(location).map(|(_, account)| account) +} + +fn parse_azdls_location(location: &str) -> Result<(Url, String)> { + let url = Url::parse(location).map_err(|error| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid ADLS storage location: {location}"), + ) + .with_source(error) + })?; + if !CloudRefresh::AZURE.schemes.contains(&url.scheme()) { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("invalid ADLS storage location scheme: {}", url.scheme()), + )); + } + let account = url + .host_str() + .and_then(|host| host.split('.').next()) + .filter(|account| !account.is_empty()) + .map(str::to_owned) + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("ADLS storage location has no account name: {location}"), + ) + })?; + Ok((url, account)) +} + fn required_nonempty(config: &HashMap, key: &str) -> Result { config .get(key) @@ -896,9 +1107,11 @@ mod tests { assert!(CloudRefresh::for_location("s3n://b/k").is_some()); assert!(CloudRefresh::for_location("gs://b/k").is_some()); assert!(CloudRefresh::for_location("gcs://b/k").is_some()); + assert!(CloudRefresh::for_location("abfs").is_some()); + assert!(CloudRefresh::for_location("abfss://fs@acct.dfs.core.windows.net/k").is_some()); + assert!(CloudRefresh::for_location("wasb://fs@acct.blob.core.windows.net/k").is_some()); + assert!(CloudRefresh::for_location("wasbs://fs@acct.blob.core.windows.net/k").is_some()); assert!(CloudRefresh::for_location("not a url").is_none()); - // Azure not supported yet -> no provider, static creds used as-is. - assert!(CloudRefresh::for_location("abfss://fs@acct.dfs.core.windows.net/k").is_none()); assert!(CloudRefresh::AWS.matches_location("s3://bucket/path")); assert!(!CloudRefresh::AWS.matches_location("s3evil://bucket/path")); } @@ -949,6 +1162,32 @@ mod tests { ); } + #[test] + fn parse_azdls_requires_account_specific_token_and_expiry() { + let prefix = "abfss://container@account1.dfs.core.windows.net/table"; + let token_key = format!("{ADLS_SAS_TOKEN_PREFIX}account1"); + let expiry_key = format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account1"); + let mut config = HashMap::new(); + + assert!(parse_azdls_credential(&config, Some(prefix.to_string())).is_err()); + config.insert(token_key, "sv=2026&sig=secret".to_string()); + assert!(parse_azdls_credential(&config, Some(prefix.to_string())).is_err()); + config.insert(expiry_key, "2500".to_string()); + + let credential = parse_azdls_credential(&config, Some(prefix.to_string())).unwrap(); + assert_eq!(credential.prefix(), Some(prefix)); + assert_eq!( + credential.expires_at(), + Some(UNIX_EPOCH + Duration::from_millis(2500)) + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=secret") + } + other => panic!("expected ADLS, got {other:?}"), + } + } + #[test] fn cached_entry_freshness() { let no_expiry = cached_s3("", "a", None); @@ -1263,6 +1502,104 @@ mod tests { mock.assert_async().await; } + #[tokio::test] + async fn fresh_azdls_seeds_are_scoped_and_served_without_fetching() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", Matcher::Any) + .expect(0) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account1"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account1"), + epoch_millis(SystemTime::now() + Duration::from_secs(3600)), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account2"), + "sv=2026&sig=other-seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account2"), + epoch_millis(SystemTime::now() + Duration::from_secs(3600)), + ), + ]); + let provider = build_vended_credential_provider( + test_client(&server.url()), + &NoopAuthManager, + &test_table(), + &server.url(), + &props, + None, + ) + .await + .unwrap() + .unwrap(); + + let credential = provider + .load_credential("abfss://container@account1.dfs.core.windows.net/table/data/a.parquet") + .await + .unwrap(); + assert_eq!( + credential.prefix(), + Some("abfss://container@account1.dfs.core.windows.net/") + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=seed") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + + let credential = provider + .load_credential("abfss://other-container@account2.dfs.core.windows.net/data/b.parquet") + .await + .unwrap(); + assert_eq!( + credential.prefix(), + Some("abfss://other-container@account2.dfs.core.windows.net/") + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=other-seed") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + mock.assert_async().await; + } + + #[tokio::test] + async fn invalid_azdls_path_is_rejected_before_refresh() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", Matcher::Any) + .expect(0) + .create_async() + .await; + let props = HashMap::from([( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + )]); + let provider = test_provider(&server.url(), &props).await; + + let error = provider + .load_credential("abfss:///data.parquet") + .await + .unwrap_err(); + + assert_eq!(error.kind(), ErrorKind::DataInvalid); + assert!(error.message().contains("no account name")); + mock.assert_async().await; + } + #[tokio::test] async fn failed_refresh_is_backed_off_while_credential_is_unexpired() { let mut server = Server::new_async().await; @@ -1349,6 +1686,7 @@ mod tests { consecutive_failures: 0, retry_not_before: None, }), + keyed_seeds: HashMap::new(), refresh: Mutex::new(()), }, ]); @@ -1427,6 +1765,7 @@ mod tests { consecutive_failures: 0, retry_not_before: None, }), + keyed_seeds: HashMap::new(), refresh: Mutex::new(()), }, ]); @@ -1564,6 +1903,58 @@ mod tests { mock.assert_async().await; } + #[tokio::test] + async fn effective_headers_override_raw_table_headers() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer local-header") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let client = test_client(&server.url()); + let props = HashMap::from([ + ( + CloudRefresh::AWS.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + "header.Authorization".to_string(), + "Bearer local-header".to_string(), + ), + ]); + let table_auth = HashMap::from([ + ("token".to_string(), "table-token".to_string()), + ( + "header.Authorization".to_string(), + "Bearer table-header".to_string(), + ), + ]); + let auth_manager = OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())); + let provider = build_vended_credential_provider( + client, + &auth_manager, + &test_table(), + &server.url(), + &props, + Some(&table_auth), + ) + .await + .unwrap() + .unwrap(); + + provider.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + #[tokio::test] async fn one_provider_refreshes_multiple_clouds() { let mut server = Server::new_async().await; @@ -1622,6 +2013,40 @@ mod tests { gcp_mock.assert_async().await; } + #[tokio::test] + async fn refreshes_azdls_credentials() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let prefix = "abfss://container@account1.dfs.core.windows.net/table"; + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"{prefix}","config":{{"adls.sas-token.account1":"sv=2026&sig=secret","adls.sas-token-expires-at-ms.account1":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/azure-credentials") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + let props = HashMap::from([( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/azure-credentials".to_string(), + )]); + let provider = test_provider(&server.url(), &props).await; + + let credential = provider + .load_credential(&format!("{prefix}/data.parquet")) + .await + .unwrap(); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=secret") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + mock.assert_async().await; + } + #[tokio::test] async fn refresh_failure_never_serves_expired_credential() { let mut server = Server::new_async().await; @@ -1645,6 +2070,50 @@ mod tests { mock.assert_async().await; } + #[tokio::test] + async fn failed_azdls_refresh_uses_account_seed_until_expiry() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(503) + .with_header("content-type", "application/json") + .with_body(r#"{"error":{"message":"unavailable","type":"test","code":503}}"#) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account2"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account2"), + epoch_millis(SystemTime::now() + Duration::from_secs(60)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + let path = "abfss://container@account2.dfs.core.windows.net/data/a.parquet"; + + for _ in 0..2 { + let credential = provider.load_credential(path).await.unwrap(); + assert_eq!( + credential.prefix(), + Some("abfss://container@account2.dfs.core.windows.net/") + ); + match credential.kind() { + StorageCredentialKind::Azdls(azdls) => { + assert_eq!(azdls.sas_token(), "sv=2026&sig=seed") + } + other => panic!("expected ADLS credential, got {other:?}"), + } + } + mock.assert_async().await; + } + #[tokio::test] async fn sub_buffer_ttl_is_refetched_on_each_sequential_operation() { let mut server = Server::new_async().await; diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index b6747c7ece..f8120568e4 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -695,6 +695,7 @@ pub async fn iceberg::inspect::SnapshotsTable<'a>::scan(&self) -> iceberg::Resul pub fn iceberg::inspect::SnapshotsTable<'a>::schema(&self) -> iceberg::spec::Schema pub mod iceberg::io pub enum iceberg::io::StorageCredentialKind +pub iceberg::io::StorageCredentialKind::Azdls(iceberg::io::AzdlsCredential) pub iceberg::io::StorageCredentialKind::Gcs(iceberg::io::GcsCredential) pub iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential) impl core::clone::Clone for iceberg::io::StorageCredentialKind @@ -731,6 +732,15 @@ impl serde_core::ser::Serialize for iceberg::io::AzdlsConfig pub fn iceberg::io::AzdlsConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::AzdlsConfig pub fn iceberg::io::AzdlsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg::io::AzdlsCredential +impl iceberg::io::AzdlsCredential +pub fn iceberg::io::AzdlsCredential::into_sas_token(self) -> alloc::string::String +pub fn iceberg::io::AzdlsCredential::new(sas_token: impl core::convert::Into) -> Self +pub fn iceberg::io::AzdlsCredential::sas_token(&self) -> &str +impl core::clone::Clone for iceberg::io::AzdlsCredential +pub fn iceberg::io::AzdlsCredential::clone(&self) -> iceberg::io::AzdlsCredential +impl core::fmt::Debug for iceberg::io::AzdlsCredential +pub fn iceberg::io::AzdlsCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::FileIO impl iceberg::io::FileIO pub fn iceberg::io::FileIO::config(&self) -> &iceberg::io::StorageConfig @@ -1048,7 +1058,11 @@ pub const iceberg::io::ADLS_AUTHORITY_HOST: &str pub const iceberg::io::ADLS_CLIENT_ID: &str pub const iceberg::io::ADLS_CLIENT_SECRET: &str pub const iceberg::io::ADLS_CONNECTION_STRING: &str +pub const iceberg::io::ADLS_REFRESH_CREDENTIALS_ENABLED: &str +pub const iceberg::io::ADLS_REFRESH_CREDENTIALS_ENDPOINT: &str pub const iceberg::io::ADLS_SAS_TOKEN: &str +pub const iceberg::io::ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX: &str +pub const iceberg::io::ADLS_SAS_TOKEN_PREFIX: &str pub const iceberg::io::ADLS_TENANT_ID: &str pub const iceberg::io::AWS_REFRESH_CREDENTIALS_ENABLED: &str pub const iceberg::io::AWS_REFRESH_CREDENTIALS_ENDPOINT: &str diff --git a/crates/iceberg/src/io/storage/config/azdls.rs b/crates/iceberg/src/io/storage/config/azdls.rs index a9541e7791..7d49a354dc 100644 --- a/crates/iceberg/src/io/storage/config/azdls.rs +++ b/crates/iceberg/src/io/storage/config/azdls.rs @@ -36,6 +36,14 @@ pub const ADLS_ACCOUNT_NAME: &str = "adls.account-name"; pub const ADLS_ACCOUNT_KEY: &str = "adls.account-key"; /// The shared access signature. pub const ADLS_SAS_TOKEN: &str = "adls.sas-token"; +/// Prefix for account-specific shared access signatures vended by a REST catalog. +pub const ADLS_SAS_TOKEN_PREFIX: &str = "adls.sas-token."; +/// Prefix for the epoch-millisecond expiration of an account-specific SAS token. +pub const ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX: &str = "adls.sas-token-expires-at-ms."; +/// Table property naming the endpoint used to refresh vended ADLS credentials. +pub const ADLS_REFRESH_CREDENTIALS_ENDPOINT: &str = "adls.refresh-credentials-endpoint"; +/// Table property controlling whether vended ADLS credentials are refreshed. +pub const ADLS_REFRESH_CREDENTIALS_ENABLED: &str = "adls.refresh-credentials-enabled"; /// The tenant-id. pub const ADLS_TENANT_ID: &str = "adls.tenant-id"; /// The client-id. diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index 70542c5fa3..e3401e03fe 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -255,6 +255,40 @@ pub enum StorageCredentialKind { S3(S3Credential), /// Google Cloud Storage credentials. Gcs(GcsCredential), + /// Azure Data Lake Storage credentials. + Azdls(AzdlsCredential), +} + +/// Temporary Azure Data Lake Storage credentials (a shared access signature). +#[derive(Clone)] +pub struct AzdlsCredential { + /// Shared access signature used to access Azure storage. + sas_token: String, +} + +impl AzdlsCredential { + /// Create an Azure Data Lake Storage credential. + pub fn new(sas_token: impl Into) -> Self { + Self { + sas_token: sas_token.into(), + } + } + + /// Return the Azure shared access signature. + pub fn sas_token(&self) -> &str { + &self.sas_token + } + + /// Consume this credential and return its shared access signature. + pub fn into_sas_token(self) -> String { + self.sas_token + } +} + +impl Debug for AzdlsCredential { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AzdlsCredential").finish_non_exhaustive() + } } /// Temporary Amazon S3 credentials. diff --git a/crates/storage/opendal/Cargo.toml b/crates/storage/opendal/Cargo.toml index e71a660434..a837206a4e 100644 --- a/crates/storage/opendal/Cargo.toml +++ b/crates/storage/opendal/Cargo.toml @@ -39,7 +39,11 @@ opendal-all = [ "opendal-hf", ] -opendal-azdls = ["opendal/services-azdls"] +opendal-azdls = [ + "opendal/services-azdls", + "reqsign-azure-storage", + "reqsign-core", +] opendal-fs = ["opendal/services-fs"] opendal-gcs = ["opendal/services-gcs", "reqsign-google", "reqsign-core"] opendal-hf = ["opendal/services-hf"] @@ -56,6 +60,7 @@ futures = { workspace = true } iceberg = { workspace = true } opendal = { workspace = true } reqsign-aws-v4 = { version = "3.0.0", optional = true } +reqsign-azure-storage = { version = "3.2.1", optional = true } reqsign-core = { version = "3.0.0", optional = true } reqsign-google = { version = "3.0.0", optional = true } serde = { workspace = true } diff --git a/crates/storage/opendal/public-api.txt b/crates/storage/opendal/public-api.txt index b253478947..4e7d89acea 100644 --- a/crates/storage/opendal/public-api.txt +++ b/crates/storage/opendal/public-api.txt @@ -4,6 +4,8 @@ pub use iceberg_storage_opendal::ProvideCredential pub enum iceberg_storage_opendal::OpenDalStorage pub iceberg_storage_opendal::OpenDalStorage::Azdls pub iceberg_storage_opendal::OpenDalStorage::Azdls::config: alloc::sync::Arc +pub iceberg_storage_opendal::OpenDalStorage::Azdls::credential_provider: core::option::Option> +pub iceberg_storage_opendal::OpenDalStorage::Azdls::sas_tokens: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Gcs pub iceberg_storage_opendal::OpenDalStorage::Gcs::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Gcs::credential_provider: core::option::Option> @@ -57,6 +59,17 @@ impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalStorageFacto pub fn iceberg_storage_opendal::OpenDalStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> +pub struct iceberg_storage_opendal::AzdlsSasTokens(_) +impl core::clone::Clone for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::clone(&self) -> iceberg_storage_opendal::AzdlsSasTokens +impl core::default::Default for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::default() -> iceberg_storage_opendal::AzdlsSasTokens +impl core::fmt::Debug for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result +impl serde_core::ser::Serialize for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer +impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::AzdlsSasTokens +pub fn iceberg_storage_opendal::AzdlsSasTokens::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg_storage_opendal::CustomAwsCredentialLoader(_) impl iceberg_storage_opendal::CustomAwsCredentialLoader pub fn iceberg_storage_opendal::CustomAwsCredentialLoader::new(provider: impl reqsign_core::api::ProvideCredential + 'static) -> Self diff --git a/crates/storage/opendal/src/azdls.rs b/crates/storage/opendal/src/azdls.rs index 6d39320df8..44f036f2b6 100644 --- a/crates/storage/opendal/src/azdls.rs +++ b/crates/storage/opendal/src/azdls.rs @@ -18,18 +18,26 @@ use std::collections::HashMap; use std::fmt::Display; use std::str::FromStr; +use std::sync::Arc; use iceberg::io::{ ADLS_ACCOUNT_KEY, ADLS_ACCOUNT_NAME, ADLS_AUTHORITY_HOST, ADLS_CLIENT_ID, ADLS_CLIENT_SECRET, - ADLS_CONNECTION_STRING, ADLS_SAS_TOKEN, ADLS_TENANT_ID, + ADLS_CONNECTION_STRING, ADLS_SAS_TOKEN, ADLS_SAS_TOKEN_PREFIX, ADLS_TENANT_ID, + StorageCredentialKind, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; use opendal::Configurator; use opendal::services::AzdlsConfig; +use reqsign_azure_storage::Credential as AzureCredential; +use reqsign_core::{ + Context, Error as ReqsignError, ProvideCredential, ProvideCredentialChain, + Result as ReqsignResult, +}; use serde::{Deserialize, Serialize}; use url::Url; -use crate::utils::from_opendal_error; +use crate::DynamicCredentialScope; +use crate::utils::{from_opendal_error, system_time_to_timestamp, validate_credential_prefix}; /// Local version of `ensure_data_valid` macro since the iceberg crate's macro /// uses `$crate::error::Error` paths that don't resolve from external crates @@ -84,6 +92,38 @@ pub(crate) fn azdls_config_parse(mut properties: HashMap) -> Res Ok(config) } +/// Account-specific ADLS SAS tokens supplied through Java-compatible storage +/// properties. +#[derive(Clone, Default, Serialize, Deserialize)] +pub struct AzdlsSasTokens(HashMap); + +impl AzdlsSasTokens { + pub(crate) fn from_properties(properties: &HashMap) -> Self { + Self( + properties + .iter() + .filter_map(|(key, value)| { + let account = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; + (!account.is_empty() && !value.is_empty()) + .then(|| (account.to_string(), value.clone())) + }) + .collect(), + ) + } + + fn for_path(&self, path: &AzureStoragePath) -> Option<&str> { + self.0.get(&path.account_name).map(String::as_str) + } +} + +impl std::fmt::Debug for AzdlsSasTokens { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AzdlsSasTokens") + .field("account_count", &self.0.len()) + .finish_non_exhaustive() + } +} + /// Builds an OpenDAL operator from the AzdlsConfig and path. /// /// The path is expected to include the scheme in a format like: @@ -91,11 +131,21 @@ pub(crate) fn azdls_config_parse(mut properties: HashMap) -> Res pub(crate) fn azdls_create_operator<'a>( absolute_path: &'a str, config: &AzdlsConfig, + sas_tokens: &AzdlsSasTokens, + credential_provider: &Option>, + credential_scope: Option<&DynamicCredentialScope>, ) -> Result<(opendal::Operator, &'a str)> { let path = absolute_path.parse::()?; match_path_with_config(&path, config)?; - let op = azdls_config_build(config, &path)?; + let op = azdls_config_build( + config, + &path, + sas_tokens, + credential_provider, + absolute_path, + credential_scope, + )?; // Paths to files in ADLS tend to be written in fully qualified form, // including their filesystem and account name. @@ -110,8 +160,9 @@ pub(crate) fn azdls_create_operator<'a>( /// Note that `abf[s]` and `wasb[s]` variants have different implications: /// - `abfs[s]` is used to refer to files in ADLS Gen2, backed by blob storage; /// paths are expected to contain the `dfs` storage service. -/// - `wasb[s]` is used to refer to files in Blob Storage directly; paths are -/// expected to contain the `blob` storage service. +/// - `wasb[s]` is accepted for compatibility with Blob Storage locations; +/// paths contain the `blob` storage service, but operations still use the +/// ADLS Gen2 `dfs` endpoint, matching Iceberg Java. #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub enum AzureStorageScheme { Abfs, @@ -121,12 +172,11 @@ pub enum AzureStorageScheme { } impl AzureStorageScheme { - // Returns the respective encrypted or plain-text HTTP scheme. + // Iceberg Java accepts the non-secure aliases for compatibility but still + // connects over TLS. SAS tokens are query parameters and must not be sent + // over a plaintext connection. pub fn as_http_scheme(&self) -> &str { - match self { - AzureStorageScheme::Abfs | AzureStorageScheme::Wasb => "http", - AzureStorageScheme::Abfss | AzureStorageScheme::Wasbs => "https", - } + "https" } } @@ -192,7 +242,14 @@ pub(crate) fn match_path_with_config(path: &AzureStoragePath, config: &AzdlsConf Ok(()) } -fn azdls_config_build(config: &AzdlsConfig, path: &AzureStoragePath) -> Result { +fn azdls_config_build( + config: &AzdlsConfig, + path: &AzureStoragePath, + sas_tokens: &AzdlsSasTokens, + credential_provider: &Option>, + absolute_path: &str, + credential_scope: Option<&DynamicCredentialScope>, +) -> Result { let mut builder = config.clone().into_builder(); if config.endpoint.is_none() { @@ -201,9 +258,92 @@ fn azdls_config_build(config: &AzdlsConfig, path: &AzureStoragePath) -> Result, + path: String, + credential_scope: Option, +} + +impl VendedAzdlsCredentialProvider { + fn new( + provider: Arc, + path: String, + credential_scope: Option, + ) -> Self { + Self { + provider, + path, + credential_scope, + } + } +} + +impl std::fmt::Debug for VendedAzdlsCredentialProvider { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("VendedAzdlsCredentialProvider") + .field("path", &self.path) + .field("credential_scope", &self.credential_scope) + .finish_non_exhaustive() + } +} + +impl ProvideCredential for VendedAzdlsCredentialProvider { + type Credential = AzureCredential; + + async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { + let credential = self + .provider + .load_credential(&self.path) + .await + .map_err(|error| { + ReqsignError::unexpected("failed to load vended ADLS credential").with_source(error) + })?; + validate_credential_prefix( + &self.path, + credential.prefix(), + self.credential_scope.as_ref(), + )?; + + let expires_at = credential + .expires_at() + .map(system_time_to_timestamp) + .transpose()?; + match credential.into_kind() { + StorageCredentialKind::Azdls(azdls) => { + let sas_token = azdls.into_sas_token(); + Ok(Some(match expires_at { + Some(expires_at) => { + AzureCredential::with_sas_token_expires_at(&sas_token, expires_at) + } + None => AzureCredential::with_sas_token(&sas_token), + })) + } + _ => Err(ReqsignError::unexpected( + "ADLS storage received a non-ADLS credential from the provider", + )), + } + } +} + /// Represents a fully qualified path to blob/ file in Azure Storage. #[derive(Debug, PartialEq)] pub(crate) struct AzureStoragePath { @@ -265,6 +405,15 @@ impl FromStr for AzureStoragePath { } } +pub(crate) fn azdls_batch_key(absolute_path: &str) -> Option { + absolute_path.parse::().ok().map(|path| { + format!( + "{}://{}@{}.{}", + path.scheme, path.filesystem, path.account_name, path.endpoint_suffix + ) + }) +} + fn parse_azure_storage_endpoint(url: &Url) -> Result<(&str, &str, &str)> { let host = url.host_str().ok_or(Error::new( ErrorKind::DataInvalid, @@ -318,10 +467,33 @@ fn validate_storage_and_scheme( #[cfg(test)] mod tests { use std::collections::HashMap; - + use std::sync::Arc; + use std::time::{Duration, SystemTime}; + + use async_trait::async_trait; + use iceberg::Result; + use iceberg::io::{ + ADLS_SAS_TOKEN_PREFIX, AzdlsCredential, StorageCredential, StorageCredentialKind, + StorageCredentialProvider, + }; use opendal::services::AzdlsConfig; + use reqsign_azure_storage::Credential as AzureCredential; + use reqsign_core::{Context, ProvideCredential}; - use super::{AzureStoragePath, AzureStorageScheme, azdls_config_parse, azdls_create_operator}; + use super::{ + AzdlsSasTokens, AzureStoragePath, AzureStorageScheme, VendedAzdlsCredentialProvider, + azdls_batch_key, azdls_config_parse, azdls_create_operator, + }; + + #[derive(Debug)] + struct FixedCredentialProvider(StorageCredential); + + #[async_trait] + impl StorageCredentialProvider for FixedCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(self.0.clone()) + } + } #[test] fn test_azdls_config_parse() { @@ -464,7 +636,8 @@ mod tests { ]; for (name, input, expected) in test_cases { - let result = azdls_create_operator(input.0, &input.1); + let result = + azdls_create_operator(input.0, &input.1, &AzdlsSasTokens::default(), &None, None); match expected { Some((expected_filesystem, expected_path)) => { assert!(result.is_ok(), "Test case {name} failed: {result:?}"); @@ -480,6 +653,87 @@ mod tests { } } + #[tokio::test] + async fn vended_provider_returns_expiring_sas_credential() { + let path = "abfss://container@account.dfs.core.windows.net/table/data.parquet"; + let expires_at = SystemTime::now() + Duration::from_secs(3600); + let credential = StorageCredential::new(StorageCredentialKind::Azdls( + AzdlsCredential::new("sv=2026&sig=secret"), + )) + .with_prefix("abfss://container@account.dfs.core.windows.net/table") + .with_expiration(expires_at); + let provider = VendedAzdlsCredentialProvider::new( + Arc::new(FixedCredentialProvider(credential)), + path.to_string(), + None, + ); + + let credential = provider + .provide_credential(&Context::new()) + .await + .unwrap() + .unwrap(); + match credential { + AzureCredential::SasToken { + token, + expires_at: actual_expires_at, + } => { + assert_eq!(token, "sv=2026&sig=secret"); + assert_eq!( + actual_expires_at, + Some(super::system_time_to_timestamp(expires_at).unwrap()) + ); + } + other => panic!("expected SAS token, got {other:?}"), + } + } + + #[test] + fn account_specific_sas_tokens_are_selected_by_storage_account() { + let properties = HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}first"), + "sv=2026&sig=first".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}second"), + "sv=2026&sig=second".to_string(), + ), + ]); + let sas_tokens = AzdlsSasTokens::from_properties(&properties); + let path = "abfss://container@second.dfs.core.windows.net/table/data.parquet" + .parse::() + .unwrap(); + + assert_eq!(sas_tokens.for_path(&path), Some("sv=2026&sig=second")); + } + + #[tokio::test] + async fn vended_provider_rejects_mismatched_prefix() { + let credential = StorageCredential::new(StorageCredentialKind::Azdls( + AzdlsCredential::new("sv=2026&sig=secret"), + )) + .with_prefix("abfss://other@account.dfs.core.windows.net/table") + .with_expiration(SystemTime::now() + Duration::from_secs(3600)); + let provider = VendedAzdlsCredentialProvider::new( + Arc::new(FixedCredentialProvider(credential)), + "abfss://container@account.dfs.core.windows.net/table/data.parquet".to_string(), + None, + ); + + assert!(provider.provide_credential(&Context::new()).await.is_err()); + } + + #[test] + fn batch_key_distinguishes_azure_filesystems_and_schemes() { + let first = azdls_batch_key("abfss://first@account.dfs.core.windows.net/table/a.parquet"); + let second = azdls_batch_key("abfss://second@account.dfs.core.windows.net/table/b.parquet"); + let blob = azdls_batch_key("wasbs://first@account.blob.core.windows.net/table/c.parquet"); + + assert_ne!(first, second); + assert_ne!(first, blob); + } + #[test] fn test_azure_storage_path_parse() { let test_cases = vec![ @@ -540,7 +794,7 @@ mod tests { "https://myaccount.dfs.core.windows.net", ), ( - "abfs uses http", + "abfs uses https for Java compatibility and SAS security", AzureStoragePath { scheme: AzureStorageScheme::Abfs, filesystem: "myfs".to_string(), @@ -548,12 +802,12 @@ mod tests { endpoint_suffix: "core.windows.net".to_string(), path: "/path/to/file.parquet".to_string(), }, - "http://myaccount.dfs.core.windows.net", + "https://myaccount.dfs.core.windows.net", ), ( "wasbs uses https and dfs", AzureStoragePath { - scheme: AzureStorageScheme::Abfss, + scheme: AzureStorageScheme::Wasbs, filesystem: "myfs".to_string(), account_name: "myaccount".to_string(), endpoint_suffix: "core.windows.net".to_string(), diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index f17b927d23..20ddc24907 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -102,11 +102,10 @@ pub(crate) fn gcs_config_build( let mut cfg = cfg.clone(); cfg.bucket = bucket.to_string(); - // When a catalog-supplied provider is present, make it the sole credential - // source. `reqsign_google` only prepends a custom provider (unlike S3, which - // replaces the chain) and its chain continues to the next provider on error, - // so without this a failed refresh would silently fall back to the stale seed - // token or ambient GCP credentials. + // `reqsign_google` continues to the next provider even when a provider returns + // an error, and OpenDAL prepends custom providers to its default chain. Disable + // every other configured and ambient source so the catalog provider is + // effectively the sole source and refresh failures cannot silently fall back. let credential_provider = credential_provider .as_ref() .filter(|provider| provider.supports_path(path)); @@ -119,6 +118,8 @@ pub(crate) fn gcs_config_build( } cfg.token = None; cfg.credential = None; + cfg.credential_path = None; + cfg.service_account = None; cfg.disable_vm_metadata = true; cfg.disable_config_load = true; } diff --git a/crates/storage/opendal/src/lib.rs b/crates/storage/opendal/src/lib.rs index fad0106f04..8791bd47d8 100644 --- a/crates/storage/opendal/src/lib.rs +++ b/crates/storage/opendal/src/lib.rs @@ -46,6 +46,7 @@ use utils::from_opendal_error; cfg_if! { if #[cfg(feature = "opendal-azdls")] { mod azdls; + pub use azdls::AzdlsSasTokens; use azdls::*; use opendal::services::AzdlsConfig; } @@ -177,6 +178,8 @@ impl StorageFactory for OpenDalStorageFactory { OpenDalStorageFactory::S3 { .. } => true, #[cfg(feature = "opendal-gcs")] OpenDalStorageFactory::Gcs => true, + #[cfg(feature = "opendal-azdls")] + OpenDalStorageFactory::Azdls => true, _ => false, }; if credential_provider.is_some() && !supports_credential_provider { @@ -213,6 +216,8 @@ impl StorageFactory for OpenDalStorageFactory { #[cfg(feature = "opendal-azdls")] OpenDalStorageFactory::Azdls => Ok(Arc::new(OpenDalStorage::Azdls { config: azdls_config_parse(config.props().clone())?.into(), + sas_tokens: Arc::new(AzdlsSasTokens::from_properties(config.props())), + credential_provider, })), #[cfg(feature = "opendal-hf")] OpenDalStorageFactory::Hf => Ok(Arc::new(OpenDalStorage::Hf { @@ -290,6 +295,12 @@ pub enum OpenDalStorage { Azdls { /// Azure DLS configuration. config: Arc, + /// Account-specific SAS tokens supplied by Java-compatible properties. + #[serde(default)] + sas_tokens: Arc, + /// Provider of refreshable vended credentials, supplied by the catalog. + #[serde(skip)] + credential_provider: Option>, }, /// HuggingFace Hub storage variant. /// @@ -448,7 +459,17 @@ impl OpenDalStorage { } } #[cfg(feature = "opendal-azdls")] - OpenDalStorage::Azdls { config } => azdls_create_operator(path, config)?, + OpenDalStorage::Azdls { + config, + sas_tokens, + credential_provider, + } => azdls_create_operator( + path, + config, + sas_tokens, + credential_provider, + credential_scope, + )?, #[cfg(feature = "opendal-hf")] OpenDalStorage::Hf { config } => hf_config_build(config, path)?, #[cfg(all( @@ -485,10 +506,13 @@ impl OpenDalStorage { /// /// For most backends the URL host (bucket name) is sufficient. For HF the host /// encodes the repo type, not the repo identity, so a more specific key is used. + #[allow(unreachable_patterns)] fn batch_key_for_path(&self, path: &str) -> String { match self { #[cfg(feature = "opendal-hf")] OpenDalStorage::Hf { .. } => hf_batch_key(path), + #[cfg(feature = "opendal-azdls")] + OpenDalStorage::Azdls { .. } => azdls_batch_key(path).unwrap_or_default(), _ => url::Url::parse(path) .ok() .and_then(|u| u.host_str().map(|s| s.to_string())) @@ -513,6 +537,11 @@ impl OpenDalStorage { credential_provider: Some(provider), .. } => Some(provider), + #[cfg(feature = "opendal-azdls")] + OpenDalStorage::Azdls { + credential_provider: Some(provider), + .. + } => Some(provider), _ => None, })?; provider.supports_path(path).then_some(provider) @@ -626,7 +655,7 @@ impl OpenDalStorage { } } #[cfg(feature = "opendal-azdls")] - OpenDalStorage::Azdls { config } => { + OpenDalStorage::Azdls { config, .. } => { let azure_path = path.parse::()?; match_path_with_config(&azure_path, config)?; let relative_path_len = azure_path.path.len(); @@ -805,13 +834,22 @@ impl FileWrite for OpenDalWriter { #[cfg(test)] mod tests { + #[allow(unused_imports)] use super::*; - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + all(feature = "opendal-memory", feature = "opendal-azdls") + ))] #[derive(Debug)] struct AlwaysSupportedCredentialProvider; - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + all(feature = "opendal-memory", feature = "opendal-azdls") + ))] #[async_trait] impl StorageCredentialProvider for AlwaysSupportedCredentialProvider { async fn load_credential(&self, path: &str) -> Result { @@ -869,7 +907,11 @@ mod tests { #[cfg(all( feature = "opendal-memory", - any(feature = "opendal-s3", feature = "opendal-gcs") + any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ) ))] #[test] fn test_factory_rejects_credentials_for_unsupported_backend() { @@ -1098,6 +1140,8 @@ mod tests { endpoint: Some("https://myaccount.dfs.core.windows.net".to_string()), ..Default::default() }), + sas_tokens: Arc::new(AzdlsSasTokens::default()), + credential_provider: None, }; assert_eq!( diff --git a/crates/storage/opendal/src/resolving.rs b/crates/storage/opendal/src/resolving.rs index 4201631f23..de61890a81 100644 --- a/crates/storage/opendal/src/resolving.rs +++ b/crates/storage/opendal/src/resolving.rs @@ -80,25 +80,35 @@ fn extract_scheme(path: &str) -> Result<&'static str> { parse_scheme(url.scheme()) } -#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] fn supports_dynamic_credentials(scheme: &str) -> bool { match scheme { #[cfg(feature = "opendal-s3")] "s3" => true, #[cfg(feature = "opendal-gcs")] "gcs" => true, + #[cfg(feature = "opendal-azdls")] + "azdls" => true, _ => false, } } /// Build an [`OpenDalStorage`] variant for the given scheme and config properties. +#[allow(unused_variables)] fn build_storage_for_scheme( scheme: &'static str, props: &HashMap, #[cfg(feature = "opendal-s3")] customized_credential_load: &Option, - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] credential_provider: &Option< - Arc, - >, + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] + credential_provider: &Option>, ) -> Result { match scheme { #[cfg(feature = "opendal-s3")] @@ -130,6 +140,8 @@ fn build_storage_for_scheme( let config = crate::azdls::azdls_config_parse(props.clone())?; Ok(OpenDalStorage::Azdls { config: Arc::new(config), + sas_tokens: Arc::new(crate::azdls::AzdlsSasTokens::from_properties(props)), + credential_provider: credential_provider.clone(), }) } #[cfg(feature = "opendal-fs")] @@ -221,7 +233,11 @@ impl StorageFactory for OpenDalResolvingStorageFactory { config: &StorageConfig, credential_provider: Option>, ) -> Result> { - #[cfg(not(any(feature = "opendal-s3", feature = "opendal-gcs")))] + #[cfg(not(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + )))] if credential_provider.is_some() { return Err(Error::new( ErrorKind::FeatureUnsupported, @@ -234,7 +250,11 @@ impl StorageFactory for OpenDalResolvingStorageFactory { storages: RwLock::new(HashMap::new()), #[cfg(feature = "opendal-s3")] customized_credential_load: self.customized_credential_load.clone(), - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] credential_provider, })) } @@ -258,7 +278,11 @@ pub struct OpenDalResolvingStorage { #[serde(skip)] customized_credential_load: Option, /// Provider of refreshable vended credentials. - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] #[serde(skip)] credential_provider: Option>, } @@ -279,7 +303,11 @@ impl OpenDalResolvingStorage { fn resolve(&self, path: &str) -> Result> { let scheme = extract_scheme(path)?; - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] if self .credential_provider .as_ref() @@ -321,7 +349,11 @@ impl OpenDalResolvingStorage { &self.props, #[cfg(feature = "opendal-s3")] &self.customized_credential_load, - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] &self.credential_provider, )?; let storage = Arc::new(storage); @@ -383,6 +415,7 @@ impl Storage for OpenDalResolvingStorage { Ok(()) } + #[allow(unreachable_code)] fn new_input(&self, path: &str) -> Result { Ok(InputFile::new( Arc::new(self.resolve(path)?.as_ref().clone()), @@ -390,6 +423,7 @@ impl Storage for OpenDalResolvingStorage { )) } + #[allow(unreachable_code)] fn new_output(&self, path: &str) -> Result { Ok(OutputFile::new( Arc::new(self.resolve(path)?.as_ref().clone()), @@ -400,11 +434,28 @@ impl Storage for OpenDalResolvingStorage { #[cfg(test)] mod tests { + #[allow(unused_imports)] use super::*; + #[cfg(any( + feature = "opendal-azdls", + not(any(feature = "opendal-s3", feature = "opendal-gcs")), + all( + feature = "opendal-memory", + any(feature = "opendal-s3", feature = "opendal-gcs") + ) + ))] #[derive(Debug)] struct AllPathsCredentialProvider; + #[cfg(any( + feature = "opendal-azdls", + not(any(feature = "opendal-s3", feature = "opendal-gcs")), + all( + feature = "opendal-memory", + any(feature = "opendal-s3", feature = "opendal-gcs") + ) + ))] #[async_trait] impl StorageCredentialProvider for AllPathsCredentialProvider { async fn load_credential(&self, _path: &str) -> Result { @@ -412,7 +463,11 @@ mod tests { } } - #[cfg(not(any(feature = "opendal-s3", feature = "opendal-gcs")))] + #[cfg(not(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + )))] #[test] fn test_factory_rejects_credentials_without_compatible_backend() { let error = OpenDalResolvingStorageFactory::new() @@ -459,8 +514,8 @@ mod tests { /// calls that don't actually hit any backend. #[cfg(any( feature = "opendal-s3", - feature = "opendal-gcs", - feature = "opendal-azdls" + feature = "opendal-azdls", + all(feature = "opendal-memory", feature = "opendal-gcs") ))] fn empty_resolving_storage() -> OpenDalResolvingStorage { OpenDalResolvingStorage { @@ -468,7 +523,11 @@ mod tests { storages: RwLock::new(HashMap::new()), #[cfg(feature = "opendal-s3")] customized_credential_load: None, - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ))] credential_provider: None, } } @@ -491,7 +550,11 @@ mod tests { #[cfg(all( feature = "opendal-memory", - any(feature = "opendal-s3", feature = "opendal-gcs") + any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ) ))] #[test] fn test_resolver_rejects_credentials_for_unsupported_backend() { @@ -507,7 +570,8 @@ mod tests { #[cfg(feature = "opendal-azdls")] #[test] fn test_resolve_azdls_aliases_share_instance() { - let storage = empty_resolving_storage(); + let mut storage = empty_resolving_storage(); + storage.credential_provider = Some(Arc::new(AllPathsCredentialProvider)); let path_for = |scheme: &str| { format!("{scheme}://myfs@myaccount.dfs.core.windows.net/path/to/file.parquet") diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index ddb4f08b2d..1f0830d6eb 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -15,9 +15,14 @@ // specific language governing permissions and limitations // under the License. -#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] use crate::DynamicCredentialScope; +#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] pub(crate) fn is_truthy(value: &str) -> bool { ["true", "t", "1", "on"].contains(&value.to_lowercase().as_str()) } @@ -34,7 +39,11 @@ pub(crate) fn from_opendal_error(e: opendal::Error) -> iceberg::Error { /// Convert a [`SystemTime`](std::time::SystemTime) credential expiry into the /// `reqsign` [`Timestamp`](reqsign_core::time::Timestamp) used on backend /// credential types (e.g. `AwsCredential::expires_in`, `google::Token::expires_at`). -#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] pub(crate) fn system_time_to_timestamp( time: std::time::SystemTime, ) -> reqsign_core::Result { @@ -55,7 +64,11 @@ pub(crate) fn system_time_to_timestamp( /// Validate that a provider's declared credential prefix covers the path for /// which the backend requested the credential. -#[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] pub(crate) fn validate_credential_prefix( path: &str, prefix: Option<&str>, From f3fa42d0da4c60074964225f21bdec0a5fcf272e Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Tue, 29 Sep 2026 17:32:14 +0100 Subject: [PATCH 6/7] fix(rest): make vended credential refresh portable and robust Rebuild the refresh provider after FileIO serialization and fix several correctness issues found in review. - Serialize the credential provider as a StorageCredentialProviderFactory that rebuilds it from the FileIO properties and connects to the catalog lazily, like Java's VendedCredentialsProvider. Catalogs with an injected AuthManager refuse to serialize; FileIO::without_credential_provider drops the provider instead. - Match credential prefixes on whole path segments and treat scheme aliases (s3a/s3n, gcs, abfs/wasb) as equal, via StorageCredential::covers. - Merge refreshed credentials prefix by prefix, keeping an unexpired cached credential when the response has no valid replacement. - Look up bulk-delete credentials by the batch scope, so a refresh during a delete can no longer abort it, and share one reqsign adapter across S3, GCS and ADLS. - Cap the refresh buffer at half the credential lifetime, so short-lived credentials are not re-fetched on every operation. - Ignore credential providers in storage factories that cannot use them instead of failing every operation, matching main and Java. - Accept ADLS SAS tokens keyed by host (adls.sas-token.) as Java and Polaris send them, and accept explicit http endpoints for abfs/wasb. - Move auth manager selection to auth::load_auth_manager and auth property helpers onto RestCatalogConfig. - Mark StorageCredentialKind and OpenDalStorage #[non_exhaustive]. --- Cargo.lock | 1 + crates/catalog/rest/Cargo.toml | 1 + crates/catalog/rest/public-api.txt | 1 - crates/catalog/rest/src/auth/mod.rs | 29 +- crates/catalog/rest/src/auth/oauth2.rs | 4 + crates/catalog/rest/src/catalog.rs | 255 ++--- crates/catalog/rest/src/client.rs | 10 +- crates/catalog/rest/src/credential.rs | 1251 ++++++++++++++--------- crates/iceberg/public-api.txt | 8 +- crates/iceberg/src/io/file_io.rs | 140 ++- crates/iceberg/src/io/storage/mod.rs | 132 ++- crates/storage/opendal/public-api.txt | 12 +- crates/storage/opendal/src/azdls.rs | 184 ++-- crates/storage/opendal/src/gcs.rs | 168 +-- crates/storage/opendal/src/lib.rs | 165 +-- crates/storage/opendal/src/resolving.rs | 77 +- crates/storage/opendal/src/s3.rs | 164 +-- crates/storage/opendal/src/utils.rs | 194 +++- 18 files changed, 1588 insertions(+), 1208 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 57febb3141..7334be0337 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2738,6 +2738,7 @@ dependencies = [ "tokio", "tracing", "typed-builder", + "typetag", "uuid", ] diff --git a/crates/catalog/rest/Cargo.toml b/crates/catalog/rest/Cargo.toml index f035f0e0ef..0f7a609786 100644 --- a/crates/catalog/rest/Cargo.toml +++ b/crates/catalog/rest/Cargo.toml @@ -43,6 +43,7 @@ serde_json = { workspace = true } tokio = { workspace = true } tracing = { workspace = true } typed-builder = { workspace = true } +typetag = { workspace = true } uuid = { workspace = true, features = ["v4"] } [dev-dependencies] diff --git a/crates/catalog/rest/public-api.txt b/crates/catalog/rest/public-api.txt index d2d8485398..eb4638d809 100644 --- a/crates/catalog/rest/public-api.txt +++ b/crates/catalog/rest/public-api.txt @@ -397,7 +397,6 @@ pub const iceberg_catalog_rest::AUTH_TYPE_NONE: &str pub const iceberg_catalog_rest::AUTH_TYPE_OAUTH2: &str pub const iceberg_catalog_rest::REST_CATALOG_PROP_AUTH_TYPE: &str pub const iceberg_catalog_rest::REST_CATALOG_PROP_DISABLE_HEADER_REDACTION: &str -pub const iceberg_catalog_rest::REST_CATALOG_PROP_SCAN_PLAN_ID: &str pub const iceberg_catalog_rest::REST_CATALOG_PROP_URI: &str pub const iceberg_catalog_rest::REST_CATALOG_PROP_WAREHOUSE: &str pub trait iceberg_catalog_rest::AuthManager: core::fmt::Debug + core::marker::Send + core::marker::Sync diff --git a/crates/catalog/rest/src/auth/mod.rs b/crates/catalog/rest/src/auth/mod.rs index a9352b1845..f6bf30ca2c 100644 --- a/crates/catalog/rest/src/auth/mod.rs +++ b/crates/catalog/rest/src/auth/mod.rs @@ -25,9 +25,10 @@ use std::fmt::Debug; use std::sync::Arc; use async_trait::async_trait; -use iceberg::{Result, TableIdent}; +use iceberg::{Error, ErrorKind, Result, TableIdent}; pub use oauth2::OAuth2Manager; +use crate::catalog::{REST_CATALOG_PROP_AUTH_TYPE, RestCatalogConfig}; use crate::client::HttpClient; use crate::request::HttpRequest; @@ -36,6 +37,32 @@ pub const AUTH_TYPE_NONE: &str = "none"; /// `rest.auth.type` value selecting OAuth2 token authentication. pub const AUTH_TYPE_OAUTH2: &str = "oauth2"; +/// Builds the auth manager selected by the `rest.auth.type` configuration, +/// like Java's `AuthManagers.loadAuthManager`. +pub(crate) fn load_auth_manager(config: &RestCatalogConfig) -> Result> { + let auth_type = config.auth_type(); + // Java parity (`AuthManagers`): make the inference visible so users + // configure the type explicitly. + if auth_type == AUTH_TYPE_OAUTH2 && !config.has_explicit_auth_type() { + tracing::warn!( + "Inferring {REST_CATALOG_PROP_AUTH_TYPE}={AUTH_TYPE_OAUTH2} from the configured \ + OAuth properties; set it explicitly to avoid this warning" + ); + } + match auth_type.as_str() { + AUTH_TYPE_NONE => Ok(Arc::new(NoopAuthManager)), + AUTH_TYPE_OAUTH2 => Ok(Arc::new(OAuth2Manager::from_config(config)?)), + other => Err(Error::new( + ErrorKind::DataInvalid, + format!( + "unknown '{REST_CATALOG_PROP_AUTH_TYPE}': {other}; use \ + `RestSessionCatalogBuilder::with_auth_manager` or \ + `RestCatalogBuilder::with_auth_manager` to inject a custom auth manager" + ), + )), + } +} + /// Creates the [`AuthSession`]s used to authenticate REST catalog requests. /// /// A manager is exclusively scoped to one catalog and must not be reused by diff --git a/crates/catalog/rest/src/auth/oauth2.rs b/crates/catalog/rest/src/auth/oauth2.rs index 3e48a75f67..39a71b53d0 100644 --- a/crates/catalog/rest/src/auth/oauth2.rs +++ b/crates/catalog/rest/src/auth/oauth2.rs @@ -146,6 +146,10 @@ impl AuthManager for OAuth2Manager { Ok(Arc::new(self.session_from(client, props).await?)) } + /// Like Java, only a `token` in the table config overrides the parent + /// session; a table-level `credential` is ignored. Unlike Java, the token is + /// used as-is and never refreshed, and token-type exchange keys are ignored, + /// because this manager implements neither token refresh nor token exchange. async fn table_session( &self, _client: &HttpClient, diff --git a/crates/catalog/rest/src/catalog.rs b/crates/catalog/rest/src/catalog.rs index c77819c715..fb5056f069 100644 --- a/crates/catalog/rest/src/catalog.rs +++ b/crates/catalog/rest/src/catalog.rs @@ -39,11 +39,11 @@ use reqwest::{Client, Method, StatusCode, Url}; use tokio::sync::OnceCell; use typed_builder::TypedBuilder; -use crate::auth::{AUTH_TYPE_NONE, AUTH_TYPE_OAUTH2, AuthManager, NoopAuthManager, OAuth2Manager}; +use crate::auth::{AUTH_TYPE_NONE, AUTH_TYPE_OAUTH2, AuthManager, load_auth_manager}; use crate::client::{ HttpClient, deserialize_catalog_response, deserialize_unexpected_catalog_error, }; -use crate::credential::build_vended_credential_provider; +use crate::credential::{RestVendedCredentialProviderFactory, build_vended_credential_provider}; use crate::endpoint::{Endpoint, V1_NAMESPACE_EXISTS, V1_TABLE_EXISTS}; use crate::request::HttpRequest; use crate::response::HttpResponse; @@ -61,7 +61,7 @@ pub const REST_CATALOG_PROP_WAREHOUSE: &str = "warehouse"; /// false for security) pub const REST_CATALOG_PROP_DISABLE_HEADER_REDACTION: &str = "disable-header-redaction"; /// Identifier for a server-side scan plan associated with credential requests. -pub const REST_CATALOG_PROP_SCAN_PLAN_ID: &str = "rest.scan.plan-id"; +pub(crate) const REST_CATALOG_PROP_SCAN_PLAN_ID: &str = "rest.scan.plan-id"; /// Authentication scheme: `none` or `oauth2`. When unset, `oauth2` is used /// if a `token`, `credential` or `oauth2-server-uri` is configured, `none` /// otherwise. @@ -292,10 +292,52 @@ impl RestCatalogConfig { /// Returns true if the `disable-header-redaction` property is set to "true". /// Defaults to false for security (headers are redacted by default). pub(crate) fn disable_header_redaction(&self) -> bool { + disable_header_redaction_from_props(&self.props).unwrap_or(false) + } + + /// The configured auth scheme: explicit `rest.auth.type` (matched + /// case-insensitively) when set; otherwise `oauth2` when a `token`, + /// `credential` or `oauth2-server-uri` is configured (preserving + /// pre-`rest.auth.type` setups), `none` when none is. + pub(crate) fn auth_type(&self) -> String { self.props - .get(REST_CATALOG_PROP_DISABLE_HEADER_REDACTION) - .map(|v| v.eq_ignore_ascii_case("true")) - .unwrap_or(false) + .get(REST_CATALOG_PROP_AUTH_TYPE) + // Matched case-insensitively, as the other flag properties are. + .map(|auth_type| auth_type.to_ascii_lowercase()) + .unwrap_or_else(|| { + if self.token().is_some() + || self.credential().is_some() + || self.explicit_oauth2_server_uri().is_some() + { + AUTH_TYPE_OAUTH2.to_string() + } else { + AUTH_TYPE_NONE.to_string() + } + }) + } + + /// Whether `rest.auth.type` is set explicitly rather than inferred. + pub(crate) fn has_explicit_auth_type(&self) -> bool { + self.props.contains_key(REST_CATALOG_PROP_AUTH_TYPE) + } + + /// The properties handed to the [`AuthManager`], with the catalog `uri` + /// and `warehouse` made explicit. + pub(crate) fn auth_props(&self) -> HashMap { + // `oauth2-server-uri` stays absent unless explicitly configured, so an + // injected manager keeps its own endpoint. The resolved `uri` and + // `warehouse` ARE passed: the builder moved them off the props, and + // the built-in manager recomputes its token endpoint from the URI. + let mut props = self.props.clone(); + props.insert(REST_CATALOG_PROP_URI.to_string(), self.uri.clone()); + if let Some(warehouse) = &self.warehouse { + // A fallback only: after the handshake the merged props hold + // the resolved warehouse, server override included. + props + .entry(REST_CATALOG_PROP_WAREHOUSE.to_string()) + .or_insert_with(|| warehouse.clone()); + } + props } /// Merge the `RestCatalogConfig` with the a [`CatalogConfig`] (fetched from the REST server). @@ -365,6 +407,13 @@ pub(crate) fn extra_headers_from_props(props: &HashMap) -> Resul Ok(headers) } +/// The `disable-header-redaction` property, when set. +pub(crate) fn disable_header_redaction_from_props(props: &HashMap) -> Option { + props + .get(REST_CATALOG_PROP_DISABLE_HEADER_REDACTION) + .map(|value| value.eq_ignore_ascii_case("true")) +} + /// The default OAuth2 token endpoint for a catalog `uri`. pub(crate) fn default_token_endpoint(uri: &str) -> String { [uri, PATH_V1, "oauth", "tokens"].join("/") @@ -449,7 +498,7 @@ impl RestClient { let init_session = auth_manager .init_session( &http_client.without_auth_session(), - &Self::auth_props(user_config), + &user_config.auth_props(), ) .await?; Self::load_config( @@ -469,10 +518,7 @@ impl RestClient { // The manager is handed an unauthenticated client: its own // requests must not be signed by the session it is deriving. let session = auth_manager - .catalog_session( - &http_client.without_auth_session(), - &Self::auth_props(&config), - ) + .catalog_session(&http_client.without_auth_session(), &config.auth_props()) .await?; Ok(Self { @@ -494,25 +540,6 @@ impl RestClient { self.http_client.query_catalog(request).await } - /// The properties handed to the [`AuthManager`], with the catalog `uri` - /// and `warehouse` made explicit. - fn auth_props(config: &RestCatalogConfig) -> HashMap { - // `oauth2-server-uri` stays absent unless explicitly configured, so an - // injected manager keeps its own endpoint. The resolved `uri` and - // `warehouse` ARE passed: the builder moved them off the props, and - // the built-in manager recomputes its token endpoint from the URI. - let mut props = config.props.clone(); - props.insert(REST_CATALOG_PROP_URI.to_string(), config.uri.clone()); - if let Some(warehouse) = &config.warehouse { - // A fallback only: after the handshake the merged props hold - // the resolved warehouse, server override included. - props - .entry(REST_CATALOG_PROP_WAREHOUSE.to_string()) - .or_insert_with(|| warehouse.clone()); - } - props - } - /// Loads the runtime config from the server using `user_config`. /// /// It's required for a REST catalog to update its config after creation. @@ -764,56 +791,12 @@ impl RestSessionCatalog { } } - /// The configured auth scheme: explicit `rest.auth.type` (matched - /// case-insensitively) when set; otherwise `oauth2` when a `token`, - /// `credential` or `oauth2-server-uri` is configured (preserving - /// pre-`rest.auth.type` setups), `none` when none is. - fn auth_type(config: &RestCatalogConfig) -> String { - config - .props - .get(REST_CATALOG_PROP_AUTH_TYPE) - // Matched case-insensitively, as the other flag properties are. - .map(|auth_type| auth_type.to_ascii_lowercase()) - .unwrap_or_else(|| { - if config.token().is_some() - || config.credential().is_some() - || config.explicit_oauth2_server_uri().is_some() - { - AUTH_TYPE_OAUTH2.to_string() - } else { - AUTH_TYPE_NONE.to_string() - } - }) - } - /// Resolves the auth manager: a `with_auth_manager` override wins, /// otherwise one is built from the `rest.auth.type` configuration. fn resolve_auth_manager(&self) -> Result> { - if let Some(auth_manager) = &self.auth_manager { - return Ok(auth_manager.clone()); - } - let config = &self.user_config; - let auth_type = Self::auth_type(config); - // Java parity (`AuthManagers`): make the inference visible so users - // configure the type explicitly. - if auth_type == AUTH_TYPE_OAUTH2 && !config.props.contains_key(REST_CATALOG_PROP_AUTH_TYPE) - { - tracing::warn!( - "Inferring {REST_CATALOG_PROP_AUTH_TYPE}={AUTH_TYPE_OAUTH2} from the configured \ - OAuth properties; set it explicitly to avoid this warning" - ); - } - match auth_type.as_str() { - AUTH_TYPE_NONE => Ok(Arc::new(NoopAuthManager)), - AUTH_TYPE_OAUTH2 => Ok(Arc::new(OAuth2Manager::from_config(config)?)), - other => Err(Error::new( - ErrorKind::DataInvalid, - format!( - "unknown '{REST_CATALOG_PROP_AUTH_TYPE}': {other}; use \ - `RestSessionCatalogBuilder::with_auth_manager` or \ - `RestCatalogBuilder::with_auth_manager` to inject a custom auth manager" - ), - )), + match &self.auth_manager { + Some(auth_manager) => Ok(auth_manager.clone()), + None => load_auth_manager(&self.user_config), } } @@ -849,17 +832,20 @@ impl RestSessionCatalog { } } + /// Builds the FileIO for `table` from the catalog properties and, when + /// loaded from a table response, its `config` overridden by the user + /// properties. async fn load_file_io( &self, table: &TableIdent, metadata_location: Option<&str>, - extra_config: Option>, - table_auth_config: Option>, + table_config: Option>, ) -> Result { let client = self.client().await?; let mut props = client.config.props.clone(); - if let Some(config) = extra_config { - props.extend(config); + if let Some(table_config) = &table_config { + props.extend(table_config.clone()); + props.extend(self.user_config.props.clone()); } // If the warehouse is a logical identifier instead of a URL we don't want @@ -889,13 +875,18 @@ impl RestSessionCatalog { // If the catalog vends refreshable credentials for this table's storage, // attach a provider so the backend re-fetches them before they expire. + // Only catalog authentication resolved from the properties can be + // rebuilt after FileIO serialization. let credential_provider = build_vended_credential_provider( - Arc::new(client.http_client.clone()), + &client.http_client, client.auth_manager.as_ref(), - table, - &client.config.uri, + RestVendedCredentialProviderFactory::new( + &client.config.uri, + table.clone(), + table_config.unwrap_or_default(), + ), &props, - table_auth_config.as_ref(), + self.auth_manager.is_none(), ) .await?; @@ -1204,20 +1195,8 @@ impl SessionCatalog for RestSessionCatalog { "Metadata location missing in `create_table` response!", ))?; - let table_config = response.config; - let config = table_config - .clone() - .into_iter() - .chain(self.user_config.props.clone()) - .collect(); - let file_io = self - .load_file_io( - &table_ident, - Some(metadata_location), - Some(config), - Some(table_config), - ) + .load_file_io(&table_ident, Some(metadata_location), Some(response.config)) .await?; let mut table_builder = Table::builder() @@ -1274,19 +1253,11 @@ impl SessionCatalog for RestSessionCatalog { } }; - let table_config = response.config; - let config = table_config - .clone() - .into_iter() - .chain(self.user_config.props.clone()) - .collect(); - let file_io = self .load_file_io( table_ident, response.metadata_location.as_deref(), - Some(config), - Some(table_config), + Some(response.config), ) .await?; @@ -1426,19 +1397,8 @@ impl SessionCatalog for RestSessionCatalog { "Metadata location missing in `register_table` response!", ))?; - let table_config = response.config; - let config = table_config - .clone() - .into_iter() - .chain(self.user_config.props.clone()) - .collect(); let file_io = self - .load_file_io( - table_ident, - Some(metadata_location), - Some(config), - Some(table_config), - ) + .load_file_io(table_ident, Some(metadata_location), Some(response.config)) .await?; let mut table_builder = Table::builder() @@ -1518,12 +1478,7 @@ impl SessionCatalog for RestSessionCatalog { }; let file_io = self - .load_file_io( - commit.identifier(), - Some(&response.metadata_location), - None, - None, - ) + .load_file_io(commit.identifier(), Some(&response.metadata_location), None) .await?; let mut table_builder = Table::builder() @@ -1733,7 +1688,7 @@ mod tests { use uuid::uuid; use super::*; - use crate::auth::AuthSession; + use crate::auth::{AuthSession, NoopAuthManager, OAuth2Manager}; use crate::request::HttpRequest; fn test_catalog(config: RestCatalogConfig) -> RestSessionCatalog { @@ -4198,6 +4153,54 @@ mod tests { load_table_mock.assert_async().await } + #[tokio::test] + async fn test_injected_auth_manager_file_io_serializes_without_provider_only() { + let mut server = Server::new_async().await; + let config_mock = create_config_mock(&mut server).await; + let mut load_response: serde_json::Value = serde_json::from_reader(BufReader::new( + File::open(format!( + "{}/testdata/load_table_response.json", + env!("CARGO_MANIFEST_DIR") + )) + .unwrap(), + )) + .unwrap(); + load_response["config"]["client.refresh-credentials-endpoint"] = + json!("/v1/namespaces/ns1/tables/test1/credentials"); + let load_table_mock = server + .mock("GET", "/v1/namespaces/ns1/tables/test1") + .with_status(200) + .with_body(load_response.to_string()) + .create_async() + .await; + + let catalog = RestCatalog::new( + SessionContext::empty(), + RestCatalogConfig::builder().uri(server.url()).build(), + Some(Box::new(NoopAuthManager)), + Some(Arc::new(LocalFsStorageFactory)), + Runtime::current(), + None, + ); + let table = catalog + .load_table(&TableIdent::from_strs(["ns1", "test1"]).unwrap()) + .await + .unwrap(); + + let error = table.file_io().serialize_all().unwrap_err(); + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + assert!( + table + .file_io() + .without_credential_provider() + .serialize_all() + .is_ok() + ); + + config_mock.assert_async().await; + load_table_mock.assert_async().await; + } + #[tokio::test] async fn test_update_table_404() { let mut server = Server::new_async().await; diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index 2d90ead06f..7a9853fe45 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -26,7 +26,7 @@ use serde::de::DeserializeOwned; use crate::auth::{AuthSession, NoopSession}; use crate::catalog::{ - REST_CATALOG_PROP_DISABLE_HEADER_REDACTION, RestCatalogConfig, explicit_headers_from_props, + RestCatalogConfig, disable_header_redaction_from_props, explicit_headers_from_props, }; use crate::request::HttpRequest; use crate::response::HttpResponse; @@ -165,15 +165,11 @@ impl HttpClient { } extra_headers.extend(table_headers); - let disable_header_redaction = props - .get(REST_CATALOG_PROP_DISABLE_HEADER_REDACTION) - .map(|value| value.eq_ignore_ascii_case("true")) - .unwrap_or(self.disable_header_redaction); - Ok(Self { client: self.client.clone(), extra_headers, - disable_header_redaction, + disable_header_redaction: disable_header_redaction_from_props(props) + .unwrap_or(self.disable_header_redaction), auth_session, }) } diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs index 24e2ed7002..ff8a8b0c0b 100644 --- a/crates/catalog/rest/src/credential.rs +++ b/crates/catalog/rest/src/credential.rs @@ -15,7 +15,7 @@ // specific language governing permissions and limitations // under the License. -//! Refresh of vended storage credentials against a REST catalog. +//! Vended storage credentials from a REST catalog. //! //! A REST catalog can vend short-lived storage credentials whose lifetime the //! client does not control. [`RestVendedCredentialProvider`] implements the @@ -34,6 +34,11 @@ //! after a failed fetch, transient failures are retried here with jittered //! exponential backoff while an unexpired credential remains available. //! +//! Like Java's `VendedCredentialsProvider`, a provider is rebuilt from the +//! FileIO properties after [`FileIO`](iceberg::io::FileIO) serialization and +//! connects to the catalog lazily in the receiving process. This requires +//! catalog authentication that can be rebuilt from `rest.auth.type`. +//! //! # Adding a cloud //! //! The refresh policy for each cloud lives in one [`CloudRefresh`] constant. To @@ -52,15 +57,17 @@ use iceberg::io::{ AWS_REFRESH_CREDENTIALS_ENDPOINT, AzdlsCredential, GCS_REFRESH_CREDENTIALS_ENABLED, GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, GcsCredential, S3_ACCESS_KEY_ID, S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, S3_SESSION_TOKEN_EXPIRES_AT_MS, - S3Credential, StorageCredential, StorageCredentialKind, StorageCredentialProvider, + S3Credential, StorageConfig, StorageCredential, StorageCredentialKind, + StorageCredentialProvider, StorageCredentialProviderFactory, storage_prefix_covers, }; use iceberg::{Error, ErrorKind, Result, TableIdent}; use rand::Rng; use reqwest::{Method, StatusCode, Url}; -use tokio::sync::Mutex; +use serde::{Deserialize, Serialize}; +use tokio::sync::{Mutex, OnceCell}; -use crate::REST_CATALOG_PROP_SCAN_PLAN_ID; -use crate::auth::AuthManager; +use crate::auth::{AuthManager, load_auth_manager}; +use crate::catalog::{REST_CATALOG_PROP_SCAN_PLAN_ID, RestCatalogConfig}; use crate::client::{HttpClient, unexpected_catalog_error_without_body}; use crate::request::HttpRequest; use crate::types::LoadCredentialsResponse; @@ -143,15 +150,6 @@ impl CloudRefresh { /// Backends with refresh support const SUPPORTED: &[Self] = &[Self::AWS, Self::GCP, Self::AZURE]; - /// The backend that serves `location`, by its URL scheme, or `None` if no - /// supported backend matches (in which case static credentials are used - /// as-is, as before). - fn for_location(location: &str) -> Option<&'static Self> { - Self::SUPPORTED - .iter() - .find(|cloud| cloud.matches_location(location)) - } - fn matches_location(&self, location: &str) -> bool { self.schemes .iter() @@ -187,7 +185,7 @@ impl CachedEntry { fn new(credential: StorageCredential, jitter_prefetch: bool) -> Self { let refresh_at = credential .expires_at() - .map(|expires_at| prefetch_time(expires_at, jitter_prefetch)); + .map(|expires_at| prefetch_time(SystemTime::now(), expires_at, jitter_prefetch)); Self { credential, refresh_at, @@ -233,6 +231,10 @@ struct CredentialError { } /// Cached credentials plus failure-backoff state. +/// +/// A fetch returns the credentials for every prefix of the table, so a path +/// the last response did not cover would not be covered by an immediate +/// re-fetch either. Backoff is therefore shared by the whole cloud. struct CacheState { entries: Vec, consecutive_failures: u32, @@ -240,6 +242,14 @@ struct CacheState { } impl CacheState { + fn new(entries: Vec) -> Self { + Self { + entries, + consecutive_failures: 0, + retry_not_before: None, + } + } + fn record_success(&mut self) { self.consecutive_failures = 0; self.retry_not_before = None; @@ -251,16 +261,25 @@ impl CacheState { Instant::now().checked_add(failure_backoff(self.consecutive_failures)); } - fn insert_fallback_if_missing(&mut self, fallback: CachedEntry, now: SystemTime) -> bool { - let has_unexpired_entry = self.entries.iter().any(|entry| { - entry.is_unexpired(now) && entry.credential.prefix() == fallback.credential.prefix() + /// Replace cached credentials with the fetched ones, prefix by prefix. + /// + /// An unexpired cached credential survives when the response carries no + /// valid replacement for its prefix, so an absent or malformed entry never + /// evicts a usable credential. Returns the prefixes the response replaced. + fn merge(&mut self, fetched: Vec, now: SystemTime) -> HashSet { + let fetched_prefixes = fetched + .iter() + .filter_map(|entry| entry.credential.prefix().map(str::to_owned)) + .collect::>(); + self.entries.retain(|entry| { + entry.is_unexpired(now) + && entry + .credential + .prefix() + .is_none_or(|prefix| !fetched_prefixes.contains(prefix)) }); - if has_unexpired_entry { - false - } else { - self.entries.push(fallback); - true - } + self.entries.extend(fetched); + fetched_prefixes } } @@ -278,6 +297,56 @@ struct ConfiguredCloud { } impl ConfiguredCloud { + fn new( + cloud: &'static CloudRefresh, + endpoint: String, + entries: Vec, + keyed_seeds: HashMap, + ) -> Self { + Self { + cloud, + endpoint, + cache: Mutex::new(CacheState::new(entries)), + keyed_seeds, + refresh: Mutex::new(()), + } + } + + /// Configure `cloud` from the FileIO properties, or `None` when they + /// advertise no enabled refresh endpoint for it. + fn configure( + cloud: &'static CloudRefresh, + base_uri: &str, + props: &HashMap, + ) -> Option { + let enabled = props + .get(cloud.enabled_key) + .is_none_or(|value| value.eq_ignore_ascii_case("true")); + let endpoint = props + .get(cloud.endpoint_key) + .filter(|endpoint| enabled && !endpoint.is_empty()) + .map(|endpoint| resolve_endpoint(base_uri, endpoint))?; + + let mut entries = Vec::new(); + let mut keyed_seeds = HashMap::new(); + match &cloud.seed_strategy { + SeedStrategy::Flat => { + if let Ok(credential) = (cloud.parse_credential)(props, None) { + entries.push(CachedEntry::seed(credential, cloud.jitter_prefetch)); + } + } + SeedStrategy::Keyed(strategy) => { + keyed_seeds = (strategy.parse_credentials)(props) + .into_iter() + .map(|(key, credential)| { + (key, CachedEntry::seed(credential, cloud.jitter_prefetch)) + }) + .collect(); + } + } + Some(Self::new(cloud, endpoint, entries, keyed_seeds)) + } + fn keyed_seed_for_path(&self, path: &str) -> Result> { let SeedStrategy::Keyed(strategy) = &self.cloud.seed_strategy else { return Ok(None); @@ -293,14 +362,114 @@ impl ConfiguredCloud { } } -/// Fetches and refreshes vended credentials from a REST catalog's table -/// credentials endpoint. +/// Serializable recipe for rebuilding a [`RestVendedCredentialProvider`] from +/// the FileIO properties, mirroring how Java rebuilds its provider on workers. +#[derive(Clone, Serialize, Deserialize)] +pub(crate) struct RestVendedCredentialProviderFactory { + /// Catalog URI, used to resolve relative refresh endpoints. + catalog_uri: String, + table: TableIdent, + /// The unmerged config returned by the table endpoint, from which the + /// auth manager derives a table session. Keeping it separate prevents + /// local FileIO overrides from masking table auth. + table_config: HashMap, +} + +impl RestVendedCredentialProviderFactory { + pub(crate) fn new( + catalog_uri: impl Into, + table: TableIdent, + table_config: HashMap, + ) -> Self { + Self { + catalog_uri: catalog_uri.into(), + table, + table_config, + } + } + + /// Connect to the catalog from the FileIO properties, as the catalog would. + async fn connect(&self, props: &HashMap) -> Result { + let config = RestCatalogConfig::builder() + .uri(self.catalog_uri.clone()) + .props(props.clone()) + .build(); + let auth_manager = load_auth_manager(&config)?; + let client = HttpClient::new(&config)?; + let session = auth_manager + .catalog_session(&client.without_auth_session(), &config.auth_props()) + .await?; + table_client( + &client.with_auth_session(session), + auth_manager.as_ref(), + &self.table, + props, + &self.table_config, + ) + .await + } +} + +impl std::fmt::Debug for RestVendedCredentialProviderFactory { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("RestVendedCredentialProviderFactory") + .field("catalog_uri", &self.catalog_uri) + .field("table", &self.table) + .finish_non_exhaustive() + } +} + +#[typetag::serde(name = "RestVendedCredentialProviderFactory")] +impl StorageCredentialProviderFactory for RestVendedCredentialProviderFactory { + fn build(&self, config: &StorageConfig) -> Result> { + let provider = + RestVendedCredentialProvider::configure(self, config.props()).ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + "FileIO configuration no longer configures vended credentials", + ) + })?; + Ok(Arc::new(RestVendedCredentialProvider { + factory: Some(self.clone()), + ..provider + })) + } +} + +/// Derive the table-scoped client used for credential requests. +async fn table_client( + catalog_client: &HttpClient, + auth_manager: &dyn AuthManager, + table: &TableIdent, + props: &HashMap, + table_config: &HashMap, +) -> Result { + let session = auth_manager + .table_session( + &catalog_client.without_auth_session(), + table, + table_config, + catalog_client.auth_session(), + ) + .await?; + catalog_client.for_table(props, session) +} + +/// Serves vended credentials for a table, refreshing them from the REST +/// catalog's table credentials endpoint. /// -/// Each cloud cache is seeded with the credential from the initial table -/// properties (when complete) and re-fetched from its endpoint as it -/// nears expiry. +/// Each cloud cache is seeded with the credentials from the initial +/// table properties (when complete) and re-fetched from its endpoint as they +/// near expiry. pub(crate) struct RestVendedCredentialProvider { - client: Arc, + /// Table-scoped catalog client. Set when the catalog builds the provider, + /// and connected on first refresh after deserialization. + client: OnceCell, + /// Rebuilds this provider in another process. `None` when the catalog + /// authentication cannot be rebuilt from properties. + factory: Option, + /// Effective FileIO properties, which supply catalog connection settings. + props: HashMap, /// Optional scan-plan identifier. plan_id: Option, /// Independently configured endpoint and cache for each backing cloud. @@ -308,12 +477,23 @@ pub(crate) struct RestVendedCredentialProvider { } impl RestVendedCredentialProvider { - fn new(client: Arc, plan_id: Option, clouds: Vec) -> Self { - Self { - client, - plan_id, + /// Configure a provider without a catalog connection, or `None` when no + /// supported cloud advertises an enabled refresh endpoint. + fn configure( + factory: &RestVendedCredentialProviderFactory, + props: &HashMap, + ) -> Option { + let clouds = CloudRefresh::SUPPORTED + .iter() + .filter_map(|cloud| ConfiguredCloud::configure(cloud, &factory.catalog_uri, props)) + .collect::>(); + (!clouds.is_empty()).then(|| Self { + client: OnceCell::new(), + factory: None, + props: props.clone(), + plan_id: props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(), clouds, - } + }) } fn configured_cloud_for_location(&self, location: &str) -> Option<&ConfiguredCloud> { @@ -322,182 +502,128 @@ impl RestVendedCredentialProvider { .find(|configured| configured.cloud.matches_location(location)) } + async fn client(&self) -> Result<&HttpClient> { + self.client + .get_or_try_init(|| async { + let factory = self.factory.as_ref().ok_or_else(|| { + Error::new( + ErrorKind::Unexpected, + "vended credential provider has no catalog connection", + ) + })?; + factory.connect(&self.props).await + }) + .await + } + /// Fetch fresh credentials from the catalog's credentials endpoint. async fn fetch(&self, configured: &ConfiguredCloud) -> Result { - let mut request = self.client.request(Method::GET, &configured.endpoint); + let cloud = configured.cloud; + let client = self.client().await?; + let mut request = client.request(Method::GET, &configured.endpoint); if let Some(plan_id) = &self.plan_id { request = request.query(&[("planId", plan_id)]); } let request = HttpRequest::build(request)?; - let response = self.client.query_catalog(request).await?; - - match response.status() { - StatusCode::OK => { - // Credential responses contain secrets. Do not include the - // response body in a deserialization error. - let parsed: LoadCredentialsResponse = serde_json::from_slice(response.body()) - .map_err(|error| { - Error::new( - ErrorKind::Unexpected, - "failed to parse vended credential response", - ) - .with_source(error) - })?; - let mut entries = Vec::new(); - let mut errors = Vec::new(); - let matching_credentials = parsed - .storage_credentials - .into_iter() - .filter(|credential| configured.cloud.matches_location(&credential.prefix)); - let parse_credential = configured.cloud.parse_credential; - let now = SystemTime::now(); - - for storage_credential in matching_credentials { - let prefix = storage_credential.prefix; - let parsed_credential = - parse_credential(&storage_credential.config, Some(prefix.clone())); - match parsed_credential { - Ok(credential) => { - let entry = - CachedEntry::new(credential, configured.cloud.jitter_prefetch); - if entry.is_unexpired(now) { - entries.push(entry); - } else { - errors.push(CredentialError { - prefix, - error: Error::new( - ErrorKind::DataInvalid, - "invalid vended credential: credential is already expired", - ), - }); - } - } - Err(error) => errors.push(CredentialError { prefix, error }), + let response = client.query_catalog(request).await?; + + if response.status() != StatusCode::OK { + return Err(unexpected_catalog_error_without_body( + response, + client.disable_header_redaction(), + )); + } + + // Credential responses contain secrets. Do not include the response + // body in a deserialization error. + let parsed: LoadCredentialsResponse = + serde_json::from_slice(response.body()).map_err(|error| { + Error::new( + ErrorKind::Unexpected, + "failed to parse vended credential response", + ) + .with_source(error) + })?; + let now = SystemTime::now(); + let mut entries = Vec::new(); + let mut errors = Vec::new(); + for credential in parsed + .storage_credentials + .into_iter() + .filter(|credential| cloud.matches_location(&credential.prefix)) + { + let prefix = credential.prefix; + let parsed = (cloud.parse_credential)(&credential.config, Some(prefix.clone())) + .map(|credential| CachedEntry::new(credential, cloud.jitter_prefetch)) + .and_then(|entry| { + if entry.is_unexpired(now) { + Ok(entry) + } else { + Err(Error::new( + ErrorKind::DataInvalid, + "invalid vended credential: credential is already expired", + )) } - } - Ok(ParsedCredentials { entries, errors }) + }); + match parsed { + Ok(entry) => entries.push(entry), + Err(error) => errors.push(CredentialError { prefix, error }), } - _ => Err(unexpected_catalog_error_without_body( - response, - self.client.disable_header_redaction(), - )), } + Ok(ParsedCredentials { entries, errors }) } async fn refresh_credential( &self, configured: &ConfiguredCloud, path: &str, - fallback: Option, + fallback: Option, ) -> Result { - match self.fetch(configured).await { + let fetched = self.fetch(configured).await; + let mut cache = configured.cache.lock().await; + let now = SystemTime::now(); + + let (fetched_prefixes, failure) = match fetched { Ok(ParsedCredentials { entries, errors }) => { - let invalid_prefixes = errors - .iter() - .map(|error| error.prefix.clone()) - .collect::>(); - let credential_error = errors + let failure = errors .into_iter() - .filter(|error| path.starts_with(&error.prefix)) - .max_by_key(|error| error.prefix.len()); - let failure = credential_error - .map(|error| error.error) - .unwrap_or_else(|| { - Error::new( - ErrorKind::Unexpected, - format!( - "no unexpired vended credential matches storage location: {path}" - ), - ) - }); - - let mut cache = configured.cache.lock().await; - if entries.is_empty() { - cache.record_failure(); - return fallback - .filter(|fallback| fallback.entry.is_unexpired(SystemTime::now())) - .map(|fallback| fallback.entry.credential) - .ok_or(failure); - } - - let now = SystemTime::now(); - let invalid_fallbacks = cache - .entries - .iter() - .filter(|entry| entry.is_unexpired(now)) - .filter(|entry| { - entry - .credential - .prefix() - .is_some_and(|prefix| invalid_prefixes.contains(prefix)) - }) - .cloned() - .collect::>(); - cache.entries = entries; - - // An invalid replacement for one prefix must not evict that - // prefix's still-valid cached credential. Iterate in reverse - // because equal-length prefix selection uses the last cached entry. - // This preserves the same credential if a bad response follows a - // response with duplicate prefixes. - let mut restored_fallback_prefixes = HashSet::new(); - for fallback in invalid_fallbacks.into_iter().rev() { - let prefix = fallback.credential.prefix().map(str::to_owned); - if cache.insert_fallback_if_missing(fallback, now) - && let Some(prefix) = prefix - { - restored_fallback_prefixes.insert(prefix); - } - } - - let selected = longest_prefix_match(&cache.entries, path) - .filter(|entry| entry.is_unexpired(now)) - .map(|entry| { - let is_restored_fallback = entry - .credential - .prefix() - .is_some_and(|prefix| restored_fallback_prefixes.contains(prefix)); - (entry.credential.clone(), is_restored_fallback) - }); - - if let Some((credential, is_restored_fallback)) = selected { - if is_restored_fallback { - cache.record_failure(); - } else { - cache.record_success(); - } - return Ok(credential); - } - - // Cache valid credentials for other prefixes. If none covers this - // path, retain its still-usable credential when its replacement was - // absent or invalid. Backoff ensures the failed path is retried promptly. - if let Some(fallback) = fallback.filter(|fallback| fallback.entry.is_unexpired(now)) - { - let credential = fallback.entry.credential.clone(); - if fallback.cacheable { - cache.insert_fallback_if_missing(fallback.entry, now); - } - cache.record_failure(); - return Ok(credential); - } - - cache.record_failure(); - Err(failure) - } - Err(fetch_error) => { - let mut cache = configured.cache.lock().await; - cache.record_failure(); - - // Graceful degradation: while the cached credential remains - // usable, serve it and retry after jittered backoff. Expired - // credentials are never served. - fallback - .filter(|fallback| fallback.entry.is_unexpired(SystemTime::now())) - .map(|fallback| fallback.entry.credential) - .ok_or(fetch_error) + .filter(|error| storage_prefix_covers(&error.prefix, path)) + .max_by_key(|error| error.prefix.len()) + .map(|error| error.error); + (cache.merge(entries, now), failure) } + Err(error) => (HashSet::new(), Some(error)), + }; + + let selected = longest_prefix_match(&cache.entries, path) + .filter(|entry| entry.is_unexpired(now)) + .cloned(); + let refreshed = selected.as_ref().is_some_and(|entry| { + entry + .credential + .prefix() + .is_some_and(|prefix| fetched_prefixes.contains(prefix)) + }); + if refreshed { + cache.record_success(); + } else { + // Graceful degradation: while a credential for this path remains + // usable, serve it and retry after jittered backoff. Expired + // credentials are never served. + cache.record_failure(); } + + selected + .or_else(|| fallback.filter(|fallback| fallback.is_unexpired(now))) + .map(|entry| entry.credential) + .ok_or_else(|| { + failure.unwrap_or_else(|| { + Error::new( + ErrorKind::Unexpected, + format!("no unexpired vended credential matches storage location: {path}"), + ) + }) + }) } } @@ -511,34 +637,10 @@ impl std::fmt::Debug for RestVendedCredentialProvider { enum CacheDecision { Use(StorageCredential), - Refresh(Option), + Refresh(Option), Backoff, } -#[derive(Clone)] -struct CredentialFallback { - entry: CachedEntry, - /// Whether this entry belongs in the ordinary prefix cache after a failed - /// replacement. Keyed seeds remain in their separately scoped map. - cacheable: bool, -} - -impl CredentialFallback { - fn cached(entry: CachedEntry) -> Self { - Self { - entry, - cacheable: true, - } - } - - fn keyed_seed(entry: CachedEntry) -> Self { - Self { - entry, - cacheable: false, - } - } -} - fn refresh_backoff_error(path: &str) -> Error { Error::new( ErrorKind::Unexpected, @@ -547,25 +649,17 @@ fn refresh_backoff_error(path: &str) -> Error { } async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> Result { - let keyed_seed = configured - .keyed_seed_for_path(path)? - .map(CredentialFallback::keyed_seed); + let keyed_seed = configured.keyed_seed_for_path(path)?; let cache = configured.cache.lock().await; let now = SystemTime::now(); - let cached = longest_prefix_match(&cache.entries, path) - .cloned() - .map(CredentialFallback::cached); - let current = match cached { - Some(cached) if cached.entry.is_unexpired(now) => Some(cached), + let current = match longest_prefix_match(&cache.entries, path).cloned() { + Some(cached) if cached.is_unexpired(now) => Some(cached), Some(expired) => keyed_seed.or(Some(expired)), None => keyed_seed, }; - if let Some(fallback) = current - .as_ref() - .filter(|fallback| fallback.entry.is_fresh(now)) - { - return Ok(CacheDecision::Use(fallback.entry.credential.clone())); + if let Some(entry) = current.as_ref().filter(|entry| entry.is_fresh(now)) { + return Ok(CacheDecision::Use(entry.credential.clone())); } if cache @@ -573,9 +667,8 @@ async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> Result Result { - if CloudRefresh::for_location(path).is_none() { - return Err(Error::new( - ErrorKind::FeatureUnsupported, - format!("no credential refresh implementation for storage location: {path}"), - )); - } let configured = self.configured_cloud_for_location(path).ok_or_else(|| { Error::new( ErrorKind::FeatureUnsupported, - format!("credential refresh is not configured for storage location: {path}"), + format!("vended credentials are not configured for storage location: {path}"), ) })?; @@ -613,11 +700,11 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { // for the in-flight refresh instead. let usable = current .as_ref() - .filter(|fallback| fallback.entry.is_unexpired(SystemTime::now())); - let _refresh_guard = if let Some(fallback) = usable { + .filter(|entry| entry.is_unexpired(SystemTime::now())); + let _refresh_guard = if let Some(entry) = usable { match configured.refresh.try_lock() { Ok(guard) => guard, - Err(_) => return Ok(fallback.entry.credential.clone()), + Err(_) => return Ok(entry.credential.clone()), } } else { configured.refresh.lock().await @@ -633,32 +720,53 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { self.refresh_credential(configured, path, current).await } + + fn factory(&self) -> Result> { + match &self.factory { + Some(factory) => Ok(Arc::new(factory.clone())), + None => Err(Error::new( + ErrorKind::FeatureUnsupported, + "the vended credential provider cannot be serialized because the REST catalog \ + uses an injected AuthManager, which cannot be rebuilt in another process; use \ + FileIO::without_credential_provider to serialize without credential refresh", + )), + } + } } /// Select the credential whose prefix is the longest match for `path`. fn longest_prefix_match<'a>(entries: &'a [CachedEntry], path: &str) -> Option<&'a CachedEntry> { entries .iter() - .filter(|entry| { - entry - .credential - .prefix() - .is_none_or(|prefix| path.starts_with(prefix)) - }) + .filter(|entry| entry.credential.covers(path)) .max_by_key(|entry| entry.credential.prefix().map_or(0, str::len)) } -/// Compute a successful credential's prefetch time. -fn prefetch_time(expires_at: SystemTime, jitter: bool) -> SystemTime { - let base = expires_at.checked_sub(REFRESH_BUFFER).unwrap_or(UNIX_EPOCH); +/// Compute the prefetch time of a credential obtained at `now`. +/// +/// Like Java, a credential is refreshed [`REFRESH_BUFFER`] before it expires. +/// Unlike Java, the buffer is capped at half the remaining lifetime: a +/// credential vended with a shorter lifetime would otherwise be due on +/// arrival, and every file operation would fetch again. +fn prefetch_time(now: SystemTime, expires_at: SystemTime, jitter: bool) -> SystemTime { + let lifetime = expires_at.duration_since(now).unwrap_or_default(); + let buffer = REFRESH_BUFFER.min(lifetime / 2); + let base = expires_at.checked_sub(buffer).unwrap_or(UNIX_EPOCH); if !jitter { return base; } - let jitter_window = REFRESH_BUFFER.saturating_sub(MIN_REFRESH_BUFFER); - let jitter_millis = rand::rng().random_range(0..jitter_window.as_millis() as u64); - base.checked_add(Duration::from_millis(jitter_millis)) - .unwrap_or(base) + // The minimum distance from expiry shrinks with a capped buffer. + let min_buffer = + buffer.mul_f64(MIN_REFRESH_BUFFER.as_secs_f64() / REFRESH_BUFFER.as_secs_f64()); + let jitter_millis = buffer.saturating_sub(min_buffer).as_millis() as u64; + if jitter_millis == 0 { + return base; + } + base.checked_add(Duration::from_millis( + rand::rng().random_range(0..jitter_millis), + )) + .unwrap_or(base) } /// Equal-jitter exponential backoff. The random lower half avoids both hot @@ -674,99 +782,45 @@ fn failure_backoff(consecutive_failures: u32) -> Duration { Duration::from_millis(rand::rng().random_range(floor_millis..=ceiling_millis)) } -/// Build a credential provider from a table's properties, -/// or `None` when no supported cloud advertises an enabled refresh endpoint. +/// Build a credential provider for a table, or `None` when no supported cloud +/// advertises an enabled refresh endpoint. /// -/// `base_uri` is the catalog URI, used to resolve a relative endpoint. +/// `catalog_client` carries the catalog session, from which `auth_manager` +/// derives a table session when the table config overrides authentication. /// `props` contains the effective FileIO properties after applying local -/// overrides, and therefore supplies headers and client policy. -/// `table_auth_props` is the unmerged config returned by the table endpoint: -/// keeping it separate prevents local FileIO overrides from masking table auth. -/// `auth_manager` derives an isolated child session when those properties -/// override authentication for `table`. +/// overrides, and therefore supplies headers and client policy. `portable` +/// states whether the catalog authentication can be rebuilt from `props` in +/// another process, which makes the provider serializable. pub(crate) async fn build_vended_credential_provider( - client: Arc, + catalog_client: &HttpClient, auth_manager: &dyn AuthManager, - table: &TableIdent, - base_uri: &str, + factory: RestVendedCredentialProviderFactory, props: &HashMap, - table_auth_props: Option<&HashMap>, + portable: bool, ) -> Result>> { - let clouds = CloudRefresh::SUPPORTED - .iter() - .filter_map(|cloud| { - let enabled = props - .get(cloud.enabled_key) - .is_none_or(|value| value.eq_ignore_ascii_case("true")); - if !enabled { - return None; - } - - let endpoint = resolve_endpoint( - base_uri, - props - .get(cloud.endpoint_key) - .filter(|value| !value.is_empty())?, - ); - let (entries, keyed_seeds) = match &cloud.seed_strategy { - SeedStrategy::Flat => { - let entries = (cloud.parse_credential)(props, None) - .ok() - .map(|credential| { - vec![CachedEntry::seed(credential, cloud.jitter_prefetch)] - }) - .unwrap_or_default(); - (entries, HashMap::new()) - } - SeedStrategy::Keyed(strategy) => { - let keyed_seeds = (strategy.parse_credentials)(props) - .into_iter() - .map(|(key, credential)| { - (key, CachedEntry::seed(credential, cloud.jitter_prefetch)) - }) - .collect(); - (Vec::new(), keyed_seeds) - } - }; - - Some(ConfiguredCloud { - cloud, - endpoint, - cache: Mutex::new(CacheState { - entries, - consecutive_failures: 0, - retry_not_before: None, - }), - keyed_seeds, - refresh: Mutex::new(()), - }) - }) - .collect::>(); - - if clouds.is_empty() { + let Some(provider) = RestVendedCredentialProvider::configure(&factory, props) else { return Ok(None); - } + }; - let empty_auth_props = HashMap::new(); - let table_auth_props = table_auth_props.unwrap_or(&empty_auth_props); - let table_session = auth_manager - .table_session( - &client.without_auth_session(), - table, - table_auth_props, - client.auth_session(), - ) - .await?; - let table_client = Arc::new(client.for_table(props, table_session)?); - let plan_id = props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(); - Ok(Some(Arc::new(RestVendedCredentialProvider::new( - table_client, - plan_id, - clouds, - )))) + let client = table_client( + catalog_client, + auth_manager, + &factory.table, + props, + &factory.table_config, + ) + .await?; + Ok(Some(Arc::new(RestVendedCredentialProvider { + client: OnceCell::new_with(Some(client)), + factory: portable.then_some(factory), + ..provider + }))) } /// Resolve a possibly-relative refresh endpoint against the catalog base URI. +/// +/// Absolute endpoints are used as-is and receive the catalog credentials, as +/// in Java. The catalog is trusted to advertise only its own endpoints. fn resolve_endpoint(base_uri: &str, endpoint: &str) -> String { if endpoint.starts_with("http://") || endpoint.starts_with("https://") { return endpoint.to_string(); @@ -784,7 +838,7 @@ fn scheme_of(location: &str) -> Option { .map(|url| url.scheme().to_string()) } -/// Parse a complete S3 credential returned by the credentials endpoint. +/// Parse a complete S3 credential supplied by the catalog. fn parse_s3_credential( config: &HashMap, prefix: Option, @@ -793,34 +847,39 @@ fn parse_s3_credential( let secret_access_key = required_nonempty(config, S3_SECRET_ACCESS_KEY)?; let session_token = required_nonempty(config, S3_SESSION_TOKEN)?; let expires_at = required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?; - let credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( - access_key_id, - secret_access_key, - Some(session_token), - ))) - .with_expiration(expires_at); - Ok(match prefix { - Some(prefix) => credential.with_prefix(prefix), - None => credential, - }) + Ok(with_prefix( + StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( + access_key_id, + secret_access_key, + Some(session_token), + ))) + .with_expiration(expires_at), + prefix, + )) } -/// Parse a complete GCS credential returned by the credentials endpoint. +/// Parse a complete GCS credential supplied by the catalog. fn parse_gcs_credential( config: &HashMap, prefix: Option, ) -> Result { let token = required_nonempty(config, GCS_TOKEN)?; let expires_at = required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?; - let credential = StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new(token))) - .with_expiration(expires_at); - Ok(match prefix { + Ok(with_prefix( + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new(token))) + .with_expiration(expires_at), + prefix, + )) +} + +fn with_prefix(credential: StorageCredential, prefix: Option) -> StorageCredential { + match prefix { Some(prefix) => credential.with_prefix(prefix), None => credential, - }) + } } -/// Parse an account-specific ADLS SAS credential returned by the catalog. +/// Parse a complete account-specific ADLS SAS credential supplied by the catalog. fn parse_azdls_credential( config: &HashMap, prefix: Option, @@ -832,12 +891,22 @@ fn parse_azdls_credential( ) })?; let account = azdls_account_name(&prefix)?; - let sas_token = required_nonempty(config, &format!("{ADLS_SAS_TOKEN_PREFIX}{account}"))?; + let (suffix, sas_token) = azdls_sas_tokens(config) + .filter(|(_, token_account, _)| *token_account == account) + .map(|(suffix, _, token)| (suffix, token)) + // Prefer the host-keyed token when both forms name the account, as + // the storage backend does. + .max() + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid vended credential: no {ADLS_SAS_TOKEN_PREFIX}* token for account {account}"), + ) + })?; let expires_at = required_epoch_millis( config, - &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{account}"), + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), )?; - Ok( StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new( sas_token, @@ -853,27 +922,41 @@ fn parse_azdls_credential( fn parse_azdls_account_seeds( config: &HashMap, ) -> HashMap { - config - .iter() - .filter_map(|(key, value)| { - let account = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; - if account.is_empty() || value.is_empty() { - return None; - } + let mut seeds = azdls_sas_tokens(config) + .filter_map(|(suffix, account, token)| { let expires_at = required_epoch_millis( config, - &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{account}"), + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), ) .ok()?; - Some(( + Some((suffix, account, token, expires_at)) + }) + .collect::>(); + // Deterministic choice when several keys name the same account: the last, + // host-keyed one wins, as in the storage backend. + seeds.sort_by(|left, right| left.0.cmp(right.0)); + seeds + .into_iter() + .map(|(_, account, token, expires_at)| { + ( account.to_string(), - StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new(value))) + StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new(token))) .with_expiration(expires_at), - )) + ) }) .collect() } +/// Account-specific SAS tokens as `(key suffix, account, token)`. Like Java, +/// keys may name the host (`account.dfs.core.windows.net`) or only the account. +fn azdls_sas_tokens(config: &HashMap) -> impl Iterator { + config.iter().filter_map(|(key, token)| { + let suffix = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; + let account = suffix.split('.').next().unwrap_or(suffix); + (!account.is_empty() && !token.is_empty()).then_some((suffix, account, token.as_str())) + }) +} + fn resolve_azdls_seed_path(location: &str) -> Result { let (mut url, account) = parse_azdls_location(location)?; url.set_path("/"); @@ -955,7 +1038,6 @@ mod tests { use mockito::{Matcher, Server}; use super::*; - use crate::RestCatalogConfig; use crate::auth::{NoopAuthManager, OAuth2Manager}; fn epoch_millis(time: SystemTime) -> String { @@ -991,15 +1073,47 @@ mod tests { } } + /// A previously cached entry: like a seed, its age is unknown, so it is due + /// once inside the nominal refresh window. fn cached_s3(prefix: &str, access_key_id: &str, expires_at: Option) -> CachedEntry { - CachedEntry::new(s3_cred(Some(prefix), access_key_id, expires_at), false) + CachedEntry::seed(s3_cred(Some(prefix), access_key_id, expires_at), false) } - fn test_client(base_uri: &str) -> Arc { + fn test_client(base_uri: &str) -> HttpClient { let config = RestCatalogConfig::builder() .uri(base_uri.to_string()) .build(); - Arc::new(HttpClient::new(&config).unwrap()) + HttpClient::new(&config).unwrap() + } + + fn test_factory( + base_uri: &str, + table_config: Option<&HashMap>, + ) -> RestVendedCredentialProviderFactory { + RestVendedCredentialProviderFactory::new( + base_uri, + test_table(), + table_config.cloned().unwrap_or_default(), + ) + } + + /// A provider whose AWS cache starts with `entries`. + fn provider_with_cached_s3( + base_uri: &str, + entries: Vec, + ) -> RestVendedCredentialProvider { + RestVendedCredentialProvider { + client: OnceCell::new_with(Some(test_client(base_uri))), + factory: None, + props: HashMap::new(), + plan_id: None, + clouds: vec![ConfiguredCloud::new( + &CloudRefresh::AWS, + format!("{base_uri}/v1/credentials"), + entries, + HashMap::new(), + )], + } } fn aws_refresh_props(endpoint: &str) -> HashMap { @@ -1031,18 +1145,17 @@ mod tests { } async fn build_test_provider( - client: Arc, + client: HttpClient, base_uri: &str, props: &HashMap, table_auth_props: Option<&HashMap>, ) -> Result>> { build_vended_credential_provider( - client, + &client, &NoopAuthManager, - &test_table(), - base_uri, + test_factory(base_uri, table_auth_props), props, - table_auth_props, + true, ) .await } @@ -1097,21 +1210,26 @@ mod tests { #[test] fn cloud_selection_accepts_root_prefixes_and_url_schemes() { + let supported = |location: &str| { + CloudRefresh::SUPPORTED + .iter() + .any(|cloud| cloud.matches_location(location)) + }; // Java uses these root prefixes for its fallback clients and accepts // credentials scoped directly to them. - assert!(CloudRefresh::for_location("s3").is_some()); - assert!(CloudRefresh::for_location("gs").is_some()); - assert!(CloudRefresh::for_location("s3://b/k").is_some()); - assert!(CloudRefresh::for_location("S3://b/k").is_some()); - assert!(CloudRefresh::for_location("s3a://b/k").is_some()); - assert!(CloudRefresh::for_location("s3n://b/k").is_some()); - assert!(CloudRefresh::for_location("gs://b/k").is_some()); - assert!(CloudRefresh::for_location("gcs://b/k").is_some()); - assert!(CloudRefresh::for_location("abfs").is_some()); - assert!(CloudRefresh::for_location("abfss://fs@acct.dfs.core.windows.net/k").is_some()); - assert!(CloudRefresh::for_location("wasb://fs@acct.blob.core.windows.net/k").is_some()); - assert!(CloudRefresh::for_location("wasbs://fs@acct.blob.core.windows.net/k").is_some()); - assert!(CloudRefresh::for_location("not a url").is_none()); + assert!(supported("s3")); + assert!(supported("gs")); + assert!(supported("s3://b/k")); + assert!(supported("S3://b/k")); + assert!(supported("s3a://b/k")); + assert!(supported("s3n://b/k")); + assert!(supported("gs://b/k")); + assert!(supported("gcs://b/k")); + assert!(supported("abfs")); + assert!(supported("abfss://fs@acct.dfs.core.windows.net/k")); + assert!(supported("wasb://fs@acct.blob.core.windows.net/k")); + assert!(supported("wasbs://fs@acct.blob.core.windows.net/k")); + assert!(!supported("not a url")); assert!(CloudRefresh::AWS.matches_location("s3://bucket/path")); assert!(!CloudRefresh::AWS.matches_location("s3evil://bucket/path")); } @@ -1124,8 +1242,13 @@ mod tests { assert!(parse_s3_credential(&config, None).is_err()); config.insert(S3_SECRET_ACCESS_KEY.to_string(), "SK".to_string()); assert!(parse_s3_credential(&config, None).is_err()); - config.insert(S3_SESSION_TOKEN.to_string(), "TOK".to_string()); + config.insert( + S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), + "not-a-timestamp".to_string(), + ); + assert!(parse_s3_credential(&config, None).is_err()); + config.insert( S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), "1500".to_string(), @@ -1212,19 +1335,39 @@ mod tests { #[test] fn prefetch_time_matches_cloud_policy() { - let expires_at = SystemTime::now() + Duration::from_secs(3600); + let now = SystemTime::now(); + let expires_at = now + Duration::from_secs(3600); assert_eq!( - prefetch_time(expires_at, false), + prefetch_time(now, expires_at, false), expires_at - REFRESH_BUFFER ); for _ in 0..16 { - let refresh_at = prefetch_time(expires_at, true); + let refresh_at = prefetch_time(now, expires_at, true); assert!(refresh_at >= expires_at - REFRESH_BUFFER); assert!(refresh_at < expires_at - MIN_REFRESH_BUFFER); } } + #[test] + fn prefetch_buffer_is_capped_at_half_the_lifetime() { + let now = SystemTime::now(); + let expires_at = now + Duration::from_secs(120); + assert_eq!( + prefetch_time(now, expires_at, false), + now + Duration::from_secs(60) + ); + + for _ in 0..16 { + let refresh_at = prefetch_time(now, expires_at, true); + assert!(refresh_at >= now + Duration::from_secs(60)); + assert!(refresh_at < expires_at - Duration::from_secs(12)); + } + + // An already expired credential is due immediately. + assert_eq!(prefetch_time(now, now, true), now); + } + #[test] fn seed_inside_nominal_window_is_immediately_due() { let credential = s3_cred( @@ -1532,17 +1675,7 @@ mod tests { epoch_millis(SystemTime::now() + Duration::from_secs(3600)), ), ]); - let provider = build_vended_credential_provider( - test_client(&server.url()), - &NoopAuthManager, - &test_table(), - &server.url(), - &props, - None, - ) - .await - .unwrap() - .unwrap(); + let provider = test_provider(&server.url(), &props).await; let credential = provider .load_credential("abfss://container@account1.dfs.core.windows.net/table/data/a.parquet") @@ -1673,23 +1806,11 @@ mod tests { .create_async() .await; - let provider = RestVendedCredentialProvider::new(test_client(&server.url()), None, vec![ - ConfiguredCloud { - cloud: &CloudRefresh::AWS, - endpoint: format!("{}/v1/credentials", server.url()), - cache: Mutex::new(CacheState { - entries: vec![cached_s3( - "s3://bucket/requested", - "SEED_AK", - Some(SystemTime::now() + Duration::from_secs(60)), - )], - consecutive_failures: 0, - retry_not_before: None, - }), - keyed_seeds: HashMap::new(), - refresh: Mutex::new(()), - }, - ]); + let provider = provider_with_cached_s3(&server.url(), vec![cached_s3( + "s3://bucket/requested", + "SEED_AK", + Some(SystemTime::now() + Duration::from_secs(60)), + )]); let fallback = provider .load_credential("s3://bucket/requested/file") @@ -1735,39 +1856,27 @@ mod tests { .create_async() .await; - let provider = RestVendedCredentialProvider::new(test_client(&server.url()), None, vec![ - ConfiguredCloud { - cloud: &CloudRefresh::AWS, - endpoint: format!("{}/v1/credentials", server.url()), - cache: Mutex::new(CacheState { - entries: vec![ - cached_s3( - "s3://bucket/table-a", - "OLD_AK", - Some(SystemTime::now() + Duration::from_secs(60)), - ), - cached_s3( - "s3://bucket/table-b", - "OLDER_BK", - Some(SystemTime::now() + Duration::from_secs(3600)), - ), - cached_s3( - "s3://bucket/table-b", - "LATEST_BK", - Some(SystemTime::now() + Duration::from_secs(3600)), - ), - cached_s3( - "s3://bucket/table-c", - "VALID_CK", - Some(SystemTime::now() + Duration::from_secs(3600)), - ), - ], - consecutive_failures: 0, - retry_not_before: None, - }), - keyed_seeds: HashMap::new(), - refresh: Mutex::new(()), - }, + let provider = provider_with_cached_s3(&server.url(), vec![ + cached_s3( + "s3://bucket/table-a", + "OLD_AK", + Some(SystemTime::now() + Duration::from_secs(60)), + ), + cached_s3( + "s3://bucket/table-b", + "OLDER_BK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + cached_s3( + "s3://bucket/table-b", + "LATEST_BK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), + cached_s3( + "s3://bucket/table-c", + "VALID_CK", + Some(SystemTime::now() + Duration::from_secs(3600)), + ), ]); let refreshed = provider @@ -1875,7 +1984,7 @@ mod tests { "catalog-token".to_string(), )])) .build(); - let client = Arc::new(HttpClient::new(&config).unwrap()); + let client = HttpClient::new(&config).unwrap(); let props = HashMap::from([ ( CloudRefresh::AWS.endpoint_key.to_string(), @@ -1888,12 +1997,11 @@ mod tests { let table_auth = HashMap::from([("token".to_string(), "table-token".to_string())]); let auth_manager = OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())); let provider = build_vended_credential_provider( - client, + &client, &auth_manager, - &test_table(), - &server.url(), + test_factory(&server.url(), Some(&table_auth)), &props, - Some(&table_auth), + true, ) .await .unwrap() @@ -1940,12 +2048,11 @@ mod tests { ]); let auth_manager = OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())); let provider = build_vended_credential_provider( - client, + &client, &auth_manager, - &test_table(), - &server.url(), + test_factory(&server.url(), Some(&table_auth)), &props, - Some(&table_auth), + true, ) .await .unwrap() @@ -2013,6 +2120,64 @@ mod tests { gcp_mock.assert_async().await; } + #[tokio::test] + async fn host_keyed_azdls_credentials_match_java() { + let mut server = Server::new_async().await; + let host = "account1.dfs.core.windows.net"; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let prefix = format!("abfss://container@{host}/table"); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"{prefix}","config":{{"adls.sas-token.{host}":"sv=2026&sig=refreshed","adls.sas-token-expires-at-ms.{host}":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/azure-credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/azure-credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + epoch_millis(SystemTime::now() + Duration::from_secs(3600)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + let sas_token = |credential: StorageCredential| match credential.into_kind() { + StorageCredentialKind::Azdls(azdls) => azdls.into_sas_token(), + other => panic!("expected ADLS credential, got {other:?}"), + }; + + // The host-keyed seed serves the account without fetching. + let seeded = provider + .load_credential(&format!("abfss://other@{host}/data.parquet")) + .await + .unwrap(); + assert_eq!(sas_token(seeded), "sv=2026&sig=seed"); + + // An account without a seed refreshes. The response is parsed with + // host-keyed tokens, and its entry then wins over the broader seed. + let uncovered = provider + .load_credential("abfss://container@account2.dfs.core.windows.net/table/a") + .await; + assert!(uncovered.is_err()); + let refreshed = provider + .load_credential(&format!("{prefix}/data.parquet")) + .await + .unwrap(); + assert_eq!(sas_token(refreshed), "sv=2026&sig=refreshed"); + mock.assert_async().await; + } + #[tokio::test] async fn refreshes_azdls_credentials() { let mut server = Server::new_async().await; @@ -2115,11 +2280,11 @@ mod tests { } #[tokio::test] - async fn sub_buffer_ttl_is_refetched_on_each_sequential_operation() { + async fn sub_buffer_ttl_is_reused_until_half_its_lifetime() { let mut server = Server::new_async().await; - // The vended TTL (60s) is shorter than REFRESH_BUFFER, so every operation - // is eligible for refresh. This matches AWS CachedSupplier: successful - // results do not receive failure backoff. + // The vended TTL (60s) is shorter than REFRESH_BUFFER. Java would treat + // it as due on arrival and fetch on every operation; the capped buffer + // keeps it for half its lifetime. let body = s3_response( "s3://bucket", "AK", @@ -2127,18 +2292,14 @@ mod tests { ); let mock = server .mock("GET", "/v1/credentials") - .expect(2) + .expect(1) .with_status(200) .with_header("content-type", "application/json") .with_body(body) .create_async() .await; - // No static creds -> no seed -> each sequential load fetches because the - // returned credential is already inside its prefetch window. - let props = aws_refresh_props("/v1/credentials"); - - let provider = test_provider(&server.url(), &props).await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; for _ in 0..2 { let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); @@ -2146,4 +2307,146 @@ mod tests { } mock.assert_async().await; } + + #[tokio::test] + async fn refreshed_credentials_must_be_complete() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body( + r#"{"storage-credentials":[{"prefix":"s3://bucket","config":{"s3.access-key-id":"AK","s3.secret-access-key":"SK"}}]}"#, + ) + .create_async() + .await; + + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + let error = provider + .load_credential("s3://bucket/x/f") + .await + .unwrap_err(); + assert!(error.message().contains("is missing or empty"), "{error}"); + mock.assert_async().await; + } + + #[tokio::test] + async fn scheme_aliases_share_prefix_scoped_credentials() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket/table", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + for path in ["s3a://bucket/table/f", "s3://bucket/table/f"] { + assert_eq!( + s3_access_key_id(&provider.load_credential(path).await.unwrap()), + "AK" + ); + } + // A prefix only covers whole path segments, so this path refreshes. + assert!( + provider + .load_credential("s3://bucket/table2/f") + .await + .is_err() + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn serialized_provider_reconnects_from_properties() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer table-token") + .match_header("x-custom", "value") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let mut props = aws_refresh_props("/v1/credentials"); + props.extend([ + ("token".to_string(), "catalog-token".to_string()), + ("header.x-custom".to_string(), "value".to_string()), + ]); + let table_auth = HashMap::from([("token".to_string(), "table-token".to_string())]); + let provider = build_vended_credential_provider( + &test_client(&server.url()), + &NoopAuthManager, + test_factory(&server.url(), Some(&table_auth)), + &props, + true, + ) + .await + .unwrap() + .unwrap(); + + let factory: Arc = + serde_json::from_str(&serde_json::to_string(&provider.factory().unwrap()).unwrap()) + .unwrap(); + let config = StorageConfig::new().with_props(props); + let rebuilt = factory.build(&config).unwrap(); + + assert_eq!( + s3_access_key_id(&rebuilt.load_credential("s3://bucket/x/f").await.unwrap()), + "AK" + ); + // The rebuilt provider is itself serializable. + assert!(rebuilt.factory().is_ok()); + mock.assert_async().await; + } + + #[tokio::test] + async fn provider_with_injected_auth_manager_is_not_serializable() { + let provider = build_vended_credential_provider( + &test_client("http://cat"), + &NoopAuthManager, + test_factory("http://cat", None), + &aws_refresh_props("/v1/creds"), + false, + ) + .await + .unwrap() + .unwrap(); + + let error = provider.factory().unwrap_err(); + assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + assert!( + error.message().contains("without_credential_provider"), + "{error}" + ); + } + + #[test] + fn factory_debug_omits_table_auth() { + let factory = test_factory( + "http://cat", + Some(&HashMap::from([( + "token".to_string(), + "secret-token".to_string(), + )])), + ); + let debug = format!("{factory:?}"); + assert!(!debug.contains("secret-token"), "{debug}"); + } } diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index dfe3ad44f0..deb821feb3 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -695,7 +695,7 @@ pub fn iceberg::inspect::SnapshotsTable<'a>::new(table: &'a iceberg::table::Tabl pub async fn iceberg::inspect::SnapshotsTable<'a>::scan(&self) -> iceberg::Result pub fn iceberg::inspect::SnapshotsTable<'a>::schema(&self) -> iceberg::spec::Schema pub mod iceberg::io -pub enum iceberg::io::StorageCredentialKind +#[non_exhaustive] pub enum iceberg::io::StorageCredentialKind pub iceberg::io::StorageCredentialKind::Azdls(iceberg::io::AzdlsCredential) pub iceberg::io::StorageCredentialKind::Gcs(iceberg::io::GcsCredential) pub iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential) @@ -755,6 +755,7 @@ pub fn iceberg::io::FileIO::new_output(&self, path: impl core::convert::AsRef Self pub fn iceberg::io::FileIO::new_with_memory() -> Self pub fn iceberg::io::FileIO::serialize_all(&self) -> iceberg::Result> +pub fn iceberg::io::FileIO::without_credential_provider(&self) -> Self impl core::clone::Clone for iceberg::io::FileIO pub fn iceberg::io::FileIO::clone(&self) -> iceberg::io::FileIO impl core::fmt::Debug for iceberg::io::FileIO @@ -1042,6 +1043,7 @@ impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg::io::StorageCredential impl iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::covers(&self, location: &str) -> bool pub fn iceberg::io::StorageCredential::expires_at(&self) -> core::option::Option pub fn iceberg::io::StorageCredential::into_kind(self) -> iceberg::io::StorageCredentialKind pub fn iceberg::io::StorageCredential::kind(&self) -> &iceberg::io::StorageCredentialKind @@ -1151,8 +1153,11 @@ pub fn iceberg::io::MemoryStorage::reader<'life0, 'life1, 'async_trait>(&'life0 pub fn iceberg::io::MemoryStorage::write<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str, bs: bytes::bytes::Bytes) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub fn iceberg::io::MemoryStorage::writer<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin>> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub trait iceberg::io::StorageCredentialProvider: core::fmt::Debug + core::marker::Send + core::marker::Sync +pub fn iceberg::io::StorageCredentialProvider::factory(&self) -> iceberg::Result> pub fn iceberg::io::StorageCredentialProvider::load_credential<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait pub fn iceberg::io::StorageCredentialProvider::supports_path(&self, _path: &str) -> bool +pub trait iceberg::io::StorageCredentialProviderFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize +pub fn iceberg::io::StorageCredentialProviderFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> pub trait iceberg::io::StorageFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize pub fn iceberg::io::StorageFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> pub fn iceberg::io::StorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> @@ -1162,6 +1167,7 @@ pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> +pub fn iceberg::io::storage_prefix_covers(prefix: &str, location: &str) -> bool pub mod iceberg::memory pub struct iceberg::memory::MemoryCatalog impl core::fmt::Debug for iceberg::memory::MemoryCatalog diff --git a/crates/iceberg/src/io/file_io.rs b/crates/iceberg/src/io/file_io.rs index 1da1310377..6769048a89 100644 --- a/crates/iceberg/src/io/file_io.rs +++ b/crates/iceberg/src/io/file_io.rs @@ -23,9 +23,9 @@ use futures::{Stream, StreamExt}; use super::storage::{ LocalFsStorageFactory, MemoryStorageFactory, Storage, StorageConfig, StorageCredentialProvider, - StorageFactory, + StorageCredentialProviderFactory, StorageFactory, }; -use crate::{Error, ErrorKind, Result}; +use crate::Result; /// FileIO implementation, used to manipulate files in underlying storage. /// @@ -77,18 +77,22 @@ mod _serde { use serde::{Deserialize, Serialize}; - use super::{StorageConfig, StorageFactory}; + use super::{StorageConfig, StorageCredentialProviderFactory, StorageFactory}; #[derive(Serialize)] pub(super) struct SerializableFileIO<'a> { pub(super) config: &'a StorageConfig, pub(super) factory: &'a Arc, + #[serde(skip_serializing_if = "Option::is_none")] + pub(super) credential_provider: Option>, } #[derive(Deserialize)] pub(super) struct DeserializedFileIO { pub(super) config: StorageConfig, pub(super) factory: Arc, + #[serde(default)] + pub(super) credential_provider: Option>, } } @@ -128,23 +132,28 @@ impl FileIO { /// /// All storage configuration properties are included in the serialized representation. These /// properties may contain credentials or other sensitive values, so the returned bytes must be - /// protected in transit and at rest by the application embedding this crate. + /// protected in transit and at rest by the application embedding this crate. A serialized + /// credential provider may likewise carry catalog authentication and vended credentials. /// /// Storage factories are serialized through [`typetag`](https://docs.rs/typetag). Third-party /// factories must use `#[typetag::serde]` on their [`StorageFactory`] implementation. - /// FileIO instances with a refreshable credential provider cannot be serialized because the - /// provider may hold process-local catalog and authentication state. + /// + /// A credential provider is serialized as the + /// [`StorageCredentialProviderFactory`] returned by + /// [`StorageCredentialProvider::factory`], and rebuilt on deserialization. Serialization fails + /// when the provider cannot be rebuilt in another process; use + /// [`FileIO::without_credential_provider`] to serialize without it. pub fn serialize_all(&self) -> Result> { - if self.credential_provider.is_some() { - return Err(Error::new( - ErrorKind::FeatureUnsupported, - "FileIO instances with credential providers cannot be serialized", - )); - } + let credential_provider = self + .credential_provider + .as_ref() + .map(|provider| provider.factory()) + .transpose()?; Ok(serde_json::to_vec(&_serde::SerializableFileIO { config: &self.config, factory: &self.factory, + credential_provider, })?) } @@ -154,15 +163,35 @@ impl FileIO { /// implementation so it is registered with `typetag`. Backend-specific requirements are /// documented by each storage factory implementation. pub fn deserialize_all(bytes: &[u8]) -> Result { - let _serde::DeserializedFileIO { config, factory } = serde_json::from_slice(bytes)?; + let _serde::DeserializedFileIO { + config, + factory, + credential_provider, + } = serde_json::from_slice(bytes)?; + let credential_provider = credential_provider + .map(|provider_factory| provider_factory.build(&config)) + .transpose()?; Ok(Self { config, factory, - credential_provider: None, + credential_provider, storage: Arc::new(OnceLock::new()), }) } + /// Returns a copy of this `FileIO` without its credential provider. + /// + /// The copy uses only the credentials in its storage configuration, which are not refreshed. + /// Use this to serialize a `FileIO` whose credential provider cannot be serialized. + pub fn without_credential_provider(&self) -> Self { + Self { + config: self.config.clone(), + factory: Arc::clone(&self.factory), + credential_provider: None, + storage: Arc::new(OnceLock::new()), + } + } + /// Get the storage configuration. pub fn config(&self) -> &StorageConfig { &self.config @@ -302,6 +331,9 @@ impl FileIOBuilder { } /// Attach a provider of refreshable, backend-specific credentials. + /// + /// Storage factories that cannot use the provider ignore it and use the + /// credentials in the configuration. pub fn with_credential_provider( mut self, provider: Arc, @@ -485,7 +517,9 @@ mod tests { use super::{FileIO, FileIOBuilder}; use crate::io::{ - LocalFsStorageFactory, MemoryStorageFactory, StorageCredential, StorageCredentialProvider, + GcsCredential, LocalFsStorageFactory, MemoryStorageFactory, StorageConfig, + StorageCredential, StorageCredentialKind, StorageCredentialProvider, + StorageCredentialProviderFactory, }; use crate::{ErrorKind, Result}; @@ -495,7 +529,38 @@ mod tests { #[async_trait::async_trait] impl StorageCredentialProvider for TestCredentialProvider { async fn load_credential(&self, _path: &str) -> Result { - unreachable!("unsupported factories must reject the provider before loading from it") + unreachable!("unsupported factories must ignore the provider") + } + } + + /// A provider rebuilt from the `FileIO` configuration, like a catalog provider. + #[derive(Debug)] + struct PortableCredentialProvider { + endpoint: Option, + } + + #[async_trait::async_trait] + impl StorageCredentialProvider for PortableCredentialProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(StorageCredential::new(StorageCredentialKind::Gcs( + GcsCredential::new(self.endpoint.clone().unwrap_or_default()), + ))) + } + + fn factory(&self) -> Result> { + Ok(Arc::new(PortableCredentialProviderFactory)) + } + } + + #[derive(Debug, serde::Serialize, serde::Deserialize)] + struct PortableCredentialProviderFactory; + + #[typetag::serde] + impl StorageCredentialProviderFactory for PortableCredentialProviderFactory { + fn build(&self, config: &StorageConfig) -> Result> { + Ok(Arc::new(PortableCredentialProvider { + endpoint: config.get("endpoint").cloned(), + })) } } @@ -646,13 +711,18 @@ mod tests { } #[tokio::test] - async fn test_file_io_rejects_credentials_for_unsupported_factory() { + async fn test_file_io_ignores_credentials_for_unsupported_factory() { let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) .with_credential_provider(Arc::new(TestCredentialProvider)) .build(); - let err = file_io.exists("memory://file").await.unwrap_err(); - assert_eq!(err.kind(), ErrorKind::FeatureUnsupported, "{err}"); + file_io + .new_output("memory://file") + .unwrap() + .write("data".into()) + .await + .unwrap(); + assert!(file_io.exists("memory://file").await.unwrap()); } #[test] @@ -663,6 +733,38 @@ mod tests { let err = file_io.serialize_all().unwrap_err(); assert_eq!(err.kind(), ErrorKind::FeatureUnsupported, "{err}"); + + let deserialized = FileIO::deserialize_all( + &file_io + .without_credential_provider() + .serialize_all() + .unwrap(), + ) + .unwrap(); + assert!(deserialized.credential_provider.is_none()); + } + + #[tokio::test] + async fn test_file_io_rebuilds_credential_provider_after_serialization() { + let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_prop("endpoint", "https://catalog/credentials") + .with_credential_provider(Arc::new(PortableCredentialProvider { endpoint: None })) + .build(); + + let deserialized = FileIO::deserialize_all(&file_io.serialize_all().unwrap()).unwrap(); + let credential = deserialized + .credential_provider + .unwrap() + .load_credential("gs://bucket/file") + .await + .unwrap(); + // The provider is rebuilt from the deserialized configuration. + match credential.kind() { + StorageCredentialKind::Gcs(gcs) => { + assert_eq!(gcs.token(), "https://catalog/credentials") + } + other => panic!("expected GCS credential, got {other:?}"), + } } #[tokio::test] diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index e3401e03fe..a193d7af44 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -143,18 +143,15 @@ pub trait StorageFactory: Debug + Send + Sync { /// Build a new Storage instance, optionally supplying a credential provider /// that the backend can call to obtain and refresh short-lived credentials. + /// + /// Backends that cannot use the provider ignore it and use the credentials + /// in `config`, as they would without one. The default does exactly that. + #[allow(unused_variables)] fn build_with_credentials( &self, config: &StorageConfig, credential_provider: Option>, ) -> Result> { - if credential_provider.is_some() { - return Err(Error::new( - ErrorKind::FeatureUnsupported, - "Storage factory does not support refreshable credential providers", - )); - } - self.build(config) } } @@ -188,8 +185,32 @@ pub trait StorageCredentialProvider: Debug + Send + Sync { /// `s3://bucket/warehouse/db/table/...`). Providers that vend distinct /// credentials per location prefix use it to select the most specific /// match. When the selected credential has a declared - /// [`StorageCredential::prefix`], it must cover `path`. + /// [`StorageCredential::prefix`], it must [cover](StorageCredential::covers) `path`. async fn load_credential(&self, path: &str) -> Result; + + /// Return a factory that rebuilds an equivalent provider in another process. + /// + /// [`FileIO::serialize_all`](crate::io::FileIO::serialize_all) serializes this + /// factory in place of the provider. The default reports that the provider + /// cannot be serialized. + fn factory(&self) -> Result> { + Err(Error::new( + ErrorKind::FeatureUnsupported, + "storage credential provider cannot be serialized", + )) + } +} + +/// Serializable recipe that rebuilds a [`StorageCredentialProvider`] after +/// [`FileIO`](crate::io::FileIO) deserialization. +/// +/// Factories are serialized through [`typetag`](https://docs.rs/typetag), so +/// implementations must use `#[typetag::serde]`, and the receiving binary must +/// link the concrete implementation. +#[typetag::serde(tag = "type")] +pub trait StorageCredentialProviderFactory: Debug + Send + Sync { + /// Build a provider for a `FileIO` with the given storage configuration. + fn build(&self, config: &StorageConfig) -> Result>; } /// A vended storage credential together with its scope and expiry. @@ -232,6 +253,20 @@ impl StorageCredential { self.prefix.as_deref() } + /// Return whether this credential applies to `location`. + /// + /// A credential without a prefix covers every location. Otherwise the + /// prefix must match whole path segments of `location`, and scheme + /// aliases (`s3a`/`s3n` for `s3`, `gcs` for `gs`, and the plain-text + /// Azure schemes for their TLS variants) are treated as equal. A prefix + /// that is only a scheme, such as `s3`, covers every location with that + /// scheme. + pub fn covers(&self, location: &str) -> bool { + self.prefix + .as_deref() + .is_none_or(|prefix| storage_prefix_covers(prefix, location)) + } + /// Return the backend-specific credential material. pub fn kind(&self) -> &StorageCredentialKind { &self.kind @@ -248,8 +283,41 @@ impl StorageCredential { } } +/// Return whether the storage-location `prefix` covers `location`, with the +/// matching rules of [`StorageCredential::covers`]. +pub fn storage_prefix_covers(prefix: &str, location: &str) -> bool { + let Some((location_scheme, location_rest)) = location.split_once("://") else { + return false; + }; + let Some((prefix_scheme, prefix_rest)) = prefix.split_once("://") else { + return !prefix.is_empty() && canonical_scheme(prefix) == canonical_scheme(location_scheme); + }; + + canonical_scheme(prefix_scheme) == canonical_scheme(location_scheme) + && location_rest + .strip_prefix(prefix_rest) + .is_some_and(|remainder| { + prefix_rest.is_empty() + || prefix_rest.ends_with('/') + || remainder.is_empty() + || remainder.starts_with('/') + }) +} + +fn canonical_scheme(scheme: &str) -> String { + let scheme = scheme.to_ascii_lowercase(); + match scheme.as_str() { + "s3a" | "s3n" => "s3".to_string(), + "gcs" => "gs".to_string(), + "abfs" => "abfss".to_string(), + "wasb" => "wasbs".to_string(), + _ => scheme, + } +} + /// Backend-specific credential material. #[derive(Clone, Debug)] +#[non_exhaustive] pub enum StorageCredentialKind { /// Amazon S3 credentials. S3(S3Credential), @@ -379,3 +447,51 @@ impl Debug for GcsCredential { f.debug_struct("GcsCredential").finish_non_exhaustive() } } + +#[cfg(test)] +mod tests { + use super::*; + + fn scoped(prefix: &str) -> StorageCredential { + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))) + .with_prefix(prefix) + } + + #[test] + fn credential_prefix_matches_whole_segments() { + let credential = scoped("s3://bucket/table"); + assert!(credential.covers("s3://bucket/table")); + assert!(credential.covers("s3://bucket/table/data/file.parquet")); + assert!(!credential.covers("s3://bucket/table2/data/file.parquet")); + assert!(!credential.covers("s3://bucket/tab")); + assert!(!credential.covers("s3://other/table/file.parquet")); + + assert!(scoped("s3://bucket/table/").covers("s3://bucket/table/file.parquet")); + assert!(scoped("s3://").covers("s3://any/file.parquet")); + } + + #[test] + fn credential_prefix_treats_scheme_aliases_as_equal() { + let credential = scoped("s3://bucket/table"); + assert!(credential.covers("s3a://bucket/table/file.parquet")); + assert!(credential.covers("S3N://bucket/table/file.parquet")); + assert!(scoped("gcs://bucket").covers("gs://bucket/file.parquet")); + assert!( + scoped("abfss://fs@account.dfs.core.windows.net/table") + .covers("abfs://fs@account.dfs.core.windows.net/table/file.parquet") + ); + assert!(!credential.covers("gs://bucket/table/file.parquet")); + } + + #[test] + fn credential_scheme_prefix_covers_the_whole_scheme() { + assert!(scoped("s3").covers("s3a://bucket/file.parquet")); + assert!(!scoped("s3").covers("gs://bucket/file.parquet")); + assert!(!scoped("").covers("s3://bucket/file.parquet")); + assert!(!scoped("s3").covers("not-a-url")); + assert!( + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))) + .covers("gs://bucket/file.parquet") + ); + } +} diff --git a/crates/storage/opendal/public-api.txt b/crates/storage/opendal/public-api.txt index 4e7d89acea..ad6c4522b8 100644 --- a/crates/storage/opendal/public-api.txt +++ b/crates/storage/opendal/public-api.txt @@ -1,21 +1,21 @@ pub mod iceberg_storage_opendal pub use iceberg_storage_opendal::AwsCredential pub use iceberg_storage_opendal::ProvideCredential -pub enum iceberg_storage_opendal::OpenDalStorage -pub iceberg_storage_opendal::OpenDalStorage::Azdls +#[non_exhaustive] pub enum iceberg_storage_opendal::OpenDalStorage +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Azdls pub iceberg_storage_opendal::OpenDalStorage::Azdls::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Azdls::credential_provider: core::option::Option> pub iceberg_storage_opendal::OpenDalStorage::Azdls::sas_tokens: alloc::sync::Arc -pub iceberg_storage_opendal::OpenDalStorage::Gcs +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Gcs pub iceberg_storage_opendal::OpenDalStorage::Gcs::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Gcs::credential_provider: core::option::Option> -pub iceberg_storage_opendal::OpenDalStorage::Hf +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Hf pub iceberg_storage_opendal::OpenDalStorage::Hf::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::LocalFs pub iceberg_storage_opendal::OpenDalStorage::Memory(opendal_core::types::operator::operator::Operator) -pub iceberg_storage_opendal::OpenDalStorage::Oss +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Oss pub iceberg_storage_opendal::OpenDalStorage::Oss::config: alloc::sync::Arc -pub iceberg_storage_opendal::OpenDalStorage::S3 +#[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::S3 pub iceberg_storage_opendal::OpenDalStorage::S3::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::S3::credential_provider: core::option::Option> pub iceberg_storage_opendal::OpenDalStorage::S3::customized_credential_load: core::option::Option diff --git a/crates/storage/opendal/src/azdls.rs b/crates/storage/opendal/src/azdls.rs index 44f036f2b6..387915d54c 100644 --- a/crates/storage/opendal/src/azdls.rs +++ b/crates/storage/opendal/src/azdls.rs @@ -36,8 +36,7 @@ use reqsign_core::{ use serde::{Deserialize, Serialize}; use url::Url; -use crate::DynamicCredentialScope; -use crate::utils::{from_opendal_error, system_time_to_timestamp, validate_credential_prefix}; +use crate::utils::{VendedCredentialSource, from_opendal_error}; /// Local version of `ensure_data_valid` macro since the iceberg crate's macro /// uses `$crate::error::Error` paths that don't resolve from external crates @@ -98,15 +97,24 @@ pub(crate) fn azdls_config_parse(mut properties: HashMap) -> Res pub struct AzdlsSasTokens(HashMap); impl AzdlsSasTokens { + /// Collect `adls.sas-token.` properties by storage account. Like + /// Java, keys may name the host (`account.dfs.core.windows.net`) or only + /// the account. pub(crate) fn from_properties(properties: &HashMap) -> Self { + let mut tokens = properties + .iter() + .filter_map(|(key, value)| { + let account = sas_token_account(key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?); + (!account.is_empty() && !value.is_empty()).then_some((key, account, value)) + }) + .collect::>(); + // Deterministic choice when several keys name the same account: the + // last, host-keyed one wins. + tokens.sort(); Self( - properties - .iter() - .filter_map(|(key, value)| { - let account = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; - (!account.is_empty() && !value.is_empty()) - .then(|| (account.to_string(), value.clone())) - }) + tokens + .into_iter() + .map(|(_, account, value)| (account.to_string(), value.clone())) .collect(), ) } @@ -116,6 +124,12 @@ impl AzdlsSasTokens { } } +/// The storage account named by the suffix of an account-specific SAS token +/// property: a host such as `account.dfs.core.windows.net`, or the account. +fn sas_token_account(suffix: &str) -> &str { + suffix.split('.').next().unwrap_or(suffix) +} + impl std::fmt::Debug for AzdlsSasTokens { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("AzdlsSasTokens") @@ -133,7 +147,7 @@ pub(crate) fn azdls_create_operator<'a>( config: &AzdlsConfig, sas_tokens: &AzdlsSasTokens, credential_provider: &Option>, - credential_scope: Option<&DynamicCredentialScope>, + credential_location: Option<&str>, ) -> Result<(opendal::Operator, &'a str)> { let path = absolute_path.parse::()?; match_path_with_config(&path, config)?; @@ -144,7 +158,7 @@ pub(crate) fn azdls_create_operator<'a>( sas_tokens, credential_provider, absolute_path, - credential_scope, + credential_location, )?; // Paths to files in ADLS tend to be written in fully qualified form, @@ -172,12 +186,19 @@ pub enum AzureStorageScheme { } impl AzureStorageScheme { - // Iceberg Java accepts the non-secure aliases for compatibility but still - // connects over TLS. SAS tokens are query parameters and must not be sent - // over a plaintext connection. + /// The HTTP scheme of an endpoint derived from a path. + /// + /// Iceberg Java accepts the non-secure aliases for compatibility but still + /// connects over TLS. SAS tokens are query parameters and must not be sent + /// over a plaintext connection unless the user configures such an endpoint. pub fn as_http_scheme(&self) -> &str { "https" } + + /// Whether an explicitly configured endpoint must use TLS. + fn requires_tls(&self) -> bool { + matches!(self, AzureStorageScheme::Abfss | AzureStorageScheme::Wasbs) + } } impl Display for AzureStorageScheme { @@ -220,12 +241,13 @@ pub(crate) fn match_path_with_config(path: &AzureStoragePath, config: &AzdlsConf } if let Some(ref configured_endpoint) = config.endpoint { - let passed_http_scheme = path.scheme.as_http_scheme(); + // An explicit plaintext endpoint, such as a local emulator, remains + // valid for the non-secure schemes. ensure_data_valid!( - configured_endpoint.starts_with(passed_http_scheme), - "Storage::Azdls: Endpoint {} does not use the expected http scheme {}.", + !path.scheme.requires_tls() || configured_endpoint.starts_with("https://"), + "Storage::Azdls: Endpoint {} does not use https, which the {} scheme requires.", configured_endpoint, - passed_http_scheme + path.scheme ); let ends_with_expected_suffix = configured_endpoint @@ -248,7 +270,7 @@ fn azdls_config_build( sas_tokens: &AzdlsSasTokens, credential_provider: &Option>, absolute_path: &str, - credential_scope: Option<&DynamicCredentialScope>, + credential_location: Option<&str>, ) -> Result { let mut builder = config.clone().into_builder(); @@ -262,10 +284,11 @@ fn azdls_config_build( .as_ref() .filter(|provider| provider.supports_path(absolute_path)); if let Some(provider) = credential_provider { - let chain = ProvideCredentialChain::new().push(VendedAzdlsCredentialProvider::new( - Arc::clone(provider), - absolute_path.to_string(), - credential_scope.cloned(), + let chain = ProvideCredentialChain::new().push(VendedAzdlsCredentialProvider( + VendedCredentialSource::new( + Arc::clone(provider), + credential_location.unwrap_or(absolute_path).to_string(), + ), )); builder = builder.credential_provider_chain(chain); } else if let Some(sas_token) = sas_tokens.for_path(path) { @@ -277,58 +300,15 @@ fn azdls_config_build( /// Adapts a generic [`StorageCredentialProvider`] into a reqsign /// [`ProvideCredential`] for expiring Azure SAS tokens. -struct VendedAzdlsCredentialProvider { - provider: Arc, - path: String, - credential_scope: Option, -} - -impl VendedAzdlsCredentialProvider { - fn new( - provider: Arc, - path: String, - credential_scope: Option, - ) -> Self { - Self { - provider, - path, - credential_scope, - } - } -} - -impl std::fmt::Debug for VendedAzdlsCredentialProvider { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("VendedAzdlsCredentialProvider") - .field("path", &self.path) - .field("credential_scope", &self.credential_scope) - .finish_non_exhaustive() - } -} +#[derive(Debug)] +struct VendedAzdlsCredentialProvider(VendedCredentialSource); impl ProvideCredential for VendedAzdlsCredentialProvider { type Credential = AzureCredential; async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { - let credential = self - .provider - .load_credential(&self.path) - .await - .map_err(|error| { - ReqsignError::unexpected("failed to load vended ADLS credential").with_source(error) - })?; - validate_credential_prefix( - &self.path, - credential.prefix(), - self.credential_scope.as_ref(), - )?; - - let expires_at = credential - .expires_at() - .map(system_time_to_timestamp) - .transpose()?; - match credential.into_kind() { - StorageCredentialKind::Azdls(azdls) => { + match self.0.load("ADLS").await? { + (StorageCredentialKind::Azdls(azdls), expires_at) => { let sas_token = azdls.into_sas_token(); Ok(Some(match expires_at { Some(expires_at) => { @@ -405,13 +385,12 @@ impl FromStr for AzureStoragePath { } } -pub(crate) fn azdls_batch_key(absolute_path: &str) -> Option { - absolute_path.parse::().ok().map(|path| { - format!( - "{}://{}@{}.{}", - path.scheme, path.filesystem, path.account_name, path.endpoint_suffix - ) - }) +pub(crate) fn azdls_batch_key(absolute_path: &str) -> Result { + let path = absolute_path.parse::()?; + Ok(format!( + "{}://{}@{}.{}", + path.scheme, path.filesystem, path.account_name, path.endpoint_suffix + )) } fn parse_azure_storage_endpoint(url: &Url) -> Result<(&str, &str, &str)> { @@ -482,7 +461,7 @@ mod tests { use super::{ AzdlsSasTokens, AzureStoragePath, AzureStorageScheme, VendedAzdlsCredentialProvider, - azdls_batch_key, azdls_config_parse, azdls_create_operator, + VendedCredentialSource, azdls_batch_key, azdls_config_parse, azdls_create_operator, }; #[derive(Debug)] @@ -581,6 +560,18 @@ mod tests { ), None, ), + ( + "plaintext endpoint for a non-secure scheme", + ( + "abfs://myfs@myaccount.dfs.core.windows.net/path/to/file.parquet", + AzdlsConfig { + account_name: Some("myaccount".to_string()), + endpoint: Some("http://myaccount.dfs.core.windows.net".to_string()), + ..Default::default() + }, + ), + Some(("myfs", "/path/to/file.parquet")), + ), ( "incompatible scheme for endpoint", ( @@ -662,11 +653,10 @@ mod tests { )) .with_prefix("abfss://container@account.dfs.core.windows.net/table") .with_expiration(expires_at); - let provider = VendedAzdlsCredentialProvider::new( + let provider = VendedAzdlsCredentialProvider(VendedCredentialSource::new( Arc::new(FixedCredentialProvider(credential)), path.to_string(), - None, - ); + )); let credential = provider .provide_credential(&Context::new()) @@ -681,7 +671,7 @@ mod tests { assert_eq!(token, "sv=2026&sig=secret"); assert_eq!( actual_expires_at, - Some(super::system_time_to_timestamp(expires_at).unwrap()) + Some(crate::utils::system_time_to_timestamp(expires_at).unwrap()) ); } other => panic!("expected SAS token, got {other:?}"), @@ -708,6 +698,23 @@ mod tests { assert_eq!(sas_tokens.for_path(&path), Some("sv=2026&sig=second")); } + #[test] + fn host_keyed_sas_tokens_match_java() { + let properties = HashMap::from([( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=host".to_string(), + )]); + let sas_tokens = AzdlsSasTokens::from_properties(&properties); + + for location in [ + "abfss://container@account.dfs.core.windows.net/table/data.parquet", + "wasbs://container@account.blob.core.windows.net/table/data.parquet", + ] { + let path = location.parse::().unwrap(); + assert_eq!(sas_tokens.for_path(&path), Some("sv=2026&sig=host")); + } + } + #[tokio::test] async fn vended_provider_rejects_mismatched_prefix() { let credential = StorageCredential::new(StorageCredentialKind::Azdls( @@ -715,11 +722,10 @@ mod tests { )) .with_prefix("abfss://other@account.dfs.core.windows.net/table") .with_expiration(SystemTime::now() + Duration::from_secs(3600)); - let provider = VendedAzdlsCredentialProvider::new( + let provider = VendedAzdlsCredentialProvider(VendedCredentialSource::new( Arc::new(FixedCredentialProvider(credential)), "abfss://container@account.dfs.core.windows.net/table/data.parquet".to_string(), - None, - ); + )); assert!(provider.provide_credential(&Context::new()).await.is_err()); } @@ -730,8 +736,12 @@ mod tests { let second = azdls_batch_key("abfss://second@account.dfs.core.windows.net/table/b.parquet"); let blob = azdls_batch_key("wasbs://first@account.blob.core.windows.net/table/c.parquet"); - assert_ne!(first, second); - assert_ne!(first, blob); + assert_ne!(first.unwrap(), second.unwrap()); + assert_ne!( + azdls_batch_key("abfss://first@account.dfs.core.windows.net/table/a.parquet").unwrap(), + blob.unwrap() + ); + assert!(azdls_batch_key("abfss:///no-account.parquet").is_err()); } #[test] diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index 20ddc24907..f374b8034c 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -30,10 +30,7 @@ use reqsign_core::{Context, Error as ReqsignError, ProvideCredential, Result as use reqsign_google::{Credential as GoogleCredential, Token as GoogleToken}; use url::Url; -use crate::DynamicCredentialScope; -use crate::utils::{ - from_opendal_error, is_truthy, system_time_to_timestamp, validate_credential_prefix, -}; +use crate::utils::{VendedCredentialSource, from_opendal_error, is_truthy}; /// Parse iceberg properties to [`GcsConfig`]. pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result { @@ -52,7 +49,7 @@ pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result>, path: &str, - credential_scope: Option<&DynamicCredentialScope>, + credential_location: Option<&str>, ) -> Result { let url = Url::parse(path)?; if !matches!(url.scheme(), "gs" | "gcs") { @@ -128,11 +125,11 @@ pub(crate) fn gcs_config_build( // A catalog-supplied provider re-fetches the vended OAuth2 token as it nears expiry if let Some(provider) = credential_provider { - builder = builder.credential_provider(VendedGcsCredentialProvider::new( - Arc::clone(provider), - path.to_string(), - credential_scope.cloned(), - )); + builder = + builder.credential_provider(VendedGcsCredentialProvider(VendedCredentialSource::new( + Arc::clone(provider), + credential_location.unwrap_or(path).to_string(), + ))); } Operator::new(builder).map_err(from_opendal_error) @@ -141,67 +138,15 @@ pub(crate) fn gcs_config_build( /// Adapts a generic [`StorageCredentialProvider`] into a `reqsign` /// [`ProvideCredential`], so the GCS signer can obtain and refresh vended OAuth2 /// tokens. -struct VendedGcsCredentialProvider { - provider: Arc, - /// Absolute path this operator serves: handed back to the provider so it can - /// select the vended credential whose prefix best matches the location. - path: String, - /// Exact scope selected when a bulk-delete operator was created. Ordinary - /// operators are unbound so they can follow the provider's best match. - credential_scope: Option, -} - -impl std::fmt::Debug for VendedGcsCredentialProvider { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("VendedGcsCredentialProvider") - .field("path", &self.path) - .field("credential_scope", &self.credential_scope) - .finish_non_exhaustive() - } -} - -impl VendedGcsCredentialProvider { - fn new( - provider: Arc, - path: String, - credential_scope: Option, - ) -> Self { - Self { - provider, - path, - credential_scope, - } - } -} +#[derive(Debug)] +struct VendedGcsCredentialProvider(VendedCredentialSource); impl ProvideCredential for VendedGcsCredentialProvider { type Credential = GoogleCredential; async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { - let credential = self - .provider - .load_credential(&self.path) - .await - .map_err(|e| { - ReqsignError::unexpected(format!( - "failed to load vended GCS credential for {}", - self.path - )) - .with_source(e) - })?; - - validate_credential_prefix( - &self.path, - credential.prefix(), - self.credential_scope.as_ref(), - )?; - - let expires_at = credential - .expires_at() - .map(system_time_to_timestamp) - .transpose()?; - match credential.into_kind() { - StorageCredentialKind::Gcs(gcs) => { + match self.0.load("GCS").await? { + (StorageCredentialKind::Gcs(gcs), expires_at) => { Ok(Some(GoogleCredential::with_token(GoogleToken { access_token: gcs.into_token(), expires_at, @@ -213,92 +158,3 @@ impl ProvideCredential for VendedGcsCredentialProvider { } } } - -#[cfg(test)] -mod tests { - use async_trait::async_trait; - use iceberg::io::{GcsCredential, StorageCredential}; - - use super::*; - - #[derive(Debug)] - struct FixedCredentialProvider(StorageCredential); - - #[async_trait] - impl StorageCredentialProvider for FixedCredentialProvider { - async fn load_credential(&self, _path: &str) -> Result { - Ok(self.0.clone()) - } - } - - fn credential(prefix: Option<&str>) -> StorageCredential { - let credential = StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new( - "access-token", - ))); - match prefix { - Some(prefix) => credential.with_prefix(prefix), - None => credential, - } - } - - fn provider( - path: &str, - returned_prefix: Option<&str>, - credential_scope: Option, - ) -> VendedGcsCredentialProvider { - VendedGcsCredentialProvider::new( - Arc::new(FixedCredentialProvider(credential(returned_prefix))), - path.to_string(), - credential_scope, - ) - } - - #[tokio::test] - async fn vended_provider_enforces_bound_credential_scope() { - let path = "gs://bucket/table/file.parquet"; - let table_prefix = "gs://bucket/table"; - - assert!( - provider( - path, - Some(table_prefix), - Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), - ) - .provide_credential(&Context::default()) - .await - .is_ok() - ); - assert!( - provider( - path, - Some("gs://bucket"), - Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), - ) - .provide_credential(&Context::default()) - .await - .is_err() - ); - assert!( - provider(path, Some("gs://bucket"), None) - .provide_credential(&Context::default()) - .await - .is_ok() - ); - assert!( - provider(path, None, Some(DynamicCredentialScope::Unscoped)) - .provide_credential(&Context::default()) - .await - .is_ok() - ); - assert!( - provider( - path, - Some(table_prefix), - Some(DynamicCredentialScope::Unscoped), - ) - .provide_credential(&Context::default()) - .await - .is_err() - ); - } -} diff --git a/crates/storage/opendal/src/lib.rs b/crates/storage/opendal/src/lib.rs index 8791bd47d8..1ef01f831a 100644 --- a/crates/storage/opendal/src/lib.rs +++ b/crates/storage/opendal/src/lib.rs @@ -172,23 +172,8 @@ impl StorageFactory for OpenDalStorageFactory { config: &StorageConfig, credential_provider: Option>, ) -> Result> { - #[allow(unreachable_patterns)] - let supports_credential_provider = match self { - #[cfg(feature = "opendal-s3")] - OpenDalStorageFactory::S3 { .. } => true, - #[cfg(feature = "opendal-gcs")] - OpenDalStorageFactory::Gcs => true, - #[cfg(feature = "opendal-azdls")] - OpenDalStorageFactory::Azdls => true, - _ => false, - }; - if credential_provider.is_some() && !supports_credential_provider { - return Err(Error::new( - ErrorKind::FeatureUnsupported, - "OpenDAL storage factory does not support refreshable credentials for this backend", - )); - } - + // Only S3, GCS and ADLS consume the provider; other backends use the + // credentials in `config`. match self { #[cfg(feature = "opendal-memory")] OpenDalStorageFactory::Memory => { @@ -248,6 +233,7 @@ fn default_memory_operator() -> Operator { /// OpenDAL-based storage implementation. #[derive(Clone, Debug, Serialize, Deserialize)] +#[non_exhaustive] pub enum OpenDalStorage { /// Memory storage variant. #[cfg(feature = "opendal-memory")] @@ -259,6 +245,7 @@ pub enum OpenDalStorage { /// /// Accepts any S3-family URL (`s3://`, `s3a://`, `s3n://`); the scheme is /// derived from the path at call time. + #[non_exhaustive] #[cfg(feature = "opendal-s3")] S3 { /// S3 configuration. @@ -271,6 +258,7 @@ pub enum OpenDalStorage { credential_provider: Option>, }, /// GCS storage variant. + #[non_exhaustive] #[cfg(feature = "opendal-gcs")] Gcs { /// GCS configuration. @@ -280,6 +268,7 @@ pub enum OpenDalStorage { credential_provider: Option>, }, /// OSS storage variant. + #[non_exhaustive] #[cfg(feature = "opendal-oss")] Oss { /// OSS configuration. @@ -291,6 +280,7 @@ pub enum OpenDalStorage { /// `abfs[s]://@.dfs./` or /// `wasb[s]://@.blob./`. /// The scheme is derived from the path at call time. + #[non_exhaustive] #[cfg(feature = "opendal-azdls")] Azdls { /// Azure DLS configuration. @@ -307,6 +297,7 @@ pub enum OpenDalStorage { /// Accepts paths of the form /// `hf:////[@]/`, /// where `` must be one of `models`, `datasets`, `spaces`, or `buckets`. + #[non_exhaustive] #[cfg(feature = "opendal-hf")] Hf { /// HuggingFace Hub configuration (token + endpoint). @@ -314,31 +305,14 @@ pub enum OpenDalStorage { }, } -#[derive(Clone, Debug, Eq, Hash, PartialEq)] -enum DeleteCredentialScope { - Static, - Dynamic(DynamicCredentialScope), -} - -#[derive(Clone, Debug, Eq, Hash, PartialEq)] -enum DynamicCredentialScope { - Unscoped, - Prefix(String), -} - -impl DynamicCredentialScope { - fn from_prefix(prefix: Option<&str>) -> Self { - match prefix { - Some(prefix) => Self::Prefix(prefix.to_string()), - None => Self::Unscoped, - } - } -} - +/// Groups bulk deletes by operator and credential scope. #[derive(Clone, Debug, Eq, Hash, PartialEq)] struct DeleteBatchKey { storage: String, - credential_scope: DeleteCredentialScope, + /// Location used to look up the dynamic credential shared by the batch: + /// the credential prefix, or the storage root for an unscoped credential. + /// `None` when the path is served by static credentials. + credential_location: Option, } impl OpenDalStorage { @@ -359,16 +333,19 @@ impl OpenDalStorage { &self, path: &'a impl AsRef, ) -> Result<(Operator, &'a str)> { - self.create_operator_with_scope(path, None) + self.create_operator_with_credential_location(path, None) } - /// Creates an operator, optionally binding dynamic credentials to the exact - /// scope used to group a bulk-delete batch. + /// Creates an operator whose dynamic credentials are looked up for + /// `credential_location` instead of `path`. + /// + /// A bulk-delete batch passes its scope location, so every credential the + /// operator loads covers all paths in the batch. #[allow(unreachable_code, unused_variables)] - fn create_operator_with_scope<'a>( + fn create_operator_with_credential_location<'a>( &self, path: &'a impl AsRef, - credential_scope: Option<&DynamicCredentialScope>, + credential_location: Option<&str>, ) -> Result<(Operator, &'a str)> { let path = path.as_ref(); let (operator, relative_path): (Operator, &str) = match self { @@ -400,7 +377,7 @@ impl OpenDalStorage { customized_credential_load, credential_provider, path, - credential_scope, + credential_location, )?; let op_info = op.info(); @@ -428,7 +405,7 @@ impl OpenDalStorage { credential_provider, } => { let operator = - gcs_config_build(config, credential_provider, path, credential_scope)?; + gcs_config_build(config, credential_provider, path, credential_location)?; let url = url::Url::parse(path).map_err(|e| { Error::new( ErrorKind::DataInvalid, @@ -468,7 +445,7 @@ impl OpenDalStorage { config, sas_tokens, credential_provider, - credential_scope, + credential_location, )?, #[cfg(feature = "opendal-hf")] OpenDalStorage::Hf { config } => hf_config_build(config, path)?, @@ -507,16 +484,16 @@ impl OpenDalStorage { /// For most backends the URL host (bucket name) is sufficient. For HF the host /// encodes the repo type, not the repo identity, so a more specific key is used. #[allow(unreachable_patterns)] - fn batch_key_for_path(&self, path: &str) -> String { + fn batch_key_for_path(&self, path: &str) -> Result { match self { #[cfg(feature = "opendal-hf")] - OpenDalStorage::Hf { .. } => hf_batch_key(path), + OpenDalStorage::Hf { .. } => Ok(hf_batch_key(path)), #[cfg(feature = "opendal-azdls")] - OpenDalStorage::Azdls { .. } => azdls_batch_key(path).unwrap_or_default(), - _ => url::Url::parse(path) + OpenDalStorage::Azdls { .. } => azdls_batch_key(path), + _ => Ok(url::Url::parse(path) .ok() .and_then(|u| u.host_str().map(|s| s.to_string())) - .unwrap_or_default(), + .unwrap_or_default()), } } @@ -551,13 +528,10 @@ impl OpenDalStorage { /// scope. Loading the credential is normally a cache hit and avoids rebuilding /// an operator for every path while preventing a batch from crossing prefixes. async fn delete_batch_key_for_path(&self, path: &str) -> Result { - let credential_scope = match self.credential_provider_for_path(path) { + let credential_location = match self.credential_provider_for_path(path) { Some(provider) => { let credential = provider.load_credential(path).await?; - if credential - .prefix() - .is_some_and(|prefix| prefix.is_empty() || !path.starts_with(prefix)) - { + if !credential.covers(path) { return Err(Error::new( ErrorKind::DataInvalid, format!( @@ -566,16 +540,17 @@ impl OpenDalStorage { ), )); } - DeleteCredentialScope::Dynamic(DynamicCredentialScope::from_prefix( - credential.prefix(), - )) + Some(match credential.prefix() { + Some(prefix) => prefix.to_string(), + None => utils::storage_root(path)?, + }) } - None => DeleteCredentialScope::Static, + None => None, }; Ok(DeleteBatchKey { - storage: self.batch_key_for_path(path), - credential_scope, + storage: self.batch_key_for_path(path)?, + credential_location, }) } @@ -761,11 +736,10 @@ impl Storage for OpenDalStorage { (self.relativize_path(&path)?.to_string(), entry.into_mut()) } Entry::Vacant(entry) => { - let credential_scope = match &entry.key().credential_scope { - DeleteCredentialScope::Static => None, - DeleteCredentialScope::Dynamic(scope) => Some(scope), - }; - let (op, rel) = self.create_operator_with_scope(&path, credential_scope)?; + let (op, rel) = self.create_operator_with_credential_location( + &path, + entry.key().credential_location.as_deref(), + )?; let rel = rel.to_string(); let deleter = op.deleter().await.map_err(from_opendal_error)?; (rel, entry.insert(deleter)) @@ -914,15 +888,15 @@ mod tests { ) ))] #[test] - fn test_factory_rejects_credentials_for_unsupported_backend() { - let error = OpenDalStorageFactory::Memory + fn test_factory_ignores_credentials_for_unsupported_backend() { + let storage = OpenDalStorageFactory::Memory .build_with_credentials( &StorageConfig::new(), Some(Arc::new(AlwaysSupportedCredentialProvider)), ) - .expect_err("memory must reject a credential provider"); + .expect("memory must ignore a credential provider"); - assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + assert!(storage.new_input("memory:/key").is_ok()); } #[cfg(feature = "opendal-memory")] @@ -994,10 +968,8 @@ mod tests { let first_key = storage.delete_batch_key_for_path(first).await.unwrap(); assert_eq!( - first_key.credential_scope, - DeleteCredentialScope::Dynamic(DynamicCredentialScope::Prefix( - "s3://bucket/table-a".to_string() - )) + first_key.credential_location.as_deref(), + Some("s3://bucket/table-a") ); assert_eq!( first_key, @@ -1012,6 +984,45 @@ mod tests { ); } + #[cfg(feature = "opendal-s3")] + #[tokio::test] + async fn test_unscoped_dynamic_credentials_batch_by_storage_root() { + #[derive(Debug)] + struct UnscopedProvider; + + #[async_trait] + impl StorageCredentialProvider for UnscopedProvider { + async fn load_credential(&self, _path: &str) -> Result { + Ok(iceberg::io::StorageCredential::new( + iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential::new( + "access-key", + "secret-key", + None, + )), + )) + } + } + + let storage = OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: None, + credential_provider: Some(Arc::new(UnscopedProvider)), + }; + + let key = storage + .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") + .await + .unwrap(); + assert_eq!(key.credential_location.as_deref(), Some("s3://bucket/")); + assert_eq!( + key, + storage + .delete_batch_key_for_path("s3://bucket/table-b/file.parquet") + .await + .unwrap() + ); + } + #[cfg(feature = "opendal-s3")] #[tokio::test] async fn test_custom_s3_credential_loader_ignores_dynamic_provider_for_batching() { @@ -1027,7 +1038,7 @@ mod tests { .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") .await .unwrap(); - assert_eq!(key.credential_scope, DeleteCredentialScope::Static); + assert_eq!(key.credential_location, None); } #[cfg(feature = "opendal-s3")] diff --git a/crates/storage/opendal/src/resolving.rs b/crates/storage/opendal/src/resolving.rs index de61890a81..29a92ad9fa 100644 --- a/crates/storage/opendal/src/resolving.rs +++ b/crates/storage/opendal/src/resolving.rs @@ -80,23 +80,6 @@ fn extract_scheme(path: &str) -> Result<&'static str> { parse_scheme(url.scheme()) } -#[cfg(any( - feature = "opendal-s3", - feature = "opendal-gcs", - feature = "opendal-azdls" -))] -fn supports_dynamic_credentials(scheme: &str) -> bool { - match scheme { - #[cfg(feature = "opendal-s3")] - "s3" => true, - #[cfg(feature = "opendal-gcs")] - "gcs" => true, - #[cfg(feature = "opendal-azdls")] - "azdls" => true, - _ => false, - } -} - /// Build an [`OpenDalStorage`] variant for the given scheme and config properties. #[allow(unused_variables)] fn build_storage_for_scheme( @@ -233,18 +216,8 @@ impl StorageFactory for OpenDalResolvingStorageFactory { config: &StorageConfig, credential_provider: Option>, ) -> Result> { - #[cfg(not(any( - feature = "opendal-s3", - feature = "opendal-gcs", - feature = "opendal-azdls" - )))] - if credential_provider.is_some() { - return Err(Error::new( - ErrorKind::FeatureUnsupported, - "OpenDAL resolving storage does not support refreshable credentials because no compatible backend is enabled", - )); - } - + // Without a compatible backend the provider is ignored, and every + // backend uses the credentials in `config`. Ok(Arc::new(OpenDalResolvingStorage { props: config.props().clone(), storages: RwLock::new(HashMap::new()), @@ -303,25 +276,6 @@ impl OpenDalResolvingStorage { fn resolve(&self, path: &str) -> Result> { let scheme = extract_scheme(path)?; - #[cfg(any( - feature = "opendal-s3", - feature = "opendal-gcs", - feature = "opendal-azdls" - ))] - if self - .credential_provider - .as_ref() - .is_some_and(|provider| provider.supports_path(path)) - && !supports_dynamic_credentials(scheme) - { - return Err(Error::new( - ErrorKind::FeatureUnsupported, - format!( - "OpenDAL resolving storage does not support refreshable credentials for scheme: {scheme}" - ), - )); - } - // Fast path: check read lock first. { let cache = self @@ -459,7 +413,7 @@ mod tests { #[async_trait] impl StorageCredentialProvider for AllPathsCredentialProvider { async fn load_credential(&self, _path: &str) -> Result { - unreachable!("unsupported backends must reject the provider before loading") + unreachable!("unsupported backends must ignore the provider") } } @@ -469,15 +423,15 @@ mod tests { feature = "opendal-azdls" )))] #[test] - fn test_factory_rejects_credentials_without_compatible_backend() { - let error = OpenDalResolvingStorageFactory::new() - .build_with_credentials( - &StorageConfig::new(), - Some(Arc::new(AllPathsCredentialProvider)), - ) - .expect_err("a provider must not be silently discarded"); - - assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + fn test_factory_ignores_credentials_without_compatible_backend() { + assert!( + OpenDalResolvingStorageFactory::new() + .build_with_credentials( + &StorageConfig::new(), + Some(Arc::new(AllPathsCredentialProvider)), + ) + .is_ok() + ); } #[cfg(feature = "opendal-s3")] @@ -557,14 +511,11 @@ mod tests { ) ))] #[test] - fn test_resolver_rejects_credentials_for_unsupported_backend() { + fn test_resolver_ignores_credentials_for_unsupported_backend() { let mut storage = empty_resolving_storage(); storage.credential_provider = Some(Arc::new(AllPathsCredentialProvider)); - let error = storage - .resolve("memory:/key") - .expect_err("memory must reject a credential provider"); - assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); + assert!(storage.resolve("memory:/key").is_ok()); } #[cfg(feature = "opendal-azdls")] diff --git a/crates/storage/opendal/src/s3.rs b/crates/storage/opendal/src/s3.rs index e06bcdefdf..ebe0518221 100644 --- a/crates/storage/opendal/src/s3.rs +++ b/crates/storage/opendal/src/s3.rs @@ -38,10 +38,7 @@ use reqsign_core::{ }; use url::Url; -use crate::DynamicCredentialScope; -use crate::utils::{ - from_opendal_error, is_truthy, system_time_to_timestamp, validate_credential_prefix, -}; +use crate::utils::{VendedCredentialSource, from_opendal_error, is_truthy}; /// Parse iceberg props to s3 config. pub(crate) fn s3_config_parse(mut m: HashMap) -> Result { @@ -138,7 +135,7 @@ pub(crate) fn s3_config_build( customized_credential_load: &Option, credential_provider: &Option>, path: &str, - credential_scope: Option<&DynamicCredentialScope>, + credential_location: Option<&str>, ) -> Result { let url = Url::parse(path)?; let bucket = url.host_str().ok_or_else(|| { @@ -171,10 +168,11 @@ pub(crate) fn s3_config_build( let chain = ProvideCredentialChain::new().push(Arc::clone(&loader.0)); builder = builder.credential_provider_chain(chain); } else if let Some(provider) = credential_provider { - let chain = ProvideCredentialChain::new().push(VendedS3CredentialProvider::new( - Arc::clone(provider), - path.to_string(), - credential_scope.cloned(), + let chain = ProvideCredentialChain::new().push(VendedS3CredentialProvider( + VendedCredentialSource::new( + Arc::clone(provider), + credential_location.unwrap_or(path).to_string(), + ), )); builder = builder.credential_provider_chain(chain); } @@ -185,67 +183,15 @@ pub(crate) fn s3_config_build( /// Adapts a generic [`StorageCredentialProvider`] into a reqsign /// [`ProvideCredential`], so the S3 signer can obtain and refresh vended /// credentials. -struct VendedS3CredentialProvider { - provider: Arc, - /// Absolute path this operator serves; handed back to the provider so it can - /// select the vended credential whose prefix best matches the location. - path: String, - /// Exact scope selected when a bulk-delete operator was created. Ordinary - /// operators are unbound so they can follow the provider's best match. - credential_scope: Option, -} - -impl std::fmt::Debug for VendedS3CredentialProvider { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("VendedS3CredentialProvider") - .field("path", &self.path) - .field("credential_scope", &self.credential_scope) - .finish_non_exhaustive() - } -} - -impl VendedS3CredentialProvider { - fn new( - provider: Arc, - path: String, - credential_scope: Option, - ) -> Self { - Self { - provider, - path, - credential_scope, - } - } -} +#[derive(Debug)] +struct VendedS3CredentialProvider(VendedCredentialSource); impl ProvideCredential for VendedS3CredentialProvider { type Credential = AwsCredential; async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { - let credential = self - .provider - .load_credential(&self.path) - .await - .map_err(|e| { - ReqsignError::unexpected(format!( - "failed to load vended S3 credential for {}", - self.path - )) - .with_source(e) - })?; - - validate_credential_prefix( - &self.path, - credential.prefix(), - self.credential_scope.as_ref(), - )?; - - let expires_in = credential - .expires_at() - .map(system_time_to_timestamp) - .transpose()?; - match credential.into_kind() { - StorageCredentialKind::S3(s3) => { + match self.0.load("S3").await? { + (StorageCredentialKind::S3(s3), expires_in) => { let (access_key_id, secret_access_key, session_token) = s3.into_parts(); Ok(Some(AwsCredential { access_key_id, @@ -292,44 +238,9 @@ impl CustomAwsCredentialLoader { mod tests { use std::collections::HashMap; - use async_trait::async_trait; - use iceberg::io::{S3_PATH_STYLE_ACCESS, S3Credential, StorageCredential}; + use iceberg::io::S3_PATH_STYLE_ACCESS; - use super::*; - - #[derive(Debug)] - struct FixedCredentialProvider(StorageCredential); - - #[async_trait] - impl StorageCredentialProvider for FixedCredentialProvider { - async fn load_credential(&self, _path: &str) -> Result { - Ok(self.0.clone()) - } - } - - fn credential(prefix: Option<&str>) -> StorageCredential { - let credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( - "access-key", - "secret-key", - Some("session-token".to_string()), - ))); - match prefix { - Some(prefix) => credential.with_prefix(prefix), - None => credential, - } - } - - fn provider( - path: &str, - returned_prefix: Option<&str>, - credential_scope: Option, - ) -> VendedS3CredentialProvider { - VendedS3CredentialProvider::new( - Arc::new(FixedCredentialProvider(credential(returned_prefix))), - path.to_string(), - credential_scope, - ) - } + use super::s3_config_parse; fn parse_with(prop: Option<&str>) -> bool { let mut props = HashMap::new(); @@ -346,53 +257,4 @@ mod tests { assert!(parse_with(Some("false"))); assert!(!parse_with(Some("true"))); } - - #[tokio::test] - async fn vended_provider_enforces_bound_credential_scope() { - let path = "s3://bucket/table/file.parquet"; - let table_prefix = "s3://bucket/table"; - - assert!( - provider( - path, - Some(table_prefix), - Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), - ) - .provide_credential(&Context::default()) - .await - .is_ok() - ); - assert!( - provider( - path, - Some("s3://bucket"), - Some(DynamicCredentialScope::Prefix(table_prefix.to_string())), - ) - .provide_credential(&Context::default()) - .await - .is_err() - ); - assert!( - provider(path, Some("s3://bucket"), None) - .provide_credential(&Context::default()) - .await - .is_ok() - ); - assert!( - provider(path, None, Some(DynamicCredentialScope::Unscoped)) - .provide_credential(&Context::default()) - .await - .is_ok() - ); - assert!( - provider( - path, - Some(table_prefix), - Some(DynamicCredentialScope::Unscoped), - ) - .provide_credential(&Context::default()) - .await - .is_err() - ); - } } diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index 1f0830d6eb..d42170c71b 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -15,13 +15,6 @@ // specific language governing permissions and limitations // under the License. -#[cfg(any( - feature = "opendal-s3", - feature = "opendal-gcs", - feature = "opendal-azdls" -))] -use crate::DynamicCredentialScope; - #[cfg(any(feature = "opendal-s3", feature = "opendal-gcs"))] pub(crate) fn is_truthy(value: &str) -> bool { ["true", "t", "1", "on"].contains(&value.to_lowercase().as_str()) @@ -62,37 +55,176 @@ pub(crate) fn system_time_to_timestamp( .map_err(|e| reqsign_core::Error::unexpected(format!("invalid credential expiry: {e}"))) } -/// Validate that a provider's declared credential prefix covers the path for -/// which the backend requested the credential. +/// Validate that a provider's credential covers the location for which the +/// backend requested it. #[cfg(any( feature = "opendal-s3", feature = "opendal-gcs", feature = "opendal-azdls" ))] pub(crate) fn validate_credential_prefix( - path: &str, - prefix: Option<&str>, - credential_scope: Option<&DynamicCredentialScope>, + location: &str, + credential: &iceberg::io::StorageCredential, ) -> reqsign_core::Result<()> { - match prefix { - Some("") => Err(reqsign_core::Error::unexpected( - "vended credential has an empty storage prefix", - )), - Some(prefix) if !path.starts_with(prefix) => Err(reqsign_core::Error::unexpected(format!( - "vended credential prefix {prefix:?} does not cover storage location {path:?}" - ))), - _ => match credential_scope { - Some(DynamicCredentialScope::Unscoped) if prefix.is_some() => { - Err(reqsign_core::Error::unexpected( - "vended credential scope changed while a bulk-delete operator was active", - )) - } - Some(DynamicCredentialScope::Prefix(expected)) if prefix != Some(expected.as_str()) => { - Err(reqsign_core::Error::unexpected( - "vended credential scope changed while a bulk-delete operator was active", + if credential.covers(location) { + Ok(()) + } else { + Err(reqsign_core::Error::unexpected(format!( + "vended credential prefix {:?} does not cover storage location {location:?}", + credential.prefix() + ))) + } +} + +/// The root of the storage location containing `path`, e.g. `s3://bucket/`. +pub(crate) fn storage_root(path: &str) -> iceberg::Result { + let url = url::Url::parse(path)?; + Ok(format!("{}://{}/", url.scheme(), url.authority())) +} + +/// Loads vended credentials for one storage location on behalf of a backend's +/// `reqsign` credential provider. +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +pub(crate) struct VendedCredentialSource { + provider: std::sync::Arc, + /// Location handed to the provider: the path an operator serves, or the + /// scope shared by a bulk-delete batch. + location: String, +} + +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +impl VendedCredentialSource { + pub(crate) fn new( + provider: std::sync::Arc, + location: String, + ) -> Self { + Self { provider, location } + } + + /// Load a credential covering the location, returning its backend-specific + /// material and expiry. + pub(crate) async fn load( + &self, + backend: &str, + ) -> reqsign_core::Result<( + iceberg::io::StorageCredentialKind, + Option, + )> { + let credential = self + .provider + .load_credential(&self.location) + .await + .map_err(|e| { + reqsign_core::Error::unexpected(format!( + "failed to load vended {backend} credential for {}", + self.location )) - } - _ => Ok(()), - }, + .with_source(e) + })?; + validate_credential_prefix(&self.location, &credential)?; + let expires_at = credential + .expires_at() + .map(system_time_to_timestamp) + .transpose()?; + Ok((credential.into_kind(), expires_at)) + } +} + +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +impl std::fmt::Debug for VendedCredentialSource { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("VendedCredentialSource") + .field("location", &self.location) + .finish_non_exhaustive() + } +} + +#[cfg(all( + test, + any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" + ) +))] +mod tests { + use std::sync::{Arc, Mutex}; + + use async_trait::async_trait; + use iceberg::io::{ + GcsCredential, StorageCredential, StorageCredentialKind, StorageCredentialProvider, + }; + + use super::*; + + /// Returns a credential scoped to `prefix` and records requested locations. + #[derive(Debug)] + struct RecordingProvider { + prefix: Option<&'static str>, + requested: Mutex>, + } + + #[async_trait] + impl StorageCredentialProvider for RecordingProvider { + async fn load_credential(&self, path: &str) -> iceberg::Result { + self.requested.lock().unwrap().push(path.to_string()); + let credential = + StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))); + Ok(match self.prefix { + Some(prefix) => credential.with_prefix(prefix), + None => credential, + }) + } + } + + async fn load(prefix: Option<&'static str>, location: &str) -> reqsign_core::Result<()> { + let provider = Arc::new(RecordingProvider { + prefix, + requested: Mutex::new(Vec::new()), + }); + let source = VendedCredentialSource::new(provider.clone(), location.to_string()); + let result = source.load("GCS").await.map(|_| ()); + assert_eq!(*provider.requested.lock().unwrap(), vec![location]); + result + } + + #[tokio::test] + async fn vended_source_requires_a_covering_credential() { + let location = "gs://bucket/table/data/file.parquet"; + assert!(load(None, location).await.is_ok()); + assert!(load(Some("gs://bucket/table"), location).await.is_ok()); + assert!(load(Some("gcs://bucket"), location).await.is_ok()); + assert!( + load(Some("gs://bucket/table/data/other"), location) + .await + .is_err() + ); + assert!(load(Some("gs://bucket/tab"), location).await.is_err()); + assert!(load(Some(""), location).await.is_err()); + } + + #[test] + fn storage_root_keeps_scheme_and_authority() { + assert_eq!( + storage_root("s3a://bucket/table/file.parquet").unwrap(), + "s3a://bucket/" + ); + assert_eq!( + storage_root("abfss://fs@account.dfs.core.windows.net/table/file.parquet").unwrap(), + "abfss://fs@account.dfs.core.windows.net/" + ); + assert!(storage_root("not a url").is_err()); } } From 9ff46f1ca7b9c8ca6bef972597826fcb8c092670 Mon Sep 17 00:00:00 2001 From: Zakariya Stasa Date: Tue, 6 Oct 2026 16:33:34 +0100 Subject: [PATCH 7/7] fix(rest): tighten vended credential refresh contracts and redact all headers - Require StorageCredentialProvider::supports_path and rename StorageFactory::build_with_credentials to build_with_credential_provider. - Pass the whole table config to AuthManager::table_session in process; use the table `token` as a static bearer only for portable providers whose auth type was inferred, so an explicit `rest.auth.type=none` ignores it. - Report non-2xx credential and token responses without echoing the body. - Redact every HTTP header value in debug output unless redaction is disabled, matching Java's HTTPHeader. - Warn once per cause when FileIO serialization drops the provider. - GCS: disable the other credential sources when a provider is attached. - Add tests for ADLS SAS signing, GCS no-fallback, single-flight loads without a usable credential, register_table, and table-client redaction. --- Cargo.lock | 3 + crates/catalog/rest/src/auth/mod.rs | 4 + crates/catalog/rest/src/auth/oauth2.rs | 61 +- crates/catalog/rest/src/catalog.rs | 159 +- crates/catalog/rest/src/client.rs | 262 +-- crates/catalog/rest/src/credential.rs | 2163 +++++++++++++++---- crates/iceberg/public-api.txt | 64 +- crates/iceberg/src/io/file_io.rs | 217 +- crates/iceberg/src/io/storage/config/gcs.rs | 15 +- crates/iceberg/src/io/storage/mod.rs | 335 +-- crates/storage/opendal/Cargo.toml | 9 +- crates/storage/opendal/public-api.txt | 17 +- crates/storage/opendal/src/azdls.rs | 384 +++- crates/storage/opendal/src/gcs.rs | 141 +- crates/storage/opendal/src/lib.rs | 444 +++- crates/storage/opendal/src/resolving.rs | 47 +- crates/storage/opendal/src/s3.rs | 130 +- crates/storage/opendal/src/utils.rs | 230 +- 18 files changed, 3430 insertions(+), 1255 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 7334be0337..53a52425fa 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2830,6 +2830,7 @@ dependencies = [ "futures", "iceberg", "iceberg_test_utils", + "mockito", "opendal", "reqsign-aws-v4", "reqsign-azure-storage", @@ -2837,8 +2838,10 @@ dependencies = [ "reqsign-google", "reqwest 0.12.28", "serde", + "serde_json", "tempfile", "tokio", + "tracing", "typetag", "url", ] diff --git a/crates/catalog/rest/src/auth/mod.rs b/crates/catalog/rest/src/auth/mod.rs index f6bf30ca2c..8498a7f4d9 100644 --- a/crates/catalog/rest/src/auth/mod.rs +++ b/crates/catalog/rest/src/auth/mod.rs @@ -27,6 +27,7 @@ use std::sync::Arc; use async_trait::async_trait; use iceberg::{Error, ErrorKind, Result, TableIdent}; pub use oauth2::OAuth2Manager; +pub(crate) use oauth2::static_token_session; use crate::catalog::{REST_CATALOG_PROP_AUTH_TYPE, RestCatalogConfig}; use crate::client::HttpClient; @@ -103,6 +104,9 @@ pub trait AuthManager: Debug + Send + Sync { /// Returns a session for requests associated with `table`. /// + /// Currently only requests for the table's vended storage credentials use + /// it; other table operations use the catalog session. + /// /// `props` are the unmerged properties returned by the table endpoint. /// The default preserves the catalog session; managers should return a /// child session only when the table properties contain an authentication diff --git a/crates/catalog/rest/src/auth/oauth2.rs b/crates/catalog/rest/src/auth/oauth2.rs index 39a71b53d0..a7beb95c3e 100644 --- a/crates/catalog/rest/src/auth/oauth2.rs +++ b/crates/catalog/rest/src/auth/oauth2.rs @@ -157,17 +157,22 @@ impl AuthManager for OAuth2Manager { props: &HashMap, parent: Arc, ) -> Result> { - let Some(token) = props.get("token") else { - return Ok(parent); - }; - - Ok(Arc::new(OAuth2Session { - token: Arc::new(Mutex::new(Some(SensitiveString::from(token.clone())))), - token_source: TokenSource::StaticToken, - })) + Ok(match props.get("token") { + Some(token) => static_token_session(token), + None => parent, + }) } } +/// A session that authenticates with `token` as a bearer token, which is never +/// refreshed. +pub(crate) fn static_token_session(token: &str) -> Arc { + Arc::new(OAuth2Session { + token: Arc::new(Mutex::new(Some(SensitiveString::from(token.to_string())))), + token_source: TokenSource::StaticToken, + }) +} + impl OAuth2Manager { /// Builds a session from the manager's options with `props` merged onto /// them, so an injected manager keeps whatever a property doesn't @@ -309,15 +314,21 @@ impl ClientCredentialsConfig { let body = response.body(); let auth_res: TokenResponse = if status == StatusCode::OK { + // A token response holds the token, and serde's errors quote the + // values they reject, so neither the body nor the error is kept. Ok(serde_json::from_slice(body).map_err(|e| { Error::new( ErrorKind::Unexpected, - "Failed to parse response from rest catalog server!", + format!( + "Failed to parse response from rest catalog server: {:?} error at \ + line {}, column {}", + e.classify(), + e.line(), + e.column() + ), ) .with_context("operation", "auth") .with_context("url", self.token_endpoint.clone()) - .with_context("json", String::from_utf8_lossy(body)) - .with_source(e) })?) } else { let e: ErrorResponse = serde_json::from_slice(body).map_err(|e| { @@ -454,4 +465,32 @@ mod tests { "Bearer table-token" ); } + + #[tokio::test] + async fn test_unparsable_token_response_is_not_quoted() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/v1/oauth/tokens") + .with_status(200) + .with_body(r#"{"access_token":"SECRET-TOKEN","token_type":7}"#) + .create_async() + .await; + let manager = OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())) + .with_credential(Some("client".to_string()), "secret".to_string()); + let session = manager + .catalog_session(&test_client(), &HashMap::new()) + .await + .unwrap(); + + let mut request = HttpRequest::new( + Client::new() + .get("https://rest.example.com/catalog") + .build() + .unwrap(), + ); + let error = session.authenticate(&mut request).await.unwrap_err(); + let error = format!("{error:?}"); + assert!(!error.contains("SECRET-TOKEN"), "{error}"); + mock.assert_async().await; + } } diff --git a/crates/catalog/rest/src/catalog.rs b/crates/catalog/rest/src/catalog.rs index b1a57ff156..86c7ba3c7c 100644 --- a/crates/catalog/rest/src/catalog.rs +++ b/crates/catalog/rest/src/catalog.rs @@ -60,8 +60,12 @@ pub const REST_CATALOG_PROP_WAREHOUSE: &str = "warehouse"; /// Disable header redaction in error logs and `Debug` output (defaults to /// false for security) pub const REST_CATALOG_PROP_DISABLE_HEADER_REDACTION: &str = "disable-header-redaction"; -/// Identifier for a server-side scan plan associated with credential requests. -pub(crate) const REST_CATALOG_PROP_SCAN_PLAN_ID: &str = "rest.scan.plan-id"; +/// Identifier of the server-side scan plan sent with credential requests as +/// `planId`, as in Java's `RESTCatalogProperties.REST_SCAN_PLAN_ID`. +pub(crate) const REST_CATALOG_PROP_SCAN_PLAN_ID: &str = "rest-scan-plan-id"; +/// Encoded view chain sent with credential requests as `referenced-by`, as in +/// Java's `RESTCatalogProperties.REST_REFERENCED_BY`. +pub(crate) const REST_CATALOG_PROP_REFERENCED_BY: &str = "rest-referenced-by"; /// Authentication scheme: `none` or `oauth2`. When unset, `oauth2` is used /// if a `token`, `credential` or `oauth2-server-uri` is configured, `none` /// otherwise. @@ -877,15 +881,19 @@ impl RestSessionCatalog { // attach a provider so the backend re-fetches them before they expire. // Only catalog authentication resolved from the properties can be // rebuilt after FileIO serialization. + let table_config = table_config.unwrap_or_default(); let credential_provider = build_vended_credential_provider( &client.http_client, client.auth_manager.as_ref(), RestVendedCredentialProviderFactory::new( &client.config.uri, table.clone(), - table_config.unwrap_or_default(), + self.user_config.auth_type(), + !self.user_config.has_explicit_auth_type(), + &table_config, ), &props, + &table_config, self.auth_manager.is_none(), ) .await?; @@ -4155,7 +4163,7 @@ mod tests { } #[tokio::test] - async fn test_injected_auth_manager_file_io_serializes_without_provider_only() { + async fn test_load_table_attaches_the_vended_credential_provider() { let mut server = Server::new_async().await; let config_mock = create_config_mock(&mut server).await; let mut load_response: serde_json::Value = serde_json::from_reader(BufReader::new( @@ -4168,17 +4176,40 @@ mod tests { .unwrap(); load_response["config"]["client.refresh-credentials-endpoint"] = json!("/v1/namespaces/ns1/tables/test1/credentials"); + load_response["config"]["token"] = json!("table-token"); let load_table_mock = server .mock("GET", "/v1/namespaces/ns1/tables/test1") .with_status(200) .with_body(load_response.to_string()) .create_async() .await; + let expires_at = (std::time::SystemTime::now() + std::time::Duration::from_secs(3600)) + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_millis(); + let credentials_mock = server + .mock("GET", "/v1/namespaces/ns1/tables/test1/credentials") + .match_header("authorization", "Bearer table-token") + .with_status(200) + .with_body( + json!({"storage-credentials": [{ + "prefix": "s3://warehouse", + "config": { + "s3.access-key-id": "VENDED_AK", + "s3.secret-access-key": "SK", + "s3.session-token": "TOK", + "s3.session-token-expires-at-ms": expires_at.to_string(), + }, + }]}) + .to_string(), + ) + .create_async() + .await; let catalog = RestCatalog::new( SessionContext::empty(), RestCatalogConfig::builder().uri(server.url()).build(), - Some(Box::new(NoopAuthManager)), + None, Some(Arc::new(LocalFsStorageFactory)), Runtime::current(), None, @@ -4188,15 +4219,80 @@ mod tests { .await .unwrap(); - let error = table.file_io().serialize_all().unwrap_err(); - assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); assert!( - table - .file_io() - .without_credential_provider() - .serialize_all() - .is_ok() + format!("{:?}", table.file_io()).contains("RestVendedCredentialProvider"), + "the table's FileIO holds the provider" + ); + // The serialized provider rebuilds, and refreshes with the table token. + let serialized: serde_json::Value = + serde_json::from_slice(&table.file_io().serialize_all().unwrap()).unwrap(); + let factory: Arc = + serde_json::from_value(serialized["credential_provider"].clone()).unwrap(); + let config = iceberg::io::StorageConfig::new().with_props( + serialized["config"]["props"] + .as_object() + .unwrap() + .iter() + .map(|(key, value)| (key.clone(), value.as_str().unwrap().to_string())), + ); + let credential = factory + .build(&config) + .unwrap() + .load_credential("s3://warehouse/database/table/data.parquet") + .await + .unwrap(); + assert_eq!( + credential + .config() + .get("s3.access-key-id") + .map(String::as_str), + Some("VENDED_AK") + ); + + config_mock.assert_async().await; + load_table_mock.assert_async().await; + credentials_mock.assert_async().await; + } + + #[tokio::test] + async fn test_injected_auth_manager_file_io_serializes_without_provider() { + let mut server = Server::new_async().await; + let config_mock = create_config_mock(&mut server).await; + let mut load_response: serde_json::Value = serde_json::from_reader(BufReader::new( + File::open(format!( + "{}/testdata/load_table_response.json", + env!("CARGO_MANIFEST_DIR") + )) + .unwrap(), + )) + .unwrap(); + load_response["config"]["client.refresh-credentials-endpoint"] = + json!("/v1/namespaces/ns1/tables/test1/credentials"); + let load_table_mock = server + .mock("GET", "/v1/namespaces/ns1/tables/test1") + .with_status(200) + .with_body(load_response.to_string()) + .create_async() + .await; + + let catalog = RestCatalog::new( + SessionContext::empty(), + RestCatalogConfig::builder().uri(server.url()).build(), + Some(Box::new(NoopAuthManager)), + Some(Arc::new(LocalFsStorageFactory)), + Runtime::current(), + None, ); + let table = catalog + .load_table(&TableIdent::from_strs(["ns1", "test1"]).unwrap()) + .await + .unwrap(); + + // The FileIO serializes with its static credentials; the provider, which + // cannot be rebuilt with an injected AuthManager, is dropped. + let serialized: serde_json::Value = + serde_json::from_slice(&table.file_io().serialize_all().unwrap()).unwrap(); + assert!(serialized.get("credential_provider").is_none()); config_mock.assert_async().await; load_table_mock.assert_async().await; @@ -4330,6 +4426,45 @@ mod tests { register_table_mock.assert_async().await; } + #[tokio::test] + async fn test_register_table_attaches_the_vended_credential_provider() { + let mut server = Server::new_async().await; + let config_mock = create_config_mock(&mut server).await; + let mut response: serde_json::Value = serde_json::from_reader(BufReader::new( + File::open(format!( + "{}/testdata/load_table_response.json", + env!("CARGO_MANIFEST_DIR") + )) + .unwrap(), + )) + .unwrap(); + response["config"]["client.refresh-credentials-endpoint"] = + json!("/v1/namespaces/ns1/tables/test1/credentials"); + let register_table_mock = server + .mock("POST", "/v1/namespaces/ns1/register") + .with_status(200) + .with_body(response.to_string()) + .create_async() + .await; + + let catalog = session_catalog(RestCatalogConfig::builder().uri(server.url()).build()); + let table = catalog + .register_table( + &SessionContext::empty(), + &TableIdent::from_strs(["ns1", "test1"]).unwrap(), + "s3://warehouse/database/table/metadata/00001-5f2f8166-244c-4eae-ac36-384ecdec81fc.gz.metadata.json".to_string(), + ) + .await + .unwrap(); + + assert!( + format!("{:?}", table.file_io()).contains("RestVendedCredentialProvider"), + "the registered table's FileIO holds the provider" + ); + config_mock.assert_async().await; + register_table_mock.assert_async().await; + } + #[tokio::test] async fn test_register_table_404() { let mut server = Server::new_async().await; diff --git a/crates/catalog/rest/src/client.rs b/crates/catalog/rest/src/client.rs index 7a9853fe45..7e568a386e 100644 --- a/crates/catalog/rest/src/client.rs +++ b/crates/catalog/rest/src/client.rs @@ -147,23 +147,15 @@ impl HttpClient { /// /// The connection pool and catalog headers are inherited, headers in the /// effective table properties override them, and `auth_session` replaces - /// the catalog session only when the auth manager selected a table-specific - /// child session. + /// the catalog session. As in Java, an explicit `header.Authorization` + /// takes precedence over the session's authentication. pub(crate) fn for_table( &self, props: &HashMap, auth_session: Arc, ) -> Result { let mut extra_headers = self.extra_headers.clone(); - let table_headers = explicit_headers_from_props(props)?; - let has_table_authorization = table_headers.contains_key(http::header::AUTHORIZATION); - - if !Arc::ptr_eq(&self.auth_session, &auth_session) && !has_table_authorization { - // An inherited catalog Authorization header would otherwise be - // applied after, and overwrite, the table session's authentication. - extra_headers.remove(http::header::AUTHORIZATION); - } - extra_headers.extend(table_headers); + extra_headers.extend(explicit_headers_from_props(props)?); Ok(Self { client: self.client.clone(), @@ -243,49 +235,23 @@ pub(crate) fn deserialize_catalog_response( }) } -/// Returns true if the header may carry a secret (matched by substring, so -/// e.g. `x-client-secret` is covered along with `authorization`). -fn is_sensitive_header(name: &str) -> bool { - let name_lower = name.to_lowercase(); - [ - "auth", - "token", - "secret", - "key", - "password", - "cookie", - "credential", - ] - .iter() - .any(|pattern| name_lower.contains(pattern)) -} - -/// Redacts sensitive headers and returns a debug-formatted string. +/// Formats headers for errors and `Debug` output. /// -/// If `disable_redaction` is true, returns all headers without redaction. -/// Otherwise, replaces sensitive header values with `[REDACTED]`. +/// Like Java, every header value is redacted, as any header may carry a +/// secret. With `disable_redaction`, values are shown as they are. pub(crate) fn format_headers_redacted(headers: &HeaderMap, disable_redaction: bool) -> String { - if disable_redaction { - // Return all headers as-is without redaction - let all: HashMap<&str, &str> = headers - .iter() - .filter_map(|(name, value)| value.to_str().ok().map(|v| (name.as_str(), v))) - .collect(); - return format!("{all:?}"); - } - - // Redact sensitive headers by replacing their values with "[REDACTED]" - let redacted: HashMap<&str, &str> = headers + let headers: HashMap<&str, &str> = headers .iter() .filter_map(|(name, value)| { - if is_sensitive_header(name.as_str()) { - Some((name.as_str(), "[REDACTED]")) + let value = if disable_redaction { + value.to_str().ok()? } else { - value.to_str().ok().map(|v| (name.as_str(), v)) - } + "[REDACTED]" + }; + Some((name.as_str(), value)) }) .collect(); - format!("{redacted:?}") + format!("{headers:?}") } /// Deserializes a unexpected catalog response into an error. @@ -323,15 +289,39 @@ fn unexpected_catalog_error( ); let bytes = response.body(); - if !include_body || bytes.is_empty() { + if bytes.is_empty() { return err; } - err.with_context("json", String::from_utf8_lossy(bytes)) + if include_body { + return err.with_context("json", String::from_utf8_lossy(bytes)); + } + // Without the body, keep only the catalog's error description. + match serde_json::from_slice::(bytes) { + Ok(CatalogErrorResponse { error }) => err + .with_context("type", error.r#type) + .with_context("code", error.code.to_string()) + .with_context("message", error.message), + Err(_) => err, + } +} + +/// The parts of a REST catalog error response that describe the error. +#[derive(serde::Deserialize)] +struct CatalogErrorResponse { + error: CatalogError, +} + +#[derive(serde::Deserialize)] +struct CatalogError { + message: String, + r#type: String, + code: u16, } #[cfg(test)] mod tests { use super::*; + use crate::catalog::REST_CATALOG_PROP_DISABLE_HEADER_REDACTION; #[derive(Debug)] struct StaticSession; @@ -417,7 +407,7 @@ mod tests { #[test] fn test_unexpected_error_carries_status_headers_and_body() { // Everything a user needs to diagnose an unexpected status, with the - // sensitive headers held back. + // header values held back. let mut headers = HeaderMap::new(); headers.insert("authorization", "Bearer leaked".parse().unwrap()); headers.insert("x-request-id", "abc123".parse().unwrap()); @@ -434,9 +424,9 @@ mod tests { assert!(err.contains("418"), "{err}"); assert!(err.contains("x-request-id"), "{err}"); - assert!(err.contains("abc123"), "{err}"); assert!(err.contains("nope"), "{err}"); assert!(!err.contains("leaked"), "{err}"); + assert!(!err.contains("abc123"), "{err}"); } #[test] @@ -494,17 +484,18 @@ mod tests { } #[tokio::test] - async fn test_table_client_does_not_inherit_conflicting_authorization_header() { + async fn test_table_client_uses_table_session_and_explicit_headers() { let mut server = mockito::Server::new_async().await; let table_session = server .mock("GET", "/session") .match_header("authorization", "Bearer table-token") + .match_header("x-catalog", "catalog") .with_status(200) .create_async() .await; - let table_header = server + let explicit_header = server .mock("GET", "/header") - .match_header("authorization", "Bearer table-header") + .match_header("authorization", "Bearer explicit-header") .with_status(200) .create_async() .await; @@ -512,45 +503,68 @@ mod tests { let config = RestCatalogConfig::builder() .uri(server.url()) .props(HashMap::from([( - "header.Authorization".to_string(), - "Bearer catalog-header".to_string(), + "header.x-catalog".to_string(), + "catalog".to_string(), )])) .build(); let catalog_client = HttpClient::new(&config).unwrap(); + let get = |client: &HttpClient, path: &str| { + HttpRequest::build(client.request(Method::GET, format!("{}{path}", server.url()))) + .unwrap() + }; + // The table session authenticates, and catalog headers are inherited. let session_client = catalog_client .for_table(&HashMap::new(), Arc::new(TableSession)) .unwrap(); session_client - .query_catalog( - HttpRequest::build( - session_client.request(Method::GET, format!("{}/session", server.url())), - ) - .unwrap(), - ) + .query_catalog(get(&session_client, "/session")) .await .unwrap(); table_session.assert_async().await; + // An explicit Authorization header in the effective properties wins + // over the session, as in Java. let header_client = catalog_client .for_table( &HashMap::from([( "header.Authorization".to_string(), - "Bearer table-header".to_string(), + "Bearer explicit-header".to_string(), )]), Arc::new(TableSession), ) .unwrap(); header_client - .query_catalog( - HttpRequest::build( - header_client.request(Method::GET, format!("{}/header", server.url())), - ) - .unwrap(), - ) + .query_catalog(get(&header_client, "/header")) .await .unwrap(); - table_header.assert_async().await; + explicit_header.assert_async().await; + } + + #[test] + fn test_table_client_redaction_follows_the_table_properties() { + let catalog_client = HttpClient::new( + &RestCatalogConfig::builder() + .uri("http://localhost".to_string()) + .build(), + ) + .unwrap(); + assert!(!catalog_client.disable_header_redaction()); + + let inherited = catalog_client + .for_table(&HashMap::new(), Arc::new(TableSession)) + .unwrap(); + assert!(!inherited.disable_header_redaction()); + let disabled = catalog_client + .for_table( + &HashMap::from([( + REST_CATALOG_PROP_DISABLE_HEADER_REDACTION.to_string(), + "true".to_string(), + )]), + Arc::new(TableSession), + ) + .unwrap(); + assert!(disabled.disable_header_redaction()); } #[test] @@ -561,17 +575,38 @@ mod tests { } #[test] - fn test_format_headers_redacted_non_sensitive() { + fn test_format_headers_redacted_redacts_every_value() { let mut headers = HeaderMap::new(); + headers.insert("authorization", "Bearer secret-token".parse().unwrap()); + headers.insert( + "set-cookie", + "CF_Authorization=sensitive-session-token; Path=/; Secure;" + .parse() + .unwrap(), + ); + headers.insert("x-custom-signature", "unmatched-secret".parse().unwrap()); headers.insert("content-type", "application/json".parse().unwrap()); - headers.insert("x-request-id", "abc123".parse().unwrap()); let result = format_headers_redacted(&headers, false); - assert!(result.contains("content-type")); - assert!(result.contains("application/json")); - assert!(result.contains("x-request-id")); - assert!(result.contains("abc123")); + // Every header name is shown, with its value redacted. + for name in [ + "authorization", + "set-cookie", + "x-custom-signature", + "content-type", + ] { + assert!(result.contains(name), "{result}"); + } + for value in [ + "secret-token", + "sensitive-session-token", + "unmatched-secret", + "application/json", + ] { + assert!(!result.contains(value), "{result}"); + } + assert!(result.contains("[REDACTED]")); } #[tokio::test] @@ -599,81 +634,6 @@ mod tests { assert!(out.contains("[REDACTED]")); } - #[test] - fn test_format_headers_redacted_filters_sensitive() { - let mut headers = HeaderMap::new(); - headers.insert("authorization", "Bearer secret-token".parse().unwrap()); - headers.insert("content-type", "application/json".parse().unwrap()); - - let result = format_headers_redacted(&headers, false); - - // Sensitive header should be present but with redacted value - assert!(result.contains("authorization")); - assert!(result.contains("[REDACTED]")); - // Sensitive value should NOT be present - assert!(!result.contains("secret-token")); - // Non-sensitive header should be present with actual value - assert!(result.contains("content-type")); - assert!(result.contains("application/json")); - } - - #[test] - fn test_format_headers_redacted_filters_set_cookie() { - let mut headers = HeaderMap::new(); - headers.insert( - "set-cookie", - "CF_Authorization=sensitive-session-token; Path=/; Secure;" - .parse() - .unwrap(), - ); - headers.insert("server", "cloudflare".parse().unwrap()); - - let result = format_headers_redacted(&headers, false); - - // Sensitive header should be present but with redacted value - assert!(result.contains("set-cookie")); - assert!(result.contains("[REDACTED]")); - // Sensitive value should NOT be present - assert!(!result.contains("sensitive-session-token")); - // Non-sensitive header should be present with actual value - assert!(result.contains("server")); - assert!(result.contains("cloudflare")); - } - - #[test] - fn test_format_headers_redacted_filters_all_sensitive() { - let mut headers = HeaderMap::new(); - headers.insert("authorization", "Bearer token".parse().unwrap()); - headers.insert("proxy-authorization", "Basic creds".parse().unwrap()); - headers.insert("set-cookie", "session=abc".parse().unwrap()); - headers.insert("cookie", "session=abc".parse().unwrap()); - headers.insert("x-api-key", "api-key-123".parse().unwrap()); - headers.insert("x-auth-token", "auth-token-456".parse().unwrap()); - headers.insert("x-request-id", "req-123".parse().unwrap()); - - let result = format_headers_redacted(&headers, false); - - // All sensitive headers should be present but with redacted values - assert!(result.contains("authorization")); - assert!(result.contains("proxy-authorization")); - assert!(result.contains("set-cookie")); - assert!(result.contains("cookie")); - assert!(result.contains("x-api-key")); - assert!(result.contains("x-auth-token")); - assert!(result.contains("[REDACTED]")); - - // Ensure no sensitive values leaked - assert!(!result.contains("Bearer token")); - assert!(!result.contains("Basic creds")); - assert!(!result.contains("session=abc")); - assert!(!result.contains("api-key-123")); - assert!(!result.contains("auth-token-456")); - - // Non-sensitive header should be present with actual value - assert!(result.contains("x-request-id")); - assert!(result.contains("req-123")); - } - #[test] fn test_format_headers_with_redaction_disabled() { let mut headers = HeaderMap::new(); diff --git a/crates/catalog/rest/src/credential.rs b/crates/catalog/rest/src/credential.rs index ff8a8b0c0b..aac08fd8a5 100644 --- a/crates/catalog/rest/src/credential.rs +++ b/crates/catalog/rest/src/credential.rs @@ -26,11 +26,12 @@ //! //! Unlike the Java client, which has one provider per cloud SDK, this is a //! single backend-agnostic provider with an independent endpoint and cache for -//! each configured cloud. The path being accessed selects the cloud cache, and -//! the returned [`StorageCredential`] enum lets the storage adapter enforce the -//! expected backend-specific type. This preserves Java's per-cloud credential -//! selection and prefetch policies while supporting mixed-cloud tables through -//! a resolving FileIO. Unlike Java's scheduled refresh, which permanently stops +//! each configured cloud. The path being accessed selects the cloud cache. Like +//! Java's `StorageCredential`, the returned [`StorageCredential`] holds the +//! backend's storage properties, which the storage adapter turns into its own +//! credential. This preserves Java's per-cloud credential selection and +//! prefetch policies while supporting mixed-cloud tables through a resolving +//! FileIO. Unlike Java's scheduled refresh, which permanently stops //! after a failed fetch, transient failures are retried here with jittered //! exponential backoff while an unexpired credential remains available. //! @@ -42,9 +43,9 @@ //! # Adding a cloud //! //! The refresh policy for each cloud lives in one [`CloudRefresh`] constant. To -//! add a backend, first add its credential type to Iceberg's storage API and -//! teach the storage adapter to consume it. Then write its `parse_*` function, -//! add a `CloudRefresh` constant, and list it in [`CloudRefresh::SUPPORTED`]. +//! add a backend, teach its storage adapter to read its credential properties +//! from a [`StorageCredential`]. Then write its `parse_*` function, add a +//! `CloudRefresh` constant, and list it in [`CloudRefresh::SUPPORTED`]. use std::collections::{HashMap, HashSet}; use std::sync::Arc; @@ -54,33 +55,43 @@ use async_trait::async_trait; use iceberg::io::{ ADLS_REFRESH_CREDENTIALS_ENABLED, ADLS_REFRESH_CREDENTIALS_ENDPOINT, ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX, ADLS_SAS_TOKEN_PREFIX, AWS_REFRESH_CREDENTIALS_ENABLED, - AWS_REFRESH_CREDENTIALS_ENDPOINT, AzdlsCredential, GCS_REFRESH_CREDENTIALS_ENABLED, - GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, GcsCredential, - S3_ACCESS_KEY_ID, S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, S3_SESSION_TOKEN_EXPIRES_AT_MS, - S3Credential, StorageConfig, StorageCredential, StorageCredentialKind, - StorageCredentialProvider, StorageCredentialProviderFactory, storage_prefix_covers, + AWS_REFRESH_CREDENTIALS_ENDPOINT, GCS_REFRESH_CREDENTIALS_ENABLED, + GCS_REFRESH_CREDENTIALS_ENDPOINT, GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, S3_ACCESS_KEY_ID, + S3_SECRET_ACCESS_KEY, S3_SESSION_TOKEN, S3_SESSION_TOKEN_EXPIRES_AT_MS, StorageConfig, + StorageCredential, StorageCredentialProvider, StorageCredentialProviderFactory, }; use iceberg::{Error, ErrorKind, Result, TableIdent}; use rand::Rng; -use reqwest::{Method, StatusCode, Url}; +use reqwest::{Method, Url}; use serde::{Deserialize, Serialize}; use tokio::sync::{Mutex, OnceCell}; -use crate::auth::{AuthManager, load_auth_manager}; -use crate::catalog::{REST_CATALOG_PROP_SCAN_PLAN_ID, RestCatalogConfig}; +use crate::auth::{AUTH_TYPE_NONE, AuthManager, load_auth_manager, static_token_session}; +use crate::catalog::{ + REST_CATALOG_PROP_AUTH_TYPE, REST_CATALOG_PROP_REFERENCED_BY, REST_CATALOG_PROP_SCAN_PLAN_ID, + RestCatalogConfig, +}; use crate::client::{HttpClient, unexpected_catalog_error_without_body}; use crate::request::HttpRequest; use crate::types::LoadCredentialsResponse; type CredentialParser = - fn(config: &HashMap, prefix: Option) -> Result; + fn(config: &HashMap, prefix: String) -> Result; type KeyedSeedCredentialParser = - fn(config: &HashMap) -> HashMap; + fn(config: &HashMap) -> HashMap; type KeyedSeedPathResolver = fn(path: &str) -> Result; +/// A credential vended by the catalog, with the expiry its config declares. +#[derive(Clone)] +struct VendedCredential { + credential: StorageCredential, + expires_at: SystemTime, +} + enum SeedStrategy { - /// One credential stored in the backend's flat properties. - Flat, + /// One credential stored in the backend's flat properties, for every + /// location with one of the scheme prefixes, like Java's root clients. + Flat { prefixes: &'static [&'static str] }, /// Credentials selected by a backend-specific key derived from each path. Keyed(KeyedSeedStrategy), } @@ -91,7 +102,8 @@ struct KeyedSeedStrategy { } struct KeyedSeedPath { - key: String, + /// Seed keys for the path, most specific first. + keys: Vec, scope: String, } @@ -107,6 +119,9 @@ struct CloudRefresh { endpoint_key: &'static str, /// Table property controlling refresh; only missing or case-insensitive `"true"` enables it. enabled_key: &'static str, + /// Table property that must be present for refresh to apply, as in Java, + /// where it decides whether the backend uses vended credentials at all. + required_key: Option<&'static str>, /// Whether to jitter successful prefetch times like AWS `CachedSupplier`. jitter_prefetch: bool, /// Parse a complete credential from catalog-supplied properties. @@ -121,24 +136,32 @@ impl CloudRefresh { schemes: &["s3", "s3a", "s3n"], endpoint_key: AWS_REFRESH_CREDENTIALS_ENDPOINT, enabled_key: AWS_REFRESH_CREDENTIALS_ENABLED, + required_key: None, jitter_prefetch: true, parse_credential: parse_s3_credential, - seed_strategy: SeedStrategy::Flat, + // `s3` also covers `s3a` and `s3n` locations. + seed_strategy: SeedStrategy::Flat { prefixes: &["s3"] }, }; /// Google Cloud Storage const GCP: Self = Self { schemes: &["gs", "gcs"], endpoint_key: GCS_REFRESH_CREDENTIALS_ENDPOINT, enabled_key: GCS_REFRESH_CREDENTIALS_ENABLED, + // Java's GCS client only refreshes a vended OAuth2 token; without + // one it uses Google's default credentials. + required_key: Some(GCS_TOKEN), jitter_prefetch: false, parse_credential: parse_gcs_credential, - seed_strategy: SeedStrategy::Flat, + seed_strategy: SeedStrategy::Flat { + prefixes: &["gs", "gcs"], + }, }; /// Azure Data Lake Storage const AZURE: Self = Self { schemes: &["abfs", "abfss", "wasb", "wasbs"], endpoint_key: ADLS_REFRESH_CREDENTIALS_ENDPOINT, enabled_key: ADLS_REFRESH_CREDENTIALS_ENABLED, + required_key: None, jitter_prefetch: false, parse_credential: parse_azdls_credential, seed_strategy: SeedStrategy::Keyed(KeyedSeedStrategy { @@ -176,46 +199,44 @@ const MAX_FAILURE_BACKOFF: Duration = Duration::from_secs(30); #[derive(Clone)] struct CachedEntry { credential: StorageCredential, - /// When this entry becomes eligible for prefetch. `None` means it does not - /// expire and therefore never needs proactive refresh. - refresh_at: Option, + expires_at: SystemTime, + /// When this entry becomes eligible for prefetch. + refresh_at: SystemTime, } impl CachedEntry { - fn new(credential: StorageCredential, jitter_prefetch: bool) -> Self { - let refresh_at = credential - .expires_at() - .map(|expires_at| prefetch_time(SystemTime::now(), expires_at, jitter_prefetch)); + fn new(vended: VendedCredential, jitter_prefetch: bool) -> Self { + let VendedCredential { + credential, + expires_at, + } = vended; Self { credential, - refresh_at, + expires_at, + refresh_at: prefetch_time(SystemTime::now(), expires_at, jitter_prefetch), } } /// Seed entries that are already inside the nominal five-minute window are /// immediately due. Otherwise AWS applies the same jitter as it does to a /// freshly fetched value. - fn seed(credential: StorageCredential, jitter_prefetch: bool) -> Self { - let due = credential.expires_at().is_some_and(|expires_at| { - SystemTime::now() - .checked_add(REFRESH_BUFFER) - .is_none_or(|refresh_boundary| refresh_boundary >= expires_at) - }); - let mut entry = Self::new(credential, jitter_prefetch); + fn seed(vended: VendedCredential, jitter_prefetch: bool) -> Self { + let due = SystemTime::now() + .checked_add(REFRESH_BUFFER) + .is_none_or(|refresh_boundary| refresh_boundary >= vended.expires_at); + let mut entry = Self::new(vended, jitter_prefetch); if due { - entry.refresh_at = Some(UNIX_EPOCH); + entry.refresh_at = UNIX_EPOCH; } entry } fn is_fresh(&self, now: SystemTime) -> bool { - self.refresh_at.is_none_or(|refresh_at| now < refresh_at) + now < self.refresh_at } fn is_unexpired(&self, now: SystemTime) -> bool { - self.credential - .expires_at() - .is_none_or(|expires_at| now < expires_at) + now < self.expires_at } } @@ -237,48 +258,97 @@ struct CredentialError { /// re-fetch either. Backoff is therefore shared by the whole cloud. struct CacheState { entries: Vec, + /// Prefixes whose credential in the last successful response was invalid. + /// Paths under them must not use a broader credential: the catalog scoped + /// them more tightly. + failed_prefixes: HashSet, consecutive_failures: u32, retry_not_before: Option, + /// Why the last refresh failed, reported while refresh is backed off. + last_failure: Option, } impl CacheState { fn new(entries: Vec) -> Self { Self { entries, + failed_prefixes: HashSet::new(), consecutive_failures: 0, retry_not_before: None, + last_failure: None, } } + /// The index of the unexpired credential with the longest prefix covering + /// `path`. + fn unexpired_match(&self, path: &str, now: SystemTime) -> Option { + self.entries + .iter() + .enumerate() + .filter(|(_, entry)| entry.is_unexpired(now) && entry.credential.covers(path)) + .max_by_key(|(_, entry)| entry.credential.prefix().len()) + .map(|(index, _)| index) + } + + /// Whether a failed prefix covering `path` is narrower than `credential`'s. + fn masks(&self, path: &str, credential: Option<&StorageCredential>) -> bool { + let prefix_len = credential.map(|credential| credential.prefix().len()); + self.failed_prefixes.iter().any(|failed| { + prefix_covers(failed, path) && prefix_len.is_none_or(|len| failed.len() > len) + }) + } + + /// Cache `entry`, replacing any entry with the same prefix, and return its + /// index. + fn insert(&mut self, entry: CachedEntry) -> usize { + let prefix = entry.credential.prefix().to_owned(); + self.entries + .retain(|cached| cached.credential.prefix() != prefix); + self.entries.push(entry); + self.entries.len() - 1 + } + fn record_success(&mut self) { self.consecutive_failures = 0; self.retry_not_before = None; + self.last_failure = None; } - fn record_failure(&mut self) { + fn record_failure(&mut self, error: &Error) { self.consecutive_failures = self.consecutive_failures.saturating_add(1); self.retry_not_before = Instant::now().checked_add(failure_backoff(self.consecutive_failures)); + self.last_failure = Some(error.to_string()); } - /// Replace cached credentials with the fetched ones, prefix by prefix. + /// Replace cached credentials with the fetched ones, prefix by prefix, + /// and record the prefixes whose credential was invalid in this response. /// /// An unexpired cached credential survives when the response carries no /// valid replacement for its prefix, so an absent or malformed entry never /// evicts a usable credential. Returns the prefixes the response replaced. - fn merge(&mut self, fetched: Vec, now: SystemTime) -> HashSet { + fn merge(&mut self, fetched: Vec, errors: &[CredentialError]) -> HashSet { let fetched_prefixes = fetched .iter() - .filter_map(|entry| entry.credential.prefix().map(str::to_owned)) + .map(|entry| entry.credential.prefix().to_owned()) + .collect::>(); + // A response names every prefix of the table, so a prefix it omits is + // no longer scoped separately: a fetched credential that covers it + // replaces its cached one. A prefix whose credential is invalid keeps + // its cached one. + let failed_prefixes = errors + .iter() + .map(|error| error.prefix.clone()) .collect::>(); self.entries.retain(|entry| { - entry.is_unexpired(now) - && entry - .credential - .prefix() - .is_none_or(|prefix| !fetched_prefixes.contains(prefix)) + let prefix = entry.credential.prefix(); + failed_prefixes.contains(prefix) + || !fetched + .iter() + .any(|fetched| fetched.credential.covers(prefix)) }); self.entries.extend(fetched); + self.failed_prefixes = failed_prefixes; fetched_prefixes } } @@ -321,7 +391,10 @@ impl ConfiguredCloud { ) -> Option { let enabled = props .get(cloud.enabled_key) - .is_none_or(|value| value.eq_ignore_ascii_case("true")); + .is_none_or(|value| value.eq_ignore_ascii_case("true")) + && cloud + .required_key + .is_none_or(|key| props.get(key).is_some_and(|value| !value.is_empty())); let endpoint = props .get(cloud.endpoint_key) .filter(|endpoint| enabled && !endpoint.is_empty()) @@ -330,10 +403,12 @@ impl ConfiguredCloud { let mut entries = Vec::new(); let mut keyed_seeds = HashMap::new(); match &cloud.seed_strategy { - SeedStrategy::Flat => { - if let Ok(credential) = (cloud.parse_credential)(props, None) { - entries.push(CachedEntry::seed(credential, cloud.jitter_prefetch)); - } + SeedStrategy::Flat { prefixes } => { + entries.extend(prefixes.iter().filter_map(|prefix| { + (cloud.parse_credential)(props, prefix.to_string()) + .ok() + .map(|vended| CachedEntry::seed(vended, cloud.jitter_prefetch)) + })); } SeedStrategy::Keyed(strategy) => { keyed_seeds = (strategy.parse_credentials)(props) @@ -352,11 +427,19 @@ impl ConfiguredCloud { return Ok(None); }; let resolved = (strategy.resolve_path)(path)?; - let Some(seed) = self.keyed_seeds.get(&resolved.key) else { + let Some(seed) = resolved + .keys + .iter() + .find_map(|key| self.keyed_seeds.get(key)) + else { return Ok(None); }; - Ok(Some(CachedEntry { - credential: seed.credential.clone().with_prefix(resolved.scope), + let credential = StorageCredential::new(resolved.scope, seed.credential.config().clone()); + // The scope is normalized, e.g. with a lowercase scheme, so it may not + // cover the path as written; the path then fetches its credential. + Ok(credential.covers(path).then_some(CachedEntry { + credential, + expires_at: seed.expires_at, refresh_at: seed.refresh_at, })) } @@ -369,30 +452,62 @@ pub(crate) struct RestVendedCredentialProviderFactory { /// Catalog URI, used to resolve relative refresh endpoints. catalog_uri: String, table: TableIdent, - /// The unmerged config returned by the table endpoint, from which the - /// auth manager derives a table session. Keeping it separate prevents - /// local FileIO overrides from masking table auth. - table_config: HashMap, + /// The `rest.auth.type` the catalog resolved, so a rebuilt provider uses + /// the same auth manager even when the merged FileIO properties would + /// infer another one. + auth_type: String, + /// Whether `auth_type` was inferred rather than configured. Like Java, + /// a table `token` then authenticates a catalog without auth. + auth_type_inferred: bool, + /// The auth-related part of the unmerged config returned by the table + /// endpoint, from which the auth manager derives a table session. Keeping + /// it separate prevents local FileIO overrides from masking table auth. + table_auth: HashMap, } +/// Table config keys that may override authentication in a table session. +const TABLE_AUTH_KEYS: &[&str] = &["token"]; + impl RestVendedCredentialProviderFactory { pub(crate) fn new( catalog_uri: impl Into, table: TableIdent, - table_config: HashMap, + auth_type: impl Into, + auth_type_inferred: bool, + table_config: &HashMap, ) -> Self { Self { catalog_uri: catalog_uri.into(), table, - table_config, + auth_type: auth_type.into(), + auth_type_inferred, + table_auth: table_config + .iter() + .filter(|(key, _)| TABLE_AUTH_KEYS.contains(&key.as_str())) + .map(|(key, value)| (key.clone(), value.clone())) + .collect(), } } + /// The table token that authenticates credential requests of a catalog + /// without auth. Java's credential providers infer OAuth2 from a `token` + /// unless `rest.auth.type` is configured. + fn bearer_table_token(&self) -> Option<&str> { + (self.auth_type == AUTH_TYPE_NONE && self.auth_type_inferred) + .then(|| self.table_auth.get("token").map(String::as_str)) + .flatten() + } + /// Connect to the catalog from the FileIO properties, as the catalog would. async fn connect(&self, props: &HashMap) -> Result { + let mut config_props = props.clone(); + config_props.insert( + REST_CATALOG_PROP_AUTH_TYPE.to_string(), + self.auth_type.clone(), + ); let config = RestCatalogConfig::builder() .uri(self.catalog_uri.clone()) - .props(props.clone()) + .props(config_props) .build(); let auth_manager = load_auth_manager(&config)?; let client = HttpClient::new(&config)?; @@ -404,7 +519,8 @@ impl RestVendedCredentialProviderFactory { auth_manager.as_ref(), &self.table, props, - &self.table_config, + &self.table_auth, + self.bearer_table_token(), ) .await } @@ -437,21 +553,30 @@ impl StorageCredentialProviderFactory for RestVendedCredentialProviderFactory { } /// Derive the table-scoped client used for credential requests. +/// +/// A `bearer_table_token` authenticates instead of a session from the auth +/// manager. async fn table_client( catalog_client: &HttpClient, auth_manager: &dyn AuthManager, table: &TableIdent, props: &HashMap, table_config: &HashMap, + bearer_table_token: Option<&str>, ) -> Result { - let session = auth_manager - .table_session( - &catalog_client.without_auth_session(), - table, - table_config, - catalog_client.auth_session(), - ) - .await?; + let session = match bearer_table_token { + Some(token) => static_token_session(token), + None => { + auth_manager + .table_session( + &catalog_client.without_auth_session(), + table, + table_config, + catalog_client.auth_session(), + ) + .await? + } + }; catalog_client.for_table(props, session) } @@ -470,8 +595,8 @@ pub(crate) struct RestVendedCredentialProvider { factory: Option, /// Effective FileIO properties, which supply catalog connection settings. props: HashMap, - /// Optional scan-plan identifier. - plan_id: Option, + /// Query parameters sent with every credentials request. + query_params: Vec<(&'static str, String)>, /// Independently configured endpoint and cache for each backing cloud. clouds: Vec, } @@ -491,7 +616,7 @@ impl RestVendedCredentialProvider { client: OnceCell::new(), factory: None, props: props.clone(), - plan_id: props.get(REST_CATALOG_PROP_SCAN_PLAN_ID).cloned(), + query_params: credentials_query_params(props), clouds, }) } @@ -520,29 +645,32 @@ impl RestVendedCredentialProvider { async fn fetch(&self, configured: &ConfiguredCloud) -> Result { let cloud = configured.cloud; let client = self.client().await?; - let mut request = client.request(Method::GET, &configured.endpoint); - if let Some(plan_id) = &self.plan_id { - request = request.query(&[("planId", plan_id)]); - } - let request = HttpRequest::build(request)?; + let url = credentials_url(&configured.endpoint, &self.query_params)?; + let request = HttpRequest::build(client.request(Method::GET, url))?; let response = client.query_catalog(request).await?; - if response.status() != StatusCode::OK { + if !response.status().is_success() { return Err(unexpected_catalog_error_without_body( response, client.disable_header_redaction(), )); } - // Credential responses contain secrets. Do not include the response - // body in a deserialization error. + // Credential responses contain secrets, and serde's errors quote the + // values they reject, so only the error's category and position are + // reported. let parsed: LoadCredentialsResponse = serde_json::from_slice(response.body()).map_err(|error| { Error::new( ErrorKind::Unexpected, - "failed to parse vended credential response", + format!( + "failed to parse vended credential response: {:?} error at line {}, \ + column {}", + error.classify(), + error.line(), + error.column() + ), ) - .with_source(error) })?; let now = SystemTime::now(); let mut entries = Vec::new(); @@ -553,8 +681,8 @@ impl RestVendedCredentialProvider { .filter(|credential| cloud.matches_location(&credential.prefix)) { let prefix = credential.prefix; - let parsed = (cloud.parse_credential)(&credential.config, Some(prefix.clone())) - .map(|credential| CachedEntry::new(credential, cloud.jitter_prefetch)) + let parsed = (cloud.parse_credential)(&credential.config, prefix.clone()) + .map(|vended| CachedEntry::new(vended, cloud.jitter_prefetch)) .and_then(|entry| { if entry.is_unexpired(now) { Ok(entry) @@ -582,48 +710,89 @@ impl RestVendedCredentialProvider { let fetched = self.fetch(configured).await; let mut cache = configured.cache.lock().await; let now = SystemTime::now(); + cache.entries.retain(|entry| entry.is_unexpired(now)); let (fetched_prefixes, failure) = match fetched { Ok(ParsedCredentials { entries, errors }) => { + let fetched_prefixes = cache.merge(entries, &errors); + // The most specific error for `path`. let failure = errors .into_iter() - .filter(|error| storage_prefix_covers(&error.prefix, path)) + .filter(|error| prefix_covers(&error.prefix, path)) .max_by_key(|error| error.prefix.len()) .map(|error| error.error); - (cache.merge(entries, now), failure) + (Some(fetched_prefixes), failure) } - Err(error) => (HashSet::new(), Some(error)), + Err(error) => (None, Some(error)), }; - let selected = longest_prefix_match(&cache.entries, path) - .filter(|entry| entry.is_unexpired(now)) - .cloned(); - let refreshed = selected.as_ref().is_some_and(|entry| { - entry - .credential - .prefix() - .is_some_and(|prefix| fetched_prefixes.contains(prefix)) + let selected = cache.unexpired_match(path, now); + // A keyed seed is only cached once served as the fallback. + let seed = selected + .is_none() + .then(|| { + fallback + .clone() + .filter(|fallback| fallback.is_unexpired(now)) + }) + .flatten(); + let candidate = selected + .map(|index| &cache.entries[index].credential) + .or(seed.as_ref().map(|seed| &seed.credential)); + if cache.masks(path, candidate) { + let error = failure.unwrap_or_else(|| masked_credential_error(path)); + cache.record_failure(&error); + return Err(error); + } + + let refreshed = selected.filter(|&index| { + fetched_prefixes + .as_ref() + .is_some_and(|fetched| fetched.contains(cache.entries[index].credential.prefix())) }); - if refreshed { + if let Some(index) = refreshed { + let entry = &mut cache.entries[index]; + // A catalog may return the same credential until shortly before + // it expires. Treat one that is no newer like no answer, so the + // checks stay few as expiry approaches. + if fallback + .as_ref() + .is_some_and(|previous| entry.expires_at <= previous.expires_at) + { + entry.refresh_at = recheck_time(now, entry.expires_at); + } cache.record_success(); - } else { + return Ok(cache.entries[index].credential.clone()); + } + + let selected = selected.or_else(|| seed.map(|seed| cache.insert(seed))); + let Some(index) = selected else { + let error = failure.unwrap_or_else(|| { + Error::new( + ErrorKind::Unexpected, + format!("no unexpired vended credential matches storage location: {path}"), + ) + }); + cache.record_failure(&error); + return Err(error); + }; + // A failed fetch always reports its error as the failure. + match failure { + None => { + // The catalog answered without a newer credential for this + // path, for example when it replaced a scheme-wide credential + // with scoped ones. Keep the current credential and check + // again later instead of on every access. + let entry = &mut cache.entries[index]; + entry.refresh_at = recheck_time(now, entry.expires_at); + cache.record_success(); + } // Graceful degradation: while a credential for this path remains // usable, serve it and retry after jittered backoff. Expired // credentials are never served. - cache.record_failure(); + Some(error) => cache.record_failure(&error), } - - selected - .or_else(|| fallback.filter(|fallback| fallback.is_unexpired(now))) - .map(|entry| entry.credential) - .ok_or_else(|| { - failure.unwrap_or_else(|| { - Error::new( - ErrorKind::Unexpected, - format!("no unexpired vended credential matches storage location: {path}"), - ) - }) - }) + Ok(cache.entries[index].credential.clone()) } } @@ -638,25 +807,37 @@ impl std::fmt::Debug for RestVendedCredentialProvider { enum CacheDecision { Use(StorageCredential), Refresh(Option), - Backoff, + /// Refresh is backed off after the failure described, if known. + Backoff(Option), } -fn refresh_backoff_error(path: &str) -> Error { +fn masked_credential_error(path: &str) -> Error { Error::new( + ErrorKind::DataInvalid, + format!("the catalog vended an invalid credential for storage location: {path}"), + ) +} + +fn refresh_backoff_error(path: &str, last_failure: Option) -> Error { + let error = Error::new( ErrorKind::Unexpected, format!("vended credential refresh is temporarily backed off for storage location: {path}"), - ) + ); + match last_failure { + Some(last_failure) => error.with_context("last_failure", last_failure), + None => error, + } } async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> Result { let keyed_seed = configured.keyed_seed_for_path(path)?; let cache = configured.cache.lock().await; let now = SystemTime::now(); - let current = match longest_prefix_match(&cache.entries, path).cloned() { - Some(cached) if cached.is_unexpired(now) => Some(cached), - Some(expired) => keyed_seed.or(Some(expired)), - None => keyed_seed, - }; + let current = cache + .unexpired_match(path, now) + .map(|index| cache.entries[index].clone()) + .or(keyed_seed) + .filter(|entry| !cache.masks(path, Some(&entry.credential))); if let Some(entry) = current.as_ref().filter(|entry| entry.is_fresh(now)) { return Ok(CacheDecision::Use(entry.credential.clone())); @@ -669,7 +850,7 @@ async fn cache_decision(configured: &ConfiguredCloud, path: &str) -> Result return Ok(credential), CacheDecision::Refresh(current) => current, - CacheDecision::Backoff => return Err(refresh_backoff_error(path)), + CacheDecision::Backoff(last_failure) => { + return Err(refresh_backoff_error(path, last_failure)); + } }; // One caller refreshes, while concurrent callers immediately keep using the @@ -715,7 +898,9 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { let current = match cache_decision(configured, path).await? { CacheDecision::Use(credential) => return Ok(credential), CacheDecision::Refresh(current) => current, - CacheDecision::Backoff => return Err(refresh_backoff_error(path)), + CacheDecision::Backoff(last_failure) => { + return Err(refresh_backoff_error(path, last_failure)); + } }; self.refresh_credential(configured, path, current).await @@ -727,21 +912,12 @@ impl StorageCredentialProvider for RestVendedCredentialProvider { None => Err(Error::new( ErrorKind::FeatureUnsupported, "the vended credential provider cannot be serialized because the REST catalog \ - uses an injected AuthManager, which cannot be rebuilt in another process; use \ - FileIO::without_credential_provider to serialize without credential refresh", + uses an injected AuthManager, which cannot be rebuilt in another process", )), } } } -/// Select the credential whose prefix is the longest match for `path`. -fn longest_prefix_match<'a>(entries: &'a [CachedEntry], path: &str) -> Option<&'a CachedEntry> { - entries - .iter() - .filter(|entry| entry.credential.covers(path)) - .max_by_key(|entry| entry.credential.prefix().map_or(0, str::len)) -} - /// Compute the prefetch time of a credential obtained at `now`. /// /// Like Java, a credential is refreshed [`REFRESH_BUFFER`] before it expires. @@ -769,6 +945,19 @@ fn prefetch_time(now: SystemTime, expires_at: SystemTime, jitter: bool) -> Syste .unwrap_or(base) } +/// When to ask the catalog again for a newer credential after it answered +/// that none exists: at the regular prefetch time, which is halfway to expiry +/// within ten minutes of it, or not at all once less than twice +/// [`MIN_REFRESH_BUFFER`] remains, so the checks stay few as expiry approaches. +fn recheck_time(now: SystemTime, expires_at: SystemTime) -> SystemTime { + let remaining = expires_at.duration_since(now).unwrap_or_default(); + if remaining <= MIN_REFRESH_BUFFER * 2 { + expires_at + } else { + prefetch_time(now, expires_at, false) + } +} + /// Equal-jitter exponential backoff. The random lower half avoids both hot /// retry loops and synchronized retries across clients. fn failure_backoff(consecutive_failures: u32) -> Duration { @@ -796,18 +985,22 @@ pub(crate) async fn build_vended_credential_provider( auth_manager: &dyn AuthManager, factory: RestVendedCredentialProviderFactory, props: &HashMap, + table_config: &HashMap, portable: bool, ) -> Result>> { let Some(provider) = RestVendedCredentialProvider::configure(&factory, props) else { return Ok(None); }; + // An injected auth manager decides on table sessions itself, given the + // whole table config. let client = table_client( catalog_client, auth_manager, &factory.table, props, - &factory.table_config, + table_config, + portable.then(|| factory.bearer_table_token()).flatten(), ) .await?; Ok(Some(Arc::new(RestVendedCredentialProvider { @@ -817,6 +1010,53 @@ pub(crate) async fn build_vended_credential_provider( }))) } +/// The `referenced-by` query parameter. Its value is already percent-encoded. +const REFERENCED_BY_QUERY_PARAMETER: &str = "referenced-by"; + +/// Query parameters for credentials requests, like Java's +/// `RESTUtil.credentialsQueryParams`. +fn credentials_query_params(props: &HashMap) -> Vec<(&'static str, String)> { + [ + (REST_CATALOG_PROP_SCAN_PLAN_ID, "planId"), + ( + REST_CATALOG_PROP_REFERENCED_BY, + REFERENCED_BY_QUERY_PARAMETER, + ), + ] + .into_iter() + .filter_map(|(property, parameter)| props.get(property).map(|value| (parameter, value.clone()))) + .collect() +} + +/// The credentials request URL. Like Java's `HTTPRequest.requestUri`, the +/// already encoded `referenced-by` value is appended verbatim instead of being +/// encoded again. +fn credentials_url(endpoint: &str, query_params: &[(&str, String)]) -> Result { + let mut url = Url::parse(endpoint).map_err(|error| { + Error::new( + ErrorKind::DataInvalid, + format!("invalid credentials endpoint: {endpoint}"), + ) + .with_source(error) + })?; + for (parameter, value) in query_params { + if *parameter != REFERENCED_BY_QUERY_PARAMETER { + url.query_pairs_mut().append_pair(parameter, value); + } + } + if let Some((_, referenced_by)) = query_params + .iter() + .find(|(parameter, _)| *parameter == REFERENCED_BY_QUERY_PARAMETER) + { + let query = match url.query() { + Some(query) => format!("{query}&{REFERENCED_BY_QUERY_PARAMETER}={referenced_by}"), + None => format!("{REFERENCED_BY_QUERY_PARAMETER}={referenced_by}"), + }; + url.set_query(Some(&query)); + } + Ok(url) +} + /// Resolve a possibly-relative refresh endpoint against the catalog base URI. /// /// Absolute endpoints are used as-is and receive the catalog credentials, as @@ -831,6 +1071,12 @@ fn resolve_endpoint(base_uri: &str, endpoint: &str) -> String { format!("{base}{separator}{endpoint}") } +/// Whether the credential `prefix` covers `location`, as in +/// [`StorageCredential::covers`]. +fn prefix_covers(prefix: &str, location: &str) -> bool { + !prefix.is_empty() && location.starts_with(prefix) +} + /// The URL scheme of `location`, lowercased (e.g. `"s3"` for `s3://bucket/k`). fn scheme_of(location: &str) -> Option { Url::parse(location) @@ -841,163 +1087,174 @@ fn scheme_of(location: &str) -> Option { /// Parse a complete S3 credential supplied by the catalog. fn parse_s3_credential( config: &HashMap, - prefix: Option, -) -> Result { - let access_key_id = required_nonempty(config, S3_ACCESS_KEY_ID)?; - let secret_access_key = required_nonempty(config, S3_SECRET_ACCESS_KEY)?; - let session_token = required_nonempty(config, S3_SESSION_TOKEN)?; - let expires_at = required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?; - Ok(with_prefix( - StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( - access_key_id, - secret_access_key, - Some(session_token), - ))) - .with_expiration(expires_at), - prefix, - )) + prefix: String, +) -> Result { + let credential_config = copy_required(config, &[ + S3_ACCESS_KEY_ID, + S3_SECRET_ACCESS_KEY, + S3_SESSION_TOKEN, + S3_SESSION_TOKEN_EXPIRES_AT_MS, + ])?; + Ok(VendedCredential { + expires_at: required_epoch_millis(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?, + credential: StorageCredential::new(prefix, credential_config), + }) } /// Parse a complete GCS credential supplied by the catalog. fn parse_gcs_credential( config: &HashMap, - prefix: Option, -) -> Result { - let token = required_nonempty(config, GCS_TOKEN)?; - let expires_at = required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?; - Ok(with_prefix( - StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new(token))) - .with_expiration(expires_at), - prefix, - )) + prefix: String, +) -> Result { + let credential_config = copy_required(config, &[GCS_TOKEN, GCS_TOKEN_EXPIRES_AT])?; + Ok(VendedCredential { + expires_at: required_epoch_millis(config, GCS_TOKEN_EXPIRES_AT)?, + credential: StorageCredential::new(prefix, credential_config), + }) } -fn with_prefix(credential: StorageCredential, prefix: Option) -> StorageCredential { - match prefix { - Some(prefix) => credential.with_prefix(prefix), - None => credential, - } +/// The non-empty values of `keys` in `config`, which must all be present. +fn copy_required( + config: &HashMap, + keys: &[&str], +) -> Result> { + keys.iter() + .map(|key| Ok((key.to_string(), required_nonempty(config, key)?))) + .collect() } /// Parse a complete account-specific ADLS SAS credential supplied by the catalog. fn parse_azdls_credential( config: &HashMap, - prefix: Option, -) -> Result { - let prefix = prefix.ok_or_else(|| { - Error::new( - ErrorKind::DataInvalid, - "invalid vended ADLS credential: storage prefix is missing", - ) - })?; - let account = azdls_account_name(&prefix)?; - let (suffix, sas_token) = azdls_sas_tokens(config) - .filter(|(_, token_account, _)| *token_account == account) - .map(|(suffix, _, token)| (suffix, token)) - // Prefer the host-keyed token when both forms name the account, as - // the storage backend does. - .max() + prefix: String, +) -> Result { + let location = AzdlsLocation::parse(&prefix)?; + let (suffix, sas_token) = location + .token_keys() + .into_iter() + .find_map(|key| azdls_sas_tokens(config).find(|(suffix, _)| *suffix == key)) .ok_or_else(|| { Error::new( ErrorKind::DataInvalid, - format!("invalid vended credential: no {ADLS_SAS_TOKEN_PREFIX}* token for account {account}"), + format!( + "invalid vended credential: no {ADLS_SAS_TOKEN_PREFIX}* token for {}", + location.host + ), ) })?; - let expires_at = required_epoch_millis( - config, - &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), - )?; - Ok( - StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new( - sas_token, - ))) - .with_prefix(prefix) - .with_expiration(expires_at), - ) + azdls_vended_credential(config, prefix, suffix, sas_token) +} + +/// The vended credential for the SAS token under `suffix`, which must have an +/// expiry under the same suffix. +fn azdls_vended_credential( + config: &HashMap, + prefix: String, + suffix: &str, + sas_token: &str, +) -> Result { + let expiry_key = format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"); + let expires_at = required_epoch_millis(config, &expiry_key)?; + let expires_at_ms = required_nonempty(config, &expiry_key)?; + let credential_config = HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{suffix}"), + sas_token.to_string(), + ), + (expiry_key, expires_at_ms), + ]); + Ok(VendedCredential { + credential: StorageCredential::new(prefix, credential_config), + expires_at, + }) } -/// Parse every complete account-qualified ADLS credential from the initial -/// table properties. Unlike a URI prefix, an Azure account occurs after the -/// filesystem in a location, so these seeds are selected by account name. +/// Parse every complete account-specific ADLS credential from the initial +/// table properties, keyed by their key suffix. Unlike a URI prefix, the host +/// of an Azure location occurs after the filesystem, so these seeds are +/// selected by host, or by account for keys that name only one. Their prefix +/// is set once a location selects them. fn parse_azdls_account_seeds( config: &HashMap, -) -> HashMap { - let mut seeds = azdls_sas_tokens(config) - .filter_map(|(suffix, account, token)| { - let expires_at = required_epoch_millis( - config, - &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), - ) - .ok()?; - Some((suffix, account, token, expires_at)) - }) - .collect::>(); - // Deterministic choice when several keys name the same account: the last, - // host-keyed one wins, as in the storage backend. - seeds.sort_by(|left, right| left.0.cmp(right.0)); - seeds - .into_iter() - .map(|(_, account, token, expires_at)| { - ( - account.to_string(), - StorageCredential::new(StorageCredentialKind::Azdls(AzdlsCredential::new(token))) - .with_expiration(expires_at), - ) +) -> HashMap { + azdls_sas_tokens(config) + .filter_map(|(suffix, token)| { + let vended = azdls_vended_credential(config, String::new(), suffix, token).ok()?; + Some((suffix.to_string(), vended)) }) .collect() } -/// Account-specific SAS tokens as `(key suffix, account, token)`. Like Java, -/// keys may name the host (`account.dfs.core.windows.net`) or only the account. -fn azdls_sas_tokens(config: &HashMap) -> impl Iterator { +/// Account-specific SAS tokens as `(key suffix, token)`. A suffix names a host +/// such as `account.dfs.core.windows.net`, as in Java, or only the account. +fn azdls_sas_tokens(config: &HashMap) -> impl Iterator { config.iter().filter_map(|(key, token)| { let suffix = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; - let account = suffix.split('.').next().unwrap_or(suffix); - (!account.is_empty() && !token.is_empty()).then_some((suffix, account, token.as_str())) + (!suffix.is_empty() && !token.is_empty()).then_some((suffix, token.as_str())) }) } fn resolve_azdls_seed_path(location: &str) -> Result { - let (mut url, account) = parse_azdls_location(location)?; - url.set_path("/"); + let location = AzdlsLocation::parse(location)?; + let keys = location.token_keys(); + let mut url = location.url; + // Without a trailing slash, the scope also covers the container root. + url.set_path(""); url.set_query(None); url.set_fragment(None); Ok(KeyedSeedPath { - key: account, + keys, scope: url.to_string(), }) } -fn azdls_account_name(location: &str) -> Result { - parse_azdls_location(location).map(|(_, account)| account) +/// An ADLS location, with the host and account that select its SAS token. +struct AzdlsLocation { + url: Url, + /// Host, e.g. `account.dfs.core.windows.net`. + host: String, + /// Storage account. + account: String, } -fn parse_azdls_location(location: &str) -> Result<(Url, String)> { - let url = Url::parse(location).map_err(|error| { - Error::new( - ErrorKind::DataInvalid, - format!("invalid ADLS storage location: {location}"), - ) - .with_source(error) - })?; - if !CloudRefresh::AZURE.schemes.contains(&url.scheme()) { - return Err(Error::new( - ErrorKind::DataInvalid, - format!("invalid ADLS storage location scheme: {}", url.scheme()), - )); - } - let account = url - .host_str() - .and_then(|host| host.split('.').next()) - .filter(|account| !account.is_empty()) - .map(str::to_owned) - .ok_or_else(|| { +impl AzdlsLocation { + fn parse(location: &str) -> Result { + let url = Url::parse(location).map_err(|error| { Error::new( ErrorKind::DataInvalid, - format!("ADLS storage location has no account name: {location}"), + format!("invalid ADLS storage location: {location}"), ) + .with_source(error) })?; - Ok((url, account)) + if !CloudRefresh::AZURE.schemes.contains(&url.scheme()) { + return Err(Error::new( + ErrorKind::DataInvalid, + format!("invalid ADLS storage location scheme: {}", url.scheme()), + )); + } + let host = url.host_str().unwrap_or_default().to_string(); + let account = host + .split('.') + .next() + .filter(|account| !account.is_empty()) + .map(str::to_owned) + .ok_or_else(|| { + Error::new( + ErrorKind::DataInvalid, + format!("ADLS storage location has no account name: {location}"), + ) + })?; + Ok(Self { url, host, account }) + } + + /// SAS token key suffixes for this location, most specific first: the + /// exact host, as in Java, then the account alone, as sent for older Java + /// versions and PyIceberg. A key that names only the account matches every + /// host of an account with that name, including one in another cloud; a + /// host-keyed token matches only its host. + fn token_keys(&self) -> Vec { + vec![self.host.clone(), self.account.clone()] + } } fn required_nonempty(config: &HashMap, key: &str) -> Result { @@ -1038,7 +1295,7 @@ mod tests { use mockito::{Matcher, Server}; use super::*; - use crate::auth::{NoopAuthManager, OAuth2Manager}; + use crate::auth::{AUTH_TYPE_NONE, AUTH_TYPE_OAUTH2, NoopAuthManager, OAuth2Manager}; fn epoch_millis(time: SystemTime) -> String { time.duration_since(UNIX_EPOCH) @@ -1047,36 +1304,46 @@ mod tests { .to_string() } - fn s3_cred( - prefix: Option<&str>, - access_key_id: &str, - expires_at: Option, - ) -> StorageCredential { - let mut credential = StorageCredential::new(StorageCredentialKind::S3(S3Credential::new( - access_key_id, - "secret", - None, - ))); - if let Some(prefix) = prefix { - credential = credential.with_prefix(prefix); - } - if let Some(expires_at) = expires_at { - credential = credential.with_expiration(expires_at); + fn s3_vended(prefix: &str, access_key_id: &str, expires_at: SystemTime) -> VendedCredential { + VendedCredential { + credential: StorageCredential::new( + prefix, + HashMap::from([ + (S3_ACCESS_KEY_ID.to_string(), access_key_id.to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "secret".to_string()), + ]), + ), + expires_at, } - credential } fn s3_access_key_id(credential: &StorageCredential) -> &str { - match credential.kind() { - StorageCredentialKind::S3(s3) => s3.access_key_id(), - other => panic!("expected S3 credential, got {other:?}"), - } + credential + .config() + .get(S3_ACCESS_KEY_ID) + .expect("an S3 credential") + } + + fn gcs_token(credential: &StorageCredential) -> &str { + credential + .config() + .get(GCS_TOKEN) + .expect("a GCS credential") + } + + fn sas_token(credential: &StorageCredential) -> &str { + credential + .config() + .iter() + .find(|(key, _)| key.starts_with(ADLS_SAS_TOKEN_PREFIX)) + .map(|(_, token)| token.as_str()) + .expect("an ADLS credential") } /// A previously cached entry: like a seed, its age is unknown, so it is due /// once inside the nominal refresh window. - fn cached_s3(prefix: &str, access_key_id: &str, expires_at: Option) -> CachedEntry { - CachedEntry::seed(s3_cred(Some(prefix), access_key_id, expires_at), false) + fn cached_s3(prefix: &str, access_key_id: &str, expires_at: SystemTime) -> CachedEntry { + CachedEntry::seed(s3_vended(prefix, access_key_id, expires_at), false) } fn test_client(base_uri: &str) -> HttpClient { @@ -1093,7 +1360,9 @@ mod tests { RestVendedCredentialProviderFactory::new( base_uri, test_table(), - table_config.cloned().unwrap_or_default(), + AUTH_TYPE_OAUTH2, + true, + &table_config.cloned().unwrap_or_default(), ) } @@ -1106,7 +1375,7 @@ mod tests { client: OnceCell::new_with(Some(test_client(base_uri))), factory: None, props: HashMap::new(), - plan_id: None, + query_params: Vec::new(), clouds: vec![ConfiguredCloud::new( &CloudRefresh::AWS, format!("{base_uri}/v1/credentials"), @@ -1155,6 +1424,7 @@ mod tests { &NoopAuthManager, test_factory(base_uri, table_auth_props), props, + table_auth_props.unwrap_or(&HashMap::new()), true, ) .await @@ -1236,53 +1506,46 @@ mod tests { #[test] fn parses_s3_credential() { + let prefix = || "s3://bucket".to_string(); let mut config = HashMap::new(); - assert!(parse_s3_credential(&config, None).is_err()); + assert!(parse_s3_credential(&config, prefix()).is_err()); config.insert(S3_ACCESS_KEY_ID.to_string(), "AK".to_string()); - assert!(parse_s3_credential(&config, None).is_err()); + assert!(parse_s3_credential(&config, prefix()).is_err()); config.insert(S3_SECRET_ACCESS_KEY.to_string(), "SK".to_string()); - assert!(parse_s3_credential(&config, None).is_err()); + assert!(parse_s3_credential(&config, prefix()).is_err()); config.insert(S3_SESSION_TOKEN.to_string(), "TOK".to_string()); config.insert( S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), "not-a-timestamp".to_string(), ); - assert!(parse_s3_credential(&config, None).is_err()); + assert!(parse_s3_credential(&config, prefix()).is_err()); config.insert( S3_SESSION_TOKEN_EXPIRES_AT_MS.to_string(), "1500".to_string(), ); - let credential = parse_s3_credential(&config, Some("s3://bucket".to_string())).unwrap(); - assert_eq!(credential.prefix(), Some("s3://bucket")); - assert_eq!( - credential.expires_at(), - Some(UNIX_EPOCH + Duration::from_millis(1500)) - ); - match credential.kind() { - StorageCredentialKind::S3(s3) => assert_eq!(s3.session_token(), Some("TOK")), - other => panic!("expected S3, got {other:?}"), - } + config.insert("unrelated".to_string(), "value".to_string()); + let vended = parse_s3_credential(&config, prefix()).unwrap(); + assert_eq!(vended.credential.prefix(), "s3://bucket"); + assert_eq!(vended.expires_at, UNIX_EPOCH + Duration::from_millis(1500)); + // Only the credential's own properties are kept. + config.remove("unrelated"); + assert_eq!(vended.credential.config(), &config); } #[test] fn parse_gcs_requires_token_and_expiry() { + let prefix = || "gs://bucket".to_string(); let mut config = HashMap::new(); - assert!(parse_gcs_credential(&config, None).is_err()); + assert!(parse_gcs_credential(&config, prefix()).is_err()); config.insert(GCS_TOKEN.to_string(), "ya29.token".to_string()); - assert!(parse_gcs_credential(&config, None).is_err()); + assert!(parse_gcs_credential(&config, prefix()).is_err()); config.insert(GCS_TOKEN_EXPIRES_AT.to_string(), "2000".to_string()); - let credential = parse_gcs_credential(&config, Some("gs://bucket".to_string())).unwrap(); - assert_eq!(credential.prefix(), Some("gs://bucket")); - match credential.kind() { - StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token(), "ya29.token"), - other => panic!("expected GCS, got {other:?}"), - } - assert_eq!( - credential.expires_at(), - Some(UNIX_EPOCH + Duration::from_millis(2000)) - ); + let vended = parse_gcs_credential(&config, prefix()).unwrap(); + assert_eq!(vended.credential.prefix(), "gs://bucket"); + assert_eq!(gcs_token(&vended.credential), "ya29.token"); + assert_eq!(vended.expires_at, UNIX_EPOCH + Duration::from_millis(2000)); } #[test] @@ -1292,41 +1555,79 @@ mod tests { let expiry_key = format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}account1"); let mut config = HashMap::new(); - assert!(parse_azdls_credential(&config, Some(prefix.to_string())).is_err()); + assert!(parse_azdls_credential(&config, prefix.to_string()).is_err()); config.insert(token_key, "sv=2026&sig=secret".to_string()); - assert!(parse_azdls_credential(&config, Some(prefix.to_string())).is_err()); + assert!(parse_azdls_credential(&config, prefix.to_string()).is_err()); config.insert(expiry_key, "2500".to_string()); - let credential = parse_azdls_credential(&config, Some(prefix.to_string())).unwrap(); - assert_eq!(credential.prefix(), Some(prefix)); + let vended = parse_azdls_credential(&config, prefix.to_string()).unwrap(); + assert_eq!(vended.credential.prefix(), prefix); + assert_eq!(vended.expires_at, UNIX_EPOCH + Duration::from_millis(2500)); + assert_eq!(sas_token(&vended.credential), "sv=2026&sig=secret"); + assert_eq!(vended.credential.config(), &config); + } + + #[test] + fn azdls_tokens_match_the_exact_host_or_the_account() { + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let config = HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}acct.dfs.core.windows.net"), + "host".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}acct.dfs.core.windows.net"), + expires.clone(), + ), + ]); + let token = |prefix: &str, config: &HashMap| { + parse_azdls_credential(config, prefix.to_string()) + .map(|vended| sas_token(&vended.credential).to_string()) + }; + assert_eq!( - credential.expires_at(), - Some(UNIX_EPOCH + Duration::from_millis(2500)) + token("abfss://c@acct.dfs.core.windows.net/t", &config).unwrap(), + "host" + ); + // The same account name in another cloud is a different account. + assert!(token("abfss://c@acct.dfs.core.usgovcloudapi.net/t", &config).is_err()); + + // A key naming only the account matches any host of it, after the + // exact host. + let mut config = config; + config.extend([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}acct"), + "account".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}acct"), + expires, + ), + ]); + assert_eq!( + token("abfss://c@acct.dfs.core.windows.net/t", &config).unwrap(), + "host" + ); + assert_eq!( + token("abfss://c@acct.dfs.core.usgovcloudapi.net/t", &config).unwrap(), + "account" ); - match credential.kind() { - StorageCredentialKind::Azdls(azdls) => { - assert_eq!(azdls.sas_token(), "sv=2026&sig=secret") - } - other => panic!("expected ADLS, got {other:?}"), - } } #[test] fn cached_entry_freshness() { - let no_expiry = cached_s3("", "a", None); let far = cached_s3( - "", + "s3", "a", - Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)), + SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60), ); // Within the buffer but not yet expired: stale for a fast-path read, but // still usable for graceful degradation. - let soon = cached_s3("", "a", Some(SystemTime::now() + Duration::from_secs(60))); - let past = cached_s3("", "a", Some(SystemTime::now() - Duration::from_secs(60))); + let soon = cached_s3("s3", "a", SystemTime::now() + Duration::from_secs(60)); + let past = cached_s3("s3", "a", SystemTime::now() - Duration::from_secs(60)); let now = SystemTime::now(); - assert!(no_expiry.is_fresh(now)); - assert!(no_expiry.is_unexpired(now)); assert!(far.is_fresh(now)); assert!(!soon.is_fresh(now)); assert!(soon.is_unexpired(now)); @@ -1370,10 +1671,10 @@ mod tests { #[test] fn seed_inside_nominal_window_is_immediately_due() { - let credential = s3_cred( - None, + let credential = s3_vended( + "s3", "a", - Some(SystemTime::now() + REFRESH_BUFFER - Duration::from_secs(1)), + SystemTime::now() + REFRESH_BUFFER - Duration::from_secs(1), ); let entry = CachedEntry::seed(credential, true); let now = SystemTime::now(); @@ -1399,26 +1700,34 @@ mod tests { } #[test] - fn longest_prefix_match_ignores_freshness() { - let far = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(3600)); - let entries = vec![ + fn unexpired_match_selects_longest_unexpired_prefix() { + let now = SystemTime::now(); + let far = now + REFRESH_BUFFER + Duration::from_secs(3600); + let cache = CacheState::new(vec![ cached_s3("s3://bucket", "wide", far), cached_s3("s3://bucket/warehouse/db", "narrow", far), - ]; - let got = longest_prefix_match(&entries, "s3://bucket/warehouse/db/t/f").unwrap(); + ]); + let got = &cache.entries[cache + .unexpired_match("s3://bucket/warehouse/db/t/f", now) + .unwrap()]; assert_eq!(s3_access_key_id(&got.credential), "narrow"); - assert_eq!(got.credential.prefix(), Some("s3://bucket/warehouse/db")); - assert!(longest_prefix_match(&entries, "s3://other/x").is_none()); - - let fresh = Some(SystemTime::now() + REFRESH_BUFFER + Duration::from_secs(60)); - let stale = Some(SystemTime::now() - Duration::from_secs(60)); - let entries = vec![ - cached_s3("s3://bucket", "wide", fresh), - cached_s3("s3://bucket/table", "narrow-stale", stale), - ]; - let selected = longest_prefix_match(&entries, "s3://bucket/table/f").unwrap(); - assert_eq!(s3_access_key_id(&selected.credential), "narrow-stale"); - assert!(!selected.is_fresh(SystemTime::now())); + assert!(cache.unexpired_match("s3://other/x", now).is_none()); + + // Freshness does not matter, but an expired narrower entry is skipped. + let cache = CacheState::new(vec![ + cached_s3("s3://bucket", "wide", far), + cached_s3("s3://bucket/due", "due", now + Duration::from_secs(60)), + cached_s3( + "s3://bucket/table", + "expired", + now - Duration::from_secs(60), + ), + ]); + let due = &cache.entries[cache.unexpired_match("s3://bucket/due/f", now).unwrap()]; + assert_eq!(s3_access_key_id(&due.credential), "due"); + assert!(!due.is_fresh(now)); + let wide = &cache.entries[cache.unexpired_match("s3://bucket/table/f", now).unwrap()]; + assert_eq!(s3_access_key_id(&wide.credential), "wide"); } #[tokio::test] @@ -1486,56 +1795,48 @@ mod tests { } #[tokio::test] - async fn refresh_includes_scan_plan_id() { + async fn unchanged_refreshed_credential_is_not_refetched_near_expiry() { let mut server = Server::new_async().await; - let body = s3_response( - "s3://bucket", - "AK", - SystemTime::now() + Duration::from_secs(3600), - ); + // Millisecond precision, as the catalog sends it. + let expires_at = UNIX_EPOCH + + Duration::from_millis( + (SystemTime::now() + Duration::from_secs(90)) + .duration_since(UNIX_EPOCH) + .unwrap() + .as_millis() as u64, + ); let mock = server .mock("GET", "/v1/credentials") - .match_query(Matcher::UrlEncoded( - "planId".to_string(), - "scan-plan-1".to_string(), - )) + .expect(1) .with_status(200) .with_header("content-type", "application/json") - .with_body(body) + .with_body(s3_response("s3://bucket", "AK", expires_at)) .create_async() .await; + // The catalog returns the credential the cache already holds. + let provider = provider_with_cached_s3(&server.url(), vec![cached_s3( + "s3://bucket", + "AK", + expires_at, + )]); - // No static creds -> no seed -> first load fetches from the endpoint. - let mut props = aws_refresh_props("/v1/credentials"); - props.insert( - REST_CATALOG_PROP_SCAN_PLAN_ID.to_string(), - "scan-plan-1".to_string(), - ); - - let provider = test_provider(&server.url(), &props).await; - let credential = provider - .load_credential("s3://bucket/warehouse/f") - .await - .unwrap(); - assert_eq!(credential.prefix(), Some("s3://bucket")); - - match credential.kind() { - StorageCredentialKind::S3(s3) => { - assert_eq!(s3.access_key_id(), "AK"); - assert_eq!(s3.session_token(), Some("TOK")); - } - other => panic!("expected S3, got {other:?}"), - } + provider.load_credential("s3://bucket/f").await.unwrap(); + let cache = provider.clouds[0].cache.lock().await; + // With under two minutes left, there is no further check. + assert_eq!(cache.entries[0].refresh_at, expires_at); + drop(cache); mock.assert_async().await; } #[tokio::test] - async fn successful_refresh_caches_entries_before_selecting_path() { + async fn malformed_narrower_credential_is_not_masked_by_broader_one() { let mut server = Server::new_async().await; - let body = s3_response( - "s3://bucket/table-a", - "AK", - SystemTime::now() + Duration::from_secs(3600), + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket","config":{{"s3.access-key-id":"BROAD_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}, + {{"prefix":"s3://bucket/table","config":{{"s3.access-key-id":"BAD"}}}} + ]}}"# ); let mock = server .mock("GET", "/v1/credentials") @@ -1545,63 +1846,668 @@ mod tests { .with_body(body) .create_async() .await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; - let props = aws_refresh_props("/v1/credentials"); - let provider = test_provider(&server.url(), &props).await; - - assert!( - provider - .load_credential("s3://bucket/table-b/file") - .await - .is_err() - ); - // A successful response with no matching credential is negatively - // cached instead of immediately hitting the endpoint again. + let error = provider + .load_credential("s3://bucket/table/f") + .await + .unwrap_err(); + assert!(error.message().contains("is missing or empty"), "{error}"); + // The cached broader credential does not mask the failure later either. assert!( provider - .load_credential("s3://bucket/table-b/file") + .load_credential("s3://bucket/table/other") .await .is_err() ); + // Paths only the broader credential covers still use it. assert_eq!( s3_access_key_id( &provider - .load_credential("s3://bucket/table-a/file") + .load_credential("s3://bucket/other/f") .await .unwrap() ), - "AK" + "BROAD_AK" ); mock.assert_async().await; } #[tokio::test] - async fn malformed_entry_does_not_discard_valid_entries() { + async fn failed_refreshes_do_not_duplicate_the_account_seed() { let mut server = Server::new_async().await; - let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); - let body = format!( - r#"{{"storage-credentials":[ - {{"prefix":"s3://bucket/invalid","config":{{"s3.access-key-id":"BAD"}}}}, - {{"prefix":"s3://bucket/valid","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}} - ]}}"# - ); let mock = server .mock("GET", "/v1/credentials") - .expect(1) - .with_status(200) - .with_header("content-type", "application/json") - .with_body(body) + .expect(2) + .with_status(503) .create_async() .await; - - let props = aws_refresh_props("/v1/credentials"); - let provider = test_provider(&server.url(), &props).await; - - assert!( - provider - .load_credential("s3://bucket/invalid/file") - .await - .is_err() + let host = "acct.dfs.core.windows.net"; + let azdls = |prefix: String, token: &str, expires_at: SystemTime| VendedCredential { + credential: StorageCredential::new( + prefix, + HashMap::from([(format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), token.to_string())]), + ), + expires_at, + }; + let seed = CachedEntry::seed( + azdls( + String::new(), + "seed", + SystemTime::now() + Duration::from_secs(60), + ), + false, + ); + let expired = CachedEntry::new( + azdls( + format!("abfss://c@{host}/data"), + "old", + SystemTime::now() - Duration::from_secs(1), + ), + false, + ); + let configured = ConfiguredCloud::new( + &CloudRefresh::AZURE, + format!("{}/v1/credentials", server.url()), + vec![expired], + HashMap::from([(host.to_string(), seed)]), + ); + let provider = RestVendedCredentialProvider { + client: OnceCell::new_with(Some(test_client(&server.url()))), + factory: None, + props: HashMap::new(), + query_params: Vec::new(), + clouds: vec![configured], + }; + let path = format!("abfss://c@{host}/data/f.parquet"); + + for _ in 0..2 { + let fallback = configured_seed(&provider, &path); + let credential = provider + .refresh_credential(&provider.clouds[0], &path, fallback) + .await + .unwrap(); + assert_eq!(sas_token(&credential), "seed"); + } + // The expired entry is gone and the seed is cached once. + assert_eq!(provider.clouds[0].cache.lock().await.entries.len(), 1); + mock.assert_async().await; + } + + fn configured_seed(provider: &RestVendedCredentialProvider, path: &str) -> Option { + provider.clouds[0].keyed_seed_for_path(path).unwrap() + } + + #[tokio::test] + async fn expired_narrower_credential_falls_back_to_scheme_wide_seed() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(503) + .create_async() + .await; + let seed = CachedEntry::seed( + s3_vended("s3", "SEED_AK", SystemTime::now() + Duration::from_secs(60)), + false, + ); + let expired = CachedEntry::new( + s3_vended( + "s3://bucket/table", + "OLD_AK", + SystemTime::now() - Duration::from_secs(1), + ), + false, + ); + let provider = provider_with_cached_s3(&server.url(), vec![seed, expired]); + + // As for ADLS account seeds, the expired narrower entry is ignored. + for _ in 0..2 { + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/table/f") + .await + .unwrap() + ), + "SEED_AK" + ); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn failed_prefix_clears_once_the_catalog_stops_naming_it() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let broad = format!( + r#"{{"prefix":"s3://bucket","config":{{"s3.access-key-id":"BROAD_AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}}"# + ); + let responses = [ + format!( + r#"{{"storage-credentials":[{broad},{{"prefix":"s3://bucket/table","config":{{"s3.access-key-id":"BAD"}}}}]}}"# + ), + // The catalog fixes the table by vending only the broader credential. + format!(r#"{{"storage-credentials":[{broad}]}}"#), + ]; + let fetches = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let counter = Arc::clone(&fetches); + let mock = server + .mock("GET", "/v1/credentials") + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body_from_request(move |_| { + let fetch = counter.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + responses[fetch.min(1)].clone().into_bytes() + }) + .create_async() + .await; + let provider = provider_with_cached_s3(&server.url(), Vec::new()); + + assert!( + provider + .load_credential("s3://bucket/table/f") + .await + .is_err() + ); + // Skip the backoff, then a later fetch no longer names the prefix. + provider.clouds[0].cache.lock().await.retry_not_before = None; + provider.clouds[0].cache.lock().await.entries[0].refresh_at = UNIX_EPOCH; + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/table/f") + .await + .unwrap() + ), + "BROAD_AK" + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn azdls_seed_is_served_consistently_when_a_broader_credential_fails() { + let mut server = Server::new_async().await; + let host = "acct.dfs.core.windows.net"; + // The catalog's credential for the whole container has no expiry. + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"abfss://c@{host}","config":{{"adls.sas-token.{host}":"sv=2026&sig=bad"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + epoch_millis(SystemTime::now() + Duration::from_secs(120)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + + // The failed prefix is not narrower than the seed, so the seed is + // served on the refresh and during the backoff that follows. + for _ in 0..2 { + let credential = provider + .load_credential(&format!("abfss://c@{host}/data.parquet")) + .await + .unwrap(); + assert_eq!(sas_token(&credential), "sv=2026&sig=seed"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn backoff_error_reports_the_last_failure() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(s3_response( + "s3://bucket/table-a", + "AK", + SystemTime::now() + Duration::from_secs(3600), + )) + .create_async() + .await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + + assert!( + provider + .load_credential("s3://bucket/table-b/f") + .await + .is_err() + ); + let error = provider + .load_credential("s3://bucket/table-b/f") + .await + .unwrap_err() + .to_string(); + assert!(error.contains("backed off"), "{error}"); + assert!( + error.contains("no unexpired vended credential matches"), + "{error}" + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn azdls_container_root_is_not_refetched_on_every_access() { + let mut server = Server::new_async().await; + let host = "acct.dfs.core.windows.net"; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"abfss://c@{host}/table","config":{{"adls.sas-token.{host}":"sv=2026&sig=table","adls.sas-token-expires-at-ms.{host}":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + // The account seed is inside its refresh window. + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + epoch_millis(SystemTime::now() + Duration::from_secs(240)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + + for _ in 0..5 { + let credential = provider + .load_credential(&format!("abfss://c@{host}")) + .await + .unwrap(); + assert_eq!(sas_token(&credential), "sv=2026&sig=seed"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn malformed_response_errors_do_not_quote_secrets() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body( + r#"{"storage-credentials":[{"prefix":"s3://b","config":"s3.secret-access-key=TOPSECRET"}]}"#, + ) + .create_async() + .await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + + for _ in 0..2 { + let error = provider + .load_credential("s3://b/f") + .await + .unwrap_err() + .to_string(); + assert!(!error.contains("TOPSECRET"), "{error}"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn catalog_error_messages_are_reported() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(403) + .with_header("content-type", "application/json") + .with_body( + r#"{"error":{"message":"not allowed to access the table","type":"ForbiddenException","code":403}}"#, + ) + .create_async() + .await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + + let error = provider + .load_credential("s3://b/f") + .await + .unwrap_err() + .to_string(); + assert!(error.contains("not allowed to access the table"), "{error}"); + assert!(error.contains("ForbiddenException"), "{error}"); + mock.assert_async().await; + } + + #[tokio::test] + async fn widened_scope_replaces_the_narrower_credential() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(s3_response( + "s3://b", + "WIDE_AK", + SystemTime::now() + Duration::from_secs(3600), + )) + .create_async() + .await; + // The narrower credential is inside its refresh window. + let provider = provider_with_cached_s3(&server.url(), vec![cached_s3( + "s3://b/t", + "NARROW_AK", + SystemTime::now() + Duration::from_secs(240), + )]); + + for _ in 0..2 { + assert_eq!( + s3_access_key_id(&provider.load_credential("s3://b/t/f").await.unwrap()), + "WIDE_AK" + ); + } + assert_eq!(provider.clouds[0].cache.lock().await.entries.len(), 1); + mock.assert_async().await; + } + + #[tokio::test] + async fn azdls_seed_is_not_served_for_a_path_it_does_not_cover() { + let mut server = Server::new_async().await; + let host = "acct.dfs.core.windows.net"; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"ABFSS://c@{host}","config":{{"adls.sas-token.{host}":"sv=2026&sig=fetched","adls.sas-token-expires-at-ms.{host}":"{expires}"}}}}]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + expires.clone(), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + + // The seed's scope has a lowercase scheme, so it does not cover this + // path, which fetches a credential that does. + let path = format!("ABFSS://c@{host}/data.parquet"); + let credential = provider.load_credential(&path).await.unwrap(); + assert!(credential.covers(&path)); + assert_eq!(sas_token(&credential), "sv=2026&sig=fetched"); + mock.assert_async().await; + } + + #[tokio::test] + async fn concurrent_loads_without_a_usable_credential_fetch_once() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + )) + .create_async() + .await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + + // Waiters for the in-flight refresh use its result instead of fetching. + let loads = (0..8).map(|_| { + let provider = Arc::clone(&provider); + tokio::spawn(async move { provider.load_credential("s3://bucket/x/f").await }) + }); + for load in futures::future::join_all(loads).await { + assert_eq!(s3_access_key_id(&load.unwrap().unwrap()), "AK"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn failed_credential_requests_do_not_report_the_body() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(500) + .with_header("content-type", "application/json") + .with_body( + r#"{"error":{"message":"internal error","type":"ServerError","code":500},"storage-credentials":[{"prefix":"s3://b","config":{"s3.secret-access-key":"TOPSECRET"}}]}"#, + ) + .create_async() + .await; + let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; + + let error = provider + .load_credential("s3://b/f") + .await + .unwrap_err() + .to_string(); + assert!(error.contains("internal error"), "{error}"); + assert!(!error.contains("TOPSECRET"), "{error}"); + mock.assert_async().await; + } + + #[tokio::test] + async fn rebuilt_provider_does_not_infer_another_auth_type() { + let mut server = Server::new_async().await; + // The merged FileIO properties hold a `credential`, which would infer + // oauth2 and exchange it; the catalog resolved no auth. + let token_mock = server + .mock("POST", "/v1/oauth/tokens") + .expect(0) + .create_async() + .await; + let credentials_mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", Matcher::Missing) + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + )) + .create_async() + .await; + let mut props = aws_refresh_props("/v1/credentials"); + props.insert("credential".to_string(), "client:secret".to_string()); + let rebuilt = RestVendedCredentialProviderFactory::new( + server.url(), + test_table(), + AUTH_TYPE_NONE, + false, + &HashMap::new(), + ) + .build(&StorageConfig::new().with_props(props)) + .unwrap(); + + rebuilt.load_credential("s3://bucket/x/f").await.unwrap(); + token_mock.assert_async().await; + credentials_mock.assert_async().await; + } + + #[tokio::test] + async fn gcs_refresh_requires_a_vended_token() { + let mut props = HashMap::from([( + CloudRefresh::GCP.endpoint_key.to_string(), + "/v1/credentials".to_string(), + )]); + // Without a vended token, like Java, GCS keeps its default credentials. + assert!( + build_test_provider(test_client("http://cat"), "http://cat", &props, None) + .await + .unwrap() + .is_none() + ); + + props.insert(GCS_TOKEN.to_string(), "TOKEN".to_string()); + let provider = test_provider("http://cat", &props).await; + assert!(provider.supports_path("gs://bucket/x")); + } + + #[tokio::test] + async fn refresh_includes_java_query_parameters() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + // `referenced-by` is sent verbatim, without encoding `%` again. + .match_query(Matcher::Exact( + "planId=scan-plan-1&referenced-by=ns%1Fview,ns%1Fother".to_string(), + )) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + // No static creds -> no seed -> first load fetches from the endpoint. + let mut props = aws_refresh_props("/v1/credentials"); + props.extend([ + ( + REST_CATALOG_PROP_SCAN_PLAN_ID.to_string(), + "scan-plan-1".to_string(), + ), + ( + REST_CATALOG_PROP_REFERENCED_BY.to_string(), + "ns%1Fview,ns%1Fother".to_string(), + ), + ]); + + let provider = test_provider(&server.url(), &props).await; + let credential = provider + .load_credential("s3://bucket/warehouse/f") + .await + .unwrap(); + assert_eq!(credential.prefix(), "s3://bucket"); + assert_eq!(s3_access_key_id(&credential), "AK"); + assert_eq!( + credential + .config() + .get(S3_SESSION_TOKEN) + .map(String::as_str), + Some("TOK") + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn successful_refresh_caches_entries_before_selecting_path() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket/table-a", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let provider = test_provider(&server.url(), &props).await; + + assert!( + provider + .load_credential("s3://bucket/table-b/file") + .await + .is_err() + ); + // A successful response with no matching credential is negatively + // cached instead of immediately hitting the endpoint again. + assert!( + provider + .load_credential("s3://bucket/table-b/file") + .await + .is_err() + ); + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/table-a/file") + .await + .unwrap() + ), + "AK" + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn malformed_entry_does_not_discard_valid_entries() { + let mut server = Server::new_async().await; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[ + {{"prefix":"s3://bucket/invalid","config":{{"s3.access-key-id":"BAD"}}}}, + {{"prefix":"s3://bucket/valid","config":{{"s3.access-key-id":"AK","s3.secret-access-key":"SK","s3.session-token":"TOK","s3.session-token-expires-at-ms":"{expires}"}}}} + ]}}"# + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let provider = test_provider(&server.url(), &props).await; + + assert!( + provider + .load_credential("s3://bucket/invalid/file") + .await + .is_err() ); assert!( provider @@ -1640,7 +2546,7 @@ mod tests { let provider = test_provider(&server.url(), &props).await; let credential = provider.load_credential("s3://bucket/x/f").await.unwrap(); - assert_eq!(credential.prefix(), None); + assert_eq!(credential.prefix(), "s3"); assert_eq!(s3_access_key_id(&credential), "SEED_AK"); mock.assert_async().await; } @@ -1683,14 +2589,9 @@ mod tests { .unwrap(); assert_eq!( credential.prefix(), - Some("abfss://container@account1.dfs.core.windows.net/") + "abfss://container@account1.dfs.core.windows.net" ); - match credential.kind() { - StorageCredentialKind::Azdls(azdls) => { - assert_eq!(azdls.sas_token(), "sv=2026&sig=seed") - } - other => panic!("expected ADLS credential, got {other:?}"), - } + assert_eq!(sas_token(&credential), "sv=2026&sig=seed"); let credential = provider .load_credential("abfss://other-container@account2.dfs.core.windows.net/data/b.parquet") @@ -1698,14 +2599,9 @@ mod tests { .unwrap(); assert_eq!( credential.prefix(), - Some("abfss://other-container@account2.dfs.core.windows.net/") + "abfss://other-container@account2.dfs.core.windows.net" ); - match credential.kind() { - StorageCredentialKind::Azdls(azdls) => { - assert_eq!(azdls.sas_token(), "sv=2026&sig=other-seed") - } - other => panic!("expected ADLS credential, got {other:?}"), - } + assert_eq!(sas_token(&credential), "sv=2026&sig=other-seed"); mock.assert_async().await; } @@ -1809,7 +2705,7 @@ mod tests { let provider = provider_with_cached_s3(&server.url(), vec![cached_s3( "s3://bucket/requested", "SEED_AK", - Some(SystemTime::now() + Duration::from_secs(60)), + SystemTime::now() + Duration::from_secs(60), )]); let fallback = provider @@ -1860,22 +2756,22 @@ mod tests { cached_s3( "s3://bucket/table-a", "OLD_AK", - Some(SystemTime::now() + Duration::from_secs(60)), + SystemTime::now() + Duration::from_secs(60), ), cached_s3( "s3://bucket/table-b", "OLDER_BK", - Some(SystemTime::now() + Duration::from_secs(3600)), + SystemTime::now() + Duration::from_secs(3600), ), cached_s3( "s3://bucket/table-b", "LATEST_BK", - Some(SystemTime::now() + Duration::from_secs(3600)), + SystemTime::now() + Duration::from_secs(3600), ), cached_s3( "s3://bucket/table-c", "VALID_CK", - Some(SystemTime::now() + Duration::from_secs(3600)), + SystemTime::now() + Duration::from_secs(3600), ), ]); @@ -1899,6 +2795,131 @@ mod tests { mock.assert_async().await; } + #[tokio::test] + async fn scoped_refresh_after_scheme_wide_seed_does_not_back_off() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket/table", + "SCOPED_AK", + SystemTime::now() + Duration::from_secs(3600), + ); + let mock = server + .mock("GET", "/v1/credentials") + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + // The scheme-wide seed is inside its refresh window. + let seed = CachedEntry::seed( + s3_vended("s3", "SEED_AK", SystemTime::now() + Duration::from_secs(60)), + false, + ); + let provider = provider_with_cached_s3(&server.url(), vec![seed]); + + assert_eq!( + s3_access_key_id( + &provider + .load_credential("s3://bucket/table/data/f.parquet") + .await + .unwrap() + ), + "SCOPED_AK" + ); + // A location only the seed covers refreshes once, finds no newer + // credential, and keeps using the seed. This is not a failure, so the + // cloud is not backed off. + for location in ["s3://bucket/", "s3://bucket/", "s3://bucket/other/f"] { + assert_eq!( + s3_access_key_id(&provider.load_credential(location).await.unwrap()), + "SEED_AK" + ); + } + let cache = provider.clouds[0].cache.lock().await; + assert_eq!(cache.consecutive_failures, 0); + assert!(cache.retry_not_before.is_none()); + drop(cache); + mock.assert_async().await; + } + + #[tokio::test] + async fn scoped_refresh_after_azdls_account_seed_does_not_back_off() { + let mut server = Server::new_async().await; + let host = "account1.dfs.core.windows.net"; + let expires = epoch_millis(SystemTime::now() + Duration::from_secs(3600)); + let body = format!( + r#"{{"storage-credentials":[{{"prefix":"abfss://container@{host}/table","config":{{"adls.sas-token.{host}":"sv=2026&sig=scoped","adls.sas-token-expires-at-ms.{host}":"{expires}"}}}}]}}"# + ); + // One fetch for the container root, which finds nothing newer, and one + // for an account without credentials. A backed-off cloud would skip + // the second. + let mock = server + .mock("GET", "/v1/credentials") + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + // The account seed is inside its refresh window. + let props = HashMap::from([ + ( + CloudRefresh::AZURE.endpoint_key.to_string(), + "/v1/credentials".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=seed".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + epoch_millis(SystemTime::now() + Duration::from_secs(240)), + ), + ]); + let provider = test_provider(&server.url(), &props).await; + let sas_token = |credential: StorageCredential| sas_token(&credential).to_string(); + + let root = format!("abfss://container@{host}/"); + for _ in 0..2 { + assert_eq!( + sas_token(provider.load_credential(&root).await.unwrap()), + "sv=2026&sig=seed" + ); + } + assert_eq!( + sas_token( + provider + .load_credential(&format!("{root}table/data.parquet")) + .await + .unwrap() + ), + "sv=2026&sig=scoped" + ); + assert!( + provider + .load_credential("abfss://container@account2.dfs.core.windows.net/a") + .await + .is_err() + ); + mock.assert_async().await; + } + + #[test] + fn recheck_time_stops_near_expiry() { + let now = SystemTime::now(); + assert_eq!( + recheck_time(now, now + Duration::from_secs(240)), + now + Duration::from_secs(120) + ); + let soon = now + MIN_REFRESH_BUFFER * 2; + assert_eq!(recheck_time(now, soon), soon); + // Longer lifetimes keep the regular prefetch time. + let later = now + Duration::from_secs(3600); + assert_eq!(recheck_time(now, later), later - REFRESH_BUFFER); + } + #[tokio::test] async fn concurrent_prefetch_serves_unexpired_credential_without_waiting() { let mut server = Server::new_async().await; @@ -2001,6 +3022,7 @@ mod tests { &auth_manager, test_factory(&server.url(), Some(&table_auth)), &props, + &table_auth, true, ) .await @@ -2052,6 +3074,7 @@ mod tests { &auth_manager, test_factory(&server.url(), Some(&table_auth)), &props, + &table_auth, true, ) .await @@ -2098,6 +3121,12 @@ mod tests { CloudRefresh::GCP.endpoint_key.to_string(), "/v1/gcp-credentials".to_string(), ), + // GCS refreshes only a vended token; this one already expired. + (GCS_TOKEN.to_string(), "EXPIRED".to_string()), + ( + GCS_TOKEN_EXPIRES_AT.to_string(), + epoch_millis(SystemTime::now() - Duration::from_secs(60)), + ), ]); let provider = test_provider(&server.url(), &props).await; @@ -2107,15 +3136,10 @@ mod tests { s3_access_key_id(&provider.load_credential("s3://bucket/x").await.unwrap()), "AK" ); - match provider - .load_credential("gs://bucket/x") - .await - .unwrap() - .kind() - { - StorageCredentialKind::Gcs(gcs) => assert_eq!(gcs.token(), "GCS"), - other => panic!("expected GCS credential, got {other:?}"), - } + assert_eq!( + gcs_token(&provider.load_credential("gs://bucket/x").await.unwrap()), + "GCS" + ); aws_mock.assert_async().await; gcp_mock.assert_async().await; } @@ -2152,10 +3176,7 @@ mod tests { ), ]); let provider = test_provider(&server.url(), &props).await; - let sas_token = |credential: StorageCredential| match credential.into_kind() { - StorageCredentialKind::Azdls(azdls) => azdls.into_sas_token(), - other => panic!("expected ADLS credential, got {other:?}"), - }; + let sas_token = |credential: StorageCredential| sas_token(&credential).to_string(); // The host-keyed seed serves the account without fetching. let seeded = provider @@ -2203,12 +3224,7 @@ mod tests { .load_credential(&format!("{prefix}/data.parquet")) .await .unwrap(); - match credential.kind() { - StorageCredentialKind::Azdls(azdls) => { - assert_eq!(azdls.sas_token(), "sv=2026&sig=secret") - } - other => panic!("expected ADLS credential, got {other:?}"), - } + assert_eq!(sas_token(&credential), "sv=2026&sig=secret"); mock.assert_async().await; } @@ -2267,14 +3283,9 @@ mod tests { let credential = provider.load_credential(path).await.unwrap(); assert_eq!( credential.prefix(), - Some("abfss://container@account2.dfs.core.windows.net/") + "abfss://container@account2.dfs.core.windows.net" ); - match credential.kind() { - StorageCredentialKind::Azdls(azdls) => { - assert_eq!(azdls.sas_token(), "sv=2026&sig=seed") - } - other => panic!("expected ADLS credential, got {other:?}"), - } + assert_eq!(sas_token(&credential), "sv=2026&sig=seed"); } mock.assert_async().await; } @@ -2332,7 +3343,7 @@ mod tests { } #[tokio::test] - async fn scheme_aliases_share_prefix_scoped_credentials() { + async fn prefixes_match_locations_like_java() { let mut server = Server::new_async().await; let body = s3_response( "s3://bucket/table", @@ -2349,16 +3360,17 @@ mod tests { .await; let provider = test_provider(&server.url(), &aws_refresh_props("/v1/credentials")).await; - for path in ["s3a://bucket/table/f", "s3://bucket/table/f"] { + // As in Java, prefixes are compared as plain strings. + for path in ["s3://bucket/table/f", "s3://bucket/table2/f"] { assert_eq!( s3_access_key_id(&provider.load_credential(path).await.unwrap()), "AK" ); } - // A prefix only covers whole path segments, so this path refreshes. + // `s3a` is another scheme, so this path is not covered and refreshes. assert!( provider - .load_credential("s3://bucket/table2/f") + .load_credential("s3a://bucket/table/f") .await .is_err() ); @@ -2377,7 +3389,7 @@ mod tests { .mock("GET", "/v1/credentials") .match_header("authorization", "Bearer table-token") .match_header("x-custom", "value") - .expect(1) + .expect(2) .with_status(200) .with_header("content-type", "application/json") .with_body(body) @@ -2392,14 +3404,16 @@ mod tests { let table_auth = HashMap::from([("token".to_string(), "table-token".to_string())]); let provider = build_vended_credential_provider( &test_client(&server.url()), - &NoopAuthManager, + &OAuth2Manager::new(format!("{}/v1/oauth/tokens", server.url())), test_factory(&server.url(), Some(&table_auth)), &props, + &table_auth, true, ) .await .unwrap() .unwrap(); + provider.load_credential("s3://bucket/x/f").await.unwrap(); let factory: Arc = serde_json::from_str(&serde_json::to_string(&provider.factory().unwrap()).unwrap()) @@ -2407,6 +3421,7 @@ mod tests { let config = StorageConfig::new().with_props(props); let rebuilt = factory.build(&config).unwrap(); + // The rebuilt provider authenticates like the original one. assert_eq!( s3_access_key_id(&rebuilt.load_credential("s3://bucket/x/f").await.unwrap()), "AK" @@ -2416,6 +3431,216 @@ mod tests { mock.assert_async().await; } + #[tokio::test] + async fn table_token_authenticates_without_catalog_auth() { + let mut server = Server::new_async().await; + let body = s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + ); + // As in Java, a table token authenticates credential requests of a + // catalog without auth, in process and after the provider is rebuilt. + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", "Bearer table-token") + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body) + .create_async() + .await; + + let props = aws_refresh_props("/v1/credentials"); + let factory = RestVendedCredentialProviderFactory::new( + server.url(), + test_table(), + AUTH_TYPE_NONE, + true, + &HashMap::from([("token".to_string(), "table-token".to_string())]), + ); + let provider = build_vended_credential_provider( + &test_client(&server.url()), + &NoopAuthManager, + factory.clone(), + &props, + &HashMap::from([("token".to_string(), "table-token".to_string())]), + true, + ) + .await + .unwrap() + .unwrap(); + provider.load_credential("s3://bucket/x/f").await.unwrap(); + + let rebuilt = factory + .build(&StorageConfig::new().with_props(props)) + .unwrap(); + rebuilt.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn explicit_no_auth_ignores_the_table_token() { + let mut server = Server::new_async().await; + // As in Java, an explicit `rest.auth.type=none` is not overridden by a + // table token, in process and after the provider is rebuilt. + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", Matcher::Missing) + .expect(2) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + )) + .create_async() + .await; + let table_config = HashMap::from([("token".to_string(), "table-token".to_string())]); + let props = aws_refresh_props("/v1/credentials"); + let factory = RestVendedCredentialProviderFactory::new( + server.url(), + test_table(), + AUTH_TYPE_NONE, + false, + &table_config, + ); + let provider = build_vended_credential_provider( + &test_client(&server.url()), + &NoopAuthManager, + factory.clone(), + &props, + &table_config, + true, + ) + .await + .unwrap() + .unwrap(); + provider.load_credential("s3://bucket/x/f").await.unwrap(); + let rebuilt = factory + .build(&StorageConfig::new().with_props(props)) + .unwrap(); + rebuilt.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn injected_auth_manager_receives_the_whole_table_config() { + #[derive(Debug, Default)] + struct RecordingAuthManager(std::sync::Mutex>); + + #[async_trait] + impl AuthManager for RecordingAuthManager { + async fn init_session( + &self, + _client: &HttpClient, + _props: &HashMap, + ) -> Result> { + unreachable!("only table sessions are created") + } + + async fn catalog_session( + &self, + _client: &HttpClient, + _props: &HashMap, + ) -> Result> { + unreachable!("only table sessions are created") + } + + async fn table_session( + &self, + _client: &HttpClient, + _table: &TableIdent, + props: &HashMap, + parent: Arc, + ) -> Result> { + let mut keys = props.keys().cloned().collect::>(); + keys.sort(); + *self.0.lock().unwrap() = keys; + Ok(parent) + } + } + + let manager = RecordingAuthManager::default(); + let table_config = HashMap::from([ + ("token".to_string(), "table-token".to_string()), + ("custom.auth".to_string(), "value".to_string()), + ]); + build_vended_credential_provider( + &test_client("http://cat"), + &manager, + test_factory("http://cat", Some(&table_config)), + &aws_refresh_props("/v1/creds"), + &table_config, + false, + ) + .await + .unwrap() + .unwrap(); + + assert_eq!(*manager.0.lock().unwrap(), vec![ + "custom.auth".to_string(), + "token".to_string() + ]); + } + + #[tokio::test] + async fn injected_auth_manager_decides_on_the_table_token() { + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/credentials") + .match_header("authorization", Matcher::Missing) + .expect(1) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(s3_response( + "s3://bucket", + "AK", + SystemTime::now() + Duration::from_secs(3600), + )) + .create_async() + .await; + let provider = build_vended_credential_provider( + &test_client(&server.url()), + &NoopAuthManager, + RestVendedCredentialProviderFactory::new( + server.url(), + test_table(), + AUTH_TYPE_NONE, + true, + &HashMap::from([("token".to_string(), "table-token".to_string())]), + ), + &aws_refresh_props("/v1/credentials"), + &HashMap::from([("token".to_string(), "table-token".to_string())]), + false, + ) + .await + .unwrap() + .unwrap(); + + provider.load_credential("s3://bucket/x/f").await.unwrap(); + mock.assert_async().await; + } + + #[test] + fn factory_keeps_only_table_auth_config() { + let factory = RestVendedCredentialProviderFactory::new( + "http://cat", + test_table(), + AUTH_TYPE_OAUTH2, + true, + &HashMap::from([ + ("token".to_string(), "table-token".to_string()), + (S3_SECRET_ACCESS_KEY.to_string(), "SECRET".to_string()), + ]), + ); + assert_eq!( + factory.table_auth, + HashMap::from([("token".to_string(), "table-token".to_string())]) + ); + } + #[tokio::test] async fn provider_with_injected_auth_manager_is_not_serializable() { let provider = build_vended_credential_provider( @@ -2423,6 +3648,7 @@ mod tests { &NoopAuthManager, test_factory("http://cat", None), &aws_refresh_props("/v1/creds"), + &HashMap::new(), false, ) .await @@ -2431,10 +3657,7 @@ mod tests { let error = provider.factory().unwrap_err(); assert_eq!(error.kind(), ErrorKind::FeatureUnsupported); - assert!( - error.message().contains("without_credential_provider"), - "{error}" - ); + assert!(error.message().contains("injected AuthManager"), "{error}"); } #[test] diff --git a/crates/iceberg/public-api.txt b/crates/iceberg/public-api.txt index ae3c2c4741..dbafce4580 100644 --- a/crates/iceberg/public-api.txt +++ b/crates/iceberg/public-api.txt @@ -695,14 +695,6 @@ pub fn iceberg::inspect::SnapshotsTable<'a>::new(table: &'a iceberg::table::Tabl pub async fn iceberg::inspect::SnapshotsTable<'a>::scan(&self) -> iceberg::Result pub fn iceberg::inspect::SnapshotsTable<'a>::schema(&self) -> iceberg::spec::Schema pub mod iceberg::io -#[non_exhaustive] pub enum iceberg::io::StorageCredentialKind -pub iceberg::io::StorageCredentialKind::Azdls(iceberg::io::AzdlsCredential) -pub iceberg::io::StorageCredentialKind::Gcs(iceberg::io::GcsCredential) -pub iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential) -impl core::clone::Clone for iceberg::io::StorageCredentialKind -pub fn iceberg::io::StorageCredentialKind::clone(&self) -> iceberg::io::StorageCredentialKind -impl core::fmt::Debug for iceberg::io::StorageCredentialKind -pub fn iceberg::io::StorageCredentialKind::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::AzdlsConfig pub iceberg::io::AzdlsConfig::account_key: core::option::Option pub iceberg::io::AzdlsConfig::account_name: core::option::Option @@ -733,15 +725,6 @@ impl serde_core::ser::Serialize for iceberg::io::AzdlsConfig pub fn iceberg::io::AzdlsConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::AzdlsConfig pub fn iceberg::io::AzdlsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> -pub struct iceberg::io::AzdlsCredential -impl iceberg::io::AzdlsCredential -pub fn iceberg::io::AzdlsCredential::into_sas_token(self) -> alloc::string::String -pub fn iceberg::io::AzdlsCredential::new(sas_token: impl core::convert::Into) -> Self -pub fn iceberg::io::AzdlsCredential::sas_token(&self) -> &str -impl core::clone::Clone for iceberg::io::AzdlsCredential -pub fn iceberg::io::AzdlsCredential::clone(&self) -> iceberg::io::AzdlsCredential -impl core::fmt::Debug for iceberg::io::AzdlsCredential -pub fn iceberg::io::AzdlsCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::FileIO impl iceberg::io::FileIO pub fn iceberg::io::FileIO::config(&self) -> &iceberg::io::StorageConfig @@ -802,15 +785,6 @@ impl serde_core::ser::Serialize for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::GcsConfig pub fn iceberg::io::GcsConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> -pub struct iceberg::io::GcsCredential -impl iceberg::io::GcsCredential -pub fn iceberg::io::GcsCredential::into_token(self) -> alloc::string::String -pub fn iceberg::io::GcsCredential::new(token: impl core::convert::Into) -> Self -pub fn iceberg::io::GcsCredential::token(&self) -> &str -impl core::clone::Clone for iceberg::io::GcsCredential -pub fn iceberg::io::GcsCredential::clone(&self) -> iceberg::io::GcsCredential -impl core::fmt::Debug for iceberg::io::GcsCredential -pub fn iceberg::io::GcsCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::HfConfig pub iceberg::io::HfConfig::endpoint: core::option::Option pub iceberg::io::HfConfig::revision: core::option::Option @@ -878,7 +852,7 @@ impl core::fmt::Debug for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::StorageFactory for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> -pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> +pub fn iceberg::io::LocalFsStorageFactory::build_with_credential_provider(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl serde_core::ser::Serialize for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::LocalFsStorageFactory @@ -917,7 +891,7 @@ impl core::fmt::Debug for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> -pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> +pub fn iceberg::io::MemoryStorageFactory::build_with_credential_provider(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl serde_core::ser::Serialize for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::MemoryStorageFactory @@ -993,17 +967,6 @@ impl serde_core::ser::Serialize for iceberg::io::S3Config pub fn iceberg::io::S3Config::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::S3Config pub fn iceberg::io::S3Config::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> -pub struct iceberg::io::S3Credential -impl iceberg::io::S3Credential -pub fn iceberg::io::S3Credential::access_key_id(&self) -> &str -pub fn iceberg::io::S3Credential::into_parts(self) -> (alloc::string::String, alloc::string::String, core::option::Option) -pub fn iceberg::io::S3Credential::new(access_key_id: impl core::convert::Into, secret_access_key: impl core::convert::Into, session_token: core::option::Option) -> Self -pub fn iceberg::io::S3Credential::secret_access_key(&self) -> &str -pub fn iceberg::io::S3Credential::session_token(&self) -> core::option::Option<&str> -impl core::clone::Clone for iceberg::io::S3Credential -pub fn iceberg::io::S3Credential::clone(&self) -> iceberg::io::S3Credential -impl core::fmt::Debug for iceberg::io::S3Credential -pub fn iceberg::io::S3Credential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result pub struct iceberg::io::StorageConfig impl iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::from_props(props: std::collections::hash::map::HashMap) -> Self @@ -1043,18 +1006,18 @@ impl<'de> serde_core::de::Deserialize<'de> for iceberg::io::StorageConfig pub fn iceberg::io::StorageConfig::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg::io::StorageCredential impl iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::config(&self) -> &std::collections::hash::map::HashMap pub fn iceberg::io::StorageCredential::covers(&self, location: &str) -> bool -pub fn iceberg::io::StorageCredential::expires_at(&self) -> core::option::Option -pub fn iceberg::io::StorageCredential::into_kind(self) -> iceberg::io::StorageCredentialKind -pub fn iceberg::io::StorageCredential::kind(&self) -> &iceberg::io::StorageCredentialKind -pub fn iceberg::io::StorageCredential::new(kind: iceberg::io::StorageCredentialKind) -> Self -pub fn iceberg::io::StorageCredential::prefix(&self) -> core::option::Option<&str> -pub fn iceberg::io::StorageCredential::with_expiration(self, expires_at: std::time::SystemTime) -> Self -pub fn iceberg::io::StorageCredential::with_prefix(self, prefix: impl core::convert::Into) -> Self +pub fn iceberg::io::StorageCredential::new(prefix: impl core::convert::Into, config: std::collections::hash::map::HashMap) -> Self +pub fn iceberg::io::StorageCredential::prefix(&self) -> &str impl core::clone::Clone for iceberg::io::StorageCredential pub fn iceberg::io::StorageCredential::clone(&self) -> iceberg::io::StorageCredential +impl core::cmp::Eq for iceberg::io::StorageCredential +impl core::cmp::PartialEq for iceberg::io::StorageCredential +pub fn iceberg::io::StorageCredential::eq(&self, other: &iceberg::io::StorageCredential) -> bool impl core::fmt::Debug for iceberg::io::StorageCredential pub fn iceberg::io::StorageCredential::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result +impl core::marker::StructuralPartialEq for iceberg::io::StorageCredential pub const iceberg::io::ADLS_ACCOUNT_KEY: &str pub const iceberg::io::ADLS_ACCOUNT_NAME: &str pub const iceberg::io::ADLS_AUTHORITY_HOST: &str @@ -1155,19 +1118,18 @@ pub fn iceberg::io::MemoryStorage::writer<'life0, 'life1, 'async_trait>(&'life0 pub trait iceberg::io::StorageCredentialProvider: core::fmt::Debug + core::marker::Send + core::marker::Sync pub fn iceberg::io::StorageCredentialProvider::factory(&self) -> iceberg::Result> pub fn iceberg::io::StorageCredentialProvider::load_credential<'life0, 'life1, 'async_trait>(&'life0 self, path: &'life1 str) -> core::pin::Pin> + core::marker::Send + 'async_trait)>> where Self: 'async_trait, 'life0: 'async_trait, 'life1: 'async_trait -pub fn iceberg::io::StorageCredentialProvider::supports_path(&self, _path: &str) -> bool +pub fn iceberg::io::StorageCredentialProvider::supports_path(&self, path: &str) -> bool pub trait iceberg::io::StorageCredentialProviderFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize pub fn iceberg::io::StorageCredentialProviderFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> pub trait iceberg::io::StorageFactory: core::fmt::Debug + core::marker::Send + core::marker::Sync + typetag::Serialize + typetag::Deserialize pub fn iceberg::io::StorageFactory::build(&self, config: &iceberg::io::StorageConfig) -> iceberg::Result> -pub fn iceberg::io::StorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> +pub fn iceberg::io::StorageFactory::build_with_credential_provider(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl iceberg::io::StorageFactory for iceberg::io::LocalFsStorageFactory pub fn iceberg::io::LocalFsStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> -pub fn iceberg::io::LocalFsStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> +pub fn iceberg::io::LocalFsStorageFactory::build_with_credential_provider(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> impl iceberg::io::StorageFactory for iceberg::io::MemoryStorageFactory pub fn iceberg::io::MemoryStorageFactory::build(&self, _config: &iceberg::io::StorageConfig) -> iceberg::Result> -pub fn iceberg::io::MemoryStorageFactory::build_with_credentials(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> -pub fn iceberg::io::storage_prefix_covers(prefix: &str, location: &str) -> bool +pub fn iceberg::io::MemoryStorageFactory::build_with_credential_provider(&self, config: &iceberg::io::StorageConfig, credential_provider: core::option::Option>) -> iceberg::Result> pub mod iceberg::memory pub struct iceberg::memory::MemoryCatalog impl core::fmt::Debug for iceberg::memory::MemoryCatalog diff --git a/crates/iceberg/src/io/file_io.rs b/crates/iceberg/src/io/file_io.rs index 300ea71f07..e61ed8de57 100644 --- a/crates/iceberg/src/io/file_io.rs +++ b/crates/iceberg/src/io/file_io.rs @@ -77,25 +77,39 @@ mod _serde { use serde::{Deserialize, Serialize}; - use super::{StorageConfig, StorageCredentialProviderFactory, StorageFactory}; + use super::{StorageConfig, StorageFactory}; #[derive(Serialize)] pub(super) struct SerializableFileIO<'a> { pub(super) config: &'a StorageConfig, pub(super) factory: &'a Arc, #[serde(skip_serializing_if = "Option::is_none")] - pub(super) credential_provider: Option>, + pub(super) credential_provider: Option, } #[derive(Deserialize)] pub(super) struct DeserializedFileIO { pub(super) config: StorageConfig, pub(super) factory: Arc, + /// Kept as JSON, so that a provider factory the receiving binary does + /// not link can be dropped instead of failing the whole `FileIO`. #[serde(default)] - pub(super) credential_provider: Option>, + pub(super) credential_provider: Option, } } +/// Warns that a `FileIO` crossing a process boundary loses its credential +/// provider. Engines may serialize a `FileIO` per task, so each cause, tracked +/// by its own `warned`, is reported once. +fn warn_once_without_provider(warned: &'static std::sync::Once, reason: impl FnOnce() -> String) { + warned.call_once(|| { + tracing::warn!( + "{}; its credentials will not be refreshed after deserialization", + reason() + ) + }); +} + impl FileIO { /// Create a new FileIO backed by in-memory storage. /// @@ -140,15 +154,25 @@ impl FileIO { /// /// A credential provider is serialized as the /// [`StorageCredentialProviderFactory`] returned by - /// [`StorageCredentialProvider::factory`], and rebuilt on deserialization. Serialization fails - /// when the provider cannot be rebuilt in another process; use - /// [`FileIO::without_credential_provider`] to serialize without it. + /// [`StorageCredentialProvider::factory`], and rebuilt on deserialization. A provider that + /// cannot be rebuilt in another process is left out with a warning, as with + /// [`FileIO::without_credential_provider`]: the deserialized `FileIO` uses the credentials in + /// its configuration, which are not refreshed. Storage factories that hold their own + /// credential sources, such as a custom AWS credential loader, still fail to serialize, as + /// leaving them out would leave no credentials at all. pub fn serialize_all(&self) -> Result> { - let credential_provider = self - .credential_provider - .as_ref() - .map(|provider| provider.factory()) - .transpose()?; + let credential_provider = self.credential_provider.as_ref().and_then(|provider| { + provider + .factory() + .and_then(|factory| Ok(serde_json::to_value(factory)?)) + .inspect_err(|error| { + static WARNED: std::sync::Once = std::sync::Once::new(); + warn_once_without_provider(&WARNED, || { + format!("serializing FileIO without its credential provider: {error}") + }) + }) + .ok() + }); Ok(serde_json::to_vec(&_serde::SerializableFileIO { config: &self.config, @@ -162,15 +186,51 @@ impl FileIO { /// The receiving binary must use a compatible crate version and link the concrete factory /// implementation so it is registered with `typetag`. Backend-specific requirements are /// documented by each storage factory implementation. + /// + /// A credential provider is rebuilt when the binary links its factory implementation, such as + /// the REST catalog's. Otherwise, or when rebuilding fails, the `FileIO` is deserialized without + /// it and logs a warning, as [`FileIO::serialize_all`] does. pub fn deserialize_all(bytes: &[u8]) -> Result { let _serde::DeserializedFileIO { config, factory, credential_provider, } = serde_json::from_slice(bytes)?; - let credential_provider = credential_provider - .map(|provider_factory| provider_factory.build(&config)) - .transpose()?; + let credential_provider = credential_provider.and_then(|provider_factory| { + // The factory may hold credentials, and serde errors quote the + // offending value, so only its type is reported. + let factory_type = provider_factory + .get("type") + .and_then(serde_json::Value::as_str) + .unwrap_or("unknown") + .to_owned(); + match serde_json::from_value::>( + provider_factory, + ) { + Ok(provider_factory) => provider_factory + .build(&config) + .inspect_err(|error| { + static WARNED: std::sync::Once = std::sync::Once::new(); + warn_once_without_provider(&WARNED, || { + format!( + "deserializing FileIO without its credential provider, which \ + could not be rebuilt: {error}" + ) + }) + }) + .ok(), + Err(_) => { + static WARNED: std::sync::Once = std::sync::Once::new(); + warn_once_without_provider(&WARNED, || { + format!( + "deserializing FileIO without its credential provider, whose \ + factory {factory_type} is not linked or does not match this version" + ) + }); + None + } + } + }); Ok(Self { config, factory, @@ -182,7 +242,8 @@ impl FileIO { /// Returns a copy of this `FileIO` without its credential provider. /// /// The copy uses only the credentials in its storage configuration, which are not refreshed. - /// Use this to serialize a `FileIO` whose credential provider cannot be serialized. + /// Use this to serialize a `FileIO` without its credential provider, even when the provider + /// could be rebuilt in another process. pub fn without_credential_provider(&self) -> Self { Self { config: self.config.clone(), @@ -211,7 +272,7 @@ impl FileIO { // support refreshable credentials can wire it into their operators. let storage = self .factory - .build_with_credentials(&self.config, self.credential_provider.clone())?; + .build_with_credential_provider(&self.config, self.credential_provider.clone())?; // Try to set it (another thread might have set it first) let _ = self.storage.set(storage.clone()); @@ -516,18 +577,21 @@ mod tests { use tempfile::TempDir; use super::{FileIO, FileIOBuilder}; + use crate::Result; use crate::io::{ - GcsCredential, LocalFsStorageFactory, MemoryStorageFactory, StorageConfig, - StorageCredential, StorageCredentialKind, StorageCredentialProvider, - StorageCredentialProviderFactory, + GCS_TOKEN, LocalFsStorageFactory, MemoryStorageFactory, StorageConfig, StorageCredential, + StorageCredentialProvider, StorageCredentialProviderFactory, }; - use crate::{ErrorKind, Result}; #[derive(Debug)] struct TestCredentialProvider; #[async_trait::async_trait] impl StorageCredentialProvider for TestCredentialProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, _path: &str) -> Result { unreachable!("unsupported factories must ignore the provider") } @@ -541,10 +605,18 @@ mod tests { #[async_trait::async_trait] impl StorageCredentialProvider for PortableCredentialProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, _path: &str) -> Result { - Ok(StorageCredential::new(StorageCredentialKind::Gcs( - GcsCredential::new(self.endpoint.clone().unwrap_or_default()), - ))) + Ok(StorageCredential::new( + "gs", + std::collections::HashMap::from([( + GCS_TOKEN.to_string(), + self.endpoint.clone().unwrap_or_default(), + )]), + )) } fn factory(&self) -> Result> { @@ -726,24 +798,95 @@ mod tests { } #[test] - fn test_file_io_with_credential_provider_serialization_fails() { + fn test_file_io_drops_unserializable_credential_provider() { let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_prop("s3.access-key-id", "static-key") .with_credential_provider(Arc::new(TestCredentialProvider)) .build(); - let err = file_io.serialize_all().unwrap_err(); - assert_eq!(err.kind(), ErrorKind::FeatureUnsupported, "{err}"); - - let deserialized = FileIO::deserialize_all( - &file_io + let deserialized = FileIO::deserialize_all(&file_io.serialize_all().unwrap()).unwrap(); + assert!(deserialized.credential_provider.is_none()); + assert_eq!( + deserialized.config().get("s3.access-key-id"), + Some(&"static-key".to_string()) + ); + assert_eq!( + file_io.serialize_all().unwrap(), + file_io .without_credential_provider() .serialize_all() - .unwrap(), - ) - .unwrap(); + .unwrap() + ); + } + + /// A provider whose factory cannot be serialized. + #[derive(Debug)] + struct UnserializableFactoryProvider; + + #[async_trait::async_trait] + impl StorageCredentialProvider for UnserializableFactoryProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + + async fn load_credential(&self, _path: &str) -> Result { + unreachable!("serialization tests never load credentials") + } + + fn factory(&self) -> Result> { + Ok(Arc::new(UnserializableFactory)) + } + } + + #[derive(Debug, serde::Deserialize)] + struct UnserializableFactory; + + impl serde::Serialize for UnserializableFactory { + fn serialize( + &self, + _serializer: S, + ) -> std::result::Result { + Err(serde::ser::Error::custom("cannot serialize")) + } + } + + #[typetag::serde] + impl StorageCredentialProviderFactory for UnserializableFactory { + fn build(&self, _config: &StorageConfig) -> Result> { + unreachable!("never serialized") + } + } + + #[test] + fn test_file_io_drops_provider_whose_factory_fails_to_serialize() { + let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_credential_provider(Arc::new(UnserializableFactoryProvider)) + .build(); + + let deserialized = FileIO::deserialize_all(&file_io.serialize_all().unwrap()).unwrap(); assert!(deserialized.credential_provider.is_none()); } + #[test] + fn test_file_io_deserializes_without_unknown_credential_provider() { + let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) + .with_prop("s3.access-key-id", "static-key") + .with_credential_provider(Arc::new(PortableCredentialProvider { endpoint: None })) + .build(); + // A receiving binary that does not link the provider's factory. + let mut serialized: serde_json::Value = + serde_json::from_slice(&file_io.serialize_all().unwrap()).unwrap(); + serialized["credential_provider"]["type"] = "UnlinkedProviderFactory".into(); + + let deserialized = + FileIO::deserialize_all(&serde_json::to_vec(&serialized).unwrap()).unwrap(); + assert!(deserialized.credential_provider.is_none()); + assert_eq!( + deserialized.config().get("s3.access-key-id"), + Some(&"static-key".to_string()) + ); + } + #[tokio::test] async fn test_file_io_rebuilds_credential_provider_after_serialization() { let file_io = FileIOBuilder::new(Arc::new(MemoryStorageFactory)) @@ -759,12 +902,10 @@ mod tests { .await .unwrap(); // The provider is rebuilt from the deserialized configuration. - match credential.kind() { - StorageCredentialKind::Gcs(gcs) => { - assert_eq!(gcs.token(), "https://catalog/credentials") - } - other => panic!("expected GCS credential, got {other:?}"), - } + assert_eq!( + credential.config().get(GCS_TOKEN).map(String::as_str), + Some("https://catalog/credentials") + ); } #[tokio::test] diff --git a/crates/iceberg/src/io/storage/config/gcs.rs b/crates/iceberg/src/io/storage/config/gcs.rs index 6d06ab6a34..4069b30894 100644 --- a/crates/iceberg/src/io/storage/config/gcs.rs +++ b/crates/iceberg/src/io/storage/config/gcs.rs @@ -111,7 +111,9 @@ impl TryFrom<&StorageConfig> for GcsConfig { } // GCS_NO_AUTH enables all anonymous/no-auth options - if props.get(GCS_NO_AUTH).is_some() { + if let Some(no_auth) = props.get(GCS_NO_AUTH) + && is_truthy(no_auth) + { cfg.allow_anonymous = true; cfg.disable_vm_metadata = true; cfg.disable_config_load = true; @@ -182,6 +184,17 @@ mod tests { assert!(gcs_config.disable_config_load); } + #[test] + fn test_gcs_config_no_auth_false() { + let storage_config = StorageConfig::new().with_prop(GCS_NO_AUTH, "false"); + + let gcs_config = GcsConfig::try_from(&storage_config).unwrap(); + + assert!(!gcs_config.allow_anonymous); + assert!(!gcs_config.disable_vm_metadata); + assert!(!gcs_config.disable_config_load); + } + #[test] fn test_gcs_config_allow_anonymous() { let storage_config = StorageConfig::new().with_prop(GCS_ALLOW_ANONYMOUS, "true"); diff --git a/crates/iceberg/src/io/storage/mod.rs b/crates/iceberg/src/io/storage/mod.rs index a193d7af44..7de00854ac 100644 --- a/crates/iceberg/src/io/storage/mod.rs +++ b/crates/iceberg/src/io/storage/mod.rs @@ -21,9 +21,9 @@ mod config; mod local_fs; mod memory; +use std::collections::HashMap; use std::fmt::Debug; use std::sync::Arc; -use std::time::SystemTime; use async_trait::async_trait; use bytes::Bytes; @@ -147,7 +147,7 @@ pub trait StorageFactory: Debug + Send + Sync { /// Backends that cannot use the provider ignore it and use the credentials /// in `config`, as they would without one. The default does exactly that. #[allow(unused_variables)] - fn build_with_credentials( + fn build_with_credential_provider( &self, config: &StorageConfig, credential_provider: Option>, @@ -162,6 +162,12 @@ pub trait StorageFactory: Debug + Send + Sync { /// storage backends can re-fetch credentials as they approach expiry instead /// of failing once the initial token's TTL runs out. /// +/// # Debug output +/// +/// [`FileIO`](crate::io::FileIO) and storage implementations print the +/// provider in their `Debug` output, so the provider's `Debug` implementation +/// must not expose credentials or other secrets. +/// /// # Caching /// /// [`load_credential`](Self::load_credential) may be called very frequently. @@ -170,29 +176,34 @@ pub trait StorageFactory: Debug + Send + Sync { /// trigger a call back to the catalog. #[async_trait] pub trait StorageCredentialProvider: Debug + Send + Sync { - /// Return whether this provider has refresh configuration for `path`. + /// Return whether this provider supplies credentials for `path`. /// - /// Backends use this before replacing their normal credential chain. The - /// default is `true` for single-backend providers; multi-backend providers - /// should return `false` for schemes they do not configure. - fn supports_path(&self, _path: &str) -> bool { - true - } + /// Backends replace their own credential chain with the provider only for + /// paths it supports, so a provider must return `false` for the storage it + /// does not serve, such as other schemes behind a resolving storage. + fn supports_path(&self, path: &str) -> bool; /// Load a fresh credential for the storage location identified by `path`. /// - /// `path` is the absolute location being accessed (e.g. - /// `s3://bucket/warehouse/db/table/...`). Providers that vend distinct + /// `path` is an absolute location: usually the file being accessed, such + /// as `s3://bucket/warehouse/db/table/data/file.parquet`, but for a bulk + /// delete the location shared by a batch, which may be a storage root + /// (`s3://bucket/`) or the prefix of a previously returned credential + /// (`s3://bucket/warehouse/db/table`). Providers that vend distinct /// credentials per location prefix use it to select the most specific - /// match. When the selected credential has a declared - /// [`StorageCredential::prefix`], it must [cover](StorageCredential::covers) `path`. + /// match, which must [cover](StorageCredential::covers) `path`. + /// + /// Backends read the credential's expiry from its config, such as + /// `s3.session-token-expires-at-ms`, and load a new one before it expires. + /// A credential without an expiry is used for as long as the backend lives. async fn load_credential(&self, path: &str) -> Result; /// Return a factory that rebuilds an equivalent provider in another process. /// /// [`FileIO::serialize_all`](crate::io::FileIO::serialize_all) serializes this - /// factory in place of the provider. The default reports that the provider - /// cannot be serialized. + /// factory in place of the provider. On an error, which the default always + /// returns, it serializes the `FileIO` without the provider and logs a + /// warning. fn factory(&self) -> Result> { Err(Error::new( ErrorKind::FeatureUnsupported, @@ -213,238 +224,61 @@ pub trait StorageCredentialProviderFactory: Debug + Send + Sync { fn build(&self, config: &StorageConfig) -> Result>; } -/// A vended storage credential together with its scope and expiry. -#[derive(Clone, Debug)] +/// A vended storage credential: the storage properties that apply to +/// locations under a prefix, like Java's `StorageCredential` and the REST +/// catalog's `storage-credentials`. +/// +/// `config` holds backend storage properties, such as `s3.access-key-id`, +/// `s3.secret-access-key`, `s3.session-token` and +/// `s3.session-token-expires-at-ms`. Backends read their credential and its +/// expiry from it. +#[derive(Clone, PartialEq, Eq)] pub struct StorageCredential { - /// Storage-location prefix this credential is scoped to. `None` represents a - /// credential without a declared scope, sourced from flat storage properties. - prefix: Option, - /// The backend-specific credential material. - kind: StorageCredentialKind, - /// When the credential expires, if known. `None` means non-expiring and - /// backends treat such a credential as always valid and never refresh it. - expires_at: Option, + /// Storage-location prefix this credential applies to, such as + /// `s3://bucket/table` or, for a whole scheme, `s3`. + prefix: String, + /// Backend storage properties holding the credential. + config: HashMap, } impl StorageCredential { - /// Create a storage credential with no declared scope or expiration. - pub fn new(kind: StorageCredentialKind) -> Self { + /// Create a credential for locations under `prefix`. + /// + /// An empty prefix covers no location, as Java rejects it. + pub fn new(prefix: impl Into, config: HashMap) -> Self { Self { - prefix: None, - kind, - expires_at: None, + prefix: prefix.into(), + config, } } - /// Set the storage-location prefix this credential is scoped to. - pub fn with_prefix(mut self, prefix: impl Into) -> Self { - self.prefix = Some(prefix.into()); - self - } - - /// Set when this credential expires. - pub fn with_expiration(mut self, expires_at: SystemTime) -> Self { - self.expires_at = Some(expires_at); - self + /// Return the storage-location prefix this credential applies to. + pub fn prefix(&self) -> &str { + &self.prefix } - /// Return the storage-location prefix this credential is scoped to. - pub fn prefix(&self) -> Option<&str> { - self.prefix.as_deref() + /// Return the storage properties holding the credential. + pub fn config(&self) -> &HashMap { + &self.config } /// Return whether this credential applies to `location`. /// - /// A credential without a prefix covers every location. Otherwise the - /// prefix must match whole path segments of `location`, and scheme - /// aliases (`s3a`/`s3n` for `s3`, `gcs` for `gs`, and the plain-text - /// Azure schemes for their TLS variants) are treated as equal. A prefix - /// that is only a scheme, such as `s3`, covers every location with that - /// scheme. + /// As in Java, `location` must start with the non-empty prefix, compared + /// as plain strings: `s3` covers every `s3://`, `s3a://` and `s3n://` + /// location, and `s3://bucket/tab` covers `s3://bucket/table`. pub fn covers(&self, location: &str) -> bool { - self.prefix - .as_deref() - .is_none_or(|prefix| storage_prefix_covers(prefix, location)) - } - - /// Return the backend-specific credential material. - pub fn kind(&self) -> &StorageCredentialKind { - &self.kind - } - - /// Consume this credential and return its backend-specific material. - pub fn into_kind(self) -> StorageCredentialKind { - self.kind - } - - /// Return when this credential expires. - pub fn expires_at(&self) -> Option { - self.expires_at - } -} - -/// Return whether the storage-location `prefix` covers `location`, with the -/// matching rules of [`StorageCredential::covers`]. -pub fn storage_prefix_covers(prefix: &str, location: &str) -> bool { - let Some((location_scheme, location_rest)) = location.split_once("://") else { - return false; - }; - let Some((prefix_scheme, prefix_rest)) = prefix.split_once("://") else { - return !prefix.is_empty() && canonical_scheme(prefix) == canonical_scheme(location_scheme); - }; - - canonical_scheme(prefix_scheme) == canonical_scheme(location_scheme) - && location_rest - .strip_prefix(prefix_rest) - .is_some_and(|remainder| { - prefix_rest.is_empty() - || prefix_rest.ends_with('/') - || remainder.is_empty() - || remainder.starts_with('/') - }) -} - -fn canonical_scheme(scheme: &str) -> String { - let scheme = scheme.to_ascii_lowercase(); - match scheme.as_str() { - "s3a" | "s3n" => "s3".to_string(), - "gcs" => "gs".to_string(), - "abfs" => "abfss".to_string(), - "wasb" => "wasbs".to_string(), - _ => scheme, - } -} - -/// Backend-specific credential material. -#[derive(Clone, Debug)] -#[non_exhaustive] -pub enum StorageCredentialKind { - /// Amazon S3 credentials. - S3(S3Credential), - /// Google Cloud Storage credentials. - Gcs(GcsCredential), - /// Azure Data Lake Storage credentials. - Azdls(AzdlsCredential), -} - -/// Temporary Azure Data Lake Storage credentials (a shared access signature). -#[derive(Clone)] -pub struct AzdlsCredential { - /// Shared access signature used to access Azure storage. - sas_token: String, -} - -impl AzdlsCredential { - /// Create an Azure Data Lake Storage credential. - pub fn new(sas_token: impl Into) -> Self { - Self { - sas_token: sas_token.into(), - } - } - - /// Return the Azure shared access signature. - pub fn sas_token(&self) -> &str { - &self.sas_token - } - - /// Consume this credential and return its shared access signature. - pub fn into_sas_token(self) -> String { - self.sas_token - } -} - -impl Debug for AzdlsCredential { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("AzdlsCredential").finish_non_exhaustive() - } -} - -/// Temporary Amazon S3 credentials. -#[derive(Clone)] -pub struct S3Credential { - /// AWS access key ID. - access_key_id: String, - /// AWS secret access key. - secret_access_key: String, - /// AWS session token, set for temporary (STS/vended) credentials. - session_token: Option, -} - -impl S3Credential { - /// Create temporary Amazon S3 credentials. - pub fn new( - access_key_id: impl Into, - secret_access_key: impl Into, - session_token: Option, - ) -> Self { - Self { - access_key_id: access_key_id.into(), - secret_access_key: secret_access_key.into(), - session_token, - } - } - - /// Return the AWS access key ID. - pub fn access_key_id(&self) -> &str { - &self.access_key_id - } - - /// Return the AWS secret access key. - pub fn secret_access_key(&self) -> &str { - &self.secret_access_key - } - - /// Return the AWS session token, if present. - pub fn session_token(&self) -> Option<&str> { - self.session_token.as_deref() - } - - /// Consume these credentials and return their component values. - pub fn into_parts(self) -> (String, String, Option) { - let Self { - access_key_id, - secret_access_key, - session_token, - } = self; - (access_key_id, secret_access_key, session_token) - } -} - -impl Debug for S3Credential { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("S3Credential").finish_non_exhaustive() - } -} - -/// Temporary Google Cloud Storage credentials (an OAuth2 access token). -#[derive(Clone)] -pub struct GcsCredential { - /// OAuth2 bearer token used to access GCS. - token: String, -} - -impl GcsCredential { - /// Create a Google Cloud Storage credential. - pub fn new(token: impl Into) -> Self { - Self { - token: token.into(), - } - } - - /// Return the OAuth2 bearer token used to access GCS. - pub fn token(&self) -> &str { - &self.token - } - - /// Consume this credential and return its OAuth2 bearer token. - pub fn into_token(self) -> String { - self.token + !self.prefix.is_empty() && location.starts_with(&self.prefix) } } -impl Debug for GcsCredential { +impl Debug for StorageCredential { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - f.debug_struct("GcsCredential").finish_non_exhaustive() + // Config values are secrets. + f.debug_struct("StorageCredential") + .field("prefix", &self.prefix) + .field("config_keys", &self.config.keys().collect::>()) + .finish() } } @@ -453,45 +287,32 @@ mod tests { use super::*; fn scoped(prefix: &str) -> StorageCredential { - StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))) - .with_prefix(prefix) + StorageCredential::new( + prefix, + HashMap::from([("k".to_string(), "secret".to_string())]), + ) } #[test] - fn credential_prefix_matches_whole_segments() { + fn credential_prefix_matches_like_java() { let credential = scoped("s3://bucket/table"); assert!(credential.covers("s3://bucket/table")); assert!(credential.covers("s3://bucket/table/data/file.parquet")); - assert!(!credential.covers("s3://bucket/table2/data/file.parquet")); + assert!(credential.covers("s3://bucket/table2/file.parquet")); assert!(!credential.covers("s3://bucket/tab")); - assert!(!credential.covers("s3://other/table/file.parquet")); + assert!(!credential.covers("s3a://bucket/table/file.parquet")); - assert!(scoped("s3://bucket/table/").covers("s3://bucket/table/file.parquet")); - assert!(scoped("s3://").covers("s3://any/file.parquet")); - } - - #[test] - fn credential_prefix_treats_scheme_aliases_as_equal() { - let credential = scoped("s3://bucket/table"); - assert!(credential.covers("s3a://bucket/table/file.parquet")); - assert!(credential.covers("S3N://bucket/table/file.parquet")); - assert!(scoped("gcs://bucket").covers("gs://bucket/file.parquet")); - assert!( - scoped("abfss://fs@account.dfs.core.windows.net/table") - .covers("abfs://fs@account.dfs.core.windows.net/table/file.parquet") - ); - assert!(!credential.covers("gs://bucket/table/file.parquet")); - } - - #[test] - fn credential_scheme_prefix_covers_the_whole_scheme() { + // A scheme prefix covers every location with that scheme, including + // `s3a` and `s3n`. assert!(scoped("s3").covers("s3a://bucket/file.parquet")); assert!(!scoped("s3").covers("gs://bucket/file.parquet")); assert!(!scoped("").covers("s3://bucket/file.parquet")); - assert!(!scoped("s3").covers("not-a-url")); - assert!( - StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))) - .covers("gs://bucket/file.parquet") - ); + } + + #[test] + fn credential_debug_omits_config_values() { + let debug = format!("{:?}", scoped("s3://bucket")); + assert!(debug.contains("s3://bucket"), "{debug}"); + assert!(!debug.contains("secret"), "{debug}"); } } diff --git a/crates/storage/opendal/Cargo.toml b/crates/storage/opendal/Cargo.toml index a837206a4e..e2e1b15231 100644 --- a/crates/storage/opendal/Cargo.toml +++ b/crates/storage/opendal/Cargo.toml @@ -59,18 +59,21 @@ cfg-if = { workspace = true } futures = { workspace = true } iceberg = { workspace = true } opendal = { workspace = true } -reqsign-aws-v4 = { version = "3.0.0", optional = true } +reqsign-aws-v4 = { version = "3.3.0", optional = true } reqsign-azure-storage = { version = "3.2.1", optional = true } -reqsign-core = { version = "3.0.0", optional = true } -reqsign-google = { version = "3.0.0", optional = true } +reqsign-core = { version = "3.3.1", optional = true } +reqsign-google = { version = "3.0.4", optional = true } serde = { workspace = true } +tracing = { workspace = true } typetag = { workspace = true } url = { workspace = true } [dev-dependencies] async-trait = { workspace = true } iceberg_test_utils = { path = "../../test_utils", features = ["tests"] } +mockito = { workspace = true } reqwest = { workspace = true } +serde_json = { workspace = true } tempfile = { workspace = true } tokio = { workspace = true, features = ["macros"] } diff --git a/crates/storage/opendal/public-api.txt b/crates/storage/opendal/public-api.txt index ad6c4522b8..9a52857b74 100644 --- a/crates/storage/opendal/public-api.txt +++ b/crates/storage/opendal/public-api.txt @@ -5,7 +5,7 @@ pub use iceberg_storage_opendal::ProvideCredential #[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Azdls pub iceberg_storage_opendal::OpenDalStorage::Azdls::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Azdls::credential_provider: core::option::Option> -pub iceberg_storage_opendal::OpenDalStorage::Azdls::sas_tokens: alloc::sync::Arc +pub iceberg_storage_opendal::OpenDalStorage::Azdls::sas_tokens: alloc::sync::Arc #[non_exhaustive] pub iceberg_storage_opendal::OpenDalStorage::Gcs pub iceberg_storage_opendal::OpenDalStorage::Gcs::config: alloc::sync::Arc pub iceberg_storage_opendal::OpenDalStorage::Gcs::credential_provider: core::option::Option> @@ -54,22 +54,11 @@ impl core::fmt::Debug for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::storage::StorageFactory for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::build(&self, config: &iceberg::io::storage::config::StorageConfig) -> iceberg::error::Result> -pub fn iceberg_storage_opendal::OpenDalStorageFactory::build_with_credentials(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> +pub fn iceberg_storage_opendal::OpenDalStorageFactory::build_with_credential_provider(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalStorageFactory pub fn iceberg_storage_opendal::OpenDalStorageFactory::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> -pub struct iceberg_storage_opendal::AzdlsSasTokens(_) -impl core::clone::Clone for iceberg_storage_opendal::AzdlsSasTokens -pub fn iceberg_storage_opendal::AzdlsSasTokens::clone(&self) -> iceberg_storage_opendal::AzdlsSasTokens -impl core::default::Default for iceberg_storage_opendal::AzdlsSasTokens -pub fn iceberg_storage_opendal::AzdlsSasTokens::default() -> iceberg_storage_opendal::AzdlsSasTokens -impl core::fmt::Debug for iceberg_storage_opendal::AzdlsSasTokens -pub fn iceberg_storage_opendal::AzdlsSasTokens::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result -impl serde_core::ser::Serialize for iceberg_storage_opendal::AzdlsSasTokens -pub fn iceberg_storage_opendal::AzdlsSasTokens::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer -impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::AzdlsSasTokens -pub fn iceberg_storage_opendal::AzdlsSasTokens::deserialize<__D>(__deserializer: __D) -> core::result::Result::Error> where __D: serde_core::de::Deserializer<'de> pub struct iceberg_storage_opendal::CustomAwsCredentialLoader(_) impl iceberg_storage_opendal::CustomAwsCredentialLoader pub fn iceberg_storage_opendal::CustomAwsCredentialLoader::new(provider: impl reqsign_core::api::ProvideCredential + 'static) -> Self @@ -108,7 +97,7 @@ impl core::fmt::Debug for iceberg_storage_opendal::OpenDalResolvingStorageFactor pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result impl iceberg::io::storage::StorageFactory for iceberg_storage_opendal::OpenDalResolvingStorageFactory pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build(&self, config: &iceberg::io::storage::config::StorageConfig) -> iceberg::error::Result> -pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build_with_credentials(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> +pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::build_with_credential_provider(&self, config: &iceberg::io::storage::config::StorageConfig, credential_provider: core::option::Option>) -> iceberg::error::Result> impl serde_core::ser::Serialize for iceberg_storage_opendal::OpenDalResolvingStorageFactory pub fn iceberg_storage_opendal::OpenDalResolvingStorageFactory::serialize<__S>(&self, __serializer: __S) -> core::result::Result<<__S as serde_core::ser::Serializer>::Ok, <__S as serde_core::ser::Serializer>::Error> where __S: serde_core::ser::Serializer impl<'de> serde_core::de::Deserialize<'de> for iceberg_storage_opendal::OpenDalResolvingStorageFactory diff --git a/crates/storage/opendal/src/azdls.rs b/crates/storage/opendal/src/azdls.rs index 387915d54c..78acd5bfb4 100644 --- a/crates/storage/opendal/src/azdls.rs +++ b/crates/storage/opendal/src/azdls.rs @@ -22,8 +22,8 @@ use std::sync::Arc; use iceberg::io::{ ADLS_ACCOUNT_KEY, ADLS_ACCOUNT_NAME, ADLS_AUTHORITY_HOST, ADLS_CLIENT_ID, ADLS_CLIENT_SECRET, - ADLS_CONNECTION_STRING, ADLS_SAS_TOKEN, ADLS_SAS_TOKEN_PREFIX, ADLS_TENANT_ID, - StorageCredentialKind, StorageCredentialProvider, + ADLS_CONNECTION_STRING, ADLS_SAS_TOKEN, ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX, + ADLS_SAS_TOKEN_PREFIX, ADLS_TENANT_ID, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; use opendal::Configurator; @@ -36,7 +36,9 @@ use reqsign_core::{ use serde::{Deserialize, Serialize}; use url::Url; -use crate::utils::{VendedCredentialSource, from_opendal_error}; +use crate::utils::{ + VendedCredentialSource, credential_expiry, from_opendal_error, required_credential_property, +}; /// Local version of `ensure_data_valid` macro since the iceberg crate's macro /// uses `$crate::error::Error` paths that don't resolve from external crates @@ -93,47 +95,54 @@ pub(crate) fn azdls_config_parse(mut properties: HashMap) -> Res /// Account-specific ADLS SAS tokens supplied through Java-compatible storage /// properties. +/// +/// Public only because [`OpenDalStorage::Azdls`](crate::OpenDalStorage::Azdls) +/// holds it; it is built from the storage configuration. +#[doc(hidden)] #[derive(Clone, Default, Serialize, Deserialize)] pub struct AzdlsSasTokens(HashMap); impl AzdlsSasTokens { - /// Collect `adls.sas-token.` properties by storage account. Like - /// Java, keys may name the host (`account.dfs.core.windows.net`) or only - /// the account. + /// Collect `adls.sas-token.` properties by key suffix. pub(crate) fn from_properties(properties: &HashMap) -> Self { - let mut tokens = properties - .iter() - .filter_map(|(key, value)| { - let account = sas_token_account(key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?); - (!account.is_empty() && !value.is_empty()).then_some((key, account, value)) - }) - .collect::>(); - // Deterministic choice when several keys name the same account: the - // last, host-keyed one wins. - tokens.sort(); Self( - tokens - .into_iter() - .map(|(_, account, value)| (account.to_string(), value.clone())) + properties + .iter() + .filter_map(|(key, value)| { + let suffix = key.strip_prefix(ADLS_SAS_TOKEN_PREFIX)?; + (!suffix.is_empty() && !value.is_empty()) + .then(|| (suffix.to_string(), value.clone())) + }) .collect(), ) } + pub(crate) fn is_empty(&self) -> bool { + self.0.is_empty() + } + + /// The token for `path`, selected as by [`sas_token_suffixes`]. fn for_path(&self, path: &AzureStoragePath) -> Option<&str> { - self.0.get(&path.account_name).map(String::as_str) + sas_token_suffixes(path) + .iter() + .find_map(|suffix| self.0.get(suffix)) + .map(String::as_str) } } -/// The storage account named by the suffix of an account-specific SAS token -/// property: a host such as `account.dfs.core.windows.net`, or the account. -fn sas_token_account(suffix: &str) -> &str { - suffix.split('.').next().unwrap_or(suffix) +/// Suffixes of the `adls.sas-token.` properties that may hold the SAS +/// token for `path`, most specific first: the exact host, as in Java, then the +/// account alone, as sent for older Java versions and PyIceberg. A key that +/// names only the account matches every host of an account with that name, +/// including one in another cloud; a host-keyed token matches only its host. +fn sas_token_suffixes(path: &AzureStoragePath) -> [String; 2] { + [path.host(), path.account_name.clone()] } impl std::fmt::Debug for AzdlsSasTokens { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("AzdlsSasTokens") - .field("account_count", &self.0.len()) + .field("token_count", &self.0.len()) .finish_non_exhaustive() } } @@ -186,15 +195,6 @@ pub enum AzureStorageScheme { } impl AzureStorageScheme { - /// The HTTP scheme of an endpoint derived from a path. - /// - /// Iceberg Java accepts the non-secure aliases for compatibility but still - /// connects over TLS. SAS tokens are query parameters and must not be sent - /// over a plaintext connection unless the user configures such an endpoint. - pub fn as_http_scheme(&self) -> &str { - "https" - } - /// Whether an explicitly configured endpoint must use TLS. fn requires_tls(&self) -> bool { matches!(self, AzureStorageScheme::Abfss | AzureStorageScheme::Wasbs) @@ -284,12 +284,13 @@ fn azdls_config_build( .as_ref() .filter(|provider| provider.supports_path(absolute_path)); if let Some(provider) = credential_provider { - let chain = ProvideCredentialChain::new().push(VendedAzdlsCredentialProvider( - VendedCredentialSource::new( + let chain = ProvideCredentialChain::new().push(VendedAzdlsCredentialProvider { + source: VendedCredentialSource::new( Arc::clone(provider), credential_location.unwrap_or(absolute_path).to_string(), ), - )); + sas_token_suffixes: sas_token_suffixes(path), + }); builder = builder.credential_provider_chain(chain); } else if let Some(sas_token) = sas_tokens.for_path(path) { builder = builder.sas_token(sas_token); @@ -301,26 +302,51 @@ fn azdls_config_build( /// Adapts a generic [`StorageCredentialProvider`] into a reqsign /// [`ProvideCredential`] for expiring Azure SAS tokens. #[derive(Debug)] -struct VendedAzdlsCredentialProvider(VendedCredentialSource); +struct VendedAzdlsCredentialProvider { + source: VendedCredentialSource, + /// Suffixes of the SAS token properties for the operator's path. + sas_token_suffixes: [String; 2], +} impl ProvideCredential for VendedAzdlsCredentialProvider { type Credential = AzureCredential; async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { - match self.0.load("ADLS").await? { - (StorageCredentialKind::Azdls(azdls), expires_at) => { - let sas_token = azdls.into_sas_token(); - Ok(Some(match expires_at { - Some(expires_at) => { - AzureCredential::with_sas_token_expires_at(&sas_token, expires_at) - } - None => AzureCredential::with_sas_token(&sas_token), - })) - } - _ => Err(ReqsignError::unexpected( - "ADLS storage received a non-ADLS credential from the provider", - )), - } + let credential = self + .source + .load("ADLS", |config| { + let suffix = self + .sas_token_suffixes + .iter() + .find(|suffix| { + config + .get(&format!("{ADLS_SAS_TOKEN_PREFIX}{suffix}")) + .is_some_and(|token| !token.is_empty()) + }) + .ok_or_else(|| { + ReqsignError::unexpected(format!( + "vended credential is missing {ADLS_SAS_TOKEN_PREFIX}{}", + self.sas_token_suffixes[0] + )) + })?; + let sas_token = required_credential_property( + config, + &format!("{ADLS_SAS_TOKEN_PREFIX}{suffix}"), + )?; + Ok( + match credential_expiry( + config, + &format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{suffix}"), + )? { + Some(expires_at) => { + AzureCredential::with_sas_token_expires_at(sas_token, expires_at) + } + None => AzureCredential::with_sas_token(sas_token), + }, + ) + }) + .await?; + Ok(Some(credential)) } } @@ -346,16 +372,24 @@ pub(crate) struct AzureStoragePath { } impl AzureStoragePath { + /// The host of the path, e.g. `account.dfs.core.windows.net`. + fn host(&self) -> String { + let service = match self.scheme { + AzureStorageScheme::Abfs | AzureStorageScheme::Abfss => "dfs", + AzureStorageScheme::Wasb | AzureStorageScheme::Wasbs => "blob", + }; + format!("{}.{service}.{}", self.account_name, self.endpoint_suffix) + } + /// Converts the AzureStoragePath into a full endpoint URL. /// /// This is possible because the path is fully qualified. + /// + /// Like Iceberg Java, the endpoint always uses TLS, also for the + /// non-secure schemes: SAS tokens are query parameters and must not be + /// sent over a plaintext connection unless the user configures one. fn as_endpoint(&self) -> String { - format!( - "{}://{}.dfs.{}", - self.scheme.as_http_scheme(), - self.account_name, - self.endpoint_suffix - ) + format!("https://{}.dfs.{}", self.account_name, self.endpoint_suffix) } } @@ -447,12 +481,11 @@ fn validate_storage_and_scheme( mod tests { use std::collections::HashMap; use std::sync::Arc; - use std::time::{Duration, SystemTime}; use async_trait::async_trait; use iceberg::Result; use iceberg::io::{ - ADLS_SAS_TOKEN_PREFIX, AzdlsCredential, StorageCredential, StorageCredentialKind, + ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX, ADLS_SAS_TOKEN_PREFIX, StorageCredential, StorageCredentialProvider, }; use opendal::services::AzdlsConfig; @@ -462,13 +495,28 @@ mod tests { use super::{ AzdlsSasTokens, AzureStoragePath, AzureStorageScheme, VendedAzdlsCredentialProvider, VendedCredentialSource, azdls_batch_key, azdls_config_parse, azdls_create_operator, + sas_token_suffixes, }; + fn adapter(credential: StorageCredential, path: &str) -> VendedAzdlsCredentialProvider { + VendedAzdlsCredentialProvider { + source: VendedCredentialSource::new( + Arc::new(FixedCredentialProvider(credential)), + path.to_string(), + ), + sas_token_suffixes: sas_token_suffixes(&path.parse().unwrap()), + } + } + #[derive(Debug)] struct FixedCredentialProvider(StorageCredential); #[async_trait] impl StorageCredentialProvider for FixedCredentialProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, _path: &str) -> Result { Ok(self.0.clone()) } @@ -647,31 +695,37 @@ mod tests { #[tokio::test] async fn vended_provider_returns_expiring_sas_credential() { let path = "abfss://container@account.dfs.core.windows.net/table/data.parquet"; - let expires_at = SystemTime::now() + Duration::from_secs(3600); - let credential = StorageCredential::new(StorageCredentialKind::Azdls( - AzdlsCredential::new("sv=2026&sig=secret"), - )) - .with_prefix("abfss://container@account.dfs.core.windows.net/table") - .with_expiration(expires_at); - let provider = VendedAzdlsCredentialProvider(VendedCredentialSource::new( - Arc::new(FixedCredentialProvider(credential)), - path.to_string(), - )); + let host = "account.dfs.core.windows.net"; + let credential = StorageCredential::new( + "abfss://container@account.dfs.core.windows.net/table", + HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}{host}"), + "sv=2026&sig=host".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_EXPIRES_AT_MS_PREFIX}{host}"), + "1500".to_string(), + ), + // The host-keyed token wins over the account-keyed one. + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account"), + "sv=2026&sig=account".to_string(), + ), + ]), + ); - let credential = provider + let credential = adapter(credential, path) .provide_credential(&Context::new()) .await .unwrap() .unwrap(); match credential { - AzureCredential::SasToken { - token, - expires_at: actual_expires_at, - } => { - assert_eq!(token, "sv=2026&sig=secret"); + AzureCredential::SasToken { token, expires_at } => { + assert_eq!(token, "sv=2026&sig=host"); assert_eq!( - actual_expires_at, - Some(crate::utils::system_time_to_timestamp(expires_at).unwrap()) + expires_at, + Some(reqsign_core::time::Timestamp::from_millisecond(1500).unwrap()) ); } other => panic!("expected SAS token, got {other:?}"), @@ -700,32 +754,170 @@ mod tests { #[test] fn host_keyed_sas_tokens_match_java() { - let properties = HashMap::from([( + let properties = HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=host".to_string(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account"), + "sv=2026&sig=account".to_string(), + ), + ]); + let sas_tokens = AzdlsSasTokens::from_properties(&properties); + let token_for = |location: &str| { + sas_tokens + .for_path(&location.parse::().unwrap()) + .map(str::to_owned) + }; + + // The exact host wins over the account. + assert_eq!( + token_for("abfss://container@account.dfs.core.windows.net/table/data.parquet") + .as_deref(), + Some("sv=2026&sig=host") + ); + // Other hosts of the account fall back to the account-keyed token. + assert_eq!( + token_for("wasbs://container@account.blob.core.windows.net/table/data.parquet") + .as_deref(), + Some("sv=2026&sig=account") + ); + + // A host-keyed token is never sent to the same account name in another + // cloud. + let host_only = AzdlsSasTokens::from_properties(&HashMap::from([( format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), "sv=2026&sig=host".to_string(), - )]); - let sas_tokens = AzdlsSasTokens::from_properties(&properties); + )])); + let other_cloud = "abfss://container@account.dfs.core.usgovcloudapi.net/data.parquet" + .parse::() + .unwrap(); + assert_eq!(host_only.for_path(&other_cloud), None); + } - for location in [ - "abfss://container@account.dfs.core.windows.net/table/data.parquet", - "wasbs://container@account.blob.core.windows.net/table/data.parquet", - ] { - let path = location.parse::().unwrap(); - assert_eq!(sas_tokens.for_path(&path), Some("sv=2026&sig=host")); + #[tokio::test] + async fn vended_provider_falls_back_to_the_account_token_when_the_host_token_is_empty() { + let path = "abfss://container@account.dfs.core.windows.net/table/data.parquet"; + let credential = StorageCredential::new( + "abfss://container@account.dfs.core.windows.net", + HashMap::from([ + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + String::new(), + ), + ( + format!("{ADLS_SAS_TOKEN_PREFIX}account"), + "sv=2026&sig=account".to_string(), + ), + ]), + ); + + match adapter(credential, path) + .provide_credential(&Context::new()) + .await + .unwrap() + .unwrap() + { + AzureCredential::SasToken { token, .. } => assert_eq!(token, "sv=2026&sig=account"), + other => panic!("expected SAS token, got {other:?}"), } } + /// Builds an ADLS operator for `path` that sends its requests to `server`. + fn operator_for( + server: &mockito::Server, + path: &str, + sas_tokens: &AzdlsSasTokens, + credential_provider: Option>, + ) -> opendal::Operator { + let config = AzdlsConfig { + endpoint: Some(server.url()), + ..Default::default() + }; + super::azdls_config_build( + &config, + &path.parse().unwrap(), + sas_tokens, + &credential_provider, + path, + None, + ) + .unwrap() + } + + #[tokio::test] + async fn operator_signs_with_the_vended_sas_token() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("HEAD", mockito::Matcher::Any) + .match_query(mockito::Matcher::Regex("sig=vended".to_string())) + .expect_at_least(1) + .with_status(404) + .create_async() + .await; + let path = "abfss://container@account.dfs.core.windows.net/table/data.parquet"; + let credential = StorageCredential::new( + "abfss://container@account.dfs.core.windows.net", + HashMap::from([( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=vended".to_string(), + )]), + ); + // A static token is not used when a provider supplies credentials. + let static_tokens = AzdlsSasTokens::from_properties(&HashMap::from([( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=static".to_string(), + )])); + + let operator = operator_for( + &server, + path, + &static_tokens, + Some(Arc::new(FixedCredentialProvider(credential))), + ); + assert!(operator.stat("table/data.parquet").await.is_err()); + mock.assert_async().await; + } + + #[tokio::test] + async fn operator_signs_with_the_static_account_sas_token() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("HEAD", mockito::Matcher::Any) + .match_query(mockito::Matcher::Regex("sig=static".to_string())) + .expect_at_least(1) + .with_status(404) + .create_async() + .await; + let static_tokens = AzdlsSasTokens::from_properties(&HashMap::from([( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=static".to_string(), + )])); + + let operator = operator_for( + &server, + "abfss://container@account.dfs.core.windows.net/table/data.parquet", + &static_tokens, + None, + ); + assert!(operator.stat("table/data.parquet").await.is_err()); + mock.assert_async().await; + } + #[tokio::test] async fn vended_provider_rejects_mismatched_prefix() { - let credential = StorageCredential::new(StorageCredentialKind::Azdls( - AzdlsCredential::new("sv=2026&sig=secret"), - )) - .with_prefix("abfss://other@account.dfs.core.windows.net/table") - .with_expiration(SystemTime::now() + Duration::from_secs(3600)); - let provider = VendedAzdlsCredentialProvider(VendedCredentialSource::new( - Arc::new(FixedCredentialProvider(credential)), - "abfss://container@account.dfs.core.windows.net/table/data.parquet".to_string(), - )); + let credential = StorageCredential::new( + "abfss://other@account.dfs.core.windows.net/table", + HashMap::from([( + format!("{ADLS_SAS_TOKEN_PREFIX}account.dfs.core.windows.net"), + "sv=2026&sig=secret".to_string(), + )]), + ); + let provider = adapter( + credential, + "abfss://container@account.dfs.core.windows.net/table/data.parquet", + ); assert!(provider.provide_credential(&Context::new()).await.is_err()); } diff --git a/crates/storage/opendal/src/gcs.rs b/crates/storage/opendal/src/gcs.rs index f374b8034c..76fe5e6404 100644 --- a/crates/storage/opendal/src/gcs.rs +++ b/crates/storage/opendal/src/gcs.rs @@ -21,16 +21,19 @@ use std::sync::Arc; use iceberg::io::{ GCS_ALLOW_ANONYMOUS, GCS_CREDENTIALS_JSON, GCS_DISABLE_CONFIG_LOAD, GCS_DISABLE_VM_METADATA, - GCS_NO_AUTH, GCS_SERVICE_HOST, GCS_TOKEN, StorageCredentialKind, StorageCredentialProvider, + GCS_NO_AUTH, GCS_SERVICE_HOST, GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; use opendal::services::GcsConfig; use opendal::{Configurator, Operator}; -use reqsign_core::{Context, Error as ReqsignError, ProvideCredential, Result as ReqsignResult}; +use reqsign_core::{Context, ProvideCredential, Result as ReqsignResult}; use reqsign_google::{Credential as GoogleCredential, Token as GoogleToken}; use url::Url; -use crate::utils::{VendedCredentialSource, from_opendal_error, is_truthy}; +use crate::utils::{ + VendedCredentialSource, credential_expiry, from_opendal_error, is_truthy, + required_credential_property, +}; /// Parse iceberg properties to [`GcsConfig`]. pub(crate) fn gcs_config_parse(mut m: HashMap) -> Result { @@ -99,10 +102,6 @@ pub(crate) fn gcs_config_build( let mut cfg = cfg.clone(); cfg.bucket = bucket.to_string(); - // `reqsign_google` continues to the next provider even when a provider returns - // an error, and OpenDAL prepends custom providers to its default chain. Disable - // every other configured and ambient source so the catalog provider is - // effectively the sole source and refresh failures cannot silently fall back. let credential_provider = credential_provider .as_ref() .filter(|provider| provider.supports_path(path)); @@ -113,12 +112,7 @@ pub(crate) fn gcs_config_build( "Invalid GCS auth settings: anonymous access cannot be combined with refreshable credentials", )); } - cfg.token = None; - cfg.credential = None; - cfg.credential_path = None; - cfg.service_account = None; - cfg.disable_vm_metadata = true; - cfg.disable_config_load = true; + disable_other_credential_sources(&mut cfg); } let mut builder = cfg.into_builder(); @@ -135,6 +129,21 @@ pub(crate) fn gcs_config_build( Operator::new(builder).map_err(from_opendal_error) } +/// Disable every configured and ambient credential source but a vended one. +/// +/// `reqsign_google` continues to the next provider even when a provider returns +/// an error, and OpenDAL prepends custom providers to its default chain. With +/// the other sources disabled, the vended provider is effectively the sole +/// source, and refresh failures cannot silently fall back. +fn disable_other_credential_sources(cfg: &mut GcsConfig) { + cfg.token = None; + cfg.credential = None; + cfg.credential_path = None; + cfg.service_account = None; + cfg.disable_vm_metadata = true; + cfg.disable_config_load = true; +} + /// Adapts a generic [`StorageCredentialProvider`] into a `reqsign` /// [`ProvideCredential`], so the GCS signer can obtain and refresh vended OAuth2 /// tokens. @@ -145,16 +154,102 @@ impl ProvideCredential for VendedGcsCredentialProvider { type Credential = GoogleCredential; async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { - match self.0.load("GCS").await? { - (StorageCredentialKind::Gcs(gcs), expires_at) => { - Ok(Some(GoogleCredential::with_token(GoogleToken { - access_token: gcs.into_token(), - expires_at, - }))) - } - _ => Err(ReqsignError::unexpected( - "GCS storage received a non-GCS credential from the provider", - )), + let token = self + .0 + .load("GCS", |config| { + Ok(GoogleToken { + access_token: required_credential_property(config, GCS_TOKEN)?.to_string(), + expires_at: credential_expiry(config, GCS_TOKEN_EXPIRES_AT)?, + }) + }) + .await?; + Ok(Some(GoogleCredential::with_token(token))) + } +} + +#[cfg(test)] +mod tests { + use async_trait::async_trait; + use iceberg::io::{S3_ACCESS_KEY_ID, StorageCredential}; + use reqsign_core::time::Timestamp; + + use super::*; + + #[derive(Debug)] + struct FixedProvider(StorageCredential); + + #[async_trait] + impl StorageCredentialProvider for FixedProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + + async fn load_credential(&self, _path: &str) -> Result { + Ok(self.0.clone()) } } + + fn adapter(credential: StorageCredential) -> VendedGcsCredentialProvider { + VendedGcsCredentialProvider(VendedCredentialSource::new( + Arc::new(FixedProvider(credential)), + "gs://bucket/table/file.parquet".to_string(), + )) + } + + fn credential(config: &[(&str, &str)]) -> StorageCredential { + StorageCredential::new( + "gs://bucket/table", + config + .iter() + .map(|(key, value)| (key.to_string(), value.to_string())) + .collect(), + ) + } + + #[tokio::test] + async fn vended_adapter_returns_expiring_oauth2_token() { + let credential = credential(&[(GCS_TOKEN, "ya29.token"), (GCS_TOKEN_EXPIRES_AT, "1500")]); + + let google = adapter(credential) + .provide_credential(&Context::new()) + .await + .unwrap() + .unwrap(); + let token = google.token.expect("a token credential"); + assert_eq!(token.access_token, "ya29.token"); + assert_eq!( + token.expires_at, + Some(Timestamp::from_millisecond(1500).unwrap()) + ); + } + + #[tokio::test] + async fn vended_adapter_rejects_other_credentials() { + assert!( + adapter(credential(&[(S3_ACCESS_KEY_ID, "AK")])) + .provide_credential(&Context::new()) + .await + .is_err() + ); + } + + #[test] + fn vended_credentials_disable_every_other_source() { + let mut cfg = gcs_config_parse(HashMap::from([ + (GCS_TOKEN.to_string(), "static-token".to_string()), + (GCS_CREDENTIALS_JSON.to_string(), "e30=".to_string()), + ])) + .unwrap(); + cfg.credential_path = Some("/path/to/credentials.json".to_string()); + cfg.service_account = Some("account@example.com".to_string()); + + disable_other_credential_sources(&mut cfg); + + assert_eq!(cfg.token, None); + assert_eq!(cfg.credential, None); + assert_eq!(cfg.credential_path, None); + assert_eq!(cfg.service_account, None); + assert!(cfg.disable_vm_metadata); + assert!(cfg.disable_config_load); + } } diff --git a/crates/storage/opendal/src/lib.rs b/crates/storage/opendal/src/lib.rs index 67a71060c4..8775504a69 100644 --- a/crates/storage/opendal/src/lib.rs +++ b/crates/storage/opendal/src/lib.rs @@ -160,14 +160,34 @@ where )) } +/// Fails serialization of a storage holding a credential provider, which would +/// otherwise silently fall back to the configured or ambient credentials. +/// [`FileIO`](iceberg::io::FileIO) serializes its provider separately. +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +pub(crate) fn serialize_credential_provider( + _provider: &Option>, + _serializer: S, +) -> std::result::Result +where + S: serde::Serializer, +{ + Err(serde::ser::Error::custom( + "storage with a credential provider cannot be serialized; serialize its FileIO instead", + )) +} + #[typetag::serde(name = "OpenDalStorageFactory")] impl StorageFactory for OpenDalStorageFactory { fn build(&self, config: &StorageConfig) -> Result> { - self.build_with_credentials(config, None) + self.build_with_credential_provider(config, None) } #[allow(unused_variables)] - fn build_with_credentials( + fn build_with_credential_provider( &self, config: &StorageConfig, credential_provider: Option>, @@ -254,7 +274,11 @@ pub enum OpenDalStorage { #[serde(skip)] customized_credential_load: Option, /// Provider of refreshable vended credentials, supplied by the catalog. - #[serde(skip)] + #[serde( + skip_deserializing, + skip_serializing_if = "Option::is_none", + serialize_with = "serialize_credential_provider" + )] credential_provider: Option>, }, /// GCS storage variant. @@ -264,7 +288,11 @@ pub enum OpenDalStorage { /// GCS configuration. config: Arc, /// Provider of refreshable vended credentials, supplied by the catalog. - #[serde(skip)] + #[serde( + skip_deserializing, + skip_serializing_if = "Option::is_none", + serialize_with = "serialize_credential_provider" + )] credential_provider: Option>, }, /// OSS storage variant. @@ -286,10 +314,14 @@ pub enum OpenDalStorage { /// Azure DLS configuration. config: Arc, /// Account-specific SAS tokens supplied by Java-compatible properties. - #[serde(default)] + #[serde(default, skip_serializing_if = "AzdlsSasTokens::is_empty")] sas_tokens: Arc, /// Provider of refreshable vended credentials, supplied by the catalog. - #[serde(skip)] + #[serde( + skip_deserializing, + skip_serializing_if = "Option::is_none", + serialize_with = "serialize_credential_provider" + )] credential_provider: Option>, }, /// HuggingFace Hub storage variant. @@ -309,10 +341,56 @@ pub enum OpenDalStorage { #[derive(Clone, Debug, Eq, Hash, PartialEq)] struct DeleteBatchKey { storage: String, - /// Location used to look up the dynamic credential shared by the batch: - /// the credential prefix, or the storage root for an unscoped credential. - /// `None` when the path is served by static credentials. - credential_location: Option, + credential: BatchCredential, +} + +impl DeleteBatchKey { + /// Open batches to flush before adding a path with this key. + /// + /// A scoped credential inside the scope of an open batch means the catalog + /// is replacing that batch's broader credential, such as a scheme-wide or + /// account-wide seed. Such batches are flushed while their credential is + /// still valid: once it expires, nothing may cover their whole scope. + fn superseded_batches<'a>( + &self, + open: impl Iterator, + ) -> Vec { + let BatchCredential::Scoped { prefix } = &self.credential else { + return Vec::new(); + }; + open.filter(|key| { + key.storage == self.storage + && key + .credential + .location() + .is_some_and(|location| location != prefix && prefix.starts_with(location)) + }) + .cloned() + .collect() + } +} + +/// The credential shared by the paths of a bulk-delete batch. +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +enum BatchCredential { + /// Served by the backend's configured credentials. + Static, + /// A dynamic credential for a whole scheme, such as `s3`, looked up for + /// the storage root. + SchemeWide { root: String }, + /// A dynamic credential scoped to `prefix`, looked up for the prefix. + Scoped { prefix: String }, +} + +impl BatchCredential { + /// Location the batch operator looks dynamic credentials up for. + fn location(&self) -> Option<&str> { + match self { + BatchCredential::Static => None, + BatchCredential::SchemeWide { root } => Some(root), + BatchCredential::Scoped { prefix } => Some(prefix), + } + } } impl OpenDalStorage { @@ -328,7 +406,6 @@ impl OpenDalStorage { /// /// * An [`opendal::Operator`] instance used to operate on file. /// * Relative path to the root uri of [`opendal::Operator`]. - #[allow(unreachable_code, unused_variables)] pub(crate) fn create_operator<'a>( &self, path: &'a impl AsRef, @@ -341,7 +418,7 @@ impl OpenDalStorage { /// /// A bulk-delete batch passes its scope location, so every credential the /// operator loads covers all paths in the batch. - #[allow(unreachable_code, unused_variables)] + #[allow(unreachable_code, unreachable_patterns, unused_variables)] fn create_operator_with_credential_location<'a>( &self, path: &'a impl AsRef, @@ -528,29 +605,34 @@ impl OpenDalStorage { /// scope. Loading the credential is normally a cache hit and avoids rebuilding /// an operator for every path while preventing a batch from crossing prefixes. async fn delete_batch_key_for_path(&self, path: &str) -> Result { - let credential_location = match self.credential_provider_for_path(path) { + let credential = match self.credential_provider_for_path(path) { Some(provider) => { let credential = provider.load_credential(path).await?; if !credential.covers(path) { return Err(Error::new( ErrorKind::DataInvalid, - format!( - "vended credential prefix {:?} does not cover storage location {path:?}", - credential.prefix() - ), + utils::uncovered_location_message(path, &credential), )); } - Some(match credential.prefix() { - Some(prefix) => prefix.to_string(), - None => utils::storage_root(path)?, - }) + let prefix = credential.prefix(); + if prefix.contains("://") { + BatchCredential::Scoped { + prefix: prefix.to_string(), + } + } else { + // A scheme-wide prefix such as `s3` is no location to look + // up, so the batch looks its credential up for the root. + BatchCredential::SchemeWide { + root: utils::storage_root(path)?, + } + } } - None => None, + None => BatchCredential::Static, }; Ok(DeleteBatchKey { storage: self.batch_key_for_path(path)?, - credential_location, + credential, }) } @@ -559,7 +641,7 @@ impl OpenDalStorage { /// This is a lightweight alternative to [`create_operator`](Self::create_operator) for cases /// where only the relative path is needed (e.g. bulk deletes where the operator is already /// available). - #[allow(unreachable_code, unused_variables)] + #[allow(unreachable_code, unreachable_patterns, unused_variables)] pub(crate) fn relativize_path<'a>(&self, path: &'a str) -> Result<&'a str> { match self { #[cfg(feature = "opendal-memory")] @@ -731,6 +813,12 @@ impl Storage for OpenDalStorage { while let Some(path) = paths.next().await { let batch_key = self.delete_batch_key_for_path(&path).await?; + for key in batch_key.superseded_batches(deleters.keys()) { + if let Some(mut deleter) = deleters.remove(&key) { + deleter.close().await.map_err(from_opendal_error)?; + } + } + let (relative_path, deleter) = match deleters.entry(batch_key) { Entry::Occupied(entry) => { (self.relativize_path(&path)?.to_string(), entry.into_mut()) @@ -738,7 +826,7 @@ impl Storage for OpenDalStorage { Entry::Vacant(entry) => { let (op, rel) = self.create_operator_with_credential_location( &path, - entry.key().credential_location.as_deref(), + entry.key().credential.location(), )?; let rel = rel.to_string(); let deleter = op.deleter().await.map_err(from_opendal_error)?; @@ -856,21 +944,45 @@ mod tests { ))] #[async_trait] impl StorageCredentialProvider for AlwaysSupportedCredentialProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, path: &str) -> Result { let prefix = if path.contains("/table-a/") { "s3://bucket/table-a" } else { "s3://bucket/table-b" }; - Ok( - iceberg::io::StorageCredential::new(iceberg::io::StorageCredentialKind::S3( - iceberg::io::S3Credential::new("access-key", "secret-key", None), - )) - .with_prefix(prefix), - ) + Ok(s3_credential(prefix, "access-key")) } } + /// An S3 credential for locations under `prefix`. + #[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + all(feature = "opendal-memory", feature = "opendal-azdls") + ))] + fn s3_credential( + prefix: impl Into, + access_key: &str, + ) -> iceberg::io::StorageCredential { + iceberg::io::StorageCredential::new( + prefix, + HashMap::from([ + ( + iceberg::io::S3_ACCESS_KEY_ID.to_string(), + access_key.to_string(), + ), + ( + iceberg::io::S3_SECRET_ACCESS_KEY.to_string(), + "secret-key".to_string(), + ), + ]), + ) + } + #[cfg(feature = "opendal-s3")] #[derive(Debug)] struct EmptyCredentialLoader; @@ -920,7 +1032,7 @@ mod tests { #[test] fn test_factory_ignores_credentials_for_unsupported_backend() { let storage = OpenDalStorageFactory::Memory - .build_with_credentials( + .build_with_credential_provider( &StorageConfig::new(), Some(Arc::new(AlwaysSupportedCredentialProvider)), ) @@ -1016,6 +1128,198 @@ mod tests { } } + #[cfg(feature = "opendal-azdls")] + #[test] + fn test_azdls_storage_omits_empty_sas_tokens() { + let storage = |sas_tokens: AzdlsSasTokens| OpenDalStorage::Azdls { + config: Arc::new(AzdlsConfig::default()), + sas_tokens: Arc::new(sas_tokens), + credential_provider: None, + }; + + let json = serde_json::to_value(storage(AzdlsSasTokens::default())).unwrap(); + assert!(json["Azdls"].get("sas_tokens").is_none(), "{json}"); + let deserialized: OpenDalStorage = serde_json::from_value(json).unwrap(); + assert!(matches!(deserialized, OpenDalStorage::Azdls { .. })); + } + + /// Vends a credential for the whole `s3` scheme until it has served + /// `table-c`, then credentials scoped to each table. The scheme-wide + /// credential expires once the `table-c` batch is signed. + #[cfg(feature = "opendal-s3")] + #[derive(Debug, Default)] + struct ChangingScopeProvider { + scoped: std::sync::atomic::AtomicBool, + scheme_wide_expired: std::sync::atomic::AtomicBool, + } + + #[cfg(feature = "opendal-s3")] + #[async_trait] + impl StorageCredentialProvider for ChangingScopeProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + + async fn load_credential(&self, path: &str) -> Result { + use std::sync::atomic::Ordering; + + if path.starts_with("s3://bucket/table-c") { + self.scoped.store(true, Ordering::SeqCst); + } + // Signing the scoped batch means the scheme-wide credential is gone. + if path == "s3://bucket/table-c" { + self.scheme_wide_expired.store(true, Ordering::SeqCst); + } + match path.strip_prefix("s3://bucket/") { + Some(rest) if self.scoped.load(Ordering::SeqCst) && !rest.is_empty() => { + let table = rest.split('/').next().unwrap_or_default(); + Ok(s3_credential(format!("s3://bucket/{table}"), "SCOPED_AK")) + } + _ if self.scheme_wide_expired.load(Ordering::SeqCst) => Err(Error::new( + ErrorKind::Unexpected, + "the scheme-wide credential expired", + )), + _ => Ok(s3_credential("s3", "SCHEME_AK")), + } + } + } + + #[cfg(feature = "opendal-s3")] + #[tokio::test] + async fn test_delete_stream_signs_each_batch_with_its_vended_credential() { + let mut server = mockito::Server::new_async().await; + let mut delete = |path: &str, access_key: &str| { + server + .mock("DELETE", path) + .match_header( + "authorization", + mockito::Matcher::Regex(format!("Credential={access_key}/")), + ) + .expect(1) + .with_status(204) + }; + // The scheme-wide batch is flushed, with its own credential, before + // the scoped credential that replaces it is used. + let scheme_wide = delete("/bucket/table-b/f.parquet", "SCHEME_AK") + .create_async() + .await; + let scoped = delete("/bucket/table-c/f.parquet", "SCOPED_AK") + .create_async() + .await; + + let mut config = S3Config::default(); + config.endpoint = Some(server.url()); + config.region = Some("us-east-1".to_string()); + config.disable_config_load = true; + config.disable_ec2_metadata = true; + let storage = OpenDalStorage::S3 { + config: Arc::new(config), + customized_credential_load: None, + credential_provider: Some(Arc::new(ChangingScopeProvider::default())), + }; + storage + .delete_stream( + futures::stream::iter([ + "s3://bucket/table-b/f.parquet".to_string(), + "s3://bucket/table-c/f.parquet".to_string(), + ]) + .boxed(), + ) + .await + .unwrap(); + + scheme_wide.assert_async().await; + scoped.assert_async().await; + } + + /// Vends `token` as a GCS credential, or fails when it is `None`. + #[cfg(feature = "opendal-gcs")] + #[derive(Debug)] + struct GcsTokenProvider(Option<&'static str>); + + #[cfg(feature = "opendal-gcs")] + #[async_trait] + impl StorageCredentialProvider for GcsTokenProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + + async fn load_credential(&self, _path: &str) -> Result { + let token = self + .0 + .ok_or_else(|| Error::new(ErrorKind::Unexpected, "refresh failed"))?; + Ok(iceberg::io::StorageCredential::new( + "gs", + HashMap::from([(iceberg::io::GCS_TOKEN.to_string(), token.to_string())]), + )) + } + } + + #[cfg(feature = "opendal-gcs")] + fn gcs_storage(server: &mockito::Server, provider: GcsTokenProvider) -> OpenDalStorage { + let config = gcs_config_parse(HashMap::from([ + (iceberg::io::GCS_SERVICE_HOST.to_string(), server.url()), + ( + iceberg::io::GCS_TOKEN.to_string(), + "static-token".to_string(), + ), + ])) + .unwrap(); + OpenDalStorage::Gcs { + config: Arc::new(config), + credential_provider: Some(Arc::new(provider)), + } + } + + #[cfg(feature = "opendal-gcs")] + #[tokio::test] + async fn test_gcs_signs_with_the_vended_token() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("GET", mockito::Matcher::Any) + .match_header("authorization", "Bearer vended-token") + .expect_at_least(1) + .with_status(404) + .create_async() + .await; + + let storage = gcs_storage(&server, GcsTokenProvider(Some("vended-token"))); + assert!(!storage.exists("gs://bucket/file.parquet").await.unwrap()); + mock.assert_async().await; + } + + #[cfg(feature = "opendal-gcs")] + #[tokio::test] + async fn test_gcs_does_not_fall_back_when_the_provider_fails() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("GET", mockito::Matcher::Any) + .expect(0) + .create_async() + .await; + + let storage = gcs_storage(&server, GcsTokenProvider(None)); + assert!(storage.exists("gs://bucket/file.parquet").await.is_err()); + mock.assert_async().await; + } + + #[cfg(feature = "opendal-s3")] + #[test] + fn test_storage_with_credential_provider_is_not_serializable() { + let storage = + |credential_provider: Option>| OpenDalStorage::S3 { + config: Arc::new(S3Config::default()), + customized_credential_load: None, + credential_provider, + }; + + assert!(serde_json::to_string(&storage(None)).is_ok()); + let error = + serde_json::to_string(&storage(Some(Arc::new(AlwaysSupportedCredentialProvider)))) + .unwrap_err(); + assert!(error.to_string().contains("credential provider"), "{error}"); + } + #[cfg(feature = "opendal-s3")] #[tokio::test] async fn test_dynamic_credentials_batch_by_prefix() { @@ -1029,10 +1333,7 @@ mod tests { let other_scope = "s3://bucket/table-b/file.parquet"; let first_key = storage.delete_batch_key_for_path(first).await.unwrap(); - assert_eq!( - first_key.credential_location.as_deref(), - Some("s3://bucket/table-a") - ); + assert_eq!(first_key.credential.location(), Some("s3://bucket/table-a")); assert_eq!( first_key, storage.delete_batch_key_for_path(same_scope).await.unwrap() @@ -1048,34 +1349,34 @@ mod tests { #[cfg(feature = "opendal-s3")] #[tokio::test] - async fn test_unscoped_dynamic_credentials_batch_by_storage_root() { + async fn test_scheme_wide_dynamic_credentials_batch_by_storage_root() { #[derive(Debug)] - struct UnscopedProvider; + struct SchemeWideProvider; #[async_trait] - impl StorageCredentialProvider for UnscopedProvider { + impl StorageCredentialProvider for SchemeWideProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, _path: &str) -> Result { - Ok(iceberg::io::StorageCredential::new( - iceberg::io::StorageCredentialKind::S3(iceberg::io::S3Credential::new( - "access-key", - "secret-key", - None, - )), - )) + Ok(s3_credential("s3", "access-key")) } } let storage = OpenDalStorage::S3 { config: Arc::new(S3Config::default()), customized_credential_load: None, - credential_provider: Some(Arc::new(UnscopedProvider)), + credential_provider: Some(Arc::new(SchemeWideProvider)), }; let key = storage .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") .await .unwrap(); - assert_eq!(key.credential_location.as_deref(), Some("s3://bucket/")); + assert_eq!(key.credential, BatchCredential::SchemeWide { + root: "s3://bucket/".to_string() + }); assert_eq!( key, storage @@ -1085,6 +1386,51 @@ mod tests { ); } + #[test] + fn test_scoped_batch_flushes_broader_batches_of_same_storage() { + let key = |storage: &str, credential: BatchCredential| DeleteBatchKey { + storage: storage.to_string(), + credential, + }; + let unscoped = key("bucket", BatchCredential::SchemeWide { + root: "s3://bucket/".to_string(), + }); + let other_storage = key("other", BatchCredential::SchemeWide { + root: "s3://other/".to_string(), + }); + let scoped = key("bucket", BatchCredential::Scoped { + prefix: "s3://bucket/table-a".to_string(), + }); + let open = [unscoped.clone(), other_storage, scoped.clone()]; + + let new_scope = key("bucket", BatchCredential::Scoped { + prefix: "s3://bucket/table-b".to_string(), + }); + assert_eq!(new_scope.superseded_batches(open.iter()), vec![ + unscoped.clone() + ]); + + // A narrower scope also supersedes a broader scoped batch, such as one + // using an account-wide ADLS seed scoped to its container. + let nested = key("bucket", BatchCredential::Scoped { + prefix: "s3://bucket/table-a/data".to_string(), + }); + let superseded = nested.superseded_batches(open.iter()); + assert_eq!(superseded.len(), 2); + assert!(superseded.contains(&unscoped) && superseded.contains(&scoped)); + // A batch never supersedes itself. + assert_eq!(scoped.superseded_batches(open.iter()), vec![ + unscoped.clone() + ]); + // Scheme-wide and static paths leave open batches alone. + assert!(unscoped.superseded_batches(open.iter()).is_empty()); + assert!( + key("bucket", BatchCredential::Static) + .superseded_batches(open.iter()) + .is_empty() + ); + } + #[cfg(feature = "opendal-s3")] #[tokio::test] async fn test_custom_s3_credential_loader_ignores_dynamic_provider_for_batching() { @@ -1100,7 +1446,7 @@ mod tests { .delete_batch_key_for_path("s3://bucket/table-a/file.parquet") .await .unwrap(); - assert_eq!(key.credential_location, None); + assert_eq!(key.credential, BatchCredential::Static); } #[cfg(feature = "opendal-s3")] diff --git a/crates/storage/opendal/src/resolving.rs b/crates/storage/opendal/src/resolving.rs index 29a92ad9fa..a8e220e8bb 100644 --- a/crates/storage/opendal/src/resolving.rs +++ b/crates/storage/opendal/src/resolving.rs @@ -207,11 +207,11 @@ impl OpenDalResolvingStorageFactory { #[typetag::serde] impl StorageFactory for OpenDalResolvingStorageFactory { fn build(&self, config: &StorageConfig) -> Result> { - self.build_with_credentials(config, None) + self.build_with_credential_provider(config, None) } #[allow(unused_variables)] - fn build_with_credentials( + fn build_with_credential_provider( &self, config: &StorageConfig, credential_provider: Option>, @@ -256,14 +256,17 @@ pub struct OpenDalResolvingStorage { feature = "opendal-gcs", feature = "opendal-azdls" ))] - #[serde(skip)] + #[serde( + skip_deserializing, + skip_serializing_if = "Option::is_none", + serialize_with = "crate::serialize_credential_provider" + )] credential_provider: Option>, } impl std::fmt::Debug for OpenDalResolvingStorage { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - // `props` can contain storage secrets, and a custom credential provider - // may carry secret state in its own Debug implementation + // `props` can contain storage secrets f.debug_struct("OpenDalResolvingStorage") .field("property_keys", &self.props.keys().collect::>()) .finish_non_exhaustive() @@ -369,7 +372,19 @@ impl Storage for OpenDalResolvingStorage { Ok(()) } - #[allow(unreachable_code)] + #[cfg_attr( + not(any( + feature = "opendal-memory", + feature = "opendal-fs", + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-oss", + feature = "opendal-azdls", + feature = "opendal-hf" + )), + // Without a backend, there is no storage to resolve. + allow(unreachable_code) + )] fn new_input(&self, path: &str) -> Result { Ok(InputFile::new( Arc::new(self.resolve(path)?.as_ref().clone()), @@ -377,7 +392,19 @@ impl Storage for OpenDalResolvingStorage { )) } - #[allow(unreachable_code)] + #[cfg_attr( + not(any( + feature = "opendal-memory", + feature = "opendal-fs", + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-oss", + feature = "opendal-azdls", + feature = "opendal-hf" + )), + // Without a backend, there is no storage to resolve. + allow(unreachable_code) + )] fn new_output(&self, path: &str) -> Result { Ok(OutputFile::new( Arc::new(self.resolve(path)?.as_ref().clone()), @@ -412,6 +439,10 @@ mod tests { ))] #[async_trait] impl StorageCredentialProvider for AllPathsCredentialProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, _path: &str) -> Result { unreachable!("unsupported backends must ignore the provider") } @@ -426,7 +457,7 @@ mod tests { fn test_factory_ignores_credentials_without_compatible_backend() { assert!( OpenDalResolvingStorageFactory::new() - .build_with_credentials( + .build_with_credential_provider( &StorageConfig::new(), Some(Arc::new(AllPathsCredentialProvider)), ) diff --git a/crates/storage/opendal/src/s3.rs b/crates/storage/opendal/src/s3.rs index ebe0518221..2fd7e1c4a1 100644 --- a/crates/storage/opendal/src/s3.rs +++ b/crates/storage/opendal/src/s3.rs @@ -22,7 +22,7 @@ use iceberg::io::{ CLIENT_REGION, S3_ACCESS_KEY_ID, S3_ALLOW_ANONYMOUS, S3_ASSUME_ROLE_ARN, S3_ASSUME_ROLE_EXTERNAL_ID, S3_ASSUME_ROLE_SESSION_NAME, S3_DISABLE_CONFIG_LOAD, S3_DISABLE_EC2_METADATA, S3_ENDPOINT, S3_PATH_STYLE_ACCESS, S3_REGION, S3_SECRET_ACCESS_KEY, - S3_SESSION_TOKEN, S3_SSE_KEY, S3_SSE_MD5, S3_SSE_TYPE, StorageCredentialKind, + S3_SESSION_TOKEN, S3_SESSION_TOKEN_EXPIRES_AT_MS, S3_SSE_KEY, S3_SSE_MD5, S3_SSE_TYPE, StorageCredentialProvider, }; use iceberg::{Error, ErrorKind, Result}; @@ -33,12 +33,14 @@ pub use reqsign_aws_v4::Credential as AwsCredential; /// Trait for types that can asynchronously supply [`AwsCredential`] to a [`CustomAwsCredentialLoader`]. pub use reqsign_core::ProvideCredential; use reqsign_core::{ - Context, Error as ReqsignError, ProvideCredentialChain, ProvideCredentialDyn, - Result as ReqsignResult, + Context, ProvideCredentialChain, ProvideCredentialDyn, Result as ReqsignResult, }; use url::Url; -use crate::utils::{VendedCredentialSource, from_opendal_error, is_truthy}; +use crate::utils::{ + VendedCredentialSource, credential_expiry, from_opendal_error, is_truthy, + required_credential_property, +}; /// Parse iceberg props to s3 config. pub(crate) fn s3_config_parse(mut m: HashMap) -> Result { @@ -190,20 +192,23 @@ impl ProvideCredential for VendedS3CredentialProvider { type Credential = AwsCredential; async fn provide_credential(&self, _ctx: &Context) -> ReqsignResult> { - match self.0.load("S3").await? { - (StorageCredentialKind::S3(s3), expires_in) => { - let (access_key_id, secret_access_key, session_token) = s3.into_parts(); - Ok(Some(AwsCredential { - access_key_id, - secret_access_key, - session_token, - expires_in, - })) - } - _ => Err(ReqsignError::unexpected( - "S3 storage received a non-S3 credential from the provider", - )), - } + let credential = self + .0 + .load("S3", |config| { + Ok(AwsCredential { + access_key_id: required_credential_property(config, S3_ACCESS_KEY_ID)? + .to_string(), + secret_access_key: required_credential_property(config, S3_SECRET_ACCESS_KEY)? + .to_string(), + session_token: config + .get(S3_SESSION_TOKEN) + .filter(|token| !token.is_empty()) + .cloned(), + expires_in: credential_expiry(config, S3_SESSION_TOKEN_EXPIRES_AT_MS)?, + }) + }) + .await?; + Ok(Some(credential)) } } @@ -237,10 +242,95 @@ impl CustomAwsCredentialLoader { #[cfg(test)] mod tests { use std::collections::HashMap; + use std::sync::Arc; + + use async_trait::async_trait; + use iceberg::io::{GCS_TOKEN, S3_PATH_STYLE_ACCESS, StorageCredential}; + use reqsign_core::time::Timestamp; + + use super::*; + + #[derive(Debug)] + struct FixedProvider(StorageCredential); + + #[async_trait] + impl StorageCredentialProvider for FixedProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + + async fn load_credential(&self, _path: &str) -> Result { + Ok(self.0.clone()) + } + } + + fn adapter(credential: StorageCredential) -> VendedS3CredentialProvider { + VendedS3CredentialProvider(VendedCredentialSource::new( + Arc::new(FixedProvider(credential)), + "s3://bucket/table/file.parquet".to_string(), + )) + } + + fn credential(prefix: &str, config: &[(&str, &str)]) -> StorageCredential { + StorageCredential::new( + prefix, + config + .iter() + .map(|(key, value)| (key.to_string(), value.to_string())) + .collect(), + ) + } + + #[tokio::test] + async fn vended_adapter_returns_expiring_aws_credential() { + let credential = credential("s3://bucket/table", &[ + (S3_ACCESS_KEY_ID, "AK"), + (S3_SECRET_ACCESS_KEY, "SK"), + (S3_SESSION_TOKEN, "TOK"), + (S3_SESSION_TOKEN_EXPIRES_AT_MS, "1500"), + ]); - use iceberg::io::S3_PATH_STYLE_ACCESS; + let aws = adapter(credential) + .provide_credential(&Context::new()) + .await + .unwrap() + .unwrap(); + assert_eq!(aws.access_key_id, "AK"); + assert_eq!(aws.secret_access_key, "SK"); + assert_eq!(aws.session_token.as_deref(), Some("TOK")); + assert_eq!( + aws.expires_in, + Some(Timestamp::from_millisecond(1500).unwrap()) + ); + } - use super::s3_config_parse; + #[tokio::test] + async fn vended_adapter_rejects_incomplete_or_uncovering_credentials() { + for credential in [ + // A credential for another backend. + credential("s3", &[(GCS_TOKEN, "token")]), + // No secret key. + credential("s3", &[(S3_ACCESS_KEY_ID, "AK")]), + // An unparsable expiry. + credential("s3", &[ + (S3_ACCESS_KEY_ID, "AK"), + (S3_SECRET_ACCESS_KEY, "SK"), + (S3_SESSION_TOKEN_EXPIRES_AT_MS, "soon"), + ]), + // A prefix that does not cover the path. + credential("s3://bucket/other", &[ + (S3_ACCESS_KEY_ID, "AK"), + (S3_SECRET_ACCESS_KEY, "SK"), + ]), + ] { + assert!( + adapter(credential) + .provide_credential(&Context::new()) + .await + .is_err() + ); + } + } fn parse_with(prop: Option<&str>) -> bool { let mut props = HashMap::new(); diff --git a/crates/storage/opendal/src/utils.rs b/crates/storage/opendal/src/utils.rs index d42170c71b..4aa4f92136 100644 --- a/crates/storage/opendal/src/utils.rs +++ b/crates/storage/opendal/src/utils.rs @@ -29,30 +29,67 @@ pub(crate) fn from_opendal_error(e: opendal::Error) -> iceberg::Error { .with_source(e) } -/// Convert a [`SystemTime`](std::time::SystemTime) credential expiry into the -/// `reqsign` [`Timestamp`](reqsign_core::time::Timestamp) used on backend -/// credential types (e.g. `AwsCredential::expires_in`, `google::Token::expires_at`). +/// The non-empty value of `key` in a vended credential's config. #[cfg(any( feature = "opendal-s3", feature = "opendal-gcs", feature = "opendal-azdls" ))] -pub(crate) fn system_time_to_timestamp( - time: std::time::SystemTime, -) -> reqsign_core::Result { - let millis = time - .duration_since(std::time::UNIX_EPOCH) - .map_err(|e| { - reqsign_core::Error::unexpected(format!( - "credential expiry precedes the UNIX epoch: {e}" - )) - })? - .as_millis(); - let millis = i64::try_from(millis).map_err(|_| { - reqsign_core::Error::unexpected("credential expiry overflows i64 milliseconds") - })?; - reqsign_core::time::Timestamp::from_millisecond(millis) - .map_err(|e| reqsign_core::Error::unexpected(format!("invalid credential expiry: {e}"))) +pub(crate) fn required_credential_property<'a>( + config: &'a std::collections::HashMap, + key: &str, +) -> reqsign_core::Result<&'a str> { + config + .get(key) + .map(String::as_str) + .filter(|value| !value.is_empty()) + .ok_or_else(|| { + reqsign_core::Error::unexpected(format!("vended credential is missing {key}")) + }) +} + +/// The latest epoch millisecond `reqsign` timestamps can represent, +/// `9999-12-30T22:00:00Z`. +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +const MAX_TIMESTAMP_MILLIS: i64 = 253_402_207_200_000; + +/// The epoch-millisecond expiry under `key` in a vended credential's config, +/// if present, as the `reqsign` timestamp backend credentials use. Later +/// expiries, such as `Long.MAX_VALUE` for a credential that never expires, +/// are clamped to the latest representable timestamp. +#[cfg(any( + feature = "opendal-s3", + feature = "opendal-gcs", + feature = "opendal-azdls" +))] +pub(crate) fn credential_expiry( + config: &std::collections::HashMap, + key: &str, +) -> reqsign_core::Result> { + config + .get(key) + .filter(|value| !value.is_empty()) + .map(|value| { + value + .parse::() + .ok() + .and_then(|millis| { + reqsign_core::time::Timestamp::from_millisecond( + millis.min(MAX_TIMESTAMP_MILLIS), + ) + .ok() + }) + .ok_or_else(|| { + reqsign_core::Error::unexpected(format!( + "vended credential has an invalid {key}" + )) + }) + }) + .transpose() } /// Validate that a provider's credential covers the location for which the @@ -69,13 +106,23 @@ pub(crate) fn validate_credential_prefix( if credential.covers(location) { Ok(()) } else { - Err(reqsign_core::Error::unexpected(format!( - "vended credential prefix {:?} does not cover storage location {location:?}", - credential.prefix() + Err(reqsign_core::Error::unexpected(uncovered_location_message( + location, credential, ))) } } +/// The error message for a vended credential that does not cover `location`. +pub(crate) fn uncovered_location_message( + location: &str, + credential: &iceberg::io::StorageCredential, +) -> String { + format!( + "vended credential prefix {:?} does not cover storage location {location:?}", + credential.prefix() + ) +} + /// The root of the storage location containing `path`, e.g. `s3://bucket/`. pub(crate) fn storage_root(path: &str) -> iceberg::Result { let url = url::Url::parse(path)?; @@ -109,32 +156,43 @@ impl VendedCredentialSource { Self { provider, location } } - /// Load a credential covering the location, returning its backend-specific - /// material and expiry. - pub(crate) async fn load( + /// Load a credential covering the location and convert its config with + /// `extract` into the backend's credential. + /// + /// Errors are logged here: reqsign's credential chain replaces them with an + /// error that only says no credential could be loaded, and reports the + /// cause only through the `log` crate. + pub(crate) async fn load( &self, backend: &str, - ) -> reqsign_core::Result<( - iceberg::io::StorageCredentialKind, - Option, - )> { + extract: impl FnOnce(&std::collections::HashMap) -> reqsign_core::Result, + ) -> reqsign_core::Result { + self.try_load(backend, extract).await.inspect_err(|error| { + tracing::warn!( + "cannot use vended {backend} credentials for {}: {error}", + self.location + ) + }) + } + + async fn try_load( + &self, + backend: &str, + extract: impl FnOnce(&std::collections::HashMap) -> reqsign_core::Result, + ) -> reqsign_core::Result { let credential = self .provider .load_credential(&self.location) .await .map_err(|e| { reqsign_core::Error::unexpected(format!( - "failed to load vended {backend} credential for {}", + "failed to load vended {backend} credential for {}: {e}", self.location )) .with_source(e) })?; validate_credential_prefix(&self.location, &credential)?; - let expires_at = credential - .expires_at() - .map(system_time_to_timestamp) - .transpose()?; - Ok((credential.into_kind(), expires_at)) + extract(credential.config()) } } @@ -160,11 +218,12 @@ impl std::fmt::Debug for VendedCredentialSource { ) ))] mod tests { + use std::collections::HashMap; use std::sync::{Arc, Mutex}; use async_trait::async_trait; use iceberg::io::{ - GcsCredential, StorageCredential, StorageCredentialKind, StorageCredentialProvider, + GCS_TOKEN, GCS_TOKEN_EXPIRES_AT, StorageCredential, StorageCredentialProvider, }; use super::*; @@ -172,30 +231,36 @@ mod tests { /// Returns a credential scoped to `prefix` and records requested locations. #[derive(Debug)] struct RecordingProvider { - prefix: Option<&'static str>, + prefix: &'static str, requested: Mutex>, } #[async_trait] impl StorageCredentialProvider for RecordingProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + async fn load_credential(&self, path: &str) -> iceberg::Result { self.requested.lock().unwrap().push(path.to_string()); - let credential = - StorageCredential::new(StorageCredentialKind::Gcs(GcsCredential::new("token"))); - Ok(match self.prefix { - Some(prefix) => credential.with_prefix(prefix), - None => credential, - }) + Ok(StorageCredential::new( + self.prefix, + HashMap::from([(GCS_TOKEN.to_string(), "token".to_string())]), + )) } } - async fn load(prefix: Option<&'static str>, location: &str) -> reqsign_core::Result<()> { + async fn load(prefix: &'static str, location: &str) -> reqsign_core::Result<()> { let provider = Arc::new(RecordingProvider { prefix, requested: Mutex::new(Vec::new()), }); let source = VendedCredentialSource::new(provider.clone(), location.to_string()); - let result = source.load("GCS").await.map(|_| ()); + let result = source + .load("GCS", |config| { + required_credential_property(config, GCS_TOKEN).map(|_| ()) + }) + .await; assert_eq!(*provider.requested.lock().unwrap(), vec![location]); result } @@ -203,16 +268,48 @@ mod tests { #[tokio::test] async fn vended_source_requires_a_covering_credential() { let location = "gs://bucket/table/data/file.parquet"; - assert!(load(None, location).await.is_ok()); - assert!(load(Some("gs://bucket/table"), location).await.is_ok()); - assert!(load(Some("gcs://bucket"), location).await.is_ok()); + assert!(load("gs", location).await.is_ok()); + assert!(load("gs://bucket/table", location).await.is_ok()); + assert!(load("gs://bucket/tab", location).await.is_ok()); assert!( - load(Some("gs://bucket/table/data/other"), location) + load("gs://bucket/table/data/other", location) .await .is_err() ); - assert!(load(Some("gs://bucket/tab"), location).await.is_err()); - assert!(load(Some(""), location).await.is_err()); + assert!(load("gcs://bucket", location).await.is_err()); + assert!(load("", location).await.is_err()); + } + + #[test] + fn credential_properties_are_parsed() { + let config = HashMap::from([ + (GCS_TOKEN.to_string(), "token".to_string()), + (GCS_TOKEN_EXPIRES_AT.to_string(), "1500".to_string()), + ]); + assert_eq!( + required_credential_property(&config, GCS_TOKEN).unwrap(), + "token" + ); + assert!(required_credential_property(&config, "missing").is_err()); + assert_eq!( + credential_expiry(&config, GCS_TOKEN_EXPIRES_AT).unwrap(), + Some(reqsign_core::time::Timestamp::from_millisecond(1500).unwrap()) + ); + assert_eq!(credential_expiry(&config, "missing").unwrap(), None); + // An expiry beyond year 9999, such as Java's `Long.MAX_VALUE`, is + // clamped rather than rejected. + let never = HashMap::from([(GCS_TOKEN_EXPIRES_AT.to_string(), i64::MAX.to_string())]); + assert_eq!( + credential_expiry(&never, GCS_TOKEN_EXPIRES_AT).unwrap(), + Some(reqsign_core::time::Timestamp::from_millisecond(MAX_TIMESTAMP_MILLIS).unwrap()) + ); + assert!(reqsign_core::time::Timestamp::from_millisecond(MAX_TIMESTAMP_MILLIS + 1).is_err()); + // The value is never quoted in the error. + let invalid = HashMap::from([(GCS_TOKEN_EXPIRES_AT.to_string(), "secret".to_string())]); + let error = credential_expiry(&invalid, GCS_TOKEN_EXPIRES_AT) + .unwrap_err() + .to_string(); + assert!(!error.contains("secret"), "{error}"); } #[test] @@ -227,4 +324,35 @@ mod tests { ); assert!(storage_root("not a url").is_err()); } + + #[derive(Debug)] + struct FailingProvider; + + #[async_trait] + impl StorageCredentialProvider for FailingProvider { + fn supports_path(&self, _path: &str) -> bool { + true + } + + async fn load_credential(&self, _path: &str) -> iceberg::Result { + Err(iceberg::Error::new( + iceberg::ErrorKind::Unexpected, + "catalog returned 503", + )) + } + } + + #[tokio::test] + async fn vended_source_errors_carry_the_cause() { + let source = VendedCredentialSource::new( + Arc::new(FailingProvider), + "gs://bucket/file.parquet".to_string(), + ); + let error = source + .load("GCS", |_| Ok(())) + .await + .unwrap_err() + .to_string(); + assert!(error.contains("catalog returned 503"), "{error}"); + } }