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",
"time",
"tokio",
"tokio-stream",
"tokio-util",
"tonic",
"tower",
@@ -10117,6 +10118,7 @@ dependencies = [
"tracing",
"transform-stream",
"url",
"uuid",
]
[[package]]
@@ -10132,6 +10134,7 @@ dependencies = [
"hotpath",
"parking_lot",
"rustfs-s3select-api",
"rustfs-test-utils",
"s3s",
"tokio",
"tracing",
+1
View File
@@ -181,6 +181,7 @@ path-absolutize = { workspace = true }
rmp.workspace = true
rmp-serde.workspace = true
tokio-util = { workspace = true, features = ["io", "compat"] }
tokio-stream = { workspace = true, features = ["sync"] }
base64 = { workspace = true }
hmac = { workspace = true }
sha1 = { workspace = true }
+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,
unregister_get_object_body_cache_hook, unregister_object_mutation_hook,
};
pub use crate::store::PreparedGetObjectReader;
pub use crate::store::{
PrepareSelectObjectSnapshotError, PreparedGetObjectReader, SelectObjectSnapshot, SelectObjectSnapshotReadError,
SnapshotConsistencyError,
};
}
pub mod rebalance {
@@ -449,6 +452,11 @@ pub mod rpc {
pub mod set_disk {
pub use crate::set_disk::{DEFAULT_READ_BUFFER_SIZE, SetDisks, get_lock_acquire_timeout, is_valid_storage_class};
#[cfg(feature = "test-util")]
pub mod test_util {
pub use crate::set_disk::{PutObjectCommitBarrier, PutObjectCommitPause};
}
}
pub mod store_list {
+98 -81
View File
@@ -1215,6 +1215,98 @@ async fn init_storage_disks_with_errors(
(disks, errs)
}
#[cfg(test)]
pub(crate) async fn make_local_two_set_sets() -> (Vec<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)]
mod tests {
use super::*;
@@ -1373,81 +1465,6 @@ mod tests {
assert_eq!(result, (Some(3), Some(1), Some(0)));
}
async fn two_set_test_sets() -> (Vec<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]
async fn heal_object_uses_explicit_set_scope() {
let (_temp_dirs, sets) = two_set_test_sets().await;
@@ -1497,7 +1514,7 @@ mod tests {
#[tokio::test]
async fn delete_prefix_surfaces_a_hard_error_from_any_set() {
let (_temp_dirs, sets) = two_set_test_sets().await;
let (_temp_dirs, sets) = make_local_two_set_sets().await;
let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await
@@ -1546,7 +1563,7 @@ mod tests {
#[tokio::test]
async fn delete_prefix_keeps_a_missing_bucket_idempotent_across_sets() {
let (_temp_dirs, sets) = two_set_test_sets().await;
let (_temp_dirs, sets) = make_local_two_set_sets().await;
let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await
@@ -1585,7 +1602,7 @@ mod tests {
#[tokio::test]
async fn delete_prefix_preserves_a_completely_missing_bucket_error() {
let (_temp_dirs, sets) = two_set_test_sets().await;
let (_temp_dirs, sets) = make_local_two_set_sets().await;
let bucket = format!("delete-prefix-missing-{}", Uuid::new_v4().simple());
let err = sets
@@ -1605,7 +1622,7 @@ mod tests {
#[tokio::test]
async fn delete_prefix_fails_when_one_set_is_entirely_offline() {
let (_temp_dirs, sets) = two_set_test_sets().await;
let (_temp_dirs, sets) = make_local_two_set_sets().await;
let bucket = format!("delete-prefix-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await
@@ -1652,7 +1669,7 @@ mod tests {
#[tokio::test]
async fn set_format_heal_accepts_quorum_from_a_nonzero_set() {
let (_temp_dirs, sets) = two_set_test_sets().await;
let (_temp_dirs, sets) = make_local_two_set_sets().await;
let (result, err) = sets.disk_set[1]
.heal_format(false)
@@ -1757,7 +1774,7 @@ mod tests {
#[serial]
async fn list_multipart_uploads_merges_all_sets_without_pagination_loss() {
let _setup_type_guard = SetupTypeGuard::switch_to(SetupType::Erasure).await;
let (_temp_dirs, sets) = two_set_test_sets().await;
let (_temp_dirs, sets) = make_local_two_set_sets().await;
let bucket = format!("multipart-list-{}", Uuid::new_v4().simple());
sets.make_bucket(&bucket, &MakeBucketOptions::default())
.await
+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.
pub fn lock_manager(&self) -> Arc<GlobalLockManager> {
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
/// [`CallCounterScope`] clears only its own object's counts on drop.
#[cfg(test)]
pub(in crate::set_disk) mod disk_call_counters {
pub(crate) mod disk_call_counters {
use std::collections::{HashMap, HashSet};
use std::sync::{Mutex, OnceLock};
+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();
mod core;
#[cfg(test)]
pub(crate) use core::io_primitives::disk_call_counters;
mod ctx;
mod metadata;
mod ops;
@@ -702,8 +704,8 @@ pub(crate) use ops::object::TransitionCleanupStoreBarrier as SetDiskTransitionCl
pub(crate) use ops::object::body_cache_plaintext_len;
#[cfg(test)]
pub(crate) use ops::object::cleanup_rejected_transition_upload_durably;
#[cfg(test)]
pub(crate) use ops::object::{PutObjectCommitBarrier, PutObjectCommitPause};
#[cfg(any(test, feature = "test-util"))]
pub use ops::object::{PutObjectCommitBarrier, PutObjectCommitPause};
mod read;
mod replication;
pub(crate) mod shard_source;
@@ -731,6 +733,10 @@ impl PreparedGetObjectMetadata {
.take()
.expect("prepared GET metadata ObjectInfo must be consumed exactly once")
}
pub(crate) fn read_semantics_identity(&self) -> [u8; 32] {
SetDisks::file_info_quorum_hash(&self.fi)
}
}
tokio::task_local! {
@@ -2553,6 +2559,53 @@ impl SetDisks {
))
}
#[cfg(any(test, feature = "test-util"))]
async fn acquire_write_lock_diag_with_pending_hook(
&self,
op: &'static str,
bucket: &str,
object: &str,
on_pending: impl FnOnce(),
) -> Result<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)]
fn log_object_lock_acquire_if_slow(
&self,
+55 -16
View File
@@ -1312,7 +1312,7 @@ impl SetDisks {
}
if !opts.no_lock && object_lock_guard.is_none() {
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeNamespace).await;
if let Some(expected_incarnation_id) = opts.expected_bucket_incarnation_id
&& opts.bucket_lifecycle_lock_fence.is_none()
@@ -1324,9 +1324,21 @@ impl SetDisks {
.await?,
);
}
object_lock_guard = Some(self.acquire_write_lock_diag("put_object_commit", bucket, object).await?);
#[cfg(any(test, feature = "test-util"))]
{
object_lock_guard = Some(
self.acquire_write_lock_diag_with_pending_hook("put_object_commit", bucket, object, || {
notify_put_object_commit_namespace_pending(bucket, object);
})
.await?,
);
}
#[cfg(not(any(test, feature = "test-util")))]
{
object_lock_guard = Some(self.acquire_write_lock_diag("put_object_commit", bucket, object).await?);
}
}
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
pause_put_object_commit(bucket, object, PutObjectCommitPause::AfterNamespace).await;
if deferred_data_movement_precondition && let Some(err) = self.check_write_precondition(bucket, object, opts).await {
@@ -2575,41 +2587,43 @@ fn remote_version_state_writer_enabled_for(requested: bool, fleet_confirmed: boo
requested && fleet_confirmed && fleet_proof_valid
}
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) enum PutObjectCommitPause {
pub enum PutObjectCommitPause {
BeforeNamespace,
AfterNamespace,
BeforeMetadata,
}
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
struct PutObjectCommitBarrierState {
bucket: String,
object: String,
pause: PutObjectCommitPause,
arrived: tokio::sync::Notify,
release: tokio::sync::Notify,
namespace_pending: tokio::sync::Notify,
}
#[cfg(test)]
pub(crate) struct PutObjectCommitBarrier {
#[cfg(any(test, feature = "test-util"))]
pub struct PutObjectCommitBarrier {
state: Arc<PutObjectCommitBarrierState>,
}
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
static PUT_OBJECT_COMMIT_BARRIER: std::sync::OnceLock<std::sync::Mutex<Vec<Arc<PutObjectCommitBarrierState>>>> =
std::sync::OnceLock::new();
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
impl PutObjectCommitBarrier {
pub(crate) fn install(bucket: &str, object: &str, pause: PutObjectCommitPause) -> Self {
pub fn install(bucket: &str, object: &str, pause: PutObjectCommitPause) -> Self {
let state = Arc::new(PutObjectCommitBarrierState {
bucket: bucket.to_string(),
object: object.to_string(),
pause,
arrived: tokio::sync::Notify::new(),
release: tokio::sync::Notify::new(),
namespace_pending: tokio::sync::Notify::new(),
});
let mut slot = PUT_OBJECT_COMMIT_BARRIER
.get_or_init(|| std::sync::Mutex::new(Vec::new()))
@@ -2626,18 +2640,27 @@ impl PutObjectCommitBarrier {
Self { state }
}
pub(crate) async fn wait_until_paused(&self) {
pub async fn wait_until_paused(&self) {
tokio::time::timeout(Duration::from_secs(30), self.state.arrived.notified())
.await
.expect("put object should reach the deterministic commit barrier");
}
pub(crate) fn release(&self) {
pub fn release(&self) {
self.state.release.notify_one();
}
pub async fn release_and_wait_until_namespace_pending(&self) {
assert_eq!(self.state.pause, PutObjectCommitPause::BeforeNamespace);
let namespace_pending = self.state.namespace_pending.notified();
self.release();
tokio::time::timeout(Duration::from_secs(5), namespace_pending)
.await
.expect("put object should wait for the namespace lock after leaving the commit barrier");
}
}
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
impl Drop for PutObjectCommitBarrier {
fn drop(&mut self) {
self.state.release.notify_one();
@@ -2649,7 +2672,7 @@ impl Drop for PutObjectCommitBarrier {
}
}
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
async fn pause_put_object_commit(bucket: &str, object: &str, pause: PutObjectCommitPause) {
let barrier = PUT_OBJECT_COMMIT_BARRIER
.get_or_init(|| std::sync::Mutex::new(Vec::new()))
@@ -2664,6 +2687,22 @@ async fn pause_put_object_commit(bucket: &str, object: &str, pause: PutObjectCom
}
}
#[cfg(any(test, feature = "test-util"))]
fn notify_put_object_commit_namespace_pending(bucket: &str, object: &str) {
let barrier = PUT_OBJECT_COMMIT_BARRIER
.get_or_init(|| std::sync::Mutex::new(Vec::new()))
.lock()
.expect("put object commit barrier mutex should not poison")
.iter()
.find(|barrier| {
barrier.bucket == bucket && barrier.object == object && barrier.pause == PutObjectCommitPause::BeforeNamespace
})
.cloned();
if let Some(barrier) = barrier {
barrier.namespace_pending.notify_one();
}
}
#[cfg(test)]
struct DeleteObjectCommitBarrierState {
bucket: String,
@@ -4414,7 +4453,7 @@ impl crate::storage_api_contracts::object::ObjectOperations for SetDisks {
self.invalidate_get_object_metadata_cache(bucket, object).await;
// Guard lock for metadata update
#[cfg(test)]
#[cfg(any(test, feature = "test-util"))]
pause_put_object_commit(bucket, object, PutObjectCommitPause::BeforeMetadata).await;
let _lock_guard = if !opts.no_lock {
Some(self.acquire_write_lock_diag("put_object_metadata", bucket, object).await?)
+4 -1
View File
@@ -152,7 +152,10 @@ mod list;
pub(crate) mod list_objects;
mod multipart;
mod object;
pub use object::PreparedGetObjectReader;
pub use object::{
PrepareSelectObjectSnapshotError, PreparedGetObjectReader, SelectObjectSnapshot, SelectObjectSnapshotReadError,
SnapshotConsistencyError,
};
mod peer;
mod rebalance;
pub(crate) mod utils;
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-util = { workspace = true, features = ["io", "compat"] }
tracing.workspace = true
uuid.workspace = true
transform-stream.workspace = true
url.workspace = true
[dev-dependencies]
rustfs-test-utils.workspace = true
rustfs-test-utils = { workspace = true, features = ["put-object-commit-barrier"] }
serial_test.workspace = true
[lib]
+26 -14
View File
@@ -15,22 +15,23 @@
#![recursion_limit = "256"]
use datafusion::{common::DataFusionError, sql::sqlparser::parser::ParserError};
use std::fmt::Display;
use std::{error::Error as StdError, fmt::Display};
use thiserror::Error;
pub mod object_store;
pub mod query;
pub mod server;
mod storage_api;
pub use storage_api::SelectObjectSnapshot;
#[cfg(test)]
mod test;
pub type QueryResult<T> = Result<T, QueryError>;
pub(crate) use storage_api::crate_boundary::{
SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError,
SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found,
select_is_err_version_not_found,
PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions,
SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError, resolve_select_object_store_handle,
select_is_err_bucket_not_found, select_is_err_object_not_found, select_is_err_version_not_found,
};
#[derive(Debug, Error)]
@@ -79,24 +80,24 @@ pub enum S3SelectPolicyError {
QueryTimeout { seconds: u64 },
}
impl S3SelectPolicyError {
fn from_error<'a>(mut err: &'a (dyn std::error::Error + 'static)) -> Option<&'a Self> {
impl QueryError {
fn source_error<T: StdError + 'static>(&self) -> Option<&T> {
let mut err: &(dyn StdError + 'static) = self;
for _ in 0..16 {
if let Some(policy_error) = err.downcast_ref::<Self>() {
return Some(policy_error);
if let Some(source) = err.downcast_ref::<T>() {
return Some(source);
}
err = err.source()?;
}
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> {
match self {
Self::Datafusion { source } => S3SelectPolicyError::from_error(source.as_ref()),
_ => None,
}
self.source_error()
}
}
@@ -230,6 +231,17 @@ mod tests {
));
}
#[test]
fn snapshot_consistency_error_is_recoverable_without_string_matching() {
let err = QueryError::Datafusion {
source: Box::new(DataFusionError::External(Box::new(SelectObjectSnapshotReadError::Consistency(
SnapshotConsistencyError::LockLost,
)))),
};
assert!(err.is_snapshot_consistency_error());
}
#[test]
fn test_query_error_from_parser_error() {
let parser_error = ParserError::ParserError("syntax error".to_string());
File diff suppressed because it is too large Load Diff
@@ -22,6 +22,7 @@ use super::{
Query,
execution::{Output, QueryStateMachine},
logical_planner::Plan,
session::QueryAdmission,
};
#[async_trait]
@@ -32,6 +33,14 @@ pub trait QueryDispatcher: Send + Sync {
async fn execute_query(&self, query: &Query) -> QueryResult<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 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 std::sync::Arc;
use crate::SelectObjectSnapshot;
pub mod analyzer;
pub mod ast;
pub mod dispatcher;
@@ -37,12 +39,26 @@ pub struct Context {
pub struct Query {
context: Context,
content: String,
snapshot: Option<Arc<SelectObjectSnapshot>>,
}
impl Query {
#[inline(always)]
pub fn new(context: Context, content: String) -> Self {
Self { context, content }
Self {
context,
content,
snapshot: None,
}
}
#[inline(always)]
pub fn new_with_snapshot(context: Context, content: String, snapshot: Arc<SelectObjectSnapshot>) -> Self {
Self {
context,
content,
snapshot: Some(snapshot),
}
}
pub fn context(&self) -> &Context {
@@ -52,4 +68,8 @@ impl Query {
pub fn content(&self) -> &str {
self.content.as_str()
}
pub fn snapshot(&self) -> Option<&Arc<SelectObjectSnapshot>> {
self.snapshot.as_ref()
}
}
+155 -14
View File
@@ -12,14 +12,16 @@
// See the License for the specific language governing permissions and
// limitations under the License.
use crate::SelectObjectSnapshot;
use crate::query::Context;
use crate::{QueryError, QueryResult, SelectStore, object_store::EcObjectStore};
use crate::{QueryError, QueryResult, object_store::EcObjectStore};
use datafusion::{
arrow::{
array::{Int32Array, StringArray},
datatypes::{DataType, Field, Schema},
record_batch::RecordBatch,
},
common::DataFusionError,
execution::{SessionStateBuilder, config::SessionConfig, context::SessionState, runtime_env::RuntimeEnvBuilder},
object_store::{ObjectStore, ObjectStoreExt, memory::InMemory, path::Path},
parquet::arrow::ArrowWriter,
@@ -39,6 +41,28 @@ use tracing::error;
pub type QueryExecutionGuard = Arc<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)]
pub struct QueryExecutionOwner {
identity: Arc<()>,
@@ -300,30 +324,30 @@ impl SessionCtxFactory {
query_tracker: QueryExecutionTracker,
memory_limit_bytes: usize,
) -> 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
}
#[cfg(test)]
async fn create_session_ctx_with_tracker_and_store(
pub async fn create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
&self,
context: &Context,
snapshot: Arc<SelectObjectSnapshot>,
query_tracker: QueryExecutionTracker,
store: Arc<SelectStore>,
memory_limit_bytes: usize,
) -> 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
}
async fn create_session_ctx_inner(
&self,
context: &Context,
snapshot: Option<Arc<SelectObjectSnapshot>>,
query_tracker: Option<QueryExecutionTracker>,
store: Option<Arc<SelectStore>>,
memory_limit_bytes: usize,
) -> QueryResult<SessionCtx> {
let df_session_ctx = self
.build_df_session_context(context, query_tracker.clone(), store, memory_limit_bytes)
.build_df_session_context(context, snapshot, query_tracker.clone(), memory_limit_bytes)
.await?;
Ok(SessionCtx {
@@ -336,8 +360,8 @@ impl SessionCtxFactory {
async fn build_df_session_context(
&self,
context: &Context,
snapshot: Option<Arc<SelectObjectSnapshot>>,
query_tracker: Option<QueryExecutionTracker>,
store: Option<Arc<SelectStore>>,
memory_limit_bytes: usize,
) -> QueryResult<SessionContext> {
let path = format!("s3://{}", context.input.bucket);
@@ -416,11 +440,13 @@ impl SessionCtxFactory {
} else {
let store: EcObjectStore = match query_tracker {
Some(query_tracker) => {
EcObjectStore::new_with_query_tracker(context.input.clone(), memory_pool, query_tracker, store)
EcObjectStore::new_with_query_tracker(context.input.clone(), memory_pool, query_tracker, snapshot)
}
None => EcObjectStore::new_with_memory_pool(context.input.clone(), memory_pool),
None => EcObjectStore::new_with_memory_pool(context.input.clone(), memory_pool, snapshot),
}
.map_err(|_| QueryError::NotImplemented { err: String::new() })?;
.map_err(|err| QueryError::Datafusion {
source: Box::new(DataFusionError::External(Box::new(err))),
})?;
df_session_state.with_object_store(&store_url, Arc::new(store)).build()
};
@@ -498,6 +524,7 @@ mod tests {
},
execution::memory_pool::MemoryLimit,
};
use http::HeaderMap;
use s3s::dto::{
CSVInput, CSVOutput, ExpressionType, InputSerialization, JSONInput, OutputSerialization, ParquetInput, ScanRange,
SelectObjectContentInput, SelectObjectContentRequest,
@@ -531,6 +558,16 @@ mod tests {
}
}
async fn prepare_test_snapshot(context: &Context) -> Arc<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]
fn session_factory_fields_remain_source_compatible() {
let factory = SessionCtxFactory {
@@ -679,10 +716,57 @@ mod tests {
));
}
#[tokio::test]
#[serial_test::serial]
async fn production_session_preserves_lazy_snapshot_entry() {
let _env = crate::storage_api::select_test_ecstore_env().await;
let session = SessionCtxFactory::new(false)
.create_session_ctx(&test_context())
.await
.expect("legacy production session should install a lazy object store");
assert_eq!(session.inner().config().target_partitions(), SessionConfig::new().target_partitions());
}
#[tokio::test]
#[serial_test::serial]
async fn legacy_tracked_production_session_preserves_lazy_snapshot_entry() {
let _env = crate::storage_api::select_test_ecstore_env().await;
let permit = Arc::new(tokio::sync::Semaphore::new(1))
.acquire_owned()
.await
.expect("query permit should be available");
let tracker = QueryExecutionTracker::new(
&QueryExecutionOwner::new(),
Arc::new(permit),
Instant::now() + std::time::Duration::from_secs(300),
300,
);
let session = SessionCtxFactory::new(false)
.create_session_ctx_with_tracker_and_memory_limit(
&test_context(),
tracker.clone(),
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
)
.await
.expect("legacy tracked session should install a lazy object store");
assert!(session.is_bound_to(&tracker));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial_test::serial]
async fn session_factory_propagates_query_guard_to_ec_store() {
let env = crate::storage_api::select_test_ecstore_env().await;
let mut context = test_context();
Arc::make_mut(&mut context.input).bucket = "s3select-query-guard-snapshot".to_string();
env.make_bucket(&context.input.bucket, false).await;
let mut reader = SelectPutObjReader::from_vec(b"id,name\n1,Alice\n".to_vec());
env.ecstore
.put_object(&context.input.bucket, &context.input.key, &mut reader, &Default::default())
.await
.expect("put query guard fixture");
let snapshot = prepare_test_snapshot(&context).await;
let admission = Arc::new(tokio::sync::Semaphore::new(1));
let permit = Arc::clone(&admission)
@@ -697,7 +781,12 @@ mod tests {
300,
);
let session = SessionCtxFactory::new(false)
.create_session_ctx_with_tracker_and_store(&test_context(), query_tracker, Arc::clone(&env.ecstore))
.create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
&context,
snapshot,
query_tracker,
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
)
.await
.expect("production session should be created with the query guard");
@@ -706,6 +795,51 @@ mod tests {
assert_eq!(Arc::strong_count(&query_guard), 1);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial_test::serial]
async fn session_factory_preserves_snapshot_binding_error_source() {
let env = crate::storage_api::select_test_ecstore_env().await;
let mut source_context = test_context();
Arc::make_mut(&mut source_context.input).bucket = "s3select-session-snapshot-identity".to_string();
Arc::make_mut(&mut source_context.input).key = "source.csv".to_string();
env.make_bucket(&source_context.input.bucket, false).await;
let mut reader = SelectPutObjReader::from_vec(b"source-marker\n".to_vec());
env.ecstore
.put_object(&source_context.input.bucket, &source_context.input.key, &mut reader, &Default::default())
.await
.expect("put snapshot identity fixture");
let snapshot = prepare_test_snapshot(&source_context).await;
let mut target_context = source_context.clone();
Arc::make_mut(&mut target_context.input).key = "different.csv".to_string();
let permit = Arc::new(tokio::sync::Semaphore::new(1))
.acquire_owned()
.await
.expect("query permit should be available");
let tracker = QueryExecutionTracker::new(
&QueryExecutionOwner::new(),
Arc::new(permit),
Instant::now() + std::time::Duration::from_secs(300),
300,
);
let error = match SessionCtxFactory::new(false)
.create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
&target_context,
snapshot,
tracker,
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
)
.await
{
Ok(_) => panic!("a session must reject a snapshot for a different object"),
Err(error) => error,
};
assert!(error.is_snapshot_consistency_error());
assert!(error.to_string().contains("snapshot consistency failure"));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial_test::serial]
async fn scan_range_is_preserved_across_large_csv_partition_boundary() {
@@ -721,6 +855,7 @@ mod tests {
assert!(data.len() > 1024 * 1024);
let mut context = test_context();
Arc::make_mut(&mut context.input).bucket = "s3select-scan-range-partition-snapshot".to_string();
let selected_start = i64::try_from(SELECTED_ROW * ROW_WIDTH).expect("selected row offset should fit in i64");
Arc::make_mut(&mut context.input).request.scan_range = Some(ScanRange {
start: Some(selected_start),
@@ -732,6 +867,7 @@ mod tests {
.put_object(&context.input.bucket, &context.input.key, &mut reader, &Default::default())
.await
.expect("put large ScanRange CSV fixture");
let snapshot = prepare_test_snapshot(&context).await;
let admission = Arc::new(tokio::sync::Semaphore::new(1));
let permit = Arc::clone(&admission)
@@ -746,7 +882,12 @@ mod tests {
);
let session = SessionCtxFactory::new(false)
.with_target_partitions(2)
.create_session_ctx_with_tracker_and_store(&context, query_tracker, Arc::clone(&env.ecstore))
.create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
&context,
snapshot,
query_tracker,
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
)
.await
.expect("create production ScanRange session");
assert!(!session.inner().config().options().optimizer.repartition_file_scans);
+10
View File
@@ -20,6 +20,7 @@ use crate::{
Query,
execution::{Output, QueryStateMachineRef},
logical_planner::Plan,
session::QueryAdmission,
},
};
@@ -44,7 +45,16 @@ impl QueryHandle {
#[async_trait]
pub trait DatabaseManagerSystem {
fn try_reserve_query(&self) -> QueryResult<QueryAdmission> {
Ok(QueryAdmission::unmanaged())
}
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_logical_plan(&self, query_state_machine: QueryStateMachineRef) -> QueryResult<Option<Plan>>;
async fn execute_logical_plan(
+22 -6
View File
@@ -22,35 +22,51 @@ use rustfs_ecstore::api::error::{
};
#[cfg(test)]
pub(crate) use rustfs_ecstore::api::object::PutObjReader as SelectPutObjReader;
pub use rustfs_ecstore::api::object::SelectObjectSnapshot;
pub(crate) use rustfs_ecstore::api::object::{
PrepareSelectObjectSnapshotError, SelectObjectSnapshotReadError, SnapshotConsistencyError,
};
use rustfs_ecstore::api::runtime::object_store_handle as resolve_select_object_store_handle_from_backend;
pub(crate) use rustfs_ecstore::api::set_disk::DEFAULT_READ_BUFFER_SIZE as SELECT_DEFAULT_READ_BUFFER_SIZE;
pub(crate) use rustfs_ecstore::api::storage::ECStore as SelectStore;
use rustfs_storage_api as storage_contracts;
#[cfg(test)]
static SELECT_TEST_OBJECT_STORE: std::sync::OnceLock<Arc<SelectStore>> = std::sync::OnceLock::new();
pub(crate) mod object_store {
pub(crate) use super::storage_contracts::{HTTPRangeSpec, ObjectIO, ObjectOperations};
pub(crate) use super::storage_contracts::HTTPRangeSpec;
#[cfg(test)]
pub(crate) use super::storage_contracts::ObjectIO;
}
pub(crate) mod crate_boundary {
pub(crate) use super::{
SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError,
SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found,
PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions,
SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError,
resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found,
select_is_err_version_not_found,
};
}
pub(crate) type SelectGetObjectReader = <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;
#[cfg(test)]
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();
ENV.get_or_init(|| async { rustfs_test_utils::TestECStoreEnv::builder().build().await })
.await
let env = ENV
.get_or_init(|| async { rustfs_test_utils::TestECStoreEnv::builder().build().await })
.await;
let _ = SELECT_TEST_OBJECT_STORE.set(Arc::clone(&env.ecstore));
env
}
pub(crate) fn resolve_select_object_store_handle() -> Option<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()
}
+3
View File
@@ -54,5 +54,8 @@ s3s = { workspace = true, features = ["minio"] }
tokio = { workspace = true, features = ["fs", "rt-multi-thread", "sync", "time"] }
tracing = { workspace = true }
[dev-dependencies]
rustfs-test-utils = { workspace = true, features = ["put-object-commit-barrier"] }
[lib]
doctest = false
+368 -26
View File
@@ -48,8 +48,8 @@ use rustfs_s3select_api::{
logical_planner::{LogicalPlanner, Plan},
parser::Parser,
session::{
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, SessionCtx,
SessionCtxFactory,
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus,
QueryExecutionTracker, SessionCtx, SessionCtxFactory,
},
},
};
@@ -120,13 +120,20 @@ impl Drop for QueryPhaseGuard<'_> {
#[async_trait]
impl QueryDispatcher for SimpleQueryDispatcher {
async fn execute_query(&self, query: &Query) -> QueryResult<Output> {
let query_state_machine = self.build_query_state_machine(query.clone()).await?;
let logical_plan = self.build_logical_plan(Arc::clone(&query_state_machine)).await?;
let Some(logical_plan) = logical_plan else {
return Ok(Output::Nil(()));
};
self.execute_query_inner(query, None).await
}
self.execute_logical_plan(logical_plan, query_state_machine).await
fn try_reserve_query(&self) -> QueryResult<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>> {
@@ -205,20 +212,64 @@ impl QueryDispatcher for SimpleQueryDispatcher {
}
async fn build_query_state_machine(&self, query: Query) -> QueryResult<Arc<QueryStateMachine>> {
let permit = self
.query_admission
.clone()
.try_acquire_owned()
.map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?;
self.build_query_state_machine_inner(query, None).await
}
}
impl SimpleQueryDispatcher {
async fn execute_query_inner(&self, query: &Query, admission: Option<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(
&self.query_execution_owner,
Arc::new(permit),
query_guard,
Instant::now() + self.query_timeout,
self.query_timeout.as_secs(),
);
let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner);
let session = self
.run_with_query_deadline(
let session = if let Some(snapshot) = query.snapshot().cloned() {
self.run_with_query_deadline(
&query_tracker,
self.session_factory
.create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
query.context(),
snapshot,
query_tracker.clone(),
self.memory_limit_bytes,
),
)
.await?
} else {
self.run_with_query_deadline(
&query_tracker,
self.session_factory.create_session_ctx_with_tracker_and_memory_limit(
query.context(),
@@ -226,7 +277,8 @@ impl QueryDispatcher for SimpleQueryDispatcher {
self.memory_limit_bytes,
),
)
.await?;
.await?
};
if !query_tracker.mark_admitted(&self.query_execution_owner) {
drop(session);
return Err(self.query_tracker_error(&query_tracker));
@@ -234,9 +286,6 @@ impl QueryDispatcher for SimpleQueryDispatcher {
phase_guard.disarm();
Ok(Arc::new(QueryStateMachine::begin_tracked(query, session, query_tracker)?))
}
}
impl SimpleQueryDispatcher {
async fn run_with_query_deadline<T>(
&self,
query_tracker: &QueryExecutionTracker,
@@ -730,15 +779,17 @@ mod tests {
use async_trait::async_trait;
use datafusion::{
arrow::{
datatypes::{Schema, SchemaRef},
array::{Int32Array, StringArray},
datatypes::{DataType, Field, Schema, SchemaRef},
record_batch::RecordBatch,
},
common::DataFusionError,
execution::object_store::ObjectStoreUrl,
object_store::{ObjectStoreExt, path::Path},
parquet::arrow::ArrowWriter,
physical_plan::{RecordBatchStream, stream::RecordBatchStreamAdapter},
};
use futures::{StreamExt, stream};
use futures::{StreamExt, TryStreamExt, stream};
use rustfs_s3select_api::{
QueryError, QueryResult, S3SelectPolicyError,
query::{
@@ -749,14 +800,15 @@ mod tests {
},
logical_planner::Plan,
session::{
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker,
SessionCtxFactory,
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus,
QueryExecutionTracker, SessionCtxFactory,
},
},
};
use rustfs_test_utils::{PutObjectCommitBarrier, TestECStoreEnv};
use s3s::dto::{
CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, OutputSerialization, SelectObjectContentInput,
SelectObjectContentRequest,
CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType,
OutputSerialization, ParquetInput, SelectObjectContentInput, SelectObjectContentRequest,
};
use std::{
pin::Pin,
@@ -998,6 +1050,260 @@ mod tests {
(dispatcher, input)
}
async fn snapshot_test_env() -> &'static TestECStoreEnv {
static ENV: tokio::sync::OnceCell<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(
permit: tokio::sync::OwnedSemaphorePermit,
deadline: Instant,
@@ -1118,6 +1424,42 @@ mod tests {
));
}
#[tokio::test]
async fn reserved_admission_is_handed_to_tracker_without_reacquiring() {
let admission = Arc::new(Semaphore::new(1));
let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300));
let reservation = dispatcher.try_reserve_query().expect("query reservation should succeed");
assert_eq!(admission.available_permits(), 0);
let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string());
let query_state_machine = dispatcher
.build_query_state_machine_inner(query, Some(reservation))
.await
.expect("reserved query should not acquire a second permit");
assert_eq!(admission.available_permits(), 0);
drop(query_state_machine);
assert_eq!(admission.available_permits(), 1);
}
#[tokio::test]
async fn reserved_admission_rejects_a_foreign_semaphore() {
let admission = Arc::new(Semaphore::new(1));
let foreign_admission = Arc::new(Semaphore::new(1));
let foreign_permit = Arc::clone(&foreign_admission)
.try_acquire_owned()
.expect("foreign permit should be available");
let reservation = QueryAdmission::new(Arc::new(foreign_permit));
let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300));
let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string());
let result = dispatcher.build_query_state_machine_inner(query, Some(reservation)).await;
assert!(matches!(result, Err(QueryError::Cancel)));
assert_eq!(admission.available_permits(), 1);
assert_eq!(foreign_admission.available_permits(), 1);
}
#[tokio::test]
async fn staged_query_rejects_when_admission_is_saturated() {
let admission = Arc::new(Semaphore::new(1));
+11 -1
View File
@@ -26,7 +26,7 @@ use rustfs_s3select_api::{
dispatcher::QueryDispatcher,
execution::QueryStateMachineRef,
logical_planner::Plan,
session::{DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES as DEFAULT_MEMORY_LIMIT_BYTES, SessionCtxFactory},
session::{DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES as DEFAULT_MEMORY_LIMIT_BYTES, QueryAdmission, SessionCtxFactory},
},
server::dbms::{DatabaseManagerSystem, QueryHandle},
};
@@ -64,12 +64,22 @@ impl<D> DatabaseManagerSystem for RustFSms<D>
where
D: QueryDispatcher,
{
fn try_reserve_query(&self) -> QueryResult<QueryAdmission> {
self.query_dispatcher.try_reserve_query()
}
async fn execute(&self, query: &Query) -> QueryResult<QueryHandle> {
let result = self.query_dispatcher.execute_query(query).await?;
Ok(QueryHandle::new(query.clone(), result))
}
async fn execute_admitted(&self, query: &Query, admission: QueryAdmission) -> QueryResult<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> {
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]
default = []
put-object-commit-barrier = ["rustfs-ecstore/test-util"]
hotpath = [
"hotpath/hotpath",
"hotpath/tokio",
@@ -51,7 +52,7 @@ hotpath-cpu = [
[dependencies]
hotpath.workspace = true
rustfs-ecstore = { workspace = true }
rustfs-ecstore.workspace = true
rustfs-storage-api = { workspace = true }
tokio = { workspace = true, features = ["fs", "rt-multi-thread"] }
tokio-util = { workspace = true, features = ["io", "compat"] }
+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::disk::endpoint::Endpoint;
pub(crate) use rustfs_ecstore::api::layout::{EndpointServerPools, Endpoints, PoolEndpoints};
pub(crate) use rustfs_ecstore::api::object::{PutObjReader, SelectObjectSnapshot};
#[cfg(feature = "put-object-commit-barrier")]
pub(crate) use rustfs_ecstore::api::set_disk as ecstore_set_disk;
pub(crate) use rustfs_ecstore::api::storage::{ECStore, init_local_disks};
pub(crate) use rustfs_storage_api::{BucketOperations, BucketOptions, MakeBucketOptions};
pub(crate) use rustfs_storage_api::{BucketOperations, BucketOptions, MakeBucketOptions, ObjectIO};
#[cfg(test)]
pub(crate) use rustfs_ecstore::api::config::com::{delete_config, read_config, save_config};
+49 -2
View File
@@ -30,14 +30,38 @@ mod ecstore_test_compat;
use std::path::PathBuf;
use std::sync::{Arc, Once};
#[cfg(feature = "put-object-commit-barrier")]
use ecstore_test_compat::fixture::ecstore_set_disk;
use ecstore_test_compat::fixture::{
BucketOperations as _, BucketOptions, ECStore, Endpoint, EndpointServerPools, Endpoints, MakeBucketOptions, PoolEndpoints,
init_bucket_metadata_sys, init_local_disks,
BucketOperations as _, BucketOptions, ECStore, Endpoint, EndpointServerPools, Endpoints, MakeBucketOptions, ObjectIO as _,
PoolEndpoints, PutObjReader, SelectObjectSnapshot, init_bucket_metadata_sys, init_local_disks,
};
use tokio_util::sync::CancellationToken;
static INIT_TRACING: Once = Once::new();
#[cfg(feature = "put-object-commit-barrier")]
pub struct PutObjectCommitBarrier(ecstore_set_disk::test_util::PutObjectCommitBarrier);
#[cfg(feature = "put-object-commit-barrier")]
impl PutObjectCommitBarrier {
pub fn before_namespace(bucket: &str, object: &str) -> Self {
Self(ecstore_set_disk::test_util::PutObjectCommitBarrier::install(
bucket,
object,
ecstore_set_disk::test_util::PutObjectCommitPause::BeforeNamespace,
))
}
pub async fn wait_until_paused(&self) {
self.0.wait_until_paused().await;
}
pub async fn release_and_wait_until_namespace_pending(&self) {
self.0.release_and_wait_until_namespace_pending().await;
}
}
/// Install the standard test tracing subscriber once per process
/// (`RUST_LOG`-driven). Safe to call from every test; later calls are no-ops.
pub fn init_tracing() {
@@ -90,6 +114,29 @@ impl TestECStoreEnv {
.await
.unwrap_or_else(|e| panic!("failed to create test bucket {bucket}: {e:?}"));
}
/// Write one complete object body through the real ECStore test backend.
pub async fn put_object_bytes(&self, bucket: &str, object: &str, bytes: Vec<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
+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::request_context::spawn_traced;
use super::storage_api::select_object::sse::{SseKmsPrincipal, authorize_sse_kms_object_read};
use super::storage_api::select_object::{get_validated_store, validate_sse_headers_for_read, validate_ssec_for_read};
use super::storage_api::select_object::{
StoragePrepareSelectObjectSnapshotError, StorageSelectObjectSnapshot, get_validated_store, validate_sse_headers_for_read,
validate_ssec_for_read,
};
use crate::app::runtime_sources::current_s3select_db;
use crate::error::ApiError;
use bytes::Bytes;
@@ -14,20 +18,29 @@ use datafusion::arrow::{
use datafusion::common::DataFusionError;
use datafusion::physical_plan::SendableRecordBatchStream;
use futures::StreamExt;
use http::{StatusCode, header::RANGE};
use http::{HeaderMap, HeaderName, HeaderValue, StatusCode, header::RANGE};
use rustfs_s3select_api::{
QueryError, S3SelectPolicyError,
object_store::{INVALID_SCAN_RANGE_MESSAGE, validate_scan_range_bounds},
query::{Context, Query},
};
use rustfs_s3select_query::instance::s3_select_query_timeout;
use rustfs_utils::http::headers::{
AMZ_ENCRYPTION_AES, AMZ_ENCRYPTION_KMS, AMZ_SERVER_SIDE_ENCRYPTION, AMZ_SERVER_SIDE_ENCRYPTION_KMS_CONTEXT,
AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID, SSEC_ALGORITHM_HEADER, SSEC_KEY_HEADER, SSEC_KEY_MD5_HEADER,
};
use s3s::dto::{
CSVOutput, CompressionType, ContinuationEvent, EndEvent, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput,
JSONOutput, JSONType, OutputSerialization, Progress, ProgressEvent, QuoteFields, RecordsEvent, SelectObjectContentEvent,
SelectObjectContentEventStream, SelectObjectContentInput, SelectObjectContentOutput, SelectObjectContentRequest, Stats,
StatsEvent,
};
use s3s::header::{
X_AMZ_SERVER_SIDE_ENCRYPTION, X_AMZ_SERVER_SIDE_ENCRYPTION_AWS_KMS_KEY_ID, X_AMZ_SERVER_SIDE_ENCRYPTION_CONTEXT,
X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5,
};
use s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio::time::{Instant, timeout_at};
@@ -38,6 +51,13 @@ const MAX_SELECT_EXPRESSION_BYTES: usize = 256 * 1024;
const RECORDS_CHUNK_TARGET: usize = 128 * 1024;
const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure";
const EMPTY_SELECT_EXPRESSION_MESSAGE: &str = "empty SQL expression";
const SELECT_MINIO_SSEC_SEALED_KEY: &str = "X-Minio-Internal-Server-Side-Encryption-Sealed-Key";
const SELECT_MINIO_S3_SEALED_KEY: &str = "X-Minio-Internal-Server-Side-Encryption-S3-Sealed-Key";
const SELECT_MINIO_KMS_SEALED_KEY: &str = "X-Minio-Internal-Server-Side-Encryption-Kms-Sealed-Key";
const SELECT_MINIO_KMS_KEY_ID: &str = "X-Minio-Internal-Server-Side-Encryption-S3-Kms-Key-Id";
const SELECT_MINIO_KMS_CONTEXT: &str = "X-Minio-Internal-Server-Side-Encryption-Context";
const SELECT_RUSTFS_KMS_KEY_ID: &str = "x-rustfs-encryption-key-id";
const SELECT_KMS_ARN_PREFIX: &str = "arn:aws:kms:";
#[derive(Clone, Debug)]
struct SelectValidation {
@@ -45,10 +65,6 @@ struct SelectValidation {
progress_enabled: bool,
}
struct SelectObjectMetadata {
size: u64,
}
#[derive(Clone, Debug)]
enum SelectOutputFormat {
Csv(CSVOutput),
@@ -60,6 +76,21 @@ enum SelectProducerOutcome {
ReceiverClosed,
}
trait SelectSnapshotFence {
fn ensure_snapshot_valid(&self) -> S3Result<()>;
}
impl SelectSnapshotFence for Arc<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(
req: S3Request<SelectObjectContentInput>,
) -> S3Result<S3Response<SelectObjectContentOutput>> {
@@ -67,19 +98,27 @@ pub async fn execute_select_object_content(
let mut input = req.input;
let validation = validate_select_request(&req.headers, &mut input)?;
log_select_request_summary(&input, &validation);
let metadata = preflight_select_object(&req.headers, &input, read_principal.as_ref()).await?;
validate_scan_range_for_object_size(&input.request, metadata.size)?;
let input = Arc::new(input);
let query_timeout = s3_select_query_timeout();
let query_deadline = Instant::now() + query_timeout;
let db = current_s3select_db((*input).clone(), false)
let input = Arc::new(input);
let db = timeout_at(query_deadline, current_s3select_db((*input).clone(), false))
.await
.map_err(|_| select_query_timeout_error(query_timeout.as_secs()))?
.map_err(map_query_error_to_s3)?;
let query = Query::new(Context { input: input.clone() }, input.request.expression.clone());
let output = db
.execute(&query)
let admission = db.try_reserve_query().map_err(map_query_error_to_s3)?;
let snapshot = timeout_at(
query_deadline,
prepare_select_object_snapshot(&req.headers, &input, read_principal.as_ref()),
)
.await
.map_err(|_| select_query_timeout_error(query_timeout.as_secs()))??;
validate_scan_range_for_object_size(&input.request, snapshot.logical_size())?;
let snapshot = Arc::new(snapshot);
let query =
Query::new_with_snapshot(Context { input: input.clone() }, input.request.expression.clone(), Arc::clone(&snapshot));
let output = timeout_at(query_deadline, db.execute_admitted(&query, admission))
.await
.map_err(|_| select_query_timeout_error(query_timeout.as_secs()))?
.map_err(map_query_error_to_s3)?
.result()
.into_record_batch_stream()
@@ -90,24 +129,213 @@ pub async fn execute_select_object_content(
.clone()
.try_reserve_owned()
.map_err(|_| s3_error!(InternalError, "can't reserve Select terminal event capacity"))?;
let response = select_object_response(rx, &snapshot.object_info().user_defined, &req.headers)?;
spawn_traced(async move {
send_select_events_until_deadline(output, tx, terminal_permit, validation, query_deadline, query_timeout.as_secs()).await;
send_select_events_until_deadline(
output,
tx,
terminal_permit,
validation,
query_deadline,
query_timeout.as_secs(),
snapshot,
)
.await;
});
Ok(S3Response::new(SelectObjectContentOutput {
payload: Some(SelectObjectContentEventStream::new(ReceiverStream::new(rx))),
}))
Ok(response)
}
async fn send_select_events_until_deadline(
fn select_object_response(
rx: mpsc::Receiver<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,
tx: mpsc::Sender<S3Result<SelectObjectContentEvent>>,
terminal_permit: mpsc::OwnedPermit<S3Result<SelectObjectContentEvent>>,
validation: SelectValidation,
deadline: Instant,
timeout_seconds: u64,
snapshot_lease: L,
) {
let outcome = match timeout_at(deadline, send_select_events(output, &tx, validation)).await {
let outcome = match timeout_at(deadline, send_select_events(output, &tx, validation, &snapshot_lease)).await {
Ok(outcome) => outcome,
Err(_) => SelectProducerOutcome::Terminal(Err(map_query_error_to_s3(
S3SelectPolicyError::QueryTimeout {
@@ -119,12 +347,14 @@ async fn send_select_events_until_deadline(
if let SelectProducerOutcome::Terminal(event) = outcome {
terminal_permit.send(event);
}
drop(snapshot_lease);
}
async fn send_select_events(
mut output: SendableRecordBatchStream,
tx: &mpsc::Sender<S3Result<SelectObjectContentEvent>>,
validation: SelectValidation,
snapshot_fence: &impl SelectSnapshotFence,
) -> SelectProducerOutcome {
let mut encoder = SelectOutputEncoder::new(validation.output_format);
let mut progress = SelectProgress::default();
@@ -180,12 +410,18 @@ async fn send_select_events(
}
}
if let Err(error) = snapshot_fence.ensure_snapshot_valid() {
return SelectProducerOutcome::Terminal(Err(error));
}
let stats = SelectObjectContentEvent::Stats(StatsEvent {
details: Some(progress.to_stats()),
});
if tx.send(Ok(stats)).await.is_err() {
return SelectProducerOutcome::ReceiverClosed;
}
if let Err(error) = snapshot_fence.ensure_snapshot_valid() {
return SelectProducerOutcome::Terminal(Err(error));
}
SelectProducerOutcome::Terminal(Ok(SelectObjectContentEvent::End(EndEvent::default())))
}
@@ -388,25 +624,35 @@ fn validate_input_delimiter_pair(field_delimiter: Option<&str>, record_delimiter
Ok(())
}
async fn preflight_select_object(
async fn prepare_select_object_snapshot(
headers: &http::HeaderMap,
input: &SelectObjectContentInput,
read_principal: Option<&SseKmsPrincipal>,
) -> S3Result<SelectObjectMetadata> {
) -> S3Result<StorageSelectObjectSnapshot> {
let opts = get_opts(&input.bucket, &input.key, None, None, headers)
.await
.map_err(ApiError::from)?;
let store = get_validated_store(&input.bucket).await?;
let info = store
.get_object_info(&input.bucket, &input.key, &opts)
let snapshot = store
.prepare_select_object_snapshot(&input.bucket, &input.key, headers, &opts)
.await
.map_err(ApiError::from)?;
.map_err(map_prepare_snapshot_error)?;
let info = snapshot.object_info();
validate_sse_headers_for_read(&info.user_defined, headers)?;
validate_ssec_for_read(&info.user_defined, input.sse_customer_key.as_ref(), input.sse_customer_key_md5.as_ref())?;
authorize_sse_kms_object_read(read_principal, &info.user_defined).await?;
Ok(SelectObjectMetadata {
size: info.size.max(0) as u64,
})
Ok(snapshot)
}
fn map_prepare_snapshot_error(err: StoragePrepareSelectObjectSnapshotError) -> S3Error {
match err {
StoragePrepareSelectObjectSnapshotError::Storage(err) => ApiError::from(err).into(),
err => {
let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, err.to_string());
s3_error.set_source(Box::new(err));
s3_error
}
}
}
fn log_select_request_summary(input: &SelectObjectContentInput, validation: &SelectValidation) {
@@ -567,6 +813,12 @@ fn clamp_i64(value: u64) -> i64 {
}
fn map_query_error_to_s3(err: QueryError) -> S3Error {
if err.is_snapshot_consistency_error() {
let message = err.to_string();
let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, message);
s3_error.set_source(Box::new(err));
return s3_error;
}
if let Some(policy_error) = err.s3_select_policy_error() {
let message = policy_error.to_string();
return match policy_error {
@@ -616,6 +868,10 @@ fn map_query_error_to_s3(err: QueryError) -> S3Error {
}
}
fn select_query_timeout_error(seconds: u64) -> S3Error {
map_query_error_to_s3(S3SelectPolicyError::QueryTimeout { seconds }.into())
}
fn looks_like_bucket_not_found(message: &str) -> bool {
message.contains("NoSuchBucket") || message.contains("bucket not found") || message.contains("BucketNotFound")
}
@@ -692,9 +948,71 @@ mod tests {
physical_plan::stream::RecordBatchStreamAdapter,
sql::sqlparser::parser::ParserError,
};
use http::HeaderMap;
use rustfs_test_utils::TestECStoreEnv;
use s3s::dto::{CSVInput, ParquetInput, ScanRange};
struct LeaseDropSignal(Option<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)]
struct CyclicError;
@@ -744,15 +1062,169 @@ mod tests {
}
}
#[test]
fn select_snapshot_sse_s3_headers_are_whitelisted() {
let metadata = HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_AES.to_string()),
(SELECT_RUSTFS_KMS_KEY_ID.to_string(), "default".to_string()),
(SELECT_MINIO_KMS_KEY_ID.to_string(), "default".to_string()),
("x-amz-meta-private".to_string(), "private-value".to_string()),
]);
let headers = select_snapshot_sse_response_headers(&metadata, &HeaderMap::new())
.expect("valid SSE-S3 snapshot metadata should project response headers");
assert_eq!(headers.len(), 1);
assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION).expect("SSE-S3 mode"), "AES256");
assert!(headers.get("x-amz-meta-private").is_none());
}
#[test]
fn select_snapshot_sse_kms_headers_use_snapshot_metadata() {
let context = "eyJ0ZW5hbnQiOiJvbmUifQ==";
for key_id in ["key-1", "arn:aws:kms:key-2"] {
let metadata = HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()),
(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), key_id.to_string()),
(SELECT_RUSTFS_KMS_KEY_ID.to_string(), key_id.to_string()),
(SELECT_MINIO_KMS_KEY_ID.to_string(), key_id.to_string()),
(SELECT_MINIO_KMS_CONTEXT.to_string(), context.to_string()),
]);
let headers = select_snapshot_sse_response_headers(&metadata, &HeaderMap::new())
.expect("valid SSE-KMS snapshot metadata should project response headers");
let expected_key_id = if key_id.starts_with(SELECT_KMS_ARN_PREFIX) {
key_id.to_string()
} else {
format!("{SELECT_KMS_ARN_PREFIX}{key_id}")
};
assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION).expect("SSE-KMS mode"), "aws:kms");
assert_eq!(
headers
.get(X_AMZ_SERVER_SIDE_ENCRYPTION_AWS_KMS_KEY_ID)
.expect("SSE-KMS key ID")
.to_str()
.expect("SSE-KMS key ID should be valid text"),
expected_key_id
);
assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CONTEXT).expect("SSE-KMS context"), context);
}
}
#[test]
fn select_snapshot_sse_c_headers_never_echo_the_customer_key() {
let key_md5 = "customer-key-md5";
let metadata = HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_AES.to_string()),
(SSEC_ALGORITHM_HEADER.to_string(), AMZ_ENCRYPTION_AES.to_string()),
(SSEC_KEY_MD5_HEADER.to_string(), key_md5.to_string()),
("x-amz-meta-private".to_string(), "private-value".to_string()),
]);
let mut request_headers = HeaderMap::new();
request_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, HeaderValue::from_static("AES256"));
request_headers.insert(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, HeaderValue::from_static(key_md5));
request_headers.insert(
http::HeaderName::from_static(SSEC_KEY_HEADER),
HeaderValue::from_static("must-not-be-returned"),
);
let headers = select_snapshot_sse_response_headers(&metadata, &request_headers)
.expect("validated SSE-C request values should project response headers");
assert_eq!(headers.len(), 2);
assert_eq!(
headers
.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM)
.expect("SSE-C algorithm"),
"AES256"
);
assert_eq!(
headers
.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5)
.expect("SSE-C key MD5"),
key_md5
);
assert!(headers.get(SSEC_KEY_HEADER).is_none());
assert!(headers.get("x-amz-meta-private").is_none());
}
#[test]
fn select_snapshot_sse_headers_fail_closed_on_corrupt_metadata() {
let invalid_context = "not-base64";
let persisted_key = "must-not-leak";
let corrupt_metadata = [
HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_ascii_lowercase(), "AES256".to_string()),
(AMZ_SERVER_SIDE_ENCRYPTION.to_ascii_uppercase(), "aws:kms".to_string()),
]),
HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()),
(SSEC_ALGORITHM_HEADER.to_string(), "AES256".to_string()),
]),
HashMap::from([(SELECT_MINIO_KMS_SEALED_KEY.to_string(), "sealed".to_string())]),
HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()),
(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), "key-1".to_string()),
(SELECT_RUSTFS_KMS_KEY_ID.to_string(), "key-2".to_string()),
]),
HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), "aws:kms".to_string()),
(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), "key-1".to_string()),
(SELECT_MINIO_KMS_CONTEXT.to_string(), invalid_context.to_string()),
]),
HashMap::from([
(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_KMS.to_string()),
(AMZ_SERVER_SIDE_ENCRYPTION_KMS_ID.to_string(), "key-1".to_string()),
(AMZ_SERVER_SIDE_ENCRYPTION_KMS_CONTEXT.to_string(), "persisted-context".to_string()),
]),
HashMap::from([
(SSEC_ALGORITHM_HEADER.to_string(), "AES256".to_string()),
(SSEC_KEY_MD5_HEADER.to_string(), "customer-key-md5".to_string()),
(SSEC_KEY_HEADER.to_string(), persisted_key.to_string()),
]),
];
for metadata in corrupt_metadata {
let error = select_snapshot_sse_response_headers(&metadata, &HeaderMap::new())
.expect_err("corrupt snapshot encryption metadata must fail closed");
assert_eq!(error.code(), &S3ErrorCode::InternalError);
assert!(!error.to_string().contains(invalid_context));
assert!(!error.to_string().contains(persisted_key));
}
}
#[test]
fn select_response_projects_snapshot_sse_headers() {
let (_tx, rx) = mpsc::channel(1);
let metadata = HashMap::from([(AMZ_SERVER_SIDE_ENCRYPTION.to_string(), AMZ_ENCRYPTION_AES.to_string())]);
let response = select_object_response(rx, &metadata, &HeaderMap::new())
.expect("valid snapshot metadata should produce a Select response");
assert_eq!(
response
.headers
.get(X_AMZ_SERVER_SIDE_ENCRYPTION)
.expect("snapshot SSE response header"),
AMZ_ENCRYPTION_AES
);
}
fn spawn_test_producer(
output: SendableRecordBatchStream,
channel_capacity: usize,
) -> (tokio::task::JoinHandle<()>, mpsc::Receiver<S3Result<SelectObjectContentEvent>>) {
) -> (
tokio::task::JoinHandle<()>,
mpsc::Receiver<S3Result<SelectObjectContentEvent>>,
tokio::sync::oneshot::Receiver<()>,
) {
let (tx, rx) = mpsc::channel(channel_capacity);
let terminal_permit = tx
.clone()
.try_reserve_owned()
.expect("test channel should reserve terminal capacity");
let (lease, lease_released) = lease_drop_signal();
let producer = tokio::spawn(send_select_events_until_deadline(
output,
tx,
@@ -760,8 +1232,9 @@ mod tests {
csv_validation(),
Instant::now() + std::time::Duration::from_secs(1),
300,
lease,
));
(producer, rx)
(producer, rx, lease_released)
}
#[test]
@@ -844,6 +1317,28 @@ mod tests {
assert_eq!(invalid_object_size.code(), &S3ErrorCode::InternalError);
}
#[test]
fn prepare_snapshot_storage_error_preserves_existing_s3_mapping() {
let err = map_prepare_snapshot_error(StoragePrepareSelectObjectSnapshotError::Storage(StorageError::ObjectNotFound(
"bucket".to_string(),
"object.csv".to_string(),
)));
assert_eq!(err.code(), &S3ErrorCode::NoSuchKey);
assert!(err.source().is_some());
}
#[test]
fn prepare_snapshot_invalid_logical_size_fails_with_internal_error_and_source() {
let err = map_prepare_snapshot_error(StoragePrepareSelectObjectSnapshotError::InvalidLogicalSize { size: -1 });
assert_eq!(err.code(), &S3ErrorCode::InternalError);
assert!(
err.source()
.is_some_and(|source| source.downcast_ref::<StoragePrepareSelectObjectSnapshotError>().is_some())
);
}
#[test]
fn error_source_matching_stops_at_the_depth_bound() {
let err = CyclicError;
@@ -867,6 +1362,7 @@ mod tests {
tx.send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default())))
.await
.expect("test channel should accept the prefilled event");
let (lease, lease_released) = lease_drop_signal();
let producer = tokio::spawn(send_select_events_until_deadline(
output,
tx,
@@ -874,6 +1370,7 @@ mod tests {
csv_validation(),
Instant::now() + std::time::Duration::from_secs(1),
300,
lease,
));
tokio::task::yield_now().await;
@@ -889,6 +1386,7 @@ mod tests {
.expect_err("terminal event should be an error");
assert_eq!(timeout_error.code(), &S3ErrorCode::Busy);
assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "timeout should release the snapshot lease");
}
#[tokio::test(start_paused = true)]
@@ -904,7 +1402,7 @@ mod tests {
};
let batches = [Ok(batch("a")), Ok(batch("b"))];
let output = Box::pin(RecordBatchStreamAdapter::new(schema, futures::stream::iter(batches)));
let (producer, mut rx) = spawn_test_producer(output, 8);
let (producer, mut rx, lease_released) = spawn_test_producer(output, 8);
producer.await.expect("producer should finish at query EOF");
@@ -921,6 +1419,7 @@ mod tests {
assert_eq!(stats.details.and_then(|details| details.bytes_returned), Some(4));
assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_)))));
assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "End should release the snapshot lease");
}
#[tokio::test(start_paused = true)]
@@ -932,7 +1431,7 @@ mod tests {
None::<(Result<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::time::advance(std::time::Duration::from_secs(1)).await;
@@ -947,6 +1446,7 @@ mod tests {
assert!(matches!(stats, SelectObjectContentEvent::Stats(_)));
assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_)))));
assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "EOF should release the snapshot lease");
}
#[tokio::test(start_paused = true)]
@@ -958,7 +1458,7 @@ mod tests {
Err(DataFusionError::External(Box::new(S3SelectPolicyError::QueryConcurrencyLimit)))
}),
));
let (producer, mut rx) = spawn_test_producer(output, 2);
let (producer, mut rx, lease_released) = spawn_test_producer(output, 2);
tokio::task::yield_now().await;
tokio::time::advance(std::time::Duration::from_secs(1)).await;
@@ -972,6 +1472,7 @@ mod tests {
.expect_err("terminal event should be an error");
assert_eq!(stream_error.code(), &S3ErrorCode::SlowDown);
assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "stream error should release the snapshot lease");
}
#[tokio::test(start_paused = true)]
@@ -983,7 +1484,7 @@ mod tests {
schema,
futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }),
));
let (producer, mut rx) = spawn_test_producer(output, 2);
let (producer, mut rx, lease_released) = spawn_test_producer(output, 2);
tokio::task::yield_now().await;
tokio::time::advance(std::time::Duration::from_secs(1)).await;
@@ -997,10 +1498,11 @@ mod tests {
.expect_err("terminal event should be an error");
assert_eq!(encoder_error.code(), &S3ErrorCode::InternalError);
assert!(rx.recv().await.is_none());
assert!(lease_released.await.is_ok(), "encoder error should release the snapshot lease");
}
#[tokio::test(start_paused = true)]
async fn producer_drops_query_stream_when_receiver_closes() {
async fn producer_drops_query_stream_and_snapshot_lease_when_receiver_closes() {
let (stream_dropped_tx, stream_dropped_rx) = tokio::sync::oneshot::channel::<()>();
let output = Box::pin(RecordBatchStreamAdapter::new(
Arc::new(Schema::empty()),
@@ -1010,7 +1512,20 @@ mod tests {
}),
));
let (tx, mut rx) = mpsc::channel(2);
let producer = send_select_events(output, &tx, csv_validation());
let terminal_permit = tx
.clone()
.try_reserve_owned()
.expect("test channel should reserve terminal capacity");
let (lease, lease_released) = lease_drop_signal();
let producer = send_select_events_until_deadline(
output,
tx,
terminal_permit,
csv_validation(),
Instant::now() + std::time::Duration::from_secs(1),
300,
lease,
);
tokio::pin!(producer);
assert!(futures::poll!(producer.as_mut()).is_pending());
@@ -1025,6 +1540,7 @@ mod tests {
stream_dropped_rx.await.is_err(),
"query stream should be dropped when the receiver closes"
);
assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease");
}
#[tokio::test(start_paused = true)]
@@ -1040,7 +1556,8 @@ mod tests {
}),
));
let (tx, mut rx) = mpsc::channel(2);
let producer = send_select_events(output, &tx, csv_validation());
let snapshot_fence = LeaseDropSignal(None);
let producer = send_select_events(output, &tx, csv_validation(), &snapshot_fence);
tokio::pin!(producer);
assert!(futures::poll!(producer.as_mut()).is_pending());
@@ -1058,6 +1575,63 @@ mod tests {
);
}
#[tokio::test]
async fn producer_rejects_successful_end_when_final_snapshot_fence_fails() {
let schema = Arc::new(Schema::new(vec![Field::new(
"value",
datafusion::arrow::datatypes::DataType::Utf8,
false,
)]));
let batch = RecordBatch::try_new(
Arc::clone(&schema),
vec![Arc::new(datafusion::arrow::array::StringArray::from(vec!["old-generation"]))],
)
.expect("test record batch should be valid");
let output = Box::pin(RecordBatchStreamAdapter::new(
schema,
futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }),
));
let (tx, mut rx) = mpsc::channel(4);
let outcome = send_select_events(output, &tx, csv_validation(), &FailingSnapshotFence).await;
let SelectProducerOutcome::Terminal(Err(error)) = outcome else {
panic!("failed final snapshot fence must produce a terminal error");
};
assert_eq!(error.code(), &S3ErrorCode::InternalError);
assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_)))));
assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Records(_)))));
assert!(rx.try_recv().is_err(), "failed final fence must not enqueue Stats or End");
}
#[tokio::test]
async fn producer_rechecks_snapshot_after_stats_backpressure() {
let output = Box::pin(RecordBatchStreamAdapter::new(
Arc::new(Schema::empty()),
futures::stream::empty::<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]
fn validate_defaults_csv_header_and_compression() {
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 contract {
pub(crate) mod object {
pub(crate) use super::super::super::storage_contracts::ObjectOperations;
}
}
pub(crate) use super::{options, request_context, sse};
pub(crate) use crate::storage::storage_api::{get_validated_store, validate_sse_headers_for_read, validate_ssec_for_read};
#[cfg(test)]
pub(crate) use crate::storage::storage_api::StorageError;
pub(crate) use crate::storage::storage_api::{
StoragePrepareSelectObjectSnapshotError, StorageSelectObjectSnapshot, get_validated_store, validate_sse_headers_for_read,
validate_ssec_for_read,
};
}
pub(crate) mod context {
+6 -3
View File
@@ -89,6 +89,8 @@ pub(crate) type StorageGetObjectReader = super::GetObjectReader;
pub(crate) type StorageObjectInfo = super::ObjectInfo;
pub(crate) type StorageObjectLockDeleteOptions = contract::object::ObjectLockDeleteOptions;
pub(crate) type StorageObjectOptions = super::ObjectOptions;
pub(crate) type StoragePrepareSelectObjectSnapshotError = ecstore_object::PrepareSelectObjectSnapshotError;
pub(crate) type StorageSelectObjectSnapshot = ecstore_object::SelectObjectSnapshot;
pub(crate) type StorageObjectToDelete = contract::object::ObjectToDelete;
pub(crate) type StoragePutObjReader = super::PutObjReader;
pub(crate) use super::ecfs_extend::{
@@ -523,9 +525,10 @@ pub(crate) mod ecstore_object {
pub(crate) use rustfs_ecstore::api::object::GetObjectBodySource;
pub(crate) use rustfs_ecstore::api::object::{
EncryptionResolutionError, EncryptionResolutionErrorKind, GetObjectBodyCacheHook, GetObjectBodyCacheHookLookup,
ObjectEncryptionResolver, ObjectMutationHook, ReadEncryptionMaterial, ReadEncryptionMode, ReadEncryptionRequest,
get_object_body_cache_plaintext_len, lookup_get_object_body_cache_hook, register_get_object_body_cache_hook,
register_object_mutation_hook, unregister_get_object_body_cache_hook, unregister_object_mutation_hook,
ObjectEncryptionResolver, ObjectMutationHook, PrepareSelectObjectSnapshotError, ReadEncryptionMaterial,
ReadEncryptionMode, ReadEncryptionRequest, SelectObjectSnapshot, get_object_body_cache_plaintext_len,
lookup_get_object_body_cache_hook, register_get_object_body_cache_hook, register_object_mutation_hook,
unregister_get_object_body_cache_hook, unregister_object_mutation_hook,
};
}
@@ -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;
});
}