diff --git a/crates/e2e_test/src/reliant/sql.rs b/crates/e2e_test/src/reliant/sql.rs index ab2d385ab..d779363e7 100644 --- a/crates/e2e_test/src/reliant/sql.rs +++ b/crates/e2e_test/src/reliant/sql.rs @@ -122,6 +122,24 @@ async fn select_json_document(client: &Client, key: &str, expression: &str) -> T process_select_response(response).await } +fn csv_select_request( + client: &Client, + key: &str, +) -> aws_sdk_s3::operation::select_object_content::builders::SelectObjectContentFluentBuilder { + client + .select_object_content() + .bucket(BUCKET) + .key(key) + .expression("SELECT * FROM S3Object") + .expression_type(ExpressionType::Sql) + .input_serialization( + InputSerialization::builder() + .csv(CsvInput::builder().file_header_info(FileHeaderInfo::Use).build()) + .build(), + ) + .output_serialization(OutputSerialization::builder().csv(CsvOutput::builder().build()).build()) +} + async fn process_select_response( mut event_stream: aws_sdk_s3::operation::select_object_content::SelectObjectContentOutput, ) -> TestResult { @@ -188,30 +206,42 @@ async fn assert_input_byte_stats( let mut last_progress: Option = None; let mut stats = None; let mut saw_end = false; - while let Some(event) = payload.recv().await? { - match event { - aws_sdk_s3::types::SelectObjectContentEventStream::Records(records) => { - if let Some(bytes) = records.payload { - records_len = records_len.saturating_add(u64::try_from(bytes.as_ref().len())?); + tokio::time::timeout(SELECT_RESPONSE_TIMEOUT, async { + // The AWS SDK validates both event-stream CRCs before yielding an event. + while let Some(event) = payload.recv().await? { + assert!(!saw_end, "Select emitted an event after End"); + match event { + aws_sdk_s3::types::SelectObjectContentEventStream::Records(records) => { + assert!(stats.is_none(), "Select emitted Records after Stats"); + if let Some(bytes) = records.payload { + records_len = records_len.saturating_add(u64::try_from(bytes.as_ref().len())?); + } } - } - aws_sdk_s3::types::SelectObjectContentEventStream::Progress(event) => { - let details = event.details.ok_or("Progress event did not contain details")?; - if let Some(previous) = last_progress.as_ref() { - assert!(details.bytes_scanned() >= previous.bytes_scanned()); - assert!(details.bytes_processed() >= previous.bytes_processed()); - assert!(details.bytes_returned() >= previous.bytes_returned()); + aws_sdk_s3::types::SelectObjectContentEventStream::Progress(event) => { + assert!(stats.is_none(), "Select emitted Progress after Stats"); + let details = event.details.ok_or("Progress event did not contain details")?; + if let Some(previous) = last_progress.as_ref() { + assert!(details.bytes_scanned() >= previous.bytes_scanned()); + assert!(details.bytes_processed() >= previous.bytes_processed()); + assert!(details.bytes_returned() >= previous.bytes_returned()); + } + last_progress = Some(details); } - last_progress = Some(details); + aws_sdk_s3::types::SelectObjectContentEventStream::Stats(event) => { + assert!(stats.is_none(), "Select emitted more than one Stats event"); + stats = event.details; + } + aws_sdk_s3::types::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"), } - aws_sdk_s3::types::SelectObjectContentEventStream::Stats(event) => stats = event.details, - aws_sdk_s3::types::SelectObjectContentEventStream::End(_) => { - saw_end = true; - break; - } - _ => {} } - } + Ok::<(), Box>(()) + }) + .await + .map_err(|_| -> Box { "Select response timed out".into() })??; let stats = stats.ok_or("Select response ended without a Stats event")?; let input_len = i64::try_from(body.len())?; @@ -219,10 +249,11 @@ async fn assert_input_byte_stats( assert_eq!(stats.bytes_processed(), Some(input_len)); assert_eq!(stats.bytes_returned(), Some(i64::try_from(records_len)?)); if progress_enabled { - let progress = last_progress.ok_or("Select response ended without a Progress event")?; - assert_eq!(progress.bytes_scanned(), stats.bytes_scanned()); - assert_eq!(progress.bytes_processed(), stats.bytes_processed()); - assert_eq!(progress.bytes_returned(), stats.bytes_returned()); + if let Some(progress) = last_progress { + assert!(stats.bytes_scanned() >= progress.bytes_scanned()); + assert!(stats.bytes_processed() >= progress.bytes_processed()); + assert!(stats.bytes_returned() >= progress.bytes_returned()); + } } else { assert!(last_progress.is_none(), "disabled request progress emitted a Progress event"); } @@ -231,7 +262,7 @@ async fn assert_input_byte_stats( } #[tokio::test(flavor = "multi_thread", worker_threads = 4)] -async fn test_select_object_content_reports_input_byte_stats() -> TestResult<()> { +async fn test_select_object_content_http_event_order_crc_and_input_byte_stats() -> TestResult<()> { const CSV_BODY: &[u8] = b"name,age\nAlice,30\nBob,25\n"; const JSON_LINES_BODY: &[u8] = b"{\"name\":\"Alice\"}\n{\"name\":\"Bob\"}\n"; const JSON_DOCUMENT_BODY: &[u8] = b"[{\"name\":\"Alice\"},{\"name\":\"Bob\"}]"; @@ -289,6 +320,60 @@ async fn test_select_object_content_reports_input_byte_stats() -> TestResult<()> Ok(()) } +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn test_select_object_content_http_disconnect_releases_query() -> TestResult<()> { + const OBJECT: &str = "disconnect.csv"; + const ROWS: usize = 16 * 1024; + const RELEASE_BACKOFF: Duration = Duration::from_millis(25); + + init_logging(); + let mut env = RustFSTestEnvironment::new().await?; + env.start_rustfs_server_with_env(vec![], &[("RUSTFS_S3SELECT_MAX_CONCURRENT_QUERIES", "1")]) + .await?; + let client = env.create_s3_client(); + setup_test_bucket(&client).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()); + } + client + .put_object() + .bucket(BUCKET) + .key(OBJECT) + .body(Bytes::from(body).into()) + .send() + .await?; + + // Leaving this response body unread fills the bounded HTTP/event channels before the query can finish. + let first = csv_select_request(&client, OBJECT).send().await?; + let saturated = csv_select_request(&client, OBJECT) + .send() + .await + .expect_err("the first HTTP stream 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 { + loop { + match csv_select_request(&client, OBJECT).send().await { + Ok(response) => return Ok::<_, Box>(response), + Err(error) if error.as_service_error().and_then(ProvideErrorMetadata::code) == Some("SlowDown") => { + tokio::time::sleep(RELEASE_BACKOFF).await; + } + Err(error) => return Err(format!("unexpected Select error after disconnect: {error}").into()), + } + } + }) + .await + .map_err(|_| -> Box { "disconnected Select did not release its query permit".into() })??; + drop(second); + + Ok(()) +} + #[tokio::test(flavor = "multi_thread", worker_threads = 4)] async fn test_select_object_content_csv_basic() -> TestResult<()> { let (_env, client) = create_test_environment().await?; diff --git a/rustfs/src/app/select_object.rs b/rustfs/src/app/select_object.rs index 68ce59f4d..48ea1ce50 100644 --- a/rustfs/src/app/select_object.rs +++ b/rustfs/src/app/select_object.rs @@ -9,11 +9,17 @@ use super::storage_api::select_object::{ }; use crate::app::runtime_sources::current_s3select_db; use crate::error::ApiError; -use bytes::Bytes; +use bytes::{Bytes, BytesMut}; use datafusion::arrow::{ - csv::{QuoteStyle, WriterBuilder as CsvWriterBuilder, writer::Terminator}, - json::{WriterBuilder as JsonWriterBuilder, writer::LineDelimited}, + array::{Array, ListLikeArray, MapArray, cast::AsArray}, + datatypes::{ + ArrowNativeType, DataType, FieldRef, Int8Type, Int16Type, Int32Type, Int64Type, UInt8Type, UInt16Type, UInt32Type, + UInt64Type, + }, + error::ArrowError, + json::writer::{EncoderOptions, NullableEncoder, make_encoder}, record_batch::RecordBatch, + util::display::{ArrayFormatter, FormatOptions}, }; #[cfg(test)] use datafusion::common::DataFusionError; @@ -33,14 +39,31 @@ use s3s::dto::{ StatsEvent, }; use s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error}; -use std::sync::Arc; +use std::{ + fmt, + future::poll_fn, + io::{self, Write}, + ops::Range, + pin::Pin, + sync::Arc, + time::Duration, +}; use tokio::sync::mpsc; -use tokio::time::{Instant, timeout_at}; +use tokio::time::{Instant, Interval, MissedTickBehavior, Sleep, timeout_at}; use tokio_stream::wrappers::ReceiverStream; +use tokio_util::sync::PollSender; use tracing::info; const MAX_SELECT_EXPRESSION_BYTES: usize = 256 * 1024; -const RECORDS_CHUNK_TARGET: usize = 128 * 1024; +const MAX_COMPAT_EVENT_STREAM_MESSAGE_BYTES: usize = 128 * 1024 - 256; +const RECORDS_EVENT_STREAM_OVERHEAD_BYTES: usize = 101; +const RECORDS_CHUNK_TARGET: usize = MAX_COMPAT_EVENT_STREAM_MESSAGE_BYTES - RECORDS_EVENT_STREAM_OVERHEAD_BYTES; +const ENCODE_TURN_TARGET_BYTES: usize = 64 * 1024; +const MAX_ENCODE_ROWS_PER_TURN: usize = 1024; +const MAX_SELECT_OUTPUT_RECORD_BYTES: usize = 1024 * 1024; +const RECORDS_FLUSH_INTERVAL: Duration = Duration::from_millis(500); +const CONTINUATION_INTERVAL: Duration = Duration::from_secs(1); +const PROGRESS_INTERVAL: Duration = Duration::from_secs(60); const DATA_SOURCE_PATH_UNSUPPORTED_CODE: &str = "DataSourcePathUnsupported"; const INVALID_QUERY_CODE: &str = "InvalidQuery"; const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure"; @@ -50,6 +73,8 @@ 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."; const UNSUPPORTED_SQL_STRUCTURE_MESSAGE: &str = "We encountered an unsupported SQL structure. Check the SQL Reference."; +const OVER_MAX_RECORD_SIZE_MESSAGE: &str = + "The length of a record in the input or result is greater than the maxCharsPerRecord limit of 1 MB."; #[derive(Clone, Debug)] struct SelectValidation { @@ -69,8 +94,14 @@ enum SelectProducerOutcome { ReceiverClosed, } +enum TerminalRecordsMode { + Complete, + PrefixBeforeError, +} + struct SelectEventChannel { tx: mpsc::Sender>, + terminal_records_permit: Option>>, terminal_permit: mpsc::OwnedPermit>, } @@ -123,7 +154,11 @@ pub async fn execute_select_object_content( .into_record_batch_stream() .map_err(map_query_error_to_s3)?; - let (tx, rx) = mpsc::channel::>(9); + let (tx, rx) = mpsc::channel::>(10); + let terminal_records_permit = tx + .clone() + .try_reserve_owned() + .map_err(|_| map_select_error_to_s3(&SelectError::InternalError))?; let terminal_permit = tx .clone() .try_reserve_owned() @@ -135,7 +170,11 @@ pub async fn execute_select_object_content( spawn_traced(async move { send_select_events_until_deadline( output, - SelectEventChannel { tx, terminal_permit }, + SelectEventChannel { + tx, + terminal_records_permit: Some(terminal_records_permit), + terminal_permit, + }, validation, input_metrics, query_deadline, @@ -150,28 +189,25 @@ pub async fn execute_select_object_content( async fn send_select_events_until_deadline( output: SendableRecordBatchStream, - event_channel: SelectEventChannel, + mut event_channel: SelectEventChannel, validation: SelectValidation, input_metrics: Arc, deadline: Instant, timeout_seconds: u64, snapshot_lease: L, ) { - let outcome = match timeout_at( + let outcome = send_select_events( + output, + &mut event_channel, + validation, + input_metrics, deadline, - send_select_events(output, &event_channel.tx, validation, input_metrics, &snapshot_lease), + timeout_seconds, + &snapshot_lease, ) - .await - { - Ok(outcome) => outcome, - Err(_) => SelectProducerOutcome::Terminal(Err(map_query_error_to_s3( - SelectError::QueryTimeout { - seconds: timeout_seconds, - } - .into(), - ))), - }; + .await; if let SelectProducerOutcome::Terminal(event) = outcome { + drop(event_channel.terminal_records_permit.take()); event_channel.terminal_permit.send(event); } drop(snapshot_lease); @@ -179,81 +215,382 @@ async fn send_select_events_until_deadline( async fn send_select_events( mut output: SendableRecordBatchStream, - tx: &mpsc::Sender>, + event_channel: &mut SelectEventChannel, validation: SelectValidation, input_metrics: Arc, + deadline: Instant, + timeout_seconds: u64, snapshot_fence: &impl SelectSnapshotFence, ) -> SelectProducerOutcome { - let mut encoder = SelectOutputEncoder::new(validation.output_format); - let mut progress = SelectProgress::new(validation.reports_input_metrics.then_some(input_metrics)); - - if tx - .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) - .await - .is_err() - { - return SelectProducerOutcome::ReceiverClosed; - } - + let SelectValidation { + output_format, + progress_enabled, + reports_input_metrics, + } = validation; + let mut encoder = SelectOutputEncoder::new(output_format); + let mut progress = SelectProgress::new(reports_input_metrics.then_some(input_metrics)); + let started_at = Instant::now(); + let records_flush = tokio::time::sleep_until(deadline); + tokio::pin!(records_flush); + let mut records_flush_armed = false; + let mut continuation = delayed_select_interval(started_at, CONTINUATION_INTERVAL); + let mut progress_interval = progress_enabled.then(|| delayed_select_interval(started_at, PROGRESS_INTERVAL)); + let deadline_sleep = tokio::time::sleep_until(deadline); + tokio::pin!(deadline_sleep); + let tx = event_channel.tx.clone(); + let mut periodic_sender = PollSender::new(tx.clone()); let receiver_closed = tx.closed(); tokio::pin!(receiver_closed); - while let Some(result) = tokio::select! { - biased; - _ = &mut receiver_closed => return SelectProducerOutcome::ReceiverClosed, - result = output.next() => result, - } { - let batch = match result { - Ok(batch) => batch, - Err(err) => { - return SelectProducerOutcome::Terminal(Err(map_query_error_to_s3(err.into()))); - } - }; - match encoder.encode_batch(&batch) { - Ok(payloads) => { - for payload in payloads { - let payload_len = payload.len(); - if tx - .send(Ok(SelectObjectContentEvent::Records(RecordsEvent { payload: Some(payload) }))) - .await + let mut records_buffer = BytesMut::new(); + let mut pending_event: Option = None; + let mut pending_batch: Option = None; + let mut pending_batch_offset = 0; + let mut continuation_due = false; + let mut progress_due = false; + + loop { + if pending_event.is_some() { + periodic_sender.abort_send(); + } + let periodic_due = progress_due || continuation_due; + let finishing_success = matches!(pending_event.as_ref(), Some(SelectObjectContentEvent::Stats(_))); + let progress_armed = progress_interval.is_some(); + + tokio::select! { + biased; + + _ = &mut receiver_closed => return SelectProducerOutcome::ReceiverClosed, + + _ = &mut deadline_sleep => { + return finish_select_with_error( + select_query_timeout_error(timeout_seconds), + event_channel, + &mut pending_event, + &mut records_buffer, + &mut progress, + ); + } + + _ = tick_optional_interval(&mut progress_interval), if progress_armed && !progress_due && !finishing_success => { + progress_due = true; + } + + _ = continuation.tick(), if !continuation_due && !finishing_success => { + continuation_due = true; + } + + permit = tx.reserve(), if pending_event.is_some() => { + let permit = match permit { + Ok(permit) => permit, + Err(_) => return SelectProducerOutcome::ReceiverClosed, + }; + let Some(event) = pending_event.take() else { + return SelectProducerOutcome::Terminal(Err(map_select_error_to_s3(&SelectError::InternalError))); + }; + let finishes_successfully = matches!(&event, SelectObjectContentEvent::Stats(_)); + if finishes_successfully + && let Err(error) = snapshot_fence.ensure_snapshot_valid() + { + return SelectProducerOutcome::Terminal(Err(error)); + } + let returned = records_payload_len(&event); + permit.send(Ok(event)); + if let Some(returned) = returned { + progress.add_returned(returned); + } + if finishes_successfully { + return SelectProducerOutcome::Terminal(Ok(SelectObjectContentEvent::End(EndEvent::default()))); + } + if !periodic_due { + schedule_buffered_records( + &mut records_buffer, + records_flush.as_mut(), + &mut records_flush_armed, + &mut pending_event, + deadline, + ); + } + } + + _ = &mut records_flush, if records_flush_armed && !finishing_success && pending_event.is_none() => { + records_flush_armed = false; + records_flush.as_mut().reset(deadline); + pending_event = take_records_payload(&mut records_buffer).map(records_event); + } + + permit = poll_fn(|cx| periodic_sender.poll_reserve(cx)), if periodic_due && pending_event.is_none() => { + if permit.is_err() { + return SelectProducerOutcome::ReceiverClosed; + } + if progress_due { + if periodic_sender + .send_item(Ok(SelectObjectContentEvent::Progress(ProgressEvent { + details: Some(progress.to_progress()), + }))) .is_err() { return SelectProducerOutcome::ReceiverClosed; } - progress.add_returned(payload_len); - if validation.progress_enabled - && tx - .send(Ok(SelectObjectContentEvent::Progress(ProgressEvent { - details: Some(progress.to_progress()), - }))) - .await - .is_err() + progress_due = false; + if let Some(interval) = progress_interval.as_mut() { + interval.reset(); + } + } else { + if periodic_sender + .send_item(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) + .is_err() { return SelectProducerOutcome::ReceiverClosed; } + continuation_due = false; + continuation.reset(); + } + if !progress_due && !continuation_due { + schedule_buffered_records( + &mut records_buffer, + records_flush.as_mut(), + &mut records_flush_armed, + &mut pending_event, + deadline, + ); } } - Err(err) => { - return SelectProducerOutcome::Terminal(Err(err)); + + _ = tokio::task::yield_now(), if pending_batch.is_some() && pending_event.is_none() && !periodic_due => { + let Some(batch) = pending_batch.as_ref() else { + return SelectProducerOutcome::Terminal(Err(map_select_error_to_s3(&SelectError::InternalError))); + }; + let remaining_rows = batch.num_rows().saturating_sub(pending_batch_offset); + if remaining_rows == 0 { + pending_batch = None; + pending_batch_offset = 0; + continue; + } + let encoded_rows = match encode_batch_turn( + &mut encoder, + batch, + pending_batch_offset, + &mut records_buffer, + ) { + Ok(encoded_rows) => encoded_rows, + Err(error) => { + return finish_select_with_error( + error, + event_channel, + &mut pending_event, + &mut records_buffer, + &mut progress, + ); + } + }; + if encoded_rows == 0 { + return SelectProducerOutcome::Terminal(Err(map_select_error_to_s3(&SelectError::InternalError))); + } + pending_batch_offset += encoded_rows; + if pending_batch_offset == batch.num_rows() { + pending_batch = None; + pending_batch_offset = 0; + } + schedule_buffered_records( + &mut records_buffer, + records_flush.as_mut(), + &mut records_flush_armed, + &mut pending_event, + deadline, + ); } + + result = output.next(), if !finishing_success && pending_event.is_none() && pending_batch.is_none() && !periodic_due => { + match result { + Some(Ok(batch)) => { + pending_batch = Some(batch); + pending_batch_offset = 0; + } + Some(Err(error)) => { + return finish_select_with_error( + map_query_error_to_s3(error.into()), + event_channel, + &mut pending_event, + &mut records_buffer, + &mut progress, + ); + } + None => { + if let Err(error) = snapshot_fence.ensure_snapshot_valid() { + return finish_select_with_error( + error, + event_channel, + &mut pending_event, + &mut records_buffer, + &mut progress, + ); + } + if let Err(error) = flush_terminal_records( + event_channel, + &mut pending_event, + &mut records_buffer, + &mut progress, + TerminalRecordsMode::Complete, + ) { + return SelectProducerOutcome::Terminal(Err(error)); + } + pending_event = Some(SelectObjectContentEvent::Stats(StatsEvent { + details: Some(progress.to_stats()), + })); + } + } + } + + } + } +} + +fn delayed_select_interval(started_at: Instant, period: Duration) -> Interval { + let mut interval = tokio::time::interval_at(started_at + period, period); + interval.set_missed_tick_behavior(MissedTickBehavior::Delay); + interval +} + +fn encode_batch_turn( + encoder: &mut SelectOutputEncoder, + batch: &RecordBatch, + offset: usize, + buffer: &mut BytesMut, +) -> S3Result { + let remaining_rows = batch.num_rows().saturating_sub(offset); + if remaining_rows == 0 { + return Ok(0); + } + + let original_len = buffer.len(); + let candidate_rows = remaining_rows.min(MAX_ENCODE_ROWS_PER_TURN); + let output_limit = if candidate_rows == 1 { + MAX_SELECT_OUTPUT_RECORD_BYTES + } else { + ENCODE_TURN_TARGET_BYTES + }; + match encoder.encode_batch_limited(batch, offset..offset + candidate_rows, buffer, output_limit) { + Ok(encoded_rows) if encoded_rows > 0 => return Ok(encoded_rows), + Ok(_) if candidate_rows > 1 => {} + Ok(_) => return Err(over_max_record_size_error()), + Err(error) => { + buffer.truncate(original_len); + return Err(error); } } - if let Err(error) = snapshot_fence.ensure_snapshot_valid() { - return SelectProducerOutcome::Terminal(Err(error)); + match encoder.encode_batch_limited(batch, offset..offset + 1, buffer, MAX_SELECT_OUTPUT_RECORD_BYTES) { + Ok(1) => Ok(1), + Ok(_) => { + buffer.truncate(original_len); + Err(over_max_record_size_error()) + } + Err(error) => { + buffer.truncate(original_len); + Err(error) + } } - let stats = SelectObjectContentEvent::Stats(StatsEvent { - details: Some(progress.to_stats()), - }); - let stats_permit = match tx.reserve().await { - Ok(permit) => permit, - Err(_) => return SelectProducerOutcome::ReceiverClosed, +} + +async fn tick_optional_interval(interval: &mut Option) { + match interval { + Some(interval) => { + interval.tick().await; + } + None => std::future::pending().await, + } +} + +fn schedule_buffered_records( + buffer: &mut BytesMut, + mut flush: Pin<&mut Sleep>, + flush_armed: &mut bool, + pending_event: &mut Option, + idle_deadline: Instant, +) { + if pending_event.is_some() || buffer.is_empty() { + return; + } + if buffer.len() >= RECORDS_CHUNK_TARGET { + *flush_armed = false; + flush.as_mut().reset(idle_deadline); + *pending_event = Some(records_event(buffer.split_to(RECORDS_CHUNK_TARGET).freeze())); + } else if !*flush_armed { + flush.as_mut().reset(Instant::now() + RECORDS_FLUSH_INTERVAL); + *flush_armed = true; + } +} + +fn take_records_payload(buffer: &mut BytesMut) -> Option { + (!buffer.is_empty()).then(|| buffer.split().freeze()) +} + +fn records_event(payload: Bytes) -> SelectObjectContentEvent { + SelectObjectContentEvent::Records(RecordsEvent { payload: Some(payload) }) +} + +fn records_payload_len(event: &SelectObjectContentEvent) -> Option { + match event { + SelectObjectContentEvent::Records(records) => records.payload.as_ref().map(Bytes::len), + _ => None, + } +} + +fn flush_terminal_records( + event_channel: &mut SelectEventChannel, + pending_event: &mut Option, + records_buffer: &mut BytesMut, + progress: &mut SelectProgress, + mode: TerminalRecordsMode, +) -> S3Result<()> { + let pending = pending_event.take(); + let pending_payload = match pending { + Some(SelectObjectContentEvent::Records(records)) => records.payload, + _ => None, }; - if let Err(error) = snapshot_fence.ensure_snapshot_valid() { + let payload = pending_payload.or_else(|| take_records_payload(records_buffer)); + if matches!(mode, TerminalRecordsMode::PrefixBeforeError) { + records_buffer.clear(); + } + let permit = event_channel.terminal_records_permit.take(); + let Some(payload) = payload else { + drop(permit); + return Ok(()); + }; + let payload = match mode { + TerminalRecordsMode::Complete if payload.len() > RECORDS_CHUNK_TARGET => { + return Err(map_select_error_to_s3(&SelectError::InternalError)); + } + TerminalRecordsMode::PrefixBeforeError if payload.len() > RECORDS_CHUNK_TARGET => payload.slice(..RECORDS_CHUNK_TARGET), + _ => payload, + }; + let Some(permit) = permit else { + return Err(map_select_error_to_s3(&SelectError::InternalError)); + }; + let returned = payload.len(); + permit.send(Ok(records_event(payload))); + progress.add_returned(returned); + Ok(()) +} + +fn finish_select_with_error( + error: S3Error, + event_channel: &mut SelectEventChannel, + pending_event: &mut Option, + records_buffer: &mut BytesMut, + progress: &mut SelectProgress, +) -> SelectProducerOutcome { + if let Err(error) = flush_terminal_records( + event_channel, + pending_event, + records_buffer, + progress, + TerminalRecordsMode::PrefixBeforeError, + ) { return SelectProducerOutcome::Terminal(Err(error)); } - stats_permit.send(Ok(stats)); - SelectProducerOutcome::Terminal(Ok(SelectObjectContentEvent::End(EndEvent::default()))) + SelectProducerOutcome::Terminal(Err(error)) } fn validate_select_request(headers: &http::HeaderMap, input: &mut SelectObjectContentInput) -> S3Result { @@ -555,91 +892,624 @@ impl SelectOutputEncoder { Self { format } } - fn encode_batch(&mut self, batch: &RecordBatch) -> S3Result> { - let bytes = match &self.format { - SelectOutputFormat::Csv(config) => encode_csv_batch(batch, config)?, - SelectOutputFormat::Json(config) => encode_json_batch(batch, config)?, - }; - Ok(split_records_payload(bytes)) - } -} - -fn encode_csv_batch(batch: &RecordBatch, config: &CSVOutput) -> S3Result> { - let mut buffer = Vec::new(); - let mut builder = CsvWriterBuilder::new().with_header(false); - if let Some(delimiter) = config.field_delimiter.as_deref() { - builder = builder.with_delimiter(delimiter.as_bytes()[0]); - } - if let Some(quote) = config.quote_character.as_deref() { - builder = builder.with_quote(quote.as_bytes()[0]); - } - if let Some(escape) = config.quote_escape_character.as_deref() { - builder = builder.with_escape(escape.as_bytes()[0]); - } - if let Some(record_delimiter) = config.record_delimiter.as_deref() { - builder = builder.with_line_terminator(csv_terminator(record_delimiter)); - } - if let Some(quote_fields) = config.quote_fields.as_ref() - && quote_fields.as_str() == QuoteFields::ALWAYS - { - builder = builder.with_quote_style(QuoteStyle::Always); - } - - let mut writer = builder.build(&mut buffer); - writer.write(batch).map_err(internal_select_error)?; - drop(writer); - Ok(buffer) -} - -fn csv_terminator(value: &str) -> Terminator { - if value == "\r\n" { - Terminator::CRLF - } else { - Terminator::Any(value.as_bytes()[0]) - } -} - -fn encode_json_batch(batch: &RecordBatch, config: &JSONOutput) -> S3Result> { - let mut buffer = Vec::new(); - let mut writer = JsonWriterBuilder::new() - .with_explicit_nulls(true) - .build::<_, LineDelimited>(&mut buffer); - writer.write(batch).map_err(internal_select_error)?; - writer.finish().map_err(internal_select_error)?; - drop(writer); - - if let Some(delimiter) = config.record_delimiter.as_deref() - && delimiter != "\n" - { - return Ok(replace_json_record_delimiter(&buffer, delimiter.as_bytes())); - } - Ok(buffer) -} - -fn replace_json_record_delimiter(buffer: &[u8], delimiter: &[u8]) -> Vec { - let mut output = Vec::with_capacity(buffer.len()); - for byte in buffer { - if *byte == b'\n' { - output.extend_from_slice(delimiter); - } else { - output.push(*byte); + fn encode_batch_limited( + &mut self, + batch: &RecordBatch, + rows: Range, + buffer: &mut BytesMut, + max_bytes: usize, + ) -> S3Result { + match &self.format { + SelectOutputFormat::Csv(config) => encode_csv_batch_limited(batch, rows, config, buffer, max_bytes), + SelectOutputFormat::Json(config) => encode_json_batch_limited(batch, rows, config, buffer, max_bytes), } } - output } -fn split_records_payload(bytes: Vec) -> Vec { - if bytes.is_empty() { - return Vec::new(); +#[cfg(test)] +fn encode_csv_batch(batch: &RecordBatch, config: &CSVOutput, buffer: &mut BytesMut) -> S3Result<()> { + if encode_csv_batch_limited(batch, 0..batch.num_rows(), config, buffer, usize::MAX)? == batch.num_rows() { + Ok(()) + } else { + Err(internal_select_error(io::Error::other("S3 Select output length overflow"))) } - let bytes = Bytes::from(bytes); - if bytes.len() <= RECORDS_CHUNK_TARGET { - return vec![bytes]; +} + +fn encode_csv_batch_limited( + batch: &RecordBatch, + rows: Range, + config: &CSVOutput, + buffer: &mut BytesMut, + max_bytes: usize, +) -> S3Result { + let options = FormatOptions::default(); + let mut formatters = Vec::new(); + let field_delimiter = config.field_delimiter.as_deref().unwrap_or(",").as_bytes()[0]; + let quote = config.quote_character.as_deref().unwrap_or("\"").as_bytes()[0]; + let quote_escape = config.quote_escape_character.as_deref().unwrap_or("\"").as_bytes()[0]; + let record_delimiter = config.record_delimiter.as_deref().unwrap_or("\n").as_bytes(); + let quote_all = config + .quote_fields + .as_ref() + .is_some_and(|quote_fields| quote_fields.as_str() == QuoteFields::ALWAYS); + let mut output = LimitedBytesWriter::new(buffer, max_bytes); + let mut field = BoundedText::default(); + let mut encoded_rows = 0; + for row in rows { + let row_start = output.checkpoint(); + for column in 0..batch.num_columns() { + if column > 0 && output.write_all(&[field_delimiter]).is_err() { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + if formatters.len() == column { + let array = batch.column(column); + if array.data_type().is_nested() { + return Err(internal_select_error(datafusion::arrow::error::ArrowError::CsvError(format!( + "Nested type {} is not supported in CSV", + array.data_type() + )))); + } + formatters.push(ArrayFormatter::try_new(array.as_ref(), &options).map_err(internal_select_error)?); + } + let formatter = formatters + .get(column) + .ok_or_else(|| internal_select_error(io::Error::other("missing S3 Select CSV formatter")))?; + field.reset(output.remaining()); + let result = formatter.value(row).write(&mut field); + if field.limit_exceeded { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + result.map_err(internal_select_error)?; + if write_csv_field(&mut output, field.value.as_bytes(), field_delimiter, quote, quote_escape, quote_all).is_err() { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + } + if output.write_all(record_delimiter).is_err() { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + encoded_rows += 1; + } + Ok(encoded_rows) +} + +fn write_csv_field( + output: &mut LimitedBytesWriter<'_>, + field: &[u8], + delimiter: u8, + quote: u8, + quote_escape: u8, + quote_all: bool, +) -> io::Result<()> { + let quote_field = quote_all || csv_field_needs_quotes(field, delimiter, quote); + if !quote_field { + return output.write_all(field); + } + + output.write_all(&[quote])?; + let mut start = 0; + while let Some(relative) = field[start..].iter().position(|byte| *byte == quote) { + let position = start + relative; + output.write_all(&field[start..position])?; + output.write_all(&[quote_escape, quote])?; + start = position + 1; + } + output.write_all(&field[start..])?; + output.write_all(&[quote]) +} + +fn csv_field_needs_quotes(field: &[u8], delimiter: u8, quote: u8) -> bool { + if field.is_empty() { + return false; + } + if field == b"\\." + || field + .iter() + .copied() + .any(|byte| matches!(byte, b'\r' | b'\n') || byte == delimiter || byte == quote) + { + return true; + } + std::str::from_utf8(field) + .ok() + .and_then(|value| value.chars().next()) + .is_some_and(char::is_whitespace) +} + +#[derive(Default)] +struct BoundedText { + value: String, + max_bytes: usize, + limit_exceeded: bool, +} + +impl BoundedText { + fn reset(&mut self, max_bytes: usize) { + self.value.clear(); + self.max_bytes = max_bytes; + self.limit_exceeded = false; + } +} + +impl fmt::Write for BoundedText { + fn write_str(&mut self, value: &str) -> fmt::Result { + let Some(len) = self.value.len().checked_add(value.len()) else { + self.limit_exceeded = true; + return Err(fmt::Error); + }; + if len > self.max_bytes { + self.limit_exceeded = true; + return Err(fmt::Error); + } + self.value.push_str(value); + Ok(()) + } +} + +struct LimitedJsonValueEncoder<'a> { + array: &'a dyn Array, + options: &'a EncoderOptions, + kind: LimitedJsonValueKind<'a>, +} + +enum LimitedJsonValueKind<'a> { + Scalar(NullableEncoder<'a>), + List { + array: &'a dyn ListLikeArray, + value_field: &'a FieldRef, + value_array: &'a dyn Array, + values: Option>>, + }, + Struct { + fields: &'a [FieldRef], + arrays: &'a [datafusion::arrow::array::ArrayRef], + values: Vec>, + }, + Dictionary { + value_index: Box usize + 'a>, + value_field: &'a FieldRef, + value_array: &'a dyn Array, + values: Option>>, + }, + Indexed { + value_index: Box usize + 'a>, + value_field: &'a FieldRef, + value_array: &'a dyn Array, + values: Option>>, + }, + Map { + array: &'a MapArray, + field: &'a FieldRef, + keys: Option>>, + values: Option>>, + }, +} + +impl<'a> LimitedJsonValueEncoder<'a> { + fn try_new(field: &'a FieldRef, array: &'a dyn Array, options: &'a EncoderOptions) -> Result { + macro_rules! dictionary { + ($key:ty) => {{ + let dictionary = array.as_dictionary::<$key>(); + LimitedJsonValueKind::Dictionary { + value_index: Box::new(move |row| dictionary.keys().value(row).as_usize()), + value_field: field, + value_array: dictionary.values().as_ref(), + values: None, + } + }}; + } + macro_rules! run_end_encoded { + ($run_end:ty) => {{ + let run = array.as_run::<$run_end>(); + LimitedJsonValueKind::Indexed { + value_index: Box::new(move |row| run.get_physical_index(row)), + value_field: field, + value_array: run.values().as_ref(), + values: None, + } + }}; + } + + let kind = match array.data_type() { + DataType::List(value_field) => { + let list = array.as_list::(); + LimitedJsonValueKind::List { + array: list, + value_field, + value_array: list.values().as_ref(), + values: None, + } + } + DataType::LargeList(value_field) => { + let list = array.as_list::(); + LimitedJsonValueKind::List { + array: list, + value_field, + value_array: list.values().as_ref(), + values: None, + } + } + DataType::ListView(value_field) => { + let list = array.as_list_view::(); + LimitedJsonValueKind::List { + array: list, + value_field, + value_array: list.values().as_ref(), + values: None, + } + } + DataType::LargeListView(value_field) => { + let list = array.as_list_view::(); + LimitedJsonValueKind::List { + array: list, + value_field, + value_array: list.values().as_ref(), + values: None, + } + } + DataType::FixedSizeList(value_field, _) => { + let list = array.as_fixed_size_list(); + LimitedJsonValueKind::List { + array: list, + value_field, + value_array: list.values().as_ref(), + values: None, + } + } + DataType::Struct(fields) => LimitedJsonValueKind::Struct { + fields: fields.as_ref(), + arrays: array.as_struct().columns(), + values: Vec::new(), + }, + DataType::Dictionary(key_type, _) => match key_type.as_ref() { + DataType::Int8 => dictionary!(Int8Type), + DataType::Int16 => dictionary!(Int16Type), + DataType::Int32 => dictionary!(Int32Type), + DataType::Int64 => dictionary!(Int64Type), + DataType::UInt8 => dictionary!(UInt8Type), + DataType::UInt16 => dictionary!(UInt16Type), + DataType::UInt32 => dictionary!(UInt32Type), + DataType::UInt64 => dictionary!(UInt64Type), + key_type => { + return Err(ArrowError::JsonError(format!( + "Unsupported dictionary key type for JSON encoding: {key_type:?}" + ))); + } + }, + DataType::RunEndEncoded(run_ends, _) => match run_ends.data_type() { + DataType::Int16 => run_end_encoded!(Int16Type), + DataType::Int32 => run_end_encoded!(Int32Type), + DataType::Int64 => run_end_encoded!(Int64Type), + run_end_type => { + return Err(ArrowError::JsonError(format!( + "Unsupported run-end type for JSON encoding: {run_end_type:?}" + ))); + } + }, + DataType::Map(_, _) => { + let map = array.as_map(); + if !matches!(map.keys().data_type(), DataType::Utf8 | DataType::LargeUtf8 | DataType::Utf8View) { + return Err(ArrowError::JsonError(format!( + "Only UTF8 keys supported by JSON MapArray Writer: got {:?}", + map.keys().data_type() + ))); + } + if map.keys().null_count() != 0 { + return Err(ArrowError::InvalidArgumentError("Encountered nulls in MapArray keys".to_string())); + } + if map.entries().nulls().is_some_and(|nulls| nulls.null_count() != 0) { + return Err(ArrowError::InvalidArgumentError("Encountered nulls in MapArray entries".to_string())); + } + LimitedJsonValueKind::Map { + array: map, + field, + keys: None, + values: None, + } + } + _ => LimitedJsonValueKind::Scalar(make_encoder(field, array, options)?), + }; + Ok(Self { array, options, kind }) + } + + fn encode(&mut self, row: usize, output: &mut LimitedBytesWriter<'_>, scratch: &mut Vec) -> Result { + if self.array.is_null(row) { + return Ok(write_limited(output, b"null")); + } + self.encode_non_null(row, output, scratch) + } + + fn encode_non_null( + &mut self, + row: usize, + output: &mut LimitedBytesWriter<'_>, + scratch: &mut Vec, + ) -> Result { + let array = self.array; + // NullArray has no physical null buffer, including behind a dictionary index. + if matches!(array.data_type(), DataType::Null) { + return Ok(write_limited(output, b"null")); + } + let options = self.options; + match &mut self.kind { + LimitedJsonValueKind::Scalar(encoder) => Ok(write_json_scalar_limited(array, encoder, row, output, scratch)), + LimitedJsonValueKind::List { + array, + value_field, + value_array, + values, + } => { + if !write_limited(output, b"[") { + return Ok(false); + } + for (index, value_row) in array.element_range(row).enumerate() { + if index > 0 && !write_limited(output, b",") { + return Ok(false); + } + let values = Self::lazy_value_encoder(values, value_field, *value_array, options)?; + if !values.encode(value_row, output, scratch)? { + return Ok(false); + } + } + Ok(write_limited(output, b"]")) + } + LimitedJsonValueKind::Struct { fields, arrays, values } => { + if !write_limited(output, b"{") { + return Ok(false); + } + for (index, field) in fields.iter().enumerate() { + if (index > 0 && !write_limited(output, b",")) + || !write_json_string_limited(output, field.name()) + || !write_limited(output, b":") + { + return Ok(false); + } + if values.len() == index { + values.push(Self::try_new(field, arrays[index].as_ref(), options)?); + } + let value = values + .get_mut(index) + .ok_or_else(|| ArrowError::JsonError("S3 Select JSON encoder state is inconsistent".to_string()))?; + if !value.encode(row, output, scratch)? { + return Ok(false); + } + } + Ok(write_limited(output, b"}")) + } + // Arrow checks dictionary key nulls but delegates value nulls to the value encoder. + LimitedJsonValueKind::Dictionary { + value_index, + value_field, + value_array, + values, + } => Self::lazy_value_encoder(values, value_field, *value_array, options)?.encode_non_null( + value_index(row), + output, + scratch, + ), + LimitedJsonValueKind::Indexed { + value_index, + value_field, + value_array, + values, + } => Self::lazy_value_encoder(values, value_field, *value_array, options)?.encode(value_index(row), output, scratch), + LimitedJsonValueKind::Map { + array, + field, + keys, + values, + } => { + if !write_limited(output, b"{") { + return Ok(false); + } + let offsets = array.value_offsets(); + let start = offsets[row].as_usize(); + let end = offsets[row + 1].as_usize(); + for (index, value_row) in (start..end).enumerate() { + if index > 0 && !write_limited(output, b",") { + return Ok(false); + } + let keys = Self::lazy_value_encoder(keys, field, array.keys(), options)?; + if !keys.encode(value_row, output, scratch)? || !write_limited(output, b":") { + return Ok(false); + } + let values = Self::lazy_value_encoder(values, field, array.values(), options)?; + if !values.encode(value_row, output, scratch)? { + return Ok(false); + } + } + Ok(write_limited(output, b"}")) + } + } + } + + fn lazy_value_encoder<'b>( + slot: &'b mut Option>>, + field: &'a FieldRef, + array: &'a dyn Array, + options: &'a EncoderOptions, + ) -> Result<&'b mut LimitedJsonValueEncoder<'a>, ArrowError> { + if slot.is_none() { + *slot = Some(Box::new(Self::try_new(field, array, options)?)); + } + slot.as_deref_mut() + .ok_or_else(|| ArrowError::JsonError("S3 Select JSON encoder state is inconsistent".to_string())) + } +} + +fn write_json_scalar_limited( + array: &dyn Array, + encoder: &mut NullableEncoder<'_>, + row: usize, + output: &mut LimitedBytesWriter<'_>, + scratch: &mut Vec, +) -> bool { + match array.data_type() { + DataType::Utf8 => write_json_string_limited(output, array.as_string::().value(row)), + DataType::LargeUtf8 => write_json_string_limited(output, array.as_string::().value(row)), + DataType::Utf8View => write_json_string_limited(output, array.as_string_view().value(row)), + DataType::Binary => write_json_binary_limited(output, array.as_binary::().value(row), scratch), + DataType::LargeBinary => write_json_binary_limited(output, array.as_binary::().value(row), scratch), + DataType::BinaryView => write_json_binary_limited(output, array.as_binary_view().value(row), scratch), + DataType::FixedSizeBinary(_) => write_json_binary_limited(output, array.as_fixed_size_binary().value(row), scratch), + _ => { + scratch.clear(); + encoder.encode(row, scratch); + write_limited(output, scratch) + } + } +} + +fn write_json_string_limited(output: &mut LimitedBytesWriter<'_>, value: &str) -> bool { + if value.len().checked_add(2).is_none_or(|minimum| minimum > output.remaining()) { + return false; + } + serde_json::to_writer(output, value).is_ok() +} + +fn write_json_binary_limited(output: &mut LimitedBytesWriter<'_>, value: &[u8], scratch: &mut Vec) -> bool { + let Some(encoded_len) = value.len().checked_mul(2).and_then(|bytes| bytes.checked_add(2)) else { + return false; + }; + if encoded_len > output.remaining() || !write_limited(output, b"\"") { + return false; + } + const HEX: &[u8; 16] = b"0123456789abcdef"; + const INPUT_CHUNK_BYTES: usize = 2048; + for chunk in value.chunks(INPUT_CHUNK_BYTES) { + scratch.clear(); + scratch.reserve(chunk.len() * 2); + for byte in chunk { + scratch.push(HEX[(byte >> 4) as usize]); + scratch.push(HEX[(byte & 0x0f) as usize]); + } + if !write_limited(output, scratch) { + return false; + } + } + write_limited(output, b"\"") +} + +fn write_limited(output: &mut LimitedBytesWriter<'_>, bytes: &[u8]) -> bool { + output.write_all(bytes).is_ok() +} + +fn json_record_delimiter(config: &JSONOutput) -> &[u8] { + if let Some(delimiter) = config.record_delimiter.as_deref() { + delimiter.as_bytes() + } else { + b"\n" + } +} + +#[cfg(test)] +fn encode_json_batch(batch: &RecordBatch, config: &JSONOutput, buffer: &mut BytesMut) -> S3Result<()> { + if encode_json_batch_limited(batch, 0..batch.num_rows(), config, buffer, usize::MAX)? == batch.num_rows() { + Ok(()) + } else { + Err(internal_select_error(io::Error::other("S3 Select output length overflow"))) + } +} + +fn encode_json_batch_limited( + batch: &RecordBatch, + rows: Range, + config: &JSONOutput, + buffer: &mut BytesMut, + max_bytes: usize, +) -> S3Result { + let options = EncoderOptions::default().with_explicit_nulls(true); + let schema = batch.schema(); + let fields = schema.fields(); + let arrays = batch.columns(); + let mut values = Vec::new(); + let delimiter = json_record_delimiter(config); + let mut output = LimitedBytesWriter::new(buffer, max_bytes); + let mut scratch = Vec::new(); + let mut encoded_rows = 0; + for row in rows { + let row_start = output.checkpoint(); + if !write_limited(&mut output, b"{") { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + for (index, field) in fields.iter().enumerate() { + if (index > 0 && !write_limited(&mut output, b",")) + || !write_json_string_limited(&mut output, field.name()) + || !write_limited(&mut output, b":") + { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + if values.len() == index { + values.push( + LimitedJsonValueEncoder::try_new(field, arrays[index].as_ref(), &options).map_err(internal_select_error)?, + ); + } + let value = values + .get_mut(index) + .ok_or_else(|| internal_select_error(io::Error::other("S3 Select JSON encoder state is inconsistent")))?; + if !value.encode(row, &mut output, &mut scratch).map_err(internal_select_error)? { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + } + if !write_limited(&mut output, b"}") || !write_limited(&mut output, delimiter) { + output.rollback_to(row_start); + return Ok(encoded_rows); + } + encoded_rows += 1; + } + Ok(encoded_rows) +} + +struct LimitedBytesWriter<'a> { + buffer: &'a mut BytesMut, + max_bytes: usize, + written: usize, +} + +impl<'a> LimitedBytesWriter<'a> { + fn new(buffer: &'a mut BytesMut, max_bytes: usize) -> Self { + Self { + buffer, + max_bytes, + written: 0, + } + } + + fn remaining(&self) -> usize { + self.max_bytes.saturating_sub(self.written) + } + + fn checkpoint(&self) -> usize { + self.buffer.len() + } + + fn rollback_to(&mut self, checkpoint: usize) { + self.buffer.truncate(checkpoint); + } +} + +impl Write for LimitedBytesWriter<'_> { + fn write(&mut self, bytes: &[u8]) -> io::Result { + let Some(written) = self.written.checked_add(bytes.len()) else { + return Err(io::Error::other("S3 Select output length overflow")); + }; + if written > self.max_bytes { + return Err(io::Error::other("S3 Select encode turn limit exceeded")); + } + self.buffer.extend_from_slice(bytes); + self.written = written; + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) } - (0..bytes.len()) - .step_by(RECORDS_CHUNK_TARGET) - .map(|start| bytes.slice(start..(start + RECORDS_CHUNK_TARGET).min(bytes.len()))) - .collect() } struct SelectProgress { @@ -735,6 +1605,10 @@ fn internal_select_error(_error: impl std::error::Error + Send + Sync + 'static) map_select_error_to_s3(&SelectError::InternalError) } +fn over_max_record_size_error() -> S3Error { + S3Error::with_message(S3ErrorCode::OverMaxRecordSize, OVER_MAX_RECORD_SIZE_MESSAGE) +} + fn custom_bad_request(code: &'static str, message: String) -> S3Error { let mut err = S3Error::with_message(S3ErrorCode::Custom(code.into()), message); err.set_status_code(StatusCode::BAD_REQUEST); @@ -756,9 +1630,15 @@ mod tests { use super::*; use datafusion::{ arrow::{ - array::{Array, ListArray, StringArray}, + array::{ + Array, ArrayRef, BinaryArray, BinaryViewArray, DictionaryArray, FixedSizeBinaryArray, Int32Array, + LargeBinaryArray, LargeListArray, LargeListViewArray, LargeStringArray, ListArray, ListViewArray, MapArray, + NullArray, RunArray, StringArray, StringDictionaryBuilder, StringViewArray, StructArray, + builder::{BooleanBuilder, FixedSizeListBuilder, Int32Builder, ListBuilder}, + }, datatypes::{DataType, Field, Int32Type, Schema}, error::ArrowError, + json::writer::{LineDelimited, WriterBuilder}, }, physical_plan::stream::RecordBatchStreamAdapter, sql::sqlparser::parser::ParserError, @@ -911,6 +1791,25 @@ mod tests { } } + fn pending_output() -> SendableRecordBatchStream { + Box::pin(RecordBatchStreamAdapter::new( + Arc::new(Schema::empty()), + futures::stream::pending::>(), + )) + } + + fn large_pending_output(chunks: usize) -> SendableRecordBatchStream { + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let value = "x".repeat(RECORDS_CHUNK_TARGET * chunks); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(StringArray::from(vec![value]))]) + .expect("test record batch should be valid"); + Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }) + .chain(futures::stream::pending::>()), + )) + } + fn spawn_test_producer( output: SendableRecordBatchStream, channel_capacity: usize, @@ -919,24 +1818,53 @@ mod tests { mpsc::Receiver>, tokio::sync::oneshot::Receiver<()>, ) { - let (tx, rx) = mpsc::channel(channel_capacity); - let terminal_permit = tx - .clone() - .try_reserve_owned() - .expect("test channel should reserve terminal capacity"); + spawn_test_producer_with(output, channel_capacity, csv_validation(), Duration::from_secs(300)) + } + + fn spawn_test_producer_with( + output: SendableRecordBatchStream, + channel_capacity: usize, + validation: SelectValidation, + deadline_after: Duration, + ) -> ( + tokio::task::JoinHandle<()>, + mpsc::Receiver>, + tokio::sync::oneshot::Receiver<()>, + ) { + let (event_channel, rx) = test_event_channel(channel_capacity); let (lease, lease_released) = lease_drop_signal(); let producer = tokio::spawn(send_select_events_until_deadline( output, - SelectEventChannel { tx, terminal_permit }, - csv_validation(), + event_channel, + validation, Arc::new(SelectInputMetrics::default()), - Instant::now() + std::time::Duration::from_secs(1), + Instant::now() + deadline_after, 300, lease, )); (producer, rx, lease_released) } + fn test_event_channel(channel_capacity: usize) -> (SelectEventChannel, mpsc::Receiver>) { + let (tx, rx) = mpsc::channel(channel_capacity); + let terminal_records_permit = tx + .clone() + .try_reserve_owned() + .expect("test channel should reserve terminal Records capacity"); + let terminal_permit = tx + .clone() + .try_reserve_owned() + .expect("test channel should reserve terminal capacity"); + ( + SelectEventChannel { + tx, + terminal_records_permit: Some(terminal_records_permit), + terminal_permit, + }, + rx, + ) + } + #[test] fn validate_rejects_http_range() { let mut input = base_input(); @@ -1150,22 +2078,17 @@ mod tests { #[tokio::test(start_paused = true)] async fn producer_deadline_cancels_backpressured_send() { - let output = Box::pin(RecordBatchStreamAdapter::new( - Arc::new(Schema::empty()), - futures::stream::pending::>(), - )); - let (tx, mut rx) = mpsc::channel(2); - let terminal_permit = tx - .clone() - .try_reserve_owned() - .expect("test channel should reserve terminal capacity"); - tx.send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) + let output = pending_output(); + let (event_channel, mut rx) = test_event_channel(3); + event_channel + .tx + .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) .await .expect("test channel should accept the prefilled event"); let (lease, lease_released) = lease_drop_signal(); let producer = tokio::spawn(send_select_events_until_deadline( output, - SelectEventChannel { tx, terminal_permit }, + event_channel, csv_validation(), Arc::new(SelectInputMetrics::default()), Instant::now() + std::time::Duration::from_secs(1), @@ -1189,6 +2112,784 @@ mod tests { assert!(lease_released.await.is_ok(), "timeout should release the snapshot lease"); } + #[tokio::test(start_paused = true)] + async fn deadline_preempts_multi_slice_batch_encoding() { + let value = "x".repeat(64 * 1024); + let mut builder = StringDictionaryBuilder::::new(); + for _ in 0..(MAX_ENCODE_ROWS_PER_TURN + 1) { + builder.append(&value).expect("dictionary value should append"); + } + let values = builder.finish(); + let schema = Arc::new(Schema::new(vec![Field::new("value", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(values)]).expect("test record batch should be valid"); + let output = Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), + )); + let (event_channel, mut rx) = test_event_channel(4); + let (lease, lease_released) = lease_drop_signal(); + let producer = send_select_events_until_deadline( + output, + event_channel, + csv_validation(), + Arc::new(SelectInputMetrics::default()), + Instant::now() + Duration::from_secs(1), + 1, + lease, + ); + tokio::pin!(producer); + + assert!(futures::poll!(producer.as_mut()).is_pending()); + assert!(futures::poll!(producer.as_mut()).is_pending()); + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + tokio::time::advance(Duration::from_secs(1)).await; + assert!(futures::poll!(producer.as_mut()).is_ready()); + + let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { + panic!("the first encoded slice should flush before the timeout"); + }; + assert_eq!(records.payload.as_ref().map(Bytes::len), Some(value.len() + 1)); + let timeout = rx + .recv() + .await + .expect("deadline should send one terminal error") + .expect_err("deadline terminal event should be an error"); + assert_eq!(timeout.code(), &S3ErrorCode::Busy); + assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "deadline should release the snapshot lease"); + } + + #[test] + fn skewed_dictionary_batch_stays_within_one_encode_turn() { + let value = "x".repeat(64 * 1024); + let mut builder = StringDictionaryBuilder::::new(); + builder.append("").expect("empty dictionary value should append"); + for _ in 0..MAX_ENCODE_ROWS_PER_TURN { + builder.append(&value).expect("large dictionary value should append"); + } + let values = builder.finish(); + let schema = Arc::new(Schema::new(vec![Field::new("value", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("dictionary rows should encode successfully"); + + assert_eq!(encoded_rows, 1, "shared dictionary values must be charged before batching rows"); + assert_eq!(buffer.as_ref(), b"\n"); + } + + #[test] + fn unused_large_dictionary_value_does_not_force_per_row_encoding() { + let row_count = MAX_ENCODE_ROWS_PER_TURN; + let keys = Int32Array::from(vec![0; row_count]); + let unused = "x".repeat(64 * 1024); + let dictionary = Arc::new(StringArray::from(vec!["", unused.as_str()])); + let values = DictionaryArray::::try_new(keys, dictionary).expect("test dictionary should be valid"); + let schema = Arc::new(Schema::new(vec![Field::new("value", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("referenced empty values should batch"); + + assert_eq!(encoded_rows, row_count); + assert_eq!(buffer.as_ref(), "\n".repeat(row_count).as_bytes()); + } + + #[test] + fn skewed_json_dictionary_keeps_the_complete_bounded_prefix() { + let row_count = MAX_ENCODE_ROWS_PER_TURN; + let mut keys = vec![0; row_count]; + keys[row_count - 1] = 1; + let large = "x".repeat(70 * 1024); + let dictionary = Arc::new(StringArray::from(vec!["a", large.as_str()])); + let values = + DictionaryArray::::try_new(Int32Array::from(keys), dictionary).expect("test dictionary should be valid"); + let schema = Arc::new(Schema::new(vec![Field::new("value", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("the bounded dictionary prefix should encode"); + + assert_eq!(encoded_rows, row_count - 1); + assert_eq!(buffer.iter().filter(|byte| **byte == b'\n').count(), encoded_rows); + } + + #[test] + fn short_rows_share_an_encode_turn() { + let row_count = MAX_ENCODE_ROWS_PER_TURN; + let values = StringArray::from(vec!["x"; row_count]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("short rows should encode successfully"); + + assert!(encoded_rows > 1, "small rows should not construct one Arrow writer per row"); + assert!(encoded_rows <= MAX_ENCODE_ROWS_PER_TURN); + assert_eq!(buffer.as_ref(), "x\n".repeat(encoded_rows).as_bytes()); + } + + #[test] + fn small_nested_json_rows_share_an_encode_turn() { + let values = ListArray::from_iter_primitive::([Some([Some(1)]), Some([Some(2)])]); + let schema = Arc::new(Schema::new(vec![Field::new("items", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("small nested JSON rows should encode"); + + assert_eq!(encoded_rows, 2); + assert_eq!(buffer.as_ref(), b"{\"items\":[1]}\n{\"items\":[2]}\n"); + } + + #[test] + fn large_second_nested_json_row_stops_after_the_complete_prefix() { + let mut builder = ListBuilder::new(BooleanBuilder::new()); + builder.values().append_value(true); + builder.append(true); + for _ in 0..(ENCODE_TURN_TARGET_BYTES / b"true,".len() + 8) { + builder.values().append_value(true); + } + builder.append(true); + let values = builder.finish(); + let schema = Arc::new(Schema::new(vec![Field::new("items", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("the complete first row should encode"); + + assert_eq!(encoded_rows, 1); + assert_eq!(buffer.as_ref(), b"{\"items\":[true]}\n"); + + buffer.clear(); + let encoded_rows = + encode_batch_turn(&mut encoder, &batch, 1, &mut buffer).expect("the larger second row should encode alone"); + assert_eq!(encoded_rows, 1); + assert!(buffer.len() > ENCODE_TURN_TARGET_BYTES); + assert!(buffer.len() <= MAX_SELECT_OUTPUT_RECORD_BYTES); + } + + #[test] + fn csv_turn_rolls_back_a_partially_escaped_second_row() { + let quoted = "\"".repeat(40 * 1024); + let values = StringArray::from(vec!["ok", quoted.as_str()]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::new(); + + let encoded_rows = + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("the complete first row should remain staged"); + + assert_eq!(encoded_rows, 1); + assert_eq!(buffer.as_ref(), b"ok\n"); + } + + #[test] + fn csv_encoder_error_discards_the_partial_current_turn() { + let values = StringArray::from(vec!["ok"]); + let nested = ListArray::from_iter_primitive::([Some([Some(1)])]); + let schema = Arc::new(Schema::new(vec![ + Field::new("value", DataType::Utf8, false), + Field::new("nested", nested.data_type().clone(), false), + ])); + let batch = + RecordBatch::try_new(schema, vec![Arc::new(values), Arc::new(nested)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::from(b"staged\n".as_slice()); + + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("nested CSV output should fail after the first column"); + + assert_eq!(error.code(), &S3ErrorCode::InternalError); + assert_eq!(buffer.as_ref(), b"staged\n"); + } + + #[test] + fn oversized_nested_json_stops_at_the_output_budget() { + let value_count = MAX_SELECT_OUTPUT_RECORD_BYTES / b"true,".len() + 1; + let mut builder = ListBuilder::new(BooleanBuilder::new()); + for _ in 0..value_count { + builder.values().append_value(true); + } + builder.append(true); + let values = builder.finish(); + let schema = Arc::new(Schema::new(vec![Field::new("items", values.data_type().clone(), false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("nested JSON larger than one MiB must stop at the output budget"); + + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + + #[test] + fn nested_json_struct_map_dictionary_and_run_end_match_arrow_semantics() { + let profile = StructArray::from(vec![ + ( + Arc::new(Field::new("name", DataType::Utf8, false)), + Arc::new(StringArray::from(vec!["a\n"])) as ArrayRef, + ), + ( + Arc::new(Field::new("count", DataType::Int32, true)), + Arc::new(Int32Array::from(vec![None])) as ArrayRef, + ), + ]); + let labels = + MapArray::from_vec_of_maps::(vec![Some(vec![("a", Some(1)), ("b", None)])], true); + let mut dictionary = StringDictionaryBuilder::::new(); + dictionary.append("small").expect("dictionary value should append"); + let dictionary = dictionary.finish(); + let run_ends = Int32Array::from(vec![1]); + let run_values = Arc::new(StringArray::from(vec!["run"])) as ArrayRef; + let run = RunArray::::try_new(&run_ends, &run_values).expect("run-end encoded value should be valid"); + let schema = Arc::new(Schema::new(vec![ + Field::new("profile", profile.data_type().clone(), false), + Field::new("labels", labels.data_type().clone(), false), + Field::new("code", dictionary.data_type().clone(), false), + Field::new("run", run.data_type().clone(), false), + ])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(profile), Arc::new(labels), Arc::new(dictionary), Arc::new(run)]) + .expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("nested JSON values should encode"), + 1 + ); + assert_eq!( + buffer.as_ref(), + b"{\"profile\":{\"name\":\"a\\n\",\"count\":null},\"labels\":{\"a\":1,\"b\":null},\"code\":\"small\",\"run\":\"run\"}\n" + ); + } + + #[test] + fn json_null_projection_encodes_without_invoking_arrows_null_encoder() { + let values = NullArray::new(1); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Null, true)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("NULL projection should encode"), + 1 + ); + assert_eq!(buffer.as_ref(), b"{\"value\":null}\n"); + } + + #[test] + fn json_dictionary_null_value_matches_arrow_writer_semantics() { + let keys = Int32Array::from(vec![0]); + let dictionary = Arc::new(StringArray::from(vec![None::<&str>])); + let values = DictionaryArray::::try_new(keys, dictionary).expect("test dictionary should be valid"); + let schema = Arc::new(Schema::new(vec![Field::new("value", values.data_type().clone(), true)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("dictionary null value should encode"), + 1 + ); + assert_eq!(buffer.as_ref(), b"{\"value\":\"\"}\n"); + } + + #[test] + fn json_dictionary_of_null_type_does_not_invoke_arrows_null_encoder() { + let values = DictionaryArray::::try_new(Int32Array::from(vec![0]), Arc::new(NullArray::new(1))) + .expect("test dictionary should be valid"); + let schema = Arc::new(Schema::new(vec![Field::new("value", values.data_type().clone(), true)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("dictionary Null value should encode"), + 1 + ); + assert_eq!(buffer.as_ref(), b"{\"value\":null}\n"); + } + + #[test] + fn json_escaped_field_name_stops_at_the_output_budget() { + let field_name = "\n".repeat(MAX_SELECT_OUTPUT_RECORD_BYTES / 2 + 1); + let values = NullArray::new(1); + let schema = Arc::new(Schema::new(vec![Field::new(field_name, DataType::Null, true)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("an escaped field name larger than one MiB must fail at the output budget"); + + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + + #[test] + fn null_json_struct_does_not_materialize_its_child_encoder_plan() { + let child_name = "\n".repeat(MAX_SELECT_OUTPUT_RECORD_BYTES); + let child = Arc::new(Field::new(child_name, DataType::Utf8, true)); + let values = StructArray::new_null(vec![child].into(), 1); + let schema = Arc::new(Schema::new(vec![Field::new("nested", values.data_type().clone(), true)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("a null struct should encode without its children"), + 1 + ); + assert_eq!(buffer.as_ref(), b"{\"nested\":null}\n"); + } + + #[test] + fn json_record_delimiter_is_in_the_encode_budget() { + let delimiter = "x".repeat(64 * 1024); + let values = StringArray::from(vec!["a", "b"]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput { + record_delimiter: Some(delimiter.clone()), + })); + let mut buffer = BytesMut::new(); + + let encoded_rows = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("bounded JSON delimiter should encode"); + + assert_eq!(encoded_rows, 1); + assert_eq!(buffer.len(), br#"{"value":"a"}"#.len() + delimiter.len()); + } + + #[test] + fn oversized_output_record_fails_before_growing_records_buffer() { + let value = "x".repeat(MAX_SELECT_OUTPUT_RECORD_BYTES + 1); + let values = StringArray::from(vec![value.as_str()]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::new(); + + let error = + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect_err("result records larger than one MiB must fail"); + + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + + #[test] + fn repeated_json_projection_is_bounded_per_field() { + let value = "x".repeat(64 * 1024); + let values = Arc::new(StringArray::from(vec![value.as_str()])); + let fields = (0..1024) + .map(|index| Field::new(format!("value_{index}"), DataType::Utf8, false)) + .collect::>(); + let columns = (0..fields.len()) + .map(|_| Arc::clone(&values) as datafusion::arrow::array::ArrayRef) + .collect::>(); + let batch = RecordBatch::try_new(Arc::new(Schema::new(fields)), columns).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("a repeated projection larger than one MiB must fail"); + + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + + #[test] + fn csv_output_limit_counts_the_record_delimiter() { + for excess in 0..=1 { + let value = "x".repeat(MAX_SELECT_OUTPUT_RECORD_BYTES - 1 + excess); + let values = StringArray::from(vec![value.as_str()]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Csv(CSVOutput::default())); + let mut buffer = BytesMut::new(); + + if excess == 0 { + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect("a complete one MiB CSV record should encode"), + 1 + ); + assert_eq!(buffer.len(), MAX_SELECT_OUTPUT_RECORD_BYTES); + } else { + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("a CSV record exceeding one MiB by its delimiter must fail"); + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + } + } + + #[test] + fn json_output_limit_counts_the_record_delimiter() { + let overhead = br#"{"value":""}"#.len() + 1; + for excess in 0..=1 { + let value = "x".repeat(MAX_SELECT_OUTPUT_RECORD_BYTES - overhead + excess); + let values = StringArray::from(vec![value.as_str()]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + if excess == 0 { + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect("a complete one MiB JSON record should encode"), + 1 + ); + assert_eq!(buffer.len(), MAX_SELECT_OUTPUT_RECORD_BYTES); + } else { + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("a JSON record exceeding one MiB by its delimiter must fail"); + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + } + } + + #[test] + fn json_binary_hex_encoding_honors_the_output_limit() { + let overhead = br#"{"value":""}"#.len() + 1; + let max_value_bytes = (MAX_SELECT_OUTPUT_RECORD_BYTES - overhead) / 2; + for excess in 0..=1 { + let value = vec![0xab; max_value_bytes + excess]; + let values = BinaryArray::from(vec![value.as_slice()]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Binary, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + if excess == 0 { + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect("the largest binary value within the output limit should encode"), + 1 + ); + assert_eq!(buffer.len(), overhead + value.len() * 2); + assert_eq!(buffer.len(), MAX_SELECT_OUTPUT_RECORD_BYTES - 1); + } else { + let error = + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer).expect_err("binary JSON exceeding one MiB must fail"); + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + } + } + + #[test] + fn json_string_escaping_honors_the_exact_output_limit() { + let overhead = br#"{"value":""}"#.len() + 1; + let escaped_budget = MAX_SELECT_OUTPUT_RECORD_BYTES - overhead; + let escaped_newlines = escaped_budget / 2; + let mut exact_value = "\n".repeat(escaped_newlines); + exact_value.push_str(&"x".repeat(escaped_budget % 2)); + + for excess in 0..=1 { + let mut value = exact_value.clone(); + value.push_str(&"x".repeat(excess)); + let values = StringArray::from(vec![value.as_str()]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut encoder = SelectOutputEncoder::new(SelectOutputFormat::Json(JSONOutput::default())); + let mut buffer = BytesMut::new(); + + if excess == 0 { + assert_eq!( + encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect("escaped JSON record of exactly one MiB should encode"), + 1 + ); + assert_eq!(buffer.len(), MAX_SELECT_OUTPUT_RECORD_BYTES); + } else { + let error = encode_batch_turn(&mut encoder, &batch, 0, &mut buffer) + .expect_err("escaped JSON record exceeding one MiB must fail"); + assert_eq!(error.code(), &S3ErrorCode::OverMaxRecordSize); + assert!(buffer.is_empty()); + } + } + } + + #[tokio::test(start_paused = true)] + async fn deadline_preempts_writable_multi_chunk_records() { + const PAYLOAD_CHUNKS: usize = 8; + + let (event_channel, mut rx) = test_event_channel(4); + let (lease, lease_released) = lease_drop_signal(); + let producer = send_select_events_until_deadline( + large_pending_output(PAYLOAD_CHUNKS), + event_channel, + csv_validation(), + Arc::new(SelectInputMetrics::default()), + Instant::now() + Duration::from_secs(1), + 1, + lease, + ); + tokio::pin!(producer); + + for _ in 0..3 { + assert!(futures::poll!(producer.as_mut()).is_pending()); + } + for _ in 0..2 { + assert!(matches!(rx.try_recv(), Ok(Ok(SelectObjectContentEvent::Records(_))))); + } + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + + tokio::time::advance(Duration::from_secs(1)).await; + assert!(futures::poll!(producer.as_mut()).is_ready()); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Records(_))))); + let timeout = rx + .recv() + .await + .expect("deadline should send one terminal error") + .expect_err("deadline terminal event should be an error"); + assert_eq!(timeout.code(), &S3ErrorCode::Busy); + assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "timeout should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn continuation_starts_at_one_second_without_query_output() { + let (producer, mut rx, lease_released) = spawn_test_producer(pending_output(), 4); + tokio::task::yield_now().await; + + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + tokio::time::advance(Duration::from_millis(999)).await; + tokio::task::yield_now().await; + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + + tokio::time::advance(Duration::from_millis(1)).await; + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn progress_starts_at_sixty_seconds_only_when_enabled() { + let mut validation = csv_validation(); + validation.progress_enabled = true; + let (producer, mut rx, lease_released) = + spawn_test_producer_with(pending_output(), 8, validation, Duration::from_secs(300)); + tokio::task::yield_now().await; + + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + tokio::time::advance(Duration::from_millis(59_999)).await; + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + + tokio::time::advance(Duration::from_millis(1)).await; + let Some(Ok(SelectObjectContentEvent::Progress(progress))) = rx.recv().await else { + panic!("enabled progress should fire at sixty seconds"); + }; + let details = progress.details.expect("Progress should contain details"); + assert_eq!(details.bytes_scanned, Some(0)); + assert_eq!(details.bytes_processed, Some(0)); + assert_eq!(details.bytes_returned, Some(0)); + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + + let (producer, mut rx, lease_released) = spawn_test_producer(pending_output(), 8); + tokio::task::yield_now().await; + tokio::time::advance(Duration::from_secs(60)).await; + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + for _ in 0..3 { + tokio::task::yield_now().await; + } + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn small_records_flush_at_five_hundred_milliseconds() { + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(StringArray::from(vec!["row"]))]) + .expect("test record batch should be valid"); + let output = Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }) + .chain(futures::stream::pending::>()), + )); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 4); + for _ in 0..3 { + tokio::task::yield_now().await; + } + + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + tokio::time::advance(Duration::from_millis(499)).await; + tokio::task::yield_now().await; + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + + tokio::time::advance(Duration::from_millis(1)).await; + let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { + panic!("small Records payload should flush at five hundred milliseconds"); + }; + assert_eq!(records.payload.as_deref(), Some(b"row\n".as_slice())); + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn full_records_payload_flushes_without_advancing_time() { + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let value = "x".repeat(RECORDS_CHUNK_TARGET); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(StringArray::from(vec![value]))]) + .expect("test record batch should be valid"); + let output = Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }) + .chain(futures::stream::pending::>()), + )); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 4); + + let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { + panic!("a full Records payload should flush without waiting for the timer"); + }; + assert_eq!(records.payload.as_ref().map(Bytes::len), Some(RECORDS_CHUNK_TARGET)); + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn delayed_intervals_do_not_burst_after_time_jump() { + let (producer, mut rx, lease_released) = spawn_test_producer(pending_output(), 16); + tokio::task::yield_now().await; + + tokio::time::advance(Duration::from_secs(10)).await; + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + for _ in 0..5 { + tokio::task::yield_now().await; + } + assert!(matches!(rx.try_recv(), Err(mpsc::error::TryRecvError::Empty))); + + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn continuous_records_do_not_starve_continuation() { + let (producer, mut rx, lease_released) = spawn_test_producer(large_pending_output(8), 16); + tokio::task::yield_now().await; + + tokio::time::advance(CONTINUATION_INTERVAL).await; + let mut saw_continuation = false; + for _ in 0..10 { + let event = rx + .recv() + .await + .expect("scheduler should emit an event") + .expect("event should not fail"); + if matches!(event, SelectObjectContentEvent::Cont(_)) { + saw_continuation = true; + break; + } + } + assert!(saw_continuation, "continuous Records must not starve the continuation ticker"); + + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn continuous_records_do_not_starve_progress() { + let mut validation = csv_validation(); + validation.progress_enabled = true; + let (producer, mut rx, lease_released) = + spawn_test_producer_with(large_pending_output(8), 16, validation, Duration::from_secs(300)); + tokio::task::yield_now().await; + + tokio::time::advance(PROGRESS_INTERVAL).await; + let mut saw_progress = false; + let mut saw_continuation = false; + for _ in 0..12 { + let event = rx + .recv() + .await + .expect("scheduler should emit an event") + .expect("event should not fail"); + match event { + SelectObjectContentEvent::Progress(progress) => { + assert!( + progress + .details + .is_some_and(|details| details.bytes_returned.is_some_and(|bytes| bytes > 0)) + ); + saw_progress = true; + } + SelectObjectContentEvent::Cont(_) => saw_continuation = true, + _ => {} + } + if saw_progress && saw_continuation { + break; + } + } + assert!(saw_progress, "continuous Records must not starve the progress ticker"); + assert!(saw_continuation, "continuous Records must not starve the continuation ticker"); + + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn backpressured_records_cannot_starve_periodic_events() { + let (producer, mut rx, lease_released) = spawn_test_producer(large_pending_output(4), 3); + tokio::task::yield_now().await; + + tokio::time::advance(CONTINUATION_INTERVAL).await; + let mut records_before_continuation = 0; + loop { + let event = rx + .recv() + .await + .expect("scheduler should emit an event") + .expect("event should not fail"); + match event { + SelectObjectContentEvent::Records(_) => { + records_before_continuation += 1; + assert!( + records_before_continuation <= 2, + "only the queued and already-pending Records events may precede a due continuation" + ); + } + SelectObjectContentEvent::Cont(_) => break, + _ => panic!("unexpected event before the due continuation"), + } + } + + let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { + panic!("buffered Records must resume immediately after the due continuation"); + }; + assert_eq!(records.payload.as_ref().map(Bytes::len), Some(RECORDS_CHUNK_TARGET)); + + drop(rx); + producer.await.expect("producer should stop after the receiver closes"); + assert!(lease_released.await.is_ok(), "receiver close should release the snapshot lease"); + } + #[tokio::test(start_paused = true)] async fn producer_preserves_finite_stream_terminal_events() { let schema = Arc::new(datafusion::arrow::datatypes::Schema::new(vec![datafusion::arrow::datatypes::Field::new( @@ -1206,13 +2907,10 @@ mod tests { producer.await.expect("producer should finish at query EOF"); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); - for expected in [b"a\n".as_slice(), b"b\n".as_slice()] { - let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { - panic!("producer should emit a records event for each batch"); - }; - assert_eq!(records.payload.as_deref(), Some(expected)); - } + let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { + panic!("producer should flush buffered records at query EOF"); + }; + assert_eq!(records.payload.as_deref(), Some(b"a\nb\n".as_slice())); let Some(Ok(SelectObjectContentEvent::Stats(stats))) = rx.recv().await else { panic!("producer should emit final stats"); }; @@ -1222,6 +2920,73 @@ mod tests { assert!(lease_released.await.is_ok(), "End should release the snapshot lease"); } + #[tokio::test(start_paused = true)] + async fn multi_slice_multi_chunk_output_is_complete_and_ordered() { + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let values = (0..(MAX_ENCODE_ROWS_PER_TURN * 2 + 1)) + .map(|index| format!("{index:04}-{}", "x".repeat(72))) + .collect::>(); + let expected = values.iter().map(|value| format!("{value}\n")).collect::(); + assert!(expected.len() > RECORDS_CHUNK_TARGET); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(StringArray::from(values))]) + .expect("test record batch should be valid"); + let output = Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), + )); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 8); + + producer.await.expect("producer should finish successfully"); + + let mut records = Vec::new(); + let mut stats_returned = None; + let mut saw_end = false; + while let Some(event) = rx.recv().await { + match event.expect("successful stream should not emit an error") { + SelectObjectContentEvent::Records(event) => { + records.extend_from_slice(event.payload.expect("Records should contain a payload").as_ref()); + } + SelectObjectContentEvent::Stats(event) => { + assert!(stats_returned.is_none(), "Stats should be emitted once"); + stats_returned = event.details.and_then(|details| details.bytes_returned); + } + SelectObjectContentEvent::End(_) => { + assert!(stats_returned.is_some(), "End must follow Stats"); + saw_end = true; + } + _ => panic!("finite query should emit only Records, Stats, and End"), + } + } + + assert_eq!(records, expected.as_bytes()); + assert_eq!( + stats_returned, + Some(i64::try_from(records.len()).expect("test output length should fit in i64")) + ); + assert!(saw_end); + assert!(lease_released.await.is_ok(), "End should release the snapshot lease"); + } + + #[tokio::test(start_paused = true)] + async fn empty_stream_emits_only_stats_then_end() { + let schema = Arc::new(Schema::empty()); + let output = Box::pin(RecordBatchStreamAdapter::new( + schema, + futures::stream::empty::>(), + )); + let (producer, mut rx, lease_released) = spawn_test_producer(output, 4); + + producer.await.expect("empty producer should finish successfully"); + + let Some(Ok(SelectObjectContentEvent::Stats(stats))) = rx.recv().await else { + panic!("empty result should start with Stats"); + }; + assert_eq!(stats.details.and_then(|details| details.bytes_returned), Some(0)); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_))))); + assert!(rx.recv().await.is_none()); + assert!(lease_released.await.is_ok(), "End should release the snapshot lease"); + } + #[tokio::test(start_paused = true)] async fn successful_stream_serializes_records_stats_and_end_without_error() { let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); @@ -1248,7 +3013,7 @@ mod tests { .find_map(|(name, value)| (name == ":event-type").then_some(value.as_str())) }) .collect::>(); - assert_eq!(event_types, ["Cont", "Records", "Stats", "End"]); + assert_eq!(event_types, ["Records", "Stats", "End"]); assert!(!messages.iter().flatten().any(|(name, value)| { (name == ":message-type" && value == "error") || name == ":error-code" || name == ":error-message" })); @@ -1256,7 +3021,7 @@ mod tests { } #[tokio::test(start_paused = true)] - async fn eof_at_deadline_uses_reserved_slot_for_stats_then_end() { + async fn deadline_wins_when_eof_becomes_ready_at_same_instant() { let output = Box::pin(RecordBatchStreamAdapter::new( Arc::new(Schema::empty()), futures::stream::unfold((), |_| async { @@ -1264,26 +3029,24 @@ mod tests { None::<(Result, ())> }), )); - let (producer, mut rx, lease_released) = spawn_test_producer(output, 3); + let (producer, mut rx, lease_released) = spawn_test_producer_with(output, 3, csv_validation(), Duration::from_secs(1)); tokio::task::yield_now().await; tokio::time::advance(std::time::Duration::from_secs(1)).await; producer.await.expect("producer should finish at the shared deadline"); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); - let stats = rx + let timeout = rx .recv() .await - .expect("successful Select should send stats") - .expect("stats event should not be an error"); - assert!(matches!(stats, SelectObjectContentEvent::Stats(_))); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::End(_))))); + .expect("deadline should send one terminal error") + .expect_err("deadline terminal event should be an error"); + assert_eq!(timeout.code(), &S3ErrorCode::Busy); assert!(rx.recv().await.is_none()); - assert!(lease_released.await.is_ok(), "EOF should release the snapshot lease"); + assert!(lease_released.await.is_ok(), "deadline should release the snapshot lease"); } #[tokio::test(start_paused = true)] - async fn stream_error_at_deadline_uses_reserved_terminal_slot() { + async fn deadline_wins_when_stream_error_becomes_ready_at_same_instant() { let output = Box::pin(RecordBatchStreamAdapter::new( Arc::new(Schema::empty()), futures::stream::once(async { @@ -1291,21 +3054,20 @@ mod tests { Err(DataFusionError::External(Box::new(SelectError::QueryConcurrencyLimit))) }), )); - let (producer, mut rx, lease_released) = spawn_test_producer(output, 2); + let (producer, mut rx, lease_released) = spawn_test_producer_with(output, 3, csv_validation(), Duration::from_secs(1)); tokio::task::yield_now().await; tokio::time::advance(std::time::Duration::from_secs(1)).await; producer.await.expect("producer should finish at the shared deadline"); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); - let stream_error = rx + let timeout = rx .recv() .await - .expect("stream failure should send one terminal error") + .expect("deadline should send one terminal error") .expect_err("terminal event should be an error"); - assert_eq!(stream_error.code(), &S3ErrorCode::SlowDown); + assert_eq!(timeout.code(), &S3ErrorCode::Busy); assert!(rx.recv().await.is_none()); - assert!(lease_released.await.is_ok(), "stream error should release the snapshot lease"); + assert!(lease_released.await.is_ok(), "deadline should release the snapshot lease"); } #[tokio::test(start_paused = true)] @@ -1458,7 +3220,7 @@ mod tests { .find_map(|(name, value)| (name == ":event-type").then_some(value.as_str())) }) .collect::>(), - ["Cont", "Records"] + ["Records"] ); let terminal_headers = messages.last().expect("stream should contain a terminal error"); assert!( @@ -1491,11 +3253,8 @@ mod tests { )); let (producer, mut rx, lease_released) = spawn_test_producer(output, 2); - tokio::task::yield_now().await; - tokio::time::advance(std::time::Duration::from_secs(1)).await; producer.await.expect("producer should not block on a terminal encoder error"); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); let encoder_error = rx .recv() .await @@ -1518,25 +3277,20 @@ mod tests { futures::future::pending::>().await }), )); - let (tx, mut rx) = mpsc::channel(2); - let terminal_permit = tx - .clone() - .try_reserve_owned() - .expect("test channel should reserve terminal capacity"); + let (event_channel, rx) = test_event_channel(2); let (lease, lease_released) = lease_drop_signal(); let producer = send_select_events_until_deadline( output, - SelectEventChannel { tx, terminal_permit }, + event_channel, csv_validation(), Arc::new(SelectInputMetrics::default()), - Instant::now() + std::time::Duration::from_secs(1), + Instant::now() + Duration::from_secs(300), 300, lease, ); tokio::pin!(producer); assert!(futures::poll!(producer.as_mut()).is_pending()); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); drop(rx); assert!( @@ -1562,14 +3316,20 @@ mod tests { Ok(RecordBatch::new_empty(Arc::new(Schema::empty()))) }), )); - let (tx, mut rx) = mpsc::channel(2); + let (mut event_channel, rx) = test_event_channel(2); let snapshot_fence = LeaseDropSignal(None); - let producer = - send_select_events(output, &tx, csv_validation(), Arc::new(SelectInputMetrics::default()), &snapshot_fence); + let producer = send_select_events( + output, + &mut event_channel, + csv_validation(), + Arc::new(SelectInputMetrics::default()), + Instant::now() + Duration::from_secs(300), + 300, + &snapshot_fence, + ); tokio::pin!(producer); assert!(futures::poll!(producer.as_mut()).is_pending()); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); drop(rx); ready_tx.send(()).expect("test should make the query stream ready"); @@ -1599,13 +3359,15 @@ mod tests { schema, futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), )); - let (tx, mut rx) = mpsc::channel(4); + let (mut event_channel, mut rx) = test_event_channel(4); let outcome = send_select_events( output, - &tx, + &mut event_channel, csv_validation(), Arc::new(SelectInputMetrics::default()), + Instant::now() + Duration::from_secs(300), + 300, &FailingSnapshotFence, ) .await; @@ -1614,30 +3376,36 @@ mod tests { panic!("failed final snapshot fence must produce a terminal error"); }; assert_eq!(error.code(), &S3ErrorCode::InternalError); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Records(_))))); assert!(rx.try_recv().is_err(), "failed final fence must not enqueue Stats or End"); } #[tokio::test] async fn producer_rechecks_snapshot_after_stats_backpressure() { + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema.clone(), vec![Arc::new(StringArray::from(vec!["row"]))]) + .expect("test record batch should be valid"); let output = Box::pin(RecordBatchStreamAdapter::new( - Arc::new(Schema::empty()), - futures::stream::empty::>(), + schema, + futures::stream::once(async move { Ok::<_, DataFusionError>(batch) }), )); - let (tx, mut rx) = mpsc::channel(2); - let _terminal_permit = tx - .clone() - .try_reserve_owned() - .expect("test channel should reserve terminal capacity"); + let (mut event_channel, mut rx) = test_event_channel(2); let snapshot_fence = FailsAfterFirstSnapshotFence(std::sync::atomic::AtomicUsize::new(0)); - let producer = - send_select_events(output, &tx, csv_validation(), Arc::new(SelectInputMetrics::default()), &snapshot_fence); + let producer = send_select_events( + output, + &mut event_channel, + csv_validation(), + Arc::new(SelectInputMetrics::default()), + Instant::now() + Duration::from_secs(300), + 300, + &snapshot_fence, + ); tokio::pin!(producer); + assert!(futures::poll!(producer.as_mut()).is_pending()); assert!(futures::poll!(producer.as_mut()).is_pending()); assert_eq!(snapshot_fence.0.load(std::sync::atomic::Ordering::Relaxed), 1); - assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Cont(_))))); + assert!(matches!(rx.recv().await, Some(Ok(SelectObjectContentEvent::Records(_))))); let SelectProducerOutcome::Terminal(Err(error)) = producer.await else { panic!("snapshot loss during Stats backpressure must reject successful End"); @@ -1891,6 +3659,22 @@ mod tests { assert_eq!(err.code(), &S3ErrorCode::InvalidRequestParameter); } + fn assert_json_encoder_matches_arrow(array: ArrayRef) { + let schema = Arc::new(Schema::new(vec![Field::new("value", array.data_type().clone(), true)])); + let batch = RecordBatch::try_new(schema, vec![array]).expect("test record batch should be valid"); + let mut expected = Vec::new(); + { + let mut writer = WriterBuilder::new() + .with_explicit_nulls(true) + .build::<_, LineDelimited>(&mut expected); + writer.write(&batch).expect("Arrow JSON reference should encode"); + writer.finish().expect("Arrow JSON reference should finish"); + } + let mut actual = BytesMut::new(); + encode_json_batch(&batch, &JSONOutput::default(), &mut actual).expect("S3 Select JSON should encode"); + assert_eq!(actual.as_ref(), expected.as_slice(), "type: {}", batch.column(0).data_type()); + } + #[test] fn json_encoder_outputs_line_delimited_records() { let schema = @@ -1907,9 +3691,9 @@ mod tests { ) .unwrap(); - let bytes = encode_json_batch(&batch, &JSONOutput::default()).unwrap(); - let output = String::from_utf8(bytes).unwrap(); - assert_eq!(output, "{\"name\":\"a\"}\n{\"name\":\"b\"}\n"); + let mut bytes = BytesMut::new(); + encode_json_batch(&batch, &JSONOutput::default(), &mut bytes).unwrap(); + assert_eq!(bytes.as_ref(), b"{\"name\":\"a\"}\n{\"name\":\"b\"}\n"); } #[test] @@ -1928,15 +3712,44 @@ mod tests { ) .unwrap(); - let bytes = encode_json_batch( + let mut bytes = BytesMut::new(); + encode_json_batch( &batch, &JSONOutput { record_delimiter: Some("|".to_string()), }, + &mut bytes, ) .unwrap(); - let output = String::from_utf8(bytes).unwrap(); - assert_eq!(output, "{\"name\":\"a\"}|{\"name\":\"b\"}|"); + assert_eq!(bytes.as_ref(), b"{\"name\":\"a\"}|{\"name\":\"b\"}|"); + } + + #[test] + fn json_encoder_matches_arrow_for_limited_encoder_variants() { + let mut fixed_list = FixedSizeListBuilder::new(Int32Builder::new(), 2); + fixed_list.values().append_value(1); + fixed_list.values().append_null(); + fixed_list.append(true); + + let arrays: Vec = vec![ + Arc::new(LargeStringArray::from(vec!["a\n"])), + Arc::new(StringViewArray::from(vec!["a\n"])), + Arc::new(BinaryArray::from(vec![b"\xab".as_slice()])), + Arc::new(LargeBinaryArray::from(vec![b"\xab".as_slice()])), + Arc::new(BinaryViewArray::from(vec![b"\xab".as_slice()])), + Arc::new( + FixedSizeBinaryArray::try_from_iter([b"\xab".as_slice()].into_iter()) + .expect("fixed binary test array should be valid"), + ), + Arc::new(LargeListArray::from_iter_primitive::([Some([Some(1), None])])), + Arc::new(ListViewArray::from_iter_primitive::([Some([Some(1), None])])), + Arc::new(LargeListViewArray::from_iter_primitive::([Some([Some(1), None])])), + Arc::new(fixed_list.finish()), + ]; + + for array in arrays { + assert_json_encoder_matches_arrow(array); + } } #[test] @@ -1954,27 +3767,149 @@ mod tests { ) .unwrap(); - let bytes = encode_csv_batch( + let mut bytes = BytesMut::new(); + encode_csv_batch( &batch, &CSVOutput { field_delimiter: Some("|".to_string()), record_delimiter: Some("\r\n".to_string()), ..Default::default() }, + &mut bytes, ) .unwrap(); - assert_eq!(String::from_utf8(bytes).unwrap(), "a|1\r\nb|2\r\n"); + assert_eq!(bytes.as_ref(), b"a|1\r\nb|2\r\n"); } #[test] - fn split_records_payload_uses_exact_returned_bytes() { - let payloads = split_records_payload(vec![b'x'; RECORDS_CHUNK_TARGET + 7]); + fn csv_encoder_matches_select_as_needed_quote_rules() { + let values = StringArray::from(vec!["", "\\.", "\u{00a0}value", "line\rbreak", "a\"b", "a|b"]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut bytes = BytesMut::new(); + + encode_csv_batch( + &batch, + &CSVOutput { + quote_escape_character: Some("\\".to_string()), + quote_fields: Some(QuoteFields::from_static(QuoteFields::ASNEEDED)), + record_delimiter: Some("|".to_string()), + ..Default::default() + }, + &mut bytes, + ) + .expect("CSV output should encode"); + + let expected = [ + b"|".as_slice(), + br#""\.""#, + b"|", + "\"\u{00a0}value\"|".as_bytes(), + b"\"line\rbreak\"|", + br#""a\"b"|"#, + b"a|b|", + ] + .concat(); + assert_eq!(bytes.as_ref(), expected); + } + + #[test] + fn csv_encoder_honors_always_and_custom_quote_characters() { + let values = StringArray::from(vec!["plain", "a'b"]); + let schema = Arc::new(Schema::new(vec![Field::new("value", DataType::Utf8, false)])); + let batch = RecordBatch::try_new(schema, vec![Arc::new(values)]).expect("test record batch should be valid"); + let mut bytes = BytesMut::new(); + + encode_csv_batch( + &batch, + &CSVOutput { + quote_character: Some("'".to_string()), + quote_escape_character: Some("\\".to_string()), + quote_fields: Some(QuoteFields::from_static(QuoteFields::ALWAYS)), + record_delimiter: Some("|".to_string()), + ..Default::default() + }, + &mut bytes, + ) + .expect("CSV output should encode"); + + assert_eq!(bytes.as_ref(), br#"'plain'|'a\'b'|"#); + } + + #[tokio::test(start_paused = true)] + async fn records_staging_preserves_payload_limit_and_returned_bytes() { + let mut buffer = BytesMut::from(vec![b'x'; RECORDS_CHUNK_TARGET + 7].as_slice()); + let flush = tokio::time::sleep(Duration::from_secs(300)); + tokio::pin!(flush); + let mut flush_armed = false; + let mut pending = None; + schedule_buffered_records( + &mut buffer, + flush.as_mut(), + &mut flush_armed, + &mut pending, + Instant::now() + Duration::from_secs(300), + ); + let Some(SelectObjectContentEvent::Records(records)) = pending else { + panic!("full payload should flush immediately"); + }; + let first = records.payload.expect("Records should contain a payload"); + assert_eq!(first.len(), RECORDS_CHUNK_TARGET); + let second = take_records_payload(&mut buffer).expect("remaining payload should stay staged"); + assert_eq!(second.len(), 7); + let mut progress = SelectProgress::new(Some(Arc::new(SelectInputMetrics::default()))); - for payload in &payloads { - progress.add_returned(payload.len()); - } + progress.add_returned(first.len()); + progress.add_returned(second.len()); assert_eq!(progress.to_stats().bytes_returned, Some((RECORDS_CHUNK_TARGET + 7) as i64)); - assert!(payloads.len() > 1); + } + + #[tokio::test] + async fn terminal_error_sends_at_most_one_compat_records_chunk() { + let (mut event_channel, mut rx) = test_event_channel(2); + let mut pending_event = None; + let mut records_buffer = BytesMut::from(vec![b'x'; RECORDS_CHUNK_TARGET + 17].as_slice()); + let mut progress = SelectProgress::new(Some(Arc::new(SelectInputMetrics::default()))); + + flush_terminal_records( + &mut event_channel, + &mut pending_event, + &mut records_buffer, + &mut progress, + TerminalRecordsMode::PrefixBeforeError, + ) + .expect("the reserved terminal Records slot should be available"); + + let Some(Ok(SelectObjectContentEvent::Records(records))) = rx.recv().await else { + panic!("the terminal prefix should be sent as Records"); + }; + assert_eq!(records.payload.as_ref().map(Bytes::len), Some(RECORDS_CHUNK_TARGET)); + assert!(records_buffer.is_empty(), "failed output after the terminal prefix must be discarded"); + assert_eq!(progress.to_stats().bytes_returned, Some(RECORDS_CHUNK_TARGET as i64)); + } + + #[tokio::test] + async fn maximum_records_payload_stays_within_compat_message_limit() { + let (tx, rx) = mpsc::channel(1); + tx.send(Ok(records_event(Bytes::from(vec![b'x'; RECORDS_CHUNK_TARGET])))) + .await + .expect("test channel should accept Records"); + drop(tx); + + let mut byte_stream = SelectObjectContentEventStream::new(ReceiverStream::new(rx)).into_byte_stream(); + let mut encoded = Vec::new(); + while let Some(chunk) = byte_stream.next().await { + encoded.extend_from_slice(&chunk.expect("Records event should serialize")); + } + + let total_len = usize::try_from(u32::from_be_bytes( + encoded[0..4] + .try_into() + .expect("event-stream message should contain a prelude"), + )) + .expect("event-stream message length should fit in usize"); + assert_eq!(total_len, encoded.len()); + assert_eq!(total_len, MAX_COMPAT_EVENT_STREAM_MESSAGE_BYTES); } #[test]