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
+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;
});
}