From 0364523dad15829a4282bf44565675af62cc3c44 Mon Sep 17 00:00:00 2001 From: GatewayJ <835269233@qq.com> Date: Sat, 25 Jul 2026 18:44:53 +0800 Subject: [PATCH] fix(s3select): enforce query and resource limits (#5028) * fix(s3select): enforce query and resource limits * fix(s3select): close query resource limit gaps * fix(s3select): preserve timeout and stream invariants * fix(s3select): enforce staged query limits * fix(s3select): preserve policy error compatibility * fix(s3select): bound error source traversal --- Cargo.lock | 1 + crates/s3select-api/Cargo.toml | 1 + crates/s3select-api/src/lib.rs | 85 + crates/s3select-api/src/object_store.rs | 640 ++++++- crates/s3select-api/src/query/execution.rs | 23 +- crates/s3select-api/src/query/session.rs | 379 ++++- crates/s3select-query/Cargo.toml | 2 +- .../s3select-query/src/dispatcher/manager.rs | 1483 ++++++++++++++++- crates/s3select-query/src/instance.rs | 109 +- crates/s3select-query/src/metadata/mod.rs | 6 +- crates/s3select-query/src/sql/planner.rs | 265 ++- .../src/test/integration_test.rs | 123 +- rustfs/src/app/select_object.rs | 333 +++- 13 files changed, 3228 insertions(+), 222 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index ac996ae88..f9bef4b33 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9868,6 +9868,7 @@ dependencies = [ "rustfs-test-utils", "s3s", "serde_json", + "serial_test", "tempfile", "thiserror 2.0.19", "tokio", diff --git a/crates/s3select-api/Cargo.toml b/crates/s3select-api/Cargo.toml index d1c678575..6138e18f9 100644 --- a/crates/s3select-api/Cargo.toml +++ b/crates/s3select-api/Cargo.toml @@ -49,6 +49,7 @@ url.workspace = true [dev-dependencies] rustfs-test-utils.workspace = true +serial_test.workspace = true tempfile.workspace = true [lib] diff --git a/crates/s3select-api/src/lib.rs b/crates/s3select-api/src/lib.rs index 47d7ec4c5..41ea340a0 100644 --- a/crates/s3select-api/src/lib.rs +++ b/crates/s3select-api/src/lib.rs @@ -64,6 +64,48 @@ pub enum QueryError { StoreError { e: String }, } +#[derive(Debug, Error)] +#[non_exhaustive] +pub enum S3SelectPolicyError { + #[error("Unsupported S3 Select SQL structure: {message}")] + UnsupportedSqlStructure { message: String }, + + #[error("S3 Select query concurrency limit reached")] + QueryConcurrencyLimit, + + #[error("S3 Select query exceeded the {seconds}-second execution limit")] + QueryTimeout { seconds: u64 }, +} + +impl S3SelectPolicyError { + fn from_error<'a>(mut err: &'a (dyn std::error::Error + 'static)) -> Option<&'a Self> { + for _ in 0..16 { + if let Some(policy_error) = err.downcast_ref::() { + return Some(policy_error); + } + err = err.source()?; + } + None + } +} + +impl QueryError { + pub fn s3_select_policy_error(&self) -> Option<&S3SelectPolicyError> { + match self { + Self::Datafusion { source } => S3SelectPolicyError::from_error(source.as_ref()), + _ => None, + } + } +} + +impl From for QueryError { + fn from(value: S3SelectPolicyError) -> Self { + Self::Datafusion { + source: Box::new(DataFusionError::External(Box::new(value))), + } + } +} + impl From for QueryError { fn from(value: DataFusionError) -> Self { match value { @@ -116,9 +158,23 @@ mod tests { }; assert_eq!(err.to_string(), "Multi-statement not allow, found num:2, sql:SELECT 1; SELECT 2;"); + let err = S3SelectPolicyError::UnsupportedSqlStructure { + message: "JOIN is not supported".to_string(), + }; + assert_eq!(err.to_string(), "Unsupported S3 Select SQL structure: JOIN is not supported"); + let err = QueryError::Cancel; assert_eq!(err.to_string(), "The query has been canceled"); + assert_eq!( + S3SelectPolicyError::QueryConcurrencyLimit.to_string(), + "S3 Select query concurrency limit reached" + ); + assert_eq!( + S3SelectPolicyError::QueryTimeout { seconds: 300 }.to_string(), + "S3 Select query exceeded the 300-second execution limit" + ); + let err = QueryError::FunctionNotExists { name: "my_func".to_string(), }; @@ -143,6 +199,35 @@ mod tests { } } + #[test] + fn query_error_variants_remain_source_compatible() { + fn exhaustive_match(err: QueryError) { + match err { + QueryError::Datafusion { .. } + | QueryError::NotImplemented { .. } + | QueryError::MultiStatement { .. } + | QueryError::BuildQueryDispatcher { .. } + | QueryError::Cancel + | QueryError::Parser { .. } + | QueryError::FunctionNotExists { .. } + | QueryError::FunctionExists { .. } + | QueryError::StoreError { .. } => {} + } + } + + exhaustive_match(QueryError::Cancel); + } + + #[test] + fn policy_error_is_recoverable_from_query_error() { + let err: QueryError = S3SelectPolicyError::QueryTimeout { seconds: 300 }.into(); + + assert!(matches!( + err.s3_select_policy_error(), + Some(S3SelectPolicyError::QueryTimeout { seconds: 300 }) + )); + } + #[test] fn test_query_error_from_parser_error() { let parser_error = ParserError::ParserError("syntax error".to_string()); diff --git a/crates/s3select-api/src/object_store.rs b/crates/s3select-api/src/object_store.rs index 4e18c8931..108e3e89b 100644 --- a/crates/s3select-api/src/object_store.rs +++ b/crates/s3select-api/src/object_store.rs @@ -14,20 +14,27 @@ use crate::{ SELECT_DEFAULT_READ_BUFFER_SIZE, SelectGetObjectReader, SelectObjectInfo, SelectObjectOptions, SelectStorageError, - SelectStore, resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, + SelectStore, + query::session::{QueryExecutionGuard, QueryExecutionTracker}, + resolve_select_object_store_handle, select_is_err_bucket_not_found, select_is_err_object_not_found, select_is_err_version_not_found, }; use async_trait::async_trait; use bytes::Bytes; use chrono::Utc; -use datafusion::object_store::{ - Attributes, CopyOptions, Error as o_Error, GetOptions, GetRange, GetResult, GetResultPayload, ListResult, MultipartUpload, - ObjectMeta, ObjectStore, PutMultipartOptions, PutOptions, PutPayload, PutResult, Result, path::Path, +use datafusion::{ + common::{DataFusionError, runtime::SpawnedTask}, + execution::memory_pool::{MemoryConsumer, MemoryPool, UnboundedMemoryPool}, + object_store::{ + Attributes, CopyOptions, Error as o_Error, GetOptions, GetRange, GetResult, GetResultPayload, ListResult, + MultipartUpload, ObjectMeta, ObjectStore, PutMultipartOptions, PutOptions, PutPayload, PutResult, Result, path::Path, + }, }; use futures::pin_mut; use futures::{Stream, StreamExt, future::ready, stream}; use futures_core::stream::BoxStream; use http::{HeaderMap, HeaderValue, header::HeaderName}; +use parking_lot::Mutex; use rustfs_common::DEFAULT_DELIMITER; use s3s::S3Result; use s3s::dto::SelectObjectContentInput; @@ -49,6 +56,13 @@ fn select_default_read_buffer_size_u64() -> u64 { u64::try_from(SELECT_DEFAULT_READ_BUFFER_SIZE).unwrap_or(u64::MAX) } +fn validated_object_size(size: i64) -> Result { + u64::try_from(size).map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + }) +} + /// Maximum allowed object size for JSON DOCUMENT mode. /// /// JSON DOCUMENT format requires loading the entire file into memory for DOM @@ -61,8 +75,13 @@ fn select_default_read_buffer_size_u64() -> u64 { /// size limit. /// /// Default: 128 MiB. This matches the AWS S3 Select limit for JSON DOCUMENT -/// inputs. +/// inputs. The query memory pool also applies: RustFS reserves 64 times the +/// input size for parsing and output. With the default 64 MiB query memory +/// limit, JSON DOCUMENT inputs larger than 1 MiB are rejected; raise +/// `RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES` to process larger inputs, up to this +/// hard cap. pub const MAX_JSON_DOCUMENT_BYTES: u64 = 128 * 1024 * 1024; +const JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER: usize = 64; pub const INVALID_SCAN_RANGE_MESSAGE: &str = "The value of a parameter in ScanRange element is invalid. Check the service API documentation and try again."; @@ -79,6 +98,8 @@ pub struct EcObjectStore { /// expression. When set, `flatten_json_document_to_ndjson` navigates to /// this key in the root JSON object before flattening. json_sub_path: Option, + memory_pool: Arc, + query_tracker: Option, store: Arc, } @@ -108,7 +129,29 @@ pub struct InvalidScanRange; impl EcObjectStore { pub fn new(input: Arc) -> S3Result { - let Some(store) = resolve_select_object_store_handle() else { + Self::build(input, Arc::new(UnboundedMemoryPool::default()), None, None) + } + + pub(crate) fn new_with_memory_pool(input: Arc, memory_pool: Arc) -> S3Result { + Self::build(input, memory_pool, None, None) + } + + pub(crate) fn new_with_query_tracker( + input: Arc, + memory_pool: Arc, + query_tracker: QueryExecutionTracker, + store: Option>, + ) -> S3Result { + Self::build(input, memory_pool, Some(query_tracker), store) + } + + fn build( + input: Arc, + memory_pool: Arc, + query_tracker: Option, + store: Option>, + ) -> S3Result { + let Some(store) = store.or_else(resolve_select_object_store_handle) else { return Err(s3_error!(InternalError, "ec store not inited")); }; @@ -151,6 +194,8 @@ impl EcObjectStore { delimiter, is_json_document, json_sub_path, + memory_pool, + query_tracker, store, }) } @@ -213,13 +258,30 @@ impl EcObjectStore { if range.is_empty() { return Ok(Bytes::new()); } - let reader = self.object_reader(Some(http_range_spec_from_range(range)), opts).await?; + let reader = self + .object_reader(Some(http_range_spec_from_range(range.clone())), opts) + .await?; + let object_size = validated_object_size(reader.object_info.size)?; + let resolved_range = GetRange::Bounded(range) + .as_range(object_size) + .map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + })?; + let expected_size = usize::try_from(resolved_range.end - resolved_range.start).map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + })?; let mut reader = reader.stream; let mut bytes = Vec::new(); reader.read_to_end(&mut bytes).await.map_err(|err| o_Error::Generic { store: "EcObjectStore", source: Box::new(err), })?; + if bytes.len() < expected_size { + return Err(incomplete_object_stream_error(expected_size - bytes.len())); + } + bytes.truncate(expected_size); Ok(Bytes::from(bytes)) } @@ -421,7 +483,7 @@ impl ObjectStore for EcObjectStore { let opts = self.object_options(&options); let needs_scan_context = options.range.is_none() && !options.head && self.input.request.scan_range.is_some(); let source_size = if needs_scan_context { - Some(self.object_info(&opts).await?.size as u64) + Some(validated_object_size(self.object_info(&opts).await?.size)?) } else { None }; @@ -442,7 +504,10 @@ impl ObjectStore for EcObjectStore { self.object_reader(range, &opts).await? }; - let original_size = source_size.unwrap_or(reader.object_info.size as u64); + let original_size = match source_size { + Some(source_size) => source_size, + None => validated_object_size(reader.object_info.size)?, + }; let etag = reader.object_info.etag; let version = reader.object_info.version_id.map(|version| version.to_string()); let attributes = Attributes::default(); @@ -473,20 +538,14 @@ impl ObjectStore for EcObjectStore { // which must load the whole file into memory; rejecting oversized // files upfront is safer than risking OOM. Users should convert // their data to JSON LINES (NDJSON) format for large files. - if original_size > MAX_JSON_DOCUMENT_BYTES { - return Err(o_Error::Generic { - store: "EcObjectStore", - source: format!( - "JSON DOCUMENT object is {original_size} bytes, which exceeds the \ - maximum allowed size of {MAX_JSON_DOCUMENT_BYTES} bytes \ - ({} MiB). Convert the input to JSON LINES (NDJSON) to process \ - large files.", - MAX_JSON_DOCUMENT_BYTES / (1024 * 1024) - ) - .into(), - }); - } - let stream = json_document_ndjson_stream(reader.stream, original_size, self.json_sub_path.clone()); + validate_json_document_size(original_size)?; + let stream = json_document_ndjson_stream( + reader.stream, + original_size, + self.json_sub_path.clone(), + Arc::clone(&self.memory_pool), + self.query_tracker.clone(), + ); GetResultPayload::Stream(stream) } else if let Some((_, scan_range)) = scan_context { let delimiter = self.record_delimiter(); @@ -503,6 +562,7 @@ impl ObjectStore for EcObjectStore { scan_range, include_header && header.is_none(), read_start, + original_size, ) .boxed(); let stream = if let Some(header) = header { @@ -664,6 +724,7 @@ struct ScanRangeState { record: Vec, pending: VecDeque, done: bool, + expected_end: u64, } fn scan_range_stream( @@ -672,6 +733,7 @@ fn scan_range_stream( range: SelectScanRange, include_header: bool, base_offset: u64, + expected_end: u64, ) -> BoxStream<'static, Result> where S: Stream> + Send + Unpin + 'static, @@ -686,6 +748,7 @@ where record: Vec::new(), pending: VecDeque::new(), done: false, + expected_end, }; stream::unfold(state, |mut state| async move { @@ -709,6 +772,10 @@ where )); } None => { + if state.offset < state.expected_end { + state.done = true; + return Some((Err(incomplete_object_stream_error(state.expected_end - state.offset)), state)); + } state.finish_pending_record(); state.done = true; } @@ -855,11 +922,57 @@ fn json_document_ndjson_stream( stream: Box, original_size: u64, json_sub_path: Option, + memory_pool: Arc, + query_tracker: Option, ) -> futures_core::stream::BoxStream<'static, Result> { + json_document_ndjson_stream_with_parser( + stream, + original_size, + json_sub_path, + memory_pool, + query_tracker, + |all_bytes, json_sub_path| parse_json_document_to_lines(&all_bytes, json_sub_path.as_deref()), + ) +} + +fn json_document_ndjson_stream_with_parser

( + stream: Box, + original_size: u64, + json_sub_path: Option, + memory_pool: Arc, + query_tracker: Option, + parser: P, +) -> futures_core::stream::BoxStream<'static, Result> +where + P: FnOnce(Vec, Option) -> std::io::Result> + Send + 'static, +{ AsyncTryStream::::new(|mut y| async move { + // Compact JSON can expand substantially into a serde_json DOM and + // per-record output buffers, so reserve a conservative upper bound + // before the source buffer is allocated. + let buffer_capacity = usize::try_from(original_size).map_err(|_| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT input size {original_size} does not fit in memory" + ))), + })?; + let reservation_bytes = buffer_capacity + .checked_mul(JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER) + .ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT memory reservation overflow for {original_size} input bytes" + ))), + })?; + let reservation = MemoryConsumer::new("S3 Select JSON document").register(&memory_pool); + reservation.try_resize(reservation_bytes).map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + })?; + pin_mut!(stream); // ── 1. Read phase (lazy: only runs when the stream is polled) ──── - let mut all_bytes = Vec::with_capacity(original_size as usize); + let mut all_bytes = Vec::with_capacity(buffer_capacity); stream .take(original_size) .read_to_end(&mut all_bytes) @@ -868,18 +981,26 @@ fn json_document_ndjson_stream( store: "EcObjectStore", source: Box::new(e), })?; + if all_bytes.len() != buffer_capacity { + return Err(incomplete_object_stream_error(buffer_capacity - all_bytes.len())); + } // ── 2. Parse phase (blocking thread pool, non-blocking runtime) ── - let lines = tokio::task::spawn_blocking(move || parse_json_document_to_lines(&all_bytes, json_sub_path.as_deref())) - .await - .map_err(|e| o_Error::Generic { - store: "EcObjectStore", - source: e.to_string().into(), - })? - .map_err(|e| o_Error::Generic { - store: "EcObjectStore", - source: Box::new(e), - })?; + let pending_query_guard = PendingQueryExecutionGuard::new(query_tracker); + let task_query_guard = pending_query_guard.task_state(); + let (lines, _reservation, _query_guard) = SpawnedTask::spawn_blocking(move || { + let query_guard = PendingQueryExecutionGuard::start(&task_query_guard)?; + parser(all_bytes, json_sub_path).map(|lines| (lines, reservation, query_guard)) + }) + .await + .map_err(|e| o_Error::Generic { + store: "EcObjectStore", + source: e.to_string().into(), + })? + .map_err(|e| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(e), + })?; // ── 3. Yield phase (one Bytes per NDJSON line) ─────────────────── for line in lines { @@ -890,6 +1011,66 @@ fn json_document_ndjson_stream( .boxed() } +struct PendingQueryExecutionGuard { + state: Arc>, +} + +enum QueryExecutionGuardState { + Pending(Option), + Started, + Cancelled, +} + +impl PendingQueryExecutionGuard { + fn new(query_tracker: Option) -> Self { + Self { + state: Arc::new(Mutex::new(QueryExecutionGuardState::Pending(query_tracker))), + } + } + + fn task_state(&self) -> Arc> { + Arc::clone(&self.state) + } + + fn start(state: &Mutex) -> std::io::Result> { + let mut state = state.lock(); + match std::mem::replace(&mut *state, QueryExecutionGuardState::Started) { + QueryExecutionGuardState::Pending(None) => Ok(None), + QueryExecutionGuardState::Pending(Some(query_tracker)) => query_tracker.query_guard().map(Some).ok_or_else(|| { + std::io::Error::new(std::io::ErrorKind::Interrupted, "JSON DOCUMENT parse was cancelled before it started") + }), + QueryExecutionGuardState::Cancelled => { + *state = QueryExecutionGuardState::Cancelled; + Err(std::io::Error::new( + std::io::ErrorKind::Interrupted, + "JSON DOCUMENT parse was cancelled before it started", + )) + } + QueryExecutionGuardState::Started => { + *state = QueryExecutionGuardState::Started; + Err(std::io::Error::other("JSON DOCUMENT parse started more than once")) + } + } + } +} + +impl Drop for PendingQueryExecutionGuard { + fn drop(&mut self) { + let query_guard = { + let mut state = self.state.lock(); + match std::mem::replace(&mut *state, QueryExecutionGuardState::Cancelled) { + QueryExecutionGuardState::Pending(query_guard) => query_guard, + QueryExecutionGuardState::Started => { + *state = QueryExecutionGuardState::Started; + None + } + QueryExecutionGuardState::Cancelled => None, + } + }; + drop(query_guard); + } +} + /// Parse a JSON DOCUMENT (a single JSON value, possibly multi-line) into a /// list of NDJSON lines – one [`Bytes`] per record. /// @@ -907,19 +1088,11 @@ fn parse_json_document_to_lines(bytes: &[u8], json_sub_path: Option<&str>) -> st // Navigate into the sub-path when the root is an object and a path was // extracted from the SQL FROM clause (e.g. `FROM s3object.employees`). - let value = if let Some(path) = json_sub_path { - if let serde_json::Value::Object(ref obj) = root { - match obj.get(path) { - Some(sub) => sub.clone(), - // Path not found – fall back to emitting the whole root object. - None => root, - } - } else { - // Root is already an array or scalar; ignore the path hint. - root + let value = match (root, json_sub_path) { + (serde_json::Value::Object(mut object), Some(path)) => { + object.remove(path).unwrap_or_else(|| serde_json::Value::Object(object)) } - } else { - root + (root, _) => root, }; let mut lines: Vec = Vec::new(); @@ -979,30 +1152,55 @@ where y.yield_ok(bytes).await; } if remaining > 0 { - return Err(o_Error::Generic { - store: "EcObjectStore", - source: Box::new(std::io::Error::new( - std::io::ErrorKind::UnexpectedEof, - format!("object stream ended with {remaining} bytes remaining"), - )), - }); + return Err(incomplete_object_stream_error(remaining)); } Ok(()) }) } +fn validate_json_document_size(original_size: u64) -> Result<()> { + if original_size <= MAX_JSON_DOCUMENT_BYTES { + return Ok(()); + } + + Err(o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT object is {original_size} bytes, which exceeds the maximum allowed size of \ + {MAX_JSON_DOCUMENT_BYTES} bytes ({} MiB). Convert the input to JSON LINES (NDJSON) to process large files.", + MAX_JSON_DOCUMENT_BYTES / (1024 * 1024) + ))), + }) +} + +fn incomplete_object_stream_error(remaining: impl std::fmt::Display) -> o_Error { + o_Error::Generic { + store: "EcObjectStore", + source: Box::new(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + format!("object stream ended with {remaining} bytes remaining"), + )), + } +} + #[cfg(test)] mod test { use super::{ - SelectScanRange, bytes_stream, convert_field_delimiter_stream, extract_json_sub_path_from_expression, find_delimiter, - flatten_json_document_to_ndjson, http_range_spec_from_get_range, replace_symbol, scan_range_from_bounds, - scan_range_read_start, scan_range_stream, select_read_headers, + EcObjectStore, JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER, SelectScanRange, bytes_stream, + convert_field_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, replace_symbol, + scan_range_from_bounds, scan_range_read_start, scan_range_stream, select_read_headers, validate_json_document_size, + validated_object_size, }; + use crate::query::session::{QueryExecutionGuard, QueryExecutionOwner, QueryExecutionTracker}; use crate::storage_api::SelectPutObjReader; use crate::storage_api::object_store::ObjectIO as _; use bytes::Bytes; - use datafusion::object_store::{self, GetRange}; - use datafusion::object_store::{GetOptions, GetResultPayload, ObjectStore as _, path::Path}; + use datafusion::{ + common::DataFusionError, + execution::memory_pool::{GreedyMemoryPool, MemoryPool}, + object_store::{self, GetOptions, GetRange, GetResultPayload, ObjectStore as _, path::Path}, + }; use futures::{StreamExt, TryStreamExt, stream}; use s3s::dto::{ CSVInput, CSVOutput, ExpressionType, InputSerialization, OutputSerialization, SelectObjectContentInput, @@ -1017,16 +1215,28 @@ mod test { atomic::{AtomicUsize, Ordering}, }; + #[test] + fn ec_object_store_constructor_remains_source_compatible() { + let _constructor: fn(Arc) -> s3s::S3Result = EcObjectStore::new; + } + use tokio::sync::Semaphore; + #[test] fn test_replace() { let result = replace_symbol(b"&&", b"dandan&&is&&best"); assert_eq!(result, b"dandan,is,best"); } + #[test] + fn test_validated_object_size_rejects_negative_metadata() { + assert_eq!(validated_object_size(0).expect("zero object size should be valid"), 0); + assert!(validated_object_size(-1).is_err()); + } + #[tokio::test] async fn test_scan_range_stream_keeps_header_and_selected_record() { let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"h1,h2\n1,a\n2,b\n3,c\n"))]); - let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(10, 11), true, 0); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(10, 11), true, 0, 18); let mut output = Vec::new(); while let Some(bytes) = stream.next().await { output.extend_from_slice(&bytes.unwrap()); @@ -1037,7 +1247,7 @@ mod test { #[tokio::test] async fn test_scan_range_stream_skips_record_when_start_is_in_middle() { let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"1,a\n2,b\n3,c\n"))]); - let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(2, 7), false, 0); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(2, 7), false, 0, 12); let mut output = Vec::new(); while let Some(bytes) = stream.next().await { output.extend_from_slice(&bytes.unwrap()); @@ -1048,7 +1258,7 @@ mod test { #[tokio::test] async fn test_scan_range_stream_keeps_record_when_end_is_in_middle() { let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"1,a\n2,b\n3,c\n"))]); - let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(0, 5), false, 0); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(0, 5), false, 0, 12); let mut output = Vec::new(); while let Some(bytes) = stream.next().await { output.extend_from_slice(&bytes.unwrap()); @@ -1059,7 +1269,7 @@ mod test { #[tokio::test] async fn test_scan_range_stream_uses_base_offset_for_range_reader() { let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"\n2,b\n3,c\n"))]); - let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(4, 7), false, 3); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(4, 7), false, 3, 12); let mut output = Vec::new(); while let Some(bytes) = stream.next().await { output.extend_from_slice(&bytes.unwrap()); @@ -1074,7 +1284,7 @@ mod test { Ok::<_, std::io::Error>(Bytes::from_static(b"\n1,a\r\n2,b\r")), Ok::<_, std::io::Error>(Bytes::from_static(b"\n3,c\r\n")), ]); - let mut stream = scan_range_stream(chunks, b"\r\n".to_vec(), SelectScanRange::new(12, 14), true, 0); + let mut stream = scan_range_stream(chunks, b"\r\n".to_vec(), SelectScanRange::new(12, 14), true, 0, 22); let mut output = Vec::new(); while let Some(bytes) = stream.next().await { output.extend_from_slice(&bytes.unwrap()); @@ -1082,6 +1292,26 @@ mod test { assert_eq!(output, b"h1,h2\r\n2,b\r\n"); } + #[tokio::test] + async fn test_scan_range_stream_rejects_early_eof() { + let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"1,a\n"))]); + let mut output = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(0, 7), false, 0, 8); + + assert_eq!(output.next().await.expect("first stream item").expect("first record"), b"1,a\n"[..]); + let err = output + .next() + .await + .expect("early EOF error") + .expect_err("short ScanRange stream must fail"); + let object_store::Error::Generic { source, .. } = err else { + panic!("expected generic object store error"); + }; + let source = source.downcast_ref::().expect("I/O error source"); + assert_eq!(source.kind(), std::io::ErrorKind::UnexpectedEof); + assert!(source.to_string().contains("4 bytes remaining")); + assert!(output.next().await.is_none()); + } + #[test] fn test_scan_range_read_start_keeps_full_delimiter_boundary() { let range = SelectScanRange::new(10, 20); @@ -1157,7 +1387,7 @@ mod test { #[tokio::test] async fn test_scan_range_output_can_convert_field_delimiter() { let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"a&&1\nb&&2\n"))]); - let stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(0, 10), false, 0); + let stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange::new(0, 10), false, 0, 10); let mut stream = convert_field_delimiter_stream(stream, Some("&&".to_string())); let mut output = Vec::new(); while let Some(bytes) = stream.next().await { @@ -1195,6 +1425,7 @@ mod test { } #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] async fn test_get_opts_validates_raw_length_before_delimiter_conversion() { let temp_root = tempfile::tempdir().expect("create s3select test temp root"); let env = rustfs_test_utils::TestECStoreEnv::builder() @@ -1242,6 +1473,8 @@ mod test { delimiter: "&&".to_string(), is_json_document: false, json_sub_path: None, + memory_pool: Arc::new(GreedyMemoryPool::new(1024)), + query_tracker: None, store: env.ecstore, }; @@ -1255,6 +1488,13 @@ mod test { let chunks: Vec = stream.try_collect().await.expect("collect converted object stream"); assert_eq!(chunks.concat(), b"a,1\n"); + + let requested_range = 3..10; + let ranges = store + .get_ranges(&Path::from(object), std::slice::from_ref(&requested_range)) + .await + .expect("bounded range past EOF should return the object remainder"); + assert_eq!(ranges, vec![Bytes::from_static(b"1\n")]); } #[tokio::test] @@ -1302,6 +1542,278 @@ mod test { assert!(output.next().await.is_none()); } + #[tokio::test] + async fn test_json_document_stream_respects_query_memory_pool() { + let input = b"{}".to_vec(); + let required = input.len() * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER; + let memory_pool: Arc = Arc::new(GreedyMemoryPool::new(required - 1)); + let mut output = json_document_ndjson_stream( + Box::new(std::io::Cursor::new(input.clone())), + input.len() as u64, + None, + memory_pool, + None, + ); + + let err = output + .next() + .await + .expect("memory error") + .expect_err("reservation should exceed the pool"); + let object_store::Error::Generic { source, .. } = err else { + panic!("expected generic object store error"); + }; + assert!(matches!( + source.downcast_ref::(), + Some(DataFusionError::ResourcesExhausted(_)) + )); + } + + #[tokio::test] + async fn test_json_document_stream_releases_memory_reservation() { + let input = b"[1,2]".to_vec(); + let required = input.len() * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER; + let memory_pool = Arc::new(GreedyMemoryPool::new(required)); + let output: Vec = json_document_ndjson_stream( + Box::new(std::io::Cursor::new(input.clone())), + input.len() as u64, + None, + memory_pool.clone(), + None, + ) + .try_collect() + .await + .expect("JSON conversion should fit the pool"); + + assert_eq!(output, vec![Bytes::from_static(b"1\n"), Bytes::from_static(b"2\n")]); + assert_eq!(memory_pool.reserved(), 0); + } + + #[tokio::test] + async fn test_json_document_stream_rejects_early_eof() { + let input = b"{}".to_vec(); + let memory_pool: Arc = Arc::new(GreedyMemoryPool::new(4 * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER)); + let mut output = json_document_ndjson_stream(Box::new(std::io::Cursor::new(input)), 4, None, memory_pool, None); + + let err = output + .next() + .await + .expect("early EOF error") + .expect_err("short JSON document must fail"); + let object_store::Error::Generic { source, .. } = err else { + panic!("expected generic object store error"); + }; + let source = source.downcast_ref::().expect("I/O error source"); + assert_eq!(source.kind(), std::io::ErrorKind::UnexpectedEof); + assert!(source.to_string().contains("2 bytes remaining")); + assert!(output.next().await.is_none()); + } + + #[test] + fn test_json_document_size_error_is_resource_exhausted() { + assert!(validate_json_document_size(super::MAX_JSON_DOCUMENT_BYTES).is_ok()); + + let err = validate_json_document_size(super::MAX_JSON_DOCUMENT_BYTES + 1).expect_err("oversized JSON document must fail"); + let object_store::Error::Generic { source, .. } = err else { + panic!("expected generic object store error"); + }; + assert!(matches!( + source.downcast_ref::(), + Some(DataFusionError::ResourcesExhausted(_)) + )); + } + + #[test] + fn test_json_document_queued_parse_releases_query_guard_when_cancelled() { + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(2) + .max_blocking_threads(1) + .enable_all() + .build() + .expect("build test runtime"); + + runtime.block_on(async { + let (blocking_started_tx, blocking_started_rx) = tokio::sync::oneshot::channel(); + let (release_blocking_tx, release_blocking_rx) = std::sync::mpsc::channel(); + let blocker = tokio::task::spawn_blocking(move || { + let _ = blocking_started_tx.send(()); + release_blocking_rx.recv().expect("release blocking worker"); + }); + blocking_started_rx.await.expect("blocking worker should start"); + + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("query permit should be available"); + let query_guard: QueryExecutionGuard = Arc::new(permit); + let query_tracker = QueryExecutionTracker::new( + &QueryExecutionOwner::new(), + query_guard, + tokio::time::Instant::now() + std::time::Duration::from_secs(30), + 30, + ); + let input = b"{}".to_vec(); + let memory_pool: Arc = + 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, + Some(query_tracker), + ); + + { + let next = output.next(); + futures::pin_mut!(next); + assert!(futures::poll!(next.as_mut()).is_pending()); + } + drop(output); + + let recovered_permit = + tokio::time::timeout(std::time::Duration::from_secs(5), Arc::clone(&admission).acquire_owned()) + .await + .expect("queued JSON parse should be cancelled") + .expect("query admission should remain open"); + release_blocking_tx.send(()).expect("release blocking worker"); + blocker.await.expect("blocking worker should finish"); + drop(recovered_permit); + assert_eq!(admission.available_permits(), 1); + }); + } + + #[test] + fn test_json_document_expired_queued_parse_does_not_start() { + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(2) + .max_blocking_threads(1) + .enable_all() + .build() + .expect("build test runtime"); + + runtime.block_on(async { + let (blocking_started_tx, blocking_started_rx) = tokio::sync::oneshot::channel(); + let (release_blocking_tx, release_blocking_rx) = std::sync::mpsc::channel(); + let blocker = tokio::task::spawn_blocking(move || { + let _ = blocking_started_tx.send(()); + release_blocking_rx.recv().expect("release blocking worker"); + }); + blocking_started_rx.await.expect("blocking worker should start"); + + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("query permit should be available"); + let owner = QueryExecutionOwner::new(); + let query_tracker = QueryExecutionTracker::new( + &owner, + Arc::new(permit), + tokio::time::Instant::now() + std::time::Duration::from_secs(30), + 30, + ); + let input = b"{}".to_vec(); + let memory_pool: Arc = + Arc::new(GreedyMemoryPool::new(input.len() * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER)); + let parser_started = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let parser_started_in_task = Arc::clone(&parser_started); + let mut output = json_document_ndjson_stream_with_parser( + Box::new(std::io::Cursor::new(input.clone())), + input.len() as u64, + None, + memory_pool, + Some(query_tracker.clone()), + move |_, _| { + parser_started_in_task.store(true, std::sync::atomic::Ordering::SeqCst); + Ok(vec![Bytes::from_static(b"{}\n")]) + }, + ); + + { + let next = output.next(); + futures::pin_mut!(next); + assert!(futures::poll!(next.as_mut()).is_pending()); + } + query_tracker.expire(&owner); + assert_eq!(admission.available_permits(), 1); + release_blocking_tx.send(()).expect("release blocking worker"); + blocker.await.expect("blocking worker should finish"); + + let err = tokio::time::timeout(std::time::Duration::from_secs(5), output.next()) + .await + .expect("queued parser should resume") + .expect("queued parser should return an error") + .expect_err("expired queued parser must not run"); + let object_store::Error::Generic { source, .. } = err else { + panic!("expected generic object store error"); + }; + let source = source.downcast_ref::().expect("I/O error source"); + assert_eq!(source.kind(), std::io::ErrorKind::Interrupted); + assert!(!parser_started.load(std::sync::atomic::Ordering::SeqCst)); + }); + } + + #[test] + fn test_json_document_started_parse_retains_query_guard_when_cancelled() { + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(2) + .max_blocking_threads(1) + .enable_all() + .build() + .expect("build test runtime"); + + runtime.block_on(async { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("query permit should be available"); + let query_guard: QueryExecutionGuard = Arc::new(permit); + let query_tracker = QueryExecutionTracker::new( + &QueryExecutionOwner::new(), + query_guard, + tokio::time::Instant::now() + std::time::Duration::from_secs(30), + 30, + ); + let input = b"{}".to_vec(); + let memory_pool: Arc = + Arc::new(GreedyMemoryPool::new(input.len() * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER)); + let (parse_started_tx, parse_started_rx) = tokio::sync::oneshot::channel(); + let (release_parse_tx, release_parse_rx) = std::sync::mpsc::channel(); + let mut output = json_document_ndjson_stream_with_parser( + Box::new(std::io::Cursor::new(input.clone())), + input.len() as u64, + None, + memory_pool, + Some(query_tracker), + move |_, _| { + let _ = parse_started_tx.send(()); + release_parse_rx.recv().expect("release JSON parser"); + Ok(vec![Bytes::from_static(b"{}\n")]) + }, + ); + + { + let next = output.next(); + futures::pin_mut!(next); + assert!(futures::poll!(next.as_mut()).is_pending()); + } + parse_started_rx.await.expect("JSON parser should start"); + drop(output); + + assert!(Arc::clone(&admission).try_acquire_owned().is_err()); + release_parse_tx.send(()).expect("release JSON parser"); + let recovered_permit = + tokio::time::timeout(std::time::Duration::from_secs(5), Arc::clone(&admission).acquire_owned()) + .await + .expect("started JSON parse should release the query guard") + .expect("query admission should remain open"); + drop(recovered_permit); + assert_eq!(admission.available_permits(), 1); + }); + } + /// A JSON array is split into one NDJSON line per element. #[test] fn test_flatten_array_produces_one_line_per_element() { diff --git a/crates/s3select-api/src/query/execution.rs b/crates/s3select-api/src/query/execution.rs index aa350c5f2..935b415d2 100644 --- a/crates/s3select-api/src/query/execution.rs +++ b/crates/s3select-api/src/query/execution.rs @@ -30,7 +30,7 @@ use crate::{QueryError, QueryResult}; use super::Query; use super::logical_planner::Plan; -use super::session::SessionCtx; +use super::session::{QueryExecutionTracker, SessionCtx}; pub struct PhaseTimer { phase_name: &'static str, @@ -172,6 +172,7 @@ pub struct QueryStateMachine { pub session: SessionCtx, pub query: Query, + query_tracker: Option, state: RwLock, start: Instant, } @@ -195,11 +196,31 @@ impl QueryStateMachine { Self { session, query, + query_tracker: None, state: RwLock::new(QueryState::ACCEPTING), start: Instant::now(), } } + pub fn begin_tracked(query: Query, session: SessionCtx, query_tracker: QueryExecutionTracker) -> QueryResult { + if !session.is_bound_to(&query_tracker) { + return Err(QueryError::Cancel); + } + let mut state_machine = Self::begin(query, session); + state_machine.query_tracker = Some(query_tracker); + Ok(state_machine) + } + + pub fn query_tracker(&self) -> Option<&QueryExecutionTracker> { + self.query_tracker.as_ref() + } + + pub fn tracker_matches_session(&self) -> bool { + self.query_tracker + .as_ref() + .is_some_and(|query_tracker| self.session.is_bound_to(query_tracker)) + } + pub fn begin_analyze(&self) { self.record_phase_timestamp("analyze", "start"); self.translate_to(QueryState::RUNNING(RUNNING::ANALYZING)); diff --git a/crates/s3select-api/src/query/session.rs b/crates/s3select-api/src/query/session.rs index 3a5879e3a..10c7e5f15 100644 --- a/crates/s3select-api/src/query/session.rs +++ b/crates/s3select-api/src/query/session.rs @@ -13,7 +13,7 @@ // limitations under the License. use crate::query::Context; -use crate::{QueryError, QueryResult, object_store::EcObjectStore}; +use crate::{QueryError, QueryResult, SelectStore, object_store::EcObjectStore}; use datafusion::{ arrow::{ array::{Int32Array, StringArray}, @@ -25,19 +25,237 @@ use datafusion::{ parquet::arrow::ArrowWriter, prelude::SessionContext, }; -use std::sync::Arc; +use parking_lot::Mutex; +use std::sync::{ + Arc, Weak, + atomic::{AtomicU8, Ordering}, +}; +use tokio::{ + sync::OwnedSemaphorePermit, + task::AbortHandle, + time::{Instant, sleep_until}, +}; use tracing::error; +pub type QueryExecutionGuard = Arc; + +#[derive(Clone, Default)] +pub struct QueryExecutionOwner { + identity: Arc<()>, +} + +impl QueryExecutionOwner { + pub fn new() -> Self { + Self { identity: Arc::new(()) } + } +} + +#[derive(Clone)] +pub struct QueryExecutionTracker { + inner: Arc, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum QueryExecutionStatus { + Active, + Finished, + TimedOut, +} + +const EXECUTION_SETTING_UP: u8 = 0; +const EXECUTION_ADMITTED: u8 = 1; +const EXECUTION_PLANNING: u8 = 2; +const EXECUTION_PLANNED: u8 = 3; +const EXECUTION_STARTING: u8 = 4; +const EXECUTION_RUNNING: u8 = 5; +const EXECUTION_FINISHED: u8 = 6; +const EXECUTION_TIMED_OUT: u8 = 7; + +struct QueryExecutionTrackerInner { + owner_identity: Arc<()>, + query_guard: Mutex>, + deadline: Instant, + timeout_seconds: u64, + state: AtomicU8, + deadline_task: Mutex>, +} + +impl QueryExecutionTracker { + pub fn new(owner: &QueryExecutionOwner, query_guard: QueryExecutionGuard, deadline: Instant, timeout_seconds: u64) -> Self { + let inner = Arc::new(QueryExecutionTrackerInner { + owner_identity: Arc::clone(&owner.identity), + query_guard: Mutex::new(Some(query_guard)), + deadline, + timeout_seconds, + state: AtomicU8::new(EXECUTION_SETTING_UP), + deadline_task: Mutex::new(None), + }); + let deadline_inner = Arc::downgrade(&inner); + let deadline_task = tokio::spawn(async move { + sleep_until(deadline).await; + if let Some(inner) = Weak::upgrade(&deadline_inner) { + inner.expire_at_deadline(); + } + }); + *inner.deadline_task.lock() = Some(deadline_task.abort_handle()); + + Self { inner } + } + + pub fn deadline(&self) -> Instant { + self.inner.deadline + } + + pub fn timeout_seconds(&self) -> u64 { + self.inner.timeout_seconds + } + + pub fn status(&self) -> QueryExecutionStatus { + match self.inner.state.load(Ordering::Acquire) { + EXECUTION_TIMED_OUT => QueryExecutionStatus::TimedOut, + state if state < EXECUTION_FINISHED => QueryExecutionStatus::Active, + _ => QueryExecutionStatus::Finished, + } + } + + pub fn is_owned_by(&self, owner: &QueryExecutionOwner) -> bool { + Arc::ptr_eq(&self.inner.owner_identity, &owner.identity) + } + + pub(crate) fn is_same_execution(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.inner, &other.inner) + } + + pub fn mark_admitted(&self, owner: &QueryExecutionOwner) -> bool { + self.transition(owner, EXECUTION_SETTING_UP, EXECUTION_ADMITTED) + } + + pub fn claim_planning(&self, owner: &QueryExecutionOwner) -> bool { + self.transition(owner, EXECUTION_ADMITTED, EXECUTION_PLANNING) + } + + pub fn mark_planned(&self, owner: &QueryExecutionOwner) -> bool { + self.transition(owner, EXECUTION_PLANNING, EXECUTION_PLANNED) + } + + pub fn claim_execution(&self, owner: &QueryExecutionOwner) -> bool { + self.transition(owner, EXECUTION_PLANNED, EXECUTION_STARTING) + } + + pub fn mark_running(&self, owner: &QueryExecutionOwner) -> bool { + self.transition(owner, EXECUTION_STARTING, EXECUTION_RUNNING) + } + + pub fn handoff_deadline(&self, owner: &QueryExecutionOwner) { + if !self.is_owned_by(owner) || self.inner.state.load(Ordering::Acquire) != EXECUTION_RUNNING { + return; + } + if let Some(deadline_task) = self.inner.deadline_task.lock().take() { + deadline_task.abort(); + } + } + + pub fn finish(&self, owner: &QueryExecutionOwner) { + if self.is_owned_by(owner) { + self.inner.finish(); + } + } + + pub fn expire(&self, owner: &QueryExecutionOwner) { + if self.is_owned_by(owner) { + self.inner.expire(); + } + } + + pub(crate) fn query_guard(&self) -> Option { + let query_guard = self.inner.query_guard.lock(); + if Instant::now() >= self.inner.deadline || self.status() != QueryExecutionStatus::Active { + return None; + } + query_guard.clone() + } + + fn transition(&self, owner: &QueryExecutionOwner, from: u8, to: u8) -> bool { + self.is_owned_by(owner) + && self + .inner + .state + .compare_exchange(from, to, Ordering::AcqRel, Ordering::Acquire) + .is_ok() + } +} + +impl QueryExecutionTrackerInner { + fn finish(&self) { + self.state.fetch_max(EXECUTION_FINISHED, Ordering::AcqRel); + self.release(); + } + + fn expire(&self) { + self.mark_timed_out(); + self.release(); + } + + fn expire_at_deadline(&self) { + if self + .mark_timed_out() + .is_some_and(|state| matches!(state, EXECUTION_ADMITTED | EXECUTION_PLANNED)) + { + self.release(); + } + } + + fn mark_timed_out(&self) -> Option { + self.state + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |state| { + (state < EXECUTION_FINISHED).then_some(EXECUTION_TIMED_OUT) + }) + .ok() + } + + fn release(&self) { + self.query_guard.lock().take(); + if let Some(deadline_task) = self.deadline_task.lock().take() { + deadline_task.abort(); + } + } +} + +impl Drop for QueryExecutionTrackerInner { + fn drop(&mut self) { + if let Some(deadline_task) = self.deadline_task.get_mut().take() { + deadline_task.abort(); + } + } +} + +impl std::fmt::Debug for QueryExecutionTracker { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("QueryExecutionTracker") + .field("deadline", &self.deadline()) + .field("timeout_seconds", &self.timeout_seconds()) + .field("status", &self.status()) + .finish_non_exhaustive() + } +} + #[derive(Clone)] pub struct SessionCtx { _desc: Arc, inner: SessionState, + query_tracker: Option, } impl SessionCtx { pub fn inner(&self) -> &SessionState { &self.inner } + + pub(crate) fn is_bound_to(&self, query_tracker: &QueryExecutionTracker) -> bool { + self.query_tracker + .as_ref() + .is_some_and(|bound_tracker| bound_tracker.is_same_execution(query_tracker)) + } } #[derive(Clone)] @@ -45,12 +263,19 @@ pub struct SessionCtxDesc { // maybe we need some info } -#[derive(Default)] pub struct SessionCtxFactory { pub is_test: bool, pub target_partitions: usize, } +pub const DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES: usize = 64 * 1024 * 1024; + +impl Default for SessionCtxFactory { + fn default() -> Self { + Self::new(false) + } +} + impl SessionCtxFactory { pub fn new(is_test: bool) -> Self { Self { @@ -65,19 +290,66 @@ impl SessionCtxFactory { } pub async fn create_session_ctx(&self, context: &Context) -> QueryResult { - let df_session_ctx = self.build_df_session_context(context).await?; + self.create_session_ctx_inner(context, None, None, DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES) + .await + } + + pub async fn create_session_ctx_with_tracker_and_memory_limit( + &self, + context: &Context, + query_tracker: QueryExecutionTracker, + memory_limit_bytes: usize, + ) -> QueryResult { + self.create_session_ctx_inner(context, Some(query_tracker), None, memory_limit_bytes) + .await + } + + #[cfg(test)] + async fn create_session_ctx_with_tracker_and_store( + &self, + context: &Context, + query_tracker: QueryExecutionTracker, + store: Arc, + ) -> QueryResult { + self.create_session_ctx_inner(context, Some(query_tracker), Some(store), DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES) + .await + } + + async fn create_session_ctx_inner( + &self, + context: &Context, + query_tracker: Option, + store: Option>, + memory_limit_bytes: usize, + ) -> QueryResult { + let df_session_ctx = self + .build_df_session_context(context, query_tracker.clone(), store, memory_limit_bytes) + .await?; Ok(SessionCtx { _desc: Arc::new(SessionCtxDesc {}), inner: df_session_ctx.state(), + query_tracker, }) } - async fn build_df_session_context(&self, context: &Context) -> QueryResult { + async fn build_df_session_context( + &self, + context: &Context, + query_tracker: Option, + store: Option>, + memory_limit_bytes: usize, + ) -> QueryResult { let path = format!("s3://{}", context.input.bucket); let store_url = url::Url::parse(&path).unwrap(); - let rt = RuntimeEnvBuilder::new().build()?; + let memory_limit_bytes = if memory_limit_bytes == 0 { + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES + } else { + memory_limit_bytes + }; + let rt = RuntimeEnvBuilder::new().with_memory_limit(memory_limit_bytes, 1.0).build()?; let config = SessionConfig::new().with_target_partitions(self.target_partitions); + let memory_pool = Arc::clone(&rt.memory_pool); let df_session_state = SessionStateBuilder::new() .with_config(config) .with_runtime_env(Arc::new(rt)) @@ -127,8 +399,13 @@ impl SessionCtxFactory { df_session_state.with_object_store(&store_url, store).build() } else { - let store: EcObjectStore = - EcObjectStore::new(context.input.clone()).map_err(|_| QueryError::NotImplemented { err: String::new() })?; + let store: EcObjectStore = match query_tracker { + Some(query_tracker) => { + EcObjectStore::new_with_query_tracker(context.input.clone(), memory_pool, query_tracker, store) + } + None => EcObjectStore::new_with_memory_pool(context.input.clone(), memory_pool), + } + .map_err(|_| QueryError::NotImplemented { err: String::new() })?; df_session_state.with_object_store(&store_url, Arc::new(store)).build() }; @@ -197,6 +474,7 @@ fn test_parquet_batch( #[cfg(test)] mod tests { use super::*; + use datafusion::execution::memory_pool::MemoryLimit; use s3s::dto::{ CSVInput, CSVOutput, ExpressionType, InputSerialization, OutputSerialization, SelectObjectContentInput, SelectObjectContentRequest, @@ -229,6 +507,17 @@ mod tests { } } + #[test] + fn session_factory_fields_remain_source_compatible() { + let factory = SessionCtxFactory { + is_test: true, + target_partitions: 0, + }; + + assert!(factory.is_test); + assert_eq!(factory.target_partitions, 0); + } + #[tokio::test] async fn session_factory_applies_target_partitions() { let factory = SessionCtxFactory::new(true).with_target_partitions(3); @@ -250,4 +539,78 @@ mod tests { assert_eq!(session.inner().config().target_partitions(), SessionConfig::new().target_partitions()); } + + #[tokio::test] + async fn session_factory_applies_memory_limit() { + let factory = SessionCtxFactory::new(true); + let session = factory + .create_session_ctx_inner(&test_context(), None, None, 1024) + .await + .expect("session should be created with a bounded memory pool"); + + assert!(matches!( + session.inner().runtime_env().memory_pool.memory_limit(), + MemoryLimit::Finite(1024) + )); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + #[serial_test::serial] + async fn session_factory_propagates_query_guard_to_ec_store() { + let temp_root = tempfile::tempdir().expect("create session test temp root"); + let env = rustfs_test_utils::TestECStoreEnv::builder() + .base_dir(temp_root.path()) + .init_bucket_metadata(false) + .build() + .await; + + let admission = Arc::new(tokio::sync::Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("query permit should be available"); + let query_guard = Arc::new(permit); + let query_tracker = QueryExecutionTracker::new( + &QueryExecutionOwner::new(), + Arc::clone(&query_guard), + Instant::now() + std::time::Duration::from_secs(300), + 300, + ); + let session = SessionCtxFactory::new(false) + .create_session_ctx_with_tracker_and_store(&test_context(), query_tracker, Arc::clone(&env.ecstore)) + .await + .expect("production session should be created with the query guard"); + + assert!(Arc::strong_count(&query_guard) > 1); + drop(session); + assert_eq!(Arc::strong_count(&query_guard), 1); + } + + #[tokio::test] + async fn elapsed_deadline_does_not_yield_query_guard_before_timer_poll() { + let admission = Arc::new(tokio::sync::Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("query permit should be available"); + let query_tracker = QueryExecutionTracker::new(&QueryExecutionOwner::new(), Arc::new(permit), Instant::now(), 0); + + assert!(query_tracker.query_guard().is_none()); + } + + #[tokio::test] + async fn session_factory_default_uses_bounded_memory() { + let factory = SessionCtxFactory::default(); + let session = SessionCtxFactory::new(true) + .create_session_ctx(&test_context()) + .await + .expect("default session should be created with a bounded memory pool"); + + assert!(!factory.is_test); + assert_eq!(factory.target_partitions, 0); + assert!(matches!( + session.inner().runtime_env().memory_pool.memory_limit(), + MemoryLimit::Finite(DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES) + )); + } } diff --git a/crates/s3select-query/Cargo.toml b/crates/s3select-query/Cargo.toml index 5346b975e..0dfe9ca7c 100644 --- a/crates/s3select-query/Cargo.toml +++ b/crates/s3select-query/Cargo.toml @@ -34,7 +34,7 @@ derive_builder = { workspace = true } futures = { workspace = true } parking_lot = { workspace = true } s3s = { workspace = true, features = ["minio"] } -tokio = { workspace = true, features = ["fs", "rt-multi-thread"] } +tokio = { workspace = true, features = ["fs", "rt-multi-thread", "sync", "time"] } tracing = { workspace = true } [lib] diff --git a/crates/s3select-query/src/dispatcher/manager.rs b/crates/s3select-query/src/dispatcher/manager.rs index 8ee239687..80c7bd584 100644 --- a/crates/s3select-query/src/dispatcher/manager.rs +++ b/crates/s3select-query/src/dispatcher/manager.rs @@ -13,10 +13,12 @@ // limitations under the License. use std::{ + future::Future, ops::Deref, pin::Pin, sync::Arc, task::{Context, Poll}, + time::Duration, }; use async_trait::async_trait; @@ -33,9 +35,10 @@ use datafusion::{ execution::{RecordBatchStream, SendableRecordBatchStream}, sql::sqlparser::parser::ParserError, }; -use futures::{Stream, StreamExt}; +use futures::Stream; +use parking_lot::Mutex; use rustfs_s3select_api::{ - QueryError, QueryResult, + QueryError, QueryResult, S3SelectPolicyError, query::{ Query, ast::ExtStatement, @@ -44,15 +47,23 @@ use rustfs_s3select_api::{ function::FuncMetaManagerRef, logical_planner::{LogicalPlanner, Plan}, parser::Parser, - session::{SessionCtx, SessionCtxFactory}, + session::{ + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, SessionCtx, + SessionCtxFactory, + }, }, }; use s3s::dto::{FileHeaderInfo, SelectObjectContentInput}; use std::sync::LazyLock; +use tokio::{ + sync::Semaphore, + time::{Instant, Sleep, sleep_until, timeout_at}, +}; use crate::{ dispatcher::parquet_table::ParquetSelectTable, execution::factory::QueryExecutionFactoryRef, + instance::{DEFAULT_MAX_CONCURRENT_QUERIES, DEFAULT_QUERY_TIMEOUT_SECS}, metadata::{ContextProviderExtension, MetadataProvider, TableHandleProviderRef, base_table::BaseTableProvider}, sql::logical::planner::DefaultLogicalPlanner, }; @@ -72,68 +83,229 @@ pub struct SimpleQueryDispatcher { // get query execution factory query_execution_factory: QueryExecutionFactoryRef, func_manager: FuncMetaManagerRef, + memory_limit_bytes: usize, + query_admission: Arc, + query_timeout: Duration, + query_execution_owner: QueryExecutionOwner, +} + +struct QueryPhaseGuard<'a> { + query_tracker: &'a QueryExecutionTracker, + query_execution_owner: &'a QueryExecutionOwner, + armed: bool, +} + +impl<'a> QueryPhaseGuard<'a> { + fn new(query_tracker: &'a QueryExecutionTracker, query_execution_owner: &'a QueryExecutionOwner) -> Self { + Self { + query_tracker, + query_execution_owner, + armed: true, + } + } + + fn disarm(mut self) { + self.armed = false; + } +} + +impl Drop for QueryPhaseGuard<'_> { + fn drop(&mut self) { + if self.armed { + self.query_tracker.finish(self.query_execution_owner); + } + } } #[async_trait] impl QueryDispatcher for SimpleQueryDispatcher { async fn execute_query(&self, query: &Query) -> QueryResult { - let query_state_machine = { self.build_query_state_machine(query.clone()).await? }; - - let logical_plan = self.build_logical_plan(query_state_machine.clone()).await?; - let logical_plan = match logical_plan { - Some(plan) => plan, - None => return Ok(Output::Nil(())), + let query_state_machine = self.build_query_state_machine(query.clone()).await?; + let logical_plan = self.build_logical_plan(Arc::clone(&query_state_machine)).await?; + let Some(logical_plan) = logical_plan else { + return Ok(Output::Nil(())); }; - let result = self.execute_logical_plan(logical_plan, query_state_machine).await?; - Ok(result) + + self.execute_logical_plan(logical_plan, query_state_machine).await } async fn build_logical_plan(&self, query_state_machine: Arc) -> QueryResult> { - let session = &query_state_machine.session; - let query = &query_state_machine.query; - - let scheme_provider = self.build_scheme_provider(session).await?; - - let logical_planner = DefaultLogicalPlanner::new(&scheme_provider); - - let statements = self.parser.parse(query.content())?; - - // not allow multi statement - if statements.len() > 1 { - return Err(QueryError::MultiStatement { - num: statements.len(), - sql: query_state_machine.query.content().to_string(), - }); + if !query_state_machine.tracker_matches_session() { + return Err(QueryError::Cancel); } - - let stmt = match statements.front() { - Some(stmt) => stmt.clone(), - None => { - return Err(QueryError::Parser { - source: ParserError::ParserError("empty SQL expression".to_string()), - }); - } - }; - + let query_tracker = query_state_machine.query_tracker().cloned().ok_or(QueryError::Cancel)?; + if !query_tracker.claim_planning(&self.query_execution_owner) { + return Err(self.query_tracker_error(&query_tracker)); + } + let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner); let logical_plan = self - .statement_to_logical_plan(stmt, &logical_planner, query_state_machine) + .run_with_query_deadline(&query_tracker, async { + let session = &query_state_machine.session; + let query = &query_state_machine.query; + + let scheme_provider = self.build_scheme_provider(session).await?; + let logical_planner = DefaultLogicalPlanner::new(&scheme_provider); + let statements = self.parser.parse(query.content())?; + + if statements.len() > 1 { + return Err(QueryError::MultiStatement { + num: statements.len(), + sql: query_state_machine.query.content().to_string(), + }); + } + + let stmt = match statements.front() { + Some(stmt) => stmt.clone(), + None => { + return Err(QueryError::Parser { + source: ParserError::ParserError("empty SQL expression".to_string()), + }); + } + }; + + let logical_plan = self + .statement_to_logical_plan(stmt, &logical_planner, query_state_machine) + .await?; + Ok(logical_plan) + }) .await?; + if !query_tracker.mark_planned(&self.query_execution_owner) { + drop(logical_plan); + return Err(self.query_tracker_error(&query_tracker)); + } + phase_guard.disarm(); Ok(Some(logical_plan)) } async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc) -> QueryResult { - self.execute_logical_plan(logical_plan, query_state_machine).await + if !query_state_machine.tracker_matches_session() { + return Err(QueryError::Cancel); + } + let query_tracker = query_state_machine.query_tracker().cloned().ok_or(QueryError::Cancel)?; + if !query_tracker.claim_execution(&self.query_execution_owner) { + return Err(self.query_tracker_error(&query_tracker)); + } + let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner); + let output = self + .run_with_query_deadline(&query_tracker, self.start_logical_plan(logical_plan, query_state_machine)) + .await?; + + match output { + Output::StreamData(stream) => { + if !query_tracker.mark_running(&self.query_execution_owner) { + drop(stream); + return Err(self.query_tracker_error(&query_tracker)); + } + let stream = TrackedRecordBatchStream::new(stream, query_tracker.clone(), self.query_execution_owner.clone()); + phase_guard.disarm(); + Ok(Output::StreamData(Box::pin(stream))) + } + Output::Nil(()) => Ok(Output::Nil(())), + } } async fn build_query_state_machine(&self, query: Query) -> QueryResult> { - let session = self.session_factory.create_session_ctx(query.context()).await?; - - let query_state_machine = Arc::new(QueryStateMachine::begin(query, session)); - Ok(query_state_machine) + let permit = self + .query_admission + .clone() + .try_acquire_owned() + .map_err(|_| QueryError::from(S3SelectPolicyError::QueryConcurrencyLimit))?; + let query_tracker = QueryExecutionTracker::new( + &self.query_execution_owner, + Arc::new(permit), + Instant::now() + self.query_timeout, + self.query_timeout.as_secs(), + ); + let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner); + let session = self + .run_with_query_deadline( + &query_tracker, + self.session_factory.create_session_ctx_with_tracker_and_memory_limit( + query.context(), + query_tracker.clone(), + self.memory_limit_bytes, + ), + ) + .await?; + if !query_tracker.mark_admitted(&self.query_execution_owner) { + drop(session); + return Err(self.query_tracker_error(&query_tracker)); + } + phase_guard.disarm(); + Ok(Arc::new(QueryStateMachine::begin_tracked(query, session, query_tracker)?)) } } impl SimpleQueryDispatcher { + async fn run_with_query_deadline( + &self, + query_tracker: &QueryExecutionTracker, + future: impl Future>, + ) -> QueryResult { + let deadline = query_tracker.deadline(); + let timeout_error = || { + S3SelectPolicyError::QueryTimeout { + seconds: query_tracker.timeout_seconds(), + } + .into() + }; + match query_tracker.status() { + QueryExecutionStatus::TimedOut => { + drop(future); + query_tracker.expire(&self.query_execution_owner); + return Err(timeout_error()); + } + QueryExecutionStatus::Finished => return Err(QueryError::Cancel), + QueryExecutionStatus::Active if Instant::now() >= deadline => { + drop(future); + query_tracker.expire(&self.query_execution_owner); + return Err(timeout_error()); + } + QueryExecutionStatus::Active => {} + } + + match timeout_at(deadline, future).await { + Ok(result) => match query_tracker.status() { + QueryExecutionStatus::TimedOut => { + drop(result); + query_tracker.expire(&self.query_execution_owner); + Err(timeout_error()) + } + QueryExecutionStatus::Finished => { + drop(result); + Err(QueryError::Cancel) + } + QueryExecutionStatus::Active if Instant::now() >= deadline => { + drop(result); + query_tracker.expire(&self.query_execution_owner); + Err(timeout_error()) + } + QueryExecutionStatus::Active => result, + }, + Err(_) => { + query_tracker.expire(&self.query_execution_owner); + Err(timeout_error()) + } + } + } + + fn query_tracker_error(&self, query_tracker: &QueryExecutionTracker) -> QueryError { + if !query_tracker.is_owned_by(&self.query_execution_owner) { + return QueryError::Cancel; + } + match query_tracker.status() { + QueryExecutionStatus::TimedOut => S3SelectPolicyError::QueryTimeout { + seconds: query_tracker.timeout_seconds(), + } + .into(), + QueryExecutionStatus::Active if Instant::now() >= query_tracker.deadline() => S3SelectPolicyError::QueryTimeout { + seconds: query_tracker.timeout_seconds(), + } + .into(), + QueryExecutionStatus::Active | QueryExecutionStatus::Finished => QueryError::Cancel, + } + } + async fn statement_to_logical_plan( &self, stmt: ExtStatement, @@ -150,17 +322,13 @@ impl SimpleQueryDispatcher { Ok(logical_plan) } - async fn execute_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc) -> QueryResult { + async fn start_logical_plan(&self, logical_plan: Plan, query_state_machine: Arc) -> QueryResult { let execution = self .query_execution_factory .create_query_execution(logical_plan, query_state_machine.clone()) .await?; - match execution.start().await { - Ok(Output::StreamData(stream)) => Ok(Output::StreamData(Box::pin(TrackedRecordBatchStream { inner: stream }))), - Ok(nil @ Output::Nil(_)) => Ok(nil), - Err(err) => Err(err), - } + execution.start().await } async fn build_scheme_provider(&self, session: &SessionCtx) -> QueryResult { @@ -287,12 +455,70 @@ impl SimpleQueryDispatcher { } pub struct TrackedRecordBatchStream { - inner: SendableRecordBatchStream, + state: Arc, + schema: SchemaRef, + deadline: Pin>, + deadline_task: tokio::task::JoinHandle<()>, + done: bool, +} + +struct TrackedRecordBatchState { + inner: Mutex>, + query_tracker: QueryExecutionTracker, + query_execution_owner: QueryExecutionOwner, +} + +impl TrackedRecordBatchState { + fn finish(&self) { + self.inner.lock().take(); + self.query_tracker.finish(&self.query_execution_owner); + } + + fn expire(&self) { + self.inner.lock().take(); + self.query_tracker.expire(&self.query_execution_owner); + } +} + +impl TrackedRecordBatchStream { + fn new( + inner: SendableRecordBatchStream, + query_tracker: QueryExecutionTracker, + query_execution_owner: QueryExecutionOwner, + ) -> Self { + let schema = inner.schema(); + let deadline = query_tracker.deadline(); + let state = Arc::new(TrackedRecordBatchState { + inner: Mutex::new(Some(inner)), + query_tracker, + query_execution_owner, + }); + let deadline_state = Arc::clone(&state); + let deadline_task = tokio::spawn(async move { + sleep_until(deadline).await; + deadline_state.expire(); + }); + state.query_tracker.handoff_deadline(&state.query_execution_owner); + Self { + state, + schema, + deadline: Box::pin(sleep_until(deadline)), + deadline_task, + done: false, + } + } +} + +impl Drop for TrackedRecordBatchStream { + fn drop(&mut self) { + self.state.finish(); + self.deadline_task.abort(); + } } impl RecordBatchStream for TrackedRecordBatchStream { fn schema(&self) -> SchemaRef { - self.inner.schema() + Arc::clone(&self.schema) } } @@ -300,10 +526,77 @@ impl Stream for TrackedRecordBatchStream { type Item = DFResult; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - self.inner.poll_next_unpin(cx) + if self.done { + return Poll::Ready(None); + } + let deadline_at = self.state.query_tracker.deadline(); + let timeout_seconds = self.state.query_tracker.timeout_seconds(); + match self.state.query_tracker.status() { + QueryExecutionStatus::TimedOut => { + self.done = true; + self.state.expire(); + self.deadline_task.abort(); + return Poll::Ready(Some(Err(query_timeout_error(timeout_seconds)))); + } + QueryExecutionStatus::Finished => { + self.done = true; + self.state.finish(); + self.deadline_task.abort(); + return Poll::Ready(Some(Err(query_cancelled_error()))); + } + QueryExecutionStatus::Active => {} + } + if self.deadline.as_mut().poll(cx).is_ready() { + self.done = true; + self.state.expire(); + self.deadline_task.abort(); + return Poll::Ready(Some(Err(query_timeout_error(timeout_seconds)))); + } + let mut inner = self.state.inner.lock(); + let poll = match inner.as_mut() { + Some(inner) => inner.as_mut().poll_next(cx), + None => Poll::Ready(None), + }; + let status = self.state.query_tracker.status(); + if status != QueryExecutionStatus::Active || Instant::now() >= deadline_at { + drop(poll); + inner.take(); + drop(inner); + self.done = true; + match status { + QueryExecutionStatus::TimedOut | QueryExecutionStatus::Active => { + self.state.query_tracker.expire(&self.state.query_execution_owner); + } + QueryExecutionStatus::Finished => { + self.state.query_tracker.finish(&self.state.query_execution_owner); + } + } + self.deadline_task.abort(); + return Poll::Ready(Some(Err(match status { + QueryExecutionStatus::TimedOut | QueryExecutionStatus::Active => query_timeout_error(timeout_seconds), + QueryExecutionStatus::Finished => query_cancelled_error(), + }))); + } + drop(inner); + if matches!(poll, Poll::Ready(None)) { + self.done = true; + self.state.finish(); + self.deadline_task.abort(); + } + poll } } +fn query_timeout_error(timeout_seconds: u64) -> datafusion::common::DataFusionError { + datafusion::common::DataFusionError::External(Box::new(S3SelectPolicyError::QueryTimeout { + seconds: timeout_seconds, + })) +} + +fn query_cancelled_error() -> datafusion::common::DataFusionError { + datafusion::common::DataFusionError::External(Box::new(QueryError::Cancel)) +} + #[derive(Default, Clone)] pub struct SimpleQueryDispatcherBuilder { input: Option>, @@ -314,6 +607,9 @@ pub struct SimpleQueryDispatcherBuilder { query_execution_factory: Option, func_manager: Option, + memory_limit_bytes: Option, + query_admission: Option>, + query_timeout: Option, } impl SimpleQueryDispatcherBuilder { @@ -346,6 +642,23 @@ impl SimpleQueryDispatcherBuilder { self } + pub fn with_memory_limit_bytes(mut self, memory_limit_bytes: usize) -> Self { + if memory_limit_bytes > 0 { + self.memory_limit_bytes = Some(memory_limit_bytes); + } + self + } + + pub fn with_query_admission(mut self, query_admission: Arc) -> Self { + self.query_admission = Some(query_admission); + self + } + + pub fn with_query_timeout(mut self, query_timeout: Duration) -> Self { + self.query_timeout = Some(query_timeout); + self + } + pub fn build(self) -> QueryResult> { let input = self.input.ok_or_else(|| QueryError::BuildQueryDispatcher { err: "lost of input".to_string(), @@ -371,6 +684,14 @@ impl SimpleQueryDispatcherBuilder { err: "lost of default_table_provider".to_string(), })?; + let memory_limit_bytes = self.memory_limit_bytes.unwrap_or(DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES); + let query_admission = self + .query_admission + .unwrap_or_else(|| Arc::new(Semaphore::new(DEFAULT_MAX_CONCURRENT_QUERIES))); + let query_timeout = self + .query_timeout + .unwrap_or_else(|| Duration::from_secs(DEFAULT_QUERY_TIMEOUT_SECS)); + let dispatcher = Arc::new(SimpleQueryDispatcher { input, _default_table_provider: default_table_provider, @@ -378,8 +699,1068 @@ impl SimpleQueryDispatcherBuilder { parser, query_execution_factory, func_manager, + memory_limit_bytes, + query_admission, + query_timeout, + query_execution_owner: QueryExecutionOwner::new(), }); Ok(dispatcher) } } + +#[cfg(test)] +mod tests { + use super::{QueryPhaseGuard, SimpleQueryDispatcher, SimpleQueryDispatcherBuilder, TrackedRecordBatchStream}; + use crate::{ + execution::{ + factory::{QueryExecutionFactoryRef, SqlQueryExecutionFactory}, + scheduler::local::LocalScheduler, + }, + function::simple_func_manager::SimpleFunctionMetadataManager, + instance::{DEFAULT_MAX_CONCURRENT_QUERIES, DEFAULT_QUERY_TIMEOUT_SECS}, + metadata::base_table::BaseTableProvider, + sql::{optimizer::CascadeOptimizerBuilder, parser::DefaultParser}, + }; + use async_trait::async_trait; + use datafusion::{ + arrow::{ + datatypes::{Schema, SchemaRef}, + record_batch::RecordBatch, + }, + common::DataFusionError, + physical_plan::{RecordBatchStream, stream::RecordBatchStreamAdapter}, + }; + use futures::{StreamExt, stream}; + use rustfs_s3select_api::{ + QueryError, QueryResult, S3SelectPolicyError, + query::{ + Context as QueryContext, Query, + dispatcher::QueryDispatcher, + execution::{ + Output, QueryExecution, QueryExecutionFactory, QueryExecutionRef, QueryStateMachine, QueryStateMachineRef, + }, + logical_planner::Plan, + session::{ + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryExecutionOwner, QueryExecutionStatus, QueryExecutionTracker, + SessionCtxFactory, + }, + }, + }; + use s3s::dto::{ + CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, OutputSerialization, SelectObjectContentInput, + SelectObjectContentRequest, + }; + use std::{ + pin::Pin, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + task::Poll, + time::Duration, + }; + use tokio::{ + sync::{Barrier, Semaphore}, + time::Instant, + }; + + async fn wait_for_query_timeout(query_tracker: &QueryExecutionTracker) { + tokio::time::timeout(Duration::from_secs(1), async { + while query_tracker.status() != QueryExecutionStatus::TimedOut { + tokio::task::yield_now().await; + } + }) + .await + .expect("query deadline should expire"); + } + + struct DropSignal(Arc); + + impl Drop for DropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } + } + + struct PanickingSchemaStream { + dropped: Arc, + _drop_guard: BlockingDrop, + } + + impl futures::Stream for PanickingSchemaStream { + type Item = Result; + + fn poll_next(self: Pin<&mut Self>, _cx: &mut std::task::Context<'_>) -> Poll> { + Poll::Pending + } + } + + impl RecordBatchStream for PanickingSchemaStream { + fn schema(&self) -> SchemaRef { + panic!("test stream schema panic"); + } + } + + struct PanickingSchemaQueryExecutionFactory { + dropped: Arc, + drop_guard: std::sync::Mutex>, + } + + #[async_trait] + impl QueryExecutionFactory for PanickingSchemaQueryExecutionFactory { + async fn create_query_execution( + &self, + _plan: Plan, + _query_state_machine: QueryStateMachineRef, + ) -> QueryResult { + let drop_guard = self + .drop_guard + .lock() + .expect("panic factory mutex should not be poisoned") + .take() + .expect("panic factory should be called once"); + Ok(Arc::new(PanickingSchemaQueryExecution { + dropped: Arc::clone(&self.dropped), + drop_guard: std::sync::Mutex::new(Some(drop_guard)), + })) + } + } + + struct PanickingSchemaQueryExecution { + dropped: Arc, + drop_guard: std::sync::Mutex>, + } + + #[async_trait] + impl QueryExecution for PanickingSchemaQueryExecution { + async fn start(&self) -> QueryResult { + let drop_guard = self + .drop_guard + .lock() + .expect("panic execution mutex should not be poisoned") + .take() + .expect("panic execution should start once"); + Ok(Output::StreamData(Box::pin(PanickingSchemaStream { + dropped: Arc::clone(&self.dropped), + _drop_guard: drop_guard, + }))) + } + + fn cancel(&self) -> QueryResult<()> { + Ok(()) + } + } + + impl Drop for PanickingSchemaStream { + fn drop(&mut self) { + self.dropped.store(true, Ordering::SeqCst); + } + } + + struct BlockingDrop { + started: std::sync::mpsc::Sender<()>, + release: std::sync::mpsc::Receiver<()>, + } + + impl Drop for BlockingDrop { + fn drop(&mut self) { + let _ = self.started.send(()); + self.release.recv().expect("release blocking drop"); + } + } + + struct BlockingError { + started: std::sync::mpsc::Sender<()>, + release: std::sync::Mutex>, + } + + impl Drop for BlockingError { + fn drop(&mut self) { + let _ = self.started.send(()); + self.release + .lock() + .expect("blocking error mutex should not be poisoned") + .recv() + .expect("release blocking error"); + } + } + + impl std::fmt::Debug for BlockingError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("BlockingError").finish_non_exhaustive() + } + } + + impl std::fmt::Display for BlockingError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("blocking drop test error") + } + } + + impl std::error::Error for BlockingError {} + + struct PendingQueryExecutionFactory; + + #[async_trait] + impl QueryExecutionFactory for PendingQueryExecutionFactory { + async fn create_query_execution( + &self, + _plan: Plan, + _query_state_machine: QueryStateMachineRef, + ) -> QueryResult { + std::future::pending().await + } + } + + struct DropBlockingPendingQueryExecutionFactory { + drop_guard: std::sync::Mutex>, + } + + #[async_trait] + impl QueryExecutionFactory for DropBlockingPendingQueryExecutionFactory { + async fn create_query_execution( + &self, + _plan: Plan, + _query_state_machine: QueryStateMachineRef, + ) -> QueryResult { + let _drop_guard = self + .drop_guard + .lock() + .expect("pending factory mutex should not be poisoned") + .take() + .expect("pending factory should be called once"); + std::future::pending().await + } + } + + fn test_input() -> SelectObjectContentInput { + SelectObjectContentInput { + bucket: "test-bucket".to_string(), + expected_bucket_owner: None, + key: "test.csv".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 { + csv: Some(CSVInput { + file_header_info: Some(FileHeaderInfo::from_static(FileHeaderInfo::USE)), + ..Default::default() + }), + ..Default::default() + }, + output_serialization: OutputSerialization { + csv: Some(CSVOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + } + } + + fn test_dispatcher( + admission: Arc, + query_timeout: Duration, + ) -> (Arc, Arc) { + let optimizer = Arc::new(CascadeOptimizerBuilder::default().build()); + let scheduler = Arc::new(LocalScheduler {}); + test_dispatcher_with_factory(admission, query_timeout, Arc::new(SqlQueryExecutionFactory::new(optimizer, scheduler))) + } + + fn test_dispatcher_with_factory( + admission: Arc, + query_timeout: Duration, + query_execution_factory: QueryExecutionFactoryRef, + ) -> (Arc, Arc) { + let input = Arc::new(test_input()); + 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(query_execution_factory) + .with_func_manager(Arc::new(SimpleFunctionMetadataManager::default())) + .with_query_admission(admission) + .with_query_timeout(query_timeout) + .build() + .expect("query dispatcher should build"); + (dispatcher, input) + } + + fn test_query_tracker( + permit: tokio::sync::OwnedSemaphorePermit, + deadline: Instant, + timeout_seconds: u64, + ) -> (QueryExecutionOwner, QueryExecutionTracker) { + let owner = QueryExecutionOwner::new(); + let query_tracker = QueryExecutionTracker::new(&owner, Arc::new(permit), deadline, timeout_seconds); + assert!(query_tracker.mark_admitted(&owner)); + assert!(query_tracker.claim_planning(&owner)); + assert!(query_tracker.mark_planned(&owner)); + assert!(query_tracker.claim_execution(&owner)); + assert!(query_tracker.mark_running(&owner)); + (owner, query_tracker) + } + + #[test] + fn builder_uses_query_limit_defaults_when_omitted() { + let optimizer = Arc::new(CascadeOptimizerBuilder::default().build()); + let scheduler = Arc::new(LocalScheduler {}); + let dispatcher = SimpleQueryDispatcherBuilder::default() + .with_input(Arc::new(test_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("legacy builder chain should use default query limits"); + + assert_eq!(dispatcher.query_admission.available_permits(), DEFAULT_MAX_CONCURRENT_QUERIES); + assert_eq!(dispatcher.query_timeout, Duration::from_secs(DEFAULT_QUERY_TIMEOUT_SECS)); + assert_eq!(dispatcher.memory_limit_bytes, DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES); + } + + #[tokio::test] + async fn rejects_query_when_admission_is_saturated() { + let admission = Arc::new(Semaphore::new(1)); + let _held_permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let (dispatcher, input) = test_dispatcher(admission, Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + + let result = dispatcher.execute_query(&query).await; + + assert!(matches!( + result, + Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryConcurrencyLimit)) + )); + } + + #[tokio::test] + async fn staged_query_rejects_when_admission_is_saturated() { + let admission = Arc::new(Semaphore::new(1)); + let _held_permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let (dispatcher, input) = test_dispatcher(admission, Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + + let result = dispatcher.build_query_state_machine(query).await; + + assert!(matches!( + result, + Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryConcurrencyLimit)) + )); + } + + #[tokio::test] + async fn untracked_query_state_machine_cannot_bypass_limits() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(admission, Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let session = SessionCtxFactory::new(true) + .create_session_ctx(query.context()) + .await + .expect("untracked session should be available for compatibility"); + let query_state_machine = Arc::new(QueryStateMachine::begin(query, session)); + + let result = dispatcher.build_logical_plan(query_state_machine).await; + + assert!(matches!(result, Err(QueryError::Cancel))); + } + + #[tokio::test] + async fn staged_query_rejects_unbound_session() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let untracked_session = SessionCtxFactory::new(true) + .create_session_ctx(query.context()) + .await + .expect("untracked session should be available"); + let mut query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let query_tracker = query_state_machine + .query_tracker() + .cloned() + .expect("staged query should retain its tracker"); + assert!(matches!( + QueryStateMachine::begin_tracked(query_state_machine.query.clone(), untracked_session.clone(), query_tracker), + Err(QueryError::Cancel) + )); + Arc::get_mut(&mut query_state_machine) + .expect("test should hold the only state machine reference") + .session = untracked_session; + + assert!(matches!( + dispatcher.build_logical_plan(query_state_machine).await, + Err(QueryError::Cancel) + )); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn staged_execution_rejects_session_replaced_after_planning() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let untracked_session = SessionCtxFactory::new(true) + .create_session_ctx(query.context()) + .await + .expect("untracked session should be available"); + let mut query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("staged query should build a logical plan") + .expect("select query should produce a logical plan"); + Arc::get_mut(&mut query_state_machine) + .expect("test should hold the only state machine reference") + .session = untracked_session; + + assert!(matches!( + dispatcher.execute_logical_plan(logical_plan, query_state_machine).await, + Err(QueryError::Cancel) + )); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn times_out_during_query_setup() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher_with_factory( + Arc::clone(&admission), + Duration::from_millis(1), + Arc::new(PendingQueryExecutionFactory), + ); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + + let result = tokio::time::timeout(Duration::from_secs(1), dispatcher.execute_query(&query)) + .await + .expect("dispatcher should enforce its query setup timeout"); + + assert!(matches!( + result, + Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::QueryTimeout { seconds: 0 })) + )); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn staged_query_execution_rejects_expired_tracker() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher_with_factory( + Arc::clone(&admission), + Duration::from_secs(300), + Arc::new(PendingQueryExecutionFactory), + ); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("staged query should build a logical plan") + .expect("select query should produce a logical plan"); + query_state_machine + .query_tracker() + .expect("staged query should retain its tracker") + .expire(&dispatcher.query_execution_owner); + + 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 })) + )); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn staged_query_stream_owns_shared_admission_tracker() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let retained_state_machine = Arc::clone(&query_state_machine); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("staged query should build a logical plan") + .expect("select query should produce a logical plan"); + let output = dispatcher + .execute_logical_plan(logical_plan, query_state_machine) + .await + .expect("staged query should start execution"); + + assert!(Arc::clone(&admission).try_acquire_owned().is_err()); + drop(output); + assert_eq!(admission.available_permits(), 1); + assert!(matches!( + dispatcher.build_logical_plan(retained_state_machine).await, + Err(QueryError::Cancel) + )); + } + + #[tokio::test] + async fn staged_query_rejects_duplicate_phase_calls() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("staged query should build a logical plan") + .expect("select query should produce a logical plan"); + + assert!(matches!( + dispatcher.build_logical_plan(Arc::clone(&query_state_machine)).await, + Err(QueryError::Cancel) + )); + + let duplicate_plan = logical_plan.clone(); + let output = dispatcher + .execute_logical_plan(logical_plan, Arc::clone(&query_state_machine)) + .await + .expect("first staged execution should start"); + assert!(matches!( + dispatcher.execute_logical_plan(duplicate_plan, query_state_machine).await, + Err(QueryError::Cancel) + )); + assert_eq!(admission.available_permits(), 0); + + drop(output); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] + async fn concurrent_planning_claim_has_single_winner() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let owner = QueryExecutionOwner::new(); + let query_tracker = QueryExecutionTracker::new(&owner, Arc::new(permit), Instant::now() + Duration::from_secs(300), 300); + assert!(query_tracker.mark_admitted(&owner)); + const CONTENDERS: usize = 16; + let barrier = Arc::new(Barrier::new(CONTENDERS + 1)); + let mut claims = Vec::with_capacity(CONTENDERS); + for _ in 0..CONTENDERS { + let task_barrier = Arc::clone(&barrier); + let task_tracker = query_tracker.clone(); + let task_owner = owner.clone(); + claims.push(tokio::spawn(async move { + task_barrier.wait().await; + task_tracker.claim_planning(&task_owner) + })); + } + barrier.wait().await; + + let mut successful_claims = 0; + for claim in claims { + successful_claims += usize::from(claim.await.expect("planning claim task should finish")); + } + assert_eq!(successful_claims, 1); + query_tracker.finish(&owner); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn staged_query_deadline_releases_retained_admission() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission semaphore should remain open"); + let query_tracker = QueryExecutionTracker::new( + &dispatcher.query_execution_owner, + Arc::new(permit), + Instant::now() + Duration::from_millis(100), + 1, + ); + let session = SessionCtxFactory::new(true) + .create_session_ctx_with_tracker_and_memory_limit( + query.context(), + query_tracker.clone(), + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, + ) + .await + .expect("tracked test session should be available"); + assert!(query_tracker.mark_admitted(&dispatcher.query_execution_owner)); + let query_state_machine = Arc::new( + QueryStateMachine::begin_tracked(query, session, query_tracker) + .expect("tracked state machine should accept its bound session"), + ); + + wait_for_query_timeout( + query_state_machine + .query_tracker() + .expect("tracked query should retain its tracker"), + ) + .await; + + 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 })) + )); + } + + #[tokio::test] + async fn planned_query_deadline_releases_retained_admission() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let owner = QueryExecutionOwner::new(); + let query_tracker = QueryExecutionTracker::new(&owner, Arc::new(permit), Instant::now() + Duration::from_millis(10), 1); + assert!(query_tracker.mark_admitted(&owner)); + assert!(query_tracker.claim_planning(&owner)); + assert!(query_tracker.mark_planned(&owner)); + + wait_for_query_timeout(&query_tracker).await; + + assert_eq!(query_tracker.status(), QueryExecutionStatus::TimedOut); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn active_phase_deadline_waits_for_phase_drop() { + let setup_admission = Arc::new(Semaphore::new(1)); + let setup_permit = Arc::clone(&setup_admission) + .acquire_owned() + .await + .expect("setup permit should be available"); + let setup_owner = QueryExecutionOwner::new(); + let setup_tracker = + QueryExecutionTracker::new(&setup_owner, Arc::new(setup_permit), Instant::now() + Duration::from_millis(10), 1); + let setup_guard = QueryPhaseGuard::new(&setup_tracker, &setup_owner); + + wait_for_query_timeout(&setup_tracker).await; + + assert_eq!(setup_tracker.status(), QueryExecutionStatus::TimedOut); + assert_eq!(setup_admission.available_permits(), 0); + drop(setup_guard); + assert_eq!(setup_admission.available_permits(), 1); + + let planning_admission = Arc::new(Semaphore::new(1)); + let planning_permit = Arc::clone(&planning_admission) + .acquire_owned() + .await + .expect("planning permit should be available"); + let planning_owner = QueryExecutionOwner::new(); + let planning_tracker = + QueryExecutionTracker::new(&planning_owner, Arc::new(planning_permit), Instant::now() + Duration::from_millis(10), 1); + assert!(planning_tracker.mark_admitted(&planning_owner)); + assert!(planning_tracker.claim_planning(&planning_owner)); + let planning_guard = QueryPhaseGuard::new(&planning_tracker, &planning_owner); + + wait_for_query_timeout(&planning_tracker).await; + + assert_eq!(planning_tracker.status(), QueryExecutionStatus::TimedOut); + assert_eq!(planning_admission.available_permits(), 0); + drop(planning_guard); + assert_eq!(planning_admission.available_permits(), 1); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn post_deadline_result_drops_before_admission_release() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, _) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let query_tracker = QueryExecutionTracker::new( + &dispatcher.query_execution_owner, + Arc::new(permit), + Instant::now() + Duration::from_millis(10), + 1, + ); + let (drop_started_tx, drop_started_rx) = std::sync::mpsc::channel(); + let (release_drop_tx, release_drop_rx) = std::sync::mpsc::channel(); + let task_dispatcher = Arc::clone(&dispatcher); + let task_tracker = query_tracker.clone(); + let task = tokio::spawn(async move { + let _phase_guard = QueryPhaseGuard::new(&task_tracker, &task_dispatcher.query_execution_owner); + task_dispatcher + .run_with_query_deadline(&task_tracker, async move { + std::thread::sleep(Duration::from_millis(20)); + Ok(BlockingDrop { + started: drop_started_tx, + release: release_drop_rx, + }) + }) + .await + }); + let drop_started = tokio::task::spawn_blocking(move || drop_started_rx.recv()) + .await + .expect("drop observer task should finish"); + drop_started.expect("post-deadline result should be dropped"); + + assert_eq!(admission.available_permits(), 0); + 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 })) + )); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn staged_query_planning_error_releases_admission() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let retained_state_machine = Arc::clone(&query_state_machine); + + assert!(matches!( + dispatcher.build_logical_plan(query_state_machine).await, + Err(QueryError::Parser { .. }) + )); + assert_eq!(admission.available_permits(), 1); + assert!(matches!( + dispatcher.build_logical_plan(retained_state_machine).await, + Err(QueryError::Cancel) + )); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn cancelled_execution_start_drops_future_before_releasing_admission() { + let admission = Arc::new(Semaphore::new(1)); + let (drop_started_tx, drop_started_rx) = std::sync::mpsc::channel(); + let (release_drop_tx, release_drop_rx) = std::sync::mpsc::channel(); + let (dispatcher, input) = test_dispatcher_with_factory( + Arc::clone(&admission), + Duration::from_secs(300), + Arc::new(DropBlockingPendingQueryExecutionFactory { + drop_guard: std::sync::Mutex::new(Some(BlockingDrop { + started: drop_started_tx, + release: release_drop_rx, + })), + }), + ); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("staged query should build a logical plan") + .expect("select query should produce a logical plan"); + let retained_state_machine = Arc::clone(&query_state_machine); + let task_dispatcher = Arc::clone(&dispatcher); + let mut execution = + Box::pin(async move { task_dispatcher.execute_logical_plan(logical_plan, query_state_machine).await }); + + assert!(futures::poll!(execution.as_mut()).is_pending()); + let drop_task = tokio::task::spawn_blocking(move || drop(execution)); + tokio::task::spawn_blocking(move || drop_started_rx.recv()) + .await + .expect("drop observer task should finish") + .expect("cancelled execution future should be dropped"); + + assert_eq!(admission.available_permits(), 0); + release_drop_tx.send(()).expect("release execution future drop"); + drop_task.await.expect("execution drop task should finish"); + assert_eq!(admission.available_permits(), 1); + assert!(matches!( + dispatcher.build_logical_plan(retained_state_machine).await, + Err(QueryError::Cancel) + )); + } + + #[tokio::test] + async fn forged_query_tracker_cannot_bypass_dispatcher_admission() { + let admission = Arc::new(Semaphore::new(1)); + let (dispatcher, input) = test_dispatcher(Arc::clone(&admission), Duration::from_secs(300)); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let forged_permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("forged permit should be available"); + let forged_owner = QueryExecutionOwner::new(); + let forged_tracker = + QueryExecutionTracker::new(&forged_owner, Arc::new(forged_permit), Instant::now() + Duration::from_secs(300), 300); + assert!(forged_tracker.mark_admitted(&forged_owner)); + forged_tracker.finish(&dispatcher.query_execution_owner); + forged_tracker.expire(&dispatcher.query_execution_owner); + assert_eq!(forged_tracker.status(), QueryExecutionStatus::Active); + assert_eq!(admission.available_permits(), 0); + let session = SessionCtxFactory::new(true) + .create_session_ctx_with_tracker_and_memory_limit( + query.context(), + forged_tracker.clone(), + DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, + ) + .await + .expect("tracked test session should be available"); + let query_state_machine = Arc::new( + QueryStateMachine::begin_tracked(query, session, forged_tracker) + .expect("forged state machine should accept its own bound session"), + ); + let retained_state_machine = Arc::clone(&query_state_machine); + + assert!(matches!( + dispatcher.build_logical_plan(query_state_machine).await, + Err(QueryError::Cancel) + )); + assert_eq!(admission.available_permits(), 0); + drop(retained_state_machine); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn query_stream_releases_permit_after_timeout() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let inner_dropped = Arc::new(AtomicBool::new(false)); + let drop_signal = DropSignal(Arc::clone(&inner_dropped)); + let inner = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + stream::poll_fn(move |_| { + let _drop_signal = &drop_signal; + Poll::Pending::>> + }), + )); + let (owner, query_tracker) = test_query_tracker(permit, Instant::now() + Duration::from_millis(10), 300); + let mut output = Box::pin(TrackedRecordBatchStream::new(inner, query_tracker, owner)); + + let err = output + .next() + .await + .expect("timeout error") + .expect_err("expired query must fail"); + let DataFusionError::External(source) = err else { + panic!("expected external query error"); + }; + assert!(matches!( + source.downcast_ref::(), + Some(S3SelectPolicyError::QueryTimeout { seconds: 300 }) + )); + assert!(inner_dropped.load(Ordering::SeqCst)); + assert_eq!(admission.available_permits(), 1); + assert!(output.next().await.is_none()); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn stream_handoff_panic_releases_admission_after_inner_drop() { + let admission = Arc::new(Semaphore::new(1)); + let inner_dropped = Arc::new(AtomicBool::new(false)); + let (drop_started_tx, drop_started_rx) = std::sync::mpsc::channel(); + let (release_drop_tx, release_drop_rx) = std::sync::mpsc::channel(); + let (dispatcher, input) = test_dispatcher_with_factory( + Arc::clone(&admission), + Duration::from_secs(300), + Arc::new(PanickingSchemaQueryExecutionFactory { + dropped: Arc::clone(&inner_dropped), + drop_guard: std::sync::Mutex::new(Some(BlockingDrop { + started: drop_started_tx, + release: release_drop_rx, + })), + }), + ); + let query = Query::new(QueryContext { input }, "SELECT * FROM S3Object".to_string()); + let query_state_machine = dispatcher + .build_query_state_machine(query) + .await + .expect("staged query should acquire admission"); + let logical_plan = dispatcher + .build_logical_plan(Arc::clone(&query_state_machine)) + .await + .expect("staged query should build a logical plan") + .expect("select query should produce a logical plan"); + let retained_state_machine = Arc::clone(&query_state_machine); + let task_dispatcher = Arc::clone(&dispatcher); + let task = tokio::spawn(async move { task_dispatcher.execute_logical_plan(logical_plan, query_state_machine).await }); + tokio::task::spawn_blocking(move || drop_started_rx.recv()) + .await + .expect("drop observer task should finish") + .expect("panicking stream should start dropping"); + + assert!(inner_dropped.load(Ordering::SeqCst)); + assert_eq!(admission.available_permits(), 0); + release_drop_tx.send(()).expect("release stream drop"); + let Err(join_error) = task.await else { + panic!("stream schema should panic"); + }; + assert!(join_error.is_panic()); + assert_eq!(admission.available_permits(), 1); + assert_eq!( + retained_state_machine + .query_tracker() + .expect("staged query should retain its tracker") + .status(), + QueryExecutionStatus::Finished + ); + } + + #[tokio::test] + async fn query_deadline_releases_resources_without_polling_stream() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let inner_dropped = Arc::new(AtomicBool::new(false)); + let drop_signal = DropSignal(Arc::clone(&inner_dropped)); + let inner = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + stream::poll_fn(move |_| { + let _drop_signal = &drop_signal; + Poll::Pending::>> + }), + )); + let (owner, query_tracker) = test_query_tracker(permit, Instant::now() + Duration::from_millis(10), 300); + let _output = TrackedRecordBatchStream::new(inner, query_tracker, owner); + + let recovered_permit = tokio::time::timeout(Duration::from_secs(5), Arc::clone(&admission).acquire_owned()) + .await + .expect("deadline should release the admission permit") + .expect("admission semaphore should remain open"); + + assert!(inner_dropped.load(Ordering::SeqCst)); + drop(recovered_permit); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn query_stream_releases_permit_after_completion() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let inner = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + stream::empty::>(), + )); + let (owner, query_tracker) = test_query_tracker(permit, Instant::now() + Duration::from_secs(300), 300); + let mut output = Box::pin(TrackedRecordBatchStream::new(inner, query_tracker, owner)); + + assert!(output.next().await.is_none()); + assert_eq!(admission.available_permits(), 1); + } + + #[tokio::test] + async fn query_timeout_during_inner_poll_returns_error() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let inner = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + stream::pending::>(), + )); + let (owner, query_tracker) = test_query_tracker(permit, Instant::now() + Duration::from_secs(300), 300); + let mut output = Box::pin(TrackedRecordBatchStream::new(inner, query_tracker, owner)); + let deadline_state = Arc::clone(&output.state); + let poll_state = Arc::clone(&deadline_state); + *deadline_state.inner.lock() = Some(Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + stream::poll_fn(move |_| { + poll_state.query_tracker.expire(&poll_state.query_execution_owner); + Poll::Ready(None::>) + }), + ))); + + let err = output + .next() + .await + .expect("timeout error") + .expect_err("timeout racing with inner poll must fail"); + let DataFusionError::External(source) = err else { + panic!("expected external query error"); + }; + assert!(matches!( + source.downcast_ref::(), + Some(S3SelectPolicyError::QueryTimeout { seconds: 300 }) + )); + assert_eq!(admission.available_permits(), 1); + assert!(output.next().await.is_none()); + } + + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] + async fn timeout_drops_late_stream_item_before_releasing_admission() { + let admission = Arc::new(Semaphore::new(1)); + let permit = Arc::clone(&admission) + .acquire_owned() + .await + .expect("admission permit should be available"); + let (drop_started_tx, drop_started_rx) = std::sync::mpsc::channel(); + let (release_drop_tx, release_drop_rx) = std::sync::mpsc::channel(); + let mut drop_guard = Some(BlockingError { + started: drop_started_tx, + release: std::sync::Mutex::new(release_drop_rx), + }); + let inner = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + stream::poll_fn(move |_| { + std::thread::sleep(Duration::from_millis(20)); + Poll::Ready(Some(Err(DataFusionError::External(Box::new( + drop_guard.take().expect("late item should be returned once"), + ))))) + }), + )); + let (owner, query_tracker) = test_query_tracker(permit, Instant::now() + Duration::from_millis(10), 1); + let mut output = Box::pin(TrackedRecordBatchStream::new(inner, query_tracker, owner)); + let task = tokio::spawn(async move { + let result = output.next().await; + (result, output) + }); + tokio::task::spawn_blocking(move || drop_started_rx.recv()) + .await + .expect("drop observer task should finish") + .expect("late stream item should start dropping"); + + assert_eq!(admission.available_permits(), 0); + release_drop_tx.send(()).expect("release late item drop"); + let (result, mut output) = task.await.expect("stream poll task should finish"); + let error = result + .expect("timeout error should be returned") + .expect_err("late item must be replaced by a timeout"); + let DataFusionError::External(source) = error else { + panic!("expected external query error"); + }; + assert!(matches!( + source.downcast_ref::(), + Some(S3SelectPolicyError::QueryTimeout { seconds: 1 }) + )); + assert_eq!(admission.available_permits(), 1); + assert!(output.next().await.is_none()); + } +} diff --git a/crates/s3select-query/src/instance.rs b/crates/s3select-query/src/instance.rs index ce9d5edd3..ca2807d46 100644 --- a/crates/s3select-query/src/instance.rs +++ b/crates/s3select-query/src/instance.rs @@ -12,18 +12,26 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::sync::Arc; +use std::{ + sync::{Arc, LazyLock}, + time::Duration, +}; use async_trait::async_trait; use derive_builder::Builder; use rustfs_s3select_api::{ QueryResult, query::{ - Query, dispatcher::QueryDispatcher, execution::QueryStateMachineRef, logical_planner::Plan, session::SessionCtxFactory, + Query, + dispatcher::QueryDispatcher, + execution::QueryStateMachineRef, + logical_planner::Plan, + session::{DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES as DEFAULT_MEMORY_LIMIT_BYTES, SessionCtxFactory}, }, server::dbms::{DatabaseManagerSystem, QueryHandle}, }; use s3s::dto::SelectObjectContentInput; +use tokio::sync::Semaphore; use crate::{ dispatcher::manager::SimpleQueryDispatcherBuilder, @@ -34,6 +42,16 @@ use crate::{ }; const ENV_RUSTFS_S3SELECT_TARGET_PARTITIONS: &str = "RUSTFS_S3SELECT_TARGET_PARTITIONS"; +const ENV_RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES: &str = "RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES"; +const ENV_RUSTFS_S3SELECT_QUERY_TIMEOUT_SECS: &str = "RUSTFS_S3SELECT_QUERY_TIMEOUT_SECS"; +const ENV_RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES: &str = "RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES"; +pub(crate) const DEFAULT_QUERY_TIMEOUT_SECS: u64 = 300; +pub(crate) const DEFAULT_MAX_CONCURRENT_QUERIES: usize = 4; +const MAX_QUERY_TIMEOUT_SECS: u64 = 24 * 60 * 60; +const TEST_MAX_CONCURRENT_QUERIES: usize = 1024; + +static QUERY_ADMISSION: LazyLock> = + LazyLock::new(|| Arc::new(Semaphore::new(S3SelectRuntimeConfig::from_env().max_concurrent_queries))); #[derive(Builder)] pub struct RustFSms { @@ -79,9 +97,23 @@ where } } -#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +#[derive(Clone, Copy, Debug, Eq, PartialEq)] struct S3SelectRuntimeConfig { target_partitions: usize, + memory_limit_bytes: usize, + query_timeout: Duration, + max_concurrent_queries: usize, +} + +impl Default for S3SelectRuntimeConfig { + fn default() -> Self { + Self { + target_partitions: 0, + memory_limit_bytes: DEFAULT_MEMORY_LIMIT_BYTES, + query_timeout: Duration::from_secs(DEFAULT_QUERY_TIMEOUT_SECS), + max_concurrent_queries: DEFAULT_MAX_CONCURRENT_QUERIES, + } + } } impl S3SelectRuntimeConfig { @@ -90,14 +122,55 @@ impl S3SelectRuntimeConfig { target_partitions: target_partitions_from_env_value( std::env::var(ENV_RUSTFS_S3SELECT_TARGET_PARTITIONS).ok().as_deref(), ), + memory_limit_bytes: bounded_usize_from_env_value( + std::env::var(ENV_RUSTFS_S3SELECT_MEMORY_LIMIT_BYTES).ok().as_deref(), + DEFAULT_MEMORY_LIMIT_BYTES, + usize::MAX, + ), + query_timeout: s3_select_query_timeout(), + max_concurrent_queries: bounded_usize_from_env_value( + std::env::var(ENV_RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES).ok().as_deref(), + DEFAULT_MAX_CONCURRENT_QUERIES, + Semaphore::MAX_PERMITS, + ), } } } +pub fn s3_select_query_timeout() -> Duration { + Duration::from_secs(bounded_u64_from_env_value( + std::env::var(ENV_RUSTFS_S3SELECT_QUERY_TIMEOUT_SECS).ok().as_deref(), + DEFAULT_QUERY_TIMEOUT_SECS, + MAX_QUERY_TIMEOUT_SECS, + )) +} + fn target_partitions_from_env_value(value: Option<&str>) -> usize { value.and_then(|value| value.parse::().ok()).unwrap_or(0) } +fn bounded_usize_from_env_value(value: Option<&str>, default: usize, max: usize) -> usize { + value + .and_then(|value| value.parse::().ok()) + .filter(|value| (1..=max).contains(value)) + .unwrap_or(default) +} + +fn bounded_u64_from_env_value(value: Option<&str>, default: u64, max: u64) -> u64 { + value + .and_then(|value| value.parse::().ok()) + .filter(|value| (1..=max).contains(value)) + .unwrap_or(default) +} + +fn query_admission(is_test: bool) -> Arc { + if is_test { + Arc::new(Semaphore::new(TEST_MAX_CONCURRENT_QUERIES)) + } else { + Arc::clone(&QUERY_ADMISSION) + } +} + pub async fn make_rustfsms(input: Arc, is_test: bool) -> QueryResult { // init Function Manager, we can define some UDF if need let func_manager = SimpleFunctionMetadataManager::default(); @@ -116,8 +189,11 @@ pub async fn make_rustfsms(input: Arc, is_test: bool) .with_func_manager(Arc::new(func_manager)) .with_default_table_provider(default_table_provider) .with_session_factory(session_factory) + .with_memory_limit_bytes(runtime_config.memory_limit_bytes) .with_parser(parser) .with_query_execution_factory(query_execution_factory) + .with_query_admission(query_admission(is_test)) + .with_query_timeout(runtime_config.query_timeout) .build()?; let mut builder = RustFSmsBuilder::default(); @@ -143,8 +219,11 @@ pub async fn make_rustfsms_with_components( .with_func_manager(func_manager) .with_default_table_provider(default_table_provider) .with_session_factory(session_factory) + .with_memory_limit_bytes(runtime_config.memory_limit_bytes) .with_parser(parser) .with_query_execution_factory(query_execution_factory) + .with_query_admission(query_admission(is_test)) + .with_query_timeout(runtime_config.query_timeout) .build()?; let mut builder = RustFSmsBuilder::default(); @@ -167,7 +246,10 @@ mod tests { use crate::get_global_db; - use super::{S3SelectRuntimeConfig, target_partitions_from_env_value}; + use super::{ + DEFAULT_MAX_CONCURRENT_QUERIES, DEFAULT_MEMORY_LIMIT_BYTES, DEFAULT_QUERY_TIMEOUT_SECS, MAX_QUERY_TIMEOUT_SECS, + S3SelectRuntimeConfig, bounded_u64_from_env_value, bounded_usize_from_env_value, target_partitions_from_env_value, + }; #[test] fn parses_target_partitions_from_env_value() { @@ -179,7 +261,24 @@ mod tests { #[test] fn default_runtime_config_uses_datafusion_default_partitions() { - assert_eq!(S3SelectRuntimeConfig::default().target_partitions, 0); + let config = S3SelectRuntimeConfig::default(); + + assert_eq!(config.target_partitions, 0); + assert_eq!(config.memory_limit_bytes, DEFAULT_MEMORY_LIMIT_BYTES); + assert_eq!(config.query_timeout.as_secs(), DEFAULT_QUERY_TIMEOUT_SECS); + assert_eq!(config.max_concurrent_queries, DEFAULT_MAX_CONCURRENT_QUERIES); + } + + #[test] + fn resource_limits_reject_invalid_and_out_of_range_values() { + assert_eq!(bounded_usize_from_env_value(Some("1024"), 64, 2048), 1024); + assert_eq!(bounded_usize_from_env_value(Some("0"), 64, 2048), 64); + assert_eq!(bounded_usize_from_env_value(Some("4096"), 64, 2048), 64); + assert_eq!(bounded_usize_from_env_value(Some("invalid"), 64, 2048), 64); + assert_eq!(bounded_u64_from_env_value(Some("30"), 300, MAX_QUERY_TIMEOUT_SECS), 30); + assert_eq!(bounded_u64_from_env_value(Some("0"), 300, MAX_QUERY_TIMEOUT_SECS), 300); + assert_eq!(bounded_u64_from_env_value(Some("86401"), 300, MAX_QUERY_TIMEOUT_SECS), 300); + assert_eq!(bounded_u64_from_env_value(None, 300, MAX_QUERY_TIMEOUT_SECS), 300); } #[tokio::test] diff --git a/crates/s3select-query/src/metadata/mod.rs b/crates/s3select-query/src/metadata/mod.rs index a682988d8..a1776248c 100644 --- a/crates/s3select-query/src/metadata/mod.rs +++ b/crates/s3select-query/src/metadata/mod.rs @@ -19,7 +19,7 @@ use datafusion::arrow::datatypes::DataType; use datafusion::common::{Result as DFResult, TableReference}; use datafusion::datasource::TableProvider; use datafusion::logical_expr::var_provider::is_system_variables; -use datafusion::logical_expr::{AggregateUDF, HigherOrderUDF, ScalarUDF, TableSource, WindowUDF}; +use datafusion::logical_expr::{AggregateUDF, HigherOrderUDF, ScalarUDF, TableSource, WindowUDF, planner::ExprPlanner}; use datafusion::variable::VarType; use datafusion::{config::ConfigOptions, sql::planner::ContextProvider}; use rustfs_s3select_api::query::{function::FuncMetaManagerRef, session::SessionCtx}; @@ -81,6 +81,10 @@ impl ContextProviderExtension for MetadataProvider { } impl ContextProvider for MetadataProvider { + fn get_expr_planners(&self) -> &[Arc] { + self.session.inner().expr_planners() + } + fn get_function_meta(&self, name: &str) -> Option> { self.func_manager .udf(name) diff --git a/crates/s3select-query/src/sql/planner.rs b/crates/s3select-query/src/sql/planner.rs index 03d96ac57..c262166f0 100644 --- a/crates/s3select-query/src/sql/planner.rs +++ b/crates/s3select-query/src/sql/planner.rs @@ -12,11 +12,18 @@ // See the License for the specific language governing permissions and // limitations under the License. +use std::ops::ControlFlow; + use async_recursion::async_recursion; use async_trait::async_trait; -use datafusion::sql::{planner::SqlToRel, sqlparser::ast::Statement}; +use datafusion::sql::{ + planner::SqlToRel, + sqlparser::ast::{ + GroupByExpr, ObjectNamePart, OrderByKind, Query, Select, SelectFlavor, SetExpr, Statement, TableFactor, Visit, Visitor, + }, +}; use rustfs_s3select_api::{ - QueryError, QueryResult, + QueryError, QueryResult, S3SelectPolicyError, query::{ ast::ExtStatement, logical_planner::{LogicalPlanner, Plan, QueryPlan}, @@ -60,6 +67,7 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> { async fn df_sql_to_plan(&self, stmt: Statement, _session: &SessionCtx) -> QueryResult { match stmt { Statement::Query(_) => { + validate_s3_select_statement(&stmt)?; let df_plan = self.df_planner.sql_statement_to_plan(stmt)?; let plan = Plan::Query(QueryPlan { df_plan, @@ -72,3 +80,256 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> { } } } + +fn validate_s3_select_statement(statement: &Statement) -> QueryResult<()> { + let Statement::Query(query) = statement else { + return Err(unsupported_structure("only SELECT queries are supported")); + }; + + if query.with.is_some() + || query.order_by.as_ref().is_some_and(|order_by| { + order_by.interpolate.is_some() + || match &order_by.kind { + OrderByKind::Expressions(expressions) => expressions.iter().any(|expression| expression.with_fill.is_some()), + OrderByKind::All(_) => true, + } + }) + || query.fetch.is_some() + || !query.locks.is_empty() + || query.for_clause.is_some() + || query.settings.is_some() + || query.format_clause.is_some() + || !query.pipe_operators.is_empty() + { + return Err(unsupported_structure("the query contains an unsupported clause")); + } + if let Some(limit_clause) = query.limit_clause.as_ref() + && !matches!( + limit_clause, + datafusion::sql::sqlparser::ast::LimitClause::LimitOffset { + limit: Some(_), + offset: None, + limit_by, + } if limit_by.is_empty() + ) + { + return Err(unsupported_structure("only LIMIT without OFFSET is supported")); + } + if let Some(datafusion::sql::sqlparser::ast::LimitClause::LimitOffset { limit: Some(limit), .. }) = + query.limit_clause.as_ref() + && limit.to_string().parse::().is_err() + { + return Err(unsupported_structure("LIMIT must be a non-negative integer")); + } + + let mut detector = SubqueryDetector { visited_root: false }; + if query.visit(&mut detector).is_break() { + return Err(unsupported_structure("subqueries are not supported")); + } + + let SetExpr::Select(select) = query.body.as_ref() else { + return Err(unsupported_structure("set operations and nested queries are not supported")); + }; + validate_select(select) +} + +fn validate_select(select: &Select) -> QueryResult<()> { + if !select.optimizer_hints.is_empty() + || select.distinct.is_some() + || select.select_modifiers.is_some() + || select.top.is_some() + || select.exclude.is_some() + || select.into.is_some() + || !select.lateral_views.is_empty() + || select.prewhere.is_some() + || !select.connect_by.is_empty() + || !select.cluster_by.is_empty() + || !select.distribute_by.is_empty() + || !select.sort_by.is_empty() + || select.having.is_some() + || !select.named_window.is_empty() + || select.qualify.is_some() + || select.value_table_mode.is_some() + || select.flavor != SelectFlavor::Standard + || !matches!(&select.group_by, GroupByExpr::Expressions(_, modifiers) if modifiers.is_empty()) + { + return Err(unsupported_structure("the SELECT contains an unsupported clause")); + } + + let [table] = select.from.as_slice() else { + return Err(unsupported_structure("exactly one S3Object source is required")); + }; + if !table.joins.is_empty() { + return Err(unsupported_structure("JOIN is not supported")); + } + let TableFactor::Table { + name, + alias, + args, + with_hints, + version, + with_ordinality, + partitions, + sample, + index_hints, + .. + } = &table.relation + else { + return Err(unsupported_structure("subqueries and table functions are not supported")); + }; + if args.is_some() + || !with_hints.is_empty() + || version.is_some() + || *with_ordinality + || !partitions.is_empty() + || sample.is_some() + || !index_hints.is_empty() + || alias.as_ref().is_some_and(|alias| !alias.columns.is_empty()) + { + return Err(unsupported_structure("the S3Object source contains unsupported modifiers")); + } + 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")); + }; + let is_s3_object = if table_name.quote_style.is_some() { + table_name.value == "S3Object" + } else { + table_name.value.eq_ignore_ascii_case("S3Object") + }; + if !is_s3_object { + return Err(unsupported_structure("the source must be S3Object")); + } + + Ok(()) +} + +fn unsupported_structure(message: &str) -> QueryError { + S3SelectPolicyError::UnsupportedSqlStructure { + message: message.to_string(), + } + .into() +} + +struct SubqueryDetector { + visited_root: bool, +} + +impl Visitor for SubqueryDetector { + type Break = (); + + fn pre_visit_query(&mut self, _query: &Query) -> ControlFlow { + if self.visited_root { + ControlFlow::Break(()) + } else { + self.visited_root = true; + ControlFlow::Continue(()) + } + } +} + +#[cfg(test)] +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}; + + fn parse_statement(sql: &str) -> Statement { + let mut statements = ExtParser::parse_sql(sql).expect("SQL should parse"); + let ExtStatement::SqlStatement(statement) = statements.pop_front().expect("one SQL statement"); + *statement + } + + #[test] + fn accepts_s3_select_query_shape() { + let statement = parse_statement("SELECT s.id FROM S3Object AS s WHERE s.id = '1' LIMIT 10"); + + assert!(validate_s3_select_statement(&statement).is_ok()); + } + + #[test] + fn accepts_json_sub_path_source() { + let statement = parse_statement("SELECT e.name FROM S3Object.employees AS e"); + + assert!(validate_s3_select_statement(&statement).is_ok()); + } + + #[test] + fn accepts_group_by_and_order_by() { + let statement = parse_statement("SELECT department, COUNT(*) FROM S3Object GROUP BY department ORDER BY department"); + + assert!(validate_s3_select_statement(&statement).is_ok()); + } + + #[test] + fn rejects_join() { + let statement = parse_statement("SELECT * FROM S3Object a JOIN S3Object b ON a.id = b.id"); + + assert!(matches!( + 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" + ) + )); + } + + #[test] + fn rejects_subquery() { + let statement = parse_statement("SELECT * FROM S3Object WHERE id IN (SELECT id FROM S3Object)"); + + assert!(matches!( + 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" + ) + )); + } + + #[test] + fn rejects_subquery_in_order_by() { + let statement = parse_statement("SELECT id FROM S3Object ORDER BY (SELECT id FROM S3Object)"); + + assert!(matches!( + 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" + ) + )); + } + + #[test] + fn rejects_non_s3_object_source() { + let statement = parse_statement("SELECT * FROM other_table"); + + assert!(matches!( + 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" + ) + )); + } + + #[test] + fn rejects_unsupported_select_clauses() { + for sql in [ + "SELECT DISTINCT id FROM S3Object", + "SELECT * FROM S3Object OFFSET 1", + "SELECT * FROM S3Object UNION SELECT * FROM S3Object", + ] { + let statement = parse_statement(sql); + assert!( + matches!( + validate_s3_select_statement(&statement), + Err(ref err) if matches!(err.s3_select_policy_error(), Some(S3SelectPolicyError::UnsupportedSqlStructure { .. })) + ), + "query should be rejected: {sql}" + ); + } + } +} diff --git a/crates/s3select-query/src/test/integration_test.rs b/crates/s3select-query/src/test/integration_test.rs index eee801261..461dbedb3 100644 --- a/crates/s3select-query/src/test/integration_test.rs +++ b/crates/s3select-query/src/test/integration_test.rs @@ -15,7 +15,10 @@ #[cfg(test)] mod integration_tests { use crate::{create_fresh_db, get_global_db, instance::make_rustfsms}; - use datafusion::arrow::array::{Array, StringArray}; + use datafusion::arrow::{ + array::{Array, Int64Array, StringArray}, + record_batch::RecordBatch, + }; use rustfs_s3select_api::{ QueryError, query::{Context, Query}, @@ -26,6 +29,52 @@ mod integration_tests { }; use std::sync::Arc; + fn assert_ages_descending(output: &[RecordBatch]) { + let ages: Vec = output + .iter() + .flat_map(|batch| { + batch + .column(1) + .as_any() + .downcast_ref::() + .expect("age column should be Int64") + .values() + .iter() + .copied() + .collect::>() + }) + .collect(); + + assert_eq!(ages, vec![40, 38, 35, 32, 30, 28, 26, 25, 24, 22]); + } + + fn assert_department_counts(output: &[RecordBatch]) { + let mut counts: Vec<(&str, i64)> = output + .iter() + .flat_map(|batch| { + let departments = batch + .column(0) + .as_any() + .downcast_ref::() + .expect("department column should be Utf8"); + let counts = batch + .column(1) + .as_any() + .downcast_ref::() + .expect("count column should be Int64"); + departments.iter().zip(counts.iter()).map(|(department, count)| { + ( + department.expect("department should not be null"), + count.expect("count should not be null"), + ) + }) + }) + .collect(); + counts.sort_unstable(); + + assert_eq!(counts, vec![("Finance", 3), ("HR", 2), ("IT", 3), ("Marketing", 2)]); + } + fn create_test_input(sql: &str) -> SelectObjectContentInput { SelectObjectContentInput { bucket: "test-bucket".to_string(), @@ -292,21 +341,21 @@ mod integration_tests { #[tokio::test] async fn test_select_with_aggregation() { - let sql = "SELECT department, COUNT(*) as count FROM S3Object GROUP BY department"; + let sql = "SELECT department, COUNT(*) FROM S3Object GROUP BY department"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); - let result = db.execute(&query).await; - // Aggregation queries might fail due to lack of actual data, which is acceptable - match result { - Ok(_) => { - // If successful, that's great - } - Err(_) => { - // Expected to fail due to no actual data source - } - } + let output = db + .execute(&query) + .await + .expect("execute grouped CSV query") + .result() + .chunk_result() + .await + .expect("collect grouped CSV output"); + + assert_department_counts(&output); } #[tokio::test] @@ -381,13 +430,22 @@ mod integration_tests { #[tokio::test] async fn test_query_with_order_by() { - let sql = "SELECT name, age FROM S3Object ORDER BY age DESC"; + let sql = "SELECT name, CAST(age AS BIGINT) AS age FROM S3Object ORDER BY CAST(age AS BIGINT) DESC"; let input = create_test_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); - let result = db.execute(&query).await; - assert!(result.is_ok()); + let output = db + .execute(&query) + .await + .expect("execute ordered CSV query") + .result() + .chunk_result() + .await + .expect("collect ordered CSV output"); + + assert_eq!(output.iter().map(|batch| batch.num_rows()).sum::(), 10); + assert_ages_descending(&output); } #[tokio::test] @@ -582,21 +640,21 @@ mod integration_tests { #[tokio::test] async fn test_select_with_aggregation_json() { - let sql = "SELECT department, COUNT(*) as count FROM S3Object GROUP BY department"; + let sql = "SELECT department, COUNT(*) FROM S3Object GROUP BY department"; let input = create_test_json_input(sql); let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); - let result = db.execute(&query).await; - // Aggregation queries may fail due to lack of actual data, which is acceptable - match result { - Ok(_) => { - // If successful, that's great - } - Err(_) => { - // Expected to fail due to no actual data source - } - } + let output = db + .execute(&query) + .await + .expect("execute grouped JSON query") + .result() + .chunk_result() + .await + .expect("collect grouped JSON output"); + + assert_department_counts(&output); } #[tokio::test] @@ -672,8 +730,17 @@ mod integration_tests { let db = get_global_db(input.clone(), true).await.unwrap(); let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); - let result = db.execute(&query).await; - assert!(result.is_ok()); + let output = db + .execute(&query) + .await + .expect("execute ordered JSON query") + .result() + .chunk_result() + .await + .expect("collect ordered JSON output"); + + assert_eq!(output.iter().map(|batch| batch.num_rows()).sum::(), 10); + assert_ages_descending(&output); } #[tokio::test] diff --git a/rustfs/src/app/select_object.rs b/rustfs/src/app/select_object.rs index db7b3f41e..5d803fb14 100644 --- a/rustfs/src/app/select_object.rs +++ b/rustfs/src/app/select_object.rs @@ -10,13 +10,16 @@ use datafusion::arrow::{ json::{WriterBuilder as JsonWriterBuilder, writer::LineDelimited}, record_batch::RecordBatch, }; +use datafusion::common::DataFusionError; +use datafusion::physical_plan::SendableRecordBatchStream; use futures::StreamExt; use http::{StatusCode, header::RANGE}; use rustfs_s3select_api::{ - QueryError, + 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 s3s::dto::{ CSVOutput, CompressionType, ContinuationEvent, EndEvent, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType, OutputSerialization, Progress, ProgressEvent, QuoteFields, RecordsEvent, SelectObjectContentEvent, @@ -26,6 +29,7 @@ use s3s::dto::{ use s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error}; use std::sync::Arc; use tokio::sync::mpsc; +use tokio::time::{Instant, timeout_at}; use tokio_stream::wrappers::ReceiverStream; use tracing::info; @@ -60,6 +64,8 @@ pub async fn execute_select_object_content( 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) .await .map_err(map_query_error_to_s3)?; @@ -72,66 +78,13 @@ pub async fn execute_select_object_content( .into_record_batch_stream() .map_err(map_query_error_to_s3)?; - let (tx, rx) = mpsc::channel::>(8); + let (tx, rx) = mpsc::channel::>(9); + let terminal_permit = tx + .clone() + .try_reserve_owned() + .map_err(|_| s3_error!(InternalError, "can't reserve Select terminal event capacity"))?; spawn_traced(async move { - let mut encoder = SelectOutputEncoder::new(validation.output_format); - let mut progress = SelectProgress::default(); - let mut output = output; - - if tx - .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) - .await - .is_err() - { - return; - } - - while let Some(result) = output.next().await { - let batch = match result { - Ok(batch) => batch, - Err(err) => { - let _ = tx.send(Err(map_query_error_to_s3(err.into()))).await; - return; - } - }; - - match encoder.encode_batch(&batch) { - Ok(payloads) => { - for payload in payloads { - progress.add_returned(payload.len()); - if tx - .send(Ok(SelectObjectContentEvent::Records(RecordsEvent { payload: Some(payload) }))) - .await - .is_err() - { - return; - } - if validation.progress_enabled - && tx - .send(Ok(SelectObjectContentEvent::Progress(ProgressEvent { - details: Some(progress.to_progress()), - }))) - .await - .is_err() - { - return; - } - } - } - Err(err) => { - let _ = tx.send(Err(err)).await; - return; - } - } - } - - let stats = SelectObjectContentEvent::Stats(StatsEvent { - details: Some(progress.to_stats()), - }); - if tx.send(Ok(stats)).await.is_err() { - return; - } - let _ = tx.send(Ok(SelectObjectContentEvent::End(EndEvent::default()))).await; + send_select_events_until_deadline(output, tx, terminal_permit, validation, query_deadline, query_timeout.as_secs()).await; }); Ok(S3Response::new(SelectObjectContentOutput { @@ -139,6 +92,91 @@ pub async fn execute_select_object_content( })) } +async fn send_select_events_until_deadline( + output: SendableRecordBatchStream, + tx: mpsc::Sender>, + terminal_permit: mpsc::OwnedPermit>, + validation: SelectValidation, + deadline: Instant, + timeout_seconds: u64, +) { + if timeout_at(deadline, send_select_events(output, &tx, validation)) + .await + .is_err() + { + terminal_permit.send(Err(map_query_error_to_s3( + S3SelectPolicyError::QueryTimeout { + seconds: timeout_seconds, + } + .into(), + ))); + } +} + +async fn send_select_events( + mut output: SendableRecordBatchStream, + tx: &mpsc::Sender>, + validation: SelectValidation, +) { + let mut encoder = SelectOutputEncoder::new(validation.output_format); + let mut progress = SelectProgress::default(); + + if tx + .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) + .await + .is_err() + { + return; + } + + while let Some(result) = output.next().await { + let batch = match result { + Ok(batch) => batch, + Err(err) => { + let _ = tx.send(Err(map_query_error_to_s3(err.into()))).await; + return; + } + }; + + match encoder.encode_batch(&batch) { + Ok(payloads) => { + for payload in payloads { + progress.add_returned(payload.len()); + if tx + .send(Ok(SelectObjectContentEvent::Records(RecordsEvent { payload: Some(payload) }))) + .await + .is_err() + { + return; + } + if validation.progress_enabled + && tx + .send(Ok(SelectObjectContentEvent::Progress(ProgressEvent { + details: Some(progress.to_progress()), + }))) + .await + .is_err() + { + return; + } + } + } + Err(err) => { + let _ = tx.send(Err(err)).await; + return; + } + } + } + + let stats = SelectObjectContentEvent::Stats(StatsEvent { + details: Some(progress.to_stats()), + }); + if tx.send(Ok(stats)).await.is_err() { + return; + } + let _ = tx.send(Ok(SelectObjectContentEvent::End(EndEvent::default()))).await; +} + fn validate_select_request(headers: &http::HeaderMap, input: &mut SelectObjectContentInput) -> S3Result { if headers.contains_key(RANGE) { return Err(S3Error::new(S3ErrorCode::UnsupportedRangeHeader)); @@ -487,11 +525,31 @@ fn clamp_i64(value: u64) -> i64 { } fn map_query_error_to_s3(err: QueryError) -> S3Error { + if let Some(policy_error) = err.s3_select_policy_error() { + let message = policy_error.to_string(); + return match policy_error { + S3SelectPolicyError::UnsupportedSqlStructure { .. } => { + S3Error::with_message(S3ErrorCode::UnsupportedSqlStructure, message) + } + S3SelectPolicyError::QueryConcurrencyLimit => S3Error::with_message(S3ErrorCode::SlowDown, message), + S3SelectPolicyError::QueryTimeout { .. } => S3Error::with_message(S3ErrorCode::Busy, message), + _ => S3Error::with_message(S3ErrorCode::InternalError, message), + }; + } let message = err.to_string(); match err { QueryError::Parser { .. } => parse_select_failure(message), QueryError::MultiStatement { .. } => S3Error::with_message(S3ErrorCode::UnsupportedSqlStructure, message), QueryError::NotImplemented { .. } => S3Error::with_message(S3ErrorCode::NotImplemented, message), + QueryError::Datafusion { source } if is_resource_exhausted(source.as_ref()) => { + S3Error::with_message(S3ErrorCode::Busy, message) + } + QueryError::Datafusion { source } if is_unexpected_eof(source.as_ref()) => { + S3Error::with_message(S3ErrorCode::InternalError, message) + } + QueryError::Datafusion { source } if is_invalid_object_size(source.as_ref()) => { + S3Error::with_message(S3ErrorCode::InternalError, message) + } QueryError::Datafusion { .. } if looks_like_invalid_scan_range(&message) => { S3Error::with_message(S3ErrorCode::InvalidRequestParameter, INVALID_SCAN_RANGE_MESSAGE.to_string()) } @@ -520,6 +578,42 @@ fn looks_like_bucket_not_found(message: &str) -> bool { message.contains("NoSuchBucket") || message.contains("bucket not found") || message.contains("BucketNotFound") } +const MAX_ERROR_SOURCE_DEPTH: usize = 16; + +fn error_chain_any( + mut err: &(dyn std::error::Error + 'static), + predicate: impl Fn(&(dyn std::error::Error + 'static)) -> bool, +) -> bool { + for _ in 0..MAX_ERROR_SOURCE_DEPTH { + if predicate(err) { + return true; + } + let Some(source) = err.source() else { + return false; + }; + err = source; + } + false +} + +fn is_resource_exhausted(err: &(dyn std::error::Error + 'static)) -> bool { + error_chain_any(err, |err| { + err.downcast_ref::() + .is_some_and(|err| matches!(err, DataFusionError::ResourcesExhausted(_))) + }) +} + +fn is_unexpected_eof(err: &(dyn std::error::Error + 'static)) -> bool { + error_chain_any(err, |err| { + err.downcast_ref::() + .is_some_and(|err| err.kind() == std::io::ErrorKind::UnexpectedEof) + }) +} + +fn is_invalid_object_size(err: &(dyn std::error::Error + 'static)) -> bool { + error_chain_any(err, |err| err.downcast_ref::().is_some()) +} + fn looks_like_object_not_found(message: &str) -> bool { message.contains("NoSuchKey") || message.contains("NoSuchVersion") @@ -548,10 +642,27 @@ fn is_json_document(json: &JSONInput) -> bool { #[cfg(test)] mod tests { use super::*; - use datafusion::sql::sqlparser::parser::ParserError; + use datafusion::{ + arrow::datatypes::Schema, physical_plan::stream::RecordBatchStreamAdapter, sql::sqlparser::parser::ParserError, + }; use http::HeaderMap; use s3s::dto::{CSVInput, ParquetInput, ScanRange}; + #[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 std::error::Error for CyclicError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + Some(self) + } + } + fn base_input() -> SelectObjectContentInput { SelectObjectContentInput { bucket: "bucket".to_string(), @@ -611,6 +722,106 @@ mod tests { assert_eq!(err.message(), Some("sql parser error: syntax error")); } + #[test] + fn map_query_policy_errors_to_s3_errors() { + let unsupported = map_query_error_to_s3( + S3SelectPolicyError::UnsupportedSqlStructure { + message: "JOIN is not supported".to_string(), + } + .into(), + ); + let saturated = map_query_error_to_s3(S3SelectPolicyError::QueryConcurrencyLimit.into()); + let timed_out = map_query_error_to_s3(S3SelectPolicyError::QueryTimeout { seconds: 300 }.into()); + let stream_timed_out = map_query_error_to_s3(QueryError::Datafusion { + source: Box::new(DataFusionError::External(Box::new(S3SelectPolicyError::QueryTimeout { seconds: 300 }))), + }); + let exhausted = map_query_error_to_s3(QueryError::Datafusion { + source: Box::new(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted("memory limit".to_string())), + }))), + }); + let truncated = map_query_error_to_s3(QueryError::Datafusion { + source: Box::new(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::Generic { + store: "EcObjectStore", + source: Box::new(std::io::Error::new(std::io::ErrorKind::UnexpectedEof, "truncated object stream")), + }))), + }); + let invalid_object_size = map_query_error_to_s3(QueryError::Datafusion { + source: Box::new(DataFusionError::ObjectStore(Box::new(datafusion::object_store::Error::Generic { + store: "EcObjectStore", + source: Box::new(u64::try_from(-1_i64).expect_err("negative size must fail conversion")), + }))), + }); + + assert_eq!(unsupported.code(), &S3ErrorCode::UnsupportedSqlStructure); + assert_eq!(unsupported.message(), Some("Unsupported S3 Select SQL structure: JOIN is not supported")); + assert_eq!(saturated.code(), &S3ErrorCode::SlowDown); + assert_eq!(saturated.message(), Some("S3 Select query concurrency limit reached")); + assert_eq!(timed_out.code(), &S3ErrorCode::Busy); + assert_eq!(timed_out.message(), Some("S3 Select query exceeded the 300-second execution limit")); + assert_eq!(stream_timed_out.code(), &S3ErrorCode::Busy); + assert_eq!( + stream_timed_out.message(), + Some("S3 Select query exceeded the 300-second execution limit") + ); + assert_eq!(exhausted.code(), &S3ErrorCode::Busy); + assert_eq!(truncated.code(), &S3ErrorCode::InternalError); + assert_eq!(invalid_object_size.code(), &S3ErrorCode::InternalError); + } + + #[test] + fn error_source_matching_stops_at_the_depth_bound() { + let err = CyclicError; + + assert!(!is_resource_exhausted(&err)); + assert!(!is_unexpected_eof(&err)); + assert!(!is_invalid_object_size(&err)); + } + + #[tokio::test(start_paused = true)] + async fn producer_deadline_cancels_backpressured_send() { + let output = Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + futures::stream::pending::>(), + )); + let (tx, mut rx) = mpsc::channel(2); + let terminal_permit = tx + .clone() + .try_reserve_owned() + .expect("test channel should reserve terminal capacity"); + tx.send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) + .await + .expect("test channel should accept the prefilled event"); + let validation = SelectValidation { + output_format: SelectOutputFormat::Csv(CSVOutput::default()), + progress_enabled: false, + }; + + let producer = tokio::spawn(send_select_events_until_deadline( + output, + tx, + terminal_permit, + validation, + Instant::now() + std::time::Duration::from_secs(1), + 300, + )); + + tokio::task::yield_now().await; + tokio::time::advance(std::time::Duration::from_secs(1)).await; + tokio::task::yield_now().await; + + producer.await.expect("producer should finish without draining the channel"); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + let timeout_error = rx + .recv() + .await + .expect("producer should send a terminal timeout error") + .expect_err("terminal event should be an error"); + assert_eq!(timeout_error.code(), &S3ErrorCode::Busy); + assert!(rx.recv().await.is_none()); + } + #[test] fn validate_defaults_csv_header_and_compression() { let mut input = base_input();