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
+308 -15
View File
@@ -14,7 +14,12 @@
#![recursion_limit = "256"]
use datafusion::{common::DataFusionError, sql::sqlparser::parser::ParserError};
use datafusion::{
arrow::error::ArrowError,
common::{DataFusionError, SchemaError},
parquet::errors::ParquetError,
sql::sqlparser::parser::ParserError,
};
use std::{error::Error as StdError, fmt::Display};
use thiserror::Error;
@@ -67,23 +72,88 @@ pub enum QueryError {
StoreError { e: String },
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum S3SelectPolicyError {
#[derive(Clone, Debug, Error, PartialEq, Eq)]
pub enum SelectError {
#[error("The file is not in a supported compression format. Only GZIP and BZIP2 are supported.")]
InvalidCompressionFormat,
#[error("The data source type is not valid. Only CSV, JSON, and Parquet are supported.")]
InvalidDataSource,
#[error(
"Object decompression failed. Check that the object is properly compressed using the format specified in the request."
)]
TruncatedInput,
#[error("An error occurred while parsing the CSV file. Check the file and try again.")]
CsvParsingError,
#[error("An error occurred while parsing the JSON file. Check the file and try again.")]
JsonParsingError,
#[error("An error occurred while parsing the Parquet file. Check the file and try again.")]
ParquetParsingError,
#[error("{message}")]
ParseSelectFailure { message: String },
#[error("The SQL expression is invalid.")]
InvalidQuery,
#[error("The SQL expression contains a data type that is not valid.")]
InvalidDataType,
#[error("An incorrect argument type was specified in a function call in the SQL expression.")]
IncorrectSqlFunctionArgumentType,
#[error("The data source path in the SQL expression is not supported.")]
DataSourcePathUnsupported,
#[error("Unsupported S3 Select SQL structure: {message}")]
UnsupportedSqlStructure { message: String },
#[error("We encountered an unsupported SQL operation.")]
UnsupportedSqlOperation,
#[error("A column name or a path provided does not exist in the SQL expression.")]
EvaluatorBindingDoesNotExist,
#[error("The field name matches to multiple fields in the file. Check the SQL expression and the file, and try again.")]
AmbiguousFieldName,
#[error("The value of a parameter in ScanRange element is invalid. Check the service API documentation and try again.")]
InvalidScanRange,
#[error("S3 Select query concurrency limit reached")]
QueryConcurrencyLimit,
#[error("S3 Select query exceeded the {seconds}-second execution limit")]
QueryTimeout { seconds: u64 },
#[error("S3 Select query resource limit exceeded")]
ResourceExhausted,
#[error("The specified bucket does not exist.")]
BucketNotFound,
#[error("The specified key does not exist.")]
ObjectNotFound,
#[error("The query was canceled")]
Canceled,
#[error("An internal error occurred.")]
InternalError,
}
pub type S3SelectPolicyError = SelectError;
const MAX_ERROR_SOURCE_DEPTH: usize = 16;
impl QueryError {
fn source_error<T: StdError + 'static>(&self) -> Option<&T> {
let mut err: &(dyn StdError + 'static) = self;
for _ in 0..16 {
for _ in 0..MAX_ERROR_SOURCE_DEPTH {
if let Some(source) = err.downcast_ref::<T>() {
return Some(source);
}
@@ -99,10 +169,113 @@ impl QueryError {
pub fn s3_select_policy_error(&self) -> Option<&S3SelectPolicyError> {
self.source_error()
}
pub fn select_error(&self) -> SelectError {
let mut err: &(dyn StdError + 'static) = match self {
Self::Datafusion { source } => source.as_ref(),
_ => self,
};
for _ in 0..MAX_ERROR_SOURCE_DEPTH {
if let Some(select_error) = classify_select_error_source(err) {
return select_error;
}
let Some(source) = err.source() else {
break;
};
err = source;
}
match self {
QueryError::NotImplemented { .. } => SelectError::UnsupportedSqlOperation,
QueryError::MultiStatement { .. } => SelectError::UnsupportedSqlStructure {
message: "multiple SQL statements are not supported".to_string(),
},
QueryError::BuildQueryDispatcher { .. } | QueryError::FunctionExists { .. } | QueryError::StoreError { .. } => {
SelectError::InternalError
}
QueryError::Cancel => SelectError::Canceled,
QueryError::FunctionNotExists { .. } => SelectError::InvalidQuery,
QueryError::Datafusion { .. } | QueryError::Parser { .. } => SelectError::InternalError,
}
}
}
impl From<S3SelectPolicyError> for QueryError {
fn from(value: S3SelectPolicyError) -> Self {
fn classify_select_error_source(err: &(dyn StdError + 'static)) -> Option<SelectError> {
if let Some(error) = err.downcast_ref::<SelectError>() {
return Some(error.clone());
}
if let Some(error) = err.downcast_ref::<object_store::SelectObjectStoreError>() {
return Some(error.select_error());
}
if let Some(error) = err.downcast_ref::<datafusion::object_store::Error>() {
return match error {
datafusion::object_store::Error::NotFound { source, .. } => Some(
source
.downcast_ref::<object_store::SelectObjectStoreError>()
.map_or(SelectError::ObjectNotFound, object_store::SelectObjectStoreError::select_error),
),
_ => None,
};
}
if let Some(error) = err.downcast_ref::<ParserError>() {
return Some(SelectError::ParseSelectFailure {
message: error.to_string(),
});
}
if let Some(error) = err.downcast_ref::<ArrowError>() {
return match error {
ArrowError::CsvError(_) => Some(SelectError::CsvParsingError),
ArrowError::JsonError(_) => Some(SelectError::JsonParsingError),
ArrowError::ParquetError(_) => Some(SelectError::ParquetParsingError),
ArrowError::CastError(_) | ArrowError::ParseError(_) => Some(SelectError::InvalidDataType),
ArrowError::MemoryError(_) => Some(SelectError::ResourceExhausted),
ArrowError::ExternalError(_) | ArrowError::IoError(_, _) => None,
_ => Some(SelectError::InternalError),
};
}
if let Some(error) = err.downcast_ref::<ParquetError>() {
return match error {
ParquetError::External(_) => None,
_ => Some(SelectError::ParquetParsingError),
};
}
if let Some(error) = err.downcast_ref::<SchemaError>() {
return Some(match error {
SchemaError::FieldNotFound { .. } => SelectError::EvaluatorBindingDoesNotExist,
SchemaError::AmbiguousReference { .. }
| SchemaError::DuplicateQualifiedField { .. }
| SchemaError::DuplicateUnqualifiedField { .. } => SelectError::AmbiguousFieldName,
});
}
if let Some(error) = err.downcast_ref::<DataFusionError>() {
return match error {
DataFusionError::NotImplemented(_) => Some(SelectError::UnsupportedSqlOperation),
DataFusionError::Plan(_) => Some(SelectError::InvalidQuery),
DataFusionError::ResourcesExhausted(_) => Some(SelectError::ResourceExhausted),
DataFusionError::Internal(_)
| DataFusionError::Execution(_)
| DataFusionError::Configuration(_)
| DataFusionError::Substrait(_)
| DataFusionError::Ffi(_) => Some(SelectError::InternalError),
DataFusionError::ArrowError(_, _)
| DataFusionError::ParquetError(_)
| DataFusionError::ObjectStore(_)
| DataFusionError::IoError(_)
| DataFusionError::SQL(_, _)
| DataFusionError::SchemaError(_, _)
| DataFusionError::ExecutionJoin(_)
| DataFusionError::External(_)
| DataFusionError::Context(_, _)
| DataFusionError::Diagnostic(_, _)
| DataFusionError::Collection(_)
| DataFusionError::Shared(_) => None,
};
}
None
}
impl From<SelectError> for QueryError {
fn from(value: SelectError) -> Self {
Self::Datafusion {
source: Box::new(DataFusionError::External(Box::new(value))),
}
@@ -161,7 +334,7 @@ mod tests {
};
assert_eq!(err.to_string(), "Multi-statement not allow, found num:2, sql:SELECT 1; SELECT 2;");
let err = S3SelectPolicyError::UnsupportedSqlStructure {
let err = SelectError::UnsupportedSqlStructure {
message: "JOIN is not supported".to_string(),
};
assert_eq!(err.to_string(), "Unsupported S3 Select SQL structure: JOIN is not supported");
@@ -170,11 +343,11 @@ mod tests {
assert_eq!(err.to_string(), "The query has been canceled");
assert_eq!(
S3SelectPolicyError::QueryConcurrencyLimit.to_string(),
SelectError::QueryConcurrencyLimit.to_string(),
"S3 Select query concurrency limit reached"
);
assert_eq!(
S3SelectPolicyError::QueryTimeout { seconds: 300 }.to_string(),
SelectError::QueryTimeout { seconds: 300 }.to_string(),
"S3 Select query exceeded the 300-second execution limit"
);
@@ -223,12 +396,132 @@ mod tests {
#[test]
fn policy_error_is_recoverable_from_query_error() {
let err: QueryError = S3SelectPolicyError::QueryTimeout { seconds: 300 }.into();
let err: QueryError = SelectError::QueryTimeout { seconds: 300 }.into();
assert!(matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::QueryTimeout { seconds: 300 })
));
assert!(matches!(err.s3_select_policy_error(), Some(SelectError::QueryTimeout { seconds: 300 })));
}
#[test]
fn query_error_classifies_data_errors_without_display_matching() {
let cases = [
(
DataFusionError::ArrowError(Box::new(ArrowError::CsvError("private csv detail".to_string())), None),
SelectError::CsvParsingError,
),
(
DataFusionError::ArrowError(Box::new(ArrowError::JsonError("private json detail".to_string())), None),
SelectError::JsonParsingError,
),
(
DataFusionError::ParquetError(Box::new(ParquetError::General("private parquet detail".to_string()))),
SelectError::ParquetParsingError,
),
(
DataFusionError::External(Box::new(SelectError::TruncatedInput)),
SelectError::TruncatedInput,
),
(
DataFusionError::ArrowError(
Box::new(ArrowError::InvalidArgumentError("private implementation detail".to_string())),
None,
),
SelectError::InternalError,
),
(
DataFusionError::ArrowError(Box::new(ArrowError::CastError("invalid cast".to_string())), None),
SelectError::InvalidDataType,
),
(
DataFusionError::ArrowError(Box::new(ArrowError::MemoryError("query memory limit".to_string())), None),
SelectError::ResourceExhausted,
),
(
DataFusionError::Execution("private execution detail".to_string()),
SelectError::InternalError,
),
(DataFusionError::Plan("invalid expression".to_string()), SelectError::InvalidQuery),
(
DataFusionError::NotImplemented("unsupported expression".to_string()),
SelectError::UnsupportedSqlOperation,
),
(
DataFusionError::SchemaError(
Box::new(SchemaError::FieldNotFound {
field: Box::new(datafusion::common::Column::from_name("missing")),
valid_fields: Vec::new(),
}),
Box::new(None),
),
SelectError::EvaluatorBindingDoesNotExist,
),
(
DataFusionError::SchemaError(
Box::new(SchemaError::AmbiguousReference {
field: Box::new(datafusion::common::Column::from_name("duplicate")),
}),
Box::new(None),
),
SelectError::AmbiguousFieldName,
),
];
for (source, expected) in cases {
let error = QueryError::from(source);
assert_eq!(error.select_error(), expected, "wrong classification for {error:?}");
}
}
#[test]
fn query_error_preserves_typed_object_store_classification() {
let bucket_error = QueryError::from(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::NotFound {
path: "private-bucket/private-object".to_string(),
source: Box::new(object_store::SelectObjectStoreError::BucketNotFound {
source: SelectStorageError::BucketNotFound("private-bucket".to_string()),
}),
})));
let object_error = QueryError::from(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::NotFound {
path: "private-bucket/private-object".to_string(),
source: Box::new(object_store::SelectObjectStoreError::ObjectNotFound {
source: SelectStorageError::ObjectNotFound("private-bucket".to_string(), "private-object".to_string()),
}),
})));
let scan_range_error =
QueryError::from(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::Generic {
store: "test",
source: Box::new(object_store::SelectObjectStoreError::InvalidScanRange),
})));
let storage_error = QueryError::from(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::Generic {
store: "test",
source: Box::new(object_store::SelectObjectStoreError::Storage {
source: SelectStorageError::LessData,
}),
})));
assert_eq!(bucket_error.select_error(), SelectError::BucketNotFound);
assert_eq!(object_error.select_error(), SelectError::ObjectNotFound);
assert_eq!(scan_range_error.select_error(), SelectError::InvalidScanRange);
assert_eq!(storage_error.select_error(), SelectError::InternalError);
}
#[test]
fn select_error_source_traversal_stops_at_the_depth_bound() {
#[derive(Debug)]
struct CyclicError;
impl std::fmt::Display for CyclicError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("cyclic error")
}
}
impl StdError for CyclicError {
fn source(&self) -> Option<&(dyn StdError + 'static)> {
Some(self)
}
}
let error = QueryError::from(DataFusionError::External(Box::new(CyclicError)));
assert_eq!(error.select_error(), SelectError::InternalError);
}
#[test]
+115 -10
View File
@@ -13,7 +13,7 @@
// limitations under the License.
use crate::{
PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectOptions,
PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectError, SelectGetObjectReader, SelectObjectOptions,
SelectObjectSnapshot, SelectObjectSnapshotReadError, SelectStorageError, SelectStore, SnapshotConsistencyError,
query::{
parser::RustFsDialect,
@@ -115,6 +115,38 @@ pub(crate) enum EcObjectStoreBuildError {
Snapshot(#[source] SnapshotConsistencyError),
}
#[derive(Debug, thiserror::Error)]
pub(crate) enum SelectObjectStoreError {
#[error("SelectObjectContent bucket does not exist")]
BucketNotFound {
#[source]
source: SelectStorageError,
},
#[error("SelectObjectContent object does not exist")]
ObjectNotFound {
#[source]
source: SelectStorageError,
},
#[error("SelectObjectContent storage failure")]
Storage {
#[source]
source: SelectStorageError,
},
#[error("SelectObjectContent ScanRange is invalid")]
InvalidScanRange,
}
impl SelectObjectStoreError {
pub(crate) fn select_error(&self) -> SelectError {
match self {
Self::BucketNotFound { .. } => SelectError::BucketNotFound,
Self::ObjectNotFound { .. } => SelectError::ObjectNotFound,
Self::InvalidScanRange => SelectError::InvalidScanRange,
Self::Storage { .. } => SelectError::InternalError,
}
}
}
#[derive(Clone, Copy, Debug)]
pub struct SelectScanRange {
start: u64,
@@ -502,8 +534,7 @@ fn map_prepare_snapshot_error(bucket: &str, object: &str, err: PrepareSelectObje
}
fn map_build_error_to_s3(error: EcObjectStoreBuildError) -> S3Error {
let message = error.to_string();
let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, message);
let mut s3_error = S3Error::with_message(S3ErrorCode::InternalError, SelectError::InternalError.to_string());
s3_error.set_source(Box::new(error));
s3_error
}
@@ -519,15 +550,21 @@ fn snapshot_read_error(bucket: &str, object: &str, err: SelectObjectSnapshotRead
}
fn map_storage_error(bucket: &str, object: &str, err: SelectStorageError) -> o_Error {
if select_is_err_bucket_not_found(&err) || select_is_err_object_not_found(&err) || select_is_err_version_not_found(&err) {
if select_is_err_bucket_not_found(&err) {
return o_Error::NotFound {
path: format!("{bucket}/{object}"),
source: Box::new(err),
source: Box::new(SelectObjectStoreError::BucketNotFound { source: err }),
};
}
if select_is_err_object_not_found(&err) || select_is_err_version_not_found(&err) {
return o_Error::NotFound {
path: format!("{bucket}/{object}"),
source: Box::new(SelectObjectStoreError::ObjectNotFound { source: err }),
};
}
o_Error::Generic {
store: "EcObjectStore",
source: Box::new(err),
source: Box::new(SelectObjectStoreError::Storage { source: err }),
}
}
@@ -602,7 +639,7 @@ fn parse_scan_range_from_bounds(
fn invalid_scan_range_store_error() -> o_Error {
o_Error::Generic {
store: "EcObjectStore",
source: format!("ScanRange: {INVALID_SCAN_RANGE_MESSAGE}").into(),
source: Box::new(SelectObjectStoreError::InvalidScanRange),
}
}
@@ -1150,7 +1187,11 @@ where
})?
.map_err(|e| o_Error::Generic {
store: "EcObjectStore",
source: Box::new(e),
source: if e.kind() == std::io::ErrorKind::InvalidData {
Box::new(SelectError::JsonParsingError)
} else {
Box::new(e)
},
})?;
// ── 3. Yield phase (one Bytes per NDJSON line) ───────────────────
@@ -1341,12 +1382,13 @@ mod test {
SELECT_DEFAULT_READ_BUFFER_SIZE, SelectObjectOptions, SelectObjectSnapshot, SelectScanRange, SnapshotConsistencyError,
bytes_stream, convert_csv_delimiter_stream, convert_field_delimiter_stream, convert_record_delimiter_stream,
extract_json_sub_path_from_expression, find_delimiter, flatten_json_document_to_ndjson, http_range_spec_from_get_range,
json_document_ndjson_stream, json_document_ndjson_stream_with_parser, scan_range_from_bounds, scan_range_stream,
select_read_headers, snapshot_last_modified, validate_json_document_size,
json_document_ndjson_stream, json_document_ndjson_stream_with_parser, map_storage_error, scan_range_from_bounds,
scan_range_stream, select_read_headers, snapshot_last_modified, validate_json_document_size,
};
use crate::query::session::{QueryExecutionGuard, QueryExecutionOwner, QueryExecutionTracker};
use crate::storage_api::SelectPutObjReader;
use crate::storage_api::object_store::ObjectIO as _;
use crate::{QueryError, SelectError, SelectStorageError};
use bytes::Bytes;
use datafusion::{
common::DataFusionError,
@@ -2688,6 +2730,69 @@ mod test {
assert!(output.next().await.is_none());
}
#[tokio::test]
async fn malformed_json_document_stream_has_typed_select_error() {
let input = b"{bad".to_vec();
let memory_pool: Arc<dyn MemoryPool> =
Arc::new(GreedyMemoryPool::new(input.len() * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER));
let mut output = json_document_ndjson_stream(
Box::new(std::io::Cursor::new(input.clone())),
input.len() as u64,
None,
memory_pool,
None,
);
let source = output
.next()
.await
.expect("malformed JSON should produce one stream error")
.expect_err("malformed JSON DOCUMENT must fail");
let error = QueryError::from(DataFusionError::ObjectStore(Box::new(source)));
assert_eq!(error.select_error(), SelectError::JsonParsingError);
assert!(output.next().await.is_none());
}
#[test]
fn storage_error_mapper_preserves_protocol_classification() {
let classify = |source| QueryError::from(DataFusionError::ObjectStore(Box::new(source))).select_error();
assert_eq!(
classify(map_storage_error(
"private-bucket",
"private-object",
SelectStorageError::BucketNotFound("private-bucket".to_string()),
)),
SelectError::BucketNotFound
);
assert_eq!(
classify(map_storage_error(
"private-bucket",
"private-object",
SelectStorageError::ObjectNotFound("private-bucket".to_string(), "private-object".to_string()),
)),
SelectError::ObjectNotFound
);
assert_eq!(
classify(map_storage_error("private-bucket", "private-object", SelectStorageError::LessData)),
SelectError::InternalError
);
assert_eq!(
classify(scan_range_from_bounds(Some(10), None, 10).expect_err("out-of-bounds range must fail")),
SelectError::InvalidScanRange
);
let parquet_source = map_storage_error(
"private-bucket",
"private-object",
SelectStorageError::ObjectNotFound("private-bucket".to_string(), "private-object".to_string()),
);
let parquet_error = QueryError::from(DataFusionError::ParquetError(Box::new(
datafusion::parquet::errors::ParquetError::External(Box::new(parquet_source)),
)));
assert_eq!(parquet_error.select_error(), SelectError::ObjectNotFound);
}
#[test]
fn test_json_document_size_error_is_resource_exhausted() {
assert!(validate_json_document_size(super::MAX_JSON_DOCUMENT_BYTES).is_ok());
+8 -12
View File
@@ -433,7 +433,7 @@ impl SessionCtxFactory {
let path = Path::from(context.input.key.clone());
store.put(&path, data_bytes.into()).await.map_err(|e| {
error!("put data into memory failed: {}", e.to_string());
QueryError::StoreError { e: e.to_string() }
QueryError::from(DataFusionError::from(e))
})?;
df_session_state.with_object_store(&store_url, store).build()
@@ -477,16 +477,11 @@ fn test_parquet_bytes() -> QueryResult<Vec<u8>> {
let mut bytes = Vec::new();
{
let mut writer =
ArrowWriter::try_new(&mut bytes, schema, None).map_err(|e| QueryError::StoreError { e: e.to_string() })?;
writer
.write(&first_batch)
.map_err(|e| QueryError::StoreError { e: e.to_string() })?;
writer.flush().map_err(|e| QueryError::StoreError { e: e.to_string() })?;
writer
.write(&second_batch)
.map_err(|e| QueryError::StoreError { e: e.to_string() })?;
writer.close().map_err(|e| QueryError::StoreError { e: e.to_string() })?;
let mut writer = ArrowWriter::try_new(&mut bytes, schema, None).map_err(DataFusionError::from)?;
writer.write(&first_batch).map_err(DataFusionError::from)?;
writer.flush().map_err(DataFusionError::from)?;
writer.write(&second_batch).map_err(DataFusionError::from)?;
writer.close().map_err(DataFusionError::from)?;
}
Ok(bytes)
}
@@ -509,7 +504,8 @@ fn test_parquet_batch(
Arc::new(Int32Array::from(salaries.to_vec())),
],
)
.map_err(|e| QueryError::StoreError { e: e.to_string() })
.map_err(DataFusionError::from)
.map_err(QueryError::from)
}
#[cfg(test)]
+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)
+26 -12
View File
@@ -23,7 +23,7 @@ use datafusion::sql::{
},
};
use rustfs_s3select_api::{
QueryError, QueryResult, S3SelectPolicyError,
QueryError, QueryResult, SelectError,
query::{
ast::ExtStatement,
logical_planner::{LogicalPlanner, Plan, QueryPlan},
@@ -68,7 +68,7 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> {
match stmt {
Statement::Query(_) => {
validate_s3_select_statement(&stmt)?;
let df_plan = self.df_planner.sql_statement_to_plan(stmt)?;
let df_plan = self.df_planner.sql_statement_to_plan(stmt).map_err(classify_planner_error)?;
let plan = Plan::Query(QueryPlan {
df_plan,
is_tag_scan: false,
@@ -76,11 +76,25 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> {
Ok(plan)
}
_ => Err(QueryError::NotImplemented { err: stmt.to_string() }),
_ => Err(unsupported_structure("only SELECT queries are supported")),
}
}
}
fn classify_planner_error(error: datafusion::common::DataFusionError) -> QueryError {
if matches!(
&error,
datafusion::common::DataFusionError::Plan(message)
if message.starts_with("Failed to coerce arguments to satisfy a call to")
|| (message.starts_with("Internal error: Function '")
&& message.contains("' failed to match any signature, errors:"))
) {
return SelectError::IncorrectSqlFunctionArgumentType.into();
}
error.into()
}
fn validate_s3_select_statement(statement: &Statement) -> QueryResult<()> {
let Statement::Query(query) = statement else {
return Err(unsupported_structure("only SELECT queries are supported"));
@@ -191,7 +205,7 @@ fn validate_select(select: &Select) -> QueryResult<()> {
let ([ObjectNamePart::Identifier(table_name)] | [ObjectNamePart::Identifier(table_name), ObjectNamePart::Identifier(_)]) =
name.0.as_slice()
else {
return Err(unsupported_structure("the source must be S3Object"));
return Err(SelectError::DataSourcePathUnsupported.into());
};
let is_s3_object = if table_name.quote_style.is_some() {
table_name.value == "S3Object"
@@ -199,14 +213,14 @@ fn validate_select(select: &Select) -> QueryResult<()> {
table_name.value.eq_ignore_ascii_case("S3Object")
};
if !is_s3_object {
return Err(unsupported_structure("the source must be S3Object"));
return Err(SelectError::DataSourcePathUnsupported.into());
}
Ok(())
}
fn unsupported_structure(message: &str) -> QueryError {
S3SelectPolicyError::UnsupportedSqlStructure {
SelectError::UnsupportedSqlStructure {
message: message.to_string(),
}
.into()
@@ -234,7 +248,7 @@ mod tests {
use super::validate_s3_select_statement;
use crate::sql::parser::ExtParser;
use datafusion::sql::sqlparser::ast::Statement;
use rustfs_s3select_api::{S3SelectPolicyError, query::ast::ExtStatement};
use rustfs_s3select_api::{SelectError, query::ast::ExtStatement};
fn parse_statement(sql: &str) -> Statement {
let mut statements = ExtParser::parse_sql(sql).expect("SQL should parse");
@@ -271,7 +285,7 @@ mod tests {
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "JOIN is not supported"
Some(SelectError::UnsupportedSqlStructure { message }) if message == "JOIN is not supported"
)
));
}
@@ -284,7 +298,7 @@ mod tests {
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "subqueries are not supported"
Some(SelectError::UnsupportedSqlStructure { message }) if message == "subqueries are not supported"
)
));
}
@@ -297,7 +311,7 @@ mod tests {
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "subqueries are not supported"
Some(SelectError::UnsupportedSqlStructure { message }) if message == "subqueries are not supported"
)
));
}
@@ -310,7 +324,7 @@ mod tests {
validate_s3_select_statement(&statement),
Err(ref err) if matches!(
err.s3_select_policy_error(),
Some(S3SelectPolicyError::UnsupportedSqlStructure { message }) if message == "the source must be S3Object"
Some(SelectError::DataSourcePathUnsupported)
)
));
}
@@ -326,7 +340,7 @@ mod tests {
assert!(
matches!(
validate_s3_select_statement(&statement),
Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::UnsupportedSqlStructure { .. }))
Err(ref err) if matches!(err.s3_select_policy_error(), Some(SelectError::UnsupportedSqlStructure { .. }))
),
"query should be rejected: {sql}"
);
@@ -16,7 +16,7 @@
mod error_handling_tests {
use crate::get_global_db;
use rustfs_s3select_api::{
QueryError,
QueryError, SelectError,
query::{Context, Query},
};
use s3s::dto::{
@@ -98,7 +98,6 @@ mod error_handling_tests {
"INSERT INTO S3Object VALUES (1, 'test')",
"UPDATE S3Object SET name = 'test'",
"DELETE FROM S3Object",
"CREATE TABLE test (id INT)",
"DROP TABLE S3Object",
];
@@ -113,6 +112,68 @@ mod error_handling_tests {
}
}
#[tokio::test]
async fn test_non_select_statement_is_typed_unsupported_structure() {
let sql = "CREATE TABLE test (id INT)";
let input = create_test_input_with_sql(sql);
let db = get_global_db(input.clone(), true).await.unwrap();
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let error = match db.execute(&query).await {
Err(error) => error,
Ok(_) => panic!("non-SELECT statement must fail"),
};
assert!(matches!(error.select_error(), SelectError::UnsupportedSqlStructure { .. }));
}
#[tokio::test]
async fn test_function_argument_coercion_failure_is_typed() {
for sql in ["SELECT ROUND(3.14, 1.1) FROM S3Object", "SELECT SQRT(1, 2) FROM S3Object"] {
let input = create_test_input_with_sql(sql);
let db = get_global_db(input.clone(), true)
.await
.expect("test database should initialize");
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let error = match db.execute(&query).await {
Err(error) => error,
Ok(_) => panic!("invalid function arguments must fail during planning: {sql}"),
};
assert_eq!(
error.select_error(),
SelectError::IncorrectSqlFunctionArgumentType,
"unexpected planner error for {sql}: {error:?}"
);
}
}
#[tokio::test]
async fn test_other_planner_failures_remain_invalid_query() {
for sql in [
"SELECT DEFINITELY_UNKNOWN_FUNCTION(1) FROM S3Object",
"SELECT 1 + 'text' FROM S3Object",
] {
let input = create_test_input_with_sql(sql);
let db = get_global_db(input.clone(), true)
.await
.expect("test database should initialize");
let query = Query::new(Context { input: Arc::new(input) }, sql.to_string());
let error = match db.execute(&query).await {
Err(error) => error,
Ok(_) => panic!("invalid query must fail during planning: {sql}"),
};
assert_eq!(
error.select_error(),
SelectError::InvalidQuery,
"unexpected planner error for {sql}: {error:?}"
);
}
}
#[tokio::test]
async fn test_invalid_column_references() {
let invalid_column_sqls = vec![