mirror of
https://github.com/rustfs/rustfs.git
synced 2026-08-13 08:36:54 +00:00
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:
@@ -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)
|
||||
|
||||
@@ -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![
|
||||
|
||||
Reference in New Issue
Block a user