fix(select): pin object snapshot for query lifetime (#5835)

This commit is contained in:
GatewayJ
2026-08-09 14:08:53 +08:00
committed by GitHub
parent b9d1ca3e4d
commit 70deb3284b
28 changed files with 3806 additions and 572 deletions
Generated
+3
View File
@@ -9378,6 +9378,7 @@ dependencies = [
"thiserror 2.0.20", "thiserror 2.0.20",
"time", "time",
"tokio", "tokio",
"tokio-stream",
"tokio-util", "tokio-util",
"tonic", "tonic",
"tower", "tower",
@@ -10117,6 +10118,7 @@ dependencies = [
"tracing", "tracing",
"transform-stream", "transform-stream",
"url", "url",
"uuid",
] ]
[[package]] [[package]]
@@ -10132,6 +10134,7 @@ dependencies = [
"hotpath", "hotpath",
"parking_lot", "parking_lot",
"rustfs-s3select-api", "rustfs-s3select-api",
"rustfs-test-utils",
"s3s", "s3s",
"tokio", "tokio",
"tracing", "tracing",
+1
View File
@@ -181,6 +181,7 @@ path-absolutize = { workspace = true }
rmp.workspace = true rmp.workspace = true
rmp-serde.workspace = true rmp-serde.workspace = true
tokio-util = { workspace = true, features = ["io", "compat"] } tokio-util = { workspace = true, features = ["io", "compat"] }
tokio-stream = { workspace = true, features = ["sync"] }
base64 = { workspace = true } base64 = { workspace = true }
hmac = { workspace = true } hmac = { workspace = true }
sha1 = { workspace = true } sha1 = { workspace = true }
+9 -1
View File
@@ -414,7 +414,10 @@ pub mod object {
lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, register_object_mutation_hook, 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, 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 { pub mod rebalance {
@@ -449,6 +452,11 @@ pub mod rpc {
pub mod set_disk { pub mod set_disk {
pub use crate::set_disk::{DEFAULT_READ_BUFFER_SIZE, SetDisks, get_lock_acquire_timeout, is_valid_storage_class}; 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 { pub mod store_list {
+98 -81
View File
@@ -1215,6 +1215,98 @@ async fn init_storage_disks_with_errors(
(disks, errs) (disks, errs)
} }
#[cfg(test)]
pub(crate) async fn make_local_two_set_sets() -> (Vec<tempfile::TempDir>, Arc<Sets>) {
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<InstanceContext>) -> (Vec<tempfile::TempDir>, Arc<Sets>) {
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<dyn rustfs_lock::LockClient>
})
.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)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
@@ -1373,81 +1465,6 @@ mod tests {
assert_eq!(result, (Some(3), Some(1), Some(0))); assert_eq!(result, (Some(3), Some(1), Some(0)));
} }
async fn two_set_test_sets() -> (Vec<tempfile::TempDir>, Arc<Sets>) {
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] #[tokio::test]
async fn heal_object_uses_explicit_set_scope() { async fn heal_object_uses_explicit_set_scope() {
let (_temp_dirs, sets) = two_set_test_sets().await; let (_temp_dirs, sets) = two_set_test_sets().await;
@@ -1497,7 +1514,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn delete_prefix_surfaces_a_hard_error_from_any_set() { 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()); let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default()) sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await .await
@@ -1546,7 +1563,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn delete_prefix_keeps_a_missing_bucket_idempotent_across_sets() { 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()); let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default()) sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await .await
@@ -1585,7 +1602,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn delete_prefix_preserves_a_completely_missing_bucket_error() { 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 bucket = format!("delete-prefix-missing-{}", Uuid::new_v4().simple());
let err = sets let err = sets
@@ -1605,7 +1622,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn delete_prefix_fails_when_one_set_is_entirely_offline() { 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()); let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default()) sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await .await
@@ -1652,7 +1669,7 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn set_format_heal_accepts_quorum_from_a_nonzero_set() { 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] let (result, err) = sets.disk_set[1]
.heal_format(false) .heal_format(false)
@@ -1757,7 +1774,7 @@ mod tests {
#[serial] #[serial]
async fn list_multipart_uploads_merges_all_sets_without_pagination_loss() { async fn list_multipart_uploads_merges_all_sets_without_pagination_loss() {
let _setup_type_guard = SetupTypeGuard::switch_to(SetupType::Erasure).await; 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()); let bucket = format!("multipart-list-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default()) sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await .await
+5
View File
@@ -208,6 +208,11 @@ impl InstanceContext {
} }
} }
#[cfg(test)]
pub(crate) fn with_lock_manager_for_test(lock_manager: Arc<GlobalLockManager>) -> Self {
Self::with_lock_manager(lock_manager)
}
/// This instance's namespace lock manager. /// This instance's namespace lock manager.
pub fn lock_manager(&self) -> Arc<GlobalLockManager> { pub fn lock_manager(&self) -> Arc<GlobalLockManager> {
self.lock_manager.clone() self.lock_manager.clone()
@@ -4651,7 +4651,7 @@ pub(in crate::set_disk) mod cleanup_fault_injection {
/// unobserved object records nothing, keeping the registry bounded, and each /// unobserved object records nothing, keeping the registry bounded, and each
/// [`CallCounterScope`] clears only its own object's counts on drop. /// [`CallCounterScope`] clears only its own object's counts on drop.
#[cfg(test)] #[cfg(test)]
pub(in crate::set_disk) mod disk_call_counters { pub(crate) mod disk_call_counters {
use std::collections::{HashMap, HashSet}; use std::collections::{HashMap, HashSet};
use std::sync::{Mutex, OnceLock}; use std::sync::{Mutex, OnceLock};
+55 -2
View File
@@ -692,6 +692,8 @@ const DEFAULT_RUSTFS_GET_MULTIPART_READER_SETUP_PREFETCH: bool = true;
static OBJECT_LOCK_DIAG_ENABLED: OnceLock<bool> = OnceLock::new(); static OBJECT_LOCK_DIAG_ENABLED: OnceLock<bool> = OnceLock::new();
mod core; mod core;
#[cfg(test)]
pub(crate) use core::io_primitives::disk_call_counters;
mod ctx; mod ctx;
mod metadata; mod metadata;
mod ops; mod ops;
@@ -702,8 +704,8 @@ pub(crate) use ops::object::TransitionCleanupStoreBarrier as SetDiskTransitionCl
pub(crate) use ops::object::body_cache_plaintext_len; pub(crate) use ops::object::body_cache_plaintext_len;
#[cfg(test)] #[cfg(test)]
pub(crate) use ops::object::cleanup_rejected_transition_upload_durably; pub(crate) use ops::object::cleanup_rejected_transition_upload_durably;
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
pub(crate) use ops::object::{PutObjectCommitBarrier, PutObjectCommitPause}; pub use ops::object::{PutObjectCommitBarrier, PutObjectCommitPause};
mod read; mod read;
mod replication; mod replication;
pub(crate) mod shard_source; pub(crate) mod shard_source;
@@ -731,6 +733,10 @@ impl PreparedGetObjectMetadata {
.take() .take()
.expect("prepared GET metadata ObjectInfo must be consumed exactly once") .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! { 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<ObjectLockDiagGuard> {
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)] #[allow(clippy::too_many_arguments)]
fn log_object_lock_acquire_if_slow( fn log_object_lock_acquire_if_slow(
&self, &self,
+55 -16
View File
@@ -1312,7 +1312,7 @@ impl SetDisks {
} }
if !opts.no_lock && object_lock_guard.is_none() { 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; pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeNamespace).await;
if let Some(expected_incarnation_id) = opts.expected_bucket_incarnation_id if let Some(expected_incarnation_id) = opts.expected_bucket_incarnation_id
&& opts.bucket_lifecycle_lock_fence.is_none() && opts.bucket_lifecycle_lock_fence.is_none()
@@ -1324,9 +1324,21 @@ impl SetDisks {
.await?, .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; 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 { 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 requested && fleet_confirmed && fleet_proof_valid
} }
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
#[derive(Clone, Copy, Debug, Eq, PartialEq)] #[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum PutObjectCommitPause { pub enum PutObjectCommitPause {
BeforeNamespace, BeforeNamespace,
AfterNamespace, AfterNamespace,
BeforeMetadata, BeforeMetadata,
} }
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
struct PutObjectCommitBarrierState { struct PutObjectCommitBarrierState {
bucket: String, bucket: String,
object: String, object: String,
pause: PutObjectCommitPause, pause: PutObjectCommitPause,
arrived: tokio::sync::Notify, arrived: tokio::sync::Notify,
release: tokio::sync::Notify, release: tokio::sync::Notify,
namespace_pending: tokio::sync::Notify,
} }
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
pub(crate) struct PutObjectCommitBarrier { pub struct PutObjectCommitBarrier {
state: Arc<PutObjectCommitBarrierState>, state: Arc<PutObjectCommitBarrierState>,
} }
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
static PUT_OBJECT_COMMIT_BARRIER: std::sync::OnceLock<std::sync::Mutex<Vec<Arc<PutObjectCommitBarrierState>>>> = static PUT_OBJECT_COMMIT_BARRIER: std::sync::OnceLock<std::sync::Mutex<Vec<Arc<PutObjectCommitBarrierState>>>> =
std::sync::OnceLock::new(); std::sync::OnceLock::new();
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
impl PutObjectCommitBarrier { 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 { let state = Arc::new(PutObjectCommitBarrierState {
bucket: bucket.to_string(), bucket: bucket.to_string(),
object: object.to_string(), object: object.to_string(),
pause, pause,
arrived: tokio::sync::Notify::new(), arrived: tokio::sync::Notify::new(),
release: tokio::sync::Notify::new(), release: tokio::sync::Notify::new(),
namespace_pending: tokio::sync::Notify::new(),
}); });
let mut slot = PUT_OBJECT_COMMIT_BARRIER let mut slot = PUT_OBJECT_COMMIT_BARRIER
.get_or_init(|| std::sync::Mutex::new(Vec::new())) .get_or_init(|| std::sync::Mutex::new(Vec::new()))
@@ -2626,18 +2640,27 @@ impl PutObjectCommitBarrier {
Self { state } 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()) tokio::time::timeout(Duration::from_secs(30), self.state.arrived.notified())
.await .await
.expect("put object should reach the deterministic commit barrier"); .expect("put object should reach the deterministic commit barrier");
} }
pub(crate) fn release(&self) { pub fn release(&self) {
self.state.release.notify_one(); 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 { impl Drop for PutObjectCommitBarrier {
fn drop(&mut self) { fn drop(&mut self) {
self.state.release.notify_one(); 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) { async fn pause_put_object_commit(bucket: &str, object: &str, pause: PutObjectCommitPause) {
let barrier = PUT_OBJECT_COMMIT_BARRIER let barrier = PUT_OBJECT_COMMIT_BARRIER
.get_or_init(|| std::sync::Mutex::new(Vec::new())) .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)] #[cfg(test)]
struct DeleteObjectCommitBarrierState { struct DeleteObjectCommitBarrierState {
bucket: String, bucket: String,
@@ -4414,7 +4453,7 @@ impl crate::storage_api_contracts::object::ObjectOperations for SetDisks {
self.invalidate_get_object_metadata_cache(bucket, object).await; self.invalidate_get_object_metadata_cache(bucket, object).await;
// Guard lock for metadata update // Guard lock for metadata update
#[cfg(test)] #[cfg(any(test, feature = "test-util"))]
pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeMetadata).await; pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeMetadata).await;
let _lock_guard = if !opts.no_lock { let _lock_guard = if !opts.no_lock {
Some(self.acquire_write_lock_diag("put_object_metadata", bucket, object).await?) Some(self.acquire_write_lock_diag("put_object_metadata", bucket, object).await?)
+4 -1
View File
@@ -152,7 +152,10 @@ mod list;
pub(crate) mod list_objects; pub(crate) mod list_objects;
mod multipart; mod multipart;
mod object; mod object;
pub use object::PreparedGetObjectReader; pub use object::{
PrepareSelectObjectSnapshotError, PreparedGetObjectReader, SelectObjectSnapshot, SelectObjectSnapshotReadError,
SnapshotConsistencyError,
};
mod peer; mod peer;
mod rebalance; mod rebalance;
pub(crate) mod utils; pub(crate) mod utils;
File diff suppressed because it is too large Load Diff
+2 -1
View File
@@ -77,11 +77,12 @@ parking_lot.workspace = true
tokio = { workspace = true, features = ["fs", "rt-multi-thread"] } tokio = { workspace = true, features = ["fs", "rt-multi-thread"] }
tokio-util = { workspace = true, features = ["io", "compat"] } tokio-util = { workspace = true, features = ["io", "compat"] }
tracing.workspace = true tracing.workspace = true
uuid.workspace = true
transform-stream.workspace = true transform-stream.workspace = true
url.workspace = true url.workspace = true
[dev-dependencies] [dev-dependencies]
rustfs-test-utils.workspace = true rustfs-test-utils = { workspace = true, features = ["put-object-commit-barrier"] }
serial_test.workspace = true serial_test.workspace = true
[lib] [lib]
+26 -14
View File
@@ -15,22 +15,23 @@
#![recursion_limit = "256"] #![recursion_limit = "256"]
use datafusion::{common::DataFusionError, sql::sqlparser::parser::ParserError}; use datafusion::{common::DataFusionError, sql::sqlparser::parser::ParserError};
use std::fmt::Display; use std::{error::Error as StdError, fmt::Display};
use thiserror::Error; use thiserror::Error;
pub mod object_store; pub mod object_store;
pub mod query; pub mod query;
pub mod server; pub mod server;
mod storage_api; mod storage_api;
pub use storage_api::SelectObjectSnapshot;
#[cfg(test)] #[cfg(test)]
mod test; mod test;
pub type QueryResult<T> = Result<T, QueryError>; pub type QueryResult<T> = Result<T, QueryError>;
pub(crate) use storage_api::crate_boundary::{ pub(crate) use storage_api::crate_boundary::{
SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError, PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions,
SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError, resolve_select_object_store_handle,
select_is_err_version_not_found, select_is_err_bucket_not_found, select_is_err_object_not_found, select_is_err_version_not_found,
}; };
#[derive(Debug, Error)] #[derive(Debug, Error)]
@@ -79,24 +80,24 @@ pub enum S3SelectPolicyError {
QueryTimeout { seconds: u64 }, QueryTimeout { seconds: u64 },
} }
impl S3SelectPolicyError { impl QueryError {
fn from_error<'a>(mut err: &'a (dyn std::error::Error + 'static)) -> Option<&'a Self> { fn source_error<T: StdError + 'static>(&self) -> Option<&T> {
let mut err: &(dyn StdError + 'static) = self;
for _ in 0..16 { for _ in 0..16 {
if let Some(policy_error) = err.downcast_ref::<Self>() { if let Some(source) = err.downcast_ref::<T>() {
return Some(policy_error); return Some(source);
} }
err = err.source()?; err = err.source()?;
} }
None None
} }
}
impl QueryError { pub fn is_snapshot_consistency_error(&self) -> bool {
self.source_error::<SnapshotConsistencyError>().is_some()
}
pub fn s3_select_policy_error(&self) -> Option<&S3SelectPolicyError> { pub fn s3_select_policy_error(&self) -> Option<&S3SelectPolicyError> {
match self { self.source_error()
Self::Datafusion { source } => S3SelectPolicyError::from_error(source.as_ref()),
_ => None,
}
} }
} }
@@ -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] #[test]
fn test_query_error_from_parser_error() { fn test_query_error_from_parser_error() {
let parser_error = ParserError::ParserError("syntax error".to_string()); let parser_error = ParserError::ParserError("syntax error".to_string());
File diff suppressed because it is too large Load Diff
@@ -22,6 +22,7 @@ use super::{
Query, Query,
execution::{Output, QueryStateMachine}, execution::{Output, QueryStateMachine},
logical_planner::Plan, logical_planner::Plan,
session::QueryAdmission,
}; };
#[async_trait] #[async_trait]
@@ -32,6 +33,14 @@ pub trait QueryDispatcher: Send + Sync {
async fn execute_query(&self, query: &Query) -> QueryResult<Output>; async fn execute_query(&self, query: &Query) -> QueryResult<Output>;
fn try_reserve_query(&self) -> QueryResult<QueryAdmission> {
Ok(QueryAdmission::unmanaged())
}
async fn execute_query_admitted(&self, query: &Query, _admission: QueryAdmission) -> QueryResult<Output> {
self.execute_query(query).await
}
async fn build_logical_plan(&self, query_state_machine: Arc<QueryStateMachine>) -> QueryResult<Option<Plan>>; async fn build_logical_plan(&self, query_state_machine: Arc<QueryStateMachine>) -> QueryResult<Option<Plan>>;
async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc<QueryStateMachine>) -> QueryResult<Output>; async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc<QueryStateMachine>) -> QueryResult<Output>;
+21 -1
View File
@@ -15,6 +15,8 @@
use s3s::dto::SelectObjectContentInput; use s3s::dto::SelectObjectContentInput;
use std::sync::Arc; use std::sync::Arc;
use crate::SelectObjectSnapshot;
pub mod analyzer; pub mod analyzer;
pub mod ast; pub mod ast;
pub mod dispatcher; pub mod dispatcher;
@@ -37,12 +39,26 @@ pub struct Context {
pub struct Query { pub struct Query {
context: Context, context: Context,
content: String, content: String,
snapshot: Option<Arc<SelectObjectSnapshot>>,
} }
impl Query { impl Query {
#[inline(always)] #[inline(always)]
pub fn new(context: Context, content: String) -> Self { 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<SelectObjectSnapshot>) -> Self {
Self {
context,
content,
snapshot: Some(snapshot),
}
} }
pub fn context(&self) -> &Context { pub fn context(&self) -> &Context {
@@ -52,4 +68,8 @@ impl Query {
pub fn content(&self) -> &str { pub fn content(&self) -> &str {
self.content.as_str() self.content.as_str()
} }
pub fn snapshot(&self) -> Option<&Arc<SelectObjectSnapshot>> {
self.snapshot.as_ref()
}
} }
+155 -14
View File
@@ -12,14 +12,16 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
use crate::SelectObjectSnapshot;
use crate::query::Context; use crate::query::Context;
use crate::{QueryError, QueryResult, SelectStore, object_store::EcObjectStore}; use crate::{QueryError, QueryResult, object_store::EcObjectStore};
use datafusion::{ use datafusion::{
arrow::{ arrow::{
array::{Int32Array, StringArray}, array::{Int32Array, StringArray},
datatypes::{DataType, Field, Schema}, datatypes::{DataType, Field, Schema},
record_batch::RecordBatch, record_batch::RecordBatch,
}, },
common::DataFusionError,
execution::{SessionStateBuilder, config::SessionConfig, context::SessionState, runtime_env::RuntimeEnvBuilder}, execution::{SessionStateBuilder, config::SessionConfig, context::SessionState, runtime_env::RuntimeEnvBuilder},
object_store::{ObjectStore, ObjectStoreExt, memory::InMemory, path::Path}, object_store::{ObjectStore, ObjectStoreExt, memory::InMemory, path::Path},
parquet::arrow::ArrowWriter, parquet::arrow::ArrowWriter,
@@ -39,6 +41,28 @@ use tracing::error;
pub type QueryExecutionGuard = Arc<OwnedSemaphorePermit>; pub type QueryExecutionGuard = Arc<OwnedSemaphorePermit>;
/// 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<QueryExecutionGuard>,
}
impl QueryAdmission {
pub fn new(query_guard: QueryExecutionGuard) -> Self {
Self {
query_guard: Some(query_guard),
}
}
pub fn into_query_guard(mut self) -> Option<QueryExecutionGuard> {
self.query_guard.take()
}
pub(crate) fn unmanaged() -> Self {
Self { query_guard: None }
}
}
#[derive(Clone, Default)] #[derive(Clone, Default)]
pub struct QueryExecutionOwner { pub struct QueryExecutionOwner {
identity: Arc<()>, identity: Arc<()>,
@@ -300,30 +324,30 @@ impl SessionCtxFactory {
query_tracker: QueryExecutionTracker, query_tracker: QueryExecutionTracker,
memory_limit_bytes: usize, memory_limit_bytes: usize,
) -> QueryResult<SessionCtx> { ) -> QueryResult<SessionCtx> {
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 .await
} }
#[cfg(test)] pub async fn create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
async fn create_session_ctx_with_tracker_and_store(
&self, &self,
context: &Context, context: &Context,
snapshot: Arc<SelectObjectSnapshot>,
query_tracker: QueryExecutionTracker, query_tracker: QueryExecutionTracker,
store: Arc<SelectStore>, memory_limit_bytes: usize,
) -> QueryResult<SessionCtx> { ) -> QueryResult<SessionCtx> {
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 .await
} }
async fn create_session_ctx_inner( async fn create_session_ctx_inner(
&self, &self,
context: &Context, context: &Context,
snapshot: Option<Arc<SelectObjectSnapshot>>,
query_tracker: Option<QueryExecutionTracker>, query_tracker: Option<QueryExecutionTracker>,
store: Option<Arc<SelectStore>>,
memory_limit_bytes: usize, memory_limit_bytes: usize,
) -> QueryResult<SessionCtx> { ) -> QueryResult<SessionCtx> {
let df_session_ctx = self 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?; .await?;
Ok(SessionCtx { Ok(SessionCtx {
@@ -336,8 +360,8 @@ impl SessionCtxFactory {
async fn build_df_session_context( async fn build_df_session_context(
&self, &self,
context: &Context, context: &Context,
snapshot: Option<Arc<SelectObjectSnapshot>>,
query_tracker: Option<QueryExecutionTracker>, query_tracker: Option<QueryExecutionTracker>,
store: Option<Arc<SelectStore>>,
memory_limit_bytes: usize, memory_limit_bytes: usize,
) -> QueryResult<SessionContext> { ) -> QueryResult<SessionContext> {
let path = format!("s3://{}", context.input.bucket); let path = format!("s3://{}", context.input.bucket);
@@ -416,11 +440,13 @@ impl SessionCtxFactory {
} else { } else {
let store: EcObjectStore = match query_tracker { let store: EcObjectStore = match query_tracker {
Some(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() df_session_state.with_object_store(&store_url, Arc::new(store)).build()
}; };
@@ -498,6 +524,7 @@ mod tests {
}, },
execution::memory_pool::MemoryLimit, execution::memory_pool::MemoryLimit,
}; };
use http::HeaderMap;
use s3s::dto::{ use s3s::dto::{
CSVInput, CSVOutput, ExpressionType, InputSerialization, JSONInput, OutputSerialization, ParquetInput, ScanRange, CSVInput, CSVOutput, ExpressionType, InputSerialization, JSONInput, OutputSerialization, ParquetInput, ScanRange,
SelectObjectContentInput, SelectObjectContentRequest, SelectObjectContentInput, SelectObjectContentRequest,
@@ -531,6 +558,16 @@ mod tests {
} }
} }
async fn prepare_test_snapshot(context: &Context) -> Arc<SelectObjectSnapshot> {
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] #[test]
fn session_factory_fields_remain_source_compatible() { fn session_factory_fields_remain_source_compatible() {
let factory = SessionCtxFactory { 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)] #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial_test::serial] #[serial_test::serial]
async fn session_factory_propagates_query_guard_to_ec_store() { async fn session_factory_propagates_query_guard_to_ec_store() {
let env = crate::storage_api::select_test_ecstore_env().await; 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 admission = Arc::new(tokio::sync::Semaphore::new(1));
let permit = Arc::clone(&admission) let permit = Arc::clone(&admission)
@@ -697,7 +781,12 @@ mod tests {
300, 300,
); );
let session = SessionCtxFactory::new(false) 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 .await
.expect("production session should be created with the query guard"); .expect("production session should be created with the query guard");
@@ -706,6 +795,51 @@ mod tests {
assert_eq!(Arc::strong_count(&query_guard), 1); 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)] #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial_test::serial] #[serial_test::serial]
async fn scan_range_is_preserved_across_large_csv_partition_boundary() { async fn scan_range_is_preserved_across_large_csv_partition_boundary() {
@@ -721,6 +855,7 @@ mod tests {
assert!(data.len() > 1024 * 1024); assert!(data.len() > 1024 * 1024);
let mut context = test_context(); 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"); 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 { Arc::make_mut(&mut context.input).request.scan_range = Some(ScanRange {
start: Some(selected_start), start: Some(selected_start),
@@ -732,6 +867,7 @@ mod tests {
.put_object(&context.input.bucket, &context.input.key, &mut reader, &Default::default()) .put_object(&context.input.bucket, &context.input.key, &mut reader, &Default::default())
.await .await
.expect("put large ScanRange CSV fixture"); .expect("put large ScanRange CSV fixture");
let snapshot = prepare_test_snapshot(&context).await;
let admission = Arc::new(tokio::sync::Semaphore::new(1)); let admission = Arc::new(tokio::sync::Semaphore::new(1));
let permit = Arc::clone(&admission) let permit = Arc::clone(&admission)
@@ -746,7 +882,12 @@ mod tests {
); );
let session = SessionCtxFactory::new(false) let session = SessionCtxFactory::new(false)
.with_target_partitions(2) .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 .await
.expect("create production ScanRange session"); .expect("create production ScanRange session");
assert!(!session.inner().config().options().optimizer.repartition_file_scans); assert!(!session.inner().config().options().optimizer.repartition_file_scans);
+10
View File
@@ -20,6 +20,7 @@ use crate::{
Query, Query,
execution::{Output, QueryStateMachineRef}, execution::{Output, QueryStateMachineRef},
logical_planner::Plan, logical_planner::Plan,
session::QueryAdmission,
}, },
}; };
@@ -44,7 +45,16 @@ impl QueryHandle {
#[async_trait] #[async_trait]
pub trait DatabaseManagerSystem { pub trait DatabaseManagerSystem {
fn try_reserve_query(&self) -> QueryResult<QueryAdmission> {
Ok(QueryAdmission::unmanaged())
}
async fn execute(&self, query: &Query) -> QueryResult<QueryHandle>; async fn execute(&self, query: &Query) -> QueryResult<QueryHandle>;
async fn execute_admitted(&self, query: &Query, _admission: QueryAdmission) -> QueryResult<QueryHandle> {
self.execute(query).await
}
async fn build_query_state_machine(&self, query: Query) -> QueryResult<QueryStateMachineRef>; async fn build_query_state_machine(&self, query: Query) -> QueryResult<QueryStateMachineRef>;
async fn build_logical_plan(&self, query_state_machine: QueryStateMachineRef) -> QueryResult<Option<Plan>>; async fn build_logical_plan(&self, query_state_machine: QueryStateMachineRef) -> QueryResult<Option<Plan>>;
async fn execute_logical_plan( async fn execute_logical_plan(
+22 -6
View File
@@ -22,35 +22,51 @@ use rustfs_ecstore::api::error::{
}; };
#[cfg(test)] #[cfg(test)]
pub(crate) use rustfs_ecstore::api::object::PutObjReader as SelectPutObjReader; 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; 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::set_disk::DEFAULT_READ_BUFFER_SIZE as SELECT_DEFAULT_READ_BUFFER_SIZE;
pub(crate) use rustfs_ecstore::api::storage::ECStore as SelectStore; pub(crate) use rustfs_ecstore::api::storage::ECStore as SelectStore;
use rustfs_storage_api as storage_contracts; use rustfs_storage_api as storage_contracts;
#[cfg(test)]
static SELECT_TEST_OBJECT_STORE: std::sync::OnceLock<Arc<SelectStore>> = std::sync::OnceLock::new();
pub(crate) mod object_store { 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) mod crate_boundary {
pub(crate) use super::{ pub(crate) use super::{
SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError, PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions,
SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, 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, select_is_err_version_not_found,
}; };
} }
pub(crate) type SelectGetObjectReader = <SelectStore as storage_contracts::ObjectIO>::GetObjectReader; pub(crate) type SelectGetObjectReader = <SelectStore as storage_contracts::ObjectIO>::GetObjectReader;
pub(crate) type SelectObjectInfo = <SelectStore as storage_contracts::ObjectOperations>::ObjectInfo;
pub(crate) type SelectObjectOptions = <SelectStore as storage_contracts::ObjectOperations>::ObjectOptions; pub(crate) type SelectObjectOptions = <SelectStore as storage_contracts::ObjectOperations>::ObjectOptions;
#[cfg(test)] #[cfg(test)]
pub(crate) async fn select_test_ecstore_env() -> &'static rustfs_test_utils::TestECStoreEnv { pub(crate) async fn select_test_ecstore_env() -> &'static rustfs_test_utils::TestECStoreEnv {
static ENV: tokio::sync::OnceCell<rustfs_test_utils::TestECStoreEnv> = tokio::sync::OnceCell::const_new(); static ENV: tokio::sync::OnceCell<rustfs_test_utils::TestECStoreEnv> = tokio::sync::OnceCell::const_new();
ENV.get_or_init(|| async { rustfs_test_utils::TestECStoreEnv::builder().build().await }) let env = ENV
.await .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<Arc<SelectStore>> { pub(crate) fn resolve_select_object_store_handle() -> Option<Arc<SelectStore>> {
#[cfg(test)]
if let Some(store) = SELECT_TEST_OBJECT_STORE.get() {
return Some(Arc::clone(store));
}
resolve_select_object_store_handle_from_backend() resolve_select_object_store_handle_from_backend()
} }
+3
View File
@@ -54,5 +54,8 @@ s3s = { workspace = true, features = ["minio"] }
tokio = { workspace = true, features = ["fs", "rt-multi-thread", "sync", "time"] } tokio = { workspace = true, features = ["fs", "rt-multi-thread", "sync", "time"] }
tracing = { workspace = true } tracing = { workspace = true }
[dev-dependencies]
rustfs-test-utils = { workspace = true, features = ["put-object-commit-barrier"] }
[lib] [lib]
doctest = false doctest = false
+368 -26
View File
@@ -48,8 +48,8 @@ use rustfs_s3select_api::{
logical_planner::{LogicalPlanner, Plan}, logical_planner::{LogicalPlanner, Plan},
parser::Parser, parser::Parser,
session::{ session::{
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, SessionCtx, DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus,
SessionCtxFactory, QueryExecutionTracker, SessionCtx, SessionCtxFactory,
}, },
}, },
}; };
@@ -120,13 +120,20 @@ impl Drop for QueryPhaseGuard<'_> {
#[async_trait] #[async_trait]
impl QueryDispatcher for SimpleQueryDispatcher { impl QueryDispatcher for SimpleQueryDispatcher {
async fn execute_query(&self, query: &Query) -> QueryResult<Output> { async fn execute_query(&self, query: &Query) -> QueryResult<Output> {
let query_state_machine = self.build_query_state_machine(query.clone()).await?; self.execute_query_inner(query, None).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 fn try_reserve_query(&self) -> QueryResult<QueryAdmission> {
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<Output> {
self.execute_query_inner(query, Some(admission)).await
} }
async fn build_logical_plan(&self, query_state_machine: Arc<QueryStateMachine>) -> QueryResult<Option<Plan>> { async fn build_logical_plan(&self, query_state_machine: Arc<QueryStateMachine>) -> QueryResult<Option<Plan>> {
@@ -205,20 +212,64 @@ impl QueryDispatcher for SimpleQueryDispatcher {
} }
async fn build_query_state_machine(&self, query: Query) -> QueryResult<Arc<QueryStateMachine>> { async fn build_query_state_machine(&self, query: Query) -> QueryResult<Arc<QueryStateMachine>> {
let permit = self self.build_query_state_machine_inner(query, None).await
.query_admission }
.clone() }
.try_acquire_owned()
.map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?; impl SimpleQueryDispatcher {
async fn execute_query_inner(&self, query: &Query, admission: Option<QueryAdmission>) -> QueryResult<Output> {
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<QueryAdmission>,
) -> QueryResult<Arc<QueryStateMachine>> {
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( let query_tracker = QueryExecutionTracker::new(
&self.query_execution_owner, &self.query_execution_owner,
Arc::new(permit), query_guard,
Instant::now() + self.query_timeout, Instant::now() + self.query_timeout,
self.query_timeout.as_secs(), self.query_timeout.as_secs(),
); );
let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner); let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner);
let session = self let session = if let Some(snapshot) = query.snapshot().cloned() {
.run_with_query_deadline( 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, &query_tracker,
self.session_factory.create_session_ctx_with_tracker_and_memory_limit( self.session_factory.create_session_ctx_with_tracker_and_memory_limit(
query.context(), query.context(),
@@ -226,7 +277,8 @@ impl QueryDispatcher for SimpleQueryDispatcher {
self.memory_limit_bytes, self.memory_limit_bytes,
), ),
) )
.await?; .await?
};
if !query_tracker.mark_admitted(&self.query_execution_owner) { if !query_tracker.mark_admitted(&self.query_execution_owner) {
drop(session); drop(session);
return Err(self.query_tracker_error(&query_tracker)); return Err(self.query_tracker_error(&query_tracker));
@@ -234,9 +286,6 @@ impl QueryDispatcher for SimpleQueryDispatcher {
phase_guard.disarm(); phase_guard.disarm();
Ok(Arc::new(QueryStateMachine::begin_tracked(query, session, query_tracker)?)) Ok(Arc::new(QueryStateMachine::begin_tracked(query, session, query_tracker)?))
} }
}
impl SimpleQueryDispatcher {
async fn run_with_query_deadline<T>( async fn run_with_query_deadline<T>(
&self, &self,
query_tracker: &QueryExecutionTracker, query_tracker: &QueryExecutionTracker,
@@ -730,15 +779,17 @@ mod tests {
use async_trait::async_trait; use async_trait::async_trait;
use datafusion::{ use datafusion::{
arrow::{ arrow::{
datatypes::{Schema, SchemaRef}, array::{Int32Array, StringArray},
datatypes::{DataType, Field, Schema, SchemaRef},
record_batch::RecordBatch, record_batch::RecordBatch,
}, },
common::DataFusionError, common::DataFusionError,
execution::object_store::ObjectStoreUrl, execution::object_store::ObjectStoreUrl,
object_store::{ObjectStoreExt, path::Path}, object_store::{ObjectStoreExt, path::Path},
parquet::arrow::ArrowWriter,
physical_plan::{RecordBatchStream, stream::RecordBatchStreamAdapter}, physical_plan::{RecordBatchStream, stream::RecordBatchStreamAdapter},
}; };
use futures::{StreamExt, stream}; use futures::{StreamExt, TryStreamExt, stream};
use rustfs_s3select_api::{ use rustfs_s3select_api::{
QueryError, QueryResult, S3SelectPolicyError, QueryError, QueryResult, S3SelectPolicyError,
query::{ query::{
@@ -749,14 +800,15 @@ mod tests {
}, },
logical_planner::Plan, logical_planner::Plan,
session::{ session::{
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus,
SessionCtxFactory, QueryExecutionTracker, SessionCtxFactory,
}, },
}, },
}; };
use rustfs_test_utils::{PutObjectCommitBarrier, TestECStoreEnv};
use s3s::dto::{ use s3s::dto::{
CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, OutputSerialization, SelectObjectContentInput, CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType,
SelectObjectContentRequest, OutputSerialization, ParquetInput, SelectObjectContentInput, SelectObjectContentRequest,
}; };
use std::{ use std::{
pin::Pin, pin::Pin,
@@ -998,6 +1050,260 @@ mod tests {
(dispatcher, input) (dispatcher, input)
} }
async fn snapshot_test_env() -> &'static TestECStoreEnv {
static ENV: tokio::sync::OnceCell<TestECStoreEnv> = tokio::sync::OnceCell::const_new();
ENV.get_or_init(|| async { TestECStoreEnv::builder().prefix("s3select_query_snapshot").build().await })
.await
}
fn production_dispatcher(input: Arc<SelectObjectContentInput>) -> Arc<SimpleQueryDispatcher> {
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<String> {
let Output::StreamData(stream) = output else {
panic!("snapshot query should return rows");
};
stream
.try_collect::<Vec<_>>()
.await
.expect("collect snapshot query output")
.iter()
.flat_map(|batch| {
batch
.column(0)
.as_any()
.downcast_ref::<StringArray>()
.expect("snapshot marker column should be Utf8")
.iter()
.map(|value| value.expect("snapshot marker should not be null").to_string())
.collect::<Vec<_>>()
})
.collect()
}
async fn run_snapshot_generation_race(
input: Arc<SelectObjectContentInput>,
old_generation: Vec<u8>,
new_generation: Vec<u8>,
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<SelectObjectContentInput> {
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<SelectObjectContentInput> {
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<u8> {
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<u8> {
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( fn test_query_tracker(
permit: tokio::sync::OwnedSemaphorePermit, permit: tokio::sync::OwnedSemaphorePermit,
deadline: Instant, 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] #[tokio::test]
async fn staged_query_rejects_when_admission_is_saturated() { async fn staged_query_rejects_when_admission_is_saturated() {
let admission = Arc::new(Semaphore::new(1)); let admission = Arc::new(Semaphore::new(1));
+11 -1
View File
@@ -26,7 +26,7 @@ use rustfs_s3select_api::{
dispatcher::QueryDispatcher, dispatcher::QueryDispatcher,
execution::QueryStateMachineRef, execution::QueryStateMachineRef,
logical_planner::Plan, 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}, server::dbms::{DatabaseManagerSystem, QueryHandle},
}; };
@@ -64,12 +64,22 @@ impl<D> DatabaseManagerSystem for RustFSms<D>
where where
D: QueryDispatcher, D: QueryDispatcher,
{ {
fn try_reserve_query(&self) -> QueryResult<QueryAdmission> {
self.query_dispatcher.try_reserve_query()
}
async fn execute(&self, query: &Query) -> QueryResult<QueryHandle> { async fn execute(&self, query: &Query) -> QueryResult<QueryHandle> {
let result = self.query_dispatcher.execute_query(query).await?; let result = self.query_dispatcher.execute_query(query).await?;
Ok(QueryHandle::new(query.clone(), result)) Ok(QueryHandle::new(query.clone(), result))
} }
async fn execute_admitted(&self, query: &Query, admission: QueryAdmission) -> QueryResult<QueryHandle> {
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<QueryStateMachineRef> { async fn build_query_state_machine(&self, query: Query) -> QueryResult<QueryStateMachineRef> {
let query_state_machine = self.query_dispatcher.build_query_state_machine(query).await?; let query_state_machine = self.query_dispatcher.build_query_state_machine(query).await?;
+2 -1
View File
@@ -27,6 +27,7 @@ documentation = "https://docs.rs/rustfs-test-utils/latest/rustfs_test_utils/"
[features] [features]
default = [] default = []
put-object-commit-barrier = ["rustfs-ecstore/test-util"]
hotpath = [ hotpath = [
"hotpath/hotpath", "hotpath/hotpath",
"hotpath/tokio", "hotpath/tokio",
@@ -51,7 +52,7 @@ hotpath-cpu = [
[dependencies] [dependencies]
hotpath.workspace = true hotpath.workspace = true
rustfs-ecstore = { workspace = true } rustfs-ecstore.workspace = true
rustfs-storage-api = { workspace = true } rustfs-storage-api = { workspace = true }
tokio = { workspace = true, features = ["fs", "rt-multi-thread"] } tokio = { workspace = true, features = ["fs", "rt-multi-thread"] }
tokio-util = { workspace = true, features = ["io", "compat"] } tokio-util = { workspace = true, features = ["io", "compat"] }
+4 -1
View File
@@ -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::bucket::metadata_sys::init_bucket_metadata_sys;
pub(crate) use rustfs_ecstore::api::disk::endpoint::Endpoint; 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::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_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)] #[cfg(test)]
pub(crate) use rustfs_ecstore::api::config::com::{delete_config, read_config, save_config}; pub(crate) use rustfs_ecstore::api::config::com::{delete_config, read_config, save_config};
+49 -2
View File
@@ -30,14 +30,38 @@ mod ecstore_test_compat;
use std::path::PathBuf; use std::path::PathBuf;
use std::sync::{Arc, Once}; use std::sync::{Arc, Once};
#[cfg(feature = "put-object-commit-barrier")]
use ecstore_test_compat::fixture::ecstore_set_disk;
use ecstore_test_compat::fixture::{ use ecstore_test_compat::fixture::{
BucketOperations as _, BucketOptions, ECStore, Endpoint, EndpointServerPools, Endpoints, MakeBucketOptions, PoolEndpoints, BucketOperations as _, BucketOptions, ECStore, Endpoint, EndpointServerPools, Endpoints, MakeBucketOptions, ObjectIO as _,
init_bucket_metadata_sys, init_local_disks, PoolEndpoints, PutObjReader, SelectObjectSnapshot, init_bucket_metadata_sys, init_local_disks,
}; };
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
static INIT_TRACING: Once = Once::new(); 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 /// Install the standard test tracing subscriber once per process
/// (`RUST_LOG`-driven). Safe to call from every test; later calls are no-ops. /// (`RUST_LOG`-driven). Safe to call from every test; later calls are no-ops.
pub fn init_tracing() { pub fn init_tracing() {
@@ -90,6 +114,29 @@ impl TestECStoreEnv {
.await .await
.unwrap_or_else(|e| panic!("failed to create test bucket {bucket}: {e:?}")); .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<u8>) {
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<SelectObjectSnapshot> {
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 /// Builder for [`TestECStoreEnv`]. Defaults reproduce the historical heal
+613 -39
View File
@@ -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::options::get_opts;
use super::storage_api::select_object::request_context::spawn_traced; 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::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::app::runtime_sources::current_s3select_db;
use crate::error::ApiError; use crate::error::ApiError;
use bytes::Bytes; use bytes::Bytes;
@@ -14,20 +18,29 @@ use datafusion::arrow::{
use datafusion::common::DataFusionError; use datafusion::common::DataFusionError;
use datafusion::physical_plan::SendableRecordBatchStream; use datafusion::physical_plan::SendableRecordBatchStream;
use futures::StreamExt; use futures::StreamExt;
use http::{StatusCode, header::RANGE}; use http::{HeaderMap, HeaderName, HeaderValue, StatusCode, header::RANGE};
use rustfs_s3select_api::{ use rustfs_s3select_api::{
QueryError, S3SelectPolicyError, QueryError, S3SelectPolicyError,
object_store::{INVALID_SCAN_RANGE_MESSAGE, validate_scan_range_bounds}, object_store::{INVALID_SCAN_RANGE_MESSAGE, validate_scan_range_bounds},
query::{Context, Query}, query::{Context, Query},
}; };
use rustfs_s3select_query::instance::s3_select_query_timeout; 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::{ use s3s::dto::{
CSVOutput, CompressionType, ContinuationEvent, EndEvent, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, CSVOutput, CompressionType, ContinuationEvent, EndEvent, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput,
JSONOutput, JSONType, OutputSerialization, Progress, ProgressEvent, QuoteFields, RecordsEvent, SelectObjectContentEvent, JSONOutput, JSONType, OutputSerialization, Progress, ProgressEvent, QuoteFields, RecordsEvent, SelectObjectContentEvent,
SelectObjectContentEventStream, SelectObjectContentInput, SelectObjectContentOutput, SelectObjectContentRequest, Stats, SelectObjectContentEventStream, SelectObjectContentInput, SelectObjectContentOutput, SelectObjectContentRequest, Stats,
StatsEvent, 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 s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error};
use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use tokio::sync::mpsc; use tokio::sync::mpsc;
use tokio::time::{Instant, timeout_at}; 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 RECORDS_CHUNK_TARGET: usize = 128 * 1024;
const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure"; const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure";
const EMPTY_SELECT_EXPRESSION_MESSAGE: &str = "empty SQL expression"; 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)] #[derive(Clone, Debug)]
struct SelectValidation { struct SelectValidation {
@@ -45,10 +65,6 @@ struct SelectValidation {
progress_enabled: bool, progress_enabled: bool,
} }
struct SelectObjectMetadata {
size: u64,
}
#[derive(Clone, Debug)] #[derive(Clone, Debug)]
enum SelectOutputFormat { enum SelectOutputFormat {
Csv(CSVOutput), Csv(CSVOutput),
@@ -60,6 +76,21 @@ enum SelectProducerOutcome {
ReceiverClosed, ReceiverClosed,
} }
trait SelectSnapshotFence {
fn ensure_snapshot_valid(&self) -> S3Result<()>;
}
impl SelectSnapshotFence for Arc<StorageSelectObjectSnapshot> {
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( pub async fn execute_select_object_content(
req: S3Request<SelectObjectContentInput>, req: S3Request<SelectObjectContentInput>,
) -> S3Result<S3Response<SelectObjectContentOutput>> { ) -> S3Result<S3Response<SelectObjectContentOutput>> {
@@ -67,19 +98,27 @@ pub async fn execute_select_object_content(
let mut input = req.input; let mut input = req.input;
let validation = validate_select_request(&req.headers, &mut input)?; let validation = validate_select_request(&req.headers, &mut input)?;
log_select_request_summary(&input, &validation); 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_timeout = s3_select_query_timeout();
let query_deadline = Instant::now() + 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 .await
.map_err(|_| select_query_timeout_error(query_timeout.as_secs()))?
.map_err(map_query_error_to_s3)?; .map_err(map_query_error_to_s3)?;
let query = Query::new(Context { input: input.clone() }, input.request.expression.clone()); let admission = db.try_reserve_query().map_err(map_query_error_to_s3)?;
let output = db let snapshot = timeout_at(
.execute(&query) 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 .await
.map_err(|_| select_query_timeout_error(query_timeout.as_secs()))?
.map_err(map_query_error_to_s3)? .map_err(map_query_error_to_s3)?
.result() .result()
.into_record_batch_stream() .into_record_batch_stream()
@@ -90,24 +129,213 @@ pub async fn execute_select_object_content(
.clone() .clone()
.try_reserve_owned() .try_reserve_owned()
.map_err(|_| s3_error!(InternalError, "can't reserve Select terminal event capacity"))?; .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 { 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 { Ok(response)
payload: Some(SelectObjectContentEventStream::new(ReceiverStream::new(rx))),
}))
} }
async fn send_select_events_until_deadline( fn select_object_response(
rx: mpsc::Receiver<S3Result<SelectObjectContentEvent>>,
metadata: &HashMap<String, String>,
request_headers: &HeaderMap,
) -> S3Result<S3Response<SelectObjectContentOutput>> {
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<String, String>, name: &str) -> S3Result<Option<&'a str>> {
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<String, String>) -> S3Result<Option<&str>> {
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<String, String>) -> S3Result<Option<SelectSnapshotSseMode>> {
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<String, String>, request_headers: &HeaderMap) -> S3Result<HeaderMap> {
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<L: SelectSnapshotFence>(
output: SendableRecordBatchStream, output: SendableRecordBatchStream,
tx: mpsc::Sender<S3Result<SelectObjectContentEvent>>, tx: mpsc::Sender<S3Result<SelectObjectContentEvent>>,
terminal_permit: mpsc::OwnedPermit<S3Result<SelectObjectContentEvent>>, terminal_permit: mpsc::OwnedPermit<S3Result<SelectObjectContentEvent>>,
validation: SelectValidation, validation: SelectValidation,
deadline: Instant, deadline: Instant,
timeout_seconds: u64, 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, Ok(outcome) => outcome,
Err(_) => SelectProducerOutcome::Terminal(Err(map_query_error_to_s3( Err(_) => SelectProducerOutcome::Terminal(Err(map_query_error_to_s3(
S3SelectPolicyError::QueryTimeout { S3SelectPolicyError::QueryTimeout {
@@ -119,12 +347,14 @@ async fn send_select_events_until_deadline(
if let SelectProducerOutcome::Terminal(event) = outcome { if let SelectProducerOutcome::Terminal(event) = outcome {
terminal_permit.send(event); terminal_permit.send(event);
} }
drop(snapshot_lease);
} }
async fn send_select_events( async fn send_select_events(
mut output: SendableRecordBatchStream, mut output: SendableRecordBatchStream,
tx: &mpsc::Sender<S3Result<SelectObjectContentEvent>>, tx: &mpsc::Sender<S3Result<SelectObjectContentEvent>>,
validation: SelectValidation, validation: SelectValidation,
snapshot_fence: &impl SelectSnapshotFence,
) -> SelectProducerOutcome { ) -> SelectProducerOutcome {
let mut encoder = SelectOutputEncoder::new(validation.output_format); let mut encoder = SelectOutputEncoder::new(validation.output_format);
let mut progress = SelectProgress::default(); 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 { let stats = SelectObjectContentEvent::Stats(StatsEvent {
details: Some(progress.to_stats()), details: Some(progress.to_stats()),
}); });
if tx.send(Ok(stats)).await.is_err() { if tx.send(Ok(stats)).await.is_err() {
return SelectProducerOutcome::ReceiverClosed; return SelectProducerOutcome::ReceiverClosed;
} }
if let Err(error) = snapshot_fence.ensure_snapshot_valid() {
return SelectProducerOutcome::Terminal(Err(error));
}
SelectProducerOutcome::Terminal(Ok(SelectObjectContentEvent::End(EndEvent::default()))) SelectProducerOutcome::Terminal(Ok(SelectObjectContentEvent::End(EndEvent::default())))
} }
@@ -388,25 +624,35 @@ fn validate_input_delimiter_pair(field_delimiter: Option<&str>, record_delimiter
Ok(()) Ok(())
} }
async fn preflight_select_object( async fn prepare_select_object_snapshot(
headers: &http::HeaderMap, headers: &http::HeaderMap,
input: &SelectObjectContentInput, input: &SelectObjectContentInput,
read_principal: Option<&SseKmsPrincipal>, read_principal: Option<&SseKmsPrincipal>,
) -> S3Result<SelectObjectMetadata> { ) -> S3Result<StorageSelectObjectSnapshot> {
let opts = get_opts(&input.bucket, &input.key, None, None, headers) let opts = get_opts(&input.bucket, &input.key, None, None, headers)
.await .await
.map_err(ApiError::from)?; .map_err(ApiError::from)?;
let store = get_validated_store(&input.bucket).await?; let store = get_validated_store(&input.bucket).await?;
let info = store let snapshot = store
.get_object_info(&input.bucket, &input.key, &opts) .prepare_select_object_snapshot(&input.bucket, &input.key, headers, &opts)
.await .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_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())?; 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?; authorize_sse_kms_object_read(read_principal, &info.user_defined).await?;
Ok(SelectObjectMetadata { Ok(snapshot)
size: info.size.max(0) as u64, }
})
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) { 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 { 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() { if let Some(policy_error) = err.s3_select_policy_error() {
let message = policy_error.to_string(); let message = policy_error.to_string();
return match policy_error { 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 { fn looks_like_bucket_not_found(message: &str) -> bool {
message.contains("NoSuchBucket") || message.contains("bucket not found") || message.contains("BucketNotFound") message.contains("NoSuchBucket") || message.contains("bucket not found") || message.contains("BucketNotFound")
} }
@@ -692,9 +948,71 @@ mod tests {
physical_plan::stream::RecordBatchStreamAdapter, physical_plan::stream::RecordBatchStreamAdapter,
sql::sqlparser::parser::ParserError, sql::sqlparser::parser::ParserError,
}; };
use http::HeaderMap; use rustfs_test_utils::TestECStoreEnv;
use s3s::dto::{CSVInput, ParquetInput, ScanRange}; use s3s::dto::{CSVInput, ParquetInput, ScanRange};
struct LeaseDropSignal(Option<tokio::sync::oneshot::Sender<()>>);
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)] #[derive(Debug)]
struct CyclicError; 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( fn spawn_test_producer(
output: SendableRecordBatchStream, output: SendableRecordBatchStream,
channel_capacity: usize, channel_capacity: usize,
) -> (tokio::task::JoinHandle<()>, mpsc::Receiver<S3Result<SelectObjectContentEvent>>) { ) -> (
tokio::task::JoinHandle<()>,
mpsc::Receiver<S3Result<SelectObjectContentEvent>>,
tokio::sync::oneshot::Receiver<()>,
) {
let (tx, rx) = mpsc::channel(channel_capacity); let (tx, rx) = mpsc::channel(channel_capacity);
let terminal_permit = tx let terminal_permit = tx
.clone() .clone()
.try_reserve_owned() .try_reserve_owned()
.expect("test channel should reserve terminal capacity"); .expect("test channel should reserve terminal capacity");
let (lease, lease_released) = lease_drop_signal();
let producer = tokio::spawn(send_select_events_until_deadline( let producer = tokio::spawn(send_select_events_until_deadline(
output, output,
tx, tx,
@@ -760,8 +1232,9 @@ mod tests {
csv_validation(), csv_validation(),
Instant::now() + std::time::Duration::from_secs(1), Instant::now() + std::time::Duration::from_secs(1),
300, 300,
lease,
)); ));
(producer, rx) (producer, rx, lease_released)
} }
#[test] #[test]
@@ -844,6 +1317,28 @@ mod tests {
assert_eq!(invalid_object_size.code(), &S3ErrorCode::InternalError); 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::<StoragePrepareSelectObjectSnapshotError>().is_some())
);
}
#[test] #[test]
fn error_source_matching_stops_at_the_depth_bound() { fn error_source_matching_stops_at_the_depth_bound() {
let err = CyclicError; let err = CyclicError;
@@ -867,6 +1362,7 @@ mod tests {
tx.send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) tx.send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default())))
.await .await
.expect("test channel should accept the prefilled event"); .expect("test channel should accept the prefilled event");
let (lease, lease_released) = lease_drop_signal();
let producer = tokio::spawn(send_select_events_until_deadline( let producer = tokio::spawn(send_select_events_until_deadline(
output, output,
tx, tx,
@@ -874,6 +1370,7 @@ mod tests {
csv_validation(), csv_validation(),
Instant::now() + std::time::Duration::from_secs(1), Instant::now() + std::time::Duration::from_secs(1),
300, 300,
lease,
)); ));
tokio::task::yield_now().await; tokio::task::yield_now().await;
@@ -889,6 +1386,7 @@ mod tests {
.expect_err("terminal event should be an error"); .expect_err("terminal event should be an error");
assert_eq!(timeout_error.code(), &S3ErrorCode::Busy); assert_eq!(timeout_error.code(), &S3ErrorCode::Busy);
assert!(rx.recv().await.is_none()); assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "timeout should release the snapshot lease");
} }
#[tokio::test(start_paused = true)] #[tokio::test(start_paused = true)]
@@ -904,7 +1402,7 @@ mod tests {
}; };
let batches = [Ok(batch("a")), Ok(batch("b"))]; let batches = [Ok(batch("a")), Ok(batch("b"))];
let output = Box::pin(RecordBatchStreamAdapter::new(schema, futures::stream::iter(batches))); 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"); 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_eq!(stats.details.and_then(|details| details.bytes_returned), Some(4));
assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_))))); assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_)))));
assert!(rx.recv().await.is_none()); assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "End should release the snapshot lease");
} }
#[tokio::test(start_paused = true)] #[tokio::test(start_paused = true)]
@@ -932,7 +1431,7 @@ mod tests {
None::<(Result<RecordBatch, DataFusionError>, ())> None::<(Result<RecordBatch, DataFusionError>, ())>
}), }),
)); ));
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::task::yield_now().await;
tokio::time::advance(std::time::Duration::from_secs(1)).await; tokio::time::advance(std::time::Duration::from_secs(1)).await;
@@ -947,6 +1446,7 @@ mod tests {
assert!(matches!(stats, SelectObjectContentEvent::Stats(_))); assert!(matches!(stats, SelectObjectContentEvent::Stats(_)));
assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_))))); assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_)))));
assert!(rx.recv().await.is_none()); assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "EOF should release the snapshot lease");
} }
#[tokio::test(start_paused = true)] #[tokio::test(start_paused = true)]
@@ -958,7 +1458,7 @@ mod tests {
Err(DataFusionError::External(Box::new(S3SelectPolicyError::QueryConcurrencyLimit))) 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::task::yield_now().await;
tokio::time::advance(std::time::Duration::from_secs(1)).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"); .expect_err("terminal event should be an error");
assert_eq!(stream_error.code(), &S3ErrorCode::SlowDown); assert_eq!(stream_error.code(), &S3ErrorCode::SlowDown);
assert!(rx.recv().await.is_none()); assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "stream error should release the snapshot lease");
} }
#[tokio::test(start_paused = true)] #[tokio::test(start_paused = true)]
@@ -983,7 +1484,7 @@ mod tests {
schema, schema,
futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), 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::task::yield_now().await;
tokio::time::advance(std::time::Duration::from_secs(1)).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"); .expect_err("terminal event should be an error");
assert_eq!(encoder_error.code(), &S3ErrorCode::InternalError); assert_eq!(encoder_error.code(), &S3ErrorCode::InternalError);
assert!(rx.recv().await.is_none()); assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "encoder error should release the snapshot lease");
} }
#[tokio::test(start_paused = true)] #[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 (stream_dropped_tx, stream_dropped_rx) = tokio::sync::oneshot::channel::<()>();
let output = Box::pin(RecordBatchStreamAdapter::new( let output = Box::pin(RecordBatchStreamAdapter::new(
Arc::new(Schema::empty()), Arc::new(Schema::empty()),
@@ -1010,7 +1512,20 @@ mod tests {
}), }),
)); ));
let (tx, mut rx) = mpsc::channel(2); 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); tokio::pin!(producer);
assert!(futures::poll!(producer.as_mut()).is_pending()); assert!(futures::poll!(producer.as_mut()).is_pending());
@@ -1025,6 +1540,7 @@ mod tests {
stream_dropped_rx.await.is_err(), stream_dropped_rx.await.is_err(),
"query stream should be dropped when the receiver closes" "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)] #[tokio::test(start_paused = true)]
@@ -1040,7 +1556,8 @@ mod tests {
}), }),
)); ));
let (tx, mut rx) = mpsc::channel(2); 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); tokio::pin!(producer);
assert!(futures::poll!(producer.as_mut()).is_pending()); 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::<Result<RecordBatch, DataFusionError>>(),
));
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] #[test]
fn validate_defaults_csv_header_and_compression() { fn validate_defaults_csv_header_and_compression() {
let mut input = base_input(); let mut input = base_input();
+6 -7
View File
@@ -1153,14 +1153,13 @@ pub(crate) mod multipart_usecase {
} }
pub(crate) mod select_object { 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 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 { pub(crate) mod context {
+6 -3
View File
@@ -89,6 +89,8 @@ pub(crate) type StorageGetObjectReader = super::GetObjectReader;
pub(crate) type StorageObjectInfo = super::ObjectInfo; pub(crate) type StorageObjectInfo = super::ObjectInfo;
pub(crate) type StorageObjectLockDeleteOptions = contract::object::ObjectLockDeleteOptions; pub(crate) type StorageObjectLockDeleteOptions = contract::object::ObjectLockDeleteOptions;
pub(crate) type StorageObjectOptions = super::ObjectOptions; 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 StorageObjectToDelete = contract::object::ObjectToDelete;
pub(crate) type StoragePutObjReader = super::PutObjReader; pub(crate) type StoragePutObjReader = super::PutObjReader;
pub(crate) use super::ecfs_extend::{ 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::GetObjectBodySource;
pub(crate) use rustfs_ecstore::api::object::{ pub(crate) use rustfs_ecstore::api::object::{
EncryptionResolutionError, EncryptionResolutionErrorKind, GetObjectBodyCacheHook, GetObjectBodyCacheHookLookup, EncryptionResolutionError, EncryptionResolutionErrorKind, GetObjectBodyCacheHook, GetObjectBodyCacheHookLookup,
ObjectEncryptionResolver, ObjectMutationHook, ReadEncryptionMaterial, ReadEncryptionMode, ReadEncryptionRequest, ObjectEncryptionResolver, ObjectMutationHook, PrepareSelectObjectSnapshotError, ReadEncryptionMaterial,
get_object_body_cache_plaintext_len, lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, ReadEncryptionMode, ReadEncryptionRequest, SelectObjectSnapshot, get_object_body_cache_plaintext_len,
register_object_mutation_hook, unregister_get_object_body_cache_hook, unregister_object_mutation_hook, 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,
}; };
} }
@@ -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<String> {
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<u8> {
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<String>, current: &Option<String>, 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<String> {
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;
});
}