From 70deb3284b715059fe545bc90a98bf41af253a30 Mon Sep 17 00:00:00 2001 From: GatewayJ <835269233@qq.com> Date: Sun, 9 Aug 2026 14:08:53 +0800 Subject: [PATCH] fix(select): pin object snapshot for query lifetime (#5835) --- Cargo.lock | 3 + crates/ecstore/Cargo.toml | 1 + crates/ecstore/src/api/mod.rs | 10 +- crates/ecstore/src/core/sets.rs | 179 +-- crates/ecstore/src/runtime/instance.rs | 5 + .../src/set_disk/core/io_primitives.rs | 2 +- crates/ecstore/src/set_disk/mod.rs | 57 +- crates/ecstore/src/set_disk/ops/object.rs | 71 +- crates/ecstore/src/store/mod.rs | 5 +- crates/ecstore/src/store/object.rs | 1211 ++++++++++++++++- crates/s3select-api/Cargo.toml | 3 +- crates/s3select-api/src/lib.rs | 40 +- crates/s3select-api/src/object_store.rs | 1184 +++++++++++----- crates/s3select-api/src/query/dispatcher.rs | 9 + crates/s3select-api/src/query/mod.rs | 22 +- crates/s3select-api/src/query/session.rs | 169 ++- crates/s3select-api/src/server/dbms.rs | 10 + crates/s3select-api/src/storage_api.rs | 28 +- crates/s3select-query/Cargo.toml | 3 + .../s3select-query/src/dispatcher/manager.rs | 394 +++++- crates/s3select-query/src/instance.rs | 12 +- crates/test-utils/Cargo.toml | 3 +- crates/test-utils/src/ecstore_test_compat.rs | 5 +- crates/test-utils/src/lib.rs | 51 +- rustfs/src/app/select_object.rs | 652 ++++++++- rustfs/src/app/storage_api.rs | 13 +- rustfs/src/storage/storage_api.rs | 9 +- rustfs/tests/embedded_select_snapshot_test.rs | 227 +++ 28 files changed, 3806 insertions(+), 572 deletions(-) create mode 100644 rustfs/tests/embedded_select_snapshot_test.rs diff --git a/Cargo.lock b/Cargo.lock index cb8ff239b..308693515 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9378,6 +9378,7 @@ dependencies = [ "thiserror 2.0.20", "time", "tokio", + "tokio-stream", "tokio-util", "tonic", "tower", @@ -10117,6 +10118,7 @@ dependencies = [ "tracing", "transform-stream", "url", + "uuid", ] [[package]] @@ -10132,6 +10134,7 @@ dependencies = [ "hotpath", "parking_lot", "rustfs-s3select-api", + "rustfs-test-utils", "s3s", "tokio", "tracing", diff --git a/crates/ecstore/Cargo.toml b/crates/ecstore/Cargo.toml index 35470aa1f..23a91ade9 100644 --- a/crates/ecstore/Cargo.toml +++ b/crates/ecstore/Cargo.toml @@ -181,6 +181,7 @@ path-absolutize = { workspace = true } rmp.workspace = true rmp-serde.workspace = true tokio-util = { workspace = true, features = ["io", "compat"] } +tokio-stream = { workspace = true, features = ["sync"] } base64 = { workspace = true } hmac = { workspace = true } sha1 = { workspace = true } diff --git a/crates/ecstore/src/api/mod.rs b/crates/ecstore/src/api/mod.rs index cc31820ec..98dfbcdbd 100644 --- a/crates/ecstore/src/api/mod.rs +++ b/crates/ecstore/src/api/mod.rs @@ -414,7 +414,10 @@ pub mod object { lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, register_object_mutation_hook, unregister_get_object_body_cache_hook, unregister_object_mutation_hook, }; - pub use crate::store::PreparedGetObjectReader; + pub use crate::store::{ + PrepareSelectObjectSnapshotError, PreparedGetObjectReader, SelectObjectSnapshot, SelectObjectSnapshotReadError, + SnapshotConsistencyError, + }; } pub mod rebalance { @@ -449,6 +452,11 @@ pub mod rpc { pub mod set_disk { pub use crate::set_disk::{DEFAULT_READ_BUFFER_SIZE, SetDisks, get_lock_acquire_timeout, is_valid_storage_class}; + + #[cfg(feature = "test-util")] + pub mod test_util { + pub use crate::set_disk::{PutObjectCommitBarrier, PutObjectCommitPause}; + } } pub mod store_list { diff --git a/crates/ecstore/src/core/sets.rs b/crates/ecstore/src/core/sets.rs index b8804963e..1d5fcbdbc 100644 --- a/crates/ecstore/src/core/sets.rs +++ b/crates/ecstore/src/core/sets.rs @@ -1215,6 +1215,98 @@ async fn init_storage_disks_with_errors( (disks, errs) } +#[cfg(test)] +pub(crate) async fn make_local_two_set_sets() -> (Vec, Arc) { + make_local_two_set_sets_with_ctx(bootstrap_ctx()).await +} + +#[cfg(test)] +pub(crate) async fn make_local_two_set_sets_with_ctx(ctx: Arc) -> (Vec, Arc) { + use crate::layout::endpoint::Endpoint; + use rustfs_lock::client::local::LocalClient; + + let format = FormatV3::new(2, 2); + let mut temp_dirs = Vec::new(); + let mut all_endpoints = Vec::new(); + let mut disk_sets = Vec::new(); + + for set_index in 0..2 { + let mut endpoints = Vec::new(); + let mut disks = Vec::new(); + for disk_index in 0..2 { + let temp_dir = tempfile::tempdir().expect("tempdir should be created"); + let mut endpoint = Endpoint::try_from(temp_dir.path().to_str().expect("tempdir path should be utf8")) + .expect("endpoint should parse"); + endpoint.set_pool_index(0); + endpoint.set_set_index(set_index); + endpoint.set_disk_index(disk_index); + let disk = new_disk( + &endpoint, + &DiskOption { + cleanup: false, + health_check: false, + }, + ) + .await + .expect("disk should be created"); + let mut disk_format = format.clone(); + disk_format.erasure.this = format.erasure.sets[set_index][disk_index]; + save_format_file(&Some(disk.clone()), &Some(disk_format)) + .await + .expect("format should be saved"); + temp_dirs.push(temp_dir); + all_endpoints.push(endpoint.clone()); + endpoints.push(endpoint); + disks.push(Some(disk)); + } + let lockers = (0..2) + .map(|_| { + Arc::new(LocalClient::with_manager(Arc::new(rustfs_lock::GlobalLockManager::Enabled(Arc::new( + rustfs_lock::FastObjectLockManager::new(), + ))))) as Arc + }) + .collect(); + disk_sets.push( + SetDisks::new_with_instance_ctx( + "test-owner".to_string(), + Arc::new(RwLock::new(disks)), + 2, + 1, + set_index, + 0, + endpoints, + format.clone(), + lockers, + Arc::clone(&ctx), + ) + .await, + ); + } + + let sets = Arc::new(Sets { + id: format.id, + disk_set: disk_sets, + pool_idx: 0, + endpoints: PoolEndpoints { + legacy: false, + set_count: 2, + drives_per_set: 2, + endpoints: Endpoints::from(all_endpoints), + cmd_line: String::new(), + platform: String::new(), + }, + format, + parity_count: 1, + set_count: 2, + set_drive_count: 2, + default_parity_count: 1, + distribution_algo: DistributionAlgoVersion::V1, + exit_signal: None, + ctx, + }); + (temp_dirs, sets) +} + #[cfg(test)] mod tests { use super::*; @@ -1373,81 +1465,6 @@ mod tests { assert_eq!(result, (Some(3), Some(1), Some(0))); } - async fn two_set_test_sets() -> (Vec, Arc) { - let format = FormatV3::new(2, 2); - let mut temp_dirs = Vec::new(); - let mut all_endpoints = Vec::new(); - let mut disk_sets = Vec::new(); - - for set_index in 0..2 { - let mut endpoints = Vec::new(); - let mut disks = Vec::new(); - for disk_index in 0..2 { - let temp_dir = tempfile::tempdir().expect("tempdir should be created"); - let mut endpoint = Endpoint::try_from(temp_dir.path().to_str().expect("tempdir path should be utf8")) - .expect("endpoint should parse"); - endpoint.set_pool_index(0); - endpoint.set_set_index(set_index); - endpoint.set_disk_index(disk_index); - let disk = new_disk( - &endpoint, - &DiskOption { - cleanup: false, - health_check: false, - }, - ) - .await - .expect("disk should be created"); - let mut disk_format = format.clone(); - disk_format.erasure.this = format.erasure.sets[set_index][disk_index]; - save_format_file(&Some(disk.clone()), &Some(disk_format)) - .await - .expect("format should be saved"); - temp_dirs.push(temp_dir); - all_endpoints.push(endpoint.clone()); - endpoints.push(endpoint); - disks.push(Some(disk)); - } - disk_sets.push( - SetDisks::new( - "test-owner".to_string(), - Arc::new(RwLock::new(disks)), - 2, - 1, - set_index, - 0, - endpoints, - format.clone(), - vec![Arc::new(LocalClient::new()), Arc::new(LocalClient::new())], - ) - .await, - ); - } - - let sets = Arc::new(Sets { - id: format.id, - disk_set: disk_sets, - pool_idx: 0, - endpoints: PoolEndpoints { - legacy: false, - set_count: 2, - drives_per_set: 2, - endpoints: Endpoints::from(all_endpoints), - cmd_line: String::new(), - platform: String::new(), - }, - format, - parity_count: 1, - set_count: 2, - set_drive_count: 2, - default_parity_count: 1, - distribution_algo: DistributionAlgoVersion::V1, - exit_signal: None, - ctx: bootstrap_ctx(), - }); - (temp_dirs, sets) - } - #[tokio::test] async fn heal_object_uses_explicit_set_scope() { let (_temp_dirs, sets) = two_set_test_sets().await; @@ -1497,7 +1514,7 @@ mod tests { #[tokio::test] async fn delete_prefix_surfaces_a_hard_error_from_any_set() { - let (_temp_dirs, sets) = two_set_test_sets().await; + let (_temp_dirs, sets) = make_local_two_set_sets().await; let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple()); sets.make_bucket(&bucket, &MakeBucketOptions::default()) .await @@ -1546,7 +1563,7 @@ mod tests { #[tokio::test] async fn delete_prefix_keeps_a_missing_bucket_idempotent_across_sets() { - let (_temp_dirs, sets) = two_set_test_sets().await; + let (_temp_dirs, sets) = make_local_two_set_sets().await; let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple()); sets.make_bucket(&bucket, &MakeBucketOptions::default()) .await @@ -1585,7 +1602,7 @@ mod tests { #[tokio::test] async fn delete_prefix_preserves_a_completely_missing_bucket_error() { - let (_temp_dirs, sets) = two_set_test_sets().await; + let (_temp_dirs, sets) = make_local_two_set_sets().await; let bucket = format!("delete-prefix-missing-{}", Uuid::new_v4().simple()); let err = sets @@ -1605,7 +1622,7 @@ mod tests { #[tokio::test] async fn delete_prefix_fails_when_one_set_is_entirely_offline() { - let (_temp_dirs, sets) = two_set_test_sets().await; + let (_temp_dirs, sets) = make_local_two_set_sets().await; let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple()); sets.make_bucket(&bucket, &MakeBucketOptions::default()) .await @@ -1652,7 +1669,7 @@ mod tests { #[tokio::test] async fn set_format_heal_accepts_quorum_from_a_nonzero_set() { - let (_temp_dirs, sets) = two_set_test_sets().await; + let (_temp_dirs, sets) = make_local_two_set_sets().await; let (result, err) = sets.disk_set[1] .heal_format(false) @@ -1757,7 +1774,7 @@ mod tests { #[serial] async fn list_multipart_uploads_merges_all_sets_without_pagination_loss() { let _setup_type_guard = SetupTypeGuard::switch_to(SetupType::Erasure).await; - let (_temp_dirs, sets) = two_set_test_sets().await; + let (_temp_dirs, sets) = make_local_two_set_sets().await; let bucket = format!("multipart-list-{}", Uuid::new_v4().simple()); sets.make_bucket(&bucket, &MakeBucketOptions::default()) .await diff --git a/crates/ecstore/src/runtime/instance.rs b/crates/ecstore/src/runtime/instance.rs index ed0bda1dd..95ff16ffe 100644 --- a/crates/ecstore/src/runtime/instance.rs +++ b/crates/ecstore/src/runtime/instance.rs @@ -208,6 +208,11 @@ impl InstanceContext { } } + #[cfg(test)] + pub(crate) fn with_lock_manager_for_test(lock_manager: Arc) -> Self { + Self::with_lock_manager(lock_manager) + } + /// This instance's namespace lock manager. pub fn lock_manager(&self) -> Arc { self.lock_manager.clone() diff --git a/crates/ecstore/src/set_disk/core/io_primitives.rs b/crates/ecstore/src/set_disk/core/io_primitives.rs index 5c67c072c..40bc2673b 100644 --- a/crates/ecstore/src/set_disk/core/io_primitives.rs +++ b/crates/ecstore/src/set_disk/core/io_primitives.rs @@ -4651,7 +4651,7 @@ pub(in crate::set_disk) mod cleanup_fault_injection { /// unobserved object records nothing, keeping the registry bounded, and each /// [`CallCounterScope`] clears only its own object's counts on drop. #[cfg(test)] -pub(in crate::set_disk) mod disk_call_counters { +pub(crate) mod disk_call_counters { use std::collections::{HashMap, HashSet}; use std::sync::{Mutex, OnceLock}; diff --git a/crates/ecstore/src/set_disk/mod.rs b/crates/ecstore/src/set_disk/mod.rs index 5f2922c5a..0a07cc1bf 100644 --- a/crates/ecstore/src/set_disk/mod.rs +++ b/crates/ecstore/src/set_disk/mod.rs @@ -692,6 +692,8 @@ const DEFAULT_RUSTFS_GET_MULTIPART_READER_SETUP_PREFETCH: bool = true; static OBJECT_LOCK_DIAG_ENABLED: OnceLock = OnceLock::new(); mod core; +#[cfg(test)] +pub(crate) use core::io_primitives::disk_call_counters; mod ctx; mod metadata; mod ops; @@ -702,8 +704,8 @@ pub(crate) use ops::object::TransitionCleanupStoreBarrier as SetDiskTransitionCl pub(crate) use ops::object::body_cache_plaintext_len; #[cfg(test)] pub(crate) use ops::object::cleanup_rejected_transition_upload_durably; -#[cfg(test)] -pub(crate) use ops::object::{PutObjectCommitBarrier, PutObjectCommitPause}; +#[cfg(any(test, feature = "test-util"))] +pub use ops::object::{PutObjectCommitBarrier, PutObjectCommitPause}; mod read; mod replication; pub(crate) mod shard_source; @@ -731,6 +733,10 @@ impl PreparedGetObjectMetadata { .take() .expect("prepared GET metadata ObjectInfo must be consumed exactly once") } + + pub(crate) fn read_semantics_identity(&self) -> [u8; 32] { + SetDisks::file_info_quorum_hash(&self.fi) + } } tokio::task_local! { @@ -2553,6 +2559,53 @@ impl SetDisks { )) } + #[cfg(any(test, feature = "test-util"))] + async fn acquire_write_lock_diag_with_pending_hook( + &self, + op: &'static str, + bucket: &str, + object: &str, + on_pending: impl FnOnce(), + ) -> Result { + crate::hp_guard!("SetDisks::acquire_write_lock"); + let diag_enabled = is_object_lock_diag_enabled(); + let ns_lock = self.new_ns_lock(bucket, object).await?; + let acquire_start = Instant::now(); + let acquire = ns_lock.get_write_lock(get_lock_acquire_timeout()); + tokio::pin!(acquire); + let mut on_pending = Some(on_pending); + let guard = futures::future::poll_fn(|cx| match std::future::Future::poll(acquire.as_mut(), cx) { + std::task::Poll::Pending => { + if let Some(on_pending) = on_pending.take() { + on_pending(); + } + std::task::Poll::Pending + } + std::task::Poll::Ready(result) => std::task::Poll::Ready(result), + }) + .await + .map_err(|e| self.map_namespace_lock_error(bucket, object, "write", e))?; + let owner = diag_enabled.then(|| ns_lock.owner().to_string()); + self.log_object_lock_acquire_if_slow( + op, + bucket, + object, + "write", + owner.as_deref(), + acquire_start.elapsed(), + diag_enabled, + ); + Ok(ObjectLockDiagGuard::new( + guard, + diag_enabled, + op, + diag_enabled.then(|| bucket.to_string()), + diag_enabled.then(|| object.to_string()), + owner, + "write", + )) + } + #[allow(clippy::too_many_arguments)] fn log_object_lock_acquire_if_slow( &self, diff --git a/crates/ecstore/src/set_disk/ops/object.rs b/crates/ecstore/src/set_disk/ops/object.rs index e51247eed..de4bed37d 100644 --- a/crates/ecstore/src/set_disk/ops/object.rs +++ b/crates/ecstore/src/set_disk/ops/object.rs @@ -1312,7 +1312,7 @@ impl SetDisks { } if !opts.no_lock && object_lock_guard.is_none() { - #[cfg(test)] + #[cfg(any(test, feature = "test-util"))] pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeNamespace).await; if let Some(expected_incarnation_id) = opts.expected_bucket_incarnation_id && opts.bucket_lifecycle_lock_fence.is_none() @@ -1324,9 +1324,21 @@ impl SetDisks { .await?, ); } - object_lock_guard = Some(self.acquire_write_lock_diag("put_object_commit", bucket, object).await?); + #[cfg(any(test, feature = "test-util"))] + { + object_lock_guard = Some( + self.acquire_write_lock_diag_with_pending_hook("put_object_commit", bucket, object, || { + notify_put_object_commit_namespace_pending(bucket, object); + }) + .await?, + ); + } + #[cfg(not(any(test, feature = "test-util")))] + { + object_lock_guard = Some(self.acquire_write_lock_diag("put_object_commit", bucket, object).await?); + } } - #[cfg(test)] + #[cfg(any(test, feature = "test-util"))] pause_put_object_commit(bucket, object, PutObjectCommitPause::AfterNamespace).await; if deferred_data_movement_precondition && let Some(err) = self.check_write_precondition(bucket, object, opts).await { @@ -2575,41 +2587,43 @@ fn remote_version_state_writer_enabled_for(requested: bool, fleet_confirmed: boo requested && fleet_confirmed && fleet_proof_valid } -#[cfg(test)] +#[cfg(any(test, feature = "test-util"))] #[derive(Clone, Copy, Debug, Eq, PartialEq)] -pub(crate) enum PutObjectCommitPause { +pub enum PutObjectCommitPause { BeforeNamespace, AfterNamespace, BeforeMetadata, } -#[cfg(test)] +#[cfg(any(test, feature = "test-util"))] struct PutObjectCommitBarrierState { bucket: String, object: String, pause: PutObjectCommitPause, arrived: tokio::sync::Notify, release: tokio::sync::Notify, + namespace_pending: tokio::sync::Notify, } -#[cfg(test)] -pub(crate) struct PutObjectCommitBarrier { +#[cfg(any(test, feature = "test-util"))] +pub struct PutObjectCommitBarrier { state: Arc, } -#[cfg(test)] +#[cfg(any(test, feature = "test-util"))] static PUT_OBJECT_COMMIT_BARRIER: std::sync::OnceLock>>> = std::sync::OnceLock::new(); -#[cfg(test)] +#[cfg(any(test, feature = "test-util"))] impl PutObjectCommitBarrier { - pub(crate) fn install(bucket: &str, object: &str, pause: PutObjectCommitPause) -> Self { + pub fn install(bucket: &str, object: &str, pause: PutObjectCommitPause) -> Self { let state = Arc::new(PutObjectCommitBarrierState { bucket: bucket.to_string(), object: object.to_string(), pause, arrived: tokio::sync::Notify::new(), release: tokio::sync::Notify::new(), + namespace_pending: tokio::sync::Notify::new(), }); let mut slot = PUT_OBJECT_COMMIT_BARRIER .get_or_init(|| std::sync::Mutex::new(Vec::new())) @@ -2626,18 +2640,27 @@ impl PutObjectCommitBarrier { Self { state } } - pub(crate) async fn wait_until_paused(&self) { + pub async fn wait_until_paused(&self) { tokio::time::timeout(Duration::from_secs(30), self.state.arrived.notified()) .await .expect("put object should reach the deterministic commit barrier"); } - pub(crate) fn release(&self) { + pub fn release(&self) { self.state.release.notify_one(); } + + pub async fn release_and_wait_until_namespace_pending(&self) { + assert_eq!(self.state.pause, PutObjectCommitPause::BeforeNamespace); + let namespace_pending = self.state.namespace_pending.notified(); + self.release(); + tokio::time::timeout(Duration::from_secs(5), namespace_pending) + .await + .expect("put object should wait for the namespace lock after leaving the commit barrier"); + } } -#[cfg(test)] +#[cfg(any(test, feature = "test-util"))] impl Drop for PutObjectCommitBarrier { fn drop(&mut self) { self.state.release.notify_one(); @@ -2649,7 +2672,7 @@ impl Drop for PutObjectCommitBarrier { } } -#[cfg(test)] +#[cfg(any(test, feature = "test-util"))] async fn pause_put_object_commit(bucket: &str, object: &str, pause: PutObjectCommitPause) { let barrier = PUT_OBJECT_COMMIT_BARRIER .get_or_init(|| std::sync::Mutex::new(Vec::new())) @@ -2664,6 +2687,22 @@ async fn pause_put_object_commit(bucket: &str, object: &str, pause: PutObjectCom } } +#[cfg(any(test, feature = "test-util"))] +fn notify_put_object_commit_namespace_pending(bucket: &str, object: &str) { + let barrier = PUT_OBJECT_COMMIT_BARRIER + .get_or_init(|| std::sync::Mutex::new(Vec::new())) + .lock() + .expect("put object commit barrier mutex should not poison") + .iter() + .find(|barrier| { + barrier.bucket == bucket && barrier.object == object && barrier.pause == PutObjectCommitPause::BeforeNamespace + }) + .cloned(); + if let Some(barrier) = barrier { + barrier.namespace_pending.notify_one(); + } +} + #[cfg(test)] struct DeleteObjectCommitBarrierState { bucket: String, @@ -4414,7 +4453,7 @@ impl crate::storage_api_contracts::object::ObjectOperations for SetDisks { self.invalidate_get_object_metadata_cache(bucket, object).await; // Guard lock for metadata update - #[cfg(test)] + #[cfg(any(test, feature = "test-util"))] pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeMetadata).await; let _lock_guard = if !opts.no_lock { Some(self.acquire_write_lock_diag("put_object_metadata", bucket, object).await?) diff --git a/crates/ecstore/src/store/mod.rs b/crates/ecstore/src/store/mod.rs index d1a7ac850..8fcf6b065 100644 --- a/crates/ecstore/src/store/mod.rs +++ b/crates/ecstore/src/store/mod.rs @@ -152,7 +152,10 @@ mod list; pub(crate) mod list_objects; mod multipart; mod object; -pub use object::PreparedGetObjectReader; +pub use object::{ + PrepareSelectObjectSnapshotError, PreparedGetObjectReader, SelectObjectSnapshot, SelectObjectSnapshotReadError, + SnapshotConsistencyError, +}; mod peer; mod rebalance; pub(crate) mod utils; diff --git a/crates/ecstore/src/store/object.rs b/crates/ecstore/src/store/object.rs index bf8c091b3..85d01d249 100644 --- a/crates/ecstore/src/store/object.rs +++ b/crates/ecstore/src/store/object.rs @@ -37,7 +37,11 @@ use crate::set_disk::{ get_lock_acquire_timeout, get_object_lock_diag_slow_acquire_threshold, get_object_lock_diag_slow_hold_threshold, is_lock_optimization_enabled, is_object_lock_diag_enabled, }; -use crate::storage_api_contracts::object::{ObjectIO as _, ObjectOperations as _}; +use crate::storage_api_contracts::{ + namespace::NamespaceLocking as _, + object::{ObjectIO as _, ObjectOperations as _}, +}; +use parking_lot::Mutex as ParkingMutex; use rustfs_io_metrics::{ record_object_lock_diag_acquire_duration, record_object_lock_diag_hold_duration, record_object_lock_diag_slow_acquire, record_object_lock_diag_slow_hold, @@ -45,6 +49,7 @@ use rustfs_io_metrics::{ use std::{ fmt, pin::Pin, + sync::atomic::{AtomicBool, Ordering}, task::{Context, Poll}, time::{Duration, Instant}, }; @@ -361,6 +366,10 @@ impl ObjectLockDiagGuard { rustfs_lock::NamespaceLockGuard::Fast(_) => None, } } + + fn is_lock_lost(&self) -> bool { + self.guard.is_lock_lost() + } } /// Opaque write-lock guard for the RestoreObject accept path; see @@ -411,6 +420,291 @@ impl Drop for ObjectLockDiagGuard { } } +/// A failure to preserve one object generation for a SelectObjectContent read. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum SnapshotConsistencyError { + #[error("namespace locking is disabled for SelectObjectContent")] + LockingDisabled, + #[error("SelectObjectContent namespace lock was lost")] + LockLost, + #[error("object read semantics changed while SelectObjectContent was running")] + ObjectChanged, +} + +/// Failure while creating a SelectObjectContent snapshot. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum PrepareSelectObjectSnapshotError { + #[error("storage failed while preparing SelectObjectContent snapshot: {0}")] + Storage(#[source] StorageError), + #[error("SelectObjectContent snapshot consistency failure: {0}")] + Consistency(#[source] SnapshotConsistencyError), + #[error("SelectObjectContent object has invalid logical size {size}")] + InvalidLogicalSize { size: i64 }, +} + +impl From for PrepareSelectObjectSnapshotError { + fn from(error: StorageError) -> Self { + Self::Storage(error) + } +} + +impl From for PrepareSelectObjectSnapshotError { + fn from(error: SnapshotConsistencyError) -> Self { + Self::Consistency(error) + } +} + +/// Failure while opening a reader from a SelectObjectContent snapshot. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum SelectObjectSnapshotReadError { + #[error("storage failed while opening SelectObjectContent snapshot reader: {0}")] + Storage(#[source] StorageError), + #[error("SelectObjectContent snapshot consistency failure: {0}")] + Consistency(#[source] SnapshotConsistencyError), +} + +impl From for SelectObjectSnapshotReadError { + fn from(error: StorageError) -> Self { + Self::Storage(error) + } +} + +impl From for SelectObjectSnapshotReadError { + fn from(error: SnapshotConsistencyError) -> Self { + Self::Consistency(error) + } +} + +struct SelectObjectSnapshotLease { + guards: Vec, + lost: Arc, + lock_loss: tokio::sync::watch::Sender, + _monitor_shutdown: tokio::sync::watch::Sender<()>, +} + +impl SelectObjectSnapshotLease { + fn new(guards: Vec) -> Self { + let signals = guards.iter().filter_map(ObjectLockDiagGuard::lock_lost_signal); + let mut waits = signals + .map(|signal| async move { signal.notified().await }) + .collect::>(); + let lost = Arc::new(AtomicBool::new(false)); + let (lock_loss, _) = tokio::sync::watch::channel(false); + let (monitor_shutdown, mut shutdown_rx) = tokio::sync::watch::channel(()); + if !waits.is_empty() { + let task_lost = Arc::clone(&lost); + let task_lock_loss = lock_loss.clone(); + tokio::spawn(async move { + tokio::select! { + lost = futures::StreamExt::next(&mut waits) => { + if lost.is_some() { + task_lost.store(true, Ordering::Release); + task_lock_loss.send_replace(true); + } + } + _ = shutdown_rx.changed() => {} + } + }); + } + Self { + guards, + lost, + lock_loss, + _monitor_shutdown: monitor_shutdown, + } + } + + fn check(&self) -> std::result::Result<(), SnapshotConsistencyError> { + if self.is_lost() || self.guards.iter().any(ObjectLockDiagGuard::is_lock_lost) { + self.lost.store(true, Ordering::Release); + self.lock_loss.send_replace(true); + return Err(SnapshotConsistencyError::LockLost); + } + Ok(()) + } + + fn is_lost(&self) -> bool { + self.lost.load(Ordering::Acquire) + } + + fn subscribe_lock_loss(&self) -> tokio::sync::watch::Receiver { + self.lock_loss.subscribe() + } +} + +/// Opaque, lock-backed object generation used by SelectObjectContent. +pub struct SelectObjectSnapshot { + pool: Arc, + bucket: String, + object: String, + headers: HeaderMap, + opts: ObjectOptions, + object_info: ObjectInfo, + logical_size: u64, + read_semantics_identity: [u8; 32], + first_metadata: ParkingMutex>, + lease: Arc, +} + +impl fmt::Debug for SelectObjectSnapshot { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct("SelectObjectSnapshot") + .field("bucket", &self.bucket) + .field("object", &self.object) + .field("logical_size", &self.logical_size) + .finish_non_exhaustive() + } +} + +impl SelectObjectSnapshot { + pub fn is_for(&self, bucket: &str, object: &str) -> bool { + self.bucket == bucket && self.object == encode_dir_object(object) + } + + pub fn object_info(&self) -> &ObjectInfo { + &self.object_info + } + + pub fn logical_size(&self) -> u64 { + self.logical_size + } + + pub fn matches_version(&self, requested: &str) -> bool { + select_snapshot_version_matches(self.object_info.version_id, requested) + } + + pub fn ensure_valid(&self) -> std::result::Result<(), SnapshotConsistencyError> { + self.lease.check() + } + + #[cfg(feature = "test-util")] + pub fn mark_lost_for_test(&self) { + self.lease.lost.store(true, Ordering::Release); + self.lease.lock_loss.send_replace(true); + } + + pub async fn open_reader( + &self, + range: Option, + ) -> std::result::Result { + self.lease.check()?; + let first_metadata = self.first_metadata.lock().take(); + let metadata = match first_metadata { + Some(metadata) => metadata, + None => { + self.pool + .prepare_get_object_reader_metadata(&self.bucket, &self.object, &self.opts) + .await? + } + }; + self.lease.check()?; + if metadata.read_semantics_identity() != self.read_semantics_identity { + return Err(SnapshotConsistencyError::ObjectChanged.into()); + } + + let mut reader = + crate::object_api::without_get_object_body_cache_hook(self.pool.get_object_reader_with_prepared_metadata( + &self.bucket, + &self.object, + range, + self.headers.clone(), + &self.opts, + metadata, + )) + .await?; + self.lease.check()?; + reader.body_source = crate::object_api::GetObjectBodySource::HookMissed; + reader.stream = Box::new(SelectObjectSnapshotReader { + inner: reader.stream, + lock_loss_wake: SelectObjectSnapshotLockLossWake::new(self.lease.subscribe_lock_loss()), + lease: Arc::clone(&self.lease), + }); + Ok(reader) + } +} + +fn select_snapshot_version_matches(actual: Option, requested: &str) -> bool { + let requested = requested.trim(); + let requested = if requested.eq_ignore_ascii_case("null") { + Uuid::nil() + } else if let Ok(requested) = Uuid::parse_str(requested) { + requested + } else { + return false; + }; + actual.unwrap_or_else(Uuid::nil) == requested +} + +struct SelectObjectSnapshotReader { + inner: Box, + lock_loss_wake: SelectObjectSnapshotLockLossWake, + lease: Arc, +} + +struct SelectObjectSnapshotLockLossWake { + stream: tokio_stream::wrappers::WatchStream, +} + +impl SelectObjectSnapshotLockLossWake { + fn new(receiver: tokio::sync::watch::Receiver) -> Self { + Self { + stream: tokio_stream::wrappers::WatchStream::new(receiver), + } + } + + fn poll_lost(&mut self, cx: &mut Context<'_>) -> bool { + loop { + match futures::Stream::poll_next(Pin::new(&mut self.stream), cx) { + Poll::Ready(Some(true)) => return true, + Poll::Ready(Some(false)) => {} + Poll::Ready(None) | Poll::Pending => return false, + } + } + } +} + +fn select_object_ssec_headers(headers: &HeaderMap) -> HeaderMap { + use rustfs_utils::http::headers::{SSEC_ALGORITHM_HEADER, SSEC_KEY_HEADER, SSEC_KEY_MD5_HEADER}; + + let mut selected = HeaderMap::new(); + for name in [SSEC_ALGORITHM_HEADER, SSEC_KEY_HEADER, SSEC_KEY_MD5_HEADER] { + if let Some(value) = headers.get(name) { + selected.insert(name, value.clone()); + } + } + selected +} + +// LockRegistry clones its canonical client Arc for each endpoint host, so an +// exact Arc set identifies one distributed namespace-lock quorum domain. +fn same_distributed_lock_domain(left: &[Arc], right: &[Arc]) -> bool { + left.iter() + .all(|left_client| right.iter().any(|right_client| Arc::ptr_eq(left_client, right_client))) + && right + .iter() + .all(|right_client| left.iter().any(|left_client| Arc::ptr_eq(left_client, right_client))) +} + +impl AsyncRead for SelectObjectSnapshotReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + if self.lock_loss_wake.poll_lost(cx) || self.lease.is_lost() { + return Poll::Ready(Err(std::io::Error::other(SnapshotConsistencyError::LockLost))); + } + let filled_before = buf.filled().len(); + let poll = Pin::new(&mut self.inner).poll_read(cx, buf); + let reached_eof = matches!(&poll, Poll::Ready(Ok(()))) && buf.filled().len() == filled_before; + if self.lease.is_lost() || (reached_eof && self.lease.check().is_err()) { + buf.set_filled(filled_before); + return Poll::Ready(Err(std::io::Error::other(SnapshotConsistencyError::LockLost))); + } + poll + } +} + fn log_object_lock_acquire_if_slow( op: &'static str, bucket: &str, @@ -790,6 +1084,66 @@ impl ECStore { /// This is an additive two-stage counterpart to `get_object_reader`. The /// existing method remains the compatibility path for callers that do not /// need a pre-reader decision point. + pub async fn prepare_select_object_snapshot( + &self, + bucket: &str, + object: &str, + headers: &HeaderMap, + opts: &ObjectOptions, + ) -> std::result::Result { + check_get_obj_args(bucket, object)?; + + let object = encode_dir_object(object); + let mut opts = opts.clone(); + opts.no_lock = false; + opts.metadata_cache_safe = false; + let read_lock_guards = self.acquire_select_object_read_locks(bucket, &object, &mut opts).await?; + if self.ctx.lock_manager().is_disabled() { + return Err(SnapshotConsistencyError::LockingDisabled.into()); + } + if read_lock_guards.iter().any(ObjectLockDiagGuard::is_lock_lost) { + return Err(SnapshotConsistencyError::LockLost.into()); + } + + let pool = if self.single_pool() { + Arc::clone(&self.pools[0]) + } else { + let (_, pool_idx) = self.get_latest_object_info_with_idx(bucket, &object, &opts).await?; + self.pools.get(pool_idx).cloned().ok_or_else(|| { + StorageError::other(format!("resolved SelectObjectContent pool index {pool_idx} is out of bounds")) + })? + }; + let mut metadata = pool.prepare_get_object_reader_metadata(bucket, &object, &opts).await?; + if read_lock_guards.iter().any(ObjectLockDiagGuard::is_lock_lost) { + return Err(SnapshotConsistencyError::LockLost.into()); + } + if let Some(error) = latest_object_access_delete_marker_error(bucket, &object, metadata.object_info(), &opts) { + return Err(error.into()); + } + + let logical_size_i64 = metadata.object_info().get_actual_size().map_err(StorageError::from)?; + let logical_size = u64::try_from(logical_size_i64) + .map_err(|_| PrepareSelectObjectSnapshotError::InvalidLogicalSize { size: logical_size_i64 })?; + let read_semantics_identity = metadata.read_semantics_identity(); + let object_info = metadata.take_object_info(); + if read_lock_guards.iter().any(ObjectLockDiagGuard::is_lock_lost) { + return Err(SnapshotConsistencyError::LockLost.into()); + } + + Ok(SelectObjectSnapshot { + pool, + bucket: bucket.to_owned(), + object, + headers: select_object_ssec_headers(headers), + opts, + object_info, + logical_size, + read_semantics_identity, + first_metadata: ParkingMutex::new(Some(metadata)), + lease: Arc::new(SelectObjectSnapshotLease::new(read_lock_guards)), + }) + } + pub async fn prepare_get_object_reader( &self, bucket: &str, @@ -984,6 +1338,67 @@ impl ECStore { ))) } + async fn acquire_select_object_read_locks( + &self, + bucket: &str, + object: &str, + opts: &mut ObjectOptions, + ) -> Result> { + let diag_enabled = is_object_lock_diag_enabled(); + let mut guards = Vec::with_capacity(self.pools.len() + 1); + + // Lock order is the store fixed domain first, then pool index ascending + // for each object's hashed set. DELETE and same-key CopyObject use the + // fixed domain, while PUT commits and data movement use the hashed set. + let distributed = self.ctx.is_dist_erasure().await; + if let Some(guard) = self + .acquire_object_read_lock_if_needed("select_object", bucket, object, opts) + .await? + { + guards.push(guard); + } + let fixed_set = Arc::clone(&self.pools[0].disk_set[0]); + let mut locked_sets = vec![fixed_set]; + + for pool in &self.pools { + let hashed_set = pool.get_disks_by_key(object); + let lock_domain_already_held = !distributed + || locked_sets + .iter() + .any(|locked_set| same_distributed_lock_domain(&locked_set.lockers, &hashed_set.lockers)); + if lock_domain_already_held { + continue; + } + let ns_lock = hashed_set.new_ns_lock(bucket, object).await?; + let acquire_start = Instant::now(); + let guard = ns_lock + .get_read_lock(get_lock_acquire_timeout()) + .await + .map_err(|err| Self::map_namespace_lock_error(bucket, object, "read", err))?; + let owner = diag_enabled.then(|| ns_lock.owner().to_string()); + log_object_lock_acquire_if_slow( + "select_object", + bucket, + object, + owner.as_deref(), + ObjectLockDiagMode::Read, + acquire_start.elapsed(), + diag_enabled, + ); + guards.push(ObjectLockDiagGuard::new( + guard, + diag_enabled, + "select_object", + diag_enabled.then(|| bucket.to_string()), + diag_enabled.then(|| object.to_string()), + owner, + ObjectLockDiagMode::Read, + )); + locked_sets.push(hashed_set); + } + Ok(guards) + } + fn attach_read_lock_guard(mut reader: GetObjectReader, guard: Option) -> GetObjectReader { if is_lock_optimization_enabled() || reader.buffered_body.is_some() { return reader; @@ -1608,9 +2023,7 @@ impl ECStore { return Err(StorageError::BucketNotFound(bucket.to_string())); } #[cfg(test)] - if current_bucket_incarnation_id.is_some() { - pause_delete_after_object_lock_snapshot(bucket).await; - } + pause_delete_after_object_lock_snapshot(bucket).await; if opts.delete_prefix && !opts.delete_prefix_object { // Prefix deletes cover multiple object keys; an exact lock on the prefix string @@ -2334,32 +2747,171 @@ mod tests { use super::*; use crate::bucket::lifecycle::core::TRANSITION_COMPLETE; use crate::bucket::lifecycle::tier_sweeper::TierDeleteJournalState; + use crate::bucket::metadata_sys::ObjectLockConfigState; use crate::bucket::replication::{ ReplicationState, ReplicationStatusType, VersionPurgeStatusType, replication_state_to_filemeta, replication_statuses_map, version_purge_statuses_map, }; - use crate::ecstore_validation_blackbox::make_local_set_disks; + use crate::core::sets::make_local_two_set_sets_with_ctx; + use crate::ecstore_validation_blackbox::{make_local_set_disks, make_local_set_disks_with_ctx}; use crate::layout::{ - endpoints::{Endpoints, PoolEndpoints}, + endpoints::{Endpoints, PoolEndpoints, SetupType}, format::FormatV3, }; use crate::object_api::{ GetObjectBodyCacheHook, GetObjectBodyCacheHookLookup, GetObjectBodySource, clear_get_object_body_cache_hook, lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, }; - use crate::set_disk::SetDisks; + use crate::set_disk::{SetDisks, disk_call_counters}; use crate::storage_api_contracts::bucket::MakeBucketOptions; use crate::storage_api_contracts::lifecycle::TransitionedObject; use bytes::Bytes; use std::io::Cursor; use std::sync::Arc; - use std::sync::atomic::{AtomicUsize, Ordering}; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; use tokio::io::AsyncReadExt; + struct WaitForLockLossReader { + inner: Cursor>, + poll_started: Option>, + resume: ParkingMutex>, + } + + impl AsyncRead for WaitForLockLossReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + let poll = Pin::new(&mut self.inner).poll_read(cx, buf); + if let Some(poll_started) = self.poll_started.take() { + let _ = poll_started.send(()); + if self.resume.lock().recv_timeout(Duration::from_secs(10)).is_err() { + return Poll::Ready(Err(std::io::Error::new( + std::io::ErrorKind::TimedOut, + "lock-loss signal was not observed during read", + ))); + } + } + poll + } + } + + struct PermanentlyPendingReader { + poll_started: Option>, + } + + impl AsyncRead for PermanentlyPendingReader { + fn poll_read(mut self: Pin<&mut Self>, _cx: &mut Context<'_>, _buf: &mut ReadBuf<'_>) -> Poll> { + if let Some(poll_started) = self.poll_started.take() { + let _ = poll_started.send(()); + } + Poll::Pending + } + } + struct CountingMissHook { calls: AtomicUsize, } + #[derive(Debug)] + struct RefreshFailureLockClient { + inner: LocalClient, + fail_refresh: AtomicBool, + } + + #[async_trait::async_trait] + impl rustfs_lock::LockClient for RefreshFailureLockClient { + async fn acquire_lock(&self, request: &rustfs_lock::LockRequest) -> rustfs_lock::Result { + rustfs_lock::LockClient::acquire_lock(&self.inner, request).await + } + + async fn release(&self, lock_id: &rustfs_lock::LockId) -> rustfs_lock::Result { + rustfs_lock::LockClient::release(&self.inner, lock_id).await + } + + async fn refresh(&self, lock_id: &rustfs_lock::LockId) -> rustfs_lock::Result { + if self.fail_refresh.load(Ordering::Acquire) { + return Ok(false); + } + rustfs_lock::LockClient::refresh(&self.inner, lock_id).await + } + + async fn force_release(&self, lock_id: &rustfs_lock::LockId) -> rustfs_lock::Result { + rustfs_lock::LockClient::force_release(&self.inner, lock_id).await + } + + async fn check_status(&self, lock_id: &rustfs_lock::LockId) -> rustfs_lock::Result> { + rustfs_lock::LockClient::check_status(&self.inner, lock_id).await + } + + async fn get_stats(&self) -> rustfs_lock::Result { + rustfs_lock::LockClient::get_stats(&self.inner).await + } + + async fn close(&self) -> rustfs_lock::Result<()> { + rustfs_lock::LockClient::close(&self.inner).await + } + + async fn is_online(&self) -> bool { + rustfs_lock::LockClient::is_online(&self.inner).await + } + + async fn is_local(&self) -> bool { + rustfs_lock::LockClient::is_local(&self.inner).await + } + } + + async fn refresh_failure_test_guard( + owner: &'static str, + ) -> ( + ObjectLockDiagGuard, + Arc, + Arc, + ) { + let manager = Arc::new(rustfs_lock::GlobalLockManager::Enabled(Arc::new( + rustfs_lock::FastObjectLockManager::new(), + ))); + let client = Arc::new(RefreshFailureLockClient { + inner: LocalClient::with_manager(manager), + fail_refresh: AtomicBool::new(false), + }); + let namespace_lock = rustfs_lock::NamespaceLock::with_clients_and_quorum( + owner.to_string(), + vec![Arc::clone(&client) as Arc], + 1, + ); + let request = + rustfs_lock::LockRequest::new(rustfs_lock::ObjectKey::new("bucket", "object"), rustfs_lock::LockType::Shared, owner) + .with_ttl(Duration::from_secs(2)) + .with_refresh_interval(Duration::from_millis(20)); + let guard = namespace_lock + .acquire_guard(&request) + .await + .expect("distributed read lock acquisition should succeed") + .expect("distributed read lock should reach quorum"); + let guard = ObjectLockDiagGuard::new( + guard, + false, + "SelectObjectContent", + Some("bucket".to_string()), + Some("object".to_string()), + Some(owner.to_string()), + ObjectLockDiagMode::Read, + ); + let signal = guard + .lock_lost_signal() + .expect("a distributed guard should expose a lock-loss signal"); + (guard, signal, client) + } + + async fn refresh_failure_test_lease( + owner: &'static str, + ) -> ( + Arc, + Arc, + Arc, + ) { + let (guard, signal, client) = refresh_failure_test_guard(owner).await; + (Arc::new(SelectObjectSnapshotLease::new(vec![guard])), signal, client) + } + #[async_trait::async_trait] impl GetObjectBodyCacheHook for CountingMissHook { async fn lookup(&self, _bucket: &str, _object: &str, _info: &ObjectInfo) -> Option { @@ -2370,6 +2922,189 @@ mod tests { struct BodyCacheHookGuard; + #[test] + fn select_snapshot_deduplicates_same_distributed_lock_domain() { + let first: Arc = Arc::new(rustfs_lock::LocalClient::new()); + let second: Arc = Arc::new(rustfs_lock::LocalClient::new()); + let other: Arc = Arc::new(rustfs_lock::LocalClient::new()); + + assert!(same_distributed_lock_domain( + &[Arc::clone(&first), Arc::clone(&second)], + &[Arc::clone(&second), Arc::clone(&first)] + )); + assert!(same_distributed_lock_domain( + &[Arc::clone(&first), Arc::clone(&first), Arc::clone(&second)], + &[Arc::clone(&second), Arc::clone(&first)] + )); + assert!(!same_distributed_lock_domain(&[first, second], &[other])); + } + + #[test] + fn select_snapshot_version_matching_normalizes_null_and_uuid_forms() { + let nil = Uuid::nil(); + for actual in [None, Some(nil)] { + assert!(select_snapshot_version_matches(actual, "null")); + assert!(select_snapshot_version_matches(actual, "NULL")); + assert!(select_snapshot_version_matches(actual, &nil.to_string())); + } + + let version = Uuid::new_v4(); + assert!(select_snapshot_version_matches(Some(version), &version.to_string().to_uppercase())); + assert!(!select_snapshot_version_matches(Some(version), "null")); + assert!(!select_snapshot_version_matches(Some(version), "not-a-version")); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn select_snapshot_reader_rolls_back_bytes_when_lock_is_lost_during_poll() { + let (lease, signal, client) = refresh_failure_test_lease("select-snapshot-lock-loss").await; + + let (poll_started_tx, poll_started_rx) = tokio::sync::oneshot::channel(); + let (resume_tx, resume_rx) = std::sync::mpsc::channel(); + let release_client = Arc::clone(&client); + let release_lease = Arc::clone(&lease); + let release_signal = Arc::clone(&signal); + let release_task = tokio::spawn(async move { + poll_started_rx.await.expect("reader poll should start"); + release_client.fail_refresh.store(true, Ordering::Release); + tokio::time::timeout(Duration::from_secs(5), release_signal.notified()) + .await + .expect("heartbeat should observe the rejected refresh"); + tokio::time::timeout(Duration::from_secs(5), async { + while !release_lease.is_lost() { + tokio::task::yield_now().await; + } + }) + .await + .expect("snapshot monitor should publish lock loss"); + resume_tx.send(()).expect("reader poll should still be waiting"); + }); + + let mut reader = SelectObjectSnapshotReader { + inner: Box::new(WaitForLockLossReader { + inner: Cursor::new(b"must-not-escape".to_vec()), + poll_started: Some(poll_started_tx), + resume: ParkingMutex::new(resume_rx), + }), + lock_loss_wake: SelectObjectSnapshotLockLossWake::new(lease.subscribe_lock_loss()), + lease, + }; + let mut output = Vec::new(); + let error = reader + .read_to_end(&mut output) + .await + .expect_err("a read that loses its lease must fail"); + release_task.await.expect("backend release task should not panic"); + + assert!(signal.is_lost(), "heartbeat should report the rejected refresh"); + assert!(output.is_empty(), "bytes read during the failed poll must be rolled back"); + assert_eq!(error.kind(), std::io::ErrorKind::Other); + } + + #[tokio::test(flavor = "current_thread")] + async fn select_snapshot_reader_checks_guards_at_eof_before_monitor_runs() { + let (guard, signal, client) = refresh_failure_test_guard("select-snapshot-eof-fence").await; + client.fail_refresh.store(true, Ordering::Release); + tokio::time::timeout(Duration::from_secs(5), signal.notified()) + .await + .expect("heartbeat should observe the rejected refresh"); + + let lease = Arc::new(SelectObjectSnapshotLease::new(vec![guard])); + assert!(!lease.is_lost(), "current-thread monitor must not run before the synchronous read"); + let mut reader = SelectObjectSnapshotReader { + inner: Box::new(Cursor::new(b"old-generation".to_vec())), + lock_loss_wake: SelectObjectSnapshotLockLossWake::new(lease.subscribe_lock_loss()), + lease, + }; + let mut output = Vec::new(); + + let error = reader + .read_to_end(&mut output) + .await + .expect_err("EOF fence must reject a lease lost before its monitor is scheduled"); + + assert_eq!(output, b"old-generation"); + assert_eq!(error.kind(), std::io::ErrorKind::Other); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn select_snapshot_lock_loss_is_broadcast_to_remaining_pending_readers() { + let (first_guard, first_signal, first_client) = refresh_failure_test_guard("select-snapshot-pending-first-lock").await; + let (second_guard, second_signal, second_client) = + refresh_failure_test_guard("select-snapshot-pending-second-lock").await; + let lease = Arc::new(SelectObjectSnapshotLease::new(vec![first_guard, second_guard])); + let dropped_reader = SelectObjectSnapshotReader { + inner: Box::new(PermanentlyPendingReader { poll_started: None }), + lock_loss_wake: SelectObjectSnapshotLockLossWake::new(lease.subscribe_lock_loss()), + lease: Arc::clone(&lease), + }; + drop(dropped_reader); + + let (first_started_tx, first_started_rx) = tokio::sync::oneshot::channel(); + let (second_started_tx, second_started_rx) = tokio::sync::oneshot::channel(); + let mut first_reader = SelectObjectSnapshotReader { + inner: Box::new(PermanentlyPendingReader { + poll_started: Some(first_started_tx), + }), + lock_loss_wake: SelectObjectSnapshotLockLossWake::new(lease.subscribe_lock_loss()), + lease: Arc::clone(&lease), + }; + let mut second_reader = SelectObjectSnapshotReader { + inner: Box::new(PermanentlyPendingReader { + poll_started: Some(second_started_tx), + }), + lock_loss_wake: SelectObjectSnapshotLockLossWake::new(lease.subscribe_lock_loss()), + lease, + }; + let first_read_task = tokio::spawn(async move { + let mut byte = [0_u8; 1]; + first_reader.read(&mut byte).await + }); + let second_read_task = tokio::spawn(async move { + let mut byte = [0_u8; 1]; + second_reader.read(&mut byte).await + }); + + first_started_rx.await.expect("first inner reader should reach Poll::Pending"); + second_started_rx + .await + .expect("second inner reader should reach Poll::Pending"); + second_client.fail_refresh.store(true, Ordering::Release); + let (first_result, second_result) = tokio::join!( + tokio::time::timeout(Duration::from_secs(5), first_read_task), + tokio::time::timeout(Duration::from_secs(5), second_read_task), + ); + for result in [first_result, second_result] { + let error = result + .expect("lock loss should wake every reader whose inner I/O remains pending") + .expect("reader task should not panic") + .expect_err("lost snapshot lease must fail every pending read"); + assert_eq!(error.kind(), std::io::ErrorKind::Other); + } + + assert!(!first_client.fail_refresh.load(Ordering::Acquire)); + assert!(!first_signal.is_lost()); + assert!(second_signal.is_lost()); + } + + #[test] + fn select_snapshot_retains_only_ssec_headers() { + use rustfs_utils::http::headers::{SSEC_ALGORITHM_HEADER, SSEC_KEY_HEADER, SSEC_KEY_MD5_HEADER}; + + let mut headers = HeaderMap::new(); + headers.insert(SSEC_ALGORITHM_HEADER, "AES256".parse().expect("valid SSE-C algorithm header")); + headers.insert(SSEC_KEY_HEADER, "secret-key".parse().expect("valid SSE-C key header")); + headers.insert(SSEC_KEY_MD5_HEADER, "key-md5".parse().expect("valid SSE-C key digest header")); + headers.insert("authorization", "credential".parse().expect("valid authorization header")); + + let selected = select_object_ssec_headers(&headers); + + assert_eq!(selected.len(), 3); + assert_eq!(selected.get(SSEC_ALGORITHM_HEADER), headers.get(SSEC_ALGORITHM_HEADER)); + assert_eq!(selected.get(SSEC_KEY_HEADER), headers.get(SSEC_KEY_HEADER)); + assert_eq!(selected.get(SSEC_KEY_MD5_HEADER), headers.get(SSEC_KEY_MD5_HEADER)); + assert!(selected.get("authorization").is_none()); + } + #[test] fn tier_delete_entry_is_prepared_and_bound_to_source_generation() { let identity = [9_u8; 32]; @@ -3048,6 +3783,13 @@ mod tests { } async fn new_prepared_reader_test_store(set_disks: &[Arc]) -> ECStore { + new_prepared_reader_test_store_with_ctx(set_disks, crate::runtime::instance::bootstrap_ctx()).await + } + + async fn new_prepared_reader_test_store_with_ctx( + set_disks: &[Arc], + ctx: Arc, + ) -> ECStore { let mut pool_configs = Vec::with_capacity(set_disks.len()); let mut pools = Vec::with_capacity(set_disks.len()); @@ -3065,31 +3807,50 @@ mod tests { platform: "test".to_string(), }; let disks = set_disks.disks.read().await.clone(); - let pool = Sets::new(disks, &pool_config, &set_disks.format, pool_idx, set_disks.default_parity_count) - .await - .expect("prepared-reader test pool should be created from local disks"); + let pool = Sets::new_with_instance_ctx( + disks, + &pool_config, + &set_disks.format, + pool_idx, + set_disks.default_parity_count, + Arc::clone(&ctx), + ) + .await + .expect("prepared-reader test pool should be created from local disks"); pool_configs.push(pool_config); pools.push(pool); } + new_prepared_reader_test_store_from_pools(pools, pool_configs, ctx) + } + + fn new_prepared_reader_test_store_from_pools( + pools: Vec>, + pool_configs: Vec, + ctx: Arc, + ) -> ECStore { let endpoint_pools = EndpointServerPools::from(pool_configs); ECStore { id: Uuid::new_v4(), disk_map: HashMap::new(), pools, - peer_sys: S3PeerSys::new(&endpoint_pools), + peer_sys: S3PeerSys::new_with_instance_ctx(&endpoint_pools, Arc::clone(&ctx)), pool_meta: RwLock::new(PoolMeta::default()), rebalance_meta: RwLock::new(None), decommission_cancelers: RwLock::new(Vec::new()), start_gate: Mutex::new(()), pool_meta_save_gate: Mutex::new(()), - ctx: crate::runtime::instance::bootstrap_ctx(), + ctx, bucket_fence_registry: std::sync::Arc::default(), } } async fn assert_prepared_reader_blocks_writer(store: &ECStore, bucket: &str, object: &str) { - let manager = Arc::clone(store.pools[0].disk_set[0].local_lock_manager_for_test()); + assert_pool_writer_is_blocked(store, 0, bucket, object).await; + } + + async fn assert_pool_writer_is_blocked(store: &ECStore, pool_idx: usize, bucket: &str, object: &str) { + let manager = Arc::clone(store.pools[pool_idx].get_disks_by_key(object).local_lock_manager_for_test()); let lock = rustfs_lock::NamespaceLock::with_local_manager("prepared-reader-writer".to_string(), manager); let err = lock .get_write_lock(rustfs_lock::ObjectKey::new(bucket, object), "competing-writer", Duration::from_millis(50)) @@ -3099,7 +3860,16 @@ mod tests { } async fn acquire_prepared_reader_writer(store: &ECStore, bucket: &str, object: &str) -> rustfs_lock::NamespaceLockGuard { - let manager = Arc::clone(store.pools[0].disk_set[0].local_lock_manager_for_test()); + acquire_pool_writer(store, 0, bucket, object).await + } + + async fn acquire_pool_writer( + store: &ECStore, + pool_idx: usize, + bucket: &str, + object: &str, + ) -> rustfs_lock::NamespaceLockGuard { + let manager = Arc::clone(store.pools[pool_idx].get_disks_by_key(object).local_lock_manager_for_test()); let lock = rustfs_lock::NamespaceLock::with_local_manager("prepared-reader-writer".to_string(), manager); lock.get_write_lock(rustfs_lock::ObjectKey::new(bucket, object), "competing-writer", Duration::from_secs(1)) .await @@ -3213,6 +3983,370 @@ mod tests { .await; } + #[tokio::test] + #[serial_test::serial(body_cache_hook)] + async fn select_snapshot_holds_namespace_lock_independent_of_get_optimization() { + temp_env::async_with_vars([(rustfs_config::ENV_OBJECT_LOCK_OPTIMIZATION_ENABLE, Some("true"))], async { + let (_dirs, set_disks) = make_local_set_disks(4, 2).await; + let store = new_prepared_reader_test_store(&[set_disks]).await; + let bucket = "select-snapshot-lock-lifetime"; + let object = "object.bin"; + let payload = b"select-snapshot-lock-lifetime-payload-".repeat(40_000); + let put_opts = ObjectOptions { + no_lock: true, + ..Default::default() + }; + + store.pools[0] + .make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created"); + let mut put_reader = PutObjReader::from_vec(payload.clone()); + store.pools[0] + .put_object(bucket, object, &mut put_reader, &put_opts) + .await + .expect("object should be written"); + + use rustfs_utils::http::headers::{SSEC_ALGORITHM_HEADER, SSEC_KEY_HEADER, SSEC_KEY_MD5_HEADER}; + let mut request_headers = HeaderMap::new(); + request_headers.insert(SSEC_ALGORITHM_HEADER, "AES256".parse().expect("valid SSE-C algorithm header")); + request_headers.insert(SSEC_KEY_HEADER, "secret-key".parse().expect("valid SSE-C key header")); + request_headers.insert(SSEC_KEY_MD5_HEADER, "key-md5".parse().expect("valid SSE-C key digest header")); + request_headers.insert("authorization", "credential".parse().expect("valid authorization header")); + let snapshot = store + .prepare_select_object_snapshot(bucket, object, &request_headers, &ObjectOptions::default()) + .await + .expect("SelectObjectContent snapshot should be prepared"); + assert_eq!(snapshot.headers.len(), 3); + assert_eq!(snapshot.headers.get(SSEC_ALGORITHM_HEADER), request_headers.get(SSEC_ALGORITHM_HEADER)); + assert_eq!(snapshot.headers.get(SSEC_KEY_HEADER), request_headers.get(SSEC_KEY_HEADER)); + assert_eq!(snapshot.headers.get(SSEC_KEY_MD5_HEADER), request_headers.get(SSEC_KEY_MD5_HEADER)); + assert!(snapshot.headers.get("authorization").is_none()); + assert_eq!( + snapshot.logical_size(), + u64::try_from(payload.len()).expect("test payload length should fit in u64") + ); + assert_eq!( + snapshot.object_info().size, + i64::try_from(payload.len()).expect("test payload length should fit in i64") + ); + assert_prepared_reader_blocks_writer(&store, bucket, object).await; + + let mut reader = snapshot.open_reader(None).await.expect("snapshot reader should open"); + let mut restored = Vec::new(); + reader + .stream + .read_to_end(&mut restored) + .await + .expect("snapshot body should stream"); + assert_eq!(restored, payload); + drop(reader); + assert_prepared_reader_blocks_writer(&store, bucket, object).await; + + drop(snapshot); + drop(acquire_prepared_reader_writer(&store, bucket, object).await); + }) + .await; + } + + #[tokio::test] + async fn select_snapshot_identity_compares_encoded_directory_object_key() { + let (_dirs, set_disks) = make_local_set_disks(4, 2).await; + let store = new_prepared_reader_test_store(&[set_disks]).await; + let bucket = "select-snapshot-directory-identity"; + let source_object = "source.bin"; + + store.pools[0] + .make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created"); + let mut put_reader = PutObjReader::from_vec(b"source".to_vec()); + store.pools[0] + .put_object( + bucket, + source_object, + &mut put_reader, + &ObjectOptions { + no_lock: true, + ..Default::default() + }, + ) + .await + .expect("source object should be written"); + + let mut snapshot = store + .prepare_select_object_snapshot(bucket, source_object, &HeaderMap::new(), &ObjectOptions::default()) + .await + .expect("source object snapshot should be prepared"); + snapshot.object = encode_dir_object("directory/"); + + assert!(snapshot.is_for(bucket, "directory/")); + assert!(!snapshot.is_for(bucket, "different/")); + } + + #[tokio::test] + #[serial_test::serial(body_cache_hook)] + async fn select_snapshot_reuses_initial_metadata_fanout_for_first_reader() { + let (_dirs, set_disks) = make_local_set_disks(4, 2).await; + let store = new_prepared_reader_test_store(&[set_disks]).await; + let bucket = "select-snapshot-initial-metadata"; + let object = "initial-metadata.bin"; + let payload = b"select snapshot initial metadata".repeat(4_000); + let no_lock_opts = ObjectOptions { + no_lock: true, + ..Default::default() + }; + + store.pools[0] + .make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created"); + let mut put_reader = PutObjReader::from_vec(payload.clone()); + store.pools[0] + .put_object(bucket, object, &mut put_reader, &no_lock_opts) + .await + .expect("object should be written"); + + let calls = disk_call_counters::observe(object); + let snapshot = store + .prepare_select_object_snapshot(bucket, object, &HeaderMap::new(), &ObjectOptions::default()) + .await + .expect("snapshot should prepare metadata"); + assert_eq!( + calls.total(disk_call_counters::KIND_READ_VERSION), + 4, + "snapshot preparation should fan out metadata once" + ); + + let mut reader = snapshot.open_reader(None).await.expect("first snapshot reader should open"); + assert_eq!( + calls.total(disk_call_counters::KIND_READ_VERSION), + 4, + "first reader should consume the metadata captured during preparation" + ); + let mut restored = Vec::new(); + reader + .stream + .read_to_end(&mut restored) + .await + .expect("snapshot body should stream"); + assert_eq!(restored, payload); + } + + #[tokio::test] + #[serial_test::serial(body_cache_hook)] + async fn select_snapshot_rejects_identity_change_before_second_reader() { + let (_dirs, set_disks) = make_local_set_disks(4, 2).await; + let store = new_prepared_reader_test_store(&[Arc::clone(&set_disks)]).await; + let bucket = "select-snapshot-identity-change"; + let object = "identity-change.bin"; + let no_lock_opts = ObjectOptions { + no_lock: true, + ..Default::default() + }; + + set_disks + .make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created"); + let mut initial_reader = PutObjReader::from_vec(b"first generation".to_vec()); + set_disks + .put_object(bucket, object, &mut initial_reader, &no_lock_opts) + .await + .expect("initial generation should be written"); + let mut snapshot = store + .prepare_select_object_snapshot(bucket, object, &HeaderMap::new(), &ObjectOptions::default()) + .await + .expect("snapshot should be prepared"); + drop( + snapshot + .open_reader(None) + .await + .expect("first snapshot reader should consume captured metadata"), + ); + snapshot.opts.metadata_cache_safe = false; + + let mut replacement_reader = PutObjReader::from_vec(b"replacement generation".to_vec()); + set_disks + .put_object(bucket, object, &mut replacement_reader, &no_lock_opts) + .await + .expect("test-only no-lock write should replace the object"); + let error = match snapshot.open_reader(None).await { + Ok(_) => panic!("a later reader must reject the replacement generation"), + Err(error) => error, + }; + assert!(matches!( + error, + SelectObjectSnapshotReadError::Consistency(SnapshotConsistencyError::ObjectChanged) + )); + } + + #[tokio::test] + #[serial_test::serial(body_cache_hook)] + async fn select_snapshot_rejects_latest_versioned_delete_marker_during_prepare() { + let (_first_dirs, first_set) = make_local_set_disks(4, 2).await; + let (_second_dirs, second_set) = make_local_set_disks(4, 2).await; + let store = new_prepared_reader_test_store(&[Arc::clone(&first_set), Arc::clone(&second_set)]).await; + let bucket = "select-snapshot-latest-delete-marker"; + let object = "versioned-object.bin"; + let versioned_opts = ObjectOptions { + no_lock: true, + versioned: true, + object_lock_config_snapshot: Some(Arc::new(ObjectLockConfigSnapshot::new(ObjectLockConfigState::ConfirmedAbsent))), + ..Default::default() + }; + + for set_disks in [&first_set, &second_set] { + set_disks + .make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created"); + } + let mut older_reader = PutObjReader::from_vec(b"older visible generation".to_vec()); + first_set + .put_object(bucket, object, &mut older_reader, &versioned_opts) + .await + .expect("older versioned object should be written"); + let mut put_reader = PutObjReader::from_vec(b"hidden generation".to_vec()); + second_set + .put_object(bucket, object, &mut put_reader, &versioned_opts) + .await + .expect("versioned object should be written"); + let marker = second_set + .delete_object(bucket, object, versioned_opts.clone()) + .await + .expect("versioned delete should create a marker"); + assert!(marker.delete_marker); + assert!(marker.version_id.is_some_and(|version_id| !version_id.is_nil())); + + let error = store + .prepare_select_object_snapshot( + bucket, + object, + &HeaderMap::new(), + &ObjectOptions { + versioned: true, + ..Default::default() + }, + ) + .await + .expect_err("latest delete marker should hide the prior object generation"); + assert!(matches!( + error, + PrepareSelectObjectSnapshotError::Storage(ref error) if is_err_object_not_found(error) + )); + } + + #[tokio::test] + #[serial_test::serial(storage_class_env)] + async fn select_snapshot_blocks_store_delete_for_object_in_nonzero_set() { + let ctx = Arc::new(crate::runtime::instance::InstanceContext::new()); + let (_dirs, sets) = make_local_two_set_sets_with_ctx(Arc::clone(&ctx)).await; + ctx.update_erasure_type(SetupType::DistErasure).await; + assert!( + sets.disk_set[0] + .lockers + .iter() + .all(|first| sets.disk_set[1].lockers.iter().all(|second| !Arc::ptr_eq(first, second))) + ); + let pool_config = sets.endpoints.clone(); + let store = Arc::new(new_prepared_reader_test_store_from_pools(vec![Arc::clone(&sets)], vec![pool_config], ctx)); + let bucket = RUSTFS_META_BUCKET; + let object = (0..1_000) + .map(|index| format!("nonzero-set-{index}.bin")) + .find(|candidate| Arc::ptr_eq(&sets.get_disks_by_key(candidate), &sets.disk_set[1])) + .expect("a key should hash to the second set"); + let mut put_reader = PutObjReader::from_vec(b"stable snapshot body".to_vec()); + sets.put_object( + bucket, + &object, + &mut put_reader, + &ObjectOptions { + no_lock: true, + ..Default::default() + }, + ) + .await + .expect("object should be written to the second set"); + assert!(Arc::ptr_eq(&sets.get_disks_by_key(&object), &sets.disk_set[1])); + + let snapshot = store + .prepare_select_object_snapshot(bucket, &object, &HeaderMap::new(), &ObjectOptions::default()) + .await + .expect("snapshot should acquire both lock domains"); + let hashed_set_writer = rustfs_lock::NamespaceLock::with_clients_and_quorum( + "select-nonzero-set-writer".to_string(), + sets.disk_set[1].lockers.clone(), + 2, + ); + let writer_error = hashed_set_writer + .get_write_lock( + rustfs_lock::ObjectKey::new(bucket, object.as_str()), + "competing-hashed-set-writer", + Duration::from_millis(50), + ) + .await + .expect_err("Select must hold the nonzero hashed-set read lock"); + assert!(matches!(writer_error, rustfs_lock::LockError::Timeout { .. })); + let barrier = DeleteAfterObjectLockSnapshotBarrier::install(bucket); + let delete_store = Arc::clone(&store); + let delete_object = object.clone(); + let mut delete = tokio::spawn(async move { + delete_store + .delete_object(bucket, &delete_object, ObjectOptions::default()) + .await + }); + barrier.wait_until_paused().await; + barrier.release(); + assert!( + tokio::time::timeout(Duration::from_millis(100), &mut delete).await.is_err(), + "store DELETE must wait for the fixed-domain Select read lock" + ); + + drop(snapshot); + let deleted = tokio::time::timeout(Duration::from_secs(10), delete) + .await + .expect("DELETE should resume after the snapshot is dropped") + .expect("DELETE task should not panic") + .expect("DELETE should complete"); + assert_eq!(deleted.name, object); + } + + #[tokio::test] + #[serial_test::serial] + async fn select_snapshot_fails_closed_when_local_locking_is_disabled() { + let ctx = Arc::new(crate::runtime::instance::InstanceContext::with_lock_manager_for_test(Arc::new( + rustfs_lock::GlobalLockManager::Disabled(rustfs_lock::DisabledLockManager::new()), + ))); + let (_dirs, set_disks) = make_local_set_disks_with_ctx(4, 2, Arc::clone(&ctx)).await; + let store = new_prepared_reader_test_store_with_ctx(&[set_disks], ctx).await; + let bucket = "select-snapshot-lock-disabled"; + let object = "object.bin"; + let mut put_reader = PutObjReader::from_vec(b"payload".to_vec()); + let no_lock_opts = ObjectOptions { + no_lock: true, + ..Default::default() + }; + + store.pools[0] + .make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created"); + store.pools[0] + .put_object(bucket, object, &mut put_reader, &no_lock_opts) + .await + .expect("object should be written without namespace locking"); + + let error = store + .prepare_select_object_snapshot(bucket, object, &HeaderMap::new(), &ObjectOptions::default()) + .await + .expect_err("SelectObjectContent must reject a disabled lock manager"); + assert!(matches!( + error, + PrepareSelectObjectSnapshotError::Consistency(SnapshotConsistencyError::LockingDisabled) + )); + } + #[tokio::test] #[serial_test::serial(body_cache_hook)] async fn prepared_object_info_releases_namespace_lock_immediately() { @@ -3286,6 +4420,53 @@ mod tests { assert_eq!(restored, payload); } + #[tokio::test] + #[serial_test::serial] + async fn select_snapshot_locks_the_hashed_set_in_every_pool() { + let (_first_dirs, first_set) = make_local_set_disks(4, 2).await; + let (_second_dirs, second_set) = make_local_set_disks(4, 2).await; + let store = new_prepared_reader_test_store(&[first_set, second_set]).await; + let bucket = "select-snapshot-all-pools"; + let object = "object.bin"; + let payload = b"second-pool-snapshot".to_vec(); + let no_lock_opts = ObjectOptions { + no_lock: true, + ..Default::default() + }; + + for pool in &store.pools { + pool.make_bucket(bucket, &MakeBucketOptions::default()) + .await + .expect("bucket should be created in each pool"); + } + let mut put_reader = PutObjReader::from_vec(payload.clone()); + store.pools[1] + .put_object(bucket, object, &mut put_reader, &no_lock_opts) + .await + .expect("object should be written only to the second pool"); + + let snapshot = store + .prepare_select_object_snapshot(bucket, object, &HeaderMap::new(), &ObjectOptions::default()) + .await + .expect("snapshot should resolve the second-pool object"); + assert_pool_writer_is_blocked(&store, 0, bucket, object).await; + assert_pool_writer_is_blocked(&store, 1, bucket, object).await; + + let mut reader = snapshot.open_reader(None).await.expect("snapshot reader should open"); + let mut restored = Vec::new(); + reader + .stream + .read_to_end(&mut restored) + .await + .expect("snapshot body should stream"); + assert_eq!(restored, payload); + drop(reader); + drop(snapshot); + + drop(acquire_pool_writer(&store, 0, bucket, object).await); + drop(acquire_pool_writer(&store, 1, bucket, object).await); + } + // Phase 5 Slice 2 (backlog#939): the instance context flows down the whole // object graph — ECStore, its Sets, and their SetDisks must all carry the // same `Arc` in a single-instance deployment. diff --git a/crates/s3select-api/Cargo.toml b/crates/s3select-api/Cargo.toml index 8a3165d98..a68ce1a1d 100644 --- a/crates/s3select-api/Cargo.toml +++ b/crates/s3select-api/Cargo.toml @@ -77,11 +77,12 @@ parking_lot.workspace = true tokio = { workspace = true, features = ["fs", "rt-multi-thread"] } tokio-util = { workspace = true, features = ["io", "compat"] } tracing.workspace = true +uuid.workspace = true transform-stream.workspace = true url.workspace = true [dev-dependencies] -rustfs-test-utils.workspace = true +rustfs-test-utils = { workspace = true, features = ["put-object-commit-barrier"] } serial_test.workspace = true [lib] diff --git a/crates/s3select-api/src/lib.rs b/crates/s3select-api/src/lib.rs index 2c8e84dfc..5c1436b7f 100644 --- a/crates/s3select-api/src/lib.rs +++ b/crates/s3select-api/src/lib.rs @@ -15,22 +15,23 @@ #![recursion_limit = "256"] use datafusion::{common::DataFusionError, sql::sqlparser::parser::ParserError}; -use std::fmt::Display; +use std::{error::Error as StdError, fmt::Display}; use thiserror::Error; pub mod object_store; pub mod query; pub mod server; mod storage_api; +pub use storage_api::SelectObjectSnapshot; #[cfg(test)] mod test; pub type QueryResult = Result; pub(crate) use storage_api::crate_boundary::{ - SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError, - SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, - select_is_err_version_not_found, + PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions, + SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError, resolve_select_object_store_handle, + select_is_err_bucket_not_found, select_is_err_object_not_found, select_is_err_version_not_found, }; #[derive(Debug, Error)] @@ -79,24 +80,24 @@ pub enum S3SelectPolicyError { QueryTimeout { seconds: u64 }, } -impl S3SelectPolicyError { - fn from_error<'a>(mut err: &'a (dyn std::error::Error + 'static)) -> Option<&'a Self> { +impl QueryError { + fn source_error(&self) -> Option<&T> { + let mut err: &(dyn StdError + 'static) = self; for _ in 0..16 { - if let Some(policy_error) = err.downcast_ref::() { - return Some(policy_error); + if let Some(source) = err.downcast_ref::() { + return Some(source); } err = err.source()?; } None } -} -impl QueryError { + pub fn is_snapshot_consistency_error(&self) -> bool { + self.source_error::().is_some() + } + pub fn s3_select_policy_error(&self) -> Option<&S3SelectPolicyError> { - match self { - Self::Datafusion { source } => S3SelectPolicyError::from_error(source.as_ref()), - _ => None, - } + self.source_error() } } @@ -230,6 +231,17 @@ mod tests { )); } + #[test] + fn snapshot_consistency_error_is_recoverable_without_string_matching() { + let err = QueryError::Datafusion { + source: Box::new(DataFusionError::External(Box::new(SelectObjectSnapshotReadError::Consistency( + SnapshotConsistencyError::LockLost, + )))), + }; + + assert!(err.is_snapshot_consistency_error()); + } + #[test] fn test_query_error_from_parser_error() { let parser_error = ParserError::ParserError("syntax error".to_string()); diff --git a/crates/s3select-api/src/object_store.rs b/crates/s3select-api/src/object_store.rs index 536a1f3bd..5a9080390 100644 --- a/crates/s3select-api/src/object_store.rs +++ b/crates/s3select-api/src/object_store.rs @@ -13,8 +13,8 @@ // limitations under the License. use crate::{ - SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError, - SelectStore, + PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions, + SelectObjectSnapshot, SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError, query::{ parser::RustFsDialect, session::{QueryExecutionGuard, QueryExecutionTracker}, @@ -24,7 +24,7 @@ use crate::{ }; use async_trait::async_trait; use bytes::Bytes; -use chrono::Utc; +use chrono::{DateTime, Utc}; use datafusion::{ common::{DataFusionError, runtime::SpawnedTask}, execution::memory_pool::{MemoryConsumer, MemoryPool, UnboundedMemoryPool}, @@ -43,33 +43,26 @@ use futures_core::stream::BoxStream; use http::{HeaderMap, HeaderValue, header::HeaderName}; use parking_lot::Mutex; use rustfs_common::DEFAULT_DELIMITER; -use s3s::S3Result; -use s3s::dto::SelectObjectContentInput; use s3s::header::{ X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, }; -use s3s::s3_error; +use s3s::{S3Error, S3ErrorCode, S3Result, dto::SelectObjectContentInput}; use std::collections::VecDeque; use std::ops::Range; use std::sync::Arc; -use tokio::io::AsyncReadExt; +#[cfg(test)] +use std::sync::atomic::{AtomicUsize, Ordering}; +use tokio::{io::AsyncReadExt, sync::OnceCell}; use tokio_util::io::ReaderStream; use transform_stream::AsyncTryStream; -use crate::storage_api::object_store::{HTTPRangeSpec, ObjectIO as _, ObjectOperations as _}; +use crate::storage_api::object_store::HTTPRangeSpec; fn select_default_read_buffer_size_u64() -> u64 { u64::try_from(SELECT_DEFAULT_READ_BUFFER_SIZE).unwrap_or(u64::MAX) } -fn validated_object_size(size: i64) -> Result { - u64::try_from(size).map_err(|err| o_Error::Generic { - store: "EcObjectStore", - source: Box::new(err), - }) -} - /// Maximum allowed object size for JSON DOCUMENT mode. /// /// JSON DOCUMENT format requires loading the entire file into memory for DOM @@ -94,7 +87,6 @@ pub const INVALID_SCAN_RANGE_MESSAGE: &str = const NORMALIZED_RECORD_DELIMITER: &[u8] = b"\r\n"; const NORMALIZED_FIELD_DELIMITER: &[u8] = &[DEFAULT_DELIMITER]; -#[derive(Debug)] pub struct EcObjectStore { input: Arc, need_convert: bool, @@ -109,38 +101,18 @@ pub struct EcObjectStore { json_sub_path: Option, memory_pool: Arc, query_tracker: Option, - - store: Arc, + store: Option>, + snapshot: OnceCell>, + #[cfg(test)] + reader_open_count: Arc, } -#[cfg(test)] -struct ScanRangeBeforeMainHook { - bucket: String, - object: String, - reached: tokio::sync::oneshot::Sender<()>, - resume: tokio::sync::oneshot::Receiver<()>, -} - -#[cfg(test)] -static SCAN_RANGE_BEFORE_MAIN_HOOK: tokio::sync::Mutex> = tokio::sync::Mutex::const_new(None); - -#[cfg(test)] -async fn run_scan_range_before_main_hook(bucket: &str, object: &str) { - let hook = { - let mut hook = SCAN_RANGE_BEFORE_MAIN_HOOK.lock().await; - if hook - .as_ref() - .is_some_and(|hook| hook.bucket == bucket && hook.object == object) - { - hook.take() - } else { - None - } - }; - if let Some(hook) = hook { - let _ = hook.reached.send(()); - let _ = hook.resume.await; - } +#[derive(Debug, thiserror::Error)] +pub(crate) enum EcObjectStoreBuildError { + #[error("ec store not inited")] + StoreUnavailable, + #[error("SelectObjectContent snapshot consistency failure: {0}")] + Snapshot(#[source] SnapshotConsistencyError), } #[derive(Clone, Copy, Debug)] @@ -168,20 +140,55 @@ pub struct InvalidScanRange; impl EcObjectStore { pub fn new(input: Arc) -> S3Result { - Self::build(input, Arc::new(UnboundedMemoryPool::default()), None, None) + Self::build_lazy(input, Arc::new(UnboundedMemoryPool::default()), None).map_err(map_build_error_to_s3) } - pub(crate) fn new_with_memory_pool(input: Arc, memory_pool: Arc) -> S3Result { - Self::build(input, memory_pool, None, None) + pub fn new_with_snapshot(input: Arc, snapshot: Arc) -> S3Result { + Self::build_with_snapshot(input, Arc::new(UnboundedMemoryPool::default()), None, snapshot).map_err(map_build_error_to_s3) + } + + pub(crate) fn new_with_memory_pool( + input: Arc, + memory_pool: Arc, + snapshot: Option>, + ) -> std::result::Result { + match snapshot { + Some(snapshot) => Self::build_with_snapshot(input, memory_pool, None, snapshot), + None => Self::build_lazy(input, memory_pool, None), + } } pub(crate) fn new_with_query_tracker( input: Arc, memory_pool: Arc, query_tracker: QueryExecutionTracker, - store: Option>, - ) -> S3Result { - Self::build(input, memory_pool, Some(query_tracker), store) + snapshot: Option>, + ) -> std::result::Result { + match snapshot { + Some(snapshot) => Self::build_with_snapshot(input, memory_pool, Some(query_tracker), snapshot), + None => Self::build_lazy(input, memory_pool, Some(query_tracker)), + } + } + + fn build_lazy( + input: Arc, + memory_pool: Arc, + query_tracker: Option, + ) -> std::result::Result { + let store = resolve_select_object_store_handle().ok_or(EcObjectStoreBuildError::StoreUnavailable)?; + Ok(Self::build(input, memory_pool, query_tracker, Some(store), None)) + } + + fn build_with_snapshot( + input: Arc, + memory_pool: Arc, + query_tracker: Option, + snapshot: Arc, + ) -> std::result::Result { + if !snapshot.is_for(&input.bucket, &input.key) { + return Err(EcObjectStoreBuildError::Snapshot(SnapshotConsistencyError::ObjectChanged)); + } + Ok(Self::build(input, memory_pool, query_tracker, None, Some(snapshot))) } fn build( @@ -189,11 +196,8 @@ impl EcObjectStore { memory_pool: Arc, query_tracker: Option, store: Option>, - ) -> S3Result { - let Some(store) = store.or_else(resolve_select_object_store_handle) else { - return Err(s3_error!(InternalError, "ec store not inited")); - }; - + snapshot: Option>, + ) -> Self { let (need_convert, delimiter) = if let Some(csv) = input.request.input_serialization.csv.as_ref() { if let Some(delimiter) = csv.field_delimiter.as_ref() { if delimiter.len() > 1 { @@ -227,7 +231,7 @@ impl EcObjectStore { None }; - Ok(Self { + Self { input, need_convert, delimiter, @@ -236,20 +240,15 @@ impl EcObjectStore { memory_pool, query_tracker, store, - }) - } - - fn object_options(&self, options: &GetOptions) -> SelectObjectOptions { - SelectObjectOptions { - version_id: options.version.clone(), - ..Default::default() + snapshot: match snapshot { + Some(snapshot) => OnceCell::new_with(Some(snapshot)), + None => OnceCell::new(), + }, + #[cfg(test)] + reader_open_count: Arc::new(AtomicUsize::new(0)), } } - fn read_headers(&self) -> HeaderMap { - select_read_headers(&self.input) - } - fn scan_range(&self, object_size: u64) -> Result> { let Some(scan_range) = self.input.request.scan_range.as_ref() else { return Ok(None); @@ -283,37 +282,60 @@ impl EcObjectStore { .is_some_and(|info| matches!(info.as_str(), "USE" | "IGNORE")) } - async fn object_info(&self, opts: &SelectObjectOptions) -> Result { - self.store - .get_object_info(&self.input.bucket, &self.input.key, opts) - .await - .map_err(|err| map_storage_error(&self.input.bucket, &self.input.key, err)) + async fn snapshot(&self, version: Option<&str>) -> Result<&Arc> { + let snapshot = self + .snapshot + .get_or_try_init(|| async { + let store = self.store.as_ref().ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: "prepared snapshot is unavailable".into(), + })?; + let opts = SelectObjectOptions { + version_id: version.map(|version| { + let version = version.trim(); + if version.eq_ignore_ascii_case("null") { + uuid::Uuid::nil().to_string() + } else { + version.to_owned() + } + }), + ..Default::default() + }; + let snapshot = store + .prepare_select_object_snapshot(&self.input.bucket, &self.input.key, &select_read_headers(&self.input), &opts) + .await + .map_err(|err| map_prepare_snapshot_error(&self.input.bucket, &self.input.key, err))?; + Ok::<_, o_Error>(Arc::new(snapshot)) + }) + .await?; + if let Some(requested) = version + && !snapshot.matches_version(requested) + { + return Err(o_Error::Generic { + store: "EcObjectStore", + source: "prepared snapshot is pinned to a different object version".into(), + }); + } + Ok(snapshot) } - async fn object_reader(&self, range: Option, opts: &SelectObjectOptions) -> Result { - let h = self.read_headers(); - self.store - .get_object_reader(&self.input.bucket, &self.input.key, range, h, opts) + async fn object_reader(&self, range: Option) -> Result { + #[cfg(test)] + self.reader_open_count.fetch_add(1, Ordering::Relaxed); + self.snapshot(None) + .await? + .open_reader(range) .await - .map_err(|err| map_storage_error(&self.input.bucket, &self.input.key, err)) + .map_err(|err| snapshot_read_error(&self.input.bucket, &self.input.key, err)) } - async fn read_raw_range_with_opts( - &self, - range: Range, - opts: &SelectObjectOptions, - expected_snapshot: Option<&SelectObjectInfo>, - ) -> Result { + async fn read_raw_range(&self, range: Range) -> Result { if range.is_empty() { return Ok(Bytes::new()); } - let reader = self - .object_reader(Some(http_range_spec_from_range(range.clone())), opts) - .await?; - if let Some(expected_snapshot) = expected_snapshot { - validate_object_snapshot(expected_snapshot, &reader.object_info)?; - } - let object_size = validated_object_size(reader.object_info.size)?; + let snapshot = self.snapshot(None).await?; + let reader = self.object_reader(Some(http_range_spec_from_range(range.clone()))).await?; + let object_size = snapshot.logical_size(); let resolved_range = GetRange::Bounded(range) .as_range(object_size) .map_err(|err| o_Error::Generic { @@ -337,25 +359,14 @@ impl EcObjectStore { Ok(Bytes::from(bytes)) } - async fn read_raw_range(&self, range: Range) -> Result { - self.read_raw_range_with_opts(range, &self.object_options(&GetOptions::new()), None) - .await - } - - async fn read_header_record( - &self, - object_size: u64, - delimiter: &[u8], - opts: &SelectObjectOptions, - expected_snapshot: &SelectObjectInfo, - ) -> Result { + async fn read_header_record(&self, object_size: u64, delimiter: &[u8]) -> Result { if object_size == 0 { return Ok(Bytes::new()); } let mut end = select_default_read_buffer_size_u64().min(object_size); loop { - let bytes = self.read_raw_range_with_opts(0..end, opts, Some(expected_snapshot)).await?; + let bytes = self.read_raw_range(0..end).await?; if let Some(pos) = find_delimiter(&bytes, delimiter) { return Ok(bytes.slice(0..pos + delimiter.len())); } @@ -366,13 +377,7 @@ impl EcObjectStore { } } - async fn scan_range_read_start( - &self, - scan_range: SelectScanRange, - delimiter: &[u8], - opts: &SelectObjectOptions, - expected_snapshot: &SelectObjectInfo, - ) -> Result { + async fn scan_range_read_start(&self, scan_range: SelectScanRange, delimiter: &[u8]) -> Result { let delimiter_len = u64::try_from(delimiter.len()).unwrap_or(u64::MAX); let fallback_start = scan_range.start().saturating_sub(delimiter_len); if delimiter.len() != 2 || delimiter[0] != delimiter[1] || scan_range.start() == 0 { @@ -380,9 +385,7 @@ impl EcObjectStore { } let context_start = scan_range.start().saturating_sub(select_default_read_buffer_size_u64()); - let context = self - .read_raw_range_with_opts(context_start..scan_range.start(), opts, Some(expected_snapshot)) - .await?; + let context = self.read_raw_range(context_start..scan_range.start()).await?; let suffix_len = context.iter().rev().take_while(|byte| **byte == delimiter[0]).count(); if suffix_len == context.len() && context_start > 0 { return Err(o_Error::Generic { @@ -397,6 +400,18 @@ impl EcObjectStore { } } +impl std::fmt::Debug for EcObjectStore { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("EcObjectStore") + .field("bucket", &self.input.bucket) + .field("object", &self.input.key) + .field("need_convert", &self.need_convert) + .field("is_json_document", &self.is_json_document) + .field("json_sub_path", &self.json_sub_path) + .finish_non_exhaustive() + } +} + impl std::fmt::Display for EcObjectStore { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.write_str("EcObjectStore") @@ -469,21 +484,6 @@ fn http_range_spec_from_start(start: u64) -> HTTPRangeSpec { } } -fn validate_object_snapshot(expected: &SelectObjectInfo, actual: &SelectObjectInfo) -> Result<()> { - if expected.size != actual.size - || expected.version_id != actual.version_id - || expected.data_dir != actual.data_dir - || expected.etag != actual.etag - || expected.mod_time != actual.mod_time - { - return Err(o_Error::Generic { - store: "EcObjectStore", - source: "object changed while preparing SelectObjectContent ScanRange".into(), - }); - } - Ok(()) -} - fn find_delimiter(bytes: &[u8], delimiter: &[u8]) -> Option { if delimiter.is_empty() { return None; @@ -491,11 +491,38 @@ fn find_delimiter(bytes: &[u8], delimiter: &[u8]) -> Option { bytes.windows(delimiter.len()).position(|window| window == delimiter) } +fn map_prepare_snapshot_error(bucket: &str, object: &str, err: PrepareSelectObjectSnapshotError) -> o_Error { + match err { + PrepareSelectObjectSnapshotError::Storage(err) => map_storage_error(bucket, object, err), + err => o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + }, + } +} + +fn map_build_error_to_s3(error: EcObjectStoreBuildError) -> S3Error { + let message = error.to_string(); + let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, message); + s3_error.set_source(Box::new(error)); + s3_error +} + +fn snapshot_read_error(bucket: &str, object: &str, err: SelectObjectSnapshotReadError) -> o_Error { + match err { + SelectObjectSnapshotReadError::Storage(err) => map_storage_error(bucket, object, err), + err => o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + }, + } +} + fn map_storage_error(bucket: &str, object: &str, err: SelectStorageError) -> o_Error { if select_is_err_bucket_not_found(&err) || select_is_err_object_not_found(&err) || select_is_err_version_not_found(&err) { return o_Error::NotFound { path: format!("{bucket}/{object}"), - source: err.to_string().into(), + source: Box::new(err), }; } o_Error::Generic { @@ -504,6 +531,17 @@ fn map_storage_error(bucket: &str, object: &str, err: SelectStorageError) -> o_E } } +fn snapshot_last_modified(snapshot: &SelectObjectSnapshot) -> Result> { + let mod_time = snapshot.object_info().mod_time.ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: std::io::Error::new(std::io::ErrorKind::InvalidData, "snapshot metadata has no modification time").into(), + })?; + DateTime::::from_timestamp(mod_time.unix_timestamp(), mod_time.nanosecond()).ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: std::io::Error::new(std::io::ErrorKind::InvalidData, "snapshot modification time is out of range").into(), + }) +} + pub fn scan_range_from_bounds(start: Option, end: Option, object_size: u64) -> Result> { parse_scan_range_from_bounds(start, end, object_size).map_err(|_| invalid_scan_range_store_error()) } @@ -579,51 +617,19 @@ impl ObjectStore for EcObjectStore { } async fn get_opts(&self, location: &Path, options: GetOptions) -> Result { - let opts = self.object_options(&options); - let record_delimiter = if options.head { - None - } else { - self.record_delimiter_for_conversion() + // SelectObjectContent has no version-id input. For direct ObjectStore + // compatibility, a version supplied on the first operation defines + // this instance's immutable snapshot; later operations reuse it. + let snapshot = self.snapshot(options.version.as_deref()).await?; + let original_size = snapshot.logical_size(); + let object_info = snapshot.object_info(); + let meta = ObjectMeta { + location: location.clone(), + last_modified: snapshot_last_modified(snapshot)?, + size: original_size, + e_tag: object_info.etag.clone(), + version: object_info.version_id.map(|version| version.to_string()), }; - let needs_scan_context = options.range.is_none() && !options.head && self.input.request.scan_range.is_some(); - let scan_context = if needs_scan_context { - let source_snapshot = self.object_info(&opts).await?; - let original_size = validated_object_size(source_snapshot.size)?; - if let Some(scan_range) = self.scan_range(original_size)? { - let delimiter = self.record_delimiter(); - let read_start = self - .scan_range_read_start(scan_range, &delimiter, &opts, &source_snapshot) - .await?; - Some((source_snapshot, scan_range, read_start)) - } else { - None - } - } else { - None - }; - - #[cfg(test)] - if scan_context.is_some() { - run_scan_range_before_main_hook(&self.input.bucket, &self.input.key).await; - } - let range = options.range.as_ref().map(http_range_spec_from_get_range); - let reader = if let Some((source_snapshot, _, read_start)) = scan_context.as_ref() { - let range = (source_snapshot.size > 0).then(|| http_range_spec_from_start(*read_start)); - self.object_reader(range, &opts).await? - } else { - self.object_reader(range, &opts).await? - }; - if let Some((source_snapshot, _, _)) = scan_context.as_ref() { - validate_object_snapshot(source_snapshot, &reader.object_info)?; - } - - let original_size = match scan_context.as_ref() { - Some((source_snapshot, _, _)) => validated_object_size(source_snapshot.size)?, - None => validated_object_size(reader.object_info.size)?, - }; - let etag = reader.object_info.etag; - let version = reader.object_info.version_id.map(|version| version.to_string()); - let attributes = Attributes::default(); let result_range = match options.range.as_ref() { Some(range) => range.as_range(original_size).map_err(|err| o_Error::Generic { store: "EcObjectStore", @@ -631,9 +637,38 @@ impl ObjectStore for EcObjectStore { })?, None => 0..original_size, }; - let payload = if options.head { - GetResultPayload::Stream(stream::empty().boxed()) - } else if options.range.is_some() { + if options.head { + return Ok(GetResult { + payload: GetResultPayload::Stream(stream::empty().boxed()), + meta, + range: result_range, + attributes: Attributes::default(), + }); + } + + let record_delimiter = self.record_delimiter_for_conversion(); + let needs_scan_context = options.range.is_none() && self.input.request.scan_range.is_some(); + let scan_context = if needs_scan_context { + if let Some(scan_range) = self.scan_range(original_size)? { + let delimiter = self.record_delimiter(); + let read_start = self.scan_range_read_start(scan_range, &delimiter).await?; + Some((scan_range, read_start)) + } else { + None + } + } else { + None + }; + + let range = options.range.as_ref().map(http_range_spec_from_get_range); + let reader = if let Some((_, read_start)) = scan_context.as_ref() { + let range = (original_size > 0).then(|| http_range_spec_from_start(*read_start)); + self.object_reader(range).await? + } else { + self.object_reader(range).await? + }; + + let payload = if options.range.is_some() { let size = usize::try_from(result_range.end - result_range.start).map_err(|err| o_Error::Generic { store: "EcObjectStore", source: Box::new(err), @@ -663,14 +698,11 @@ impl ObjectStore for EcObjectStore { self.query_tracker.clone(), ); GetResultPayload::Stream(stream) - } else if let Some((source_snapshot, scan_range, read_start)) = scan_context { + } else if let Some((scan_range, read_start)) = scan_context { let delimiter = self.record_delimiter(); let include_header = self.csv_has_header(); let header = if include_header && read_start > 0 { - Some( - self.read_header_record(original_size, &delimiter, &opts, &source_snapshot) - .await?, - ) + Some(self.read_header_record(original_size, &delimiter).await?) } else { None }; @@ -694,10 +726,11 @@ impl ObjectStore for EcObjectStore { self.need_convert.then(|| self.delimiter.clone()), )) } else { - let stream = bytes_stream( - ReaderStream::with_capacity(reader.stream, SELECT_DEFAULT_READ_BUFFER_SIZE), - original_size as usize, - ); + let stream_size = usize::try_from(original_size).map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + })?; + let stream = bytes_stream(ReaderStream::with_capacity(reader.stream, SELECT_DEFAULT_READ_BUFFER_SIZE), stream_size); GetResultPayload::Stream(convert_csv_delimiter_stream( stream, record_delimiter, @@ -705,19 +738,11 @@ impl ObjectStore for EcObjectStore { )) }; - let meta = ObjectMeta { - location: location.clone(), - last_modified: Utc::now(), - size: original_size, - e_tag: etag, - version, - }; - Ok(GetResult { payload, meta, range: result_range, - attributes, + attributes: Attributes::default(), }) } @@ -1312,12 +1337,12 @@ fn incomplete_object_stream_error(remaining: impl std::fmt::Display) -> o_Error #[cfg(test)] mod test { use super::{ - EcObjectStore, JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER, SCAN_RANGE_BEFORE_MAIN_HOOK, SELECT_DEFAULT_READ_BUFFER_SIZE, - ScanRangeBeforeMainHook, SelectScanRange, bytes_stream, convert_csv_delimiter_stream, convert_field_delimiter_stream, - convert_record_delimiter_stream, extract_json_sub_path_from_expression, find_delimiter, flatten_json_document_to_ndjson, - http_range_spec_from_get_range, json_document_ndjson_stream, json_document_ndjson_stream_with_parser, - scan_range_from_bounds, scan_range_stream, select_read_headers, validate_json_document_size, validate_object_snapshot, - validated_object_size, + EcObjectStore, EcObjectStoreBuildError, JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER, OnceCell, + SELECT_DEFAULT_READ_BUFFER_SIZE, SelectObjectOptions, SelectObjectSnapshot, SelectScanRange, SnapshotConsistencyError, + bytes_stream, convert_csv_delimiter_stream, convert_field_delimiter_stream, convert_record_delimiter_stream, + extract_json_sub_path_from_expression, find_delimiter, flatten_json_document_to_ndjson, http_range_spec_from_get_range, + json_document_ndjson_stream, json_document_ndjson_stream_with_parser, scan_range_from_bounds, scan_range_stream, + select_read_headers, snapshot_last_modified, validate_json_document_size, }; use crate::query::session::{QueryExecutionGuard, QueryExecutionOwner, QueryExecutionTracker}; use crate::storage_api::SelectPutObjReader; @@ -1332,6 +1357,9 @@ mod test { prelude::CsvReadOptions, }; use futures::{StreamExt, TryStreamExt, stream}; + use http::HeaderMap; + use rustfs_test_utils::PutObjectCommitBarrier; + use s3s::S3ErrorCode; use s3s::dto::{ CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, OutputSerialization, ScanRange, SelectObjectContentInput, SelectObjectContentRequest, @@ -1345,38 +1373,596 @@ mod test { atomic::{AtomicUsize, Ordering}, }; - #[test] - fn ec_object_store_constructor_remains_source_compatible() { - let _constructor: fn(Arc) -> s3s::S3Result = EcObjectStore::new; - } - use tokio::sync::Semaphore; + use tokio::{io::AsyncReadExt, sync::Semaphore}; - #[test] - fn test_validated_object_size_rejects_negative_metadata() { - assert_eq!(validated_object_size(0).expect("zero object size should be valid"), 0); - assert!(validated_object_size(-1).is_err()); + fn csv_input(bucket: &str, object: &str) -> Arc { + Arc::new(SelectObjectContentInput { + bucket: bucket.to_string(), + expected_bucket_owner: None, + key: object.to_string(), + sse_customer_algorithm: None, + sse_customer_key: None, + sse_customer_key_md5: None, + request: SelectObjectContentRequest { + expression: "SELECT * FROM s3object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + csv: Some(CSVInput::default()), + ..Default::default() + }, + output_serialization: OutputSerialization { + csv: Some(CSVOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + }) } #[test] - fn test_scan_range_snapshot_validation_rejects_changed_object() { - let expected = crate::SelectObjectInfo::default(); - let mut actual = expected.clone(); - assert!(validate_object_snapshot(&expected, &actual).is_ok()); + fn lazy_snapshot_headers_preserve_ssec_context() { + let mut input = (*csv_input("bucket", "object.csv")).clone(); + input.sse_customer_algorithm = Some("AES256".to_string()); + input.sse_customer_key = Some("customer-key".to_string()); + input.sse_customer_key_md5 = Some("customer-key-md5".to_string()); - actual.size = 1; - assert!(validate_object_snapshot(&expected, &actual).is_err()); - actual = expected.clone(); - actual.version_id = Some("00000000-0000-0000-0000-000000000001".parse().expect("valid version UUID")); - assert!(validate_object_snapshot(&expected, &actual).is_err()); - actual = expected.clone(); - actual.data_dir = Some("00000000-0000-0000-0000-000000000002".parse().expect("valid data-dir UUID")); - assert!(validate_object_snapshot(&expected, &actual).is_err()); - actual = expected.clone(); - actual.etag = Some("changed".to_string()); - assert!(validate_object_snapshot(&expected, &actual).is_err()); - actual = expected.clone(); - actual.mod_time = Some(std::time::SystemTime::UNIX_EPOCH.into()); - assert!(validate_object_snapshot(&expected, &actual).is_err()); + let headers = select_read_headers(&input); + + assert_eq!( + headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM) + .and_then(|value| value.to_str().ok()), + Some("AES256") + ); + assert_eq!( + headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY) + .and_then(|value| value.to_str().ok()), + Some("customer-key") + ); + assert_eq!( + headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5) + .and_then(|value| value.to_str().ok()), + Some("customer-key-md5") + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn legacy_constructor_retries_snapshot_after_not_found() { + let env = crate::storage_api::select_test_ecstore_env().await; + let bucket = "s3select-lazy-snapshot-retry"; + let object = "input.csv"; + env.make_bucket(bucket, false).await; + let store = EcObjectStore::new(csv_input(bucket, object)).expect("legacy constructor should resolve the global store"); + + let error = store + .get_opts( + &Path::from(object), + GetOptions { + head: true, + ..Default::default() + }, + ) + .await + .expect_err("missing object should remain a typed not-found error"); + assert!(matches!(error, object_store::Error::NotFound { .. })); + + env.put_object_bytes(bucket, object, b"id,name\n1,Alice\n".to_vec()).await; + let result = store + .get_opts( + &Path::from(object), + GetOptions { + head: true, + ..Default::default() + }, + ) + .await + .expect("failed snapshot initialization must not be cached"); + assert_eq!(result.meta.size, 16); + } + + #[tokio::test] + #[serial_test::serial] + async fn legacy_constructor_maps_missing_bucket_to_not_found() { + let _env = crate::storage_api::select_test_ecstore_env().await; + let bucket = "s3select-lazy-snapshot-missing-bucket"; + let object = "input.csv"; + let store = EcObjectStore::new(csv_input(bucket, object)).expect("legacy constructor should resolve the global store"); + + let error = store + .get_opts( + &Path::from(object), + GetOptions { + head: true, + ..Default::default() + }, + ) + .await + .expect_err("missing bucket should remain a typed not-found error"); + + assert!(matches!(error, object_store::Error::NotFound { .. })); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn legacy_constructor_reuses_head_snapshot_for_body() { + const BUCKET: &str = "s3select-lazy-snapshot-head-body"; + const OBJECT: &str = "input.csv"; + const OLD_DATA: &[u8] = b"id,name\n1,old\n"; + const NEW_DATA: &[u8] = b"id,name\n1,new\n"; + + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + env.put_object_bytes(BUCKET, OBJECT, OLD_DATA.to_vec()).await; + let store = EcObjectStore::new(csv_input(BUCKET, OBJECT)).expect("legacy constructor should resolve the global store"); + let head = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + head: true, + ..Default::default() + }, + ) + .await + .expect("HEAD should lazily prepare the snapshot"); + assert_eq!(head.meta.size, u64::try_from(OLD_DATA.len()).expect("fixture length should fit in u64")); + + let commit_barrier = PutObjectCommitBarrier::before_namespace(BUCKET, OBJECT); + let writer = tokio::spawn(async move { + env.put_object_bytes(BUCKET, OBJECT, NEW_DATA.to_vec()).await; + }); + commit_barrier.wait_until_paused().await; + commit_barrier.release_and_wait_until_namespace_pending().await; + assert!(!writer.is_finished(), "overwrite must wait for the lazy snapshot"); + + let result = store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect("body should reuse the HEAD snapshot"); + let GetResultPayload::Stream(stream) = result.payload else { + panic!("expected streaming snapshot body"); + }; + let bytes = stream + .try_collect::>() + .await + .expect("collect lazy snapshot body") + .concat(); + assert_eq!(bytes, OLD_DATA); + assert_eq!(store.reader_open_count.load(Ordering::Relaxed), 1); + + drop(store); + tokio::time::timeout(std::time::Duration::from_secs(5), writer) + .await + .expect("overwrite should finish after the lazy snapshot is released") + .expect("overwrite task should join"); + assert_eq!(read_current_object(BUCKET, OBJECT).await, NEW_DATA); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn legacy_constructor_first_version_pins_later_reads() { + const BUCKET: &str = "s3select-lazy-snapshot-version"; + const OBJECT: &str = "input.csv"; + const OLD_DATA: &[u8] = b"old-marker\n"; + const NEW_DATA: &[u8] = b"new-poison-value\n"; + + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, true).await; + let versioned_opts = SelectObjectOptions { + versioned: true, + ..Default::default() + }; + let mut old_reader = SelectPutObjReader::from_vec(OLD_DATA.to_vec()); + let old_info = env + .ecstore + .put_object(BUCKET, OBJECT, &mut old_reader, &versioned_opts) + .await + .expect("put old version fixture"); + let old_version = old_info + .version_id + .expect("versioned PUT should return a version ID") + .to_string(); + let mut new_reader = SelectPutObjReader::from_vec(NEW_DATA.to_vec()); + let new_info = env + .ecstore + .put_object(BUCKET, OBJECT, &mut new_reader, &versioned_opts) + .await + .expect("put latest version poison fixture"); + let new_version = new_info + .version_id + .expect("versioned PUT should return a version ID") + .to_string(); + + let store = EcObjectStore::new(csv_input(BUCKET, OBJECT)).expect("legacy constructor should resolve the global store"); + let head = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + head: true, + version: Some(old_version.to_uppercase()), + ..Default::default() + }, + ) + .await + .expect("first HEAD should bind the requested old version"); + assert_eq!(head.meta.version.as_deref(), Some(old_version.as_str())); + + let mismatch = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + head: true, + version: Some(new_version), + ..Default::default() + }, + ) + .await + .expect_err("an explicit different version must not reuse the pinned snapshot"); + assert!(mismatch.to_string().contains("different object version")); + + let range = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + range: Some(GetRange::Bounded(0..3)), + ..Default::default() + }, + ) + .await + .expect("later range should reuse the old-version snapshot"); + let GetResultPayload::Stream(range_stream) = range.payload else { + panic!("expected ranged snapshot stream"); + }; + assert_eq!( + range_stream + .try_collect::>() + .await + .expect("collect old-version range") + .concat(), + b"old" + ); + + let body = store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect("later body should reuse the old-version snapshot"); + let GetResultPayload::Stream(body_stream) = body.payload else { + panic!("expected full snapshot stream"); + }; + assert_eq!( + body_stream + .try_collect::>() + .await + .expect("collect old-version body") + .concat(), + OLD_DATA + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn legacy_constructor_normalizes_null_version_before_snapshot_prepare() { + const BUCKET: &str = "s3select-lazy-snapshot-null-version"; + const OBJECT: &str = "input.csv"; + const DATA: &[u8] = b"null-version-marker\n"; + + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + env.put_object_bytes(BUCKET, OBJECT, DATA.to_vec()).await; + let store = EcObjectStore::new(csv_input(BUCKET, OBJECT)).expect("legacy constructor should resolve the global store"); + + let head = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + head: true, + version: Some("NULL".to_string()), + ..Default::default() + }, + ) + .await + .expect("null version should prepare an unversioned snapshot"); + assert!(head.meta.version.is_none()); + + let body = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + version: Some(uuid::Uuid::nil().to_string()), + ..Default::default() + }, + ) + .await + .expect("nil UUID should match the pinned null-version snapshot"); + let GetResultPayload::Stream(body_stream) = body.payload else { + panic!("expected null-version snapshot stream"); + }; + assert_eq!( + body_stream + .try_collect::>() + .await + .expect("collect null-version body") + .concat(), + DATA + ); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn prepared_snapshot_rejects_a_different_query_object() { + const BUCKET: &str = "s3select-snapshot-identity"; + const OBJECT: &str = "source.csv"; + + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + env.put_object_bytes(BUCKET, OBJECT, b"source-marker\n".to_vec()).await; + let snapshot = Arc::new( + env.ecstore + .prepare_select_object_snapshot(BUCKET, OBJECT, &HeaderMap::new(), &Default::default()) + .await + .expect("prepare source snapshot"), + ); + + let error = EcObjectStore::new_with_snapshot(csv_input(BUCKET, "different.csv"), snapshot) + .expect_err("a snapshot must remain bound to its source object"); + + assert_eq!(error.code(), &S3ErrorCode::InternalError); + assert!(error.source().is_some_and(|source| { + source + .downcast_ref::() + .is_some_and(|error| matches!(error, EcObjectStoreBuildError::Snapshot(SnapshotConsistencyError::ObjectChanged))) + })); + } + + async fn prepare_test_snapshot(bucket: &str, object: &str) -> Arc { + let env = crate::storage_api::select_test_ecstore_env().await; + Arc::new( + env.ecstore + .prepare_select_object_snapshot(bucket, object, &HeaderMap::new(), &Default::default()) + .await + .expect("prepare SelectObjectContent snapshot"), + ) + } + + fn scan_range_csv_store( + bucket: &str, + object: &str, + snapshot: Arc, + record_delimiter: &str, + file_header_info: Option, + start: i64, + end: i64, + ) -> EcObjectStore { + EcObjectStore::new_with_snapshot( + Arc::new(SelectObjectContentInput { + bucket: bucket.to_string(), + expected_bucket_owner: None, + key: object.to_string(), + sse_customer_algorithm: None, + sse_customer_key: None, + sse_customer_key_md5: None, + request: SelectObjectContentRequest { + expression: "SELECT * FROM s3object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + csv: Some(CSVInput { + record_delimiter: Some(record_delimiter.to_string()), + file_header_info, + ..Default::default() + }), + ..Default::default() + }, + output_serialization: OutputSerialization { + csv: Some(CSVOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: Some(ScanRange { + start: Some(start), + end: Some(end), + }), + }, + }), + snapshot, + ) + .expect("snapshot should match SelectObjectContent input") + } + + async fn read_current_object(bucket: &str, object: &str) -> Vec { + let snapshot = prepare_test_snapshot(bucket, object).await; + let mut reader = snapshot.open_reader(None).await.expect("current object reader should open"); + let mut bytes = Vec::new(); + reader + .stream + .read_to_end(&mut bytes) + .await + .expect("current object should be readable"); + bytes + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn head_uses_snapshot_metadata_without_opening_body() { + let env = crate::storage_api::select_test_ecstore_env().await; + let bucket = "s3select-snapshot-head"; + let object = "input.csv"; + env.make_bucket(bucket, false).await; + let mut reader = SelectPutObjReader::from_vec(b"id,name\n1,Alice\n".to_vec()); + env.ecstore + .put_object(bucket, object, &mut reader, &Default::default()) + .await + .expect("put HEAD fixture"); + + let snapshot = prepare_test_snapshot(bucket, object).await; + let expected_modified = snapshot_last_modified(&snapshot).expect("snapshot modification time"); + let expected_size = snapshot.logical_size(); + let expected_etag = snapshot.object_info().etag.clone(); + let expected_version = snapshot.object_info().version_id.map(|version| version.to_string()); + let input = Arc::new(SelectObjectContentInput { + bucket: bucket.to_string(), + expected_bucket_owner: None, + key: object.to_string(), + sse_customer_algorithm: Some("secret-algorithm".to_string()), + sse_customer_key: Some("secret-customer-key".to_string()), + sse_customer_key_md5: Some("secret-customer-key-md5".to_string()), + request: SelectObjectContentRequest { + expression: "SELECT * FROM s3object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + csv: Some(CSVInput::default()), + ..Default::default() + }, + output_serialization: OutputSerialization { + csv: Some(CSVOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + }); + let store = EcObjectStore::new_with_snapshot(input, snapshot).expect("snapshot should match SelectObjectContent input"); + let debug = format!("{store:?}"); + assert!(!debug.contains("secret-customer-key")); + + let result = store + .get_opts( + &Path::from(object), + GetOptions { + head: true, + ..Default::default() + }, + ) + .await + .expect("HEAD from snapshot metadata"); + + assert_eq!(result.meta.last_modified, expected_modified); + assert_eq!(result.meta.size, expected_size); + assert_eq!(result.meta.e_tag, expected_etag); + assert_eq!(result.meta.version, expected_version); + assert_eq!(store.reader_open_count.load(Ordering::Relaxed), 0); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn snapshot_keeps_csv_header_and_body_on_one_generation_during_overwrite() { + const BUCKET: &str = "s3select-snapshot-header-body-race"; + const OBJECT: &str = "input.csv"; + const OLD_DATA: &[u8] = b"old_header,value\nskip_old,0\nold_body,1\n"; + const NEW_DATA: &[u8] = b"new_header,value\nskip_new,0\nnew_body,1\n"; + + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + let mut reader = SelectPutObjReader::from_vec(OLD_DATA.to_vec()); + env.ecstore + .put_object(BUCKET, OBJECT, &mut reader, &Default::default()) + .await + .expect("put old CSV header/body fixture"); + + let selected_start = i64::try_from(b"old_header,value\nskip_old,0\n".len()).expect("fixture offset should fit in i64"); + let snapshot = prepare_test_snapshot(BUCKET, OBJECT).await; + let store = scan_range_csv_store( + BUCKET, + OBJECT, + snapshot, + "\n", + Some(FileHeaderInfo::from_static(FileHeaderInfo::USE)), + selected_start, + selected_start, + ); + let commit_barrier = PutObjectCommitBarrier::before_namespace(BUCKET, OBJECT); + let writer = tokio::spawn(async move { + env.put_object_bytes(BUCKET, OBJECT, NEW_DATA.to_vec()).await; + }); + commit_barrier.wait_until_paused().await; + commit_barrier.release_and_wait_until_namespace_pending().await; + assert!( + !writer.is_finished(), + "overwrite must remain blocked while the SelectObjectContent snapshot is alive" + ); + + let result = store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect("read CSV header and body from one snapshot"); + let GetResultPayload::Stream(stream) = result.payload else { + panic!("expected streaming CSV header/body payload"); + }; + let bytes = stream + .try_collect::>() + .await + .expect("collect CSV header/body snapshot") + .concat(); + + assert_eq!(bytes, b"old_header,value\nold_body,1\n"); + assert_eq!(store.reader_open_count.load(Ordering::Relaxed), 2); + assert!(!writer.is_finished(), "overwrite must remain blocked after both snapshot readers finish"); + + drop(store); + tokio::time::timeout(std::time::Duration::from_secs(5), writer) + .await + .expect("overwrite should finish after the snapshot is released") + .expect("overwrite task should join"); + assert_eq!(read_current_object(BUCKET, OBJECT).await, NEW_DATA); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn snapshot_keeps_scan_range_context_and_main_reader_on_one_generation_during_overwrite() { + const BUCKET: &str = "s3select-snapshot-scan-context-race"; + const OBJECT: &str = "input.csv"; + const OLD_DATA: &[u8] = b"111aaa222aa333aa"; + const NEW_DATA: &[u8] = b"999aaa888aa777aa"; + + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + let mut reader = SelectPutObjReader::from_vec(OLD_DATA.to_vec()); + env.ecstore + .put_object(BUCKET, OBJECT, &mut reader, &Default::default()) + .await + .expect("put old ScanRange context fixture"); + + let snapshot = prepare_test_snapshot(BUCKET, OBJECT).await; + let store = scan_range_csv_store(BUCKET, OBJECT, snapshot, "aa", None, 4, 5); + let commit_barrier = PutObjectCommitBarrier::before_namespace(BUCKET, OBJECT); + let writer = tokio::spawn(async move { + env.put_object_bytes(BUCKET, OBJECT, NEW_DATA.to_vec()).await; + }); + commit_barrier.wait_until_paused().await; + commit_barrier.release_and_wait_until_namespace_pending().await; + assert!( + !writer.is_finished(), + "overwrite must remain blocked while the SelectObjectContent snapshot is alive" + ); + + let result = store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect("read ScanRange context and main body from one snapshot"); + let GetResultPayload::Stream(stream) = result.payload else { + panic!("expected streaming ScanRange context payload"); + }; + let bytes = stream + .try_collect::>() + .await + .expect("collect ScanRange context snapshot") + .concat(); + + assert_eq!(bytes, b"a222\r\n"); + assert_eq!(store.reader_open_count.load(Ordering::Relaxed), 2); + assert!( + !writer.is_finished(), + "overwrite must remain blocked after context and main readers finish" + ); + + drop(store); + tokio::time::timeout(std::time::Duration::from_secs(5), writer) + .await + .expect("overwrite should finish after the snapshot is released") + .expect("overwrite task should join"); + assert_eq!(read_current_object(BUCKET, OBJECT).await, NEW_DATA); } #[tokio::test] @@ -1595,37 +2181,6 @@ mod test { assert_eq!(range.end, 19); } - #[test] - fn test_select_read_headers_preserves_ssec_context() { - let input = SelectObjectContentInput { - bucket: "bucket".to_string(), - expected_bucket_owner: None, - key: "object.csv".to_string(), - sse_customer_algorithm: Some("AES256".to_string()), - sse_customer_key: Some("customer-key".to_string()), - sse_customer_key_md5: Some("customer-key-md5".to_string()), - request: SelectObjectContentRequest { - expression: "SELECT * FROM s3object".to_string(), - expression_type: ExpressionType::from_static(ExpressionType::SQL), - input_serialization: InputSerialization { - csv: Some(CSVInput::default()), - ..Default::default() - }, - output_serialization: OutputSerialization { - csv: Some(CSVOutput::default()), - ..Default::default() - }, - request_progress: None, - scan_range: None, - }, - }; - - let headers = select_read_headers(&input); - assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM).unwrap(), "AES256"); - assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY).unwrap(), "customer-key"); - assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5).unwrap(), "customer-key-md5"); - } - #[tokio::test] async fn test_scan_range_output_can_convert_field_delimiter() { let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"a&&1\nb&&2\n"))]); @@ -1679,7 +2234,7 @@ mod test { .await .expect("put self-overlapping delimiter ScanRange fixture"); - let make_store = |start, end, file_header_info| EcObjectStore { + let make_store = |start, end, file_header_info, snapshot| EcObjectStore { input: Arc::new(SelectObjectContentInput { bucket: bucket.to_string(), expected_bucket_owner: None, @@ -1715,10 +2270,12 @@ mod test { json_sub_path: None, memory_pool: Arc::new(GreedyMemoryPool::new(1024)), query_tracker: None, - store: Arc::clone(&env.ecstore), + store: None, + snapshot: OnceCell::new_with(Some(snapshot)), + reader_open_count: Arc::new(AtomicUsize::new(0)), }; - let store = make_store(6, 6, None); + let store = make_store(6, 6, None, prepare_test_snapshot(bucket, object).await); let result = store .get_opts(&Path::from(object), GetOptions::default()) .await @@ -1728,8 +2285,9 @@ mod test { }; let chunks: Vec = stream.try_collect().await.expect("collect record-data ScanRange output"); assert!(chunks.concat().is_empty()); + drop(store); - let store = make_store(4, 5, None); + let store = make_store(4, 5, None, prepare_test_snapshot(bucket, object).await); let result = store .get_opts(&Path::from(object), GetOptions::default()) .await @@ -1742,13 +2300,14 @@ mod test { .await .expect("collect overlapping-delimiter ScanRange output"); assert_eq!(chunks.concat(), b"a222\r\n"); + drop(store); let mut reader = SelectPutObjReader::from_vec(b"111aa222aa333aa".to_vec()); env.ecstore .put_object(bucket, object, &mut reader, &Default::default()) .await .expect("put exact review ScanRange fixture"); - let result = make_store(6, 6, None) + let result = make_store(6, 6, None, prepare_test_snapshot(bucket, object).await) .get_opts(&Path::from(object), GetOptions::default()) .await .expect("read exact review ScanRange fixture"); @@ -1758,7 +2317,7 @@ mod test { let chunks: Vec = stream.try_collect().await.expect("collect exact review ScanRange output"); assert!(chunks.concat().is_empty()); - let result = make_store(5, 5, None) + let result = make_store(5, 5, None, prepare_test_snapshot(bucket, object).await) .get_opts(&Path::from(object), GetOptions::default()) .await .expect("read ScanRange starting after an even delimiter run"); @@ -1773,40 +2332,27 @@ mod test { .put_object(bucket, object, &mut reader, &Default::default()) .await .expect("put ScanRange header snapshot fixture"); - let result = make_store(8, 8, Some(FileHeaderInfo::from_static(FileHeaderInfo::USE))) - .get_opts(&Path::from(object), GetOptions::default()) - .await - .expect("read ScanRange with a separate header read"); + let result = make_store( + 8, + 8, + Some(FileHeaderInfo::from_static(FileHeaderInfo::USE)), + prepare_test_snapshot(bucket, object).await, + ) + .get_opts(&Path::from(object), GetOptions::default()) + .await + .expect("read ScanRange with a separate header read"); let GetResultPayload::Stream(stream) = result.payload else { panic!("expected streaming ScanRange header payload"); }; let chunks: Vec = stream.try_collect().await.expect("collect ScanRange header output"); assert_eq!(chunks.concat(), b"h1\r\nv2\r\n"); - let header_store = make_store(8, 8, Some(FileHeaderInfo::from_static(FileHeaderInfo::USE))); - let header_opts = header_store.object_options(&GetOptions::new()); - let header_snapshot = header_store - .object_info(&header_opts) - .await - .expect("read header snapshot before overwrite"); - let header_size = validated_object_size(header_snapshot.size).expect("header fixture size should be valid"); - let mut reader = SelectPutObjReader::from_vec(b"q9aaz8aaz7aa".to_vec()); - env.ecstore - .put_object(bucket, object, &mut reader, &Default::default()) - .await - .expect("overwrite header snapshot fixture"); - let err = header_store - .read_header_record(header_size, b"aa", &header_opts, &header_snapshot) - .await - .expect_err("stale header snapshot must fail closed"); - assert!(err.to_string().contains("object changed")); - let mut reader = SelectPutObjReader::from_vec(b"aaa222aa".to_vec()); env.ecstore .put_object(bucket, object, &mut reader, &Default::default()) .await .expect("put object-start delimiter context fixture"); - let result = make_store(3, 3, None) + let result = make_store(3, 3, None, prepare_test_snapshot(bucket, object).await) .get_opts(&Path::from(object), GetOptions::default()) .await .expect("read delimiter context that reaches the object start"); @@ -1825,7 +2371,7 @@ mod test { .await .expect("put large self-overlapping delimiter ScanRange fixture"); let scan_start = i64::try_from(run_start + 3).expect("fixture offset should fit in i64"); - let result = make_store(scan_start, scan_start, None) + let result = make_store(scan_start, scan_start, None, prepare_test_snapshot(bucket, object).await) .get_opts(&Path::from(object), GetOptions::default()) .await .expect("read large ScanRange with bounded delimiter context"); @@ -1844,57 +2390,11 @@ mod test { .put_object(bucket, object, &mut reader, &Default::default()) .await .expect("put oversized delimiter context fixture"); - let err = make_store(scan_start, scan_start, None) + let err = make_store(scan_start, scan_start, None, prepare_test_snapshot(bucket, object).await) .get_opts(&Path::from(object), GetOptions::default()) .await .expect_err("oversized self-overlapping delimiter context must fail closed"); assert!(err.to_string().contains("bounded ScanRange context")); - - let store = make_store(0, 0, None); - let opts = store.object_options(&GetOptions::new()); - let snapshot = store.object_info(&opts).await.expect("read snapshot before overwrite"); - let snapshot_size = usize::try_from(snapshot.size).expect("fixture size should fit in usize"); - let mut reader = SelectPutObjReader::from_vec(vec![b'x'; snapshot_size]); - env.ecstore - .put_object(bucket, object, &mut reader, &Default::default()) - .await - .expect("overwrite snapshot fixture"); - let err = store - .read_raw_range_with_opts(0..1, &opts, Some(&snapshot)) - .await - .expect_err("stale ScanRange snapshot must fail closed"); - assert!(err.to_string().contains("object changed")); - - let original = b"111aa222aa333aa"; - let mut reader = SelectPutObjReader::from_vec(original.to_vec()); - env.ecstore - .put_object(bucket, object, &mut reader, &Default::default()) - .await - .expect("restore context-to-main race fixture"); - let (reached_tx, reached_rx) = tokio::sync::oneshot::channel(); - let (resume_tx, resume_rx) = tokio::sync::oneshot::channel(); - *SCAN_RANGE_BEFORE_MAIN_HOOK.lock().await = Some(ScanRangeBeforeMainHook { - bucket: bucket.to_string(), - object: object.to_string(), - reached: reached_tx, - resume: resume_rx, - }); - let store = make_store(6, 6, None); - let read_task = tokio::spawn(async move { store.get_opts(&Path::from("input.csv"), GetOptions::default()).await }); - reached_rx - .await - .expect("ScanRange read should pause before opening its main reader"); - let mut reader = SelectPutObjReader::from_vec(b"999aa888aa777aa".to_vec()); - env.ecstore - .put_object(bucket, object, &mut reader, &Default::default()) - .await - .expect("overwrite between ScanRange context and main reads"); - resume_tx.send(()).expect("resume ScanRange main read"); - let err = read_task - .await - .expect("ScanRange read task should join") - .expect_err("context-to-main overwrite must fail closed"); - assert!(err.to_string().contains("object changed")); } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] @@ -1947,6 +2447,7 @@ mod test { scan_range: None, }, }); + let snapshot = prepare_test_snapshot(bucket, object).await; let store = Arc::new(EcObjectStore { input, need_convert: false, @@ -1955,7 +2456,9 @@ mod test { json_sub_path: None, memory_pool: Arc::new(GreedyMemoryPool::new(32 * 1024 * 1024)), query_tracker: None, - store: Arc::clone(&env.ecstore), + store: None, + snapshot: OnceCell::new_with(Some(snapshot)), + reader_open_count: Arc::new(AtomicUsize::new(0)), }); let config = SessionConfig::new() @@ -2040,6 +2543,7 @@ mod test { scan_range: None, }, }); + let snapshot = prepare_test_snapshot(bucket, object).await; let store = super::EcObjectStore { input, need_convert: true, @@ -2048,7 +2552,9 @@ mod test { json_sub_path: None, memory_pool: Arc::new(GreedyMemoryPool::new(1024)), query_tracker: None, - store: Arc::clone(&env.ecstore), + store: None, + snapshot: OnceCell::new_with(Some(snapshot)), + reader_open_count: Arc::new(AtomicUsize::new(0)), }; let result = store diff --git a/crates/s3select-api/src/query/dispatcher.rs b/crates/s3select-api/src/query/dispatcher.rs index 7ee78b52e..8a3e8a908 100644 --- a/crates/s3select-api/src/query/dispatcher.rs +++ b/crates/s3select-api/src/query/dispatcher.rs @@ -22,6 +22,7 @@ use super::{ Query, execution::{Output, QueryStateMachine}, logical_planner::Plan, + session::QueryAdmission, }; #[async_trait] @@ -32,6 +33,14 @@ pub trait QueryDispatcher: Send + Sync { async fn execute_query(&self, query: &Query) -> QueryResult; + fn try_reserve_query(&self) -> QueryResult { + Ok(QueryAdmission::unmanaged()) + } + + async fn execute_query_admitted(&self, query: &Query, _admission: QueryAdmission) -> QueryResult { + self.execute_query(query).await + } + async fn build_logical_plan(&self, query_state_machine: Arc) -> QueryResult>; async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc) -> QueryResult; diff --git a/crates/s3select-api/src/query/mod.rs b/crates/s3select-api/src/query/mod.rs index d83af94b5..11fd92ebf 100644 --- a/crates/s3select-api/src/query/mod.rs +++ b/crates/s3select-api/src/query/mod.rs @@ -15,6 +15,8 @@ use s3s::dto::SelectObjectContentInput; use std::sync::Arc; +use crate::SelectObjectSnapshot; + pub mod analyzer; pub mod ast; pub mod dispatcher; @@ -37,12 +39,26 @@ pub struct Context { pub struct Query { context: Context, content: String, + snapshot: Option>, } impl Query { #[inline(always)] pub fn new(context: Context, content: String) -> Self { - Self { context, content } + Self { + context, + content, + snapshot: None, + } + } + + #[inline(always)] + pub fn new_with_snapshot(context: Context, content: String, snapshot: Arc) -> Self { + Self { + context, + content, + snapshot: Some(snapshot), + } } pub fn context(&self) -> &Context { @@ -52,4 +68,8 @@ impl Query { pub fn content(&self) -> &str { self.content.as_str() } + + pub fn snapshot(&self) -> Option<&Arc> { + self.snapshot.as_ref() + } } diff --git a/crates/s3select-api/src/query/session.rs b/crates/s3select-api/src/query/session.rs index 81824a027..39a70e383 100644 --- a/crates/s3select-api/src/query/session.rs +++ b/crates/s3select-api/src/query/session.rs @@ -12,14 +12,16 @@ // See the License for the specific language governing permissions and // limitations under the License. +use crate::SelectObjectSnapshot; use crate::query::Context; -use crate::{QueryError, QueryResult, SelectStore, object_store::EcObjectStore}; +use crate::{QueryError, QueryResult, object_store::EcObjectStore}; use datafusion::{ arrow::{ array::{Int32Array, StringArray}, datatypes::{DataType, Field, Schema}, record_batch::RecordBatch, }, + common::DataFusionError, execution::{SessionStateBuilder, config::SessionConfig, context::SessionState, runtime_env::RuntimeEnvBuilder}, object_store::{ObjectStore, ObjectStoreExt, memory::InMemory, path::Path}, parquet::arrow::ArrowWriter, @@ -39,6 +41,28 @@ use tracing::error; pub type QueryExecutionGuard = Arc; +/// A one-shot query admission reservation handed from the request boundary to +/// the dispatcher that owns the corresponding concurrency semaphore. +pub struct QueryAdmission { + query_guard: Option, +} + +impl QueryAdmission { + pub fn new(query_guard: QueryExecutionGuard) -> Self { + Self { + query_guard: Some(query_guard), + } + } + + pub fn into_query_guard(mut self) -> Option { + self.query_guard.take() + } + + pub(crate) fn unmanaged() -> Self { + Self { query_guard: None } + } +} + #[derive(Clone, Default)] pub struct QueryExecutionOwner { identity: Arc<()>, @@ -300,30 +324,30 @@ impl SessionCtxFactory { query_tracker: QueryExecutionTracker, memory_limit_bytes: usize, ) -> QueryResult { - self.create_session_ctx_inner(context, Some(query_tracker), None, memory_limit_bytes) + self.create_session_ctx_inner(context, None, Some(query_tracker), memory_limit_bytes) .await } - #[cfg(test)] - async fn create_session_ctx_with_tracker_and_store( + pub async fn create_session_ctx_with_snapshot_and_tracker_and_memory_limit( &self, context: &Context, + snapshot: Arc, query_tracker: QueryExecutionTracker, - store: Arc, + memory_limit_bytes: usize, ) -> QueryResult { - self.create_session_ctx_inner(context, Some(query_tracker), Some(store), DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES) + self.create_session_ctx_inner(context, Some(snapshot), Some(query_tracker), memory_limit_bytes) .await } async fn create_session_ctx_inner( &self, context: &Context, + snapshot: Option>, query_tracker: Option, - store: Option>, memory_limit_bytes: usize, ) -> QueryResult { let df_session_ctx = self - .build_df_session_context(context, query_tracker.clone(), store, memory_limit_bytes) + .build_df_session_context(context, snapshot, query_tracker.clone(), memory_limit_bytes) .await?; Ok(SessionCtx { @@ -336,8 +360,8 @@ impl SessionCtxFactory { async fn build_df_session_context( &self, context: &Context, + snapshot: Option>, query_tracker: Option, - store: Option>, memory_limit_bytes: usize, ) -> QueryResult { let path = format!("s3://{}", context.input.bucket); @@ -416,11 +440,13 @@ impl SessionCtxFactory { } else { let store: EcObjectStore = match query_tracker { Some(query_tracker) => { - EcObjectStore::new_with_query_tracker(context.input.clone(), memory_pool, query_tracker, store) + EcObjectStore::new_with_query_tracker(context.input.clone(), memory_pool, query_tracker, snapshot) } - None => EcObjectStore::new_with_memory_pool(context.input.clone(), memory_pool), + None => EcObjectStore::new_with_memory_pool(context.input.clone(), memory_pool, snapshot), } - .map_err(|_| QueryError::NotImplemented { err: String::new() })?; + .map_err(|err| QueryError::Datafusion { + source: Box::new(DataFusionError::External(Box::new(err))), + })?; df_session_state.with_object_store(&store_url, Arc::new(store)).build() }; @@ -498,6 +524,7 @@ mod tests { }, execution::memory_pool::MemoryLimit, }; + use http::HeaderMap; use s3s::dto::{ CSVInput, CSVOutput, ExpressionType, InputSerialization, JSONInput, OutputSerialization, ParquetInput, ScanRange, SelectObjectContentInput, SelectObjectContentRequest, @@ -531,6 +558,16 @@ mod tests { } } + async fn prepare_test_snapshot(context: &Context) -> Arc { + let env = crate::storage_api::select_test_ecstore_env().await; + Arc::new( + env.ecstore + .prepare_select_object_snapshot(&context.input.bucket, &context.input.key, &HeaderMap::new(), &Default::default()) + .await + .expect("prepare SelectObjectContent snapshot"), + ) + } + #[test] fn session_factory_fields_remain_source_compatible() { let factory = SessionCtxFactory { @@ -679,10 +716,57 @@ mod tests { )); } + #[tokio::test] + #[serial_test::serial] + async fn production_session_preserves_lazy_snapshot_entry() { + let _env = crate::storage_api::select_test_ecstore_env().await; + let session = SessionCtxFactory::new(false) + .create_session_ctx(&test_context()) + .await + .expect("legacy production session should install a lazy object store"); + + assert_eq!(session.inner().config().target_partitions(), SessionConfig::new().target_partitions()); + } + + #[tokio::test] + #[serial_test::serial] + async fn legacy_tracked_production_session_preserves_lazy_snapshot_entry() { + let _env = crate::storage_api::select_test_ecstore_env().await; + let permit = Arc::new(tokio::sync::Semaphore::new(1)) + .acquire_owned() + .await + .expect("query permit should be available"); + let tracker = QueryExecutionTracker::new( + &QueryExecutionOwner::new(), + Arc::new(permit), + Instant::now() + std::time::Duration::from_secs(300), + 300, + ); + let session = SessionCtxFactory::new(false) + .create_session_ctx_with_tracker_and_memory_limit( + &test_context(), + tracker.clone(), + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, + ) + .await + .expect("legacy tracked session should install a lazy object store"); + + assert!(session.is_bound_to(&tracker)); + } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial_test::serial] async fn session_factory_propagates_query_guard_to_ec_store() { let env = crate::storage_api::select_test_ecstore_env().await; + let mut context = test_context(); + Arc::make_mut(&mut context.input).bucket = "s3select-query-guard-snapshot".to_string(); + env.make_bucket(&context.input.bucket, false).await; + let mut reader = SelectPutObjReader::from_vec(b"id,name\n1,Alice\n".to_vec()); + env.ecstore + .put_object(&context.input.bucket, &context.input.key, &mut reader, &Default::default()) + .await + .expect("put query guard fixture"); + let snapshot = prepare_test_snapshot(&context).await; let admission = Arc::new(tokio::sync::Semaphore::new(1)); let permit = Arc::clone(&admission) @@ -697,7 +781,12 @@ mod tests { 300, ); let session = SessionCtxFactory::new(false) - .create_session_ctx_with_tracker_and_store(&test_context(), query_tracker, Arc::clone(&env.ecstore)) + .create_session_ctx_with_snapshot_and_tracker_and_memory_limit( + &context, + snapshot, + query_tracker, + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, + ) .await .expect("production session should be created with the query guard"); @@ -706,6 +795,51 @@ mod tests { assert_eq!(Arc::strong_count(&query_guard), 1); } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn session_factory_preserves_snapshot_binding_error_source() { + let env = crate::storage_api::select_test_ecstore_env().await; + let mut source_context = test_context(); + Arc::make_mut(&mut source_context.input).bucket = "s3select-session-snapshot-identity".to_string(); + Arc::make_mut(&mut source_context.input).key = "source.csv".to_string(); + env.make_bucket(&source_context.input.bucket, false).await; + let mut reader = SelectPutObjReader::from_vec(b"source-marker\n".to_vec()); + env.ecstore + .put_object(&source_context.input.bucket, &source_context.input.key, &mut reader, &Default::default()) + .await + .expect("put snapshot identity fixture"); + let snapshot = prepare_test_snapshot(&source_context).await; + + let mut target_context = source_context.clone(); + Arc::make_mut(&mut target_context.input).key = "different.csv".to_string(); + let permit = Arc::new(tokio::sync::Semaphore::new(1)) + .acquire_owned() + .await + .expect("query permit should be available"); + let tracker = QueryExecutionTracker::new( + &QueryExecutionOwner::new(), + Arc::new(permit), + Instant::now() + std::time::Duration::from_secs(300), + 300, + ); + + let error = match SessionCtxFactory::new(false) + .create_session_ctx_with_snapshot_and_tracker_and_memory_limit( + &target_context, + snapshot, + tracker, + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, + ) + .await + { + Ok(_) => panic!("a session must reject a snapshot for a different object"), + Err(error) => error, + }; + + assert!(error.is_snapshot_consistency_error()); + assert!(error.to_string().contains("snapshot consistency failure")); + } + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial_test::serial] async fn scan_range_is_preserved_across_large_csv_partition_boundary() { @@ -721,6 +855,7 @@ mod tests { assert!(data.len() > 1024 * 1024); let mut context = test_context(); + Arc::make_mut(&mut context.input).bucket = "s3select-scan-range-partition-snapshot".to_string(); let selected_start = i64::try_from(SELECTED_ROW * ROW_WIDTH).expect("selected row offset should fit in i64"); Arc::make_mut(&mut context.input).request.scan_range = Some(ScanRange { start: Some(selected_start), @@ -732,6 +867,7 @@ mod tests { .put_object(&context.input.bucket, &context.input.key, &mut reader, &Default::default()) .await .expect("put large ScanRange CSV fixture"); + let snapshot = prepare_test_snapshot(&context).await; let admission = Arc::new(tokio::sync::Semaphore::new(1)); let permit = Arc::clone(&admission) @@ -746,7 +882,12 @@ mod tests { ); let session = SessionCtxFactory::new(false) .with_target_partitions(2) - .create_session_ctx_with_tracker_and_store(&context, query_tracker, Arc::clone(&env.ecstore)) + .create_session_ctx_with_snapshot_and_tracker_and_memory_limit( + &context, + snapshot, + query_tracker, + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, + ) .await .expect("create production ScanRange session"); assert!(!session.inner().config().options().optimizer.repartition_file_scans); diff --git a/crates/s3select-api/src/server/dbms.rs b/crates/s3select-api/src/server/dbms.rs index 2186fe31d..6f688bcd1 100644 --- a/crates/s3select-api/src/server/dbms.rs +++ b/crates/s3select-api/src/server/dbms.rs @@ -20,6 +20,7 @@ use crate::{ Query, execution::{Output, QueryStateMachineRef}, logical_planner::Plan, + session::QueryAdmission, }, }; @@ -44,7 +45,16 @@ impl QueryHandle { #[async_trait] pub trait DatabaseManagerSystem { + fn try_reserve_query(&self) -> QueryResult { + Ok(QueryAdmission::unmanaged()) + } + async fn execute(&self, query: &Query) -> QueryResult; + + async fn execute_admitted(&self, query: &Query, _admission: QueryAdmission) -> QueryResult { + self.execute(query).await + } + async fn build_query_state_machine(&self, query: Query) -> QueryResult; async fn build_logical_plan(&self, query_state_machine: QueryStateMachineRef) -> QueryResult>; async fn execute_logical_plan( diff --git a/crates/s3select-api/src/storage_api.rs b/crates/s3select-api/src/storage_api.rs index bb9e9ed66..4cacf0571 100644 --- a/crates/s3select-api/src/storage_api.rs +++ b/crates/s3select-api/src/storage_api.rs @@ -22,35 +22,51 @@ use rustfs_ecstore::api::error::{ }; #[cfg(test)] pub(crate) use rustfs_ecstore::api::object::PutObjReader as SelectPutObjReader; +pub use rustfs_ecstore::api::object::SelectObjectSnapshot; +pub(crate) use rustfs_ecstore::api::object::{ + PrepareSelectObjectSnapshotError, SelectObjectSnapshotReadError, SnapshotConsistencyError, +}; use rustfs_ecstore::api::runtime::object_store_handle as resolve_select_object_store_handle_from_backend; pub(crate) use rustfs_ecstore::api::set_disk::DEFAULT_READ_BUFFER_SIZE as SELECT_DEFAULT_READ_BUFFER_SIZE; pub(crate) use rustfs_ecstore::api::storage::ECStore as SelectStore; use rustfs_storage_api as storage_contracts; +#[cfg(test)] +static SELECT_TEST_OBJECT_STORE: std::sync::OnceLock> = std::sync::OnceLock::new(); + pub(crate) mod object_store { - pub(crate) use super::storage_contracts::{HTTPRangeSpec, ObjectIO, ObjectOperations}; + pub(crate) use super::storage_contracts::HTTPRangeSpec; + #[cfg(test)] + pub(crate) use super::storage_contracts::ObjectIO; } pub(crate) mod crate_boundary { pub(crate) use super::{ - SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError, - SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, + PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions, + SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError, + resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, select_is_err_version_not_found, }; } pub(crate) type SelectGetObjectReader = ::GetObjectReader; -pub(crate) type SelectObjectInfo = ::ObjectInfo; pub(crate) type SelectObjectOptions = ::ObjectOptions; #[cfg(test)] pub(crate) async fn select_test_ecstore_env() -> &'static rustfs_test_utils::TestECStoreEnv { static ENV: tokio::sync::OnceCell = tokio::sync::OnceCell::const_new(); - ENV.get_or_init(|| async { rustfs_test_utils::TestECStoreEnv::builder().build().await }) - .await + let env = ENV + .get_or_init(|| async { rustfs_test_utils::TestECStoreEnv::builder().build().await }) + .await; + let _ = SELECT_TEST_OBJECT_STORE.set(Arc::clone(&env.ecstore)); + env } pub(crate) fn resolve_select_object_store_handle() -> Option> { + #[cfg(test)] + if let Some(store) = SELECT_TEST_OBJECT_STORE.get() { + return Some(Arc::clone(store)); + } resolve_select_object_store_handle_from_backend() } diff --git a/crates/s3select-query/Cargo.toml b/crates/s3select-query/Cargo.toml index 61c687a50..dfec9520b 100644 --- a/crates/s3select-query/Cargo.toml +++ b/crates/s3select-query/Cargo.toml @@ -54,5 +54,8 @@ s3s = { workspace = true, features = ["minio"] } tokio = { workspace = true, features = ["fs", "rt-multi-thread", "sync", "time"] } tracing = { workspace = true } +[dev-dependencies] +rustfs-test-utils = { workspace = true, features = ["put-object-commit-barrier"] } + [lib] doctest = false diff --git a/crates/s3select-query/src/dispatcher/manager.rs b/crates/s3select-query/src/dispatcher/manager.rs index a47661179..c9398db2b 100644 --- a/crates/s3select-query/src/dispatcher/manager.rs +++ b/crates/s3select-query/src/dispatcher/manager.rs @@ -48,8 +48,8 @@ use rustfs_s3select_api::{ logical_planner::{LogicalPlanner, Plan}, parser::Parser, session::{ - DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, SessionCtx, - SessionCtxFactory, + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus, + QueryExecutionTracker, SessionCtx, SessionCtxFactory, }, }, }; @@ -120,13 +120,20 @@ impl Drop for QueryPhaseGuard<'_> { #[async_trait] impl QueryDispatcher for SimpleQueryDispatcher { async fn execute_query(&self, query: &Query) -> QueryResult { - let query_state_machine = self.build_query_state_machine(query.clone()).await?; - let logical_plan = self.build_logical_plan(Arc::clone(&query_state_machine)).await?; - let Some(logical_plan) = logical_plan else { - return Ok(Output::Nil(())); - }; + self.execute_query_inner(query, None).await + } - self.execute_logical_plan(logical_plan, query_state_machine).await + fn try_reserve_query(&self) -> QueryResult { + let permit = self + .query_admission + .clone() + .try_acquire_owned() + .map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?; + Ok(QueryAdmission::new(Arc::new(permit))) + } + + async fn execute_query_admitted(&self, query: &Query, admission: QueryAdmission) -> QueryResult { + self.execute_query_inner(query, Some(admission)).await } async fn build_logical_plan(&self, query_state_machine: Arc) -> QueryResult> { @@ -205,20 +212,64 @@ impl QueryDispatcher for SimpleQueryDispatcher { } async fn build_query_state_machine(&self, query: Query) -> QueryResult> { - let permit = self - .query_admission - .clone() - .try_acquire_owned() - .map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?; + self.build_query_state_machine_inner(query, None).await + } +} + +impl SimpleQueryDispatcher { + async fn execute_query_inner(&self, query: &Query, admission: Option) -> QueryResult { + let query_state_machine = self.build_query_state_machine_inner(query.clone(), admission).await?; + let logical_plan = self.build_logical_plan(Arc::clone(&query_state_machine)).await?; + let Some(logical_plan) = logical_plan else { + return Ok(Output::Nil(())); + }; + + self.execute_logical_plan(logical_plan, query_state_machine).await + } + + async fn build_query_state_machine_inner( + &self, + query: Query, + admission: Option, + ) -> QueryResult> { + let query_guard = match admission { + Some(admission) => { + let query_guard = admission.into_query_guard().ok_or(QueryError::Cancel)?; + if !Arc::ptr_eq(query_guard.semaphore(), &self.query_admission) { + return Err(QueryError::Cancel); + } + query_guard + } + None => { + let permit = self + .query_admission + .clone() + .try_acquire_owned() + .map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?; + Arc::new(permit) + } + }; let query_tracker = QueryExecutionTracker::new( &self.query_execution_owner, - Arc::new(permit), + query_guard, Instant::now() + self.query_timeout, self.query_timeout.as_secs(), ); let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner); - let session = self - .run_with_query_deadline( + let session = if let Some(snapshot) = query.snapshot().cloned() { + self.run_with_query_deadline( + &query_tracker, + self.session_factory + .create_session_ctx_with_snapshot_and_tracker_and_memory_limit( + query.context(), + snapshot, + query_tracker.clone(), + self.memory_limit_bytes, + ), + ) + .await? + } else { + self.run_with_query_deadline( &query_tracker, self.session_factory.create_session_ctx_with_tracker_and_memory_limit( query.context(), @@ -226,7 +277,8 @@ impl QueryDispatcher for SimpleQueryDispatcher { self.memory_limit_bytes, ), ) - .await?; + .await? + }; if !query_tracker.mark_admitted(&self.query_execution_owner) { drop(session); return Err(self.query_tracker_error(&query_tracker)); @@ -234,9 +286,6 @@ impl QueryDispatcher for SimpleQueryDispatcher { phase_guard.disarm(); Ok(Arc::new(QueryStateMachine::begin_tracked(query, session, query_tracker)?)) } -} - -impl SimpleQueryDispatcher { async fn run_with_query_deadline( &self, query_tracker: &QueryExecutionTracker, @@ -730,15 +779,17 @@ mod tests { use async_trait::async_trait; use datafusion::{ arrow::{ - datatypes::{Schema, SchemaRef}, + array::{Int32Array, StringArray}, + datatypes::{DataType, Field, Schema, SchemaRef}, record_batch::RecordBatch, }, common::DataFusionError, execution::object_store::ObjectStoreUrl, object_store::{ObjectStoreExt, path::Path}, + parquet::arrow::ArrowWriter, physical_plan::{RecordBatchStream, stream::RecordBatchStreamAdapter}, }; - use futures::{StreamExt, stream}; + use futures::{StreamExt, TryStreamExt, stream}; use rustfs_s3select_api::{ QueryError, QueryResult, S3SelectPolicyError, query::{ @@ -749,14 +800,15 @@ mod tests { }, logical_planner::Plan, session::{ - DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, - SessionCtxFactory, + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus, + QueryExecutionTracker, SessionCtxFactory, }, }, }; + use rustfs_test_utils::{PutObjectCommitBarrier, TestECStoreEnv}; use s3s::dto::{ - CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, OutputSerialization, SelectObjectContentInput, - SelectObjectContentRequest, + CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType, + OutputSerialization, ParquetInput, SelectObjectContentInput, SelectObjectContentRequest, }; use std::{ pin::Pin, @@ -998,6 +1050,260 @@ mod tests { (dispatcher, input) } + async fn snapshot_test_env() -> &'static TestECStoreEnv { + static ENV: tokio::sync::OnceCell = tokio::sync::OnceCell::const_new(); + ENV.get_or_init(|| async { TestECStoreEnv::builder().prefix("s3select_query_snapshot").build().await }) + .await + } + + fn production_dispatcher(input: Arc) -> Arc { + let optimizer = Arc::new(CascadeOptimizerBuilder::default().build()); + let scheduler = Arc::new(LocalScheduler {}); + SimpleQueryDispatcherBuilder::default() + .with_input(input) + .with_default_table_provider(Arc::new(BaseTableProvider::default())) + .with_session_factory(Arc::new(SessionCtxFactory::new(false))) + .with_parser(Arc::new(DefaultParser::default())) + .with_query_execution_factory(Arc::new(SqlQueryExecutionFactory::new(optimizer, scheduler))) + .with_func_manager(Arc::new(SimpleFunctionMetadataManager::default())) + .build() + .expect("production query dispatcher should build") + } + + async fn collect_utf8_output(output: Output) -> Vec { + let Output::StreamData(stream) = output else { + panic!("snapshot query should return rows"); + }; + stream + .try_collect::>() + .await + .expect("collect snapshot query output") + .iter() + .flat_map(|batch| { + batch + .column(0) + .as_any() + .downcast_ref::() + .expect("snapshot marker column should be Utf8") + .iter() + .map(|value| value.expect("snapshot marker should not be null").to_string()) + .collect::>() + }) + .collect() + } + + async fn run_snapshot_generation_race( + input: Arc, + old_generation: Vec, + new_generation: Vec, + expected_old_markers: &[&str], + ) { + let env = snapshot_test_env().await; + env.make_bucket(&input.bucket, false).await; + env.put_object_bytes(&input.bucket, &input.key, old_generation).await; + + let snapshot = env.prepare_select_object_snapshot(&input.bucket, &input.key).await; + let dispatcher = production_dispatcher(Arc::clone(&input)); + let query = Query::new_with_snapshot( + QueryContext { + input: Arc::clone(&input), + }, + input.request.expression.clone(), + snapshot, + ); + let query_state_machine = dispatcher + .build_query_state_machine(query.clone()) + .await + .expect("build production snapshot session"); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("infer old-generation schema") + .expect("SELECT should produce a logical plan"); + + let commit_barrier = PutObjectCommitBarrier::before_namespace(&input.bucket, &input.key); + let writer_env = env; + let writer_bucket = input.bucket.clone(); + let writer_object = input.key.clone(); + let writer = tokio::spawn(async move { + writer_env + .put_object_bytes(&writer_bucket, &writer_object, new_generation) + .await; + }); + commit_barrier.wait_until_paused().await; + commit_barrier.release_and_wait_until_namespace_pending().await; + assert!(!writer.is_finished(), "overwrite must wait for the SelectObjectContent snapshot"); + + let output = dispatcher + .execute_logical_plan(logical_plan, query_state_machine) + .await + .expect("scan old-generation rows"); + let values = collect_utf8_output(output).await; + assert_eq!(values, expected_old_markers); + assert!(!writer.is_finished(), "overwrite must remain blocked while Query owns the snapshot"); + + drop(query); + tokio::time::timeout(Duration::from_secs(5), writer) + .await + .expect("overwrite should finish after snapshot release") + .expect("overwrite task should join"); + } + + fn json_snapshot_input() -> Arc { + Arc::new(SelectObjectContentInput { + bucket: "s3select-json-snapshot-race".to_string(), + expected_bucket_owner: None, + key: "input.jsonl".to_string(), + sse_customer_algorithm: None, + sse_customer_key: None, + sse_customer_key_md5: None, + request: SelectObjectContentRequest { + expression: "SELECT old_marker FROM S3Object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + json: Some(JSONInput { + type_: Some(JSONType::from_static(JSONType::LINES)), + }), + ..Default::default() + }, + output_serialization: OutputSerialization { + json: Some(JSONOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + }) + } + + fn parquet_snapshot_input() -> Arc { + Arc::new(SelectObjectContentInput { + bucket: "s3select-parquet-snapshot-race".to_string(), + expected_bucket_owner: None, + key: "input.parquet".to_string(), + sse_customer_algorithm: None, + sse_customer_key: None, + sse_customer_key_md5: None, + request: SelectObjectContentRequest { + expression: "SELECT old_marker FROM S3Object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + parquet: Some(ParquetInput {}), + ..Default::default() + }, + output_serialization: OutputSerialization { + json: Some(JSONOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + }) + } + + fn old_parquet_generation() -> Vec { + let schema = Arc::new(Schema::new(vec![Field::new("old_marker", DataType::Utf8, false)])); + let first = RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(StringArray::from(vec!["parquet-old-footer"]))]) + .expect("build first old-generation parquet row group"); + let second = RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(StringArray::from(vec!["parquet-old-row-group"]))]) + .expect("build second old-generation parquet row group"); + let mut bytes = Vec::new(); + let mut writer = ArrowWriter::try_new(&mut bytes, schema, None).expect("create old-generation parquet writer"); + writer.write(&first).expect("write first old-generation parquet row group"); + writer.flush().expect("flush first old-generation parquet row group"); + writer.write(&second).expect("write second old-generation parquet row group"); + writer.close().expect("close old-generation parquet writer"); + bytes + } + + fn new_parquet_generation() -> Vec { + let schema = Arc::new(Schema::new(vec![Field::new("new_schema_poison", DataType::Int32, false)])); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(Int32Array::from(vec![1629]))]) + .expect("build new-generation parquet poison row group"); + let mut bytes = Vec::new(); + let mut writer = ArrowWriter::try_new(&mut bytes, schema, None).expect("create new-generation parquet writer"); + writer.write(&batch).expect("write new-generation parquet poison row group"); + writer.close().expect("close new-generation parquet writer"); + bytes + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn json_schema_inference_and_scan_use_one_snapshot_generation() { + run_snapshot_generation_race( + json_snapshot_input(), + b"{\"old_marker\":\"json-old-schema\"}\n{\"old_marker\":\"json-old-scan\"}\n".to_vec(), + b"{\"new_schema_poison\":1629}\n".to_vec(), + &["json-old-schema", "json-old-scan"], + ) + .await; + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn parquet_footer_and_row_group_reads_use_one_snapshot_generation() { + run_snapshot_generation_race( + parquet_snapshot_input(), + old_parquet_generation(), + new_parquet_generation(), + &["parquet-old-footer", "parquet-old-row-group"], + ) + .await; + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn planner_failure_drops_snapshot_and_unblocks_overwrite() { + let mut input = json_snapshot_input(); + let input_mut = Arc::make_mut(&mut input); + input_mut.bucket = "s3select-planner-failure-snapshot".to_string(); + input_mut.request.expression = "SELECT missing_binding FROM S3Object".to_string(); + + let env = snapshot_test_env().await; + env.make_bucket(&input.bucket, false).await; + env.put_object_bytes(&input.bucket, &input.key, b"{\"old_marker\":\"old\"}\n".to_vec()) + .await; + let snapshot = env.prepare_select_object_snapshot(&input.bucket, &input.key).await; + let dispatcher = production_dispatcher(Arc::clone(&input)); + let query = Query::new_with_snapshot( + QueryContext { + input: Arc::clone(&input), + }, + input.request.expression.clone(), + snapshot, + ); + let query_state_machine = dispatcher + .build_query_state_machine(query.clone()) + .await + .expect("build production snapshot session"); + + let commit_barrier = PutObjectCommitBarrier::before_namespace(&input.bucket, &input.key); + let writer_env = env; + let writer_bucket = input.bucket.clone(); + let writer_object = input.key.clone(); + let writer = tokio::spawn(async move { + writer_env + .put_object_bytes(&writer_bucket, &writer_object, b"{\"new_schema_poison\":1629}\n".to_vec()) + .await; + }); + commit_barrier.wait_until_paused().await; + commit_barrier.release_and_wait_until_namespace_pending().await; + assert!(!writer.is_finished(), "overwrite must wait for the planner's snapshot"); + + let Err(error) = dispatcher.build_logical_plan(Arc::clone(&query_state_machine)).await else { + panic!("missing schema binding must fail planning"); + }; + assert!(error.to_string().contains("missing_binding")); + assert!( + !writer.is_finished(), + "planner failure must not release snapshots still owned by Query/qsm" + ); + + drop(query_state_machine); + drop(query); + tokio::time::timeout(Duration::from_secs(5), writer) + .await + .expect("overwrite should finish after failed-plan snapshot release") + .expect("overwrite task should join"); + } + fn test_query_tracker( permit: tokio::sync::OwnedSemaphorePermit, deadline: Instant, @@ -1118,6 +1424,42 @@ mod tests { )); } + #[tokio::test] + async fn reserved_admission_is_handed_to_tracker_without_reacquiring() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let reservation = dispatcher.try_reserve_query().expect("query reservation should succeed"); + assert_eq!(admission.available_permits(), 0); + + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine_inner(query, Some(reservation)) + .await + .expect("reserved query should not acquire a second permit"); + + assert_eq!(admission.available_permits(), 0); + drop(query_state_machine); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn reserved_admission_rejects_a_foreign_semaphore() { + let admission = Arc::new(Semaphore::new(1)); + let foreign_admission = Arc::new(Semaphore::new(1)); + let foreign_permit = Arc::clone(&foreign_admission) + .try_acquire_owned() + .expect("foreign permit should be available"); + let reservation = QueryAdmission::new(Arc::new(foreign_permit)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + + let result = dispatcher.build_query_state_machine_inner(query, Some(reservation)).await; + + assert!(matches!(result, Err(QueryError::Cancel))); + assert_eq!(admission.available_permits(), 1); + assert_eq!(foreign_admission.available_permits(), 1); + } + #[tokio::test] async fn staged_query_rejects_when_admission_is_saturated() { let admission = Arc::new(Semaphore::new(1)); diff --git a/crates/s3select-query/src/instance.rs b/crates/s3select-query/src/instance.rs index ca2807d46..e33a73a71 100644 --- a/crates/s3select-query/src/instance.rs +++ b/crates/s3select-query/src/instance.rs @@ -26,7 +26,7 @@ use rustfs_s3select_api::{ dispatcher::QueryDispatcher, execution::QueryStateMachineRef, logical_planner::Plan, - session::{DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES as DEFAULT_MEMORY_LIMIT_BYTES, SessionCtxFactory}, + session::{DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES as DEFAULT_MEMORY_LIMIT_BYTES, QueryAdmission, SessionCtxFactory}, }, server::dbms::{DatabaseManagerSystem, QueryHandle}, }; @@ -64,12 +64,22 @@ impl DatabaseManagerSystem for RustFSms where D: QueryDispatcher, { + fn try_reserve_query(&self) -> QueryResult { + self.query_dispatcher.try_reserve_query() + } + async fn execute(&self, query: &Query) -> QueryResult { let result = self.query_dispatcher.execute_query(query).await?; Ok(QueryHandle::new(query.clone(), result)) } + async fn execute_admitted(&self, query: &Query, admission: QueryAdmission) -> QueryResult { + let result = self.query_dispatcher.execute_query_admitted(query, admission).await?; + + Ok(QueryHandle::new(query.clone(), result)) + } + async fn build_query_state_machine(&self, query: Query) -> QueryResult { let query_state_machine = self.query_dispatcher.build_query_state_machine(query).await?; diff --git a/crates/test-utils/Cargo.toml b/crates/test-utils/Cargo.toml index bf56d069a..8d9bcb6fb 100644 --- a/crates/test-utils/Cargo.toml +++ b/crates/test-utils/Cargo.toml @@ -27,6 +27,7 @@ documentation = "https://docs.rs/rustfs-test-utils/latest/rustfs_test_utils/" [features] default = [] +put-object-commit-barrier = ["rustfs-ecstore/test-util"] hotpath = [ "hotpath/hotpath", "hotpath/tokio", @@ -51,7 +52,7 @@ hotpath-cpu = [ [dependencies] hotpath.workspace = true -rustfs-ecstore = { workspace = true } +rustfs-ecstore.workspace = true rustfs-storage-api = { workspace = true } tokio = { workspace = true, features = ["fs", "rt-multi-thread"] } tokio-util = { workspace = true, features = ["io", "compat"] } diff --git a/crates/test-utils/src/ecstore_test_compat.rs b/crates/test-utils/src/ecstore_test_compat.rs index e21d534d8..6ed8bfaa1 100644 --- a/crates/test-utils/src/ecstore_test_compat.rs +++ b/crates/test-utils/src/ecstore_test_compat.rs @@ -26,8 +26,11 @@ pub(crate) mod fixture { pub(crate) use rustfs_ecstore::api::bucket::metadata_sys::init_bucket_metadata_sys; pub(crate) use rustfs_ecstore::api::disk::endpoint::Endpoint; pub(crate) use rustfs_ecstore::api::layout::{EndpointServerPools, Endpoints, PoolEndpoints}; + pub(crate) use rustfs_ecstore::api::object::{PutObjReader, SelectObjectSnapshot}; + #[cfg(feature = "put-object-commit-barrier")] + pub(crate) use rustfs_ecstore::api::set_disk as ecstore_set_disk; pub(crate) use rustfs_ecstore::api::storage::{ECStore, init_local_disks}; - pub(crate) use rustfs_storage_api::{BucketOperations, BucketOptions, MakeBucketOptions}; + pub(crate) use rustfs_storage_api::{BucketOperations, BucketOptions, MakeBucketOptions, ObjectIO}; #[cfg(test)] pub(crate) use rustfs_ecstore::api::config::com::{delete_config, read_config, save_config}; diff --git a/crates/test-utils/src/lib.rs b/crates/test-utils/src/lib.rs index 1511f2aa2..46f215a9d 100644 --- a/crates/test-utils/src/lib.rs +++ b/crates/test-utils/src/lib.rs @@ -30,14 +30,38 @@ mod ecstore_test_compat; use std::path::PathBuf; use std::sync::{Arc, Once}; +#[cfg(feature = "put-object-commit-barrier")] +use ecstore_test_compat::fixture::ecstore_set_disk; use ecstore_test_compat::fixture::{ - BucketOperations as _, BucketOptions, ECStore, Endpoint, EndpointServerPools, Endpoints, MakeBucketOptions, PoolEndpoints, - init_bucket_metadata_sys, init_local_disks, + BucketOperations as _, BucketOptions, ECStore, Endpoint, EndpointServerPools, Endpoints, MakeBucketOptions, ObjectIO as _, + PoolEndpoints, PutObjReader, SelectObjectSnapshot, init_bucket_metadata_sys, init_local_disks, }; use tokio_util::sync::CancellationToken; static INIT_TRACING: Once = Once::new(); +#[cfg(feature = "put-object-commit-barrier")] +pub struct PutObjectCommitBarrier(ecstore_set_disk::test_util::PutObjectCommitBarrier); + +#[cfg(feature = "put-object-commit-barrier")] +impl PutObjectCommitBarrier { + pub fn before_namespace(bucket: &str, object: &str) -> Self { + Self(ecstore_set_disk::test_util::PutObjectCommitBarrier::install( + bucket, + object, + ecstore_set_disk::test_util::PutObjectCommitPause::BeforeNamespace, + )) + } + + pub async fn wait_until_paused(&self) { + self.0.wait_until_paused().await; + } + + pub async fn release_and_wait_until_namespace_pending(&self) { + self.0.release_and_wait_until_namespace_pending().await; + } +} + /// Install the standard test tracing subscriber once per process /// (`RUST_LOG`-driven). Safe to call from every test; later calls are no-ops. pub fn init_tracing() { @@ -90,6 +114,29 @@ impl TestECStoreEnv { .await .unwrap_or_else(|e| panic!("failed to create test bucket {bucket}: {e:?}")); } + + /// Write one complete object body through the real ECStore test backend. + pub async fn put_object_bytes(&self, bucket: &str, object: &str, bytes: Vec) { + let mut reader = PutObjReader::from_vec(bytes); + self.ecstore + .put_object(bucket, object, &mut reader, &Default::default()) + .await + .unwrap_or_else(|e| panic!("failed to write test object {bucket}/{object}: {e:?}")); + } + + /// Prepare the lock-backed object snapshot used by SelectObjectContent tests. + /// + /// The concrete ECStore snapshot type stays behind this crate's test + /// compatibility boundary; consumers can pass the inferred value directly + /// to the S3 Select API without importing ECStore facade paths. + pub async fn prepare_select_object_snapshot(&self, bucket: &str, object: &str) -> Arc { + Arc::new( + self.ecstore + .prepare_select_object_snapshot(bucket, object, &Default::default(), &Default::default()) + .await + .unwrap_or_else(|e| panic!("failed to prepare test object snapshot {bucket}/{object}: {e:?}")), + ) + } } /// Builder for [`TestECStoreEnv`]. Defaults reproduce the historical heal diff --git a/rustfs/src/app/select_object.rs b/rustfs/src/app/select_object.rs index 7d5454e03..0b4904893 100644 --- a/rustfs/src/app/select_object.rs +++ b/rustfs/src/app/select_object.rs @@ -1,8 +1,12 @@ -use super::storage_api::select_object::contract::object::ObjectOperations as _; +#[cfg(test)] +use super::storage_api::select_object::StorageError; use super::storage_api::select_object::options::get_opts; use super::storage_api::select_object::request_context::spawn_traced; use super::storage_api::select_object::sse::{SseKmsPrincipal, authorize_sse_kms_object_read}; -use super::storage_api::select_object::{get_validated_store, validate_sse_headers_for_read, validate_ssec_for_read}; +use super::storage_api::select_object::{ + StoragePrepareSelectObjectSnapshotError, StorageSelectObjectSnapshot, get_validated_store, validate_sse_headers_for_read, + validate_ssec_for_read, +}; use crate::app::runtime_sources::current_s3select_db; use crate::error::ApiError; use bytes::Bytes; @@ -14,20 +18,29 @@ use datafusion::arrow::{ use datafusion::common::DataFusionError; use datafusion::physical_plan::SendableRecordBatchStream; use futures::StreamExt; -use http::{StatusCode, header::RANGE}; +use http::{HeaderMap, HeaderName, HeaderValue, StatusCode, header::RANGE}; use rustfs_s3select_api::{ QueryError, S3SelectPolicyError, object_store::{INVALID_SCAN_RANGE_MESSAGE, validate_scan_range_bounds}, query::{Context, Query}, }; use rustfs_s3select_query::instance::s3_select_query_timeout; +use rustfs_utils::http::headers::{ + AMZ_ENCRYPTION_AES, AMZ_ENCRYPTION_KMS, AMZ_SERVER_SIDE_ENCRYPTION, AMZ_SERVER_SIDE_ENCRYPTION_KMS_CONTEXT, + AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, SSEC_ALGORITHM_HEADER, SSEC_KEY_HEADER, SSEC_KEY_MD5_HEADER, +}; use s3s::dto::{ CSVOutput, CompressionType, ContinuationEvent, EndEvent, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType, OutputSerialization, Progress, ProgressEvent, QuoteFields, RecordsEvent, SelectObjectContentEvent, SelectObjectContentEventStream, SelectObjectContentInput, SelectObjectContentOutput, SelectObjectContentRequest, Stats, StatsEvent, }; +use s3s::header::{ + X_AMZ_SERVER_SIDE_ENCRYPTION, X_AMZ_SERVER_SIDE_ENCRYPTION_AWS_KMS_KEY_ID, X_AMZ_SERVER_SIDE_ENCRYPTION_CONTEXT, + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, +}; use s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error}; +use std::collections::HashMap; use std::sync::Arc; use tokio::sync::mpsc; use tokio::time::{Instant, timeout_at}; @@ -38,6 +51,13 @@ const MAX_SELECT_EXPRESSION_BYTES: usize = 256 * 1024; const RECORDS_CHUNK_TARGET: usize = 128 * 1024; const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure"; const EMPTY_SELECT_EXPRESSION_MESSAGE: &str = "empty SQL expression"; +const SELECT_MINIO_SSEC_SEALED_KEY: &str = "X-Minio-Internal-Server-Side-Encryption-Sealed-Key"; +const SELECT_MINIO_S3_SEALED_KEY: &str = "X-Minio-Internal-Server-Side-Encryption-S3-Sealed-Key"; +const SELECT_MINIO_KMS_SEALED_KEY: &str = "X-Minio-Internal-Server-Side-Encryption-Kms-Sealed-Key"; +const SELECT_MINIO_KMS_KEY_ID: &str = "X-Minio-Internal-Server-Side-Encryption-S3-Kms-Key-Id"; +const SELECT_MINIO_KMS_CONTEXT: &str = "X-Minio-Internal-Server-Side-Encryption-Context"; +const SELECT_RUSTFS_KMS_KEY_ID: &str = "x-rustfs-encryption-key-id"; +const SELECT_KMS_ARN_PREFIX: &str = "arn:aws:kms:"; #[derive(Clone, Debug)] struct SelectValidation { @@ -45,10 +65,6 @@ struct SelectValidation { progress_enabled: bool, } -struct SelectObjectMetadata { - size: u64, -} - #[derive(Clone, Debug)] enum SelectOutputFormat { Csv(CSVOutput), @@ -60,6 +76,21 @@ enum SelectProducerOutcome { ReceiverClosed, } +trait SelectSnapshotFence { + fn ensure_snapshot_valid(&self) -> S3Result<()>; +} + +impl SelectSnapshotFence for Arc { + fn ensure_snapshot_valid(&self) -> S3Result<()> { + self.ensure_valid().map_err(|error| { + let message = error.to_string(); + let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, message); + s3_error.set_source(Box::new(error)); + s3_error + }) + } +} + pub async fn execute_select_object_content( req: S3Request, ) -> S3Result> { @@ -67,19 +98,27 @@ pub async fn execute_select_object_content( let mut input = req.input; let validation = validate_select_request(&req.headers, &mut input)?; log_select_request_summary(&input, &validation); - let metadata = preflight_select_object(&req.headers, &input, read_principal.as_ref()).await?; - validate_scan_range_for_object_size(&input.request, metadata.size)?; - - let input = Arc::new(input); let query_timeout = s3_select_query_timeout(); let query_deadline = Instant::now() + query_timeout; - let db = current_s3select_db((*input).clone(), false) + let input = Arc::new(input); + let db = timeout_at(query_deadline, current_s3select_db((*input).clone(), false)) .await + .map_err(|_| select_query_timeout_error(query_timeout.as_secs()))? .map_err(map_query_error_to_s3)?; - let query = Query::new(Context { input: input.clone() }, input.request.expression.clone()); - let output = db - .execute(&query) + let admission = db.try_reserve_query().map_err(map_query_error_to_s3)?; + let snapshot = timeout_at( + query_deadline, + prepare_select_object_snapshot(&req.headers, &input, read_principal.as_ref()), + ) + .await + .map_err(|_| select_query_timeout_error(query_timeout.as_secs()))??; + validate_scan_range_for_object_size(&input.request, snapshot.logical_size())?; + let snapshot = Arc::new(snapshot); + let query = + Query::new_with_snapshot(Context { input: input.clone() }, input.request.expression.clone(), Arc::clone(&snapshot)); + let output = timeout_at(query_deadline, db.execute_admitted(&query, admission)) .await + .map_err(|_| select_query_timeout_error(query_timeout.as_secs()))? .map_err(map_query_error_to_s3)? .result() .into_record_batch_stream() @@ -90,24 +129,213 @@ pub async fn execute_select_object_content( .clone() .try_reserve_owned() .map_err(|_| s3_error!(InternalError, "can't reserve Select terminal event capacity"))?; + let response = select_object_response(rx, &snapshot.object_info().user_defined, &req.headers)?; spawn_traced(async move { - send_select_events_until_deadline(output, tx, terminal_permit, validation, query_deadline, query_timeout.as_secs()).await; + send_select_events_until_deadline( + output, + tx, + terminal_permit, + validation, + query_deadline, + query_timeout.as_secs(), + snapshot, + ) + .await; }); - Ok(S3Response::new(SelectObjectContentOutput { - payload: Some(SelectObjectContentEventStream::new(ReceiverStream::new(rx))), - })) + Ok(response) } -async fn send_select_events_until_deadline( +fn select_object_response( + rx: mpsc::Receiver>, + metadata: &HashMap, + request_headers: &HeaderMap, +) -> S3Result> { + let response_headers = select_snapshot_sse_response_headers(metadata, request_headers)?; + let mut response = S3Response::new(SelectObjectContentOutput { + payload: Some(SelectObjectContentEventStream::new(ReceiverStream::new(rx))), + }); + response.headers = response_headers; + Ok(response) +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum SelectSnapshotSseMode { + S3, + Kms, + Customer, +} + +fn invalid_select_snapshot_sse_metadata() -> S3Error { + S3Error::with_message( + S3ErrorCode::InternalError, + "Persisted SelectObjectContent encryption metadata is invalid.", + ) +} + +fn select_metadata_value<'a>(metadata: &'a HashMap, name: &str) -> S3Result> { + let mut values = metadata + .iter() + .filter_map(|(key, value)| key.eq_ignore_ascii_case(name).then_some(value.as_str())); + let Some(value) = values.next() else { + return Ok(None); + }; + if values.any(|candidate| candidate != value) { + return Err(invalid_select_snapshot_sse_metadata()); + } + Ok(Some(value)) +} + +fn select_snapshot_kms_key_id(metadata: &HashMap) -> S3Result> { + let values = [ + select_metadata_value(metadata, AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID)?, + select_metadata_value(metadata, SELECT_RUSTFS_KMS_KEY_ID)?, + select_metadata_value(metadata, SELECT_MINIO_KMS_KEY_ID)?, + ]; + let mut resolved = None; + for value in values.into_iter().flatten() { + if resolved.is_some_and(|current| current != value) { + return Err(invalid_select_snapshot_sse_metadata()); + } + resolved = Some(value); + } + Ok(resolved) +} + +fn select_snapshot_sse_mode(metadata: &HashMap) -> S3Result> { + let public_mode = select_metadata_value(metadata, AMZ_SERVER_SIDE_ENCRYPTION)?; + let customer_algorithm = select_metadata_value(metadata, SSEC_ALGORITHM_HEADER)?; + let has_ssec_marker = select_metadata_value(metadata, SELECT_MINIO_SSEC_SEALED_KEY)?.is_some(); + let has_s3_marker = select_metadata_value(metadata, SELECT_MINIO_S3_SEALED_KEY)?.is_some(); + let has_kms_marker = select_metadata_value(metadata, SELECT_MINIO_KMS_SEALED_KEY)?.is_some(); + + let public_mode = match public_mode { + Some(AMZ_ENCRYPTION_AES) => Some(SelectSnapshotSseMode::S3), + Some(AMZ_ENCRYPTION_KMS) => Some(SelectSnapshotSseMode::Kms), + Some(_) => return Err(invalid_select_snapshot_sse_metadata()), + None => None, + }; + if customer_algorithm.is_some_and(|algorithm| algorithm != AMZ_ENCRYPTION_AES) { + return Err(invalid_select_snapshot_sse_metadata()); + } + + let resolved = if customer_algorithm.is_some() { + if public_mode == Some(SelectSnapshotSseMode::Kms) { + return Err(invalid_select_snapshot_sse_metadata()); + } + Some(SelectSnapshotSseMode::Customer) + } else { + public_mode + }; + let internal_modes = [ + has_ssec_marker.then_some(SelectSnapshotSseMode::Customer), + has_s3_marker.then_some(SelectSnapshotSseMode::S3), + has_kms_marker.then_some(SelectSnapshotSseMode::Kms), + ]; + for mode in internal_modes.into_iter().flatten() { + if resolved != Some(mode) { + return Err(invalid_select_snapshot_sse_metadata()); + } + } + if resolved.is_none() + && metadata + .keys() + .any(|key| rustfs_utils::http::is_object_encryption_marker(key)) + { + return Err(invalid_select_snapshot_sse_metadata()); + } + Ok(resolved) +} + +fn insert_select_snapshot_header(headers: &mut HeaderMap, name: HeaderName, value: &str) -> S3Result<()> { + let value = HeaderValue::from_str(value).map_err(|_| invalid_select_snapshot_sse_metadata())?; + headers.insert(name, value); + Ok(()) +} + +fn select_snapshot_sse_response_headers(metadata: &HashMap, request_headers: &HeaderMap) -> S3Result { + if select_metadata_value(metadata, SSEC_KEY_HEADER)?.is_some() + || select_metadata_value(metadata, AMZ_SERVER_SIDE_ENCRYPTION_KMS_CONTEXT)?.is_some() + { + return Err(invalid_select_snapshot_sse_metadata()); + } + let Some(mode) = select_snapshot_sse_mode(metadata)? else { + return Ok(HeaderMap::new()); + }; + let kms_key_id = select_snapshot_kms_key_id(metadata)?; + + let mut response_headers = HeaderMap::with_capacity(3); + match mode { + SelectSnapshotSseMode::S3 => { + if select_metadata_value(metadata, AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID)?.is_some() + || select_metadata_value(metadata, SELECT_MINIO_KMS_CONTEXT)?.is_some() + || select_metadata_value(metadata, SSEC_KEY_MD5_HEADER)?.is_some() + { + return Err(invalid_select_snapshot_sse_metadata()); + } + response_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static(AMZ_ENCRYPTION_AES)); + } + SelectSnapshotSseMode::Kms => { + if select_metadata_value(metadata, SSEC_KEY_MD5_HEADER)?.is_some() { + return Err(invalid_select_snapshot_sse_metadata()); + } + let key_id = kms_key_id + .filter(|key_id| !key_id.is_empty()) + .ok_or_else(invalid_select_snapshot_sse_metadata)?; + response_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION, HeaderValue::from_static(AMZ_ENCRYPTION_KMS)); + if key_id.starts_with(SELECT_KMS_ARN_PREFIX) { + insert_select_snapshot_header(&mut response_headers, X_AMZ_SERVER_SIDE_ENCRYPTION_AWS_KMS_KEY_ID, key_id)?; + } else { + insert_select_snapshot_header( + &mut response_headers, + X_AMZ_SERVER_SIDE_ENCRYPTION_AWS_KMS_KEY_ID, + &format!("{SELECT_KMS_ARN_PREFIX}{key_id}"), + )?; + } + if let Some(context) = select_metadata_value(metadata, SELECT_MINIO_KMS_CONTEXT)? { + let context = HeaderValue::from_str(context).map_err(|_| invalid_select_snapshot_sse_metadata())?; + let mut validation_headers = HeaderMap::with_capacity(1); + validation_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION_CONTEXT, context.clone()); + super::storage_api::select_object::sse::extract_ssekms_context_from_headers(&validation_headers) + .map_err(|_| invalid_select_snapshot_sse_metadata())?; + response_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION_CONTEXT, context); + } + } + SelectSnapshotSseMode::Customer => { + if kms_key_id.is_some() || select_metadata_value(metadata, SELECT_MINIO_KMS_CONTEXT)?.is_some() { + return Err(invalid_select_snapshot_sse_metadata()); + } + let algorithm = request_headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM) + .and_then(|value| value.to_str().ok()) + .filter(|algorithm| *algorithm == AMZ_ENCRYPTION_AES) + .ok_or_else(invalid_select_snapshot_sse_metadata)?; + let key_md5 = request_headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5) + .and_then(|value| value.to_str().ok()) + .ok_or_else(invalid_select_snapshot_sse_metadata)?; + let stored_md5 = + select_metadata_value(metadata, SSEC_KEY_MD5_HEADER)?.ok_or_else(invalid_select_snapshot_sse_metadata)?; + if stored_md5 != key_md5 { + return Err(invalid_select_snapshot_sse_metadata()); + } + insert_select_snapshot_header(&mut response_headers, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, algorithm)?; + insert_select_snapshot_header(&mut response_headers, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, key_md5)?; + } + } + Ok(response_headers) +} + +async fn send_select_events_until_deadline( output: SendableRecordBatchStream, tx: mpsc::Sender>, terminal_permit: mpsc::OwnedPermit>, validation: SelectValidation, deadline: Instant, timeout_seconds: u64, + snapshot_lease: L, ) { - let outcome = match timeout_at(deadline, send_select_events(output, &tx, validation)).await { + let outcome = match timeout_at(deadline, send_select_events(output, &tx, validation, &snapshot_lease)).await { Ok(outcome) => outcome, Err(_) => SelectProducerOutcome::Terminal(Err(map_query_error_to_s3( S3SelectPolicyError::QueryTimeout { @@ -119,12 +347,14 @@ async fn send_select_events_until_deadline( if let SelectProducerOutcome::Terminal(event) = outcome { terminal_permit.send(event); } + drop(snapshot_lease); } async fn send_select_events( mut output: SendableRecordBatchStream, tx: &mpsc::Sender>, validation: SelectValidation, + snapshot_fence: &impl SelectSnapshotFence, ) -> SelectProducerOutcome { let mut encoder = SelectOutputEncoder::new(validation.output_format); let mut progress = SelectProgress::default(); @@ -180,12 +410,18 @@ async fn send_select_events( } } + if let Err(error) = snapshot_fence.ensure_snapshot_valid() { + return SelectProducerOutcome::Terminal(Err(error)); + } let stats = SelectObjectContentEvent::Stats(StatsEvent { details: Some(progress.to_stats()), }); if tx.send(Ok(stats)).await.is_err() { return SelectProducerOutcome::ReceiverClosed; } + if let Err(error) = snapshot_fence.ensure_snapshot_valid() { + return SelectProducerOutcome::Terminal(Err(error)); + } SelectProducerOutcome::Terminal(Ok(SelectObjectContentEvent::End(EndEvent::default()))) } @@ -388,25 +624,35 @@ fn validate_input_delimiter_pair(field_delimiter: Option<&str>, record_delimiter Ok(()) } -async fn preflight_select_object( +async fn prepare_select_object_snapshot( headers: &http::HeaderMap, input: &SelectObjectContentInput, read_principal: Option<&SseKmsPrincipal>, -) -> S3Result { +) -> S3Result { let opts = get_opts(&input.bucket, &input.key, None, None, headers) .await .map_err(ApiError::from)?; let store = get_validated_store(&input.bucket).await?; - let info = store - .get_object_info(&input.bucket, &input.key, &opts) + let snapshot = store + .prepare_select_object_snapshot(&input.bucket, &input.key, headers, &opts) .await - .map_err(ApiError::from)?; + .map_err(map_prepare_snapshot_error)?; + let info = snapshot.object_info(); validate_sse_headers_for_read(&info.user_defined, headers)?; validate_ssec_for_read(&info.user_defined, input.sse_customer_key.as_ref(), input.sse_customer_key_md5.as_ref())?; authorize_sse_kms_object_read(read_principal, &info.user_defined).await?; - Ok(SelectObjectMetadata { - size: info.size.max(0) as u64, - }) + Ok(snapshot) +} + +fn map_prepare_snapshot_error(err: StoragePrepareSelectObjectSnapshotError) -> S3Error { + match err { + StoragePrepareSelectObjectSnapshotError::Storage(err) => ApiError::from(err).into(), + err => { + let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, err.to_string()); + s3_error.set_source(Box::new(err)); + s3_error + } + } } fn log_select_request_summary(input: &SelectObjectContentInput, validation: &SelectValidation) { @@ -567,6 +813,12 @@ fn clamp_i64(value: u64) -> i64 { } fn map_query_error_to_s3(err: QueryError) -> S3Error { + if err.is_snapshot_consistency_error() { + let message = err.to_string(); + let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, message); + s3_error.set_source(Box::new(err)); + return s3_error; + } if let Some(policy_error) = err.s3_select_policy_error() { let message = policy_error.to_string(); return match policy_error { @@ -616,6 +868,10 @@ fn map_query_error_to_s3(err: QueryError) -> S3Error { } } +fn select_query_timeout_error(seconds: u64) -> S3Error { + map_query_error_to_s3(S3SelectPolicyError::QueryTimeout { seconds }.into()) +} + fn looks_like_bucket_not_found(message: &str) -> bool { message.contains("NoSuchBucket") || message.contains("bucket not found") || message.contains("BucketNotFound") } @@ -692,9 +948,71 @@ mod tests { physical_plan::stream::RecordBatchStreamAdapter, sql::sqlparser::parser::ParserError, }; - use http::HeaderMap; + use rustfs_test_utils::TestECStoreEnv; use s3s::dto::{CSVInput, ParquetInput, ScanRange}; + struct LeaseDropSignal(Option>); + + impl Drop for LeaseDropSignal { + fn drop(&mut self) { + if let Some(tx) = self.0.take() { + let _ = tx.send(()); + } + } + } + + impl SelectSnapshotFence for LeaseDropSignal { + fn ensure_snapshot_valid(&self) -> S3Result<()> { + Ok(()) + } + } + + struct FailingSnapshotFence; + + impl SelectSnapshotFence for FailingSnapshotFence { + fn ensure_snapshot_valid(&self) -> S3Result<()> { + Err(S3Error::with_message(S3ErrorCode::InternalError, "snapshot lease was lost")) + } + } + + struct FailsAfterFirstSnapshotFence(std::sync::atomic::AtomicUsize); + + impl SelectSnapshotFence for FailsAfterFirstSnapshotFence { + fn ensure_snapshot_valid(&self) -> S3Result<()> { + if self.0.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 { + Ok(()) + } else { + Err(S3Error::with_message(S3ErrorCode::InternalError, "snapshot lease was lost")) + } + } + } + + fn lease_drop_signal() -> (LeaseDropSignal, tokio::sync::oneshot::Receiver<()>) { + let (tx, rx) = tokio::sync::oneshot::channel(); + (LeaseDropSignal(Some(tx)), rx) + } + + #[tokio::test] + async fn storage_snapshot_fence_forwards_lost_lease() { + let env = TestECStoreEnv::builder() + .prefix("select_snapshot_fence_adapter") + .build() + .await; + env.make_bucket("select-snapshot-fence-adapter", false).await; + env.put_object_bytes("select-snapshot-fence-adapter", "input.csv", b"value\nold\n".to_vec()) + .await; + let snapshot = env + .prepare_select_object_snapshot("select-snapshot-fence-adapter", "input.csv") + .await; + snapshot.mark_lost_for_test(); + + let error = SelectSnapshotFence::ensure_snapshot_valid(&snapshot) + .expect_err("production fence adapter must reject a lost storage snapshot"); + + assert_eq!(error.code(), &S3ErrorCode::InternalError); + assert!(error.to_string().contains("namespace lock was lost")); + } + #[derive(Debug)] struct CyclicError; @@ -744,15 +1062,169 @@ mod tests { } } + #[test] + fn select_snapshot_sse_s3_headers_are_whitelisted() { + let metadata = HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_AES.to_string()), + (SELECT_RUSTFS_KMS_KEY_ID.to_string(), "default".to_string()), + (SELECT_MINIO_KMS_KEY_ID.to_string(), "default".to_string()), + ("x-amz-meta-private".to_string(), "private-value".to_string()), + ]); + + let headers = select_snapshot_sse_response_headers(&metadata, &HeaderMap::new()) + .expect("valid SSE-S3 snapshot metadata should project response headers"); + + assert_eq!(headers.len(), 1); + assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION).expect("SSE-S3 mode"), "AES256"); + assert!(headers.get("x-amz-meta-private").is_none()); + } + + #[test] + fn select_snapshot_sse_kms_headers_use_snapshot_metadata() { + let context = "eyJ0ZW5hbnQiOiJvbmUifQ=="; + for key_id in ["key-1", "arn:aws:kms:key-2"] { + let metadata = HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()), + (AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), key_id.to_string()), + (SELECT_RUSTFS_KMS_KEY_ID.to_string(), key_id.to_string()), + (SELECT_MINIO_KMS_KEY_ID.to_string(), key_id.to_string()), + (SELECT_MINIO_KMS_CONTEXT.to_string(), context.to_string()), + ]); + + let headers = select_snapshot_sse_response_headers(&metadata, &HeaderMap::new()) + .expect("valid SSE-KMS snapshot metadata should project response headers"); + let expected_key_id = if key_id.starts_with(SELECT_KMS_ARN_PREFIX) { + key_id.to_string() + } else { + format!("{SELECT_KMS_ARN_PREFIX}{key_id}") + }; + + assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION).expect("SSE-KMS mode"), "aws:kms"); + assert_eq!( + headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_AWS_KMS_KEY_ID) + .expect("SSE-KMS key ID") + .to_str() + .expect("SSE-KMS key ID should be valid text"), + expected_key_id + ); + assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CONTEXT).expect("SSE-KMS context"), context); + } + } + + #[test] + fn select_snapshot_sse_c_headers_never_echo_the_customer_key() { + let key_md5 = "customer-key-md5"; + let metadata = HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_AES.to_string()), + (SSEC_ALGORITHM_HEADER.to_string(), AMZ_ENCRYPTION_AES.to_string()), + (SSEC_KEY_MD5_HEADER.to_string(), key_md5.to_string()), + ("x-amz-meta-private".to_string(), "private-value".to_string()), + ]); + let mut request_headers = HeaderMap::new(); + request_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, HeaderValue::from_static("AES256")); + request_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, HeaderValue::from_static(key_md5)); + request_headers.insert( + http::HeaderName::from_static(SSEC_KEY_HEADER), + HeaderValue::from_static("must-not-be-returned"), + ); + + let headers = select_snapshot_sse_response_headers(&metadata, &request_headers) + .expect("validated SSE-C request values should project response headers"); + + assert_eq!(headers.len(), 2); + assert_eq!( + headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM) + .expect("SSE-C algorithm"), + "AES256" + ); + assert_eq!( + headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5) + .expect("SSE-C key MD5"), + key_md5 + ); + assert!(headers.get(SSEC_KEY_HEADER).is_none()); + assert!(headers.get("x-amz-meta-private").is_none()); + } + + #[test] + fn select_snapshot_sse_headers_fail_closed_on_corrupt_metadata() { + let invalid_context = "not-base64"; + let persisted_key = "must-not-leak"; + let corrupt_metadata = [ + HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_ascii_lowercase(), "AES256".to_string()), + (AMZ_SERVER_SIDE_ENCRYPTION.to_ascii_uppercase(), "aws:kms".to_string()), + ]), + HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()), + (SSEC_ALGORITHM_HEADER.to_string(), "AES256".to_string()), + ]), + HashMap::from([(SELECT_MINIO_KMS_SEALED_KEY.to_string(), "sealed".to_string())]), + HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()), + (AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), "key-1".to_string()), + (SELECT_RUSTFS_KMS_KEY_ID.to_string(), "key-2".to_string()), + ]), + HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()), + (AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), "key-1".to_string()), + (SELECT_MINIO_KMS_CONTEXT.to_string(), invalid_context.to_string()), + ]), + HashMap::from([ + (AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_KMS.to_string()), + (AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), "key-1".to_string()), + (AMZ_SERVER_SIDE_ENCRYPTION_KMS_CONTEXT.to_string(), "persisted-context".to_string()), + ]), + HashMap::from([ + (SSEC_ALGORITHM_HEADER.to_string(), "AES256".to_string()), + (SSEC_KEY_MD5_HEADER.to_string(), "customer-key-md5".to_string()), + (SSEC_KEY_HEADER.to_string(), persisted_key.to_string()), + ]), + ]; + + for metadata in corrupt_metadata { + let error = select_snapshot_sse_response_headers(&metadata, &HeaderMap::new()) + .expect_err("corrupt snapshot encryption metadata must fail closed"); + assert_eq!(error.code(), &S3ErrorCode::InternalError); + assert!(!error.to_string().contains(invalid_context)); + assert!(!error.to_string().contains(persisted_key)); + } + } + + #[test] + fn select_response_projects_snapshot_sse_headers() { + let (_tx, rx) = mpsc::channel(1); + let metadata = HashMap::from([(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_AES.to_string())]); + + let response = select_object_response(rx, &metadata, &HeaderMap::new()) + .expect("valid snapshot metadata should produce a Select response"); + + assert_eq!( + response + .headers + .get(X_AMZ_SERVER_SIDE_ENCRYPTION) + .expect("snapshot SSE response header"), + AMZ_ENCRYPTION_AES + ); + } + fn spawn_test_producer( output: SendableRecordBatchStream, channel_capacity: usize, - ) -> (tokio::task::JoinHandle<()>, mpsc::Receiver>) { + ) -> ( + tokio::task::JoinHandle<()>, + mpsc::Receiver>, + tokio::sync::oneshot::Receiver<()>, + ) { let (tx, rx) = mpsc::channel(channel_capacity); let terminal_permit = tx .clone() .try_reserve_owned() .expect("test channel should reserve terminal capacity"); + let (lease, lease_released) = lease_drop_signal(); let producer = tokio::spawn(send_select_events_until_deadline( output, tx, @@ -760,8 +1232,9 @@ mod tests { csv_validation(), Instant::now() + std::time::Duration::from_secs(1), 300, + lease, )); - (producer, rx) + (producer, rx, lease_released) } #[test] @@ -844,6 +1317,28 @@ mod tests { assert_eq!(invalid_object_size.code(), &S3ErrorCode::InternalError); } + #[test] + fn prepare_snapshot_storage_error_preserves_existing_s3_mapping() { + let err = map_prepare_snapshot_error(StoragePrepareSelectObjectSnapshotError::Storage(StorageError::ObjectNotFound( + "bucket".to_string(), + "object.csv".to_string(), + ))); + + assert_eq!(err.code(), &S3ErrorCode::NoSuchKey); + assert!(err.source().is_some()); + } + + #[test] + fn prepare_snapshot_invalid_logical_size_fails_with_internal_error_and_source() { + let err = map_prepare_snapshot_error(StoragePrepareSelectObjectSnapshotError::InvalidLogicalSize { size: -1 }); + + assert_eq!(err.code(), &S3ErrorCode::InternalError); + assert!( + err.source() + .is_some_and(|source| source.downcast_ref::().is_some()) + ); + } + #[test] fn error_source_matching_stops_at_the_depth_bound() { let err = CyclicError; @@ -867,6 +1362,7 @@ mod tests { tx.send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) .await .expect("test channel should accept the prefilled event"); + let (lease, lease_released) = lease_drop_signal(); let producer = tokio::spawn(send_select_events_until_deadline( output, tx, @@ -874,6 +1370,7 @@ mod tests { csv_validation(), Instant::now() + std::time::Duration::from_secs(1), 300, + lease, )); tokio::task::yield_now().await; @@ -889,6 +1386,7 @@ mod tests { .expect_err("terminal event should be an error"); assert_eq!(timeout_error.code(), &S3ErrorCode::Busy); assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "timeout should release the snapshot lease"); } #[tokio::test(start_paused = true)] @@ -904,7 +1402,7 @@ mod tests { }; let batches = [Ok(batch("a")), Ok(batch("b"))]; let output = Box::pin(RecordBatchStreamAdapter::new(schema, futures::stream::iter(batches))); - let (producer, mut rx) = spawn_test_producer(output, 8); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 8); producer.await.expect("producer should finish at query EOF"); @@ -921,6 +1419,7 @@ mod tests { assert_eq!(stats.details.and_then(|details| details.bytes_returned), Some(4)); assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_))))); assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "End should release the snapshot lease"); } #[tokio::test(start_paused = true)] @@ -932,7 +1431,7 @@ mod tests { None::<(Result, ())> }), )); - let (producer, mut rx) = spawn_test_producer(output, 3); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 3); tokio::task::yield_now().await; tokio::time::advance(std::time::Duration::from_secs(1)).await; @@ -947,6 +1446,7 @@ mod tests { assert!(matches!(stats, SelectObjectContentEvent::Stats(_))); assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_))))); assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "EOF should release the snapshot lease"); } #[tokio::test(start_paused = true)] @@ -958,7 +1458,7 @@ mod tests { Err(DataFusionError::External(Box::new(S3SelectPolicyError::QueryConcurrencyLimit))) }), )); - let (producer, mut rx) = spawn_test_producer(output, 2); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 2); tokio::task::yield_now().await; tokio::time::advance(std::time::Duration::from_secs(1)).await; @@ -972,6 +1472,7 @@ mod tests { .expect_err("terminal event should be an error"); assert_eq!(stream_error.code(), &S3ErrorCode::SlowDown); assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "stream error should release the snapshot lease"); } #[tokio::test(start_paused = true)] @@ -983,7 +1484,7 @@ mod tests { schema, futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), )); - let (producer, mut rx) = spawn_test_producer(output, 2); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 2); tokio::task::yield_now().await; tokio::time::advance(std::time::Duration::from_secs(1)).await; @@ -997,10 +1498,11 @@ mod tests { .expect_err("terminal event should be an error"); assert_eq!(encoder_error.code(), &S3ErrorCode::InternalError); assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "encoder error should release the snapshot lease"); } #[tokio::test(start_paused = true)] - async fn producer_drops_query_stream_when_receiver_closes() { + async fn producer_drops_query_stream_and_snapshot_lease_when_receiver_closes() { let (stream_dropped_tx, stream_dropped_rx) = tokio::sync::oneshot::channel::<()>(); let output = Box::pin(RecordBatchStreamAdapter::new( Arc::new(Schema::empty()), @@ -1010,7 +1512,20 @@ mod tests { }), )); let (tx, mut rx) = mpsc::channel(2); - let producer = send_select_events(output, &tx, csv_validation()); + let terminal_permit = tx + .clone() + .try_reserve_owned() + .expect("test channel should reserve terminal capacity"); + let (lease, lease_released) = lease_drop_signal(); + let producer = send_select_events_until_deadline( + output, + tx, + terminal_permit, + csv_validation(), + Instant::now() + std::time::Duration::from_secs(1), + 300, + lease, + ); tokio::pin!(producer); assert!(futures::poll!(producer.as_mut()).is_pending()); @@ -1025,6 +1540,7 @@ mod tests { stream_dropped_rx.await.is_err(), "query stream should be dropped when the receiver closes" ); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); } #[tokio::test(start_paused = true)] @@ -1040,7 +1556,8 @@ mod tests { }), )); let (tx, mut rx) = mpsc::channel(2); - let producer = send_select_events(output, &tx, csv_validation()); + let snapshot_fence = LeaseDropSignal(None); + let producer = send_select_events(output, &tx, csv_validation(), &snapshot_fence); tokio::pin!(producer); assert!(futures::poll!(producer.as_mut()).is_pending()); @@ -1058,6 +1575,63 @@ mod tests { ); } + #[tokio::test] + async fn producer_rejects_successful_end_when_final_snapshot_fence_fails() { + let schema = Arc::new(Schema::new(vec![Field::new( + "value", + datafusion::arrow::datatypes::DataType::Utf8, + false, + )])); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(datafusion::arrow::array::StringArray::from(vec!["old-generation"]))], + ) + .expect("test record batch should be valid"); + let output = Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), + )); + let (tx, mut rx) = mpsc::channel(4); + + let outcome = send_select_events(output, &tx, csv_validation(), &FailingSnapshotFence).await; + + let SelectProducerOutcome::Terminal(Err(error)) = outcome else { + panic!("failed final snapshot fence must produce a terminal error"); + }; + assert_eq!(error.code(), &S3ErrorCode::InternalError); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Records(_))))); + assert!(rx.try_recv().is_err(), "failed final fence must not enqueue Stats or End"); + } + + #[tokio::test] + async fn producer_rechecks_snapshot_after_stats_backpressure() { + let output = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + futures::stream::empty::>(), + )); + let (tx, mut rx) = mpsc::channel(2); + let _terminal_permit = tx + .clone() + .try_reserve_owned() + .expect("test channel should reserve terminal capacity"); + let snapshot_fence = FailsAfterFirstSnapshotFence(std::sync::atomic::AtomicUsize::new(0)); + let producer = send_select_events(output, &tx, csv_validation(), &snapshot_fence); + tokio::pin!(producer); + + assert!(futures::poll!(producer.as_mut()).is_pending()); + assert_eq!(snapshot_fence.0.load(std::sync::atomic::Ordering::Relaxed), 1); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + + let SelectProducerOutcome::Terminal(Err(error)) = producer.await else { + panic!("snapshot loss during Stats backpressure must reject successful End"); + }; + assert_eq!(error.code(), &S3ErrorCode::InternalError); + assert_eq!(snapshot_fence.0.load(std::sync::atomic::Ordering::Relaxed), 2); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Stats(_))))); + assert!(rx.try_recv().is_err(), "snapshot loss after Stats must not enqueue End"); + } + #[test] fn validate_defaults_csv_header_and_compression() { let mut input = base_input(); diff --git a/rustfs/src/app/storage_api.rs b/rustfs/src/app/storage_api.rs index 38d96fbb1..91579de7d 100644 --- a/rustfs/src/app/storage_api.rs +++ b/rustfs/src/app/storage_api.rs @@ -1153,14 +1153,13 @@ pub(crate) mod multipart_usecase { } pub(crate) mod select_object { - pub(crate) mod contract { - pub(crate) mod object { - pub(crate) use super::super::super::storage_contracts::ObjectOperations; - } - } - pub(crate) use super::{options, request_context, sse}; - pub(crate) use crate::storage::storage_api::{get_validated_store, validate_sse_headers_for_read, validate_ssec_for_read}; + #[cfg(test)] + pub(crate) use crate::storage::storage_api::StorageError; + pub(crate) use crate::storage::storage_api::{ + StoragePrepareSelectObjectSnapshotError, StorageSelectObjectSnapshot, get_validated_store, validate_sse_headers_for_read, + validate_ssec_for_read, + }; } pub(crate) mod context { diff --git a/rustfs/src/storage/storage_api.rs b/rustfs/src/storage/storage_api.rs index ef7dac0ad..3b82b0966 100644 --- a/rustfs/src/storage/storage_api.rs +++ b/rustfs/src/storage/storage_api.rs @@ -89,6 +89,8 @@ pub(crate) type StorageGetObjectReader = super::GetObjectReader; pub(crate) type StorageObjectInfo = super::ObjectInfo; pub(crate) type StorageObjectLockDeleteOptions = contract::object::ObjectLockDeleteOptions; pub(crate) type StorageObjectOptions = super::ObjectOptions; +pub(crate) type StoragePrepareSelectObjectSnapshotError = ecstore_object::PrepareSelectObjectSnapshotError; +pub(crate) type StorageSelectObjectSnapshot = ecstore_object::SelectObjectSnapshot; pub(crate) type StorageObjectToDelete = contract::object::ObjectToDelete; pub(crate) type StoragePutObjReader = super::PutObjReader; pub(crate) use super::ecfs_extend::{ @@ -523,9 +525,10 @@ pub(crate) mod ecstore_object { pub(crate) use rustfs_ecstore::api::object::GetObjectBodySource; pub(crate) use rustfs_ecstore::api::object::{ EncryptionResolutionError, EncryptionResolutionErrorKind, GetObjectBodyCacheHook, GetObjectBodyCacheHookLookup, - ObjectEncryptionResolver, ObjectMutationHook, ReadEncryptionMaterial, ReadEncryptionMode, ReadEncryptionRequest, - get_object_body_cache_plaintext_len, lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, - register_object_mutation_hook, unregister_get_object_body_cache_hook, unregister_object_mutation_hook, + ObjectEncryptionResolver, ObjectMutationHook, PrepareSelectObjectSnapshotError, ReadEncryptionMaterial, + ReadEncryptionMode, ReadEncryptionRequest, SelectObjectSnapshot, get_object_body_cache_plaintext_len, + lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, register_object_mutation_hook, + unregister_get_object_body_cache_hook, unregister_object_mutation_hook, }; } diff --git a/rustfs/tests/embedded_select_snapshot_test.rs b/rustfs/tests/embedded_select_snapshot_test.rs new file mode 100644 index 000000000..ea15dac6d --- /dev/null +++ b/rustfs/tests/embedded_select_snapshot_test.rs @@ -0,0 +1,227 @@ +// Copyright 2024 RustFS Team +// +// Licensed 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. + +use aws_sdk_s3::config::{Credentials, Region}; +use aws_sdk_s3::primitives::ByteStream; +use aws_sdk_s3::types::{ + BucketVersioningStatus, CsvInput, CsvOutput, ExpressionType, FileHeaderInfo, InputSerialization, OutputSerialization, + VersioningConfiguration, +}; +use aws_sdk_s3::{Client, Config}; +use bytes::Bytes; +use rustfs::embedded::{RustFSServerBuilder, find_available_port}; +use rustfs_ecstore::api::set_disk::test_util::{PutObjectCommitBarrier, PutObjectCommitPause}; +use std::time::Duration; + +mod common; + +const OBJECT: &str = "snapshot-race.csv"; +const CSV_HEADER: &[u8] = b"generation\n"; +const OLD_GENERATION: &[u8] = b"OLD_GENERATION_POISON_1629"; +const NEW_GENERATION: &[u8] = b"NEW_GENERATION_POISON_1629"; +const ROW_COUNT: usize = 200_000; +const OPERATION_TIMEOUT: Duration = Duration::from_secs(90); + +fn s3_client(endpoint: &str, access_key: &str, secret_key: &str) -> Client { + let credentials = Credentials::new(access_key, secret_key, None, None, "select-snapshot-test"); + let config = Config::builder() + .credentials_provider(credentials) + .region(Region::new("us-east-1")) + .endpoint_url(endpoint) + .force_path_style(true) + .behavior_version_latest() + .build(); + Client::from_conf(config) +} + +fn snapshot_csv(generation: &[u8]) -> Bytes { + let mut body = Vec::with_capacity(CSV_HEADER.len() + ROW_COUNT * (generation.len() + 1)); + body.extend_from_slice(CSV_HEADER); + for _ in 0..ROW_COUNT { + body.extend_from_slice(generation); + body.push(b'\n'); + } + Bytes::from(body) +} + +async fn put_generation(client: &Client, bucket: &str, body: Bytes) -> Option { + tokio::time::timeout( + OPERATION_TIMEOUT, + client + .put_object() + .bucket(bucket) + .key(OBJECT) + .body(ByteStream::from(body)) + .send(), + ) + .await + .expect("snapshot PUT should finish before the test timeout") + .expect("snapshot PUT should succeed") + .version_id +} + +async fn start_select(client: &Client, bucket: &str) -> aws_sdk_s3::operation::select_object_content::SelectObjectContentOutput { + tokio::time::timeout( + OPERATION_TIMEOUT, + client + .select_object_content() + .bucket(bucket) + .key(OBJECT) + .expression("SELECT * FROM S3Object") + .expression_type(ExpressionType::Sql) + .input_serialization( + InputSerialization::builder() + .csv(CsvInput::builder().file_header_info(FileHeaderInfo::Use).build()) + .build(), + ) + .output_serialization( + OutputSerialization::builder() + .csv(CsvOutput::builder().record_delimiter("\n").field_delimiter(",").build()) + .build(), + ) + .send(), + ) + .await + .expect("SelectObjectContent should start before the test timeout") + .expect("SelectObjectContent should start") +} + +async fn collect_select(mut response: aws_sdk_s3::operation::select_object_content::SelectObjectContentOutput) -> Vec { + tokio::time::timeout(OPERATION_TIMEOUT, async { + let mut output = Vec::new(); + let mut saw_end = false; + while let Some(event) = response.payload.recv().await.expect("Select event should decode") { + match event { + aws_sdk_s3::types::SelectObjectContentEventStream::Records(records) => { + if let Some(payload) = records.payload { + output.extend_from_slice(payload.as_ref()); + } + } + aws_sdk_s3::types::SelectObjectContentEventStream::End(_) => { + saw_end = true; + break; + } + _ => {} + } + } + assert!(saw_end, "Select response should contain an End event"); + output + }) + .await + .expect("Select response should finish before the test timeout") +} + +fn assert_version_advanced(previous: &Option, current: &Option, versioned: bool) { + if !versioned { + return; + } + assert!( + previous.as_deref().is_some_and(|version| !version.is_empty()), + "versioned fixture PUT should return a version ID" + ); + assert!( + current.as_deref().is_some_and(|version| !version.is_empty()), + "versioned overwrite should return a version ID" + ); + assert_ne!(previous, current, "versioned overwrite should create a new generation"); +} + +async fn pending_overwrite_keeps_select_on_one_generation( + client: &Client, + bucket: &str, + expected_body: &Bytes, + replacement_body: Bytes, +) -> Option { + let barrier = PutObjectCommitBarrier::install(bucket, OBJECT, PutObjectCommitPause::BeforeNamespace); + let overwrite_client = client.clone(); + let overwrite_bucket = bucket.to_string(); + let overwrite = tokio::spawn(async move { put_generation(&overwrite_client, &overwrite_bucket, replacement_body).await }); + + barrier.wait_until_paused().await; + let response = start_select(client, bucket).await; + barrier.release_and_wait_until_namespace_pending().await; + + let output = collect_select(response).await; + assert_eq!( + output.as_slice(), + &expected_body[CSV_HEADER.len()..], + "Select should return only the generation captured before overwrite" + ); + + tokio::time::timeout(OPERATION_TIMEOUT, overwrite) + .await + .expect("overwrite should finish after Select releases its snapshot") + .expect("overwrite task should not panic") +} + +async fn run_snapshot_case(client: &Client, versioned: bool) { + let bucket = if versioned { + "select-snapshot-versioned" + } else { + "select-snapshot-unversioned" + }; + client + .create_bucket() + .bucket(bucket) + .send() + .await + .expect("create snapshot test bucket"); + if versioned { + client + .put_bucket_versioning() + .bucket(bucket) + .versioning_configuration( + VersioningConfiguration::builder() + .status(BucketVersioningStatus::Enabled) + .build(), + ) + .send() + .await + .expect("enable snapshot test bucket versioning"); + } + + let old_body = snapshot_csv(OLD_GENERATION); + let new_body = snapshot_csv(NEW_GENERATION); + let initial_version = put_generation(client, bucket, old_body.clone()).await; + let overwrite_version = pending_overwrite_keeps_select_on_one_generation(client, bucket, &old_body, new_body.clone()).await; + assert_version_advanced(&initial_version, &overwrite_version, versioned); +} + +#[test] +fn select_snapshot_is_stable_during_http_overwrite() { + common::run_embedded_test(|| async { + let port = match find_available_port() { + Ok(port) => port, + Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => return, + Err(error) => panic!("find free port for Select snapshot test: {error}"), + }; + let server = RustFSServerBuilder::new() + .address(format!("127.0.0.1:{port}")) + .access_key("select-snapshot-access") + .secret_key("select-snapshot-secret") + .build() + .await + .expect("start embedded RustFS for Select snapshot test"); + assert!( + server.endpoint().ends_with(&format!(":{port}")), + "embedded server should bind the requested port" + ); + let client = s3_client(&server.endpoint(), server.access_key(), server.secret_key()); + + run_snapshot_case(&client, false).await; + run_snapshot_case(&client, true).await; + + server.shutdown().await; + }); +}