fix(select): enforce typed S3 Select error semantics (#5942)

* fix(select): enforce typed S3 Select error semantics

* fix(select): classify function argument planner errors

---------

Co-authored-by: overtrue <anzhengchao@gmail.com>
This commit is contained in:
GatewayJ
2026-08-11 23:59:38 +08:00
committed by GitHub
parent a206a0779e
commit e9728192e2
8 changed files with 1159 additions and 299 deletions
+64 -29
View File
@@ -38,7 +38,7 @@ use datafusion::{
use futures::Stream;
use parking_lot::Mutex;
use rustfs_s3select_api::{
QueryError, QueryResult, S3SelectPolicyError,
QueryError, QueryResult, SelectError,
query::{
Query,
ast::ExtStatement,
@@ -128,7 +128,7 @@ impl QueryDispatcher for SimpleQueryDispatcher {
.query_admission
.clone()
.try_acquire_owned()
.map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?;
.map_err(|_| QueryError::from(SelectError::QueryConcurrencyLimit))?;
Ok(QueryAdmission::new(Arc::new(permit)))
}
@@ -245,7 +245,7 @@ impl SimpleQueryDispatcher {
.query_admission
.clone()
.try_acquire_owned()
.map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?;
.map_err(|_| QueryError::from(SelectError::QueryConcurrencyLimit))?;
Arc::new(permit)
}
};
@@ -293,7 +293,7 @@ impl SimpleQueryDispatcher {
) -> QueryResult<T> {
let deadline = query_tracker.deadline();
let timeout_error = || {
S3SelectPolicyError::QueryTimeout {
SelectError::QueryTimeout {
seconds: query_tracker.timeout_seconds(),
}
.into()
@@ -343,11 +343,11 @@ impl SimpleQueryDispatcher {
return QueryError::Cancel;
}
match query_tracker.status() {
QueryExecutionStatus::TimedOut => S3SelectPolicyError::QueryTimeout {
QueryExecutionStatus::TimedOut => SelectError::QueryTimeout {
seconds: query_tracker.timeout_seconds(),
}
.into(),
QueryExecutionStatus::Active if Instant::now() >= query_tracker.deadline() => S3SelectPolicyError::QueryTimeout {
QueryExecutionStatus::Active if Instant::now() >= query_tracker.deadline() => SelectError::QueryTimeout {
seconds: query_tracker.timeout_seconds(),
}
.into(),
@@ -430,15 +430,11 @@ impl SimpleQueryDispatcher {
} else if *info == *USE {
file_format = file_format.with_has_header(true);
} else {
return Err(QueryError::NotImplemented {
err: "unsupported FileHeaderInfo".to_string(),
});
return Err(SelectError::InvalidDataSource.into());
}
}
_ => {
return Err(QueryError::NotImplemented {
err: "unsupported FileHeaderInfo".to_string(),
});
return Err(SelectError::InvalidDataSource.into());
}
}
if let Some(quote) = csv.quote_character.as_ref() {
@@ -462,9 +458,7 @@ impl SimpleQueryDispatcher {
.unwrap_or_else(|| ".json".to_string());
(ListingOptions::new(Arc::new(file_format)).with_file_extension(file_ext), false, false)
} else {
return Err(QueryError::NotImplemented {
err: "not support this file type".to_string(),
});
return Err(SelectError::InvalidDataSource.into());
};
let resolve_schema = listing_options.infer_schema(session.inner(), &table_path).await?;
@@ -642,7 +636,7 @@ impl Stream for TrackedRecordBatchStream {
}
fn query_timeout_error(timeout_seconds: u64) -> datafusion::common::DataFusionError {
datafusion::common::DataFusionError::External(Box::new(S3SelectPolicyError::QueryTimeout {
datafusion::common::DataFusionError::External(Box::new(SelectError::QueryTimeout {
seconds: timeout_seconds,
}))
}
@@ -791,7 +785,7 @@ mod tests {
};
use futures::{StreamExt, TryStreamExt, stream};
use rustfs_s3select_api::{
QueryError, QueryResult, S3SelectPolicyError,
QueryError, QueryResult, SelectError,
query::{
Context as QueryContext, Query,
dispatcher::QueryDispatcher,
@@ -1338,6 +1332,47 @@ mod tests {
assert_eq!(dispatcher.memory_limit_bytes, DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES);
}
#[tokio::test]
async fn invalid_csv_header_info_is_typed_invalid_data_source() {
let mut input = test_input();
input
.request
.input_serialization
.csv
.as_mut()
.expect("test input should use CSV")
.file_header_info = Some(FileHeaderInfo::from_static("INVALID"));
let input = Arc::new(input);
let optimizer = Arc::new(CascadeOptimizerBuilder::default().build());
let scheduler = Arc::new(LocalScheduler {});
let dispatcher = SimpleQueryDispatcherBuilder::default()
.with_input(Arc::clone(&input))
.with_default_table_provider(Arc::new(BaseTableProvider::default()))
.with_session_factory(Arc::new(SessionCtxFactory::new(true)))
.with_parser(Arc::new(DefaultParser::default()))
.with_query_execution_factory(Arc::new(SqlQueryExecutionFactory::new(optimizer, scheduler)))
.with_func_manager(Arc::new(SimpleFunctionMetadataManager::default()))
.build()
.expect("query dispatcher should build");
let query = Query::new(
QueryContext {
input: Arc::clone(&input),
},
input.request.expression.clone(),
);
let query_state_machine = dispatcher
.build_query_state_machine(query)
.await
.expect("query should acquire admission");
let error = match dispatcher.build_logical_plan(query_state_machine).await {
Err(error) => error,
Ok(_) => panic!("invalid FileHeaderInfo must fail while building the provider"),
};
assert_eq!(error.select_error(), SelectError::InvalidDataSource);
}
#[tokio::test]
async fn csv_query_uses_custom_record_delimiter_across_file_partitions() {
const ROW_COUNT: usize = 200_000;
@@ -1420,7 +1455,7 @@ mod tests {
assert!(matches!(
result,
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryConcurrencyLimit))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::QueryConcurrencyLimit))
));
}
@@ -1474,7 +1509,7 @@ mod tests {
assert!(matches!(
result,
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryConcurrencyLimit))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::QueryConcurrencyLimit))
));
}
@@ -1571,7 +1606,7 @@ mod tests {
assert!(matches!(
result,
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryTimeout { seconds: 0 }))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::QueryTimeout { seconds: 0 }))
));
assert_eq!(admission.available_permits(), 1);
}
@@ -1601,7 +1636,7 @@ mod tests {
assert!(matches!(
dispatcher.execute_logical_plan(logical_plan, query_state_machine).await,
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryTimeout { seconds: 300 }))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::QueryTimeout { seconds: 300 }))
));
assert_eq!(admission.available_permits(), 1);
}
@@ -1742,7 +1777,7 @@ mod tests {
assert_eq!(admission.available_permits(), 1);
assert!(matches!(
dispatcher.build_logical_plan(query_state_machine).await,
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryTimeout { seconds: 1 }))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::QueryTimeout { seconds: 1 }))
));
}
@@ -1843,7 +1878,7 @@ mod tests {
release_drop_tx.send(()).expect("release result drop");
assert!(matches!(
task.await.expect("deadline task should finish"),
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryTimeout { seconds: 1 }))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::QueryTimeout { seconds: 1 }))
));
assert_eq!(admission.available_permits(), 1);
}
@@ -1985,8 +2020,8 @@ mod tests {
panic!("expected external query error");
};
assert!(matches!(
source.downcast_ref::<S3SelectPolicyError>(),
Some(S3SelectPolicyError::QueryTimeout { seconds: 300 })
source.downcast_ref::<SelectError>(),
Some(SelectError::QueryTimeout { seconds: 300 })
));
assert!(inner_dropped.load(Ordering::SeqCst));
assert_eq!(admission.available_permits(), 1);
@@ -2124,8 +2159,8 @@ mod tests {
panic!("expected external query error");
};
assert!(matches!(
source.downcast_ref::<S3SelectPolicyError>(),
Some(S3SelectPolicyError::QueryTimeout { seconds: 300 })
source.downcast_ref::<SelectError>(),
Some(SelectError::QueryTimeout { seconds: 300 })
));
assert_eq!(admission.available_permits(), 1);
assert!(output.next().await.is_none());
@@ -2174,8 +2209,8 @@ mod tests {
panic!("expected external query error");
};
assert!(matches!(
source.downcast_ref::<S3SelectPolicyError>(),
Some(S3SelectPolicyError::QueryTimeout { seconds: 1 })
source.downcast_ref::<SelectError>(),
Some(SelectError::QueryTimeout { seconds: 1 })
));
assert_eq!(admission.available_permits(), 1);
assert!(output.next().await.is_none());
@@ -38,7 +38,7 @@ use datafusion::{
};
use futures::{FutureExt, TryFutureExt, future::BoxFuture};
use rustfs_s3select_api::{
QueryError, QueryResult,
QueryResult,
object_store::{SelectScanRange, scan_range_from_bounds},
};
use s3s::dto::SelectObjectContentInput;
@@ -106,7 +106,10 @@ impl ParquetSelectTable {
let object_store_url = table_path.object_store();
let object_location = Path::from(input.key.clone());
let store = state.runtime_env().object_store(&object_store_url)?;
let object_meta = store.head(&object_location).await.map_err(query_store_error)?;
let object_meta = store
.head(&object_location)
.await
.map_err(datafusion::common::DataFusionError::from)?;
let reader = ObjectStoreParquetReader {
store: Arc::clone(&store),
@@ -115,7 +118,7 @@ impl ParquetSelectTable {
};
let builder = ParquetRecordBatchStreamBuilder::new(reader)
.await
.map_err(query_store_error)?;
.map_err(datafusion::common::DataFusionError::from)?;
let schema = Arc::clone(builder.schema());
let metadata = Arc::clone(builder.metadata());
let access_plan = parquet_access_plan(input, object_meta.size, metadata.as_ref())?;
@@ -180,7 +183,8 @@ fn parquet_access_plan(
let Some(scan_range) = input.request.scan_range.as_ref() else {
return Ok(None);
};
let scan_range = scan_range_from_bounds(scan_range.start, scan_range.end, object_size).map_err(query_store_error)?;
let scan_range = scan_range_from_bounds(scan_range.start, scan_range.end, object_size)
.map_err(datafusion::common::DataFusionError::from)?;
Ok(scan_range.map(|range| Arc::new(access_plan_for_scan_range(range, metadata))))
}
@@ -214,10 +218,6 @@ fn parquet_store_error(err: ObjectStoreError) -> ParquetError {
ParquetError::External(Box::new(err))
}
fn query_store_error(err: impl fmt::Display) -> QueryError {
QueryError::StoreError { e: err.to_string() }
}
#[cfg(test)]
mod tests {
use super::*;
@@ -227,7 +227,13 @@ mod tests {
datatypes::{DataType, Field, Schema, SchemaRef},
record_batch::RecordBatch,
},
object_store::memory::InMemory,
parquet::arrow::{ArrowWriter, arrow_reader::ParquetRecordBatchReaderBuilder},
prelude::SessionContext,
};
use rustfs_s3select_api::SelectError;
use s3s::dto::{
CSVOutput, ExpressionType, InputSerialization, OutputSerialization, ParquetInput, ScanRange, SelectObjectContentRequest,
};
use std::{
fs::File,
@@ -275,6 +281,85 @@ mod tests {
assert!(!plan.should_scan(1));
}
#[test]
fn parquet_access_plan_has_typed_invalid_scan_range_error() {
let metadata = two_row_group_metadata();
let mut input = parquet_input("test.parquet");
input.request.scan_range = Some(ScanRange {
start: Some(10),
end: None,
});
let error = parquet_access_plan(&input, 10, metadata.as_ref()).expect_err("out-of-bounds range must fail");
assert_eq!(error.select_error(), SelectError::InvalidScanRange);
}
#[tokio::test]
async fn try_new_preserves_missing_object_error() {
let store = Arc::new(InMemory::new());
let context = parquet_session(store);
let state = context.state();
let error = match ParquetSelectTable::try_new(&state, &parquet_input("missing.parquet")).await {
Ok(_) => panic!("missing parquet object must fail"),
Err(error) => error,
};
assert_eq!(error.select_error(), SelectError::ObjectNotFound);
}
#[tokio::test]
async fn try_new_preserves_parquet_metadata_error() {
let store = Arc::new(InMemory::new());
let object = Path::from("corrupt.parquet");
store
.put(&object, Bytes::from_static(b"not a parquet file").into())
.await
.expect("put corrupt parquet object");
let context = parquet_session(store);
let state = context.state();
let error = match ParquetSelectTable::try_new(&state, &parquet_input(object.as_ref())).await {
Ok(_) => panic!("corrupt parquet metadata must fail"),
Err(error) => error,
};
assert_eq!(error.select_error(), SelectError::ParquetParsingError);
}
fn parquet_session(store: Arc<dyn ObjectStore>) -> SessionContext {
let context = SessionContext::new();
let store_url = ObjectStoreUrl::parse("s3://test-bucket").expect("valid test object store URL");
context.register_object_store(store_url.as_ref(), store);
context
}
fn parquet_input(key: &str) -> SelectObjectContentInput {
SelectObjectContentInput {
bucket: "test-bucket".to_string(),
expected_bucket_owner: None,
key: key.to_string(),
sse_customer_algorithm: None,
sse_customer_key: None,
sse_customer_key_md5: None,
request: SelectObjectContentRequest {
expression: "SELECT * FROM S3Object".to_string(),
expression_type: ExpressionType::from_static(ExpressionType::SQL),
input_serialization: InputSerialization {
parquet: Some(ParquetInput::default()),
..Default::default()
},
output_serialization: OutputSerialization {
csv: Some(CSVOutput::default()),
..Default::default()
},
request_progress: None,
scan_range: None,
},
}
}
fn two_row_group_metadata() -> Arc<ParquetMetaData> {
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)