From 589a9544780a40421823e8478bbd98f49b8991d8 Mon Sep 17 00:00:00 2001 From: GatewayJ <835269233@qq.com> Date: Mon, 31 Aug 2026 13:35:59 +0800 Subject: [PATCH] feat(s3select): support compressed CSV and JSON input (#6915) --- Cargo.lock | 6 + crates/e2e_test/Cargo.toml | 1 + crates/e2e_test/src/reliant/mod.rs | 1 + .../src/reliant/s3_select_compression.rs | 351 ++++ crates/s3select-api/Cargo.toml | 5 + crates/s3select-api/src/input_stream.rs | 1570 +++++++++++++++++ crates/s3select-api/src/lib.rs | 10 + crates/s3select-api/src/metrics.rs | 107 +- crates/s3select-api/src/object_store.rs | 941 ++++++++-- crates/s3select-api/src/query/session.rs | 24 + .../s3select-query/src/dispatcher/manager.rs | 40 +- rustfs/src/app/select_object.rs | 145 +- 12 files changed, 3057 insertions(+), 144 deletions(-) create mode 100644 crates/e2e_test/src/reliant/s3_select_compression.rs create mode 100644 crates/s3select-api/src/input_stream.rs diff --git a/Cargo.lock b/Cargo.lock index 1f6087f27..1375bf3b3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -3926,6 +3926,7 @@ dependencies = [ "aws-sdk-s3", "aws-sdk-sts", "aws-smithy-http-client", + "aws-smithy-types", "base64-simd", "bytes", "chrono", @@ -10557,10 +10558,14 @@ dependencies = [ name = "rustfs-s3select-api" version = "1.0.0-rc.4" dependencies = [ + "arc-swap", + "async-compression", "async-trait", "bytes", "chrono", + "crc-fast", "datafusion", + "flate2", "futures", "futures-core", "hotpath", @@ -10576,6 +10581,7 @@ dependencies = [ "serial_test", "thiserror 2.0.20", "tokio", + "tokio-stream", "tokio-util", "tracing", "transform-stream", diff --git a/crates/e2e_test/Cargo.toml b/crates/e2e_test/Cargo.toml index 75a539561..0a0a70872 100644 --- a/crates/e2e_test/Cargo.toml +++ b/crates/e2e_test/Cargo.toml @@ -100,6 +100,7 @@ aws-sdk-s3 = { workspace = true, default-features = false, features = ["sigv4a", aws-sdk-sts = { workspace = true, default-features = false, features = ["default-https-client", "rt-tokio"] } aws-config = { workspace = true } aws-smithy-http-client = { workspace = true, default-features = false, features = ["rustls-aws-lc"] } +aws-smithy-types.workspace = true async-compression = { workspace = true, features = ["tokio", "bzip2", "xz"] } async-trait = { workspace = true } flate2.workspace = true diff --git a/crates/e2e_test/src/reliant/mod.rs b/crates/e2e_test/src/reliant/mod.rs index ddc42a855..80a18ac11 100644 --- a/crates/e2e_test/src/reliant/mod.rs +++ b/crates/e2e_test/src/reliant/mod.rs @@ -21,5 +21,6 @@ mod head_tls_bodyless_test; mod lifecycle; mod lock; mod node_interact_test; +mod s3_select_compression; mod sql; mod tiering; diff --git a/crates/e2e_test/src/reliant/s3_select_compression.rs b/crates/e2e_test/src/reliant/s3_select_compression.rs new file mode 100644 index 000000000..379352fe2 --- /dev/null +++ b/crates/e2e_test/src/reliant/s3_select_compression.rs @@ -0,0 +1,351 @@ +#![cfg(test)] +// Copyright 2024 RustFS Team +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use crate::common::{RustFSTestEnvironment, init_logging}; +use async_compression::tokio::write::BzEncoder; +use aws_sdk_s3::{ + Client, + error::ProvideErrorMetadata, + operation::select_object_content::{SelectObjectContentOutput, builders::SelectObjectContentFluentBuilder}, + types::{ + CompressionType, CsvInput, CsvOutput, ExpressionType, FileHeaderInfo, InputSerialization, JsonInput, JsonOutput, + JsonType, OutputSerialization, SelectObjectContentEventStream, + }, +}; +use aws_smithy_types::event_stream::RawMessage; +use bytes::Bytes; +use flate2::{Compression, write::GzEncoder}; +use std::{error::Error, io::Cursor, time::Duration}; +use tokio::io::AsyncWriteExt; + +const BUCKET: &str = "s3-select-compression"; +const SELECT_RESPONSE_TIMEOUT: Duration = Duration::from_secs(30); + +type TestResult = Result>; + +async fn create_test_environment(extra_env: &[(&str, &str)]) -> TestResult<(RustFSTestEnvironment, Client)> { + init_logging(); + let mut env = RustFSTestEnvironment::new().await?; + env.start_rustfs_server_with_env(vec![], extra_env).await?; + let client = env.create_s3_client(); + client.create_bucket().bucket(BUCKET).send().await?; + Ok((env, client)) +} + +async fn put_object(client: &Client, key: &str, body: &[u8]) -> TestResult<()> { + client + .put_object() + .bucket(BUCKET) + .key(key) + .body(Bytes::copy_from_slice(body).into()) + .send() + .await?; + Ok(()) +} + +fn gzip(input: &[u8]) -> TestResult> { + let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); + std::io::Write::write_all(&mut encoder, input)?; + Ok(encoder.finish()?) +} + +async fn bzip2(input: &[u8]) -> TestResult> { + let mut encoder = BzEncoder::new(Cursor::new(Vec::new())); + encoder.write_all(input).await?; + encoder.shutdown().await?; + Ok(encoder.into_inner().into_inner()) +} + +fn csv_select_request( + client: &Client, + key: &str, + compression: CompressionType, + expression: &str, +) -> SelectObjectContentFluentBuilder { + client + .select_object_content() + .bucket(BUCKET) + .key(key) + .expression(expression) + .expression_type(ExpressionType::Sql) + .input_serialization( + InputSerialization::builder() + .compression_type(compression) + .csv(CsvInput::builder().file_header_info(FileHeaderInfo::Use).build()) + .build(), + ) + .output_serialization(OutputSerialization::builder().csv(CsvOutput::builder().build()).build()) +} + +fn json_select_request( + client: &Client, + key: &str, + compression: CompressionType, + json_type: JsonType, +) -> SelectObjectContentFluentBuilder { + client + .select_object_content() + .bucket(BUCKET) + .key(key) + .expression("SELECT name FROM S3Object") + .expression_type(ExpressionType::Sql) + .input_serialization( + InputSerialization::builder() + .compression_type(compression) + .json(JsonInput::builder().set_type(Some(json_type)).build()) + .build(), + ) + .output_serialization(OutputSerialization::builder().json(JsonOutput::builder().build()).build()) +} + +async fn collect_success( + mut response: SelectObjectContentOutput, + compressed_bytes: usize, + processed_bytes: usize, +) -> TestResult> { + tokio::time::timeout(SELECT_RESPONSE_TIMEOUT, async move { + let mut records = Vec::new(); + let mut stats = None; + let mut saw_end = false; + + while let Some(event) = response.payload.recv().await? { + assert!(!saw_end, "Select emitted an event after End"); + match event { + SelectObjectContentEventStream::Records(event) => { + assert!(stats.is_none(), "Select emitted Records after Stats"); + if let Some(payload) = event.payload { + records.extend_from_slice(payload.as_ref()); + } + } + SelectObjectContentEventStream::Stats(event) => { + assert!(stats.is_none(), "Select emitted more than one Stats event"); + stats = event.details; + } + SelectObjectContentEventStream::End(_) => { + assert!(stats.is_some(), "Select emitted End before Stats"); + saw_end = true; + } + _ => assert!(stats.is_none(), "Select emitted a non-terminal event after Stats"), + } + } + + let stats = stats.ok_or("Select response ended without a Stats event")?; + assert_eq!(stats.bytes_scanned(), Some(i64::try_from(compressed_bytes)?)); + assert_eq!(stats.bytes_processed(), Some(i64::try_from(processed_bytes)?)); + assert_eq!(stats.bytes_returned(), Some(i64::try_from(records.len())?)); + assert!(saw_end, "Select response ended without an End event"); + Ok::<_, Box>(records) + }) + .await + .map_err(|_| -> Box { "Select response timed out".into() })? +} + +async fn assert_truncated_stream_failure(mut response: SelectObjectContentOutput) -> TestResult<()> { + tokio::time::timeout(SELECT_RESPONSE_TIMEOUT, async move { + loop { + match response.payload.recv().await { + Err(error) => { + // S3 Select request-level errors use `error` frames, which this SDK version exposes as raw response errors. + if let Some(code) = error.code() { + assert_eq!(code, "TruncatedInput", "unexpected modeled event-stream error: {error:?}"); + } else if let aws_sdk_s3::error::SdkError::ResponseError(context) = &error + && let RawMessage::Decoded(message) = context.raw() + { + let header = |name: &str| { + message + .headers() + .iter() + .find(|header| header.name().as_str() == name) + .and_then(|header| header.value().as_string().ok()) + .map(|value| value.as_str()) + }; + assert_eq!(header(":message-type"), Some("error")); + assert_eq!(header(":error-code"), Some("TruncatedInput")); + } else { + panic!("unexpected event-stream error: {error:?}"); + } + return Ok(()); + } + Ok(Some(SelectObjectContentEventStream::Stats(_))) | Ok(Some(SelectObjectContentEventStream::End(_))) => { + return Err("truncated compressed input reached a success terminal event".into()); + } + Ok(Some(_)) => {} + Ok(None) => return Err("truncated compressed input ended without an error event".into()), + } + } + }) + .await + .map_err(|_| -> Box { "truncated Select response timed out".into() })? +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn test_select_object_content_compressed_csv_and_json() -> TestResult<()> { + const CSV: &[u8] = b"name,age\nAlice,30\nBob,25\n"; + const JSON_LINES: &[u8] = b"{\"name\":\"Alice\"}\n{\"name\":\"Bob\"}\n"; + const JSON_DOCUMENT: &[u8] = br#"[{"name":"Alice"},{"name":"Bob"}]"#; + + let (_env, client) = create_test_environment(&[]).await?; + + let gzip_csv = gzip(CSV)?; + put_object(&client, "records.csv.gz", &gzip_csv).await?; + let gzip_csv_records = collect_success( + csv_select_request(&client, "records.csv.gz", CompressionType::Gzip, "SELECT * FROM S3Object") + .send() + .await?, + gzip_csv.len(), + CSV.len(), + ) + .await?; + assert_eq!(gzip_csv_records, b"Alice,30\nBob,25\n"); + + let bzip_csv = bzip2(CSV).await?; + put_object(&client, "records.csv.bz2", &bzip_csv).await?; + let bzip_csv_records = collect_success( + csv_select_request(&client, "records.csv.bz2", CompressionType::Bzip2, "SELECT * FROM S3Object") + .send() + .await?, + bzip_csv.len(), + CSV.len(), + ) + .await?; + assert_eq!(bzip_csv_records, gzip_csv_records); + + let gzip_json_lines = gzip(JSON_LINES)?; + put_object(&client, "json-lines", &gzip_json_lines).await?; + let gzip_json_records = collect_success( + json_select_request(&client, "json-lines", CompressionType::Gzip, JsonType::Lines) + .send() + .await?, + gzip_json_lines.len(), + JSON_LINES.len(), + ) + .await?; + assert_eq!(gzip_json_records, JSON_LINES); + + let bzip_json_lines = bzip2(JSON_LINES).await?; + put_object(&client, "records.jsonl.bz2", &bzip_json_lines).await?; + let bzip_json_records = collect_success( + json_select_request(&client, "records.jsonl.bz2", CompressionType::Bzip2, JsonType::Lines) + .send() + .await?, + bzip_json_lines.len(), + JSON_LINES.len(), + ) + .await?; + assert_eq!(bzip_json_records, gzip_json_records); + + let gzip_json_document = gzip(JSON_DOCUMENT)?; + put_object(&client, "document.json.gz", &gzip_json_document).await?; + let document_records = collect_success( + json_select_request(&client, "document.json.gz", CompressionType::Gzip, JsonType::Document) + .send() + .await?, + gzip_json_document.len(), + JSON_DOCUMENT.len(), + ) + .await?; + assert_eq!(document_records, JSON_LINES); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn test_select_object_content_invalid_compressed_stream_fails() -> TestResult<()> { + const CSV: &[u8] = b"name\nAlice\n"; + + let (_env, client) = create_test_environment(&[]).await?; + + put_object(&client, "invalid.csv.gz", CSV).await?; + let invalid = csv_select_request(&client, "invalid.csv.gz", CompressionType::Gzip, "SELECT * FROM S3Object") + .send() + .await + .expect_err("invalid GZIP header must fail before streaming"); + assert_eq!( + invalid.as_service_error().and_then(ProvideErrorMetadata::code), + Some("InvalidCompressionFormat") + ); + + put_object(&client, "empty.csv.gz", b"").await?; + let empty = csv_select_request(&client, "empty.csv.gz", CompressionType::Gzip, "SELECT * FROM S3Object") + .send() + .await + .expect_err("empty GZIP input must fail as truncated"); + assert_eq!(empty.as_service_error().and_then(ProvideErrorMetadata::code), Some("TruncatedInput")); + + let mut truncated = bzip2(CSV).await?; + truncated.pop(); + put_object(&client, "truncated.csv.bz2", &truncated).await?; + let truncated = csv_select_request(&client, "truncated.csv.bz2", CompressionType::Bzip2, "SELECT * FROM S3Object") + .send() + .await?; + assert_truncated_stream_failure(truncated).await?; + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn test_select_object_content_compressed_disconnect_releases_query() -> TestResult<()> { + const OBJECT: &str = "disconnect.csv.gz"; + const ROWS: usize = 16 * 1024; + const RELEASE_ATTEMPTS: usize = 20; + const RELEASE_BACKOFF: Duration = Duration::from_millis(25); + + let (_env, client) = create_test_environment(&[("RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES", "1")]).await?; + let row = format!("{}\n", "x".repeat(1023)); + let mut body = Vec::with_capacity("value\n".len() + ROWS * row.len()); + body.extend_from_slice(b"value\n"); + for _ in 0..ROWS { + body.extend_from_slice(row.as_bytes()); + } + let compressed = gzip(&body)?; + put_object(&client, OBJECT, &compressed).await?; + + let first = csv_select_request(&client, OBJECT, CompressionType::Gzip, "SELECT * FROM S3Object") + .send() + .await?; + let saturated = csv_select_request(&client, OBJECT, CompressionType::Gzip, "SELECT * FROM S3Object") + .send() + .await + .expect_err("the unread compressed response should retain the only query permit"); + assert_eq!(saturated.as_service_error().and_then(ProvideErrorMetadata::code), Some("SlowDown")); + + drop(first); + let second = tokio::time::timeout(Duration::from_secs(5), async { + for attempt in 0..RELEASE_ATTEMPTS { + match csv_select_request(&client, OBJECT, CompressionType::Gzip, "SELECT * FROM S3Object") + .send() + .await + { + Ok(response) => return Ok::<_, Box>(response), + Err(error) + if error.as_service_error().and_then(ProvideErrorMetadata::code) == Some("SlowDown") + && attempt + 1 < RELEASE_ATTEMPTS => + { + tokio::time::sleep(RELEASE_BACKOFF).await; + } + Err(error) if error.as_service_error().and_then(ProvideErrorMetadata::code) == Some("SlowDown") => { + return Err("disconnected compressed Select retained its query permit".into()); + } + Err(error) => return Err(format!("unexpected Select error after disconnect: {error}").into()), + } + } + Err("query permit release retry loop ended unexpectedly".into()) + }) + .await + .map_err(|_| -> Box { "compressed Select did not release its query permit".into() })??; + drop(second); + + Ok(()) +} diff --git a/crates/s3select-api/Cargo.toml b/crates/s3select-api/Cargo.toml index a68ce1a1d..38c9eabf5 100644 --- a/crates/s3select-api/Cargo.toml +++ b/crates/s3select-api/Cargo.toml @@ -60,21 +60,26 @@ hotpath-cpu = [ [dependencies] hotpath.workspace = true metrics = { workspace = true } +async-compression = { workspace = true, features = ["tokio", "gzip", "bzip2"] } async-trait.workspace = true +arc-swap.workspace = true bytes = { workspace = true, features = ["serde"] } chrono = { workspace = true, features = ["serde"] } +crc-fast.workspace = true rustfs-common.workspace = true datafusion = { workspace = true, default-features = false, features = ["parquet", "recursive_protection", "sql"] } rustfs-ecstore.workspace = true rustfs-storage-api.workspace = true futures = { workspace = true } futures-core = { workspace = true } +flate2.workspace = true http.workspace = true s3s = { workspace = true, features = ["minio"] } serde_json = { workspace = true, features = ["raw_value"] } thiserror = { workspace = true } parking_lot.workspace = true tokio = { workspace = true, features = ["fs", "rt-multi-thread"] } +tokio-stream.workspace = true tokio-util = { workspace = true, features = ["io", "compat"] } tracing.workspace = true uuid.workspace = true diff --git a/crates/s3select-api/src/input_stream.rs b/crates/s3select-api/src/input_stream.rs new file mode 100644 index 000000000..ff2b501dd --- /dev/null +++ b/crates/s3select-api/src/input_stream.rs @@ -0,0 +1,1570 @@ +// Copyright 2024 RustFS Team +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +use crate::{ + MAX_ERROR_SOURCE_DEPTH, SelectError, SelectInputMetrics, metrics::SelectInputMetricsRecorder, + query::session::QueryExecutionGuard, +}; +use async_compression::tokio::bufread::BzDecoder; +use bytes::{Buf as _, Bytes}; +use datafusion::object_store::{Error as ObjectStoreError, Result as ObjectStoreResult}; +use flate2::bufread::GzDecoder; +use futures::{StreamExt, stream}; +use futures_core::stream::BoxStream; +use std::{ + error::Error as StdError, + io::{self, BufRead as _, Read as _}, + pin::Pin, + sync::Arc, + task::{Context, Poll}, +}; +use tokio::{ + io::{AsyncRead, AsyncReadExt, BufReader, ReadBuf}, + sync::{mpsc, oneshot}, +}; +use tokio_stream::wrappers::ReceiverStream; +use tokio_util::io::{ReaderStream, StreamReader}; + +pub(crate) const MAX_SELECT_RECORD_BYTES: usize = 1024 * 1024; +const MAX_SELECT_PROCESSED_BYTES: u64 = 5 * 1024 * 1024 * 1024 * 1024; +pub(crate) const SELECT_DECODE_CHUNK_BYTES: usize = 64 * 1024; +const DECOMPRESSION_CHANNEL_CAPACITY: usize = 2; + +pub(crate) type SelectInputReader = Box; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) enum CompressionFormat { + Gzip, + Bzip2, +} + +impl CompressionFormat { + fn name(self) -> &'static str { + match self { + Self::Gzip => "GZIP", + Self::Bzip2 => "BZIP2", + } + } + + fn invalid_header_error(self) -> SelectError { + SelectError::InvalidCompressionFormatForObject { + compression: self.name(), + } + } +} + +pub(crate) fn processed_bytes_limit() -> u64 { + // Processed throughput is independent of compression ratio and live + // memory; use the S3 Select object-size ceiling as the absolute bound. + MAX_SELECT_PROCESSED_BYTES +} + +pub(crate) fn compressed_input_reader( + reader: SelectInputReader, + compressed_size: u64, + format: CompressionFormat, + input_metrics: Arc, + max_processed_bytes: u64, + query_guard: Option, +) -> SelectInputReader { + // rustfs-zip does not expose Select's streaming metrics, member validation, + // typed errors, or cancellation contract, so the protocol adapter lives here. + let input_metrics = input_metrics.recorder(); + let reader = ScannedReader::new(reader, compressed_size, input_metrics.clone()); + let reader = CooperativeReader::new(reader); + let decoder = match format { + CompressionFormat::Gzip => blocking_gzip_reader(Box::new(reader), query_guard), + CompressionFormat::Bzip2 => blocking_bzip2_reader(Box::new(Bzip2HeaderValidatingReader::new(reader)), query_guard), + }; + Box::new(CooperativeReader::new(ProcessedReader::new(decoder, input_metrics, max_processed_bytes))) +} + +fn blocking_gzip_reader(reader: SelectInputReader, query_guard: Option) -> SelectInputReader { + // The bounded bridge keeps RFC 1952 decoding off Tokio workers while + // preserving member validation, backpressure, and reader cancellation. + let (compressed_tx, compressed_rx) = mpsc::channel(DECOMPRESSION_CHANNEL_CAPACITY); + let (decoded_tx, decoded_rx) = mpsc::channel(DECOMPRESSION_CHANNEL_CAPACITY); + + let decoded_closed = decoded_tx.clone(); + drop(tokio::spawn(async move { + let mut stream = ReaderStream::with_capacity(reader, SELECT_DECODE_CHUNK_BYTES); + loop { + let item = tokio::select! { + biased; + _ = decoded_closed.closed() => break, + item = stream.next() => item, + }; + let Some(item) = item else { + break; + }; + let sent = tokio::select! { + biased; + _ = decoded_closed.closed() => false, + result = compressed_tx.send(item) => result.is_ok(), + }; + if !sent { + break; + } + } + })); + + let blocking_output = decoded_tx.clone(); + spawn_decoder_thread("s3select-gzip", decoded_tx, query_guard, move || { + decode_gzip(compressed_rx, blocking_output) + }); + + Box::new(StreamReader::new(ReceiverStream::new(decoded_rx))) +} + +fn blocking_bzip2_reader(reader: SelectInputReader, query_guard: Option) -> SelectInputReader { + let (decoded_tx, decoded_rx) = mpsc::channel(DECOMPRESSION_CHANNEL_CAPACITY); + let runtime = tokio::runtime::Handle::current(); + let blocking_output = decoded_tx.clone(); + // A dedicated thread avoids occupying Tokio's blocking pool while the + // decoder waits for asynchronous object reads. Query admission bounds the + // number of concurrent Select decoder threads. + spawn_decoder_thread("s3select-bzip2", decoded_tx, query_guard, move || { + runtime.block_on(decode_bzip2(reader, blocking_output)) + }); + Box::new(StreamReader::new(ReceiverStream::new(decoded_rx))) +} + +fn spawn_decoder_thread( + name: &'static str, + decoded: mpsc::Sender>, + query_guard: Option, + task: impl FnOnce() -> io::Result<()> + Send + 'static, +) { + let (finished_tx, finished_rx) = oneshot::channel(); + let spawn_result = std::thread::Builder::new().name(name.to_string()).spawn(move || { + let _query_guard = query_guard; + let _ = finished_tx.send(task()); + }); + if spawn_result.is_err() { + let _ = decoded.try_send(Err(io::Error::other(SelectError::InternalError))); + return; + } + + drop(tokio::spawn(async move { + let result = finished_rx + .await + .unwrap_or_else(|_| Err(io::Error::other(SelectError::InternalError))); + if let Err(error) = result { + let _ = decoded.send(Err(error)).await; + } + })); +} + +async fn decode_bzip2(reader: SelectInputReader, decoded: mpsc::Sender>) -> io::Result<()> { + let reader = BufReader::with_capacity(SELECT_DECODE_CHUNK_BYTES, reader); + let mut decoder = BzDecoder::new(reader); + decoder.multiple_members(true); + let mut buffer = vec![0; SELECT_DECODE_CHUNK_BYTES]; + loop { + let read = tokio::select! { + biased; + _ = decoded.closed() => return Ok(()), + result = decoder.read(&mut buffer) => result?, + }; + if read == 0 { + return Ok(()); + } + let bytes = Bytes::copy_from_slice(&buffer[..read]); + tokio::select! { + biased; + _ = decoded.closed() => return Ok(()), + result = decoded.send(Ok(bytes)) => { + if result.is_err() { + return Ok(()); + } + } + } + } +} + +fn decode_gzip(compressed: mpsc::Receiver>, decoded: mpsc::Sender>) -> io::Result<()> { + let mut reader = io::BufReader::with_capacity(SELECT_DECODE_CHUNK_BYTES, BlockingChannelReader::new(compressed)); + let mut buffer = vec![0; SELECT_DECODE_CHUNK_BYTES]; + let mut decoded_member = false; + loop { + if decoded.is_closed() { + return Ok(()); + } + + if reader.fill_buf()?.is_empty() { + return if decoded_member { + Ok(()) + } else { + Err(io::Error::new(io::ErrorKind::UnexpectedEof, SelectError::TruncatedInput)) + }; + } + let header = read_gzip_header(&mut reader).map_err(|error| { + if decoded_member + && find_error_source::(&error) + .is_some_and(|source| matches!(source, SelectError::InvalidCompressionFormatForObject { .. })) + { + io::Error::new(error.kind(), SelectError::TruncatedInput) + } else { + error + } + })?; + let mut decoder = GzDecoder::new(std::io::Read::chain(io::Cursor::new(header), reader)); + loop { + let read = decoder.read(&mut buffer)?; + if read == 0 { + break; + } + if decoded.blocking_send(Ok(Bytes::copy_from_slice(&buffer[..read]))).is_err() { + return Ok(()); + } + } + let (_, remaining) = decoder.into_inner().into_inner(); + reader = remaining; + decoded_member = true; + } +} + +fn read_gzip_header(reader: &mut R) -> io::Result<[u8; 10]> { + let mut fixed = [0; 10]; + read_gzip_exact(reader, &mut fixed)?; + if fixed[..3] != [0x1f, 0x8b, 0x08] || fixed[3] & 0xe0 != 0 { + return Err(invalid_gzip_header_error()); + } + let flags = fixed[3]; + let mut header_crc = (flags & GZIP_FLAG_HEADER_CRC != 0).then(|| crc_fast::Digest::new(crc_fast::CrcAlgorithm::Crc32IsoHdlc)); + update_gzip_header_crc(&mut header_crc, &fixed); + + if flags & GZIP_FLAG_EXTRA != 0 { + let mut length = [0; 2]; + read_gzip_exact(reader, &mut length)?; + update_gzip_header_crc(&mut header_crc, &length); + let extra_len = usize::from(u16::from_le_bytes(length)); + read_gzip_header_bytes(reader, extra_len, &mut header_crc)?; + } + if flags & GZIP_FLAG_NAME != 0 { + read_gzip_text_field(reader, &mut header_crc)?; + } + if flags & GZIP_FLAG_COMMENT != 0 { + read_gzip_text_field(reader, &mut header_crc)?; + } + if let Some(digest) = header_crc { + let expected = + u16::try_from(digest.finalize() & u64::from(u16::MAX)).map_err(|_| io::Error::other(SelectError::InternalError))?; + let mut actual = [0; 2]; + read_gzip_exact(reader, &mut actual)?; + if u16::from_le_bytes(actual) != expected { + return Err(invalid_gzip_header_error()); + } + } + + fixed[3] &= GZIP_FLAG_TEXT; + Ok(fixed) +} + +fn read_gzip_exact(reader: &mut R, bytes: &mut [u8]) -> io::Result<()> { + reader.read_exact(bytes).map_err(|error| { + if error.kind() == io::ErrorKind::UnexpectedEof && !error_chain_contains::(&error) { + io::Error::new(io::ErrorKind::UnexpectedEof, SelectError::TruncatedInput) + } else { + error + } + }) +} + +fn read_gzip_header_bytes( + reader: &mut R, + mut remaining: usize, + header_crc: &mut Option, +) -> io::Result<()> { + while remaining > 0 { + let available = reader.fill_buf()?; + if available.is_empty() { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, SelectError::TruncatedInput)); + } + let consumed = remaining.min(available.len()); + update_gzip_header_crc(header_crc, &available[..consumed]); + reader.consume(consumed); + remaining -= consumed; + } + Ok(()) +} + +fn read_gzip_text_field(reader: &mut R, header_crc: &mut Option) -> io::Result<()> { + loop { + let available = reader.fill_buf()?; + if available.is_empty() { + return Err(io::Error::new(io::ErrorKind::UnexpectedEof, SelectError::TruncatedInput)); + } + let terminator = available.iter().position(|byte| *byte == 0); + let consumed = terminator.map_or(available.len(), |position| position + 1); + update_gzip_header_crc(header_crc, &available[..consumed]); + reader.consume(consumed); + if terminator.is_some() { + return Ok(()); + } + } +} + +fn update_gzip_header_crc(header_crc: &mut Option, bytes: &[u8]) { + if let Some(digest) = header_crc { + digest.update(bytes); + } +} + +fn invalid_gzip_header_error() -> io::Error { + io::Error::new(io::ErrorKind::InvalidData, CompressionFormat::Gzip.invalid_header_error()) +} + +struct BlockingChannelReader { + receiver: mpsc::Receiver>, + current: Bytes, +} + +impl BlockingChannelReader { + fn new(receiver: mpsc::Receiver>) -> Self { + Self { + receiver, + current: Bytes::new(), + } + } +} + +impl io::Read for BlockingChannelReader { + fn read(&mut self, buffer: &mut [u8]) -> io::Result { + while self.current.is_empty() { + match self.receiver.blocking_recv() { + Some(Ok(bytes)) => self.current = bytes, + Some(Err(error)) => return Err(error), + None => return Ok(0), + } + } + let read = buffer.len().min(self.current.len()); + buffer[..read].copy_from_slice(&self.current[..read]); + self.current.advance(read); + Ok(read) + } +} + +pub(crate) fn compressed_input_stream( + reader: SelectInputReader, + compressed_size: u64, + format: CompressionFormat, + input_metrics: Arc, + record_delimiter: Vec, + max_processed_bytes: u64, + query_guard: Option, +) -> ObjectStoreResult>> { + let record_size = RecordSizeTracker::new(record_delimiter).map_err(select_object_store_error)?; + let reader = compressed_input_reader(reader, compressed_size, format, input_metrics, max_processed_bytes, query_guard); + let stream = ReaderStream::with_capacity(reader, SELECT_DECODE_CHUNK_BYTES); + Ok(stream::try_unfold((stream, record_size), |(mut stream, mut record_size)| async move { + match stream.next().await { + Some(Ok(bytes)) => { + record_size.observe(&bytes).map_err(select_object_store_error)?; + Ok(Some((bytes, (stream, record_size)))) + } + Some(Err(error)) => Err(input_io_error(error)), + None => { + record_size.finish().map_err(select_object_store_error)?; + Ok(None) + } + } + }) + .boxed()) +} + +pub(crate) fn input_io_error(source: io::Error) -> ObjectStoreError { + let source: Box = match find_error_source::(&source) { + Some(error) => Box::new(error.clone()), + None => Box::new(source), + }; + ObjectStoreError::Generic { + store: "EcObjectStore", + source, + } +} + +fn select_object_store_error(source: SelectError) -> ObjectStoreError { + ObjectStoreError::Generic { + store: "EcObjectStore", + source: Box::new(source), + } +} + +#[derive(Debug, thiserror::Error)] +#[error("compressed object source read failed")] +struct CompressedSourceReadError { + #[source] + source: io::Error, +} + +struct ScannedReader { + inner: tokio::io::Take, + input_metrics: SelectInputMetricsRecorder, +} + +impl ScannedReader { + fn new(reader: R, compressed_size: u64, input_metrics: SelectInputMetricsRecorder) -> Self { + Self { + inner: reader.take(compressed_size), + input_metrics, + } + } +} + +impl AsyncRead for ScannedReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + let before = buf.filled().len(); + match Pin::new(&mut self.inner).poll_read(cx, buf) { + Poll::Pending => Poll::Pending, + Poll::Ready(Err(source)) => Poll::Ready(Err(source_read_error(source))), + Poll::Ready(Ok(())) => { + let read = buf.filled().len() - before; + if read == 0 && self.inner.limit() > 0 { + let source = io::Error::new( + io::ErrorKind::UnexpectedEof, + format!("compressed object stream ended with {} bytes remaining", self.inner.limit()), + ); + return Poll::Ready(Err(source_read_error(source))); + } + self.input_metrics.record_scanned(read); + Poll::Ready(Ok(())) + } + } + } +} + +struct CooperativeReader { + inner: R, + bytes_since_yield: usize, + yield_pending: bool, +} + +impl CooperativeReader { + fn new(inner: R) -> Self { + Self { + inner, + bytes_since_yield: 0, + yield_pending: false, + } + } +} + +impl AsyncRead for CooperativeReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + if self.yield_pending { + self.yield_pending = false; + self.bytes_since_yield = 0; + cx.waker().wake_by_ref(); + return Poll::Pending; + } + if buf.remaining() == 0 { + return Poll::Ready(Ok(())); + } + + let remaining_budget = SELECT_DECODE_CHUNK_BYTES - self.bytes_since_yield; + let read_limit = remaining_budget.min(buf.remaining()); + let read = { + let unfilled = buf.initialize_unfilled_to(read_limit); + let mut limited = ReadBuf::new(unfilled); + match Pin::new(&mut self.inner).poll_read(cx, &mut limited) { + Poll::Pending => return Poll::Pending, + Poll::Ready(Err(error)) => return Poll::Ready(Err(error)), + Poll::Ready(Ok(())) => limited.filled().len(), + } + }; + buf.advance(read); + self.bytes_since_yield += read; + self.yield_pending = read > 0 && self.bytes_since_yield == SELECT_DECODE_CHUNK_BYTES; + Poll::Ready(Ok(())) + } +} + +fn source_read_error(source: io::Error) -> io::Error { + let kind = source.kind(); + io::Error::new(kind, CompressedSourceReadError { source }) +} + +struct Bzip2HeaderValidatingReader { + inner: R, + position: usize, + pending_error: Option, +} + +impl Bzip2HeaderValidatingReader { + fn new(inner: R) -> Self { + Self { + inner, + position: 0, + pending_error: None, + } + } +} + +impl AsyncRead for Bzip2HeaderValidatingReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + if let Some(error) = self.pending_error.take() { + return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, error))); + } + if buf.remaining() == 0 { + return Poll::Ready(Ok(())); + } + if self.position == 4 { + return Pin::new(&mut self.inner).poll_read(cx, buf); + } + + let before = buf.filled().len(); + match Pin::new(&mut self.inner).poll_read(cx, buf) { + Poll::Pending => Poll::Pending, + Poll::Ready(Err(error)) => Poll::Ready(Err(error)), + Poll::Ready(Ok(())) => { + let after = buf.filled().len(); + if after == before { + return Poll::Ready(Err(io::Error::new(io::ErrorKind::UnexpectedEof, SelectError::TruncatedInput))); + } + if let Err(error) = validate_bzip2_header(&mut self.position, &buf.filled()[before..after]) { + buf.set_filled(before + error.offset); + if error.offset == 0 { + return Poll::Ready(Err(io::Error::new(io::ErrorKind::InvalidData, error.source))); + } + self.pending_error = Some(error.source); + } + Poll::Ready(Ok(())) + } + } + } +} + +struct HeaderValidationError { + offset: usize, + source: SelectError, +} + +fn validate_bzip2_header(position: &mut usize, bytes: &[u8]) -> Result<(), HeaderValidationError> { + for (offset, byte) in bytes.iter().copied().enumerate() { + let valid = match *position { + 0 => byte == b'B', + 1 => byte == b'Z', + 2 => byte == b'h', + 3 => matches!(byte, b'1'..=b'9'), + _ => break, + }; + if !valid { + return Err(HeaderValidationError { + offset, + source: CompressionFormat::Bzip2.invalid_header_error(), + }); + } + *position += 1; + } + Ok(()) +} + +const GZIP_FLAG_HEADER_CRC: u8 = 0x02; +const GZIP_FLAG_EXTRA: u8 = 0x04; +const GZIP_FLAG_NAME: u8 = 0x08; +const GZIP_FLAG_COMMENT: u8 = 0x10; +const GZIP_FLAG_TEXT: u8 = 0x01; + +struct ProcessedReader { + inner: R, + input_metrics: SelectInputMetricsRecorder, + processed_bytes: u64, + max_processed_bytes: u64, +} + +impl ProcessedReader { + fn new(inner: R, input_metrics: SelectInputMetricsRecorder, max_processed_bytes: u64) -> Self { + Self { + inner, + input_metrics, + processed_bytes: 0, + max_processed_bytes, + } + } +} + +impl AsyncRead for ProcessedReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + let before = buf.filled().len(); + match Pin::new(&mut self.inner).poll_read(cx, buf) { + Poll::Pending => Poll::Pending, + Poll::Ready(Err(error)) => Poll::Ready(Err(classify_decoder_error(error))), + Poll::Ready(Ok(())) => { + let read = buf.filled().len() - before; + let read = u64::try_from(read).unwrap_or(u64::MAX); + let Some(processed_bytes) = self.processed_bytes.checked_add(read) else { + return Poll::Ready(Err(processed_bytes_limit_error())); + }; + if processed_bytes > self.max_processed_bytes { + return Poll::Ready(Err(processed_bytes_limit_error())); + } + self.processed_bytes = processed_bytes; + self.input_metrics.record_processed(buf.filled().len() - before); + Poll::Ready(Ok(())) + } + } + } +} + +fn processed_bytes_limit_error() -> io::Error { + io::Error::new(io::ErrorKind::OutOfMemory, SelectError::ResourceExhausted) +} + +fn classify_decoder_error(error: io::Error) -> io::Error { + if error_chain_contains::(&error) || error_chain_contains::(&error) { + return error; + } + + let select_error = if error.kind() == io::ErrorKind::OutOfMemory { + SelectError::ResourceExhausted + } else { + SelectError::TruncatedInput + }; + io::Error::new(error.kind(), select_error) +} + +fn error_chain_contains(error: &(dyn StdError + 'static)) -> bool { + find_error_source::(error).is_some() +} + +fn find_error_source<'a, T: StdError + 'static>(error: &'a (dyn StdError + 'static)) -> Option<&'a T> { + let mut current = Some(error); + for _ in 0..MAX_ERROR_SOURCE_DEPTH { + let Some(error) = current else { + break; + }; + if let Some(error) = error.downcast_ref::() { + return Some(error); + } + current = error + .downcast_ref::() + .and_then(|error| error.get_ref()) + .map(|source| source as &(dyn StdError + 'static)) + .or_else(|| error.source()); + } + None +} + +struct RecordSizeTracker { + delimiter: Vec, + prefix: Vec, + record_bytes: usize, + matched: usize, +} + +impl RecordSizeTracker { + fn new(delimiter: Vec) -> Result { + if delimiter.is_empty() { + return Err(SelectError::InvalidDataSource); + } + + let mut prefix = vec![0; delimiter.len()]; + let mut matched = 0; + for index in 1..delimiter.len() { + while matched > 0 && delimiter[index] != delimiter[matched] { + matched = prefix[matched - 1]; + } + if delimiter[index] == delimiter[matched] { + matched += 1; + } + prefix[index] = matched; + } + + Ok(Self { + delimiter, + prefix, + record_bytes: 0, + matched: 0, + }) + } + + fn observe(&mut self, bytes: &[u8]) -> Result<(), SelectError> { + if self.delimiter.len() == 1 { + return self.observe_single_byte_delimiter(bytes); + } + + for &byte in bytes { + self.record_bytes = self.record_bytes.checked_add(1).ok_or(SelectError::OverMaxRecordSize)?; + while self.matched > 0 && byte != self.delimiter[self.matched] { + self.matched = self.prefix[self.matched - 1]; + } + if byte == self.delimiter[self.matched] { + self.matched += 1; + } + if self.matched == self.delimiter.len() { + let payload_bytes = self.record_bytes - self.delimiter.len(); + if payload_bytes > MAX_SELECT_RECORD_BYTES { + return Err(SelectError::OverMaxRecordSize); + } + self.record_bytes = 0; + self.matched = 0; + } else if self.record_bytes - self.matched > MAX_SELECT_RECORD_BYTES { + return Err(SelectError::OverMaxRecordSize); + } + } + Ok(()) + } + + fn observe_single_byte_delimiter(&mut self, mut bytes: &[u8]) -> Result<(), SelectError> { + let delimiter = self.delimiter[0]; + while let Some(index) = bytes.iter().position(|byte| *byte == delimiter) { + let payload_bytes = self.record_bytes.checked_add(index).ok_or(SelectError::OverMaxRecordSize)?; + if payload_bytes > MAX_SELECT_RECORD_BYTES { + return Err(SelectError::OverMaxRecordSize); + } + self.record_bytes = 0; + bytes = &bytes[index + 1..]; + } + self.record_bytes = self + .record_bytes + .checked_add(bytes.len()) + .ok_or(SelectError::OverMaxRecordSize)?; + if self.record_bytes > MAX_SELECT_RECORD_BYTES { + return Err(SelectError::OverMaxRecordSize); + } + Ok(()) + } + + fn finish(&self) -> Result<(), SelectError> { + if self.record_bytes > MAX_SELECT_RECORD_BYTES { + Err(SelectError::OverMaxRecordSize) + } else { + Ok(()) + } + } +} + +#[cfg(test)] +pub(crate) async fn encode_compressed_fixture(format: CompressionFormat, input: &[u8]) -> Vec { + use tokio::io::AsyncWriteExt as _; + + let cursor = std::io::Cursor::new(Vec::new()); + match format { + CompressionFormat::Gzip => { + let mut encoder = async_compression::tokio::write::GzipEncoder::new(cursor); + encoder.write_all(input).await.expect("gzip fixture should encode"); + encoder.shutdown().await.expect("gzip fixture should finish"); + encoder.into_inner().into_inner() + } + CompressionFormat::Bzip2 => { + let mut encoder = async_compression::tokio::write::BzEncoder::new(cursor); + encoder.write_all(input).await.expect("bzip2 fixture should encode"); + encoder.shutdown().await.expect("bzip2 fixture should finish"); + encoder.into_inner().into_inner() + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::SelectInputMetricsSnapshot; + use futures::TryStreamExt; + use std::{io::Cursor, sync::Mutex as StdMutex, thread::ThreadId}; + use tokio::io::{AsyncWriteExt, DuplexStream}; + + struct ThreadRecordingReader { + inner: Cursor>, + thread_id: Arc>>, + } + + impl AsyncRead for ThreadRecordingReader { + fn poll_read(self: Pin<&mut Self>, cx: &mut Context<'_>, buffer: &mut ReadBuf<'_>) -> Poll> { + let this = self.get_mut(); + let mut thread_id = this.thread_id.lock().expect("thread recorder mutex should not be poisoned"); + thread_id.get_or_insert_with(|| std::thread::current().id()); + drop(thread_id); + Pin::new(&mut this.inner).poll_read(cx, buffer) + } + } + + struct ErrorAfterReader { + inner: Cursor>, + end: u64, + failed: bool, + } + + impl AsyncRead for ErrorAfterReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buffer: &mut ReadBuf<'_>) -> Poll> { + if self.inner.position() < self.end { + return Pin::new(&mut self.inner).poll_read(cx, buffer); + } + if !self.failed { + self.failed = true; + return Poll::Ready(Err(io::Error::new(io::ErrorKind::ConnectionReset, "injected source failure"))); + } + Poll::Ready(Ok(())) + } + } + + async fn decode(format: CompressionFormat, compressed: Vec) -> (ObjectStoreResult>, SelectInputMetricsSnapshot) { + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let metrics = Arc::new(SelectInputMetrics::default()); + let stream = compressed_input_stream( + Box::new(Cursor::new(compressed)), + compressed_len, + format, + Arc::clone(&metrics), + b"\n".to_vec(), + u64::MAX, + None, + ) + .expect("record delimiter should be valid"); + let result = stream.try_collect::>().await.map(|chunks| chunks.concat()); + (result, metrics.snapshot()) + } + + fn select_error(error: &ObjectStoreError) -> Option { + find_error_source::(error).cloned() + } + + #[test] + fn protocol_limits_match_the_s3_select_contract() { + assert_eq!(MAX_SELECT_RECORD_BYTES, 1_048_576); + assert_eq!(MAX_SELECT_PROCESSED_BYTES, 5_497_558_138_880); + } + + #[tokio::test] + async fn gzip_and_bzip2_preserve_bytes_and_metric_boundaries() { + const INPUT: &[u8] = b"name,age\nAlice,30\n"; + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let compressed = encode_compressed_fixture(format, INPUT).await; + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let (decoded, metrics) = decode(format, compressed).await; + + assert_eq!(decoded.expect("valid compressed input should decode"), INPUT); + assert_eq!(metrics.bytes_scanned, compressed_len); + assert_eq!( + metrics.bytes_processed, + u64::try_from(INPUT.len()).expect("fixture length should fit in u64") + ); + } + } + + #[tokio::test] + async fn source_read_errors_are_not_reclassified_as_truncated_input() { + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let compressed = encode_compressed_fixture(format, b"name\nAlice\n").await; + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let expected_size = u64::try_from(compressed.len() + 1).expect("fixture length should fit in u64"); + let stream = compressed_input_stream( + Box::new(ErrorAfterReader { + inner: Cursor::new(compressed), + end: compressed_len, + failed: false, + }), + expected_size, + format, + Arc::new(SelectInputMetrics::default()), + b"\n".to_vec(), + u64::MAX, + None, + ) + .expect("record delimiter should be valid"); + + let error = stream + .try_collect::>() + .await + .expect_err("a storage read failure must terminate decoding"); + + assert_eq!(select_error(&error), None, "{format:?} must not report TruncatedInput: {error:?}"); + assert!( + find_error_source::(&error).is_some(), + "{format:?} must preserve the source error class" + ); + } + } + + #[tokio::test] + async fn source_read_error_inside_gzip_header_is_not_truncated_input() { + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, b"name\nAlice\n").await; + let partial_header = compressed[..5].to_vec(); + let stream = compressed_input_stream( + Box::new(Cursor::new(partial_header)), + 10, + CompressionFormat::Gzip, + Arc::new(SelectInputMetrics::default()), + b"\n".to_vec(), + u64::MAX, + None, + ) + .expect("record delimiter should be valid"); + + let error = stream + .try_collect::>() + .await + .expect_err("a source failure inside the GZIP header must terminate decoding"); + + assert_eq!(select_error(&error), None, "source failure must not report TruncatedInput: {error:?}"); + assert!( + find_error_source::(&error).is_some(), + "the source error class must survive GZIP header validation" + ); + } + + #[tokio::test] + async fn concatenated_members_decode_in_order() { + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let mut compressed = encode_compressed_fixture(format, b"a\n").await; + compressed.extend_from_slice(&encode_compressed_fixture(format, b"b\n").await); + let (decoded, _) = decode(format, compressed).await; + assert_eq!(decoded.expect("valid concatenated members should decode"), b"a\nb\n"); + } + + let mut compressed = encode_compressed_fixture(CompressionFormat::Gzip, b"a\n").await; + let (second, _) = gzip_with_optional_header_crc(b"b\n").await; + compressed.extend_from_slice(&second); + let (decoded, _) = decode(CompressionFormat::Gzip, compressed).await; + assert_eq!(decoded.expect("optional headers must work in later GZIP members"), b"a\nb\n"); + } + + #[tokio::test] + async fn invalid_initial_headers_are_typed_as_invalid_compression() { + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let (result, _) = decode(format, b"not compressed\n".to_vec()).await; + let error = result.expect_err("invalid compression header must fail"); + let compression = format.name(); + assert_eq!(select_error(&error), Some(SelectError::InvalidCompressionFormatForObject { compression })); + assert_eq!( + select_error(&error).expect("typed compression error").to_string(), + format!("{compression} is not applicable to the queried object. Please correct the request and try again.") + ); + } + } + + #[tokio::test] + async fn fixed_header_variants_are_validated_before_decoding() { + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, b"name\nAlice\n").await; + for (position, invalid) in [(2, 0), (3, 0x20)] { + let mut mutated = encoded.clone(); + mutated[position] = invalid; + let (result, _) = decode(CompressionFormat::Gzip, mutated).await; + assert_eq!( + select_error(&result.expect_err("invalid GZIP fixed header must fail")), + Some(CompressionFormat::Gzip.invalid_header_error()) + ); + } + + let (result, _) = decode(CompressionFormat::Bzip2, b"BZh0".to_vec()).await; + assert_eq!( + select_error(&result.expect_err("invalid BZIP2 block size must fail")), + Some(CompressionFormat::Bzip2.invalid_header_error()) + ); + } + + #[tokio::test] + async fn incomplete_valid_headers_are_typed_as_truncated_input() { + for (format, prefixes) in [ + (CompressionFormat::Gzip, vec![b"".as_slice(), b"\x1f", b"\x1f\x8b", b"\x1f\x8b\x08"]), + (CompressionFormat::Bzip2, vec![b"".as_slice(), b"B", b"BZ", b"BZh"]), + ] { + for prefix in prefixes { + let (result, _) = decode(format, prefix.to_vec()).await; + assert_eq!( + select_error(&result.expect_err("incomplete compression header must fail")), + Some(SelectError::TruncatedInput) + ); + } + } + } + + async fn gzip_with_optional_header_crc(input: &[u8]) -> (Vec, usize) { + const EXTRA: &[u8] = b"s3-select"; + const FILE_NAME: &[u8] = b"select.csv"; + const COMMENT: &[u8] = b"fixture"; + + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, input).await; + let mut header = encoded[..10].to_vec(); + header[3] = GZIP_FLAG_EXTRA | GZIP_FLAG_NAME | GZIP_FLAG_COMMENT | GZIP_FLAG_HEADER_CRC; + header.extend_from_slice( + &u16::try_from(EXTRA.len()) + .expect("GZIP extra fixture should fit its length field") + .to_le_bytes(), + ); + header.extend_from_slice(EXTRA); + header.extend_from_slice(FILE_NAME); + header.push(0); + header.extend_from_slice(COMMENT); + header.push(0); + let mut digest = crc_fast::Digest::new(crc_fast::CrcAlgorithm::Crc32IsoHdlc); + digest.update(&header); + let crc = digest.finalize(); + let crc = u16::try_from(crc & u64::from(u16::MAX)).expect("masked header CRC should fit in u16"); + let crc_offset = header.len(); + header.extend_from_slice(&crc.to_le_bytes()); + header.extend_from_slice(&encoded[10..]); + (header, crc_offset) + } + + async fn gzip_with_text_header(input: &[u8], flag: u8, field_bytes: usize) -> Vec { + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, input).await; + let mut member = encoded[..10].to_vec(); + member[3] = flag; + member.extend(std::iter::repeat_n(b'x', field_bytes)); + member.push(0); + member.extend_from_slice(&encoded[10..]); + member + } + + async fn gzip_with_text_header_crc(input: &[u8], flag: u8, field_bytes: usize) -> (Vec, usize) { + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, input).await; + let mut member = encoded[..10].to_vec(); + member[3] = flag | GZIP_FLAG_HEADER_CRC; + member.extend(std::iter::repeat_n(b'x', field_bytes)); + member.push(0); + let mut digest = crc_fast::Digest::new(crc_fast::CrcAlgorithm::Crc32IsoHdlc); + digest.update(&member); + let crc = u16::try_from(digest.finalize() & u64::from(u16::MAX)).expect("masked header CRC should fit in u16"); + let crc_offset = member.len(); + member.extend_from_slice(&crc.to_le_bytes()); + member.extend_from_slice(&encoded[10..]); + (member, crc_offset) + } + + async fn gzip_with_extra_header_crc(input: &[u8], extra_bytes: usize) -> (Vec, usize) { + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, input).await; + let mut member = encoded[..10].to_vec(); + member[3] = GZIP_FLAG_EXTRA | GZIP_FLAG_HEADER_CRC; + member.extend_from_slice( + &u16::try_from(extra_bytes) + .expect("GZIP extra fixture should fit its length field") + .to_le_bytes(), + ); + member.extend(std::iter::repeat_n(b'x', extra_bytes)); + let mut digest = crc_fast::Digest::new(crc_fast::CrcAlgorithm::Crc32IsoHdlc); + digest.update(&member); + let crc = u16::try_from(digest.finalize() & u64::from(u16::MAX)).expect("masked header CRC should fit in u16"); + let crc_offset = member.len(); + member.extend_from_slice(&crc.to_le_bytes()); + member.extend_from_slice(&encoded[10..]); + (member, crc_offset) + } + + #[tokio::test] + async fn gzip_optional_header_crc_is_validated_without_error_strings() { + const INPUT: &[u8] = b"name\nAlice\n"; + + let (encoded, crc_offset) = gzip_with_optional_header_crc(INPUT).await; + let (decoded, _) = decode(CompressionFormat::Gzip, encoded.clone()).await; + assert_eq!(decoded.expect("valid optional GZIP header should decode"), INPUT); + + let mut corrupt = encoded; + corrupt[crc_offset] ^= 0xff; + let (result, _) = decode(CompressionFormat::Gzip, corrupt).await; + assert_eq!( + select_error(&result.expect_err("invalid GZIP header CRC must fail")), + Some(SelectError::InvalidCompressionFormatForObject { + compression: CompressionFormat::Gzip.name(), + }) + ); + + let mut truncated = encode_compressed_fixture(CompressionFormat::Gzip, INPUT).await; + truncated.truncate(10); + truncated[3] = GZIP_FLAG_NAME; + truncated.extend_from_slice(b"unterminated-name"); + let (result, _) = decode(CompressionFormat::Gzip, truncated).await; + assert_eq!( + select_error(&result.expect_err("incomplete optional GZIP header must fail")), + Some(SelectError::TruncatedInput) + ); + } + + #[tokio::test] + async fn gzip_optional_header_can_cross_source_chunks() { + const INPUT: &[u8] = b"name\nAlice\n"; + const LARGE_GZIP_TEXT_FIELD_BYTES: usize = SELECT_DECODE_CHUNK_BYTES * 3 + 17; + + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, INPUT).await; + let mut with_name = encoded[..10].to_vec(); + with_name[3] = GZIP_FLAG_NAME; + with_name.extend(std::iter::repeat_n(b'x', LARGE_GZIP_TEXT_FIELD_BYTES)); + with_name.push(0); + with_name.extend_from_slice(&encoded[10..]); + let expected_scanned = u64::try_from(with_name.len()).expect("fixture length should fit in u64"); + let (mut writer, reader) = tokio::io::duplex(257); + let writer = tokio::spawn(async move { + writer + .write_all(&with_name) + .await + .expect("fixture source should accept bytes"); + writer.shutdown().await.expect("fixture source should close"); + }); + let metrics = Arc::new(SelectInputMetrics::default()); + let stream = compressed_input_stream( + Box::new(reader), + expected_scanned, + CompressionFormat::Gzip, + Arc::clone(&metrics), + b"\n".to_vec(), + u64::MAX, + None, + ) + .expect("record delimiter should be valid"); + let decoded = stream + .try_collect::>() + .await + .map(|chunks| chunks.concat()) + .expect("chunked optional GZIP names should decode"); + writer.await.expect("fixture writer should complete"); + + assert_eq!(decoded, INPUT); + assert_eq!(metrics.snapshot().bytes_scanned, expected_scanned); + } + + #[tokio::test] + async fn valid_512_byte_text_fields_decode_in_every_gzip_member() { + for flag in [GZIP_FLAG_NAME, GZIP_FLAG_COMMENT] { + let mut encoded = encode_compressed_fixture(CompressionFormat::Gzip, b"a\n").await; + encoded.extend_from_slice(&gzip_with_text_header(b"b\n", flag, 512).await); + let (decoded, _) = decode(CompressionFormat::Gzip, encoded).await; + assert_eq!(decoded.expect("RFC 1952 does not limit zero-terminated text fields"), b"a\nb\n"); + } + } + + #[tokio::test] + async fn long_text_field_header_crc_accumulates_across_source_chunks() { + const FIELD_BYTES: usize = SELECT_DECODE_CHUNK_BYTES * 3 + 17; + + for flag in [GZIP_FLAG_NAME, GZIP_FLAG_COMMENT] { + let first = encode_compressed_fixture(CompressionFormat::Gzip, b"a\n").await; + let (second, crc_offset) = gzip_with_text_header_crc(b"b\n", flag, FIELD_BYTES).await; + + let mut valid = first.clone(); + valid.extend_from_slice(&second); + let (decoded, _) = decode(CompressionFormat::Gzip, valid).await; + assert_eq!(decoded.expect("multi-chunk GZIP header CRC should validate"), b"a\nb\n"); + + let mut corrupt_second = second; + corrupt_second[crc_offset] ^= 0xff; + let mut corrupt = first; + corrupt.extend_from_slice(&corrupt_second); + let (result, _) = decode(CompressionFormat::Gzip, corrupt).await; + assert_eq!( + select_error(&result.expect_err("corrupt multi-chunk GZIP header CRC must fail")), + Some(SelectError::TruncatedInput) + ); + } + } + + #[tokio::test] + async fn maximum_extra_field_header_crc_accumulates_across_source_chunks() { + const INPUT: &[u8] = b"name\nAlice\n"; + let (valid, crc_offset) = gzip_with_extra_header_crc(INPUT, usize::from(u16::MAX)).await; + + let (decoded, _) = decode(CompressionFormat::Gzip, valid.clone()).await; + assert_eq!(decoded.expect("maximum GZIP extra field should decode"), INPUT); + + let mut corrupt = valid; + corrupt[crc_offset] ^= 0xff; + let (result, _) = decode(CompressionFormat::Gzip, corrupt).await; + assert_eq!( + select_error(&result.expect_err("corrupt multi-chunk GZIP extra-field CRC must fail")), + Some(CompressionFormat::Gzip.invalid_header_error()) + ); + } + + #[test] + fn processed_byte_limit_is_independent_of_live_memory_budget() { + assert_eq!(processed_bytes_limit(), MAX_SELECT_PROCESSED_BYTES); + assert!(processed_bytes_limit() > 80 * 1024 * 1024); + } + + #[test] + fn cooperative_reader_yields_after_bounded_ready_input() { + let source = Cursor::new(vec![0_u8; SELECT_DECODE_CHUNK_BYTES + 1]); + let mut reader = CooperativeReader::new(source); + let mut output = vec![0_u8; SELECT_DECODE_CHUNK_BYTES + 1]; + let mut read_buf = ReadBuf::new(&mut output); + let waker = futures::task::noop_waker_ref(); + let mut context = Context::from_waker(waker); + + assert!(Pin::new(&mut reader).poll_read(&mut context, &mut read_buf).is_ready()); + assert_eq!(read_buf.filled().len(), SELECT_DECODE_CHUNK_BYTES); + assert!(Pin::new(&mut reader).poll_read(&mut context, &mut read_buf).is_pending()); + assert!(Pin::new(&mut reader).poll_read(&mut context, &mut read_buf).is_ready()); + assert_eq!(read_buf.filled().len(), SELECT_DECODE_CHUNK_BYTES + 1); + } + + #[tokio::test] + async fn bzip2_decoded_output_yields_at_the_cooperative_quantum() { + let input = vec![b'x'; SELECT_DECODE_CHUNK_BYTES * 2]; + let compressed = encode_compressed_fixture(CompressionFormat::Bzip2, &input).await; + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let mut reader = compressed_input_reader( + Box::new(Cursor::new(compressed)), + compressed_len, + CompressionFormat::Bzip2, + Arc::new(SelectInputMetrics::default()), + processed_bytes_limit(), + None, + ); + let mut output = vec![0; SELECT_DECODE_CHUNK_BYTES * 2]; + let mut read_buf = ReadBuf::new(&mut output); + while read_buf.filled().len() < SELECT_DECODE_CHUNK_BYTES { + futures::future::poll_fn(|cx| Pin::new(&mut reader).poll_read(cx, &mut read_buf)) + .await + .expect("valid BZIP2 input should decode"); + } + assert_eq!(read_buf.filled().len(), SELECT_DECODE_CHUNK_BYTES); + + let waker = futures::task::noop_waker_ref(); + let mut context = Context::from_waker(waker); + assert!(Pin::new(&mut reader).poll_read(&mut context, &mut read_buf).is_pending()); + } + + #[tokio::test(flavor = "current_thread")] + async fn bzip2_decoder_polls_source_off_runtime_thread() { + let input = b"name\nAlice\n"; + let compressed = encode_compressed_fixture(CompressionFormat::Bzip2, input).await; + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let runtime_thread = std::thread::current().id(); + let source_thread = Arc::new(StdMutex::new(None)); + let source = ThreadRecordingReader { + inner: Cursor::new(compressed), + thread_id: Arc::clone(&source_thread), + }; + let mut reader = compressed_input_reader( + Box::new(source), + compressed_len, + CompressionFormat::Bzip2, + Arc::new(SelectInputMetrics::default()), + processed_bytes_limit(), + None, + ); + let mut decoded = Vec::new(); + + reader + .read_to_end(&mut decoded) + .await + .expect("valid BZIP2 input should decode off the runtime thread"); + + assert_eq!(decoded, input); + assert_ne!( + source_thread + .lock() + .expect("thread recorder mutex should not be poisoned") + .expect("compressed source should be polled"), + runtime_thread, + "BZIP2 decoding must not poll codec work on a Tokio runtime worker" + ); + } + + #[tokio::test] + async fn decoder_thread_holds_query_admission_until_exit() { + let admission = Arc::new(tokio::sync::Semaphore::new(1)); + let permit = Arc::new( + Arc::clone(&admission) + .try_acquire_owned() + .expect("query admission should be available"), + ); + let (started_tx, started_rx) = oneshot::channel(); + let (release_tx, release_rx) = std::sync::mpsc::channel(); + let (decoded_tx, _decoded_rx) = mpsc::channel(1); + + spawn_decoder_thread("s3select-guard-test", decoded_tx, Some(permit), move || { + let _ = started_tx.send(()); + release_rx.recv().map_err(io::Error::other) + }); + + started_rx.await.expect("decoder thread should start"); + assert!(Arc::clone(&admission).try_acquire_owned().is_err()); + release_tx.send(()).expect("decoder thread should be releasable"); + let recovered = tokio::time::timeout(std::time::Duration::from_secs(1), Arc::clone(&admission).acquire_owned()) + .await + .expect("decoder thread should release admission promptly") + .expect("query admission should remain open"); + drop(recovered); + } + + #[test] + fn bzip2_decoder_does_not_wait_for_tokio_blocking_pool() { + let runtime = tokio::runtime::Builder::new_multi_thread() + .worker_threads(1) + .max_blocking_threads(1) + .enable_all() + .build() + .expect("test runtime should build"); + + runtime.block_on(async { + let (blocker_started_tx, blocker_started_rx) = tokio::sync::oneshot::channel(); + let (release_tx, release_rx) = std::sync::mpsc::channel(); + let blocker = tokio::task::spawn_blocking(move || { + let _ = blocker_started_tx.send(()); + let _ = release_rx.recv(); + }); + blocker_started_rx.await.expect("blocking worker should be occupied"); + + let input = b"name\nAlice\n"; + let compressed = encode_compressed_fixture(CompressionFormat::Bzip2, input).await; + let decode_result = + tokio::time::timeout(std::time::Duration::from_secs(2), decode(CompressionFormat::Bzip2, compressed)).await; + + let _ = release_tx.send(()); + blocker.await.expect("blocking worker should finish"); + + let (decoded, _) = decode_result.expect("BZIP2 decoder must not queue behind Tokio blocking work"); + assert_eq!(decoded.expect("valid BZIP2 input should decode"), input); + }); + } + + #[tokio::test] + async fn processed_reader_streams_past_the_default_memory_budget() { + let expected = 64 * 1024 * 1024 + 1; + let metrics = Arc::new(SelectInputMetrics::default()); + let source = tokio::io::repeat(b'x').take(expected); + let mut reader = ProcessedReader::new(source, metrics.recorder(), processed_bytes_limit()); + let mut buffer = vec![0; SELECT_DECODE_CHUNK_BYTES]; + let mut total = 0_u64; + + loop { + let read = reader + .read(&mut buffer) + .await + .expect("streaming input within the expansion budget should pass"); + if read == 0 { + break; + } + total += u64::try_from(read).expect("read buffer length fits in u64"); + } + + assert_eq!(total, expected); + assert_eq!(metrics.snapshot().bytes_processed, expected); + } + + #[tokio::test] + async fn truncated_and_corrupt_gzip_are_typed_as_truncated_input() { + let encoded = encode_compressed_fixture(CompressionFormat::Gzip, b"name\nAlice\n").await; + + let mut truncated = encoded.clone(); + truncated.truncate(truncated.len() - 1); + let (result, _) = decode(CompressionFormat::Gzip, truncated).await; + assert_eq!( + select_error(&result.expect_err("truncated gzip trailer must fail")), + Some(SelectError::TruncatedInput) + ); + + let mut corrupt = encoded; + let checksum_index = corrupt.len() - 8; + corrupt[checksum_index] ^= 0xff; + let (result, _) = decode(CompressionFormat::Gzip, corrupt).await; + assert_eq!( + select_error(&result.expect_err("corrupt gzip checksum must fail")), + Some(SelectError::TruncatedInput) + ); + + let mut corrupt_size = encode_compressed_fixture(CompressionFormat::Gzip, b"name\nAlice\n").await; + let size_index = corrupt_size.len() - 1; + corrupt_size[size_index] ^= 0xff; + let (result, _) = decode(CompressionFormat::Gzip, corrupt_size).await; + assert_eq!( + select_error(&result.expect_err("corrupt gzip uncompressed size must fail")), + Some(SelectError::TruncatedInput) + ); + } + + #[tokio::test] + async fn truncated_bzip2_is_typed_as_truncated_input() { + let mut encoded = encode_compressed_fixture(CompressionFormat::Bzip2, b"name\nAlice\n").await; + encoded.truncate(encoded.len() - 1); + let (result, _) = decode(CompressionFormat::Bzip2, encoded).await; + assert_eq!( + select_error(&result.expect_err("truncated bzip2 trailer must fail")), + Some(SelectError::TruncatedInput) + ); + + let mut corrupt = encode_compressed_fixture(CompressionFormat::Bzip2, b"name\nAlice\n").await; + let checksum_index = corrupt.len() - 2; + corrupt[checksum_index] ^= 0xff; + let (result, _) = decode(CompressionFormat::Bzip2, corrupt).await; + assert_eq!( + select_error(&result.expect_err("corrupt bzip2 checksum must fail")), + Some(SelectError::TruncatedInput) + ); + } + + #[tokio::test] + async fn trailing_non_member_bytes_fail_closed() { + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let mut encoded = encode_compressed_fixture(format, b"name\nAlice\n").await; + encoded.extend_from_slice(b"trailing garbage"); + let (result, _) = decode(format, encoded).await; + assert_eq!( + select_error(&result.expect_err("trailing bytes must not be ignored")), + Some(SelectError::TruncatedInput) + ); + } + } + + #[tokio::test] + async fn oversized_compressed_record_fails_before_unbounded_buffering() { + let mut input = vec![b'x'; MAX_SELECT_RECORD_BYTES + 1]; + input.push(b'\n'); + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, &input).await; + let (result, _) = decode(CompressionFormat::Gzip, compressed).await; + assert_eq!( + select_error(&result.expect_err("oversized compressed record must fail")), + Some(SelectError::OverMaxRecordSize) + ); + } + + #[tokio::test] + async fn one_megabyte_record_is_accepted_with_or_without_delimiter() { + for terminated in [false, true] { + let mut input = vec![b'x'; MAX_SELECT_RECORD_BYTES]; + if terminated { + input.push(b'\n'); + } + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, &input).await; + let (decoded, _) = decode(CompressionFormat::Gzip, compressed).await; + assert_eq!(decoded.expect("record at the protocol limit should decode"), input); + } + } + + #[tokio::test] + async fn processed_byte_limit_rejects_many_small_records() { + let input = b"{}\n".repeat(1024); + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, &input).await; + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let exact_stream = compressed_input_stream( + Box::new(Cursor::new(compressed.clone())), + compressed_len, + CompressionFormat::Gzip, + Arc::new(SelectInputMetrics::default()), + b"\n".to_vec(), + u64::try_from(input.len()).expect("fixture length should fit in u64"), + None, + ) + .expect("record delimiter should be valid"); + assert_eq!( + exact_stream + .try_collect::>() + .await + .expect("decoded bytes at the processed limit should pass") + .concat(), + input + ); + + let stream = compressed_input_stream( + Box::new(Cursor::new(compressed)), + compressed_len, + CompressionFormat::Gzip, + Arc::new(SelectInputMetrics::default()), + b"\n".to_vec(), + u64::try_from(input.len() - 1).expect("fixture length should fit in u64"), + None, + ) + .expect("record delimiter should be valid"); + + let error = stream + .try_collect::>() + .await + .expect_err("decoded bytes over the decompression budget must fail"); + assert_eq!(select_error(&error), Some(SelectError::ResourceExhausted)); + } + + #[tokio::test] + async fn oversized_final_record_with_partial_delimiter_fails_at_eof() { + let mut input = vec![b'x'; MAX_SELECT_RECORD_BYTES + 1]; + input.push(b'\r'); + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, &input).await; + let compressed_len = u64::try_from(compressed.len()).expect("fixture length should fit in u64"); + let metrics = Arc::new(SelectInputMetrics::default()); + let stream = compressed_input_stream( + Box::new(Cursor::new(compressed)), + compressed_len, + CompressionFormat::Gzip, + metrics, + b"\r\n".to_vec(), + u64::MAX, + None, + ) + .expect("record delimiter should be valid"); + + let error = stream + .try_collect::>() + .await + .expect_err("oversized unterminated record must fail"); + assert_eq!(select_error(&error), Some(SelectError::OverMaxRecordSize)); + } + + struct DropObservedReader { + inner: DuplexStream, + dropped: Arc, + } + + impl AsyncRead for DropObservedReader { + fn poll_read(mut self: Pin<&mut Self>, cx: &mut Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_read(cx, buf) + } + } + + impl Drop for DropObservedReader { + fn drop(&mut self) { + self.dropped.store(true, std::sync::atomic::Ordering::Release); + } + } + + #[tokio::test] + async fn dropping_decoder_stream_releases_source_without_background_work() { + let input = b"a\n".repeat(MAX_SELECT_RECORD_BYTES); + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let compressed = encode_compressed_fixture(format, &input).await; + let compressed_len = u64::try_from(compressed.len() + 1).expect("fixture length should fit in u64"); + let (source, mut peer) = tokio::io::duplex(compressed.len()); + peer.write_all(&compressed).await.expect("write compressed fixture"); + let dropped = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let admission = Arc::new(tokio::sync::Semaphore::new(1)); + let permit = Arc::new( + Arc::clone(&admission) + .try_acquire_owned() + .expect("query admission should be available"), + ); + let mut stream = compressed_input_stream( + Box::new(DropObservedReader { + inner: source, + dropped: Arc::clone(&dropped), + }), + compressed_len, + format, + Arc::new(SelectInputMetrics::default()), + b"\n".to_vec(), + u64::MAX, + Some(permit), + ) + .expect("record delimiter should be valid"); + let first = tokio::time::timeout(std::time::Duration::from_secs(1), stream.next()) + .await + .expect("decoder should produce a chunk before source EOF") + .expect("decoder stream should produce a chunk") + .expect("valid partial decode should succeed"); + assert!(!first.is_empty()); + + drop(stream); + + tokio::time::timeout(std::time::Duration::from_secs(1), async { + while !dropped.load(std::sync::atomic::Ordering::Acquire) { + tokio::task::yield_now().await; + } + }) + .await + .expect("dropping decoded output must cancel the blocked source read"); + let recovered = tokio::time::timeout(std::time::Duration::from_secs(1), Arc::clone(&admission).acquire_owned()) + .await + .expect("decoder exit should release query admission") + .expect("query admission should remain open"); + drop(recovered); + } + } +} diff --git a/crates/s3select-api/src/lib.rs b/crates/s3select-api/src/lib.rs index 7662b4a78..a443a07a1 100644 --- a/crates/s3select-api/src/lib.rs +++ b/crates/s3select-api/src/lib.rs @@ -23,6 +23,7 @@ use datafusion::{ use std::{error::Error as StdError, fmt::Display}; use thiserror::Error; +mod input_stream; mod metrics; pub mod object_store; pub mod query; @@ -79,6 +80,9 @@ pub enum SelectError { #[error("The file is not in a supported compression format. Only GZIP and BZIP2 are supported.")] InvalidCompressionFormat, + #[error("{compression} is not applicable to the queried object. Please correct the request and try again.")] + InvalidCompressionFormatForObject { compression: &'static str }, + #[error("The data source type is not valid. Only CSV, JSON, and Parquet are supported.")] InvalidDataSource, @@ -87,6 +91,9 @@ pub enum SelectError { )] TruncatedInput, + #[error("Scan range queries are not supported on this type of object.")] + UnsupportedScanRangeInput, + #[error("An error occurred while parsing the CSV file. Check the file and try again.")] CsvParsingError, @@ -96,6 +103,9 @@ pub enum SelectError { #[error("An error occurred while parsing the Parquet file. Check the file and try again.")] ParquetParsingError, + #[error("The length of a record in the input or result is greater than the maxCharsPerRecord limit of 1 MB.")] + OverMaxRecordSize, + #[error("{message}")] ParseSelectFailure { message: String }, diff --git a/crates/s3select-api/src/metrics.rs b/crates/s3select-api/src/metrics.rs index d38eefbc8..619f1e895 100644 --- a/crates/s3select-api/src/metrics.rs +++ b/crates/s3select-api/src/metrics.rs @@ -12,7 +12,11 @@ // See the License for the specific language governing permissions and // limitations under the License. -use std::sync::atomic::{AtomicU64, Ordering}; +use arc_swap::ArcSwap; +use std::sync::{ + Arc, + atomic::{AtomicU64, Ordering}, +}; #[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] pub struct SelectInputMetricsSnapshot { @@ -20,33 +24,72 @@ pub struct SelectInputMetricsSnapshot { pub bytes_processed: u64, } -#[derive(Debug, Default)] +#[derive(Debug)] pub struct SelectInputMetrics { + active: ArcSwap, +} + +#[derive(Debug, Default)] +struct SelectInputMetricBank { uncompressed_bytes: AtomicU64, + compressed_bytes_scanned: AtomicU64, + compressed_bytes_processed: AtomicU64, +} + +#[derive(Clone, Debug)] +pub(crate) struct SelectInputMetricsRecorder { + bank: Arc, +} + +impl Default for SelectInputMetrics { + fn default() -> Self { + Self { + active: ArcSwap::from_pointee(SelectInputMetricBank::default()), + } + } } impl SelectInputMetrics { pub fn snapshot(&self) -> SelectInputMetricsSnapshot { - let uncompressed_bytes = self.uncompressed_bytes.load(Ordering::Relaxed); + let bank = self.active.load(); + let uncompressed_bytes = bank.uncompressed_bytes.load(Ordering::Relaxed); SelectInputMetricsSnapshot { - bytes_scanned: uncompressed_bytes, - bytes_processed: uncompressed_bytes, + bytes_scanned: uncompressed_bytes.saturating_add(bank.compressed_bytes_scanned.load(Ordering::Relaxed)), + bytes_processed: uncompressed_bytes.saturating_add(bank.compressed_bytes_processed.load(Ordering::Relaxed)), } } - pub(crate) fn record_uncompressed(&self, bytes: usize) { - let increment = u64::try_from(bytes).unwrap_or(u64::MAX); - let _ = self - .uncompressed_bytes - .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| Some(current.saturating_add(increment))); + pub(crate) fn recorder(&self) -> SelectInputMetricsRecorder { + SelectInputMetricsRecorder { + bank: self.active.load_full(), + } } - /// Clears planner-only reads before query execution begins. + /// Publishes a fresh bank so late planner writes remain isolated. pub fn reset(&self) { - self.uncompressed_bytes.store(0, Ordering::Relaxed); + self.active.store(Arc::new(SelectInputMetricBank::default())); } } +impl SelectInputMetricsRecorder { + pub(crate) fn record_uncompressed(&self, bytes: usize) { + saturating_add(&self.bank.uncompressed_bytes, bytes); + } + + pub(crate) fn record_scanned(&self, bytes: usize) { + saturating_add(&self.bank.compressed_bytes_scanned, bytes); + } + + pub(crate) fn record_processed(&self, bytes: usize) { + saturating_add(&self.bank.compressed_bytes_processed, bytes); + } +} + +fn saturating_add(counter: &AtomicU64, bytes: usize) { + let increment = u64::try_from(bytes).unwrap_or(u64::MAX); + let _ = counter.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| Some(current.saturating_add(increment))); +} + #[cfg(test)] mod tests { use super::*; @@ -54,7 +97,7 @@ mod tests { #[test] fn records_uncompressed_input_at_both_boundaries() { let metrics = SelectInputMetrics::default(); - metrics.record_uncompressed(7); + metrics.recorder().record_uncompressed(7); assert_eq!( metrics.snapshot(), @@ -68,21 +111,51 @@ mod tests { #[test] fn counters_saturate_instead_of_wrapping() { let metrics = SelectInputMetrics::default(); - metrics.uncompressed_bytes.store(u64::MAX - 1, Ordering::Relaxed); + metrics + .active + .load() + .uncompressed_bytes + .store(u64::MAX - 1, Ordering::Relaxed); - metrics.record_uncompressed(2); + metrics.recorder().record_uncompressed(2); assert_eq!(metrics.snapshot().bytes_scanned, u64::MAX); assert_eq!(metrics.snapshot().bytes_processed, u64::MAX); } + #[test] + fn compressed_boundaries_are_counted_independently() { + let metrics = SelectInputMetrics::default(); + let recorder = metrics.recorder(); + recorder.record_scanned(39); + recorder.record_processed(19); + + assert_eq!( + metrics.snapshot(), + SelectInputMetricsSnapshot { + bytes_scanned: 39, + bytes_processed: 19, + } + ); + } + #[test] fn reset_clears_schema_inference_bytes() { let metrics = SelectInputMetrics::default(); - metrics.record_uncompressed(9); + let planning = metrics.recorder(); + planning.record_uncompressed(9); metrics.reset(); + planning.record_uncompressed(5); + let execution = metrics.recorder(); + execution.record_uncompressed(3); - assert_eq!(metrics.snapshot(), SelectInputMetricsSnapshot::default()); + assert_eq!( + metrics.snapshot(), + SelectInputMetricsSnapshot { + bytes_scanned: 3, + bytes_processed: 3, + } + ); } } diff --git a/crates/s3select-api/src/object_store.rs b/crates/s3select-api/src/object_store.rs index 20d165623..c135000aa 100644 --- a/crates/s3select-api/src/object_store.rs +++ b/crates/s3select-api/src/object_store.rs @@ -13,9 +13,14 @@ // limitations under the License. use crate::{ - PrepareSelectObjectSnapshotError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectError, SelectGetObjectReader, SelectInputMetrics, - SelectObjectOptions, SelectObjectSnapshot, SelectObjectSnapshotReadError, SelectStorageError, SelectStore, - SnapshotConsistencyError, + PrepareSelectObjectSnapshotError, QueryError, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectError, SelectGetObjectReader, + SelectInputMetrics, SelectObjectOptions, SelectObjectSnapshot, SelectObjectSnapshotReadError, SelectStorageError, + SelectStore, SnapshotConsistencyError, + input_stream::{ + CompressionFormat, MAX_SELECT_RECORD_BYTES, SELECT_DECODE_CHUNK_BYTES, SelectInputReader, compressed_input_reader, + compressed_input_stream, input_io_error, processed_bytes_limit, + }, + metrics::SelectInputMetricsRecorder, query::{ ast::{JsonPathSegment, JsonSource}, parser::RustFsDialect, @@ -52,12 +57,15 @@ use s3s::header::{ X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, }; -use s3s::{S3Error, S3ErrorCode, S3Result, dto::SelectObjectContentInput}; +use s3s::{ + S3Error, S3ErrorCode, S3Result, + dto::{CompressionType, InputSerialization, ScanRange, SelectObjectContentInput}, +}; use std::collections::VecDeque; use std::ops::Range; -use std::sync::Arc; #[cfg(test)] use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, atomic::AtomicBool}; use tokio::{io::AsyncReadExt, sync::OnceCell}; use tokio_util::io::ReaderStream; use transform_stream::AsyncTryStream; @@ -68,6 +76,15 @@ fn select_default_read_buffer_size_u64() -> u64 { u64::try_from(SELECT_DEFAULT_READ_BUFFER_SIZE).unwrap_or(u64::MAX) } +fn compression_format(input: &InputSerialization) -> Result, SelectError> { + match input.compression_type.as_ref().map(|value| value.as_str()) { + None | Some(CompressionType::NONE) => Ok(None), + Some(CompressionType::GZIP) => Ok(Some(CompressionFormat::Gzip)), + Some(CompressionType::BZIP2) => Ok(Some(CompressionFormat::Bzip2)), + Some(_) => Err(SelectError::InvalidCompressionFormat), + } +} + /// Maximum allowed object size for JSON DOCUMENT mode. /// /// JSON DOCUMENT format requires loading the entire file into memory for DOM @@ -89,11 +106,18 @@ fn select_default_read_buffer_size_u64() -> u64 { pub const MAX_JSON_DOCUMENT_BYTES: u64 = 128 * 1024 * 1024; const JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER: usize = 64; const JSON_SCALAR_COLUMN_MEMORY_RESERVATION_MULTIPLIER: usize = 14; +const JSON_CANCELLATION_CHECK_BYTES: usize = 64 * 1024; +const JSON_CANCELLATION_CHECK_KEYS: usize = 1024; 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."; const NORMALIZED_RECORD_DELIMITER: &[u8] = b"\r\n"; const NORMALIZED_FIELD_DELIMITER: &[u8] = &[DEFAULT_DELIMITER]; +/// Returns true for the MinIO-compatible full-scan range marker. +pub fn is_noop_scan_range(scan_range: &ScanRange) -> bool { + scan_range.start == Some(0) && scan_range.end.is_none() +} + pub struct EcObjectStore { input: Arc, need_convert: bool, @@ -312,6 +336,9 @@ impl EcObjectStore { let Some(scan_range) = self.input.request.scan_range.as_ref() else { return Ok(None); }; + if is_noop_scan_range(scan_range) { + return Ok(None); + } scan_range_from_bounds(scan_range.start, scan_range.end, object_size) } @@ -766,6 +793,22 @@ impl ObjectStore for EcObjectStore { // this instance's immutable snapshot; later operations reuse it. let snapshot = self.snapshot(options.version.as_deref()).await?; let original_size = snapshot.logical_size(); + let compression = compression_format(&self.input.request.input_serialization).map_err(|source| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(source), + })?; + let has_effective_request_range = self + .input + .request + .scan_range + .as_ref() + .is_some_and(|scan_range| !is_noop_scan_range(scan_range)); + if compression.is_some() && (options.range.is_some() || has_effective_request_range) { + return Err(o_Error::Generic { + store: "EcObjectStore", + source: Box::new(SelectError::UnsupportedScanRangeInput), + }); + } let object_info = snapshot.object_info(); let meta = ObjectMeta { location: location.clone(), @@ -791,7 +834,7 @@ impl ObjectStore for EcObjectStore { } let record_delimiter = self.record_delimiter_for_conversion(); - let needs_scan_context = options.range.is_none() && self.input.request.scan_range.is_some(); + let needs_scan_context = options.range.is_none() && has_effective_request_range; let scan_context = if needs_scan_context { if let Some(scan_range) = self.scan_range(original_size)? { let delimiter = self.record_delimiter(); @@ -813,7 +856,51 @@ impl ObjectStore for EcObjectStore { }; let meter_input = self.input.request.input_serialization.parquet.is_none(); - let payload = if options.range.is_some() { + let payload = if let Some(compression) = compression { + let max_processed_bytes = processed_bytes_limit(); + let query_guard = match self.query_tracker.as_ref() { + Some(query_tracker) => Some(query_tracker.query_guard().ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(QueryError::Cancel), + })?), + None => None, + }; + if self.is_json_document { + let reader = compressed_input_reader( + reader.stream, + original_size, + compression, + Arc::clone(&self.input_metrics), + max_processed_bytes, + query_guard, + ); + let stream = compressed_json_document_ndjson_stream( + reader, + self.json_source.clone(), + Arc::clone(&self.memory_pool), + self.query_tracker.clone(), + ); + GetResultPayload::Stream(stream) + } else { + let input_record_delimiter = if self.input.request.input_serialization.csv.is_some() { + self.record_delimiter() + } else { + b"\n".to_vec() + }; + let stream = compressed_input_stream( + reader.stream, + original_size, + compression, + Arc::clone(&self.input_metrics), + input_record_delimiter, + max_processed_bytes, + query_guard, + )?; + let stream = + convert_csv_delimiter_stream(stream, record_delimiter, self.need_convert.then(|| self.delimiter.clone())); + GetResultPayload::Stream(stream) + } + } else if options.range.is_some() { let size = usize::try_from(result_range.end - result_range.start).map_err(|err| o_Error::Generic { store: "EcObjectStore", source: Box::new(err), @@ -851,8 +938,9 @@ impl ObjectStore for EcObjectStore { let delimiter = self.record_delimiter(); let include_header = self.csv_has_header(); let header = if include_header && read_start > 0 { + let input_metrics = self.input_metrics.recorder(); let header = self.read_header_record(original_size, &delimiter).await?; - self.input_metrics.record_uncompressed(header.len()); + input_metrics.record_uncompressed(header.len()); Some(header) } else { None @@ -1197,7 +1285,7 @@ impl ScanRangeState { /// scalar/object root) is yielded as a separate [`Bytes`] chunk, so /// DataFusion can pipeline row processing as lines arrive. fn json_document_ndjson_stream( - stream: Box, + stream: SelectInputReader, original_size: u64, json_source: JsonSource, input_metrics: Arc, @@ -1206,60 +1294,107 @@ fn json_document_ndjson_stream( ) -> futures_core::stream::BoxStream<'static, Result> { json_document_ndjson_stream_with_parser( stream, - original_size, + JsonDocumentReadMode::Exact { + original_size, + input_metrics: input_metrics.recorder(), + }, json_source, - input_metrics, memory_pool, query_tracker, - |all_bytes, json_source| parse_json_document_to_lines(&all_bytes, &json_source), + |all_bytes, json_source, cancellation| { + parse_json_document_to_lines_cancellable(&all_bytes, &json_source, cancellation.as_ref()) + }, ) } -fn json_document_ndjson_stream_with_parser

( - stream: Box, - original_size: u64, +fn compressed_json_document_ndjson_stream( + stream: SelectInputReader, + json_source: JsonSource, + memory_pool: Arc, + query_tracker: Option, +) -> futures_core::stream::BoxStream<'static, Result> { + json_document_ndjson_stream_with_parser( + stream, + JsonDocumentReadMode::Bounded, + json_source, + memory_pool, + query_tracker, + |all_bytes, json_source, cancellation| { + parse_json_document_to_lines_cancellable(&all_bytes, &json_source, cancellation.as_ref()) + }, + ) +} + +enum JsonDocumentReadMode { + Exact { + original_size: u64, + input_metrics: SelectInputMetricsRecorder, + }, + Bounded, +} + +fn json_document_ndjson_stream_with_parser

( + stream: SelectInputReader, + read_mode: JsonDocumentReadMode, json_source: JsonSource, - input_metrics: Arc, memory_pool: Arc, query_tracker: Option, parser: P, ) -> futures_core::stream::BoxStream<'static, Result> where - P: FnOnce(Vec, JsonSource) -> std::io::Result> + Send + 'static, + P: FnOnce(Vec, JsonSource, Arc) -> 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 = - json_document_memory_reservation_bytes(buffer_capacity, &json_source).map_err(|source| o_Error::Generic { - store: "EcObjectStore", - source: Box::new(source), - })?; 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), - })?; // ── 1. Read phase (lazy: only runs when the stream is polled) ──── pin_mut!(stream); - let mut all_bytes = Vec::with_capacity(buffer_capacity); - let read_result = stream.take(original_size).read_to_end(&mut all_bytes).await; - input_metrics.record_uncompressed(all_bytes.len()); - read_result.map_err(|e| o_Error::Generic { - store: "EcObjectStore", - source: Box::new(e), - })?; - if all_bytes.len() != buffer_capacity { - return Err(incomplete_object_stream_error(buffer_capacity - all_bytes.len())); - } + let all_bytes = match read_mode { + JsonDocumentReadMode::Exact { + original_size, + input_metrics, + } => { + 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" + ))), + })?; + resize_json_document_reservation(&reservation, buffer_capacity, &json_source)?; + let mut all_bytes = Vec::with_capacity(buffer_capacity); + let read_result = stream.take(original_size).read_to_end(&mut all_bytes).await; + input_metrics.record_uncompressed(all_bytes.len()); + read_result.map_err(input_io_error)?; + if all_bytes.len() != buffer_capacity { + return Err(incomplete_object_stream_error(buffer_capacity - all_bytes.len())); + } + all_bytes + } + JsonDocumentReadMode::Bounded => { + let mut all_bytes = Vec::new(); + let mut buffer = vec![0; SELECT_DECODE_CHUNK_BYTES]; + loop { + let read = stream.read(&mut buffer).await.map_err(input_io_error)?; + if read == 0 { + break; + } + let new_len = all_bytes.len().checked_add(read).ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(json_document_memory_reservation_overflow(all_bytes.len())), + })?; + let new_len_u64 = u64::try_from(new_len).map_err(|_| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT input size {new_len} does not fit in the object size type" + ))), + })?; + validate_json_document_size(new_len_u64)?; + grow_json_document_buffer(&mut all_bytes, new_len, &reservation, &json_source)?; + all_bytes.extend_from_slice(&buffer[..read]); + } + all_bytes + } + }; // ── 2. Parse phase (blocking thread pool, non-blocking runtime) ── let queued_query_guard = match query_tracker.as_ref() { @@ -1275,11 +1410,24 @@ where }; let pending_query_guard = PendingQueryExecutionGuard::new(query_tracker); let task_query_guard = pending_query_guard.task_state(); + let cancellation = Arc::new(AtomicBool::new(false)); + let queued_task = Arc::new(Mutex::new(Some(JsonDocumentParseTask { + parser, + all_bytes, + json_source, + task_resources, + }))); + let _cancel_on_drop = JsonDocumentCancellation::new(Arc::clone(&cancellation), Arc::clone(&queued_task)); let (lines, _task_resources) = SpawnedTask::spawn_blocking(move || { + let JsonDocumentParseTask { + parser, + all_bytes, + json_source, + mut task_resources, + } = queued_task.lock().take().ok_or_else(json_document_parse_interrupted_error)?; let query_guard = PendingQueryExecutionGuard::start(&task_query_guard)?; - let mut task_resources = task_resources; task_resources.query_guard = query_guard; - parser(all_bytes, json_source).map(|lines| (lines, task_resources)) + parser(all_bytes, json_source, cancellation).map(|lines| (lines, task_resources)) }) .await .map_err(|e| o_Error::Generic { @@ -1300,12 +1448,116 @@ where .boxed() } +fn grow_json_document_buffer( + buffer: &mut Vec, + required_len: usize, + reservation: &MemoryReservation, + json_source: &JsonSource, +) -> Result<()> { + if required_len <= buffer.capacity() { + return Ok(()); + } + + let max_capacity = usize::try_from(MAX_JSON_DOCUMENT_BYTES).map_err(|_| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted( + "JSON DOCUMENT size limit does not fit in memory".to_string(), + )), + })?; + let target_capacity = required_len + .checked_next_power_of_two() + .ok_or_else(|| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT buffer capacity overflow at {required_len} bytes" + ))), + })? + .min(max_capacity); + if target_capacity < required_len { + return Err(o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT input size {required_len} exceeds the maximum buffer capacity" + ))), + }); + } + + resize_json_document_reservation(reservation, target_capacity, json_source)?; + buffer + .try_reserve_exact(target_capacity - buffer.len()) + .map_err(|_| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(DataFusionError::ResourcesExhausted(format!( + "JSON DOCUMENT input buffer allocation failed at {target_capacity} bytes" + ))), + })?; + resize_json_document_reservation(reservation, buffer.capacity(), json_source) +} + +fn resize_json_document_reservation(reservation: &MemoryReservation, input_bytes: usize, json_source: &JsonSource) -> Result<()> { + let reservation_bytes = + json_document_memory_reservation_bytes(input_bytes, json_source).map_err(|source| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(source), + })?; + reservation.try_resize(reservation_bytes).map_err(|source| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(source), + }) +} + struct JsonDocumentTaskResources { // Struct fields drop in declaration order, so admission covers the reservation through teardown. _reservation: MemoryReservation, query_guard: Option, } +struct JsonDocumentParseTask

{ + parser: P, + all_bytes: Vec, + json_source: JsonSource, + task_resources: JsonDocumentTaskResources, +} + +struct JsonDocumentCancellation { + cancelled: Arc, + queued: Arc>>, +} + +impl JsonDocumentCancellation { + fn new(cancelled: Arc, queued: Arc>>) -> Self { + Self { cancelled, queued } + } +} + +impl Drop for JsonDocumentCancellation { + fn drop(&mut self) { + self.cancelled.store(true, std::sync::atomic::Ordering::Release); + let queued = self.queued.lock().take(); + drop(queued); + } +} + +struct CancellableJsonReader<'a> { + inner: std::io::Cursor<&'a [u8]>, + cancelled: &'a AtomicBool, +} + +impl std::io::Read for CancellableJsonReader<'_> { + fn read(&mut self, buffer: &mut [u8]) -> std::io::Result { + ensure_json_parse_active(self.cancelled)?; + std::io::Read::read(&mut self.inner, buffer) + } +} + +fn ensure_json_parse_active(cancelled: &AtomicBool) -> std::io::Result<()> { + if cancelled.load(std::sync::atomic::Ordering::Acquire) { + Err(std::io::Error::new(std::io::ErrorKind::Interrupted, SelectError::Canceled)) + } else { + Ok(()) + } +} + fn json_document_memory_reservation_bytes(input_bytes: usize, json_source: &JsonSource) -> datafusion::common::Result { let base = input_bytes .checked_mul(JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER) @@ -1409,11 +1661,33 @@ impl Drop for PendingQueryExecutionGuard { /// /// - A JSON array → one line per element. /// - A JSON object or scalar root → one line. +#[cfg(test)] fn parse_json_document_to_lines(bytes: &[u8], json_source: &JsonSource) -> std::io::Result> { - let root: serde_json::Value = - serde_json::from_slice(bytes).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?; + parse_json_document_to_lines_cancellable(bytes, json_source, &AtomicBool::new(false)) +} + +fn parse_json_document_to_lines_cancellable( + bytes: &[u8], + json_source: &JsonSource, + cancelled: &AtomicBool, +) -> std::io::Result> { + let reader = std::io::BufReader::with_capacity( + JSON_CANCELLATION_CHECK_BYTES, + CancellableJsonReader { + inner: std::io::Cursor::new(bytes), + cancelled, + }, + ); + let root: serde_json::Value = match serde_json::from_reader(reader) { + Ok(root) => root, + Err(error) => { + ensure_json_parse_active(cancelled)?; + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, error)); + } + }; + ensure_json_parse_active(cancelled)?; let json_source_path = json_source.path(); - let values = expand_json_source(root, json_source_path)?; + let values = expand_json_source(root, json_source_path, cancelled)?; // Preserve the two pre-path-AST forms that flattened arrays implicitly. // Explicit indexes and wildcards already identify the intended records // and must not flatten an array-valued result a second time. @@ -1424,19 +1698,25 @@ fn parse_json_document_to_lines(bytes: &[u8], json_source: &JsonSource) -> std:: }); let mut lines: Vec = Vec::new(); for value in values { + ensure_json_parse_active(cancelled)?; match value { serde_json::Value::Array(array) if implicitly_expand_arrays => { for item in array { - lines.push(json_value_to_line(item, scalar_column)?); + ensure_json_parse_active(cancelled)?; + lines.push(json_value_to_line_cancellable(item, scalar_column, cancelled)?); } } - other => lines.push(json_value_to_line(other, scalar_column)?), + other => lines.push(json_value_to_line_cancellable(other, scalar_column, cancelled)?), } } Ok(lines) } -fn expand_json_source(root: serde_json::Value, json_source_path: &[JsonPathSegment]) -> std::io::Result> { +fn expand_json_source( + root: serde_json::Value, + json_source_path: &[JsonPathSegment], + cancelled: &AtomicBool, +) -> std::io::Result> { // S3Object[*] identifies the input record stream. JSON DOCUMENT already // presents the root value as that stream, so the leading marker is not a // lookup against the root object. @@ -1454,11 +1734,13 @@ fn expand_json_source(root: serde_json::Value, json_source_path: &[JsonPathSegme }; for segment in path { + ensure_json_parse_active(cancelled)?; let mut expanded = Vec::new(); for value in values { + ensure_json_parse_active(cancelled)?; match (segment, value) { (JsonPathSegment::Key { name, quoted }, serde_json::Value::Object(mut object)) => { - if let Some(value) = remove_json_source_key(&mut object, name, *quoted)? { + if let Some(value) = remove_json_source_key(&mut object, name, *quoted, cancelled)? { expanded.push(value); } } @@ -1495,21 +1777,69 @@ fn remove_json_source_key( object: &mut serde_json::Map, name: &str, quoted: bool, + cancelled: &AtomicBool, +) -> std::io::Result> { + let mut checkpoint = || ensure_json_parse_active(cancelled); + remove_json_source_key_with_checkpoint(object, name, quoted, &mut checkpoint) +} + +fn remove_json_source_key_with_checkpoint( + object: &mut serde_json::Map, + name: &str, + quoted: bool, + checkpoint: &mut impl FnMut() -> std::io::Result<()>, ) -> std::io::Result> { if quoted { return Ok(object.remove(name)); } - let mut matches = object.keys().filter(|key| key.eq_ignore_ascii_case(name)); - let matched = matches.next().cloned(); - if matches.next().is_some() { - return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::AmbiguousFieldName)); + let mut matched = None; + for (index, key) in object.keys().enumerate() { + if index % JSON_CANCELLATION_CHECK_KEYS == 0 { + checkpoint()?; + } + if json_key_eq_ignore_ascii_case_with_checkpoint(key, name, checkpoint)? { + if matched.is_some() { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::AmbiguousFieldName)); + } + matched = Some(key.clone()); + } } - drop(matches); Ok(matched.and_then(|key| object.remove(&key))) } +fn json_key_eq_ignore_ascii_case_with_checkpoint( + key: &str, + expected: &str, + checkpoint: &mut impl FnMut() -> std::io::Result<()>, +) -> std::io::Result { + if key.len() != expected.len() { + return Ok(false); + } + + for (key_chunk, expected_chunk) in key + .as_bytes() + .chunks(JSON_CANCELLATION_CHECK_BYTES) + .zip(expected.as_bytes().chunks(JSON_CANCELLATION_CHECK_BYTES)) + { + checkpoint()?; + if !key_chunk.eq_ignore_ascii_case(expected_chunk) { + return Ok(false); + } + } + Ok(true) +} + +#[cfg(test)] fn json_value_to_line(value: serde_json::Value, scalar_column: &str) -> std::io::Result { + json_value_to_line_cancellable(value, scalar_column, &AtomicBool::new(false)) +} + +fn json_value_to_line_cancellable( + value: serde_json::Value, + scalar_column: &str, + cancelled: &AtomicBool, +) -> std::io::Result { let value = match value { value @ serde_json::Value::Object(_) => value, value => { @@ -1518,11 +1848,62 @@ fn json_value_to_line(value: serde_json::Value, scalar_column: &str) -> std::io: serde_json::Value::Object(row) } }; - let mut line = serde_json::to_vec(&value).map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))?; + let mut line = Vec::new(); + let (serialize_result, limit_exceeded) = { + let mut writer = CancellableJsonWriter { + inner: &mut line, + cancelled, + bytes_since_check: 0, + limit_exceeded: false, + }; + let serialize_result = serde_json::to_writer(&mut writer, &value); + (serialize_result, writer.limit_exceeded) + }; + if limit_exceeded { + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::OverMaxRecordSize)); + } + if let Err(error) = serialize_result { + ensure_json_parse_active(cancelled)?; + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, error)); + } + ensure_json_parse_active(cancelled)?; line.push(b'\n'); Ok(Bytes::from(line)) } +struct CancellableJsonWriter<'a> { + inner: &'a mut Vec, + cancelled: &'a AtomicBool, + bytes_since_check: usize, + limit_exceeded: bool, +} + +impl std::io::Write for CancellableJsonWriter<'_> { + fn write(&mut self, buffer: &[u8]) -> std::io::Result { + let Some(new_len) = self.inner.len().checked_add(buffer.len()) else { + self.limit_exceeded = true; + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::OverMaxRecordSize)); + }; + if new_len > MAX_SELECT_RECORD_BYTES { + self.limit_exceeded = true; + return Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::OverMaxRecordSize)); + } + self.bytes_since_check = self.bytes_since_check.saturating_add(buffer.len()); + if self.bytes_since_check >= JSON_CANCELLATION_CHECK_BYTES { + ensure_json_parse_active(self.cancelled)?; + self.bytes_since_check %= JSON_CANCELLATION_CHECK_BYTES; + } + self.inner.extend_from_slice(buffer); + Ok(buffer.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + ensure_json_parse_active(self.cancelled)?; + self.bytes_since_check = 0; + Ok(()) + } +} + fn invalid_json_source_path(message: &'static str) -> std::io::Error { std::io::Error::new(std::io::ErrorKind::InvalidData, message) } @@ -1552,6 +1933,7 @@ where S: Stream> + Send + 'static, E: Send + 'static, { + let input_metrics = input_metrics.recorder(); stream.inspect_ok(move |bytes| input_metrics.record_uncompressed(bytes.len())) } @@ -1611,14 +1993,16 @@ fn incomplete_object_stream_error(remaining: impl std::fmt::Display) -> o_Error #[cfg(test)] mod test { use super::{ - EcObjectStore, EcObjectStoreBuildError, JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER, OnceCell, + EcObjectStore, EcObjectStoreBuildError, JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER, JsonDocumentReadMode, OnceCell, SELECT_DEFAULT_READ_BUFFER_SIZE, SelectObjectOptions, SelectObjectSnapshot, SelectScanRange, SnapshotConsistencyError, - bytes_stream, convert_csv_delimiter_stream, convert_field_delimiter_stream, convert_record_delimiter_stream, - find_delimiter, flatten_json_document_to_ndjson, http_range_spec_from_get_range, json_document_ndjson_stream, - json_document_ndjson_stream_with_parser, legacy_json_source_from_input, map_storage_error, - meter_uncompressed_input_stream, scan_range_from_bounds, scan_range_stream, select_read_headers, snapshot_last_modified, - validate_json_document_size, + bytes_stream, compressed_json_document_ndjson_stream, convert_csv_delimiter_stream, convert_field_delimiter_stream, + convert_record_delimiter_stream, find_delimiter, flatten_json_document_to_ndjson, grow_json_document_buffer, + http_range_spec_from_get_range, json_document_ndjson_stream, json_document_ndjson_stream_with_parser, + json_key_eq_ignore_ascii_case_with_checkpoint, legacy_json_source_from_input, map_storage_error, + meter_uncompressed_input_stream, remove_json_source_key_with_checkpoint, scan_range_from_bounds, scan_range_stream, + select_read_headers, snapshot_last_modified, validate_json_document_size, }; + use crate::input_stream::{CompressionFormat, MAX_SELECT_RECORD_BYTES, compressed_input_reader, encode_compressed_fixture}; use crate::query::ast::{JsonPathSegment, JsonSource}; use crate::query::session::{QueryExecutionGuard, QueryExecutionOwner, QueryExecutionTracker}; use crate::storage_api::SelectPutObjReader; @@ -1627,7 +2011,7 @@ mod test { use bytes::Bytes; use datafusion::{ common::DataFusionError, - execution::memory_pool::{GreedyMemoryPool, MemoryLimit, MemoryPool, MemoryReservation}, + execution::memory_pool::{GreedyMemoryPool, MemoryConsumer, MemoryLimit, MemoryPool, MemoryReservation}, execution::{config::SessionConfig, context::SessionContext}, object_store::{self, GetOptions, GetRange, GetResultPayload, ObjectStore as _, path::Path}, physical_plan::ExecutionPlanProperties, @@ -1639,8 +2023,8 @@ mod test { use rustfs_test_utils::PutObjectCommitBarrier; use s3s::S3ErrorCode; use s3s::dto::{ - CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType, - OutputSerialization, ScanRange, SelectObjectContentInput, SelectObjectContentRequest, + CSVInput, CSVOutput, CompressionType, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, + JSONType, OutputSerialization, ScanRange, SelectObjectContentInput, SelectObjectContentRequest, }; use s3s::header::{ X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY, @@ -3128,7 +3512,7 @@ mod test { Arc::new(GreedyMemoryPool::new(1024 * 1024)), None, Arc::clone(&input_metrics), - snapshot, + Arc::clone(&snapshot), JsonSource::default(), ) .expect("build metrics-aware object store"); @@ -3175,6 +3559,152 @@ mod test { assert_eq!(input_metrics.snapshot().bytes_processed, 2); } + #[tokio::test] + async fn compressed_object_uses_one_full_stream_and_rejects_internal_ranges() { + const BUCKET: &str = "s3select-compressed-object"; + const OBJECT: &str = "input.csv"; + const DATA: &[u8] = b"id,name\n1,Alice\n2,Bob\n"; + + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, DATA).await; + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + let mut reader = SelectPutObjReader::from_vec(compressed.clone()); + env.ecstore + .put_object(BUCKET, OBJECT, &mut reader, &Default::default()) + .await + .expect("put compressed CSV fixture"); + let snapshot = prepare_test_snapshot(BUCKET, OBJECT).await; + let mut input = (*csv_input(BUCKET, OBJECT)).clone(); + input.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + input.request.scan_range = Some(ScanRange { + start: Some(0), + end: None, + }); + let input_metrics = Arc::new(SelectInputMetrics::default()); + let store = EcObjectStore::build_with_snapshot( + Arc::new(input), + Arc::new(GreedyMemoryPool::new(1024 * 1024)), + None, + Arc::clone(&input_metrics), + Arc::clone(&snapshot), + JsonSource::default(), + ) + .expect("build compressed object store"); + + let error = store + .get_opts( + &Path::from(OBJECT), + GetOptions { + range: Some(GetRange::Bounded(0..1)), + ..Default::default() + }, + ) + .await + .expect_err("compressed input must reject DataFusion byte ranges"); + assert_eq!( + QueryError::from(DataFusionError::ObjectStore(Box::new(error))).select_error(), + SelectError::UnsupportedScanRangeInput + ); + assert_eq!(store.reader_open_count.load(Ordering::SeqCst), 0); + + let mut scan_input = (*csv_input(BUCKET, OBJECT)).clone(); + scan_input.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + scan_input.request.scan_range = Some(ScanRange { + start: Some(1), + end: None, + }); + let scan_store = EcObjectStore::build_with_snapshot( + Arc::new(scan_input), + Arc::new(GreedyMemoryPool::new(1024 * 1024)), + None, + Arc::new(SelectInputMetrics::default()), + snapshot, + JsonSource::default(), + ) + .expect("build compressed ScanRange object store"); + let error = scan_store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect_err("compressed request ScanRange must fail before object I/O"); + assert_eq!( + QueryError::from(DataFusionError::ObjectStore(Box::new(error))).select_error(), + SelectError::UnsupportedScanRangeInput + ); + assert_eq!(scan_store.reader_open_count.load(Ordering::SeqCst), 0); + + let result = store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect("open full compressed object stream"); + let GetResultPayload::Stream(stream) = result.payload else { + panic!("expected compressed stream payload"); + }; + let body = stream + .try_collect::>() + .await + .expect("decode compressed object") + .concat(); + + assert_eq!(body, DATA); + assert_eq!(store.reader_open_count.load(Ordering::SeqCst), 1); + assert_eq!( + input_metrics.snapshot().bytes_scanned, + u64::try_from(compressed.len()).expect("compressed fixture length should fit in u64") + ); + assert_eq!( + input_metrics.snapshot().bytes_processed, + u64::try_from(DATA.len()).expect("input fixture length should fit in u64") + ); + } + + #[tokio::test] + async fn compressed_stream_throughput_is_independent_of_query_memory_pool() { + const BUCKET: &str = "s3select-compressed-throughput"; + const OBJECT: &str = "input.csv.gz"; + + let data = b"a\n".repeat(1024 * 1024); + let compressed = encode_compressed_fixture(CompressionFormat::Gzip, &data).await; + let env = crate::storage_api::select_test_ecstore_env().await; + env.make_bucket(BUCKET, false).await; + let mut reader = SelectPutObjReader::from_vec(compressed); + env.ecstore + .put_object(BUCKET, OBJECT, &mut reader, &Default::default()) + .await + .expect("put compressed throughput fixture"); + let snapshot = prepare_test_snapshot(BUCKET, OBJECT).await; + let mut input = (*csv_input(BUCKET, OBJECT)).clone(); + input.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + let metrics = Arc::new(SelectInputMetrics::default()); + let store = EcObjectStore::build_with_snapshot( + Arc::new(input), + Arc::new(GreedyMemoryPool::new(1)), + None, + Arc::clone(&metrics), + snapshot, + JsonSource::default(), + ) + .expect("build compressed throughput store"); + + let result = store + .get_opts(&Path::from(OBJECT), GetOptions::default()) + .await + .expect("open compressed throughput stream"); + let GetResultPayload::Stream(stream) = result.payload else { + panic!("expected compressed stream payload"); + }; + let decoded = stream + .try_collect::>() + .await + .expect("streamed decoded bytes should not consume the query memory pool") + .concat(); + + assert_eq!(decoded, data); + assert_eq!( + metrics.snapshot().bytes_processed, + u64::try_from(data.len()).expect("fixture length should fit in u64") + ); + } + #[tokio::test] async fn dropping_real_object_stream_counts_only_consumed_bytes() { const BUCKET: &str = "s3select-partial-input-metrics"; @@ -3353,6 +3883,125 @@ mod test { )); } + #[tokio::test] + async fn compressed_json_document_uses_decoded_size_and_metric_boundaries() { + const INPUT: &[u8] = br#"[{"id":1},{"id":2}]"#; + const EXPECTED: &[u8] = b"{\"id\":1}\n{\"id\":2}\n"; + + for format in [CompressionFormat::Gzip, CompressionFormat::Bzip2] { + let compressed = encode_compressed_fixture(format, INPUT).await; + let compressed_len = u64::try_from(compressed.len()).expect("compressed fixture length should fit in u64"); + let input_metrics = Arc::new(SelectInputMetrics::default()); + let reader = compressed_input_reader( + Box::new(std::io::Cursor::new(compressed)), + compressed_len, + format, + Arc::clone(&input_metrics), + u64::MAX, + None, + ); + let output = compressed_json_document_ndjson_stream( + reader, + JsonSource::default(), + Arc::new(GreedyMemoryPool::new(1024 * 1024)), + None, + ) + .try_collect::>() + .await + .expect("compressed JSON DOCUMENT should decode and parse") + .concat(); + + assert_eq!(output, EXPECTED); + assert_eq!(input_metrics.snapshot().bytes_scanned, compressed_len); + assert_eq!( + input_metrics.snapshot().bytes_processed, + u64::try_from(INPUT.len()).expect("JSON fixture length should fit in u64") + ); + } + } + + async fn compressed_json_document_select_error( + format: CompressionFormat, + compressed: Vec, + max_processed_bytes: u64, + ) -> SelectError { + let compressed_len = u64::try_from(compressed.len()).expect("compressed fixture length should fit in u64"); + let reader = compressed_input_reader( + Box::new(std::io::Cursor::new(compressed)), + compressed_len, + format, + Arc::new(SelectInputMetrics::default()), + max_processed_bytes, + None, + ); + let mut output = compressed_json_document_ndjson_stream( + reader, + JsonSource::default(), + Arc::new(GreedyMemoryPool::new(1024 * 1024)), + None, + ); + let source = output + .next() + .await + .expect("decoder failure should produce one stream error") + .expect_err("compressed JSON DOCUMENT decoding must fail"); + assert!(output.next().await.is_none()); + QueryError::from(DataFusionError::ObjectStore(Box::new(source))).select_error() + } + + #[tokio::test] + async fn compressed_json_document_preserves_decoder_select_errors() { + const INPUT: &[u8] = br#"[{"id":1}]"#; + + for (format, compression) in [(CompressionFormat::Gzip, "GZIP"), (CompressionFormat::Bzip2, "BZIP2")] { + assert_eq!( + compressed_json_document_select_error(format, b"not compressed".to_vec(), u64::MAX).await, + SelectError::InvalidCompressionFormatForObject { compression } + ); + + let mut truncated = encode_compressed_fixture(format, INPUT).await; + truncated.truncate(truncated.len() - 1); + assert_eq!( + compressed_json_document_select_error(format, truncated, u64::MAX).await, + SelectError::TruncatedInput + ); + + let compressed = encode_compressed_fixture(format, INPUT).await; + let max_processed_bytes = u64::try_from(INPUT.len() - 1).expect("fixture length should fit in u64"); + assert_eq!( + compressed_json_document_select_error(format, compressed, max_processed_bytes).await, + SelectError::ResourceExhausted + ); + } + } + + #[test] + fn compressed_json_document_buffer_grows_amortized_and_reserves_capacity() { + let memory_pool: Arc = Arc::new(GreedyMemoryPool::new(1024 * 1024)); + let reservation = MemoryConsumer::new("compressed JSON document test").register(&memory_pool); + let json_source = JsonSource::default(); + let mut buffer = Vec::new(); + let mut capacity_growths = 0; + + for _ in 0..1025 { + let old_capacity = buffer.capacity(); + let required_len = buffer.len() + 1; + grow_json_document_buffer(&mut buffer, required_len, &reservation, &json_source) + .expect("bounded JSON buffer should grow"); + if buffer.capacity() != old_capacity { + capacity_growths += 1; + } + buffer.push(0); + assert_eq!( + reservation.size(), + super::json_document_memory_reservation_bytes(buffer.capacity(), &json_source) + .expect("test reservation should fit") + ); + } + + assert!(capacity_growths <= 12, "power-of-two growth should stay logarithmic"); + } + #[tokio::test] async fn scalar_alias_expansion_is_in_the_query_memory_reservation() { let input = b"[0,0]".to_vec(); @@ -3489,12 +4138,14 @@ mod test { Arc::new(GreedyMemoryPool::new(input.len() * JSON_DOCUMENT_MEMORY_RESERVATION_MULTIPLIER)); let mut output = json_document_ndjson_stream_with_parser( Box::new(std::io::Cursor::new(input.clone())), - input.len() as u64, + JsonDocumentReadMode::Exact { + original_size: input.len() as u64, + input_metrics: SelectInputMetrics::default().recorder(), + }, JsonSource::default(), - Arc::new(SelectInputMetrics::default()), memory_pool, None, - |_, _| Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::AmbiguousFieldName)), + |_, _, _| Err(std::io::Error::new(std::io::ErrorKind::InvalidData, SelectError::AmbiguousFieldName)), ); let source = output @@ -3562,7 +4213,7 @@ mod test { } #[test] - fn test_json_document_queued_parse_retains_query_guard_until_dequeued() { + fn test_json_document_cancelled_queued_parse_releases_before_dequeue() { let runtime = tokio::runtime::Builder::new_multi_thread() .worker_threads(2) .max_blocking_threads(1) @@ -3614,29 +4265,22 @@ mod test { } drop(output); - assert_eq!(admission.available_permits(), 0); - assert!(memory_pool.reserved() > 0); assert!( + tokio::time::timeout(std::time::Duration::from_millis(100), reservation_released) + .await + .expect("queued parse cancellation should release its memory immediately") + .expect("memory reservation release observer should remain open"), + "query admission must cover the memory reservation through teardown" + ); + assert_eq!(memory_pool.reserved(), 0); + let recovered_permit = tokio::time::timeout(std::time::Duration::from_millis(100), Arc::clone(&admission).acquire_owned()) .await - .is_err(), - "queued parse resources must remain covered by admission" - ); + .expect("queued parse cancellation should release admission before a worker is available") + .expect("query admission should remain open"); + release_blocking_tx.send(()).expect("release blocking worker"); blocker.await.expect("blocking worker should finish"); - assert!( - tokio::time::timeout(std::time::Duration::from_secs(5), reservation_released) - .await - .expect("cancelled JSON parse should release its memory reservation") - .expect("memory reservation release observer should remain open"), - "query admission must cover the memory reservation through task teardown" - ); - let recovered_permit = - tokio::time::timeout(std::time::Duration::from_secs(5), Arc::clone(&admission).acquire_owned()) - .await - .expect("cancelled JSON parse should be dequeued") - .expect("query admission should remain open"); - assert_eq!(memory_pool.reserved(), 0); drop(recovered_permit); assert_eq!(admission.available_permits(), 1); }); @@ -3679,12 +4323,14 @@ mod test { 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, + JsonDocumentReadMode::Exact { + original_size: input.len() as u64, + input_metrics: SelectInputMetrics::default().recorder(), + }, JsonSource::default(), - Arc::new(SelectInputMetrics::default()), Arc::clone(&memory_pool), Some(query_tracker.clone()), - move |_, _| { + move |_, _, _| { parser_started_in_task.store(true, std::sync::atomic::Ordering::SeqCst); Ok(vec![Bytes::from_static(b"{}\n")]) }, @@ -3756,12 +4402,14 @@ mod test { 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, + JsonDocumentReadMode::Exact { + original_size: input.len() as u64, + input_metrics: SelectInputMetrics::default().recorder(), + }, JsonSource::default(), - Arc::new(SelectInputMetrics::default()), Arc::clone(&memory_pool), Some(query_tracker), - move |_, _| { + move |_, _, _| { parser_started_in_task.store(true, std::sync::atomic::Ordering::SeqCst); Ok(vec![Bytes::from_static(b"{}\n")]) }, @@ -3787,7 +4435,7 @@ mod test { } #[test] - fn test_json_document_started_parse_retains_query_guard_when_cancelled() { + fn test_json_document_started_parse_cancels_and_releases_query_guard() { let runtime = tokio::runtime::Builder::new_multi_thread() .worker_threads(2) .max_blocking_threads(1) @@ -3812,18 +4460,21 @@ mod test { 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, + JsonDocumentReadMode::Exact { + original_size: input.len() as u64, + input_metrics: SelectInputMetrics::default().recorder(), + }, JsonSource::default(), - Arc::new(SelectInputMetrics::default()), memory_pool, Some(query_tracker), - move |_, _| { + move |_, _, cancellation| { let _ = parse_started_tx.send(()); - release_parse_rx.recv().expect("release JSON parser"); - Ok(vec![Bytes::from_static(b"{}\n")]) + while !cancellation.load(std::sync::atomic::Ordering::Acquire) { + std::thread::yield_now(); + } + Err(std::io::Error::new(std::io::ErrorKind::Interrupted, SelectError::Canceled)) }, ); @@ -3833,14 +4484,13 @@ mod test { assert!(futures::poll!(next.as_mut()).is_pending()); } parse_started_rx.await.expect("JSON parser should start"); + assert!(Arc::clone(&admission).try_acquire_owned().is_err()); 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("cancelled JSON parse should release the query guard without an external unblock") .expect("query admission should remain open"); drop(recovered_permit); assert_eq!(admission.available_permits(), 1); @@ -3931,6 +4581,32 @@ mod test { assert_eq!(err.kind(), std::io::ErrorKind::InvalidData); } + #[test] + fn json_document_logical_record_enforces_one_megabyte_limit() { + const OBJECT_OVERHEAD: usize = br#"{"v":""}"#.len(); + + let at_limit = serde_json::json!({"v": "x".repeat(MAX_SELECT_RECORD_BYTES - OBJECT_OVERHEAD)}); + let line = super::json_value_to_line(at_limit, "_1").expect("one-megabyte logical record should be accepted"); + assert_eq!(line.len(), MAX_SELECT_RECORD_BYTES + 1); + + let over_limit = serde_json::json!({"v": "x".repeat(MAX_SELECT_RECORD_BYTES + 1 - OBJECT_OVERHEAD)}); + let error = super::json_value_to_line(over_limit, "_1").expect_err("oversized logical record must fail"); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(error.get_ref().is_some_and(|source| { + source + .downcast_ref::() + .is_some_and(|error| error == &SelectError::OverMaxRecordSize) + })); + + let escaped = serde_json::json!({"v": "\0".repeat(MAX_SELECT_RECORD_BYTES / 6)}); + let error = super::json_value_to_line(escaped, "_1").expect_err("escaped output must be bounded while it is serialized"); + assert!(error.get_ref().is_some_and(|source| { + source + .downcast_ref::() + .is_some_and(|error| error == &SelectError::OverMaxRecordSize) + })); + } + /// Completely empty input returns an error (not valid JSON). #[test] fn test_flatten_empty_input_returns_error() { @@ -4156,6 +4832,55 @@ mod test { ); } + #[test] + fn unquoted_source_key_scan_honors_cancellation() { + let mut object = serde_json::Map::new(); + for index in 0..=super::JSON_CANCELLATION_CHECK_KEYS { + object.insert(format!("field-{index}"), serde_json::Value::Null); + } + let mut checkpoints = 0; + let mut cancel_on_second_checkpoint = || { + checkpoints += 1; + if checkpoints == 2 { + Err(std::io::Error::new(std::io::ErrorKind::Interrupted, SelectError::Canceled)) + } else { + Ok(()) + } + }; + + let error = remove_json_source_key_with_checkpoint(&mut object, "x", false, &mut cancel_on_second_checkpoint) + .expect_err("key scan should stop at its second cancellation checkpoint"); + + assert_eq!(error.kind(), std::io::ErrorKind::Interrupted); + assert_eq!(checkpoints, 2); + assert!( + error + .get_ref() + .and_then(|source| source.downcast_ref::()) + .is_some_and(|error| *error == SelectError::Canceled) + ); + } + + #[test] + fn long_json_key_comparison_honors_cancellation() { + let key = "x".repeat(super::JSON_CANCELLATION_CHECK_BYTES * 2); + let mut checkpoints = 0; + let mut cancel_on_second_checkpoint = || { + checkpoints += 1; + if checkpoints == 2 { + Err(std::io::Error::new(std::io::ErrorKind::Interrupted, SelectError::Canceled)) + } else { + Ok(()) + } + }; + + let error = json_key_eq_ignore_ascii_case_with_checkpoint(&key, &key, &mut cancel_on_second_checkpoint) + .expect_err("long-key comparison should stop at its second cancellation checkpoint"); + + assert_eq!(error.kind(), std::io::ErrorKind::Interrupted); + assert_eq!(checkpoints, 2); + } + #[test] fn missing_source_key_produces_no_records() { let input = br#"{"employees":[]}"#; diff --git a/crates/s3select-api/src/query/session.rs b/crates/s3select-api/src/query/session.rs index dc3a66dc0..7df3cf60a 100644 --- a/crates/s3select-api/src/query/session.rs +++ b/crates/s3select-api/src/query/session.rs @@ -30,6 +30,7 @@ use datafusion::{ prelude::SessionContext, }; use parking_lot::Mutex; +use s3s::dto::CompressionType; use std::sync::{ Arc, Weak, atomic::{AtomicU8, Ordering}, @@ -446,11 +447,19 @@ impl SessionCtxFactory { let scan_range_requires_single_file_scan = context.input.request.scan_range.is_some() && context.input.request.input_serialization.parquet.is_none(); let json_document_requires_single_file_scan = is_json_document_input(&context.input); + let compressed_input_requires_single_file_scan = context + .input + .request + .input_serialization + .compression_type + .as_ref() + .is_some_and(|compression| compression.as_str() != CompressionType::NONE); let metered_input_requires_single_file_scan = input_metrics.is_some() && context.input.request.input_serialization.parquet.is_none(); let config = if custom_two_byte_record_delimiter || scan_range_requires_single_file_scan || json_document_requires_single_file_scan + || compressed_input_requires_single_file_scan || metered_input_requires_single_file_scan { config.with_repartition_file_scans(false) @@ -847,6 +856,21 @@ mod tests { assert!(session.inner().config().options().optimizer.repartition_file_scans); } + #[tokio::test] + async fn compressed_input_disables_file_repartitioning_without_metrics() { + let mut context = test_context(); + Arc::make_mut(&mut context.input).request.input_serialization.compression_type = + Some(CompressionType::from_static(CompressionType::GZIP)); + + let session = SessionCtxFactory::new(true) + .with_target_partitions(2) + .create_session_ctx(&context) + .await + .expect("compressed session should be created"); + + assert!(!session.inner().config().options().optimizer.repartition_file_scans); + } + #[tokio::test] async fn json_document_disables_file_repartitioning() { let mut context = test_context(); diff --git a/crates/s3select-query/src/dispatcher/manager.rs b/crates/s3select-query/src/dispatcher/manager.rs index 72117ce61..0f69949ff 100644 --- a/crates/s3select-query/src/dispatcher/manager.rs +++ b/crates/s3select-query/src/dispatcher/manager.rs @@ -53,7 +53,7 @@ use rustfs_s3select_api::{ }, }, }; -use s3s::dto::{FileHeaderInfo, JSONType, SelectObjectContentInput}; +use s3s::dto::{CompressionType, FileHeaderInfo, JSONType, SelectObjectContentInput}; use std::sync::LazyLock; use tokio::{ sync::Semaphore, @@ -72,6 +72,7 @@ use crate::{ static IGNORE: LazyLock = LazyLock::new(|| FileHeaderInfo::from_static(FileHeaderInfo::IGNORE)); static NONE: LazyLock = LazyLock::new(|| FileHeaderInfo::from_static(FileHeaderInfo::NONE)); static USE: LazyLock = LazyLock::new(|| FileHeaderInfo::from_static(FileHeaderInfo::USE)); +const EXACT_OBJECT_FILE_EXTENSION: &str = ""; #[derive(Clone)] pub struct SimpleQueryDispatcher { @@ -416,6 +417,13 @@ impl SimpleQueryDispatcher { let path = format!("s3://{}/{}", self.input.bucket, self.input.key); let table_path = ListingTableUrl::parse(path)?; + let compressed_input = self + .input + .request + .input_serialization + .compression_type + .as_ref() + .is_some_and(|compression| compression.as_str() != CompressionType::NONE); let (listing_options, need_rename_volume_name, need_ignore_volume_name) = if let Some(csv) = self.input.request.input_serialization.csv.as_ref() { let mut need_rename_volume_name = false; @@ -465,22 +473,30 @@ impl SimpleQueryDispatcher { file_format = file_format.with_quote(quote.as_bytes().first().copied().unwrap_or_default()); } ( - ListingOptions::new(Arc::new(file_format)).with_file_extension(".csv"), + ListingOptions::new(Arc::new(file_format)).with_file_extension(if compressed_input { + EXACT_OBJECT_FILE_EXTENSION + } else { + ".csv" + }), need_rename_volume_name, need_ignore_volume_name, ) } else if self.input.request.input_serialization.json.is_some() { let file_format = JsonFormat::default(); - // Use the actual file extension from the object key so that files stored - // with a `.jsonl` suffix (newline-delimited JSON) are also matched by - // DataFusion's listing/schema-inference logic. Falling back to ".json" - // preserves behaviour for keys that have no extension. - let file_ext = std::path::Path::new(&self.input.key) - .extension() - .and_then(|e| e.to_str()) - .map(|e| format!(".{e}")) - .unwrap_or_else(|| ".json".to_string()); - (ListingOptions::new(Arc::new(file_format)).with_file_extension(file_ext), false, false) + let file_extension = if compressed_input { + EXACT_OBJECT_FILE_EXTENSION.to_string() + } else { + std::path::Path::new(&self.input.key) + .extension() + .and_then(|extension| extension.to_str()) + .map(|extension| format!(".{extension}")) + .unwrap_or_else(|| ".json".to_string()) + }; + ( + ListingOptions::new(Arc::new(file_format)).with_file_extension(file_extension), + false, + false, + ) } else { return Err(SelectError::InvalidDataSource.into()); }; diff --git a/rustfs/src/app/select_object.rs b/rustfs/src/app/select_object.rs index b4aa98408..68ce59f4d 100644 --- a/rustfs/src/app/select_object.rs +++ b/rustfs/src/app/select_object.rs @@ -22,7 +22,7 @@ use futures::StreamExt; use http::{HeaderMap, StatusCode, header::RANGE}; use rustfs_s3select_api::{ QueryError, SelectError, SelectInputMetrics, - object_store::{INVALID_SCAN_RANGE_MESSAGE, validate_scan_range_bounds}, + object_store::{INVALID_SCAN_RANGE_MESSAGE, is_noop_scan_range, validate_scan_range_bounds}, query::{Context, Query}, }; use rustfs_s3select_query::instance::s3_select_query_timeout; @@ -44,6 +44,8 @@ const RECORDS_CHUNK_TARGET: usize = 128 * 1024; const DATA_SOURCE_PATH_UNSUPPORTED_CODE: &str = "DataSourcePathUnsupported"; const INVALID_QUERY_CODE: &str = "InvalidQuery"; const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure"; +const INVALID_REQUEST_PARAMETER_MESSAGE: &str = + "The value of a parameter in the SelectRequest element is invalid. Check the service API documentation and try again."; const BUSY_MESSAGE: &str = "The service is unavailable. Try again later."; const EMPTY_SELECT_EXPRESSION_MESSAGE: &str = "empty SQL expression"; const SLOW_DOWN_MESSAGE: &str = "Reduce your request rate."; @@ -103,7 +105,11 @@ pub async fn execute_select_object_content( ) .await .map_err(|_| select_query_timeout_error(query_timeout.as_secs()))??; - validate_scan_range_for_object_size(&input.request, snapshot.logical_size())?; + let object_size = snapshot.logical_size(); + validate_scan_range_for_object_size(&input.request, object_size)?; + if object_size == 0 && is_compressed_input(&input.request.input_serialization) { + return Err(map_select_error_to_s3(&SelectError::TruncatedInput)); + } let snapshot = Arc::new(snapshot); let query = Query::new_with_snapshot(Context { input: input.clone() }, input.request.expression.clone(), Arc::clone(&snapshot)); @@ -262,7 +268,14 @@ fn validate_select_request(headers: &http::HeaderMap, input: &mut SelectObjectCo } normalize_input_serialization(&mut input.request.input_serialization)?; + let compressed_input = is_compressed_input(&input.request.input_serialization); + if compressed_input && input.request.scan_range.as_ref().is_some_and(is_noop_scan_range) { + input.request.scan_range = None; + } validate_scan_range(&input.request)?; + if compressed_input && input.request.scan_range.is_some() { + return Err(map_select_error_to_s3(&SelectError::UnsupportedScanRangeInput)); + } let output_format = normalize_output_serialization(&mut input.request.output_serialization)?; if input.request.expression.trim().is_empty() { @@ -284,6 +297,13 @@ fn validate_select_request(headers: &http::HeaderMap, input: &mut SelectObjectCo }) } +fn is_compressed_input(input: &InputSerialization) -> bool { + input + .compression_type + .as_ref() + .is_some_and(|compression| compression.as_str() != CompressionType::NONE) +} + fn normalize_input_serialization(input: &mut InputSerialization) -> S3Result<()> { let format_count = usize::from(input.csv.is_some()) + usize::from(input.json.is_some()) + usize::from(input.parquet.is_some()); @@ -298,15 +318,19 @@ fn normalize_input_serialization(input: &mut InputSerialization) -> S3Result<()> match compression.as_str() { CompressionType::NONE => {} CompressionType::GZIP | CompressionType::BZIP2 => { - return Err(s3_error!( - NotImplemented, - "SelectObjectContent currently supports only uncompressed input" - )); + if input.parquet.is_some() { + return Err(S3Error::with_message( + S3ErrorCode::InvalidRequestParameter, + INVALID_REQUEST_PARAMETER_MESSAGE, + )); + } } _ => return Err(map_select_error_to_s3(&SelectError::InvalidCompressionFormat)), } } - input.compression_type = Some(CompressionType::from_static(CompressionType::NONE)); + input + .compression_type + .get_or_insert_with(|| CompressionType::from_static(CompressionType::NONE)); if let Some(csv) = input.csv.as_mut() { if csv.allow_quoted_record_delimiter.unwrap_or(false) { @@ -667,11 +691,16 @@ fn map_query_error_to_s3(err: QueryError) -> S3Error { fn map_select_error_to_s3(err: &SelectError) -> S3Error { match err { SelectError::InvalidCompressionFormat => S3Error::with_message(S3ErrorCode::InvalidCompressionFormat, err.to_string()), + SelectError::InvalidCompressionFormatForObject { .. } => { + S3Error::with_message(S3ErrorCode::InvalidCompressionFormat, err.to_string()) + } SelectError::InvalidDataSource => S3Error::with_message(S3ErrorCode::InvalidDataSource, err.to_string()), SelectError::TruncatedInput => S3Error::with_message(S3ErrorCode::TruncatedInput, err.to_string()), + SelectError::UnsupportedScanRangeInput => S3Error::with_message(S3ErrorCode::UnsupportedScanRangeInput, err.to_string()), SelectError::CsvParsingError => S3Error::with_message(S3ErrorCode::CSVParsingError, err.to_string()), SelectError::JsonParsingError => S3Error::with_message(S3ErrorCode::JSONParsingError, err.to_string()), SelectError::ParquetParsingError => S3Error::with_message(S3ErrorCode::ParquetParsingError, err.to_string()), + SelectError::OverMaxRecordSize => S3Error::with_message(S3ErrorCode::OverMaxRecordSize, err.to_string()), SelectError::ParseSelectFailure { message } => custom_bad_request(PARSE_SELECT_FAILURE_CODE, message.clone()), SelectError::InvalidQuery => custom_bad_request(INVALID_QUERY_CODE, err.to_string()), SelectError::InvalidDataType => S3Error::with_message(S3ErrorCode::InvalidDataType, err.to_string()), @@ -995,8 +1024,20 @@ mod tests { S3ErrorCode::InvalidCompressionFormat, StatusCode::BAD_REQUEST, ), + ( + SelectError::InvalidCompressionFormatForObject { + compression: CompressionType::GZIP, + }, + S3ErrorCode::InvalidCompressionFormat, + StatusCode::BAD_REQUEST, + ), (SelectError::InvalidDataSource, S3ErrorCode::InvalidDataSource, StatusCode::BAD_REQUEST), (SelectError::TruncatedInput, S3ErrorCode::TruncatedInput, StatusCode::BAD_REQUEST), + ( + SelectError::UnsupportedScanRangeInput, + S3ErrorCode::UnsupportedScanRangeInput, + StatusCode::BAD_REQUEST, + ), (SelectError::CsvParsingError, S3ErrorCode::CSVParsingError, StatusCode::BAD_REQUEST), (SelectError::JsonParsingError, S3ErrorCode::JSONParsingError, StatusCode::BAD_REQUEST), ( @@ -1004,6 +1045,7 @@ mod tests { S3ErrorCode::ParquetParsingError, StatusCode::BAD_REQUEST, ), + (SelectError::OverMaxRecordSize, S3ErrorCode::OverMaxRecordSize, StatusCode::BAD_REQUEST), ( SelectError::ParseSelectFailure { message: "invalid SELECT expression".to_string(), @@ -1372,6 +1414,12 @@ mod tests { assert_eq!(compression_status, StatusCode::BAD_REQUEST); assert!(compression_body.contains("InvalidCompressionFormat")); assert!(compression_body.contains("")); + + let scan_range_error = map_select_error_to_s3(&SelectError::UnsupportedScanRangeInput); + let (scan_range_status, scan_range_body) = http_xml_error(scan_range_error).await; + assert_eq!(scan_range_status, StatusCode::BAD_REQUEST); + assert!(scan_range_body.contains("UnsupportedScanRangeInput")); + assert!(scan_range_body.contains("Scan range queries are not supported on this type of object.")); } #[tokio::test(start_paused = true)] @@ -1625,6 +1673,89 @@ mod tests { ); } + #[test] + fn validate_preserves_supported_compression_for_csv_and_json_lines() { + for compression in [CompressionType::GZIP, CompressionType::BZIP2] { + let mut csv_input = base_input(); + csv_input.request.input_serialization.compression_type = Some(CompressionType::from_static(compression)); + validate_select_request(&HeaderMap::new(), &mut csv_input).expect("compressed CSV should be accepted"); + assert_eq!( + csv_input + .request + .input_serialization + .compression_type + .as_ref() + .map(|value| value.as_str()), + Some(compression) + ); + + let mut json_input = base_input(); + json_input.request.input_serialization.csv = None; + json_input.request.input_serialization.json = Some(JSONInput { + type_: Some(JSONType::from_static(JSONType::LINES)), + }); + json_input.request.input_serialization.compression_type = Some(CompressionType::from_static(compression)); + validate_select_request(&HeaderMap::new(), &mut json_input).expect("compressed JSON LINES should be accepted"); + assert_eq!( + json_input + .request + .input_serialization + .compression_type + .as_ref() + .map(|value| value.as_str()), + Some(compression) + ); + } + } + + #[test] + fn validate_rejects_parquet_compression_with_select_request_error() { + let mut input = base_input(); + input.request.input_serialization.csv = None; + input.request.input_serialization.parquet = Some(ParquetInput {}); + input.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + + let error = validate_select_request(&HeaderMap::new(), &mut input).expect_err("compressed Parquet must fail"); + + assert_eq!(error.code(), &S3ErrorCode::InvalidRequestParameter); + assert_eq!(error.message(), Some(INVALID_REQUEST_PARAMETER_MESSAGE)); + } + + #[test] + fn validate_normalizes_noop_compressed_scan_range_and_rejects_real_ranges() { + let mut noop = base_input(); + noop.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + noop.request.scan_range = Some(ScanRange { + start: Some(0), + end: None, + }); + validate_select_request(&HeaderMap::new(), &mut noop).expect("zero-start full scan should be normalized"); + assert!(noop.request.scan_range.is_none()); + + let mut ranged = base_input(); + ranged.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + ranged.request.scan_range = Some(ScanRange { + start: Some(1), + end: None, + }); + let error = validate_select_request(&HeaderMap::new(), &mut ranged) + .expect_err("compressed input with an effective ScanRange must fail before object I/O"); + assert_eq!(error.code(), &S3ErrorCode::UnsupportedScanRangeInput); + assert_eq!(error.status_code(), Some(StatusCode::BAD_REQUEST)); + assert_eq!(error.message(), Some("Scan range queries are not supported on this type of object.")); + + let mut malformed = base_input(); + malformed.request.input_serialization.compression_type = Some(CompressionType::from_static(CompressionType::GZIP)); + malformed.request.scan_range = Some(ScanRange { + start: Some(10), + end: Some(1), + }); + let error = validate_select_request(&HeaderMap::new(), &mut malformed) + .expect_err("malformed ScanRange must fail before compression compatibility validation"); + assert_eq!(error.code(), &S3ErrorCode::InvalidRequestParameter); + assert_eq!(error.message(), Some(INVALID_SCAN_RANGE_MESSAGE)); + } + #[test] fn validate_rejects_unknown_csv_header_mode_before_streaming() { let mut input = base_input();