mirror of
https://github.com/rustfs/rustfs.git
synced 2026-08-16 01:48:21 +00:00
fix(select): pin object snapshot for query lifetime (#5835)
This commit is contained in:
+613
-39
@@ -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();
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;
|
||||
});
|
||||
}
|
||||
Reference in New Issue
Block a user