diff --git a/docs/source/about/versioning_policy.md b/docs/source/about/versioning_policy.md index e5cb832f446..286a3023748 100644 --- a/docs/source/about/versioning_policy.md +++ b/docs/source/about/versioning_policy.md @@ -209,13 +209,17 @@ here, because a vendor jar built against one Comet release is loaded by another. The SPI consists of: - `CometS3CredentialProvider`, the interface a vendor implements. +- `CometS3LocationScopedCredentialProvider`, an optional extension of it for buckets whose credentials + differ by location. - `CometS3Credentials`, the value a provider returns. - `CometS3CredentialContext` and `CometS3AccessMode`, describing the request being served. Additive changes are allowed in a minor release, for example a new accessor on `CometS3CredentialContext`, because a vendor jar compiled against an earlier `1.x` continues to load and run. Any change that would break such a jar, including adding an abstract method to -`CometS3CredentialProvider` without a default implementation, requires a major release. +`CometS3CredentialProvider` or `CometS3LocationScopedCredentialProvider` without a default +implementation, requires a major release. The same holds for changing the meaning of an existing +method, such as how `CometS3LocationScopedCredentialProvider` matches a path to a location. `CometS3CredentialDispatcher` is the JNI entry point Comet uses to reach a provider. It is internal despite living in the same package, and vendors must not call it. diff --git a/docs/source/contributor-guide/s3-credential-provider-design.md b/docs/source/contributor-guide/s3-credential-provider-design.md index 25d1951eca2..8fa449dd555 100644 --- a/docs/source/contributor-guide/s3-credential-provider-design.md +++ b/docs/source/contributor-guide/s3-credential-provider-design.md @@ -73,6 +73,23 @@ Comet's bridge does not maintain a TTL cache, schedule refresh, or broadcast cat A Comet-side cache would have to either expose a tuning knob (TTL, max size, eviction policy) and grow over time, or be hardcoded and surprise vendors whose policies disagree. The bridge intentionally has neither and forwards every call. +## Location-scoped credentials on the Parquet path + +`object_store::CredentialProvider::get_credential` receives no request path, so one `AmazonS3` store presents one credential, and the process-wide `object_store` cache holds one store per `(scheme://bucket, config_hash, hdfs_backend)`. A base provider therefore gets one credential per bucket, requested with the path of the first file Comet reads. A bucket whose policies differ by location (one for `warehouse/sales`, another for `warehouse/finance`) needs a store per location and something that picks among them for each request. + +`CometS3LocationScopedCredentialProvider` supplies that. `getPolicyLocations(bucket)` returns every location in the bucket that has its own policy, and `create_store` returns a `LocationScopedObjectStore` for the bucket instead of a plain store: + +- It is cached and registered like any other S3 store, one per key, so later scans on the bucket share it and the cache and registry keep one store per identity. As for any store, scans that miss the cache at the same moment each build one, and the last one cached wins. +- Each request is served by the store of the longest location covering its path, matched one segment at a time after percent-decoding, with the bucket root as an implicit location. Routing is per request, so a partition whose files span several locations reads each file with its own location's credential. +- A location's store is an `AmazonS3` whose bridge is bound to the location itself, so the vendor sees one stable path per location. The bridge is derived from the bucket's bridge, sharing its provider registration, so creating it calls no `ensureInitialized` and loads no classes on the Tokio worker that usually creates it. It is built on first use, usually inside an async read, from an `S3StoreTemplate` that resolved the region when the bucket's store was created, so building never blocks on the Tokio runtime. +- The locations are a snapshot. A 403 is either a real denial or a location added or removed since the snapshot, so the store fetches the locations again and retries the request once if its path now routes elsewhere. Every attempt, successful or not, starts a new snapshot generation, so requests routed from the same generation share one attempt and a failed attempt fails them all instead of each calling the provider in turn. The retry budget is per request, with no state that outlives it. + +The locations come from asking for the bucket's whole list rather than which prefixes one session covers. Asking per session leaves Comet to discover the other scopes from 403s, which needs mutable per-store state and cannot route a partition that spans scopes. With the whole list up front, routing is a function of the path. + +This is consistent with [Why no Comet-side cache](#why-no-comet-side-cache): the store caches no credentials, and every request still calls `getCredentialsForPath`. What it keeps is the provider's location list and a store for each location that has been read, so it grows with the vendor's policy list, not with the paths read. The list is only fetched again after a 403, so a change that causes none is not seen until the executor builds a new store. + +The dispatcher returns `null` for a provider that does not implement the interface without calling it, and Comet builds the same plain store as before. The Iceberg path does not use locations. Operations other than reads route by path without the 403 retry, since Comet only reads through these stores. + ## Path-specific behavior `object_store::CredentialProvider` and `reqsign_core::ProvideCredential` differ in what they consume: diff --git a/docs/source/user-guide/latest/s3-credential-providers.md b/docs/source/user-guide/latest/s3-credential-providers.md index 33777c497af..df3d76005bb 100644 --- a/docs/source/user-guide/latest/s3-credential-providers.md +++ b/docs/source/user-guide/latest/s3-credential-providers.md @@ -201,6 +201,35 @@ public CometS3Credentials getCredentialsForPath(CometS3CredentialContext ctx) th } ``` +### Credentials per location + +On the Parquet path, a `CometS3CredentialProvider` gets one credential per bucket: Comet requests it with the path of the first file it reads from the bucket and uses it for every file there. If your policies differ by location within a bucket, for example one policy for `warehouse/sales` and another for `warehouse/finance`, implement `CometS3LocationScopedCredentialProvider` and tell Comet where those locations are: + +```java +public final class MyLocationProvider implements CometS3LocationScopedCredentialProvider { + @Override + public List getPolicyLocations(String bucket) throws Exception { + return policyService.locationsWithPolicies(bucket); // ["warehouse/sales", "warehouse/finance"] + } + + @Override + public CometS3Credentials getCredentialsForPath(CometS3CredentialContext ctx) throws Exception { + // ctx.getPath() is a returned location with a leading slash, or "/" for the bucket root. + return sessionForLocation(ctx.getBucket(), ctx.getPath(), ctx.getMode()); + } +} +``` + +Comet serves each request with the credential of the longest location that covers its path. A location covers a path when the path is the location itself or lies below it, compared one `/`-separated segment at a time, so `warehouse/sales` covers `warehouse/sales/part-0.parquet` but not `warehouse/sales_eu/part-0.parquet`. The bucket root covers every path that no returned location covers, and an empty list serves the whole bucket with the root's credential. Write locations the way `CometS3CredentialContext.getPath()` writes paths: percent-encoded, without the scheme or bucket name. A literal `%` must be written as `%25`; other characters may be left unencoded, and a leading or trailing `/` is optional. When several locations decode to the same path, Comet keeps the first. + +Comet requests a location's credential by calling `getCredentialsForPath` with the location as the path, as you returned it but with a leading slash. Every request under a location shares that credential, so it must authorize every path the location is the longest match for, and your cache can key on the location. Locations apply to Comet's native Parquet reads only; Iceberg reads call `getCredentialsForPath` as they do for any provider. + +**When Comet asks.** Comet calls `getPolicyLocations` when it creates the store for a bucket on an executor and keeps the answer for later reads of that bucket with the same S3 configuration. Reads that start at the same moment may each create a store and call it. If a read then fails with 403, Comet asks again, once for all the reads that failed on the same answer, and retries each read once if its path now falls under a different location, so a location added while a job runs is picked up. A location added or removed without causing a 403 is not seen until the executor creates a new store. Make `getPolicyLocations` thread-safe and independent of where it runs; it may be called on the driver or on executors. + +**Failures.** If `getPolicyLocations` throws or returns `null`, or returns a location that is `null` or invalid, the read fails. A location is invalid if, once decoded, it is not valid UTF-8 or has a segment that is empty, `.`, `..`, or contains a control character, so a URI such as `s3://bucket/a` is invalid too. Comet does not fall back to a broader credential. + +**Backward compatibility.** Providers that implement only `CometS3CredentialProvider` are unaffected. Comet calls nothing new on them and keeps one credential per bucket, as before. + ### Composing multiple credential backends A single configured provider class is the dispatcher. If a vendor needs to route across several credential backends (per bucket, per path prefix, per tenant), the dispatch lives inside the vendor's class: diff --git a/native/core/src/cloud/s3/credential_bridge.rs b/native/core/src/cloud/s3/credential_bridge.rs index 9fc5b562cbc..415dcf2ae47 100644 --- a/native/core/src/cloud/s3/credential_bridge.rs +++ b/native/core/src/cloud/s3/credential_bridge.rs @@ -23,7 +23,7 @@ use crate::execution::operators::ExecutionError; use crate::jvm_bridge::{jni_new_global_ref, jni_static_call, JVMClasses}; use async_trait::async_trait; use iceberg_storage_opendal::AwsCredential as IcebergAwsCredential; -use jni::objects::{Global, JFieldID, JObject, JString, JValue}; +use jni::objects::{Global, JFieldID, JObject, JObjectArray, JString, JValue}; use jni::signature::{Primitive, ReturnType}; use jni::strings::JNIString; use jni::sys::jint; @@ -64,7 +64,9 @@ pub enum AccessMode { /// Granularity: although the JVM SPI accepts `(bucket, path)`, neither /// `object_store::CredentialProvider::get_credential` nor /// `reqsign_core::ProvideCredential::provide_credential` carries a per-request path, so the -/// effective identity is per-bucket (Parquet) or per-table-location (Iceberg). +/// effective identity is per-bucket (Parquet) or per-table-location (Iceberg). A Parquet provider +/// that implements `CometS3LocationScopedCredentialProvider` gets one bridge per policy location +/// instead; see `parquet::objectstore::location_scoped`. pub struct CometS3CredentialBridge { provider_class: String, dispatch_key: String, @@ -134,6 +136,31 @@ impl CometS3CredentialBridge { }) } + /// Returns a bridge to the same provider registration for another path in the bucket. It + /// shares this bridge's handle and bucket string and creates only the path string, so it makes + /// no `ensureInitialized` call and needs no class loading on the calling thread. + pub fn for_path(&self, path: impl Into) -> Result { + let path = path.into(); + let path_jstr = JVMClasses::with_env(|env| -> Result<_, ExecutionError> { + let p = env + .new_string(&path) + .map_err(|e| ExecutionError::GeneralError(format!("new_string(path): {e}")))?; + Ok(Arc::new(jni_new_global_ref!(env, p).map_err(|e| { + ExecutionError::GeneralError(format!("global_ref(path): {e}")) + })?)) + })?; + Ok(Self { + provider_class: self.provider_class.clone(), + dispatch_key: self.dispatch_key.clone(), + bucket: self.bucket.clone(), + path, + mode: self.mode, + handle: self.handle, + bucket_jstr: Arc::clone(&self.bucket_jstr), + path_jstr, + }) + } + fn fetch_raw(&self) -> Result { JVMClasses::with_env(|env| -> Result { let mode = self.mode as jint; @@ -183,6 +210,46 @@ impl CometS3CredentialBridge { }) }) } + + /// Returns the bucket's policy locations when the provider implements + /// `CometS3LocationScopedCredentialProvider`, or `None` for any other provider. The + /// dispatcher copies the provider's list into a `String[]`, so provider code, including a lazy + /// list, runs inside the checked JNI call and its exceptions come back as errors here. + pub fn policy_locations(&self) -> Result>, ExecutionError> { + JVMClasses::with_env(|env| -> Result>, ExecutionError> { + let locations: JObject = unsafe { + jni_static_call!(env, + comet_s3_credential_dispatcher.get_policy_locations( + self.handle, + self.bucket_jstr.as_obj() + ) -> JObject + )? + }; + if locations.is_null() { + return Ok(None); + } + // SAFETY: `getPolicyLocations` is declared to return `String[]`, and the dispatcher + // rejects null elements, so every element is a non-null `java.lang.String`. + let locations = unsafe { JObjectArray::::from_raw(env, locations.into_raw()) }; + let len = locations.len(env).map_err(|e| { + ExecutionError::GeneralError(format!("policy locations length: {e}")) + })?; + let mut out = Vec::with_capacity(len); + for i in 0..len { + let element = locations.get_element(env, i).map_err(|e| { + ExecutionError::GeneralError(format!("policy location {i}: {e}")) + })?; + let element = unsafe { JString::from_raw(&*env, element.into_raw()) }; + let location = element.try_to_string(env).map_err(|e| { + ExecutionError::GeneralError(format!("policy location {i}: {e}")) + })?; + // A bucket can have more locations than the local frame holds, so free each one. + env.delete_local_ref(element); + out.push(location); + } + Ok(Some(out)) + }) + } } fn ensure_initialized( diff --git a/native/core/src/parquet/objectstore/location_scoped.rs b/native/core/src/parquet/objectstore/location_scoped.rs new file mode 100644 index 00000000000..7f368615e42 --- /dev/null +++ b/native/core/src/parquet/objectstore/location_scoped.rs @@ -0,0 +1,879 @@ +// 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. + +//! The object store for a `CometS3LocationScopedCredentialProvider`. +//! +//! `object_store::CredentialProvider::get_credential` receives no request path, so one S3 store +//! presents one credential. A bucket whose policies differ by location needs a store per location +//! and something that picks the right one for each request. [`LocationScopedObjectStore`] is +//! registered once per bucket in place of a plain S3 store. It keeps the provider's policy +//! locations and serves each request with the store of the longest location that covers the +//! request's path, compared one path segment at a time. The bucket root is an implicit location +//! that covers every other path. A location's store is built the first time a request needs it and +//! kept for the life of this store. +//! +//! The locations are a snapshot, so a 403 can mean a location was added or removed after it was +//! taken. A read (`get_opts` or `get_ranges`) that gets a 403 fetches the locations again, unless +//! another read already tried since this one was routed, and retries once if its path now routes to +//! a different location; otherwise the 403 is returned. A failed fetch fails every read that shared +//! it. Other operations route by path without retrying, because Comet only reads through this store. + +use std::collections::HashMap; +use std::fmt; +use std::ops::Range; +use std::sync::{Arc, PoisonError, RwLock}; + +use async_trait::async_trait; +use bytes::Bytes; +use futures::stream::{self, BoxStream, StreamExt, TryStreamExt}; +use object_store::path::Path; +use object_store::{ + CopyOptions, Error, GetOptions, GetResult, ListResult, MultipartUpload, ObjectMeta, + ObjectStore, ObjectStoreExt, PutMultipartOptions, PutOptions, PutPayload, PutResult, + RenameOptions, Result, +}; +use tokio::sync::Mutex; + +const STORE: &str = "LocationScopedS3"; + +/// The credential path used for paths that no returned location covers. +const ROOT_CREDENTIAL_PATH: &str = "/"; + +/// Fetches the provider's current policy locations for the bucket. +pub(crate) type LocationSource = Arc Result> + Send + Sync>; + +/// Builds the store for one location from the path passed to `getCredentialsForPath`. +pub(crate) type LocationStoreFactory = + Arc Result> + Send + Sync>; + +/// One snapshot of the provider's locations. +struct LocationIndex { + /// Counts refresh attempts, including failed ones. + generation: u64, + /// Canonical location to credential path: the location as the provider returned it, with a + /// leading slash. + locations: Arc>, + /// Why the attempt that produced this snapshot failed, when it kept the previous locations. + failure: Option>, +} + +impl LocationIndex { + fn new(generation: u64, locations: Vec) -> Result { + let mut index = HashMap::with_capacity(locations.len()); + for location in locations { + // Request paths are percent-decoded the same way, so both sides compare as raw keys. + let canonical = Path::from_url_path(&location).map_err(|e| Error::Generic { + store: STORE, + source: format!("Invalid policy location {location:?}: {e}").into(), + })?; + // A duplicate keeps the first spelling, which is the path the provider is given. + index + .entry(canonical) + .or_insert_with(|| credential_path(&location)); + } + Ok(Self { + generation, + locations: Arc::new(index), + failure: None, + }) + } + + /// The snapshot after a failed refresh: the same locations under the next generation, with the + /// failure recorded for the requests routed before it. + fn after_failure(&self, failure: &Error) -> Self { + Self { + generation: self.generation + 1, + locations: Arc::clone(&self.locations), + failure: Some(failure.to_string().into()), + } + } + + /// Returns the credential path of the longest location that covers `path`. + fn route(&self, path: &Path) -> &str { + let mut longest = self + .locations + .get(&Path::default()) + .map_or(ROOT_CREDENTIAL_PATH, String::as_str); + let mut prefix = Path::default(); + for part in path.parts() { + prefix = prefix.join(part); + if let Some(credential_path) = self.locations.get(&prefix) { + longest = credential_path; + } + } + longest + } +} + +fn credential_path(location: &str) -> String { + if location.starts_with('/') { + location.to_string() + } else { + format!("/{location}") + } +} + +fn is_forbidden(err: &Error) -> bool { + matches!(err, Error::PermissionDenied { .. }) +} + +/// The store chosen for one request, and the snapshot it was chosen from. +struct Route { + generation: u64, + credential_path: String, + store: Arc, +} + +struct Inner { + bucket: String, + source: LocationSource, + factory: LocationStoreFactory, + index: RwLock>, + /// Serializes refreshes, so the 403s from one snapshot share one attempt. + refresh_lock: Mutex<()>, + /// Location stores by credential path. + stores: RwLock>>, +} + +impl Inner { + fn index(&self) -> Arc { + Arc::clone(&self.index.read().unwrap_or_else(PoisonError::into_inner)) + } + + fn route(&self, path: &Path) -> Result { + let index = self.index(); + let credential_path = index.route(path).to_string(); + let store = self.store(&credential_path)?; + Ok(Route { + generation: index.generation, + credential_path, + store, + }) + } + + fn store(&self, credential_path: &str) -> Result> { + if let Some(store) = self + .stores + .read() + .unwrap_or_else(PoisonError::into_inner) + .get(credential_path) + { + return Ok(Arc::clone(store)); + } + // Build outside the lock, since building creates a bridge through JNI. When two requests + // race to build the same location, the first insert wins and the other store is dropped. + let store = (self.factory)(credential_path)?; + let mut stores = self.stores.write().unwrap_or_else(PoisonError::into_inner); + Ok(Arc::clone( + stores.entry(credential_path.to_string()).or_insert(store), + )) + } + + /// Returns the route to retry on after `failed` returned the 403 `err` for `path`, or the error + /// the request should return. + async fn retry_route(&self, path: &Path, failed: &Route, err: Error) -> Result { + if let Err(refresh) = self.refresh(failed).await { + return Err(Error::Generic { + store: STORE, + source: format!( + "{err}; fetching the policy locations for bucket {} again failed: {refresh}", + self.bucket + ) + .into(), + }); + } + let route = self.route(path)?; + if route.credential_path == failed.credential_path { + return Err(err); + } + Ok(route) + } + + /// Fetches the locations again, unless a refresh was attempted after `failed` was routed, in + /// which case this request shares that attempt's outcome. + async fn refresh(&self, failed: &Route) -> Result<()> { + let _refresh = self.refresh_lock.lock().await; + let current = self.index(); + if current.generation != failed.generation { + return match ¤t.failure { + Some(failure) => Err(Error::Generic { + store: STORE, + source: failure.to_string().into(), + }), + None => Ok(()), + }; + } + let next = (self.source)() + .and_then(|locations| LocationIndex::new(current.generation + 1, locations)); + let (index, result) = match next { + Ok(index) => (index, Ok(())), + Err(e) => (current.after_failure(&e), Err(e)), + }; + *self.index.write().unwrap_or_else(PoisonError::into_inner) = Arc::new(index); + result + } +} + +/// Serves each request with the store of the longest policy location covering its path. See the +/// module documentation. +pub struct LocationScopedObjectStore { + inner: Arc, +} + +impl LocationScopedObjectStore { + /// `locations` is the provider's first answer for `bucket`; `source` fetches it again after a + /// 403, and `factory` builds a location's store on first use. Fails if a location is not a + /// valid path. + pub(crate) fn new( + bucket: String, + locations: Vec, + source: LocationSource, + factory: LocationStoreFactory, + ) -> Result { + let index = LocationIndex::new(0, locations)?; + Ok(Self { + inner: Arc::new(Inner { + bucket, + source, + factory, + index: RwLock::new(Arc::new(index)), + refresh_lock: Mutex::new(()), + stores: RwLock::new(HashMap::new()), + }), + }) + } +} + +impl fmt::Debug for LocationScopedObjectStore { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("LocationScopedObjectStore") + .field("bucket", &self.inner.bucket) + .field("locations", &self.inner.index().locations.len()) + .finish() + } +} + +impl fmt::Display for LocationScopedObjectStore { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "LocationScopedS3({})", self.inner.bucket) + } +} + +#[async_trait] +impl ObjectStore for LocationScopedObjectStore { + async fn put_opts( + &self, + location: &Path, + payload: PutPayload, + opts: PutOptions, + ) -> Result { + let route = self.inner.route(location)?; + route.store.put_opts(location, payload, opts).await + } + + async fn put_multipart_opts( + &self, + location: &Path, + opts: PutMultipartOptions, + ) -> Result> { + let route = self.inner.route(location)?; + route.store.put_multipart_opts(location, opts).await + } + + async fn get_opts(&self, location: &Path, options: GetOptions) -> Result { + let route = self.inner.route(location)?; + match route.store.get_opts(location, options.clone()).await { + Err(e) if is_forbidden(&e) => { + let retry = self.inner.retry_route(location, &route, e).await?; + retry.store.get_opts(location, options).await + } + other => other, + } + } + + async fn get_ranges(&self, location: &Path, ranges: &[Range]) -> Result> { + let route = self.inner.route(location)?; + match route.store.get_ranges(location, ranges).await { + Err(e) if is_forbidden(&e) => { + let retry = self.inner.retry_route(location, &route, e).await?; + retry.store.get_ranges(location, ranges).await + } + other => other, + } + } + + fn delete_stream( + &self, + locations: BoxStream<'static, Result>, + ) -> BoxStream<'static, Result> { + let inner = Arc::clone(&self.inner); + locations + .and_then(move |location| { + let route = inner.route(&location); + async move { + route?.store.delete(&location).await?; + Ok(location) + } + }) + .boxed() + } + + fn list(&self, prefix: Option<&Path>) -> BoxStream<'static, Result> { + match self.inner.route(prefix.unwrap_or(&Path::default())) { + Ok(route) => route.store.list(prefix), + Err(e) => stream::once(async move { Err(e) }).boxed(), + } + } + + fn list_with_offset( + &self, + prefix: Option<&Path>, + offset: &Path, + ) -> BoxStream<'static, Result> { + match self.inner.route(prefix.unwrap_or(&Path::default())) { + Ok(route) => route.store.list_with_offset(prefix, offset), + Err(e) => stream::once(async move { Err(e) }).boxed(), + } + } + + async fn list_with_delimiter(&self, prefix: Option<&Path>) -> Result { + let route = self.inner.route(prefix.unwrap_or(&Path::default()))?; + route.store.list_with_delimiter(prefix).await + } + + async fn copy_opts(&self, from: &Path, to: &Path, options: CopyOptions) -> Result<()> { + let route = self.inner.route(from)?; + route.store.copy_opts(from, to, options).await + } + + async fn rename_opts(&self, from: &Path, to: &Path, options: RenameOptions) -> Result<()> { + let route = self.inner.route(from)?; + route.store.rename_opts(from, to, options).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; + use std::sync::Mutex as StdMutex; + use std::time::Duration; + use tokio::sync::Barrier; + + /// What S3 allows for one credential: reads under `allowed` succeed and every other read gets a + /// 403. A success is reported as `NotFound` carrying the credential path, which shows which + /// location served the request without producing data. + #[derive(Debug)] + struct CredentialView { + credential_path: String, + allowed: Vec, + gets: AtomicUsize, + /// When set, each read waits here first, so a test can hold requests in flight together. + gate: Option>, + } + + impl fmt::Display for CredentialView { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "CredentialView({})", self.credential_path) + } + } + + #[async_trait] + impl ObjectStore for CredentialView { + async fn put_opts( + &self, + _location: &Path, + _payload: PutPayload, + _opts: PutOptions, + ) -> Result { + unimplemented!("reads only") + } + + async fn put_multipart_opts( + &self, + _location: &Path, + _opts: PutMultipartOptions, + ) -> Result> { + unimplemented!("reads only") + } + + async fn get_opts(&self, location: &Path, _options: GetOptions) -> Result { + self.gets.fetch_add(1, Ordering::SeqCst); + if let Some(gate) = &self.gate { + gate.wait().await; + } + let source = self.credential_path.clone().into(); + let path = location.to_string(); + if self.allowed.iter().any(|p| location.prefix_matches(p)) { + Err(Error::NotFound { path, source }) + } else { + Err(Error::PermissionDenied { path, source }) + } + } + + fn delete_stream( + &self, + _locations: BoxStream<'static, Result>, + ) -> BoxStream<'static, Result> { + unimplemented!("reads only") + } + + fn list(&self, _prefix: Option<&Path>) -> BoxStream<'static, Result> { + unimplemented!("reads only") + } + + async fn list_with_delimiter(&self, _prefix: Option<&Path>) -> Result { + unimplemented!("reads only") + } + + async fn copy_opts(&self, _from: &Path, _to: &Path, _options: CopyOptions) -> Result<()> { + unimplemented!("reads only") + } + } + + /// A provider with changeable locations whose credentials read exactly the prefixes in + /// `grants`. + struct Provider { + locations: StdMutex>, + grants: HashMap>, + fail_refresh: AtomicBool, + /// How long fetching the locations takes, so concurrent refreshes overlap. + refresh_delay: Option, + refreshes: AtomicUsize, + /// Credential path whose reads wait at a shared barrier. + gate: Option<(String, Arc)>, + /// Credential path whose store cannot be built. + fail_build: Option, + views: StdMutex>>, + builds: AtomicUsize, + } + + impl Provider { + fn new(locations: &[&str], grants: &[(&str, &[&str])]) -> Self { + Self { + locations: StdMutex::new(locations.iter().map(|l| l.to_string()).collect()), + grants: grants + .iter() + .map(|(credential_path, prefixes)| { + let prefixes = prefixes.iter().map(|p| Path::from(*p)).collect(); + (credential_path.to_string(), prefixes) + }) + .collect(), + fail_refresh: AtomicBool::new(false), + refresh_delay: None, + refreshes: AtomicUsize::new(0), + gate: None, + fail_build: None, + views: StdMutex::new(HashMap::new()), + builds: AtomicUsize::new(0), + } + } + + fn with_gate(mut self, credential_path: &str, parties: usize) -> Self { + self.gate = Some((credential_path.to_string(), Arc::new(Barrier::new(parties)))); + self + } + + fn with_slow_refresh(mut self, delay: Duration) -> Self { + self.refresh_delay = Some(delay); + self + } + + fn with_failed_build(mut self, credential_path: &str) -> Self { + self.fail_build = Some(credential_path.to_string()); + self + } + + fn set_locations(&self, locations: &[&str]) { + *self.locations.lock().unwrap() = locations.iter().map(|l| l.to_string()).collect(); + } + + fn store(self: &Arc) -> LocationScopedObjectStore { + let provider = Arc::clone(self); + let source: LocationSource = Arc::new(move || { + provider.refreshes.fetch_add(1, Ordering::SeqCst); + if let Some(delay) = provider.refresh_delay { + std::thread::sleep(delay); + } + if provider.fail_refresh.load(Ordering::SeqCst) { + return Err(Error::Generic { + store: "test", + source: "policy service unavailable".into(), + }); + } + Ok(provider.locations.lock().unwrap().clone()) + }); + let provider = Arc::clone(self); + let factory: LocationStoreFactory = Arc::new(move |credential_path: &str| { + provider.builds.fetch_add(1, Ordering::SeqCst); + if provider.fail_build.as_deref() == Some(credential_path) { + return Err(Error::Generic { + store: "test", + source: "bridge init failed".into(), + }); + } + let gate = provider + .gate + .as_ref() + .filter(|(gated, _)| gated == credential_path) + .map(|(_, barrier)| Arc::clone(barrier)); + let view = Arc::new(CredentialView { + credential_path: credential_path.to_string(), + allowed: provider + .grants + .get(credential_path) + .cloned() + .unwrap_or_default(), + gets: AtomicUsize::new(0), + gate, + }); + provider + .views + .lock() + .unwrap() + .insert(credential_path.to_string(), Arc::clone(&view)); + Ok(view as Arc) + }); + let locations = self.locations.lock().unwrap().clone(); + LocationScopedObjectStore::new("bucket".to_string(), locations, source, factory) + .unwrap() + } + + fn gets(&self, credential_path: &str) -> usize { + self.views.lock().unwrap()[credential_path] + .gets + .load(Ordering::SeqCst) + } + + fn built(&self) -> Vec { + let mut built: Vec = self.views.lock().unwrap().keys().cloned().collect(); + built.sort(); + built + } + } + + /// Returns the credential path that served a read, or panics if the read failed. + fn served_by(result: Result) -> String { + match result { + Err(Error::NotFound { source, .. }) => source.to_string(), + Err(e) => panic!("expected the read to be served, got {e}"), + Ok(_) => panic!("test stores never return data"), + } + } + + async fn get(store: &LocationScopedObjectStore, path: &str) -> Result { + store + .get_opts(&Path::from(path), GetOptions::default()) + .await + } + + /// The two ranges are close enough to be coalesced into one request. + async fn get_ranges(store: &LocationScopedObjectStore, path: &str) -> Result> { + store.get_ranges(&Path::from(path), &[0..4, 8..12]).await + } + + #[test] + fn routes_to_the_longest_covering_location() { + let index = LocationIndex::new( + 0, + vec![ + "warehouse/sales".into(), + "warehouse/sales/eu/".into(), + "/warehouse/finance".into(), + ], + ) + .unwrap(); + let route = |path: &str| index.route(&Path::from(path)).to_string(); + assert_eq!(route("warehouse/sales/a.parquet"), "/warehouse/sales"); + assert_eq!( + route("warehouse/sales/eu/b.parquet"), + "/warehouse/sales/eu/" + ); + assert_eq!(route("warehouse/sales"), "/warehouse/sales"); + assert_eq!(route("warehouse/finance/c.parquet"), "/warehouse/finance"); + // Locations match whole segments, so a sibling that shares a name prefix is not covered. + assert_eq!(route("warehouse/sales_eu/d.parquet"), "/"); + assert_eq!(route("warehouse/e.parquet"), "/"); + assert_eq!(route("other/f.parquet"), "/"); + } + + #[test] + fn compares_locations_and_paths_percent_decoded() { + // A URI path escapes '%' as %25, so Spark's %3A partition escape arrives as %253A. + let index = LocationIndex::new(0, vec!["tbl/ts=2024-01-01%2000%253A00".into()]).unwrap(); + let spark_path = Path::from_url_path("/tbl/ts=2024-01-01%2000%253A00/part-0.parquet"); + assert_eq!( + index.route(&spark_path.unwrap()), + "/tbl/ts=2024-01-01%2000%253A00", + "the provider is given the location as it was returned" + ); + // Decoded once, %3A is ':', which names a different key. + let other_key = Path::from_url_path("/tbl/ts=2024-01-01%2000%3A00/part-0.parquet"); + assert_eq!(index.route(&other_key.unwrap()), "/"); + } + + #[test] + fn keeps_the_first_spelling_of_a_duplicate_location() { + let index = LocationIndex::new(0, vec!["a/b".into(), "/a/b/".into()]).unwrap(); + assert_eq!(index.locations.len(), 1); + assert_eq!(index.route(&Path::from("a/b/c")), "/a/b"); + } + + #[test] + fn accepts_the_bucket_root_as_a_location() { + let index = LocationIndex::new(0, vec!["".into(), "a".into()]).unwrap(); + assert!(index.locations.contains_key(&Path::default())); + assert_eq!(index.route(&Path::from("x/y")), "/"); + assert_eq!(index.route(&Path::from("a/y")), "/a"); + } + + #[test] + fn rejects_locations_that_are_not_bucket_paths() { + // Empty and relative segments, a URI (its "//" is an empty segment), a control character + // once decoded, and bytes that do not decode to UTF-8. + for location in ["a//b", "a/../b", "s3://bucket/a", "a/%0Ab", "a/%FF"] { + assert!( + LocationIndex::new(0, vec![location.into()]).is_err(), + "{location:?} should be rejected" + ); + } + } + + /// Several locations can be read through one store, in any order, which is what a scan does + /// when one partition holds files from several locations. + #[tokio::test] + async fn routes_each_read_to_its_own_location() { + let provider = Arc::new(Provider::new( + &["a", "b", "c"], + &[("/a", &["a"]), ("/b", &["b"]), ("/c", &["c"])], + )); + let store = provider.store(); + assert!(provider.built().is_empty(), "stores are built on first use"); + + for (path, location) in [("a/1", "/a"), ("b/1", "/b"), ("a/2", "/a")] { + assert_eq!(served_by(get(&store, path).await), location); + } + assert_eq!(provider.built(), ["/a", "/b"]); + for (path, location) in [("c/1", "/c"), ("b/2", "/b"), ("a/3", "/a")] { + assert_eq!(served_by(get_ranges(&store, path).await), location); + } + assert_eq!(provider.built(), ["/a", "/b", "/c"]); + assert_eq!(provider.builds.load(Ordering::SeqCst), 3); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 0); + } + + #[tokio::test] + async fn returns_a_403_when_the_locations_are_unchanged() { + let provider = Arc::new(Provider::new(&["a"], &[("/a", &["a/public"])])); + let store = provider.store(); + + let err = get(&store, "a/private/1").await.unwrap_err(); + assert!(matches!(err, Error::PermissionDenied { .. }), "got {err}"); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + assert_eq!(provider.gets("/a"), 1, "not retried on the same location"); + } + + #[tokio::test] + async fn retries_on_a_location_added_after_the_snapshot() { + let provider = Arc::new(Provider::new( + &["warehouse"], + &[ + ("/warehouse", &["warehouse/sales"]), + ("/warehouse/finance", &["warehouse/finance"]), + ], + )); + let store = provider.store(); + provider.set_locations(&["warehouse", "warehouse/finance"]); + + assert_eq!( + served_by(get(&store, "warehouse/finance/1").await), + "/warehouse/finance" + ); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + assert_eq!( + served_by(get(&store, "warehouse/finance/2").await), + "/warehouse/finance" + ); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + assert_eq!(provider.gets("/warehouse"), 1); + } + + /// Two reads that get a 403 from the same snapshot fetch the locations once, and both retry on + /// the new location instead of one of them returning its 403. + async fn concurrent_403s_share_one_refresh(use_ranges: bool) { + let provider = Arc::new( + Provider::new( + &["warehouse"], + &[ + ("/warehouse", &["warehouse/sales"]), + ("/warehouse/finance", &["warehouse/finance"]), + ], + ) + .with_gate("/warehouse", 2), + ); + let store = provider.store(); + provider.set_locations(&["warehouse", "warehouse/finance"]); + + let (first, second) = if use_ranges { + let (first, second) = tokio::join!( + get_ranges(&store, "warehouse/finance/1"), + get_ranges(&store, "warehouse/finance/2") + ); + (served_by(first), served_by(second)) + } else { + let (first, second) = tokio::join!( + get(&store, "warehouse/finance/1"), + get(&store, "warehouse/finance/2") + ); + (served_by(first), served_by(second)) + }; + assert_eq!(first, "/warehouse/finance"); + assert_eq!(second, "/warehouse/finance"); + assert_eq!(provider.gets("/warehouse"), 2, "both reads were in flight"); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn concurrent_get_opts_403s_share_one_refresh() { + concurrent_403s_share_one_refresh(false).await; + } + + #[tokio::test] + async fn concurrent_get_ranges_403s_share_one_refresh() { + concurrent_403s_share_one_refresh(true).await; + } + + #[tokio::test] + async fn fails_the_read_when_fetching_the_locations_again_fails() { + let provider = Arc::new(Provider::new(&["a"], &[("/a", &["a/public"])])); + let store = provider.store(); + provider.fail_refresh.store(true, Ordering::SeqCst); + + let err = get(&store, "a/private/1").await.unwrap_err(); + let message = err.to_string(); + assert!(matches!(err, Error::Generic { .. }), "got {message}"); + assert!(message.contains("policy service unavailable"), "{message}"); + assert!(message.contains("a/private/1"), "{message}"); + } + + /// Reads routed before a failed refresh share its failure instead of each asking the provider + /// in turn. + #[tokio::test] + async fn concurrent_403s_share_one_failed_refresh() { + let provider = Arc::new(Provider::new(&["a"], &[("/a", &["a/public"])]).with_gate("/a", 2)); + let store = provider.store(); + provider.fail_refresh.store(true, Ordering::SeqCst); + + let (first, second) = tokio::join!(get(&store, "a/private/1"), get(&store, "a/private/2")); + for result in [first, second] { + let message = result.unwrap_err().to_string(); + assert!(message.contains("policy service unavailable"), "{message}"); + } + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + } + + /// A read routed after a failed refresh asks the provider again. + #[tokio::test] + async fn asks_again_after_a_failed_refresh() { + let provider = Arc::new(Provider::new(&["a"], &[("/a", &["a/public"])])); + let store = provider.store(); + provider.fail_refresh.store(true, Ordering::SeqCst); + assert!(get(&store, "a/private/1").await.is_err()); + + provider.fail_refresh.store(false, Ordering::SeqCst); + let err = get(&store, "a/private/2").await.unwrap_err(); + assert!(matches!(err, Error::PermissionDenied { .. }), "got {err}"); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 2); + } + + /// Reads that get a 403 while another read is fetching the locations wait for that fetch + /// instead of starting their own. Each read runs on its own thread and runtime, so the two + /// refreshes overlap; tasks on one multi-thread runtime can end up serialized. + #[test] + fn waits_for_a_refresh_in_progress() { + let provider = Arc::new( + Provider::new( + &["warehouse"], + &[ + ("/warehouse", &["warehouse/sales"]), + ("/warehouse/finance", &["warehouse/finance"]), + ], + ) + .with_gate("/warehouse", 2) + .with_slow_refresh(Duration::from_millis(100)), + ); + let store = Arc::new(provider.store()); + provider.set_locations(&["warehouse", "warehouse/finance"]); + + let reads: Vec<_> = ["warehouse/finance/1", "warehouse/finance/2"] + .into_iter() + .map(|path| { + let store = Arc::clone(&store); + std::thread::spawn(move || { + let runtime = tokio::runtime::Builder::new_current_thread() + .build() + .unwrap(); + served_by(runtime.block_on(get(&store, path))) + }) + }) + .collect(); + for read in reads { + assert_eq!(read.join().unwrap(), "/warehouse/finance"); + } + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + } + + /// The retry gets one attempt: if the new location's credential is denied too, that 403 is + /// returned without another refresh. + #[tokio::test] + async fn returns_the_retry_403_without_retrying_again() { + let provider = Arc::new(Provider::new( + &["warehouse"], + &[ + ("/warehouse", &["warehouse/sales"]), + ("/warehouse/finance", &[]), + ], + )); + let store = provider.store(); + provider.set_locations(&["warehouse", "warehouse/finance"]); + + let err = get(&store, "warehouse/finance/1").await.unwrap_err(); + assert!(matches!(err, Error::PermissionDenied { .. }), "got {err}"); + assert_eq!(provider.refreshes.load(Ordering::SeqCst), 1); + assert_eq!(provider.gets("/warehouse"), 1); + assert_eq!(provider.gets("/warehouse/finance"), 1); + } + + /// A store that cannot be built after a successful refresh reports its own error, not a failed + /// refresh. + #[tokio::test] + async fn reports_a_failed_build_after_a_refresh_as_itself() { + let provider = Arc::new( + Provider::new(&["warehouse"], &[("/warehouse", &["warehouse/sales"])]) + .with_failed_build("/warehouse/finance"), + ); + let store = provider.store(); + provider.set_locations(&["warehouse", "warehouse/finance"]); + + let message = get(&store, "warehouse/finance/1") + .await + .unwrap_err() + .to_string(); + assert!(message.contains("bridge init failed"), "{message}"); + assert!(!message.contains("policy locations"), "{message}"); + } +} diff --git a/native/core/src/parquet/objectstore/mod.rs b/native/core/src/parquet/objectstore/mod.rs index 1bf42156492..91552a199b1 100644 --- a/native/core/src/parquet/objectstore/mod.rs +++ b/native/core/src/parquet/objectstore/mod.rs @@ -16,5 +16,6 @@ // under the License. pub mod azure; +pub mod location_scoped; pub mod s3; pub mod s3_blob_fs_support; diff --git a/native/core/src/parquet/objectstore/s3.rs b/native/core/src/parquet/objectstore/s3.rs index c10fc60816e..620bbdbdda5 100644 --- a/native/core/src/parquet/objectstore/s3.rs +++ b/native/core/src/parquet/objectstore/s3.rs @@ -22,6 +22,9 @@ use url::Url; use crate::cloud::s3::credential_bridge::{AccessMode, CometS3CredentialBridge}; use crate::execution::jni_api::get_runtime; +use crate::parquet::objectstore::location_scoped::{ + LocationScopedObjectStore, LocationSource, LocationStoreFactory, +}; use async_trait::async_trait; use aws_config::{ ecs::EcsCredentialsProvider, environment::EnvironmentVariableCredentialsProvider, @@ -35,7 +38,7 @@ use aws_credential_types::{ Credentials, }; use object_store::{ - aws::{AmazonS3Builder, AmazonS3ConfigKey, AwsCredential}, + aws::{AmazonS3, AmazonS3Builder, AmazonS3ConfigKey, AwsCredential, AwsCredentialProvider}, path::Path, CredentialProvider, ObjectStore, ObjectStoreScheme, }; @@ -47,6 +50,10 @@ use std::{ /// Creates an S3 object store using options specified as Hadoop S3A configurations. /// +/// When the configured `CometS3CredentialProvider` implements +/// `CometS3LocationScopedCredentialProvider`, the store is a [`LocationScopedObjectStore`] that +/// serves each of the provider's policy locations with its own credential. +/// /// # Arguments /// /// * `url` - The URL of the S3 object to access. @@ -71,9 +78,6 @@ pub fn create_store( } let path = Path::parse(path)?; - let mut builder = AmazonS3Builder::new() - .with_url(url.to_string()) - .with_allow_http(true); let bucket = url.host_str().ok_or_else(|| object_store::Error::Generic { store: "S3", source: "Missing bucket name in S3 URL".into(), @@ -81,7 +85,7 @@ pub fn create_store( // Parquet path: catalog_properties is empty; vendors here read from Hadoop conf. let empty_props: HashMap = HashMap::new(); - builder = match lookup_provider_class(configs, bucket) { + let credentials = match lookup_provider_class(configs, bucket) { Some(provider_class) => { // Fail rather than fall back to the default chain, which could resolve to the wrong // identity for a user who explicitly named a provider. @@ -97,48 +101,153 @@ pub fn create_store( store: "S3", source: format!("CometS3CredentialBridge init failed for {bucket}: {e}").into(), })?; - builder.with_credentials(Arc::new(bridge)) + let locations = + bridge + .policy_locations() + .map_err(|e| object_store::Error::Generic { + store: "S3", + source: format!("Failed to get policy locations for {bucket}: {e}").into(), + })?; + if let Some(locations) = locations { + let template = S3StoreTemplate::new(url, configs, bucket)?; + let store = location_scoped_store(template, bucket, bridge, locations)?; + return Ok((Box::new(store), path)); + } + S3Credentials::Provider(Arc::new(bridge)) } None => { match get_runtime().block_on(build_credential_provider(configs, bucket, min_ttl))? { - Some(provider) => builder.with_credentials(Arc::new(provider)), - None => builder.with_skip_signature(true), + Some(provider) => S3Credentials::Provider(Arc::new(provider)), + None => S3Credentials::SkipSignature, } } }; - let s3_configs = extract_s3_config_options(configs, bucket); - debug!("S3 configs for bucket {bucket}: {s3_configs:?}"); + let object_store = S3StoreTemplate::new(url, configs, bucket)?.build(credentials)?; - // When using the default AWS S3 endpoint (no custom endpoint configured), a valid region - // is required. If no region is explicitly configured, attempt to auto-resolve it by - // making a HeadBucket request to determine the bucket's region. - if !s3_configs.contains_key(&AmazonS3ConfigKey::Endpoint) - && !s3_configs.contains_key(&AmazonS3ConfigKey::Region) - { - let region = get_runtime() - .block_on(resolve_bucket_region(bucket)) - .map_err(|e| object_store::Error::Generic { - store: "S3", - source: format!( - "Failed to resolve region: {e}. If '{bucket}' is on a non-AWS S3-compatible \ - service, set fs.s3a.endpoint (and optionally fs.s3a.endpoint.region, \ - fs.s3a.path.style.access) or the per-bucket variants \ - fs.s3a.bucket.{bucket}.endpoint[.region] so Comet skips the AWS HEAD probe." - ) - .into(), - })?; - debug!("resolved region: {region:?}"); - builder = builder.with_config(AmazonS3ConfigKey::Region, region.to_string()); - } + Ok((Box::new(object_store), path)) +} + +/// How a store built from an [`S3StoreTemplate`] signs its requests. +enum S3Credentials { + Provider(AwsCredentialProvider), + SkipSignature, +} - for (key, value) in s3_configs { - builder = builder.with_config(key, value); +/// Builder settings shared by every store for one bucket. Creating a template may block on a +/// region lookup; building a store from it does not, so location-scoped stores can be built from +/// async code on a Tokio worker. +struct S3StoreTemplate { + url: String, + region: Option, + s3_configs: HashMap, +} + +impl S3StoreTemplate { + fn new( + url: &Url, + configs: &HashMap, + bucket: &str, + ) -> Result { + let s3_configs = extract_s3_config_options(configs, bucket); + debug!("S3 configs for bucket {bucket}: {s3_configs:?}"); + + // When using the default AWS S3 endpoint (no custom endpoint configured), a valid region + // is required. If no region is explicitly configured, attempt to auto-resolve it by + // making a HeadBucket request to determine the bucket's region. + let region = if !s3_configs.contains_key(&AmazonS3ConfigKey::Endpoint) + && !s3_configs.contains_key(&AmazonS3ConfigKey::Region) + { + let region = get_runtime() + .block_on(resolve_bucket_region(bucket)) + .map_err(|e| object_store::Error::Generic { + store: "S3", + source: format!( + "Failed to resolve region: {e}. If '{bucket}' is on a non-AWS S3-compatible \ + service, set fs.s3a.endpoint (and optionally fs.s3a.endpoint.region, \ + fs.s3a.path.style.access) or the per-bucket variants \ + fs.s3a.bucket.{bucket}.endpoint[.region] so Comet skips the AWS HEAD probe." + ) + .into(), + })?; + debug!("resolved region: {region:?}"); + Some(region) + } else { + None + }; + + Ok(Self { + url: url.to_string(), + region, + s3_configs, + }) } - let object_store = builder.build()?; + fn build(&self, credentials: S3Credentials) -> Result { + let builder = AmazonS3Builder::new() + .with_url(self.url.clone()) + .with_allow_http(true); + let mut builder = match credentials { + S3Credentials::Provider(provider) => builder.with_credentials(provider), + S3Credentials::SkipSignature => builder.with_skip_signature(true), + }; + if let Some(region) = &self.region { + builder = builder.with_config(AmazonS3ConfigKey::Region, region.clone()); + } + for (key, value) in &self.s3_configs { + builder = builder.with_config(*key, value.clone()); + } + builder.build() + } +} - Ok((Box::new(object_store), path)) +/// Builds the store for a `CometS3LocationScopedCredentialProvider`. `bridge` was created on this +/// thread, which registered the provider. It fetches the locations again after a 403, and each +/// location's bridge is derived from it on first use, often on a Tokio worker, so every location +/// shares the bucket's provider registration without another `ensureInitialized` call. +fn location_scoped_store( + template: S3StoreTemplate, + bucket: &str, + bridge: CometS3CredentialBridge, + locations: Vec, +) -> Result { + let bridge = Arc::new(bridge); + + let source_bridge = Arc::clone(&bridge); + let source_bucket = bucket.to_string(); + let source: LocationSource = Arc::new(move || { + let locations = + source_bridge + .policy_locations() + .map_err(|e| object_store::Error::Generic { + store: "S3", + source: format!("Failed to get policy locations for {source_bucket}: {e}") + .into(), + })?; + locations.ok_or_else(|| object_store::Error::Generic { + store: "S3", + source: format!("The provider for {source_bucket} stopped returning policy locations") + .into(), + }) + }); + + let factory_bucket = bucket.to_string(); + let factory: LocationStoreFactory = Arc::new(move |credential_path: &str| { + let location_bridge = + bridge + .for_path(credential_path) + .map_err(|e| object_store::Error::Generic { + store: "S3", + source: format!( + "CometS3CredentialBridge init failed for {factory_bucket}: {e}" + ) + .into(), + })?; + let store = template.build(S3Credentials::Provider(Arc::new(location_bridge)))?; + Ok(Arc::new(store) as Arc) + }); + + LocationScopedObjectStore::new(bucket.to_string(), locations, source, factory) } /// Process-wide cache of resolved S3 bucket regions, keyed by bucket name. @@ -995,6 +1104,26 @@ mod tests { ); } + /// A location-scoped store builds each location's store on first use, usually inside an async + /// read on a Tokio worker, so building from a template must not block on the runtime. The + /// template resolves the region when it is created; this bucket's region is already cached, so + /// no request is made. + #[test] + fn builds_from_a_template_inside_the_runtime() { + let bucket = "comet-template-test-bucket"; + region_cache() + .write() + .unwrap() + .insert(bucket.to_string(), "us-west-2".to_string()); + let url = Url::parse(&format!("s3a://{bucket}/warehouse/sales/part-0.parquet")).unwrap(); + // With no endpoint or region configured, creating the template resolves the region. + let template = S3StoreTemplate::new(&url, &HashMap::new(), bucket).unwrap(); + assert_eq!(template.region.as_deref(), Some("us-west-2")); + + let store = get_runtime().block_on(async { template.build(S3Credentials::SkipSignature) }); + assert!(store.is_ok(), "{:?}", store.err()); + } + #[test] fn test_get_config_trimmed() { let configs = TestConfigBuilder::new() diff --git a/native/core/src/parquet/parquet_support.rs b/native/core/src/parquet/parquet_support.rs index 1964f531746..42661df1f5f 100644 --- a/native/core/src/parquet/parquet_support.rs +++ b/native/core/src/parquet/parquet_support.rs @@ -726,6 +726,11 @@ type ObjectStoreCache = RwLock /// connection pool + DNS resolution), and there is no meaningful benefit from eviction, so /// no eviction policy is applied. /// +/// A provider that implements `CometS3LocationScopedCredentialProvider` gets one entry per +/// bucket as well: a `LocationScopedObjectStore` that holds an S3 store for each of the +/// provider's locations that has been read, so it grows with the provider's location list, not +/// with the number of paths read. +/// /// ## Credential invalidation /// /// Object stores that use dynamic credentials (IMDS, WebIdentity, ECS role, STS assume-role) diff --git a/native/jni-bridge/src/comet_s3_credential_dispatcher.rs b/native/jni-bridge/src/comet_s3_credential_dispatcher.rs index b38ca0348af..210a8db89dc 100644 --- a/native/jni-bridge/src/comet_s3_credential_dispatcher.rs +++ b/native/jni-bridge/src/comet_s3_credential_dispatcher.rs @@ -34,6 +34,10 @@ pub struct CometS3CredentialDispatcher<'a> { pub method_ensure_initialized_ret: ReturnType, pub method_get_credentials_for_path: JStaticMethodID, pub method_get_credentials_for_path_ret: ReturnType, + /// Returns a bucket's policy locations as a `String[]`, or null when the provider does not + /// implement `CometS3LocationScopedCredentialProvider`. + pub method_get_policy_locations: JStaticMethodID, + pub method_get_policy_locations_ret: ReturnType, pub field_access_key_id: JFieldID, pub field_secret_access_key: JFieldID, pub field_session_token: JFieldID, @@ -63,6 +67,12 @@ impl<'a> CometS3CredentialDispatcher<'a> { ), )?, method_get_credentials_for_path_ret: ReturnType::Object, + method_get_policy_locations: env.get_static_method_id( + JNIString::new(Self::JVM_CLASS), + jni::jni_str!("getPolicyLocations"), + jni::jni_sig!("(JLjava/lang/String;)[Ljava/lang/String;"), + )?, + method_get_policy_locations_ret: ReturnType::Array, field_access_key_id: env.get_field_id( &credentials_class, jni::jni_str!("accessKeyId"), diff --git a/spark/src/main/java/org/apache/comet/cloud/s3/CometS3CredentialDispatcher.java b/spark/src/main/java/org/apache/comet/cloud/s3/CometS3CredentialDispatcher.java index 311b36a64a5..5b2284f1237 100644 --- a/spark/src/main/java/org/apache/comet/cloud/s3/CometS3CredentialDispatcher.java +++ b/spark/src/main/java/org/apache/comet/cloud/s3/CometS3CredentialDispatcher.java @@ -22,6 +22,7 @@ import java.lang.reflect.InvocationTargetException; import java.util.Collections; import java.util.HashMap; +import java.util.List; import java.util.Map; import java.util.Objects; import java.util.concurrent.ConcurrentHashMap; @@ -120,6 +121,50 @@ public static CometS3Credentials getCredentialsForPath( new CometS3CredentialContext(bucket, path, accessMode)); } + /** + * Invoked by native code when it creates the object store for {@code bucket}, and again after a + * read through that store fails with 403. Returns {@code null} when the provider behind {@code + * handle} does not implement {@link CometS3LocationScopedCredentialProvider}, which leaves it + * with one credential per bucket. Otherwise returns a copy of the provider's locations. + * + *

Copying the list here runs any lazy {@code List} code inside this call, so its exceptions + * reach native code as ordinary Java exceptions, and a non-{@code String} element fails with + * {@link ArrayStoreException}. A {@code null} list or location is a contract violation and + * throws, so the read fails instead of using a broader credential. + */ + public static String[] getPolicyLocations(long handle, String bucket) throws Exception { + RegisteredProvider registered = INSTANCES.get(handle); + if (registered == null) { + throw new IllegalStateException( + "CometS3CredentialProvider handle " + + handle + + " was not initialized; " + + "ensureInitialized must be called before getPolicyLocations"); + } + if (!(registered.provider instanceof CometS3LocationScopedCredentialProvider)) { + return null; + } + List locations = + ((CometS3LocationScopedCredentialProvider) registered.provider).getPolicyLocations(bucket); + if (locations == null) { + throw new IllegalStateException( + registered.key.providerClassName + + ".getPolicyLocations returned null for bucket " + + bucket + + "; return an empty list when the bucket has no locations"); + } + String[] copy = locations.toArray(new String[0]); + for (String location : copy) { + if (location == null) { + throw new IllegalStateException( + registered.key.providerClassName + + ".getPolicyLocations returned a null location for bucket " + + bucket); + } + } + return copy; + } + private static CometS3CredentialProvider instantiate(String providerClassName) { Class clazz; try { diff --git a/spark/src/main/java/org/apache/comet/cloud/s3/CometS3LocationScopedCredentialProvider.java b/spark/src/main/java/org/apache/comet/cloud/s3/CometS3LocationScopedCredentialProvider.java new file mode 100644 index 00000000000..f97195c74da --- /dev/null +++ b/spark/src/main/java/org/apache/comet/cloud/s3/CometS3LocationScopedCredentialProvider.java @@ -0,0 +1,76 @@ +/* + * 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. + */ + +package org.apache.comet.cloud.s3; + +import java.util.List; + +import org.apache.comet.annotation.Public; + +/** + * Opt-in extension of {@link CometS3CredentialProvider} for buckets whose credentials differ by + * location, for example one policy for {@code warehouse/sales} and another for {@code + * warehouse/finance}. Implementing it tells Comet to request a credential per location rather than + * one per bucket. + * + *

Each request is served with the credential of the longest location that covers its path. A + * location covers a path when the path is the location itself or lies below it, compared one path + * segment at a time: {@code warehouse/sales} covers {@code warehouse/sales/part-0.parquet} but not + * {@code warehouse/sales_eu/part-0.parquet}. When locations nest, the longer one wins. The bucket + * root is an implicit location that covers every path no returned location covers. + * + *

Comet requests a location's credential by calling {@link + * #getCredentialsForPath(CometS3CredentialContext)} with the location as the context's path, as + * returned but with a leading slash ({@code /} for the bucket root). The credential must authorize + * every path for which that location is the longest covering one. + * + *

Locations apply to Comet's native Parquet reads. Iceberg reads do not use them: there Comet + * calls {@link #getCredentialsForPath(CometS3CredentialContext)} as it does for any provider. + * + *

Providers that implement only {@link CometS3CredentialProvider} are unaffected and keep one + * credential per bucket. + */ +@Public +public interface CometS3LocationScopedCredentialProvider extends CometS3CredentialProvider { + + /** + * Returns every location in {@code bucket} that has its own credential policy. + * + *

Locations are written like {@link CometS3CredentialContext#getPath()}: the path within the + * bucket, percent-encoded as in an {@code s3://} URI, without the scheme or bucket name. A + * literal {@code %} must be written as {@code %25}; other characters may be left unencoded. A + * leading or trailing {@code /} is optional. Comet percent-decodes locations and request paths + * before comparing them, and when several locations decode to the same path it keeps the first. + * + *

A location is invalid if, once decoded, it is not valid UTF-8 or has a segment that is + * empty, {@code .}, {@code ..}, or contains a control character, so a URI such as {@code + * s3://bucket/a} is invalid too. An invalid or {@code null} location fails the read, as does a + * {@code null} list. + * + *

Comet treats the result as a snapshot and may keep it for the life of an executor. It asks + * again after a request fails with 403, so a location added or removed later may not take effect + * until then. It may call this on the driver or on executors, and from several threads at once. + * + * @param bucket the S3 bucket name, without scheme or path + * @return the bucket's locations, or an empty list if every path uses the bucket-root credential + * @throws Exception if the locations cannot be determined. Comet then fails the read rather than + * fall back to a broader credential. + */ + List getPolicyLocations(String bucket) throws Exception; +} diff --git a/spark/src/test/java/org/apache/comet/cloud/s3/CometS3LocationScopedCredentialProviderTest.java b/spark/src/test/java/org/apache/comet/cloud/s3/CometS3LocationScopedCredentialProviderTest.java new file mode 100644 index 00000000000..f542bcd39f9 --- /dev/null +++ b/spark/src/test/java/org/apache/comet/cloud/s3/CometS3LocationScopedCredentialProviderTest.java @@ -0,0 +1,176 @@ +/* + * 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. + */ + +package org.apache.comet.cloud.s3; + +import java.util.AbstractList; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Collections; +import java.util.List; + +import org.junit.Before; +import org.junit.Test; + +import static org.junit.Assert.assertArrayEquals; +import static org.junit.Assert.assertEquals; +import static org.junit.Assert.assertNull; +import static org.junit.Assert.assertSame; +import static org.junit.Assert.assertThrows; +import static org.junit.Assert.assertTrue; + +/** + * Covers {@link CometS3CredentialDispatcher#getPolicyLocations(long, String)}, the entry point + * native code uses to read a {@link CometS3LocationScopedCredentialProvider}'s locations. + */ +public class CometS3LocationScopedCredentialProviderTest { + + private static final String BASE_PROVIDER = TestCometS3CredentialProvider.class.getName(); + private static final String SCOPED_PROVIDER = + TestCometS3LocationScopedCredentialProvider.class.getName(); + private static final String DK = "location-scoped-test-dispatch-key"; + + @Before + public void resetTestProviders() { + CometS3CredentialDispatcher.closeAll(); + TestCometS3CredentialProvider.reset(); + TestCometS3LocationScopedCredentialProvider.reset(); + } + + private static long scopedHandle() { + return CometS3CredentialDispatcher.ensureInitialized( + SCOPED_PROVIDER, DK, Collections.emptyMap()); + } + + @Test + public void baseProviderReturnsNullWithoutCallingIt() throws Exception { + long handle = + CometS3CredentialDispatcher.ensureInitialized(BASE_PROVIDER, DK, Collections.emptyMap()); + + assertNull(CometS3CredentialDispatcher.getPolicyLocations(handle, "b")); + assertEquals(0, TestCometS3CredentialProvider.callCount.get()); + } + + @Test + public void returnsTheProvidersLocationsForTheBucket() throws Exception { + long handle = scopedHandle(); + TestCometS3LocationScopedCredentialProvider.nextLocations.set( + List.of("warehouse/sales", "/warehouse/finance/")); + + String[] locations = CometS3CredentialDispatcher.getPolicyLocations(handle, "my-bucket"); + + assertArrayEquals(new String[] {"warehouse/sales", "/warehouse/finance/"}, locations); + assertEquals("my-bucket", TestCometS3LocationScopedCredentialProvider.lastBucket); + assertEquals(1, TestCometS3LocationScopedCredentialProvider.locationCallCount.get()); + assertEquals( + "getPolicyLocations must not fetch credentials", + 0, + TestCometS3LocationScopedCredentialProvider.callCount.get()); + } + + @Test + public void emptyListReturnsEmptyArray() throws Exception { + String[] locations = CometS3CredentialDispatcher.getPolicyLocations(scopedHandle(), "b"); + + assertEquals(0, locations.length); + } + + @Test + public void nullListIsRejected() { + long handle = scopedHandle(); + TestCometS3LocationScopedCredentialProvider.nextLocations.set(null); + + IllegalStateException thrown = + assertThrows( + IllegalStateException.class, + () -> CometS3CredentialDispatcher.getPolicyLocations(handle, "b")); + assertTrue(thrown.getMessage(), thrown.getMessage().contains("returned null")); + } + + @Test + public void nullLocationIsRejected() { + long handle = scopedHandle(); + TestCometS3LocationScopedCredentialProvider.nextLocations.set(Arrays.asList("a", null)); + + IllegalStateException thrown = + assertThrows( + IllegalStateException.class, + () -> CometS3CredentialDispatcher.getPolicyLocations(handle, "b")); + assertTrue(thrown.getMessage(), thrown.getMessage().contains("null location")); + } + + @Test + @SuppressWarnings({"unchecked", "rawtypes"}) + public void nonStringLocationIsRejected() { + long handle = scopedHandle(); + List raw = new ArrayList(); + raw.add("a"); + raw.add(42); + TestCometS3LocationScopedCredentialProvider.nextLocations.set(raw); + + assertThrows( + ArrayStoreException.class, + () -> CometS3CredentialDispatcher.getPolicyLocations(handle, "b")); + } + + /** A lazy list's exception surfaces from the dispatcher call, not later in native code. */ + @Test + public void lazyListExceptionsPropagate() { + long handle = scopedHandle(); + IllegalStateException boom = new IllegalStateException("policy source unavailable"); + TestCometS3LocationScopedCredentialProvider.nextLocations.set( + new AbstractList() { + @Override + public String get(int index) { + throw boom; + } + + @Override + public int size() { + throw boom; + } + }); + + Exception thrown = + assertThrows( + Exception.class, () -> CometS3CredentialDispatcher.getPolicyLocations(handle, "b")); + assertSame(boom, thrown); + } + + @Test + public void providerExceptionsPropagate() { + long handle = scopedHandle(); + IllegalStateException boom = new IllegalStateException("simulated policy source failure"); + TestCometS3LocationScopedCredentialProvider.throwOnNextLocationCall = boom; + + Exception thrown = + assertThrows( + Exception.class, () -> CometS3CredentialDispatcher.getPolicyLocations(handle, "b")); + assertSame(boom, thrown); + } + + @Test + public void unknownHandleRejected() { + IllegalStateException thrown = + assertThrows( + IllegalStateException.class, + () -> CometS3CredentialDispatcher.getPolicyLocations(Long.MAX_VALUE, "b")); + assertTrue(thrown.getMessage().contains("not initialized")); + } +} diff --git a/spark/src/test/java/org/apache/comet/cloud/s3/MinioLocationScopedCredentialProvider.java b/spark/src/test/java/org/apache/comet/cloud/s3/MinioLocationScopedCredentialProvider.java new file mode 100644 index 00000000000..dc6389065ba --- /dev/null +++ b/spark/src/test/java/org/apache/comet/cloud/s3/MinioLocationScopedCredentialProvider.java @@ -0,0 +1,100 @@ +/* + * 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. + */ + +package org.apache.comet.cloud.s3; + +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.Set; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; + +/** + * Test {@link CometS3LocationScopedCredentialProvider} for suites that run against Minio. It + * returns the locations installed with {@link #installLocations} and Minio's static credentials, + * and records the path of every credential request, so a suite can assert which location served + * each read. Minio does not enforce per-prefix policies here, so the recorded paths, not a 403, + * show that routing worked. + */ +public final class MinioLocationScopedCredentialProvider + implements CometS3LocationScopedCredentialProvider { + + private static final AtomicReference CREDS = new AtomicReference<>(); + private static final AtomicReference> LOCATIONS = + new AtomicReference<>(Collections.emptyList()); + private static final AtomicInteger LOCATION_CALL_COUNT = new AtomicInteger(0); + + /** + * Counts provider instances. Not cleared by {@link #resetCounters}: the dispatcher keeps one + * instance per registration for the life of the JVM. + */ + private static final AtomicInteger INIT_COUNT = new AtomicInteger(0); + + private static final Set CREDENTIAL_PATHS = ConcurrentHashMap.newKeySet(); + + public static void installCredentials(String accessKeyId, String secretAccessKey) { + CREDS.set(new CometS3Credentials(accessKeyId, secretAccessKey, null, 0L)); + } + + public static void installLocations(List locations) { + LOCATIONS.set(List.copyOf(locations)); + } + + public static int locationCallCount() { + return LOCATION_CALL_COUNT.get(); + } + + public static int initCount() { + return INIT_COUNT.get(); + } + + /** The paths passed to {@link #getCredentialsForPath} since the last reset. */ + public static Set credentialPaths() { + return Set.copyOf(CREDENTIAL_PATHS); + } + + public static void resetCounters() { + LOCATION_CALL_COUNT.set(0); + CREDENTIAL_PATHS.clear(); + } + + @Override + public void initialize(Map catalogProperties) { + INIT_COUNT.incrementAndGet(); + } + + @Override + public CometS3Credentials getCredentialsForPath(CometS3CredentialContext context) { + CREDENTIAL_PATHS.add(context.getPath()); + CometS3Credentials creds = CREDS.get(); + if (creds == null) { + throw new IllegalStateException( + "MinioLocationScopedCredentialProvider.installCredentials was not called"); + } + return creds; + } + + @Override + public List getPolicyLocations(String bucket) { + LOCATION_CALL_COUNT.incrementAndGet(); + return LOCATIONS.get(); + } +} diff --git a/spark/src/test/java/org/apache/comet/cloud/s3/TestCometS3LocationScopedCredentialProvider.java b/spark/src/test/java/org/apache/comet/cloud/s3/TestCometS3LocationScopedCredentialProvider.java new file mode 100644 index 00000000000..083eab7fcf3 --- /dev/null +++ b/spark/src/test/java/org/apache/comet/cloud/s3/TestCometS3LocationScopedCredentialProvider.java @@ -0,0 +1,70 @@ +/* + * 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. + */ + +package org.apache.comet.cloud.s3; + +import java.util.Collections; +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; + +/** + * Test-only {@link CometS3LocationScopedCredentialProvider}. State is static because the dispatcher + * caches one instance per (FQCN, dispatchKey) for the JVM lifetime. + */ +public class TestCometS3LocationScopedCredentialProvider + implements CometS3LocationScopedCredentialProvider { + + static final AtomicInteger callCount = new AtomicInteger(0); + static final AtomicInteger locationCallCount = new AtomicInteger(0); + static final AtomicReference> nextLocations = + new AtomicReference<>(Collections.emptyList()); + static volatile String lastBucket; + static volatile Exception throwOnNextLocationCall; + + static void reset() { + callCount.set(0); + locationCallCount.set(0); + nextLocations.set(Collections.emptyList()); + lastBucket = null; + throwOnNextLocationCall = null; + } + + @Override + public void initialize(Map catalogProperties) {} + + @Override + public CometS3Credentials getCredentialsForPath(CometS3CredentialContext context) { + callCount.incrementAndGet(); + return new CometS3Credentials("AKIASCOPED", "secret", "session-tok", 0L); + } + + @Override + public List getPolicyLocations(String bucket) throws Exception { + locationCallCount.incrementAndGet(); + lastBucket = bucket; + Exception toThrow = throwOnNextLocationCall; + if (toThrow != null) { + throwOnNextLocationCall = null; + throw toThrow; + } + return nextLocations.get(); + } +} diff --git a/spark/src/test/scala/org/apache/comet/CometPublicApiSuite.scala b/spark/src/test/scala/org/apache/comet/CometPublicApiSuite.scala index d16b09b3605..f6e15cfa8e6 100644 --- a/spark/src/test/scala/org/apache/comet/CometPublicApiSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometPublicApiSuite.scala @@ -44,7 +44,8 @@ class CometPublicApiSuite extends AnyFunSuite { "org.apache.comet.cloud.s3.CometS3AccessMode", "org.apache.comet.cloud.s3.CometS3CredentialContext", "org.apache.comet.cloud.s3.CometS3CredentialProvider", - "org.apache.comet.cloud.s3.CometS3Credentials") + "org.apache.comet.cloud.s3.CometS3Credentials", + "org.apache.comet.cloud.s3.CometS3LocationScopedCredentialProvider") test("the public API is exactly the set enumerated in the versioning policy") { val classesDir = mainClassesDir() diff --git a/spark/src/test/scala/org/apache/comet/cloud/s3/CometS3CredentialBridgeSuite.scala b/spark/src/test/scala/org/apache/comet/cloud/s3/CometS3CredentialBridgeSuite.scala index 2bb0b173342..d6b8ebb4b0e 100644 --- a/spark/src/test/scala/org/apache/comet/cloud/s3/CometS3CredentialBridgeSuite.scala +++ b/spark/src/test/scala/org/apache/comet/cloud/s3/CometS3CredentialBridgeSuite.scala @@ -41,11 +41,22 @@ class CometS3CredentialBridgeSuite override protected val testBucketName = "bridge-test-bucket" + /** + * Read through [[MinioLocationScopedCredentialProvider]] by a per-bucket override, so the rest + * of the suite keeps exercising the base provider and this bucket's cache entry is its own. + */ + private val scopedBucket = "bridge-scoped-bucket" + private val scopedLocations = + java.util.List.of("warehouse/sales", "warehouse/sales/eu/", "/warehouse/finance") + override protected def sparkConf: SparkConf = { val conf = super.sparkConf val providerClassName = classOf[MinioCometS3CredentialProvider].getName // Activate the bridge for the Parquet (object_store) path via the Hadoop S3A namespace. conf.set("spark.hadoop.fs.s3a.comet.credential.provider.class", providerClassName) + conf.set( + s"spark.hadoop.fs.s3a.bucket.$scopedBucket.comet.credential.provider.class", + classOf[MinioLocationScopedCredentialProvider].getName) // Activate the bridge for the Iceberg (opendal) path via the per-catalog s3 namespace. conf.set("spark.sql.catalog.s3_catalog", "org.apache.iceberg.spark.SparkCatalog") conf.set("spark.sql.catalog.s3_catalog.type", "hadoop") @@ -59,6 +70,8 @@ class CometS3CredentialBridgeSuite override def beforeAll(): Unit = { super.beforeAll() MinioCometS3CredentialProvider.installCredentials(userName, password) + MinioLocationScopedCredentialProvider.installCredentials(userName, password) + createBucketIfNotExists(scopedBucket) } private def assertHasCometParquetScan(plan: SparkPlan): Unit = @@ -86,6 +99,10 @@ class CometS3CredentialBridgeSuite assert( MinioCometS3CredentialProvider.lastBucket() == testBucketName, s"Bridge received unexpected bucket: ${MinioCometS3CredentialProvider.lastBucket()}") + // A base provider keeps one credential per bucket, requested with the path of a file it reads. + assert( + MinioCometS3CredentialProvider.lastPath().startsWith("/data/bridge-parquet.parquet/"), + s"Bridge received unexpected path: ${MinioCometS3CredentialProvider.lastPath()}") } test("Iceberg read on S3 routes credentials through CometS3CredentialProvider") { @@ -229,4 +246,73 @@ class CometS3CredentialBridgeSuite spark.sql("DROP TABLE iso_b.db.t") } } + + // Minio does not enforce per-prefix policies here, so these tests check the path of each + // credential request: a read must request the credential of its longest covering location. + + test("location-scoped provider: one scan requests each file's longest covering location") { + MinioLocationScopedCredentialProvider.installLocations(scopedLocations) + val files = Seq( + "warehouse/sales/a.parquet" -> 10L, + "warehouse/sales/eu/b.parquet" -> 20L, + // Shares a name prefix with warehouse/sales but not a segment, so the bucket root covers it. + // It is the only file under the root, so "/" below shows it was not routed by name prefix. + "warehouse/sales_eu/c.parquet" -> 30L, + "warehouse/finance/d.parquet" -> 40L) + val paths = files.map { case (key, _) => s"s3a://$scopedBucket/$key" } + files.zip(paths).foreach { case ((_, rows), path) => + spark.range(0, rows).write.format("parquet").mode(SaveMode.Overwrite).save(path) + } + + MinioLocationScopedCredentialProvider.resetCounters() + // Pack every file into one partition, so one store serves every location in a single task. + withSQLConf( + "spark.sql.files.openCostInBytes" -> "1", + "spark.sql.files.minPartitionNum" -> "1") { + val df = spark.read.format("parquet").load(paths: _*) + assert(df.rdd.getNumPartitions == 1) + val total = df.agg(sum(col("id"))) + assertHasCometParquetScan(total.queryExecution.executedPlan) + assert(total.first().getLong(0) == files.map { case (_, rows) => (0L until rows).sum }.sum) + } + + assert( + MinioLocationScopedCredentialProvider.credentialPaths() == + java.util.Set.of("/warehouse/sales", "/warehouse/sales/eu/", "/warehouse/finance", "/"), + s"Unexpected credential paths: ${MinioLocationScopedCredentialProvider.credentialPaths()}") + // Every location's bridge shares the bucket's provider registration, whatever properties the + // bucket's store was created with, so one provider instance serves all four locations. + assert( + MinioLocationScopedCredentialProvider.initCount() == 1, + s"Provider initialized ${MinioLocationScopedCredentialProvider.initCount()} times") + } + + test("location-scoped provider: later scans reuse the bucket's locations") { + MinioLocationScopedCredentialProvider.installLocations(scopedLocations) + val path = s"s3a://$scopedBucket/warehouse/finance/reuse.parquet" + spark.range(0, 100).write.format("parquet").mode(SaveMode.Overwrite).save(path) + val expectedSum = (0L until 100L).sum + // Creates the bucket's store unless an earlier test did. Its tasks can each miss the store + // cache at once and fetch the locations, so the count is only checked after this scan. + assert( + spark.read.format("parquet").load(path).agg(sum(col("id"))).first().getLong(0) == + expectedSum) + + MinioLocationScopedCredentialProvider.resetCounters() + for (_ <- 1 to 3) { + val df = spark.read.format("parquet").load(path).agg(sum(col("id"))) + assertHasCometParquetScan(df.queryExecution.executedPlan) + assert(df.first().getLong(0) == expectedSum) + } + + assert( + MinioLocationScopedCredentialProvider.locationCallCount() == 0, + s"Locations fetched ${MinioLocationScopedCredentialProvider.locationCallCount()} times") + // Collections.singleton, not Set.of: Scala 2.12 cannot choose between Set.of(E) and + // Set.of(E...) for a single argument. + assert( + MinioLocationScopedCredentialProvider.credentialPaths() == java.util.Collections.singleton( + "/warehouse/finance"), + s"Unexpected credential paths: ${MinioLocationScopedCredentialProvider.credentialPaths()}") + } }