diff --git a/crates/s3select-api/src/object_store.rs b/crates/s3select-api/src/object_store.rs index cdb859e02..60d744ebc 100644 --- a/crates/s3select-api/src/object_store.rs +++ b/crates/s3select-api/src/object_store.rs @@ -18,30 +18,34 @@ use chrono::Utc; use futures::pin_mut; use futures::{Stream, StreamExt, future::ready, stream}; use futures_core::stream::BoxStream; -use http::HeaderMap; +use http::{HeaderMap, HeaderValue, header::HeaderName}; use object_store::{ - Attributes, CopyOptions, Error as o_Error, GetOptions, GetResult, ListResult, MultipartUpload, ObjectMeta, ObjectStore, - PutMultipartOptions, PutOptions, PutPayload, PutResult, Result, path::Path, + Attributes, CopyOptions, Error as o_Error, GetOptions, GetRange, GetResult, GetResultPayload, ListResult, MultipartUpload, + ObjectMeta, ObjectStore, PutMultipartOptions, PutOptions, PutPayload, PutResult, Result, path::Path, }; use pin_project_lite::pin_project; use rustfs_common::DEFAULT_DELIMITER; +use rustfs_ecstore::error::{StorageError, is_err_bucket_not_found, is_err_object_not_found, is_err_version_not_found}; use rustfs_ecstore::new_object_layer_fn; use rustfs_ecstore::set_disk::DEFAULT_READ_BUFFER_SIZE; use rustfs_ecstore::store::ECStore; -use rustfs_ecstore::store_api::ObjectIO; -use rustfs_ecstore::store_api::ObjectOptions; +use rustfs_ecstore::store_api::{GetObjectReader, HTTPRangeSpec, ObjectIO, ObjectOperations, ObjectOptions}; use s3s::S3Result; use s3s::dto::SelectObjectContentInput; +use s3s::header::{ + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY, + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, +}; use s3s::s3_error; +use std::collections::VecDeque; use std::ops::Range; use std::pin::Pin; use std::sync::Arc; use std::task::Poll; use std::task::ready; -use tokio::io::AsyncRead; use tokio::io::AsyncReadExt; +use tokio::io::{AsyncRead, ReadBuf}; use tokio_util::io::ReaderStream; -use tracing::info; use transform_stream::AsyncTryStream; /// Maximum allowed object size for JSON DOCUMENT mode. @@ -58,6 +62,8 @@ use transform_stream::AsyncTryStream; /// Default: 128 MiB. This matches the AWS S3 Select limit for JSON DOCUMENT /// inputs. pub const MAX_JSON_DOCUMENT_BYTES: u64 = 128 * 1024 * 1024; +pub const INVALID_SCAN_RANGE_MESSAGE: &str = + "The value of a parameter in ScanRange element is invalid. Check the service API documentation and try again."; #[derive(Debug)] pub struct EcObjectStore { @@ -75,6 +81,16 @@ pub struct EcObjectStore { store: Arc, } + +#[derive(Clone, Copy, Debug)] +struct SelectScanRange { + start: u64, + end: u64, +} + +#[derive(Clone, Copy, Debug)] +pub struct InvalidScanRange; + impl EcObjectStore { pub fn new(input: Arc) -> S3Result { let Some(store) = new_object_layer_fn() else { @@ -123,6 +139,97 @@ impl EcObjectStore { store, }) } + + fn object_options(&self, options: &GetOptions) -> ObjectOptions { + ObjectOptions { + version_id: options.version.clone(), + ..Default::default() + } + } + + fn read_headers(&self) -> HeaderMap { + select_read_headers(&self.input) + } + + fn scan_range(&self, object_size: u64) -> Result> { + let Some(scan_range) = self.input.request.scan_range.as_ref() else { + return Ok(None); + }; + scan_range_from_bounds(scan_range.start, scan_range.end, object_size) + } + + fn record_delimiter(&self) -> Vec { + self.input + .request + .input_serialization + .csv + .as_ref() + .and_then(|csv| csv.record_delimiter.as_ref()) + .map(|delimiter| delimiter.as_bytes().to_vec()) + .unwrap_or_else(|| b"\n".to_vec()) + } + + fn csv_has_header(&self) -> bool { + self.input + .request + .input_serialization + .csv + .as_ref() + .and_then(|csv| csv.file_header_info.as_ref()) + .is_some_and(|info| matches!(info.as_str(), "USE" | "IGNORE")) + } + + async fn object_info(&self, opts: &ObjectOptions) -> Result { + self.store + .get_object_info(&self.input.bucket, &self.input.key, opts) + .await + .map_err(|err| map_storage_error(&self.input.bucket, &self.input.key, err)) + } + + async fn object_reader(&self, range: Option, opts: &ObjectOptions) -> Result { + let h = self.read_headers(); + self.store + .get_object_reader(&self.input.bucket, &self.input.key, range, h, opts) + .await + .map_err(|err| map_storage_error(&self.input.bucket, &self.input.key, err)) + } + + async fn read_raw_range_with_opts(&self, range: Range, opts: &ObjectOptions) -> Result { + if range.is_empty() { + return Ok(Bytes::new()); + } + let reader = self.object_reader(Some(http_range_spec_from_range(range)), opts).await?; + let mut reader = reader.stream; + let mut bytes = Vec::new(); + reader.read_to_end(&mut bytes).await.map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + })?; + Ok(Bytes::from(bytes)) + } + + async fn read_raw_range(&self, range: Range) -> Result { + self.read_raw_range_with_opts(range, &self.object_options(&GetOptions::new())) + .await + } + + async fn read_header_record(&self, object_size: u64, delimiter: &[u8], opts: &ObjectOptions) -> Result { + if object_size == 0 { + return Ok(Bytes::new()); + } + + let mut end = (DEFAULT_READ_BUFFER_SIZE as u64).min(object_size); + loop { + let bytes = self.read_raw_range_with_opts(0..end, opts).await?; + if let Some(pos) = find_delimiter(&bytes, delimiter) { + return Ok(bytes.slice(0..pos + delimiter.len())); + } + if end == object_size { + return Ok(bytes); + } + end = end.saturating_mul(2).min(object_size); + } + } } impl std::fmt::Display for EcObjectStore { @@ -141,6 +248,150 @@ fn unsupported_store_error(op: &str) -> o_Error { } } +fn insert_header(headers: &mut HeaderMap, name: HeaderName, value: Option<&str>) { + if let Some(value) = value + && let Ok(value) = HeaderValue::from_str(value) + { + headers.insert(name, value); + } +} + +fn select_read_headers(input: &SelectObjectContentInput) -> HeaderMap { + let mut headers = HeaderMap::new(); + insert_header( + &mut headers, + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, + input.sse_customer_algorithm.as_deref(), + ); + insert_header(&mut headers, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY, input.sse_customer_key.as_deref()); + insert_header( + &mut headers, + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, + input.sse_customer_key_md5.as_deref(), + ); + headers +} + +fn http_range_spec_from_get_range(range: &GetRange) -> HTTPRangeSpec { + match range { + GetRange::Bounded(range) => http_range_spec_from_range(range.clone()), + GetRange::Offset(start) => HTTPRangeSpec { + is_suffix_length: false, + start: *start as i64, + end: -1, + }, + GetRange::Suffix(length) => HTTPRangeSpec { + is_suffix_length: true, + start: *length as i64, + end: -1, + }, + } +} + +fn http_range_spec_from_range(range: Range) -> HTTPRangeSpec { + HTTPRangeSpec { + is_suffix_length: false, + start: range.start as i64, + end: range.end.saturating_sub(1) as i64, + } +} + +fn http_range_spec_from_start(start: u64) -> HTTPRangeSpec { + HTTPRangeSpec { + is_suffix_length: false, + start: start as i64, + end: -1, + } +} + +fn scan_range_read_start(scan_range: SelectScanRange, delimiter: &[u8]) -> u64 { + scan_range.start.saturating_sub(delimiter.len() as u64) +} + +fn find_delimiter(bytes: &[u8], delimiter: &[u8]) -> Option { + if delimiter.is_empty() { + return None; + } + bytes.windows(delimiter.len()).position(|window| window == delimiter) +} + +fn map_storage_error(bucket: &str, object: &str, err: StorageError) -> o_Error { + if is_err_bucket_not_found(&err) || is_err_object_not_found(&err) || is_err_version_not_found(&err) { + return o_Error::NotFound { + path: format!("{bucket}/{object}"), + source: err.to_string().into(), + }; + } + o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + } +} + +fn scan_range_from_bounds(start: Option, end: Option, object_size: u64) -> Result> { + parse_scan_range_from_bounds(start, end, object_size).map_err(|_| invalid_scan_range_store_error()) +} + +pub fn validate_scan_range_bounds( + start: Option, + end: Option, + object_size: u64, +) -> std::result::Result<(), InvalidScanRange> { + parse_scan_range_from_bounds(start, end, object_size).map(|_| ()) +} + +fn parse_scan_range_from_bounds( + start: Option, + end: Option, + object_size: u64, +) -> std::result::Result, InvalidScanRange> { + if start.is_none() && end.is_none() { + return Ok(None); + } + if start.is_some_and(|value| value < 0) || end.is_some_and(|value| value < 0) { + return Err(InvalidScanRange); + } + if let (Some(start), Some(end)) = (start, end) + && start > end + { + return Err(InvalidScanRange); + } + if let Some(start) = start { + let start = start as u64; + if object_size == 0 { + if start > 0 { + return Err(InvalidScanRange); + } + return Ok(Some(SelectScanRange { start: 0, end: 0 })); + } + if start >= object_size { + return Err(InvalidScanRange); + } + } + if object_size == 0 { + return Ok(Some(SelectScanRange { start: 0, end: 0 })); + } + + let last_byte = object_size - 1; + let (start, end) = match (start, end) { + (Some(start), Some(end)) => (start as u64, (end as u64).min(last_byte)), + (Some(start), None) => (start as u64, last_byte), + (None, Some(suffix_len)) => { + let suffix_len = suffix_len as u64; + (object_size.saturating_sub(suffix_len), last_byte) + } + (None, None) => return Ok(None), + }; + Ok(Some(SelectScanRange { start, end })) +} + +fn invalid_scan_range_store_error() -> o_Error { + o_Error::Generic { + store: "EcObjectStore", + source: format!("ScanRange: {INVALID_SCAN_RANGE_MESSAGE}").into(), + } +} + #[async_trait] impl ObjectStore for EcObjectStore { async fn put_opts(&self, _location: &Path, _payload: PutPayload, _opts: PutOptions) -> Result { @@ -151,24 +402,50 @@ impl ObjectStore for EcObjectStore { Err(unsupported_store_error("put_multipart_opts")) } - async fn get_opts(&self, location: &Path, _options: GetOptions) -> Result { - info!("{:?}", location); - let opts = ObjectOptions::default(); - let h = HeaderMap::new(); - let reader = self - .store - .get_object_reader(&self.input.bucket, &self.input.key, None, h, &opts) - .await - .map_err(|_| o_Error::NotFound { - path: format!("{}/{}", self.input.bucket, self.input.key), - source: "can not get object info".into(), - })?; + async fn get_opts(&self, location: &Path, options: GetOptions) -> Result { + let opts = self.object_options(&options); + let needs_scan_context = options.range.is_none() && !options.head && self.input.request.scan_range.is_some(); + let source_size = if needs_scan_context { + Some(self.object_info(&opts).await?.size as u64) + } else { + None + }; + let scan_context = if needs_scan_context { + let original_size = source_size.expect("source size is loaded when scan range is present"); + self.scan_range(original_size)?.map(|scan_range| (original_size, scan_range)) + } else { + None + }; - let original_size = reader.object_info.size as u64; + let range = options.range.as_ref().map(http_range_spec_from_get_range); + let reader = if let Some((original_size, scan_range)) = scan_context.as_ref() { + let delimiter = self.record_delimiter(); + let read_start = scan_range_read_start(*scan_range, &delimiter); + let range = (*original_size > 0).then(|| http_range_spec_from_start(read_start)); + self.object_reader(range, &opts).await? + } else { + self.object_reader(range, &opts).await? + }; + + let original_size = source_size.unwrap_or(reader.object_info.size as u64); let etag = reader.object_info.etag; + let version = reader.object_info.version_id.map(|version| version.to_string()); let attributes = Attributes::default(); + let result_range = match options.range.as_ref() { + Some(range) => range.as_range(original_size).map_err(|err| o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + })?, + None => 0..original_size, + }; - let (payload, size) = if self.is_json_document { + let payload = if options.head { + GetResultPayload::Stream(stream::empty().boxed()) + } else if options.range.is_some() { + let size = (result_range.end - result_range.start) as usize; + let stream = bytes_stream(ReaderStream::with_capacity(reader.stream, DEFAULT_READ_BUFFER_SIZE), size).boxed(); + GetResultPayload::Stream(stream) + } else if self.is_json_document { // JSON DOCUMENT mode: gate on object size before doing any I/O. // // Small files (<= MAX_JSON_DOCUMENT_BYTES): build a lazy stream @@ -195,41 +472,68 @@ impl ObjectStore for EcObjectStore { }); } let stream = json_document_ndjson_stream(reader.stream, original_size, self.json_sub_path.clone()); - (object_store::GetResultPayload::Stream(stream), original_size) + GetResultPayload::Stream(stream) + } else if let Some((_, scan_range)) = scan_context { + let delimiter = self.record_delimiter(); + let include_header = self.csv_has_header(); + let read_start = scan_range_read_start(scan_range, &delimiter); + let header = if include_header && scan_range.start > 0 { + Some(self.read_header_record(original_size, &delimiter, &opts).await?) + } else { + None + }; + let stream = scan_range_stream( + ReaderStream::with_capacity(reader.stream, DEFAULT_READ_BUFFER_SIZE), + delimiter, + scan_range, + include_header && header.is_none(), + read_start, + ) + .boxed(); + let stream = if let Some(header) = header { + stream::once(ready(Ok(header))).chain(stream).boxed() + } else { + stream + }; + GetResultPayload::Stream(convert_field_delimiter_stream(stream, self.need_convert.then(|| self.delimiter.clone()))) } else if self.need_convert { let stream = bytes_stream( ReaderStream::with_capacity(ConvertStream::new(reader.stream, self.delimiter.clone()), DEFAULT_READ_BUFFER_SIZE), original_size as usize, ) .boxed(); - (object_store::GetResultPayload::Stream(stream), original_size) + GetResultPayload::Stream(stream) } else { let stream = bytes_stream( ReaderStream::with_capacity(reader.stream, DEFAULT_READ_BUFFER_SIZE), original_size as usize, ) .boxed(); - (object_store::GetResultPayload::Stream(stream), original_size) + GetResultPayload::Stream(stream) }; let meta = ObjectMeta { location: location.clone(), last_modified: Utc::now(), - size, + size: original_size, e_tag: etag, - version: None, + version, }; Ok(GetResult { payload, meta, - range: 0..size, + range: result_range, attributes, }) } - async fn get_ranges(&self, _location: &Path, _ranges: &[Range]) -> Result> { - Err(unsupported_store_error("get_ranges")) + async fn get_ranges(&self, _location: &Path, ranges: &[Range]) -> Result> { + let mut out = Vec::with_capacity(ranges.len()); + for range in ranges { + out.push(self.read_raw_range(range.clone()).await?); + } + Ok(out) } fn delete_stream(&self, _locations: BoxStream<'static, Result>) -> BoxStream<'static, Result> { @@ -252,7 +556,11 @@ impl ObjectStore for EcObjectStore { pin_project! { struct ConvertStream { inner: R, - delimiter: Vec, + converter: DelimiterConverter, + read_buf: Vec, + pending: Vec, + pending_pos: usize, + eof: bool, } } @@ -260,29 +568,120 @@ impl ConvertStream { fn new(inner: R, delimiter: String) -> Self { ConvertStream { inner, - delimiter: delimiter.as_bytes().to_vec(), + converter: DelimiterConverter::new(delimiter.into_bytes()), + read_buf: vec![0; DEFAULT_READ_BUFFER_SIZE], + pending: Vec::new(), + pending_pos: 0, + eof: false, } } } impl AsyncRead for ConvertStream { #[tracing::instrument(level = "debug", skip_all)] - fn poll_read( - self: Pin<&mut Self>, - cx: &mut std::task::Context<'_>, - buf: &mut tokio::io::ReadBuf<'_>, - ) -> Poll> { - let me = self.project(); - ready!(Pin::new(&mut *me.inner).poll_read(cx, buf))?; - let bytes = buf.filled(); - let replaced = replace_symbol(me.delimiter, bytes); - buf.clear(); - buf.put_slice(&replaced); - Poll::Ready(Ok(())) + fn poll_read(self: Pin<&mut Self>, cx: &mut std::task::Context<'_>, buf: &mut ReadBuf<'_>) -> Poll> { + if buf.remaining() == 0 { + return Poll::Ready(Ok(())); + } + let this = self.project(); + loop { + if drain_pending(this.pending, this.pending_pos, buf) || *this.eof { + return Poll::Ready(Ok(())); + } + + let read_len = DEFAULT_READ_BUFFER_SIZE.min(buf.remaining().max(1)); + let bytes_read = { + let mut read_buf = ReadBuf::new(&mut this.read_buf[..read_len]); + ready!(Pin::new(&mut *this.inner).poll_read(cx, &mut read_buf))?; + read_buf.filled().len() + }; + if bytes_read == 0 { + *this.eof = true; + *this.pending = this.converter.finish(); + *this.pending_pos = 0; + continue; + } + + *this.pending = this.converter.convert_chunk(&this.read_buf[..bytes_read]); + *this.pending_pos = 0; + } } } +struct DelimiterConverter { + delimiter: Vec, + carry: Vec, +} + +impl DelimiterConverter { + fn new(delimiter: Vec) -> Self { + Self { + delimiter, + carry: Vec::new(), + } + } + + fn convert_chunk(&mut self, chunk: &[u8]) -> Vec { + if self.delimiter.is_empty() { + return chunk.to_vec(); + } + + let mut combined = Vec::with_capacity(self.carry.len() + chunk.len()); + combined.extend_from_slice(&self.carry); + combined.extend_from_slice(chunk); + + let safe_end = combined.len().saturating_sub(self.delimiter.len().saturating_sub(1)); + let mut converted = Vec::with_capacity(combined.len()); + let mut pos = 0; + while pos < safe_end { + if combined[pos..].starts_with(&self.delimiter) { + converted.push(DEFAULT_DELIMITER); + pos += self.delimiter.len(); + } else { + converted.push(combined[pos]); + pos += 1; + } + } + self.carry.clear(); + self.carry.extend_from_slice(&combined[pos..]); + converted + } + + fn finish(&mut self) -> Vec { + if self.delimiter.is_empty() { + return std::mem::take(&mut self.carry); + } + let converted = replace_symbol(&self.delimiter, &self.carry); + self.carry.clear(); + converted + } +} + +fn drain_pending(pending: &mut Vec, pending_pos: &mut usize, buf: &mut ReadBuf<'_>) -> bool { + if *pending_pos >= pending.len() { + pending.clear(); + *pending_pos = 0; + return false; + } + if buf.remaining() == 0 { + return false; + } + + let len = buf.remaining().min(pending.len() - *pending_pos); + buf.put_slice(&pending[*pending_pos..*pending_pos + len]); + *pending_pos += len; + if *pending_pos >= pending.len() { + pending.clear(); + *pending_pos = 0; + } + true +} + fn replace_symbol(delimiter: &[u8], slice: &[u8]) -> Vec { + if delimiter.is_empty() { + return slice.to_vec(); + } + let mut result = Vec::with_capacity(slice.len()); let mut i = 0; while i < slice.len() { @@ -297,6 +696,148 @@ fn replace_symbol(delimiter: &[u8], slice: &[u8]) -> Vec { result } +fn convert_field_delimiter_stream(stream: S, delimiter: Option) -> BoxStream<'static, Result> +where + S: Stream> + Send + 'static, +{ + let Some(delimiter) = delimiter else { + return stream.boxed(); + }; + AsyncTryStream::::new(|mut y| async move { + let mut converter = DelimiterConverter::new(delimiter.into_bytes()); + pin_mut!(stream); + while let Some(result) = stream.next().await { + let bytes = result?; + let converted = converter.convert_chunk(&bytes); + if !converted.is_empty() { + y.yield_ok(Bytes::from(converted)).await; + } + } + let converted = converter.finish(); + if !converted.is_empty() { + y.yield_ok(Bytes::from(converted)).await; + } + Ok(()) + }) + .boxed() +} + +struct ScanRangeState { + stream: S, + delimiter: Vec, + range: SelectScanRange, + include_header: bool, + offset: u64, + record_start: u64, + record: Vec, + pending: VecDeque, + done: bool, +} + +fn scan_range_stream( + stream: S, + delimiter: Vec, + range: SelectScanRange, + include_header: bool, + base_offset: u64, +) -> BoxStream<'static, Result> +where + S: Stream> + Send + Unpin + 'static, +{ + let state = ScanRangeState { + stream, + delimiter, + range, + include_header, + offset: base_offset, + record_start: base_offset, + record: Vec::new(), + pending: VecDeque::new(), + done: false, + }; + + stream::unfold(state, |mut state| async move { + loop { + if let Some(bytes) = state.pending.pop_front() { + return Some((Ok(bytes), state)); + } + if state.done { + return None; + } + match state.stream.next().await { + Some(Ok(bytes)) => state.push_chunk(&bytes), + Some(Err(err)) => { + state.done = true; + return Some(( + Err(o_Error::Generic { + store: "EcObjectStore", + source: Box::new(err), + }), + state, + )); + } + None => { + state.finish_pending_record(); + state.done = true; + } + } + } + }) + .boxed() +} + +impl ScanRangeState { + fn push_chunk(&mut self, bytes: &[u8]) { + if bytes.is_empty() || self.done { + return; + } + if self.record.is_empty() { + self.record_start = self.offset; + } + let search_start = self.record.len().saturating_sub(self.delimiter.len().saturating_sub(1)); + self.record.extend_from_slice(bytes); + self.offset = self.offset.saturating_add(bytes.len() as u64); + + let mut search_start = search_start; + while let Some(pos) = find_delimiter(&self.record[search_start..], &self.delimiter) { + let record_end = search_start + pos + self.delimiter.len(); + self.finish_record(record_end); + if self.done { + break; + } + search_start = 0; + } + } + + fn finish_record(&mut self, record_end: usize) { + let record = self.record.drain(..record_end).collect::>(); + let record_start = self.record_start; + self.record_start = self.record_start.saturating_add(record_end as u64); + self.push_record(record, record_start); + } + + fn finish_pending_record(&mut self) { + if self.record.is_empty() { + return; + } + let record = std::mem::take(&mut self.record); + let record_start = self.record_start; + self.push_record(record, record_start); + } + + fn push_record(&mut self, record: Vec, record_start: u64) { + let include_header = self.include_header && record_start == 0; + let include_record = record_start >= self.range.start && record_start <= self.range.end; + if include_header || include_record { + self.pending.push_back(Bytes::from(record)); + } else { + if record_start > self.range.end { + self.done = true; + } + } + } +} + /// Extract the JSON sub-path from a SQL expression's FROM clause. /// /// Given `SELECT e.name FROM s3object.employees e WHERE …` this returns @@ -508,19 +1049,226 @@ where #[cfg(test)] mod test { - use super::{extract_json_sub_path_from_expression, flatten_json_document_to_ndjson, replace_symbol}; + use super::{ + ConvertStream, SelectScanRange, convert_field_delimiter_stream, extract_json_sub_path_from_expression, find_delimiter, + flatten_json_document_to_ndjson, http_range_spec_from_get_range, replace_symbol, scan_range_from_bounds, + scan_range_read_start, scan_range_stream, select_read_headers, + }; + use bytes::Bytes; + use futures::{StreamExt, stream}; + use object_store::GetRange; + use s3s::dto::{ + CSVInput, CSVOutput, ExpressionType, InputSerialization, OutputSerialization, SelectObjectContentInput, + SelectObjectContentRequest, + }; + use s3s::header::{ + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM, X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY, + X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5, + }; + use tokio::io::AsyncReadExt; + use tokio_util::io::StreamReader; #[test] fn test_replace() { - let ss = String::from("dandan&&is&&best"); - let slice = ss.as_bytes(); - let delimiter = b"&&"; - println!("len: {}", "╦".len()); - let result = replace_symbol(delimiter, slice); - match String::from_utf8(result) { - Ok(s) => println!("slice: {s}"), - Err(e) => eprintln!("Error converting to string: {e}"), + let result = replace_symbol(b"&&", b"dandan&&is&&best"); + assert_eq!(result, b"dandan,is,best"); + } + + #[tokio::test] + async fn test_convert_stream_replaces_delimiter_across_chunks() { + let chunks = stream::iter(vec![ + Ok::<_, std::io::Error>(Bytes::from_static(b"a&")), + Ok::<_, std::io::Error>(Bytes::from_static(b"&b&&c")), + ]); + let reader = StreamReader::new(chunks); + let mut reader = ConvertStream::new(reader, "&&".to_string()); + let mut output = Vec::new(); + reader.read_to_end(&mut output).await.unwrap(); + assert_eq!(output, b"a,b,c"); + } + + #[tokio::test] + async fn test_convert_stream_replaces_delimiter_at_stream_end() { + let chunks = stream::iter(vec![ + Ok::<_, std::io::Error>(Bytes::from_static(b"a&")), + Ok::<_, std::io::Error>(Bytes::from_static(b"&")), + ]); + let reader = StreamReader::new(chunks); + let mut reader = ConvertStream::new(reader, "&&".to_string()); + let mut output = Vec::new(); + reader.read_to_end(&mut output).await.unwrap(); + assert_eq!(output, b"a,"); + } + + #[tokio::test] + async fn test_scan_range_stream_keeps_header_and_selected_record() { + let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"h1,h2\n1,a\n2,b\n3,c\n"))]); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange { start: 10, end: 11 }, true, 0); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); } + assert_eq!(output, b"h1,h2\n2,b\n"); + } + + #[tokio::test] + async fn test_scan_range_stream_skips_record_when_start_is_in_middle() { + let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"1,a\n2,b\n3,c\n"))]); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange { start: 2, end: 7 }, false, 0); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"2,b\n"); + } + + #[tokio::test] + async fn test_scan_range_stream_keeps_record_when_end_is_in_middle() { + let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"1,a\n2,b\n3,c\n"))]); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange { start: 0, end: 5 }, false, 0); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"1,a\n2,b\n"); + } + + #[tokio::test] + async fn test_scan_range_stream_uses_base_offset_for_range_reader() { + let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"\n2,b\n3,c\n"))]); + let mut stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange { start: 4, end: 7 }, false, 3); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"2,b\n"); + } + + #[tokio::test] + async fn test_scan_range_stream_handles_delimiter_split_across_chunks() { + let chunks = stream::iter(vec![ + Ok::<_, std::io::Error>(Bytes::from_static(b"h1,h2\r")), + Ok::<_, std::io::Error>(Bytes::from_static(b"\n1,a\r\n2,b\r")), + Ok::<_, std::io::Error>(Bytes::from_static(b"\n3,c\r\n")), + ]); + let mut stream = scan_range_stream(chunks, b"\r\n".to_vec(), SelectScanRange { start: 12, end: 14 }, true, 0); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"h1,h2\r\n2,b\r\n"); + } + + #[test] + fn test_scan_range_read_start_keeps_full_delimiter_boundary() { + let range = SelectScanRange { start: 10, end: 20 }; + assert_eq!(scan_range_read_start(range, b"\n"), 9); + assert_eq!(scan_range_read_start(range, b"\r\n"), 8); + assert_eq!(scan_range_read_start(range, b"abcdef"), 4); + } + + #[test] + fn test_find_delimiter_handles_multi_byte_delimiter() { + assert_eq!(find_delimiter(b"one\r\ntwo", b"\r\n"), Some(3)); + assert_eq!(find_delimiter(b"one\ntwo", b"\r\n"), None); + } + + #[test] + fn test_scan_range_end_only_uses_aws_suffix_semantics() { + let range = scan_range_from_bounds(None, Some(35), 100).unwrap().unwrap(); + assert_eq!(range.start, 65); + assert_eq!(range.end, 99); + } + + #[test] + fn test_scan_range_start_after_object_is_rejected_before_reader() { + let err = scan_range_from_bounds(Some(100), None, 100).unwrap_err(); + assert!(err.to_string().contains("ScanRange")); + } + + #[test] + fn test_scan_range_start_after_end_is_rejected() { + let err = scan_range_from_bounds(Some(20), Some(10), 100).unwrap_err(); + assert!(err.to_string().contains("ScanRange")); + } + + #[test] + fn test_get_range_conversion_for_parquet_bounded_ranges() { + let range = http_range_spec_from_get_range(&GetRange::Bounded(10..20)); + assert!(!range.is_suffix_length); + assert_eq!(range.start, 10); + assert_eq!(range.end, 19); + } + + #[test] + fn test_select_read_headers_preserves_ssec_context() { + let input = SelectObjectContentInput { + bucket: "bucket".to_string(), + expected_bucket_owner: None, + key: "object.csv".to_string(), + sse_customer_algorithm: Some("AES256".to_string()), + sse_customer_key: Some("customer-key".to_string()), + sse_customer_key_md5: Some("customer-key-md5".to_string()), + request: SelectObjectContentRequest { + expression: "SELECT * FROM s3object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + csv: Some(CSVInput::default()), + ..Default::default() + }, + output_serialization: OutputSerialization { + csv: Some(CSVOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + }; + + let headers = select_read_headers(&input); + assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_ALGORITHM).unwrap(), "AES256"); + assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY).unwrap(), "customer-key"); + assert_eq!(headers.get(X_AMZ_SERVER_SIDE_ENCRYPTION_CUSTOMER_KEY_MD5).unwrap(), "customer-key-md5"); + } + + #[tokio::test] + async fn test_scan_range_output_can_convert_field_delimiter() { + let chunks = stream::iter(vec![Ok::<_, std::io::Error>(Bytes::from_static(b"a&&1\nb&&2\n"))]); + let stream = scan_range_stream(chunks, b"\n".to_vec(), SelectScanRange { start: 0, end: 10 }, false, 0); + let mut stream = convert_field_delimiter_stream(stream, Some("&&".to_string())); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"a,1\nb,2\n"); + } + + #[tokio::test] + async fn test_field_delimiter_stream_converts_delimiter_split_across_chunks() { + let chunks = stream::iter(vec![ + Ok::<_, object_store::Error>(Bytes::from_static(b"a&")), + Ok::<_, object_store::Error>(Bytes::from_static(b"&1\nb&&2\n")), + ]); + let mut stream = convert_field_delimiter_stream(chunks, Some("&&".to_string())); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"a,1\nb,2\n"); + } + + #[tokio::test] + async fn test_field_delimiter_stream_converts_delimiter_at_stream_end() { + let chunks = stream::iter(vec![ + Ok::<_, object_store::Error>(Bytes::from_static(b"a&")), + Ok::<_, object_store::Error>(Bytes::from_static(b"&")), + ]); + let mut stream = convert_field_delimiter_stream(chunks, Some("&&".to_string())); + let mut output = Vec::new(); + while let Some(bytes) = stream.next().await { + output.extend_from_slice(&bytes.unwrap()); + } + assert_eq!(output, b"a,"); } /// A JSON array is split into one NDJSON line per element. diff --git a/crates/s3select-api/src/query/execution.rs b/crates/s3select-api/src/query/execution.rs index 9a1148fba..7e4aa62d6 100644 --- a/crates/s3select-api/src/query/execution.rs +++ b/crates/s3select-api/src/query/execution.rs @@ -117,6 +117,15 @@ impl Output { } } + pub fn into_record_batch_stream(self) -> QueryResult { + match self { + Self::StreamData(stream) => Ok(stream), + Self::Nil(_) => Err(QueryError::NotImplemented { + err: "empty select output stream".to_string(), + }), + } + } + pub async fn num_rows(self) -> usize { match self.chunk_result().await { Ok(rb) => rb.iter().map(|e| e.num_rows()).sum(), diff --git a/crates/s3select-api/src/query/session.rs b/crates/s3select-api/src/query/session.rs index 73437d2b5..220be2c53 100644 --- a/crates/s3select-api/src/query/session.rs +++ b/crates/s3select-api/src/query/session.rs @@ -15,7 +15,13 @@ use crate::query::Context; use crate::{QueryError, QueryResult, object_store::EcObjectStore}; use datafusion::{ + arrow::{ + array::{Int32Array, StringArray}, + datatypes::{DataType, Field, Schema}, + record_batch::RecordBatch, + }, execution::{SessionStateBuilder, context::SessionState, runtime_env::RuntimeEnvBuilder}, + parquet::arrow::ArrowWriter, prelude::SessionContext, }; use object_store::{ObjectStore, ObjectStoreExt, memory::InMemory, path::Path}; @@ -66,7 +72,9 @@ impl SessionCtxFactory { let store: Arc = Arc::new(InMemory::new()); // Choose test data format based on what the request serialization specifies. - let data_bytes: &[u8] = if context.input.request.input_serialization.json.is_some() { + let data_bytes: Vec = if context.input.request.input_serialization.parquet.is_some() { + test_parquet_bytes()? + } else if context.input.request.input_serialization.json.is_some() { // NDJSON: one JSON object per line — usable for both LINES and DOCUMENT // requests (DOCUMENT inputs are converted to NDJSON by EcObjectStore, but // in test mode we bypass EcObjectStore, so we put NDJSON here directly). @@ -80,6 +88,7 @@ impl SessionCtxFactory { {\"id\":8,\"name\":\"Henry\",\"age\":32,\"department\":\"IT\",\"salary\":6200}\n\ {\"id\":9,\"name\":\"Ivy\",\"age\":24,\"department\":\"Marketing\",\"salary\":4800}\n\ {\"id\":10,\"name\":\"Jack\",\"age\":38,\"department\":\"Finance\",\"salary\":7500}\n" + .to_vec() } else { b"id,name,age,department,salary 1,Alice,25,HR,5000 @@ -92,6 +101,7 @@ impl SessionCtxFactory { 8,Henry,32,IT,6200 9,Ivy,24,Marketing,4800 10,Jack,38,Finance,7500" + .to_vec() }; let path = Path::from(context.input.key.clone()); @@ -112,3 +122,35 @@ impl SessionCtxFactory { Ok(df_session_ctx) } } + +fn test_parquet_bytes() -> QueryResult> { + let schema = Arc::new(Schema::new(vec![ + Field::new("id", DataType::Int32, false), + Field::new("name", DataType::Utf8, false), + Field::new("age", DataType::Int32, false), + Field::new("department", DataType::Utf8, false), + Field::new("salary", DataType::Int32, false), + ])); + let batch = RecordBatch::try_new( + schema.clone(), + vec![ + Arc::new(Int32Array::from(vec![1, 2, 3, 4, 5])), + Arc::new(StringArray::from(vec!["Alice", "Bob", "Charlie", "Diana", "Eve"])), + Arc::new(Int32Array::from(vec![25, 30, 35, 22, 28])), + Arc::new(StringArray::from(vec!["HR", "IT", "Finance", "Marketing", "IT"])), + Arc::new(Int32Array::from(vec![5000, 6000, 7000, 4500, 5500])), + ], + ) + .map_err(|e| QueryError::StoreError { e: e.to_string() })?; + + let mut bytes = Vec::new(); + { + let mut writer = + ArrowWriter::try_new(&mut bytes, schema, None).map_err(|e| QueryError::StoreError { e: e.to_string() })?; + writer + .write(&batch) + .map_err(|e| QueryError::StoreError { e: e.to_string() })?; + writer.close().map_err(|e| QueryError::StoreError { e: e.to_string() })?; + } + Ok(bytes) +} diff --git a/crates/s3select-query/src/dispatcher/manager.rs b/crates/s3select-query/src/dispatcher/manager.rs index 7fc1bfe33..46d977c1d 100644 --- a/crates/s3select-query/src/dispatcher/manager.rs +++ b/crates/s3select-query/src/dispatcher/manager.rs @@ -31,6 +31,7 @@ use datafusion::{ }, error::Result as DFResult, execution::{RecordBatchStream, SendableRecordBatchStream}, + sql::sqlparser::parser::ParserError, }; use futures::{Stream, StreamExt}; use rustfs_s3select_api::{ @@ -106,7 +107,11 @@ impl QueryDispatcher for SimpleQueryDispatcher { let stmt = match statements.front() { Some(stmt) => stmt.clone(), - None => return Ok(None), + None => { + return Err(QueryError::Parser { + source: ParserError::ParserError("empty SQL expression".to_string()), + }); + } }; let logical_plan = self diff --git a/crates/s3select-query/src/test/error_handling_test.rs b/crates/s3select-query/src/test/error_handling_test.rs index 289cdcc89..72a2144fe 100644 --- a/crates/s3select-query/src/test/error_handling_test.rs +++ b/crates/s3select-query/src/test/error_handling_test.rs @@ -178,15 +178,10 @@ mod error_handling_tests { let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); let result = db.execute(&query).await; - // Empty queries might be handled differently by the parser - match result { - Ok(_) => { - // Some parsers might accept empty queries - } - Err(_) => { - // Expected to fail for empty SQL - } - } + assert!( + matches!(result, Err(QueryError::Parser { .. })), + "Expected parser error for empty SQL: {sql:?}" + ); } } diff --git a/crates/s3select-query/src/test/integration_test.rs b/crates/s3select-query/src/test/integration_test.rs index e52353d8d..833b61de0 100644 --- a/crates/s3select-query/src/test/integration_test.rs +++ b/crates/s3select-query/src/test/integration_test.rs @@ -21,7 +21,7 @@ mod integration_tests { }; use s3s::dto::{ CSVInput, CSVOutput, ExpressionType, FileHeaderInfo, InputSerialization, JSONInput, JSONOutput, JSONType, - OutputSerialization, SelectObjectContentInput, SelectObjectContentRequest, + OutputSerialization, ParquetInput, SelectObjectContentInput, SelectObjectContentRequest, }; use std::sync::Arc; @@ -83,6 +83,31 @@ mod integration_tests { } } + fn create_test_parquet_input(sql: &str) -> SelectObjectContentInput { + SelectObjectContentInput { + bucket: "test-bucket".to_string(), + expected_bucket_owner: None, + key: "test.parquet".to_string(), + sse_customer_algorithm: None, + sse_customer_key: None, + sse_customer_key_md5: None, + request: SelectObjectContentRequest { + expression: sql.to_string(), + expression_type: ExpressionType::from_static("SQL"), + input_serialization: InputSerialization { + parquet: Some(ParquetInput {}), + ..Default::default() + }, + output_serialization: OutputSerialization { + json: Some(JSONOutput::default()), + ..Default::default() + }, + request_progress: None, + scan_range: None, + }, + } + } + #[tokio::test] async fn test_database_creation() { let input = create_test_input("SELECT * FROM S3Object"); @@ -290,6 +315,21 @@ mod integration_tests { assert!(output.is_ok()); } + #[tokio::test] + async fn test_simple_select_query_parquet() { + let sql = "SELECT name, age FROM S3Object WHERE age > 25"; + let input = create_test_parquet_input(sql); + let db = get_global_db(input.clone(), true).await.unwrap(); + let query = Query::new(Context { input: Arc::new(input) }, sql.to_string()); + + let result = db.execute(&query).await; + assert!(result.is_ok()); + + let output = result.unwrap().result().chunk_result().await.unwrap(); + let total_rows: usize = output.iter().map(|batch| batch.num_rows()).sum(); + assert_eq!(total_rows, 3); + } + #[tokio::test] async fn test_select_with_where_clause_json() { let sql = "SELECT name, age FROM S3Object WHERE age > 30"; diff --git a/rustfs/src/app/mod.rs b/rustfs/src/app/mod.rs index 3084c1090..27e62f66b 100644 --- a/rustfs/src/app/mod.rs +++ b/rustfs/src/app/mod.rs @@ -20,6 +20,7 @@ pub mod bucket_usecase; pub mod context; pub mod multipart_usecase; pub mod object_usecase; +mod select_object; #[cfg(test)] mod capacity_dirty_scope_test; diff --git a/rustfs/src/app/object_usecase.rs b/rustfs/src/app/object_usecase.rs index 8b72016f3..267ff93be 100644 --- a/rustfs/src/app/object_usecase.rs +++ b/rustfs/src/app/object_usecase.rs @@ -35,9 +35,6 @@ use crate::storage::sse::{SSEType, build_ssec_read_headers, encryption_material_ use crate::storage::timeout_wrapper::{GetObjectTimeoutPolicy, RequestTimeoutWrapper}; use crate::storage::*; use bytes::Bytes; -use datafusion::arrow::{ - csv::WriterBuilder as CsvWriterBuilder, json::WriterBuilder as JsonWriterBuilder, json::writer::JsonArray, -}; use futures::StreamExt; use http::{HeaderMap, HeaderValue, StatusCode}; use md5::Context as Md5Context; @@ -87,11 +84,7 @@ use rustfs_notify::EventArgsBuilder; use rustfs_policy::policy::action::{Action, S3Action}; use rustfs_rio::{CompressReader, DynReader, EncryptReader, HashReader, wrap_reader}; use rustfs_s3_ops::{S3Operation, delete_event_name_for_marker, put_event_name_for_post_object}; -use rustfs_s3select_api::{ - object_store::bytes_stream, - query::{Context, Query}, -}; -use rustfs_s3select_query::get_global_db; +use rustfs_s3select_api::object_store::bytes_stream; use rustfs_targets::{ EventName, extract_params_header, extract_resp_elements, get_request_host, get_request_port, get_request_user_agent, }; @@ -126,8 +119,6 @@ use std::time::Duration; use time::{OffsetDateTime, format_description::well_known::Rfc3339}; use tokio::io::{AsyncRead, ReadBuf}; use tokio::sync::RwLock; -use tokio::sync::mpsc; -use tokio_stream::wrappers::ReceiverStream; use tokio_tar::Archive; use tokio_util::io::{ReaderStream, StreamReader}; use tracing::{debug, error, info, instrument, warn}; @@ -3949,74 +3940,7 @@ impl DefaultObjectUsecase { let _ = context.object_store(); } - info!("handle select_object_content"); - - let input = Arc::new(req.input); - info!("{:?}", input); - - let db = get_global_db((*input).clone(), false).await.map_err(|e| { - error!("get global db failed, {}", e.to_string()); - s3_error!(InternalError, "{}", e.to_string()) - })?; - let query = Query::new(Context { input: input.clone() }, input.request.expression.clone()); - let result = db - .execute(&query) - .await - .map_err(|e| s3_error!(InternalError, "{}", e.to_string()))?; - - let results = result - .result() - .chunk_result() - .await - .map_err(|e| s3_error!(InternalError, "{}", e.to_string()))? - .to_vec(); - - let mut buffer = Vec::new(); - if input.request.output_serialization.csv.is_some() { - let mut csv_writer = CsvWriterBuilder::new().with_header(false).build(&mut buffer); - for batch in results { - csv_writer - .write(&batch) - .map_err(|e| s3_error!(InternalError, "can't encode output to csv. e: {}", e.to_string()))?; - } - } else if input.request.output_serialization.json.is_some() { - let mut json_writer = JsonWriterBuilder::new() - .with_explicit_nulls(true) - .build::<_, JsonArray>(&mut buffer); - for batch in results { - json_writer - .write(&batch) - .map_err(|e| s3_error!(InternalError, "can't encode output to json. e: {}", e.to_string()))?; - } - json_writer - .finish() - .map_err(|e| s3_error!(InternalError, "writer output into json error, e: {}", e.to_string()))?; - } else { - return Err(s3_error!( - InvalidArgument, - "Unsupported output format. Supported formats are CSV and JSON" - )); - } - - let (tx, rx) = mpsc::channel::>(2); - let stream = ReceiverStream::new(rx); - spawn_traced(async move { - let _ = tx - .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) - .await; - let _ = tx - .send(Ok(SelectObjectContentEvent::Records(RecordsEvent { - payload: Some(Bytes::from(buffer)), - }))) - .await; - let _ = tx.send(Ok(SelectObjectContentEvent::End(EndEvent::default()))).await; - - drop(tx); - }); - - Ok(S3Response::new(SelectObjectContentOutput { - payload: Some(SelectObjectContentEventStream::new(stream)), - })) + crate::app::select_object::execute_select_object_content(req).await } #[instrument(level = "debug", skip(self, req))] diff --git a/rustfs/src/app/select_object.rs b/rustfs/src/app/select_object.rs new file mode 100644 index 000000000..3edac6f78 --- /dev/null +++ b/rustfs/src/app/select_object.rs @@ -0,0 +1,797 @@ +use crate::error::ApiError; +use crate::storage::options::get_opts; +use crate::storage::request_context::spawn_traced; +use crate::storage::{get_validated_store, validate_sse_headers_for_read, validate_ssec_for_read}; +use bytes::Bytes; +use datafusion::arrow::{ + csv::{QuoteStyle, WriterBuilder as CsvWriterBuilder, writer::Terminator}, + json::{WriterBuilder as JsonWriterBuilder, writer::LineDelimited}, + record_batch::RecordBatch, +}; +use futures::StreamExt; +use http::{StatusCode, header::RANGE}; +use rustfs_ecstore::store_api::ObjectOperations; +use rustfs_s3select_api::{ + QueryError, + object_store::{INVALID_SCAN_RANGE_MESSAGE, validate_scan_range_bounds}, + query::{Context, Query}, +}; +use rustfs_s3select_query::get_global_db; +use s3s::dto::*; +use s3s::{S3Error, S3ErrorCode, S3Request, S3Response, S3Result, s3_error}; +use std::sync::Arc; +use tokio::sync::mpsc; +use tokio_stream::wrappers::ReceiverStream; +use tracing::info; + +const MAX_SELECT_EXPRESSION_BYTES: usize = 256 * 1024; +const RECORDS_CHUNK_TARGET: usize = 128 * 1024; +const PARSE_SELECT_FAILURE_CODE: &str = "ParseSelectFailure"; +const EMPTY_SELECT_EXPRESSION_MESSAGE: &str = "empty SQL expression"; + +#[derive(Clone, Debug)] +struct SelectValidation { + output_format: SelectOutputFormat, + progress_enabled: bool, +} + +struct SelectObjectMetadata { + size: u64, +} + +#[derive(Clone, Debug)] +enum SelectOutputFormat { + Csv(CSVOutput), + Json(JSONOutput), +} + +pub async fn execute_select_object_content( + req: S3Request, +) -> S3Result> { + let mut input = req.input; + let validation = validate_select_request(&req.headers, &mut input)?; + log_select_request_summary(&input, &validation); + let metadata = preflight_select_object(&req.headers, &input).await?; + validate_scan_range_for_object_size(&input.request, metadata.size)?; + + let input = Arc::new(input); + let db = get_global_db((*input).clone(), false).await.map_err(map_query_error_to_s3)?; + let query = Query::new(Context { input: input.clone() }, input.request.expression.clone()); + let output = db + .execute(&query) + .await + .map_err(map_query_error_to_s3)? + .result() + .into_record_batch_stream() + .map_err(map_query_error_to_s3)?; + + let (tx, rx) = mpsc::channel::>(8); + spawn_traced(async move { + let mut encoder = SelectOutputEncoder::new(validation.output_format); + let mut progress = SelectProgress::default(); + let mut output = output; + + if tx + .send(Ok(SelectObjectContentEvent::Cont(ContinuationEvent::default()))) + .await + .is_err() + { + return; + } + + while let Some(result) = output.next().await { + let batch = match result { + Ok(batch) => batch, + Err(err) => { + let _ = tx.send(Err(map_query_error_to_s3(err.into()))).await; + return; + } + }; + + match encoder.encode_batch(&batch) { + Ok(payloads) => { + for payload in payloads { + progress.add_returned(payload.len()); + if tx + .send(Ok(SelectObjectContentEvent::Records(RecordsEvent { payload: Some(payload) }))) + .await + .is_err() + { + return; + } + if validation.progress_enabled + && tx + .send(Ok(SelectObjectContentEvent::Progress(ProgressEvent { + details: Some(progress.to_progress()), + }))) + .await + .is_err() + { + return; + } + } + } + Err(err) => { + let _ = tx.send(Err(err)).await; + return; + } + } + } + + let stats = SelectObjectContentEvent::Stats(StatsEvent { + details: Some(progress.to_stats()), + }); + if tx.send(Ok(stats)).await.is_err() { + return; + } + let _ = tx.send(Ok(SelectObjectContentEvent::End(EndEvent::default()))).await; + }); + + Ok(S3Response::new(SelectObjectContentOutput { + payload: Some(SelectObjectContentEventStream::new(ReceiverStream::new(rx))), + })) +} + +fn validate_select_request(headers: &http::HeaderMap, input: &mut SelectObjectContentInput) -> S3Result { + if headers.contains_key(RANGE) { + return Err(S3Error::new(S3ErrorCode::UnsupportedRangeHeader)); + } + if input.request.expression.len() > MAX_SELECT_EXPRESSION_BYTES { + return Err(S3Error::new(S3ErrorCode::ExpressionTooLong)); + } + if input.request.expression_type.as_str() != ExpressionType::SQL { + return Err(S3Error::new(S3ErrorCode::InvalidExpressionType)); + } + + normalize_input_serialization(&mut input.request.input_serialization)?; + validate_scan_range(&input.request)?; + + let output_format = normalize_output_serialization(&mut input.request.output_serialization)?; + if input.request.expression.trim().is_empty() { + return Err(parse_select_failure(EMPTY_SELECT_EXPRESSION_MESSAGE)); + } + let progress_enabled = input + .request + .request_progress + .as_ref() + .and_then(|progress| progress.enabled) + .unwrap_or(false); + + Ok(SelectValidation { + output_format, + progress_enabled, + }) +} + +fn normalize_input_serialization(input: &mut InputSerialization) -> S3Result<()> { + let format_count = + usize::from(input.csv.is_some()) + usize::from(input.json.is_some()) + usize::from(input.parquet.is_some()); + if format_count == 0 { + return Err(S3Error::new(S3ErrorCode::MissingRequiredParameter)); + } + if format_count > 1 { + return Err(S3Error::new(S3ErrorCode::ObjectSerializationConflict)); + } + + if let Some(compression) = input.compression_type.as_ref() + && compression.as_str() != CompressionType::NONE + { + return Err(s3_error!( + NotImplemented, + "SelectObjectContent currently supports only uncompressed input" + )); + } + input.compression_type = Some(CompressionType::from_static(CompressionType::NONE)); + + if let Some(csv) = input.csv.as_mut() { + if csv.allow_quoted_record_delimiter.unwrap_or(false) { + return Err(s3_error!( + NotImplemented, + "CSV AllowQuotedRecordDelimiter is not supported by SelectObjectContent" + )); + } + csv.file_header_info + .get_or_insert_with(|| FileHeaderInfo::from_static(FileHeaderInfo::NONE)); + validate_single_byte(csv.comments.as_deref(), S3ErrorCode::InvalidRequestParameter)?; + validate_single_byte(csv.quote_character.as_deref(), S3ErrorCode::InvalidRequestParameter)?; + validate_single_byte(csv.quote_escape_character.as_deref(), S3ErrorCode::InvalidRequestParameter)?; + validate_record_delimiter(csv.record_delimiter.as_deref())?; + } + + if let Some(json) = input.json.as_mut() { + let json_type = json.type_.get_or_insert_with(|| JSONType::from_static(JSONType::LINES)); + if !matches!(json_type.as_str(), JSONType::DOCUMENT | JSONType::LINES) { + return Err(S3Error::new(S3ErrorCode::InvalidJsonType)); + } + } + + Ok(()) +} + +fn normalize_output_serialization(output: &mut OutputSerialization) -> S3Result { + let format_count = usize::from(output.csv.is_some()) + usize::from(output.json.is_some()); + if format_count == 0 { + return Err(S3Error::new(S3ErrorCode::MissingRequiredParameter)); + } + if format_count > 1 { + return Err(S3Error::new(S3ErrorCode::ObjectSerializationConflict)); + } + + if let Some(csv) = output.csv.as_ref() { + validate_single_byte(csv.field_delimiter.as_deref(), S3ErrorCode::InvalidRequestParameter)?; + validate_single_byte(csv.quote_character.as_deref(), S3ErrorCode::InvalidRequestParameter)?; + validate_single_byte(csv.quote_escape_character.as_deref(), S3ErrorCode::InvalidRequestParameter)?; + validate_record_delimiter(csv.record_delimiter.as_deref())?; + if let Some(quote_fields) = csv.quote_fields.as_ref() + && !matches!(quote_fields.as_str(), QuoteFields::ALWAYS | QuoteFields::ASNEEDED) + { + return Err(S3Error::new(S3ErrorCode::InvalidQuoteFields)); + } + return Ok(SelectOutputFormat::Csv(csv.clone())); + } + + let json = output.json.as_ref().expect("checked exactly one output format"); + Ok(SelectOutputFormat::Json(json.clone())) +} + +fn validate_scan_range(request: &SelectObjectContentRequest) -> S3Result<()> { + let Some(scan_range) = request.scan_range.as_ref() else { + return Ok(()); + }; + let start = scan_range.start; + let end = scan_range.end; + if start.is_none() && end.is_none() { + return Err(invalid_scan_range_error()); + } + if validate_scan_range_bounds(start, end, u64::MAX).is_err() { + return Err(invalid_scan_range_error()); + } + if request.input_serialization.parquet.is_some() || request.input_serialization.json.as_ref().is_some_and(is_json_document) { + return Err(invalid_scan_range_error()); + } + Ok(()) +} + +fn validate_scan_range_for_object_size(request: &SelectObjectContentRequest, object_size: u64) -> S3Result<()> { + let Some(scan_range) = request.scan_range.as_ref() else { + return Ok(()); + }; + if validate_scan_range_bounds(scan_range.start, scan_range.end, object_size).is_err() { + return Err(invalid_scan_range_error()); + } + Ok(()) +} + +fn invalid_scan_range_error() -> S3Error { + S3Error::with_message(S3ErrorCode::InvalidRequestParameter, INVALID_SCAN_RANGE_MESSAGE.to_string()) +} + +fn parse_select_failure(message: impl Into) -> S3Error { + let mut err = S3Error::with_message(S3ErrorCode::Custom(PARSE_SELECT_FAILURE_CODE.into()), message.into()); + err.set_status_code(StatusCode::BAD_REQUEST); + err +} + +fn validate_single_byte(value: Option<&str>, code: S3ErrorCode) -> S3Result<()> { + if let Some(value) = value + && value.len() != 1 + { + return Err(S3Error::new(code)); + } + Ok(()) +} + +fn validate_record_delimiter(value: Option<&str>) -> S3Result<()> { + if let Some(value) = value + && value.len() != 1 + && value != "\r\n" + { + return Err(S3Error::new(S3ErrorCode::InvalidRequestParameter)); + } + Ok(()) +} + +async fn preflight_select_object(headers: &http::HeaderMap, input: &SelectObjectContentInput) -> S3Result { + let opts = get_opts(&input.bucket, &input.key, None, None, headers) + .await + .map_err(ApiError::from)?; + let store = get_validated_store(&input.bucket).await?; + let info = store + .get_object_info(&input.bucket, &input.key, &opts) + .await + .map_err(ApiError::from)?; + validate_sse_headers_for_read(&info.user_defined, headers)?; + validate_ssec_for_read(&info.user_defined, input.sse_customer_key.as_ref(), input.sse_customer_key_md5.as_ref())?; + Ok(SelectObjectMetadata { + size: info.size.max(0) as u64, + }) +} + +fn log_select_request_summary(input: &SelectObjectContentInput, validation: &SelectValidation) { + let output_format = match &validation.output_format { + SelectOutputFormat::Csv(_) => "csv", + SelectOutputFormat::Json(_) => "json", + }; + let input_format = if input.request.input_serialization.csv.is_some() { + "csv" + } else if input.request.input_serialization.json.is_some() { + "json" + } else { + "parquet" + }; + info!( + bucket = %input.bucket, + key = %input.key, + expression_len = input.request.expression.len(), + input_format, + output_format, + has_scan_range = input.request.scan_range.is_some(), + has_sse_customer_key = input.sse_customer_key.is_some(), + "handle select_object_content" + ); +} + +struct SelectOutputEncoder { + format: SelectOutputFormat, +} + +impl SelectOutputEncoder { + fn new(format: SelectOutputFormat) -> Self { + 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(|err| s3_error!(InternalError, "can't encode Select output to CSV: {}", err))?; + 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(|err| s3_error!(InternalError, "can't encode Select output to JSON: {}", err))?; + writer + .finish() + .map_err(|err| s3_error!(InternalError, "can't finish Select JSON output: {}", err))?; + 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); + } + } + output +} + +fn split_records_payload(bytes: Vec) -> Vec { + if bytes.is_empty() { + return Vec::new(); + } + let bytes = Bytes::from(bytes); + if bytes.len() <= RECORDS_CHUNK_TARGET { + return vec![bytes]; + } + (0..bytes.len()) + .step_by(RECORDS_CHUNK_TARGET) + .map(|start| bytes.slice(start..(start + RECORDS_CHUNK_TARGET).min(bytes.len()))) + .collect() +} + +#[derive(Default)] +struct SelectProgress { + bytes_returned: u64, +} + +impl SelectProgress { + fn add_returned(&mut self, bytes: usize) { + self.bytes_returned = self.bytes_returned.saturating_add(bytes as u64); + } + + fn to_progress(&self) -> Progress { + Progress { + bytes_processed: None, + bytes_returned: Some(clamp_i64(self.bytes_returned)), + bytes_scanned: None, + } + } + + fn to_stats(&self) -> Stats { + Stats { + bytes_processed: None, + bytes_returned: Some(clamp_i64(self.bytes_returned)), + bytes_scanned: None, + } + } +} + +fn clamp_i64(value: u64) -> i64 { + value.min(i64::MAX as u64) as i64 +} + +fn map_query_error_to_s3(err: QueryError) -> S3Error { + let message = err.to_string(); + match err { + QueryError::Parser { .. } => parse_select_failure(message), + QueryError::MultiStatement { .. } => S3Error::with_message(S3ErrorCode::UnsupportedSqlStructure, message), + QueryError::NotImplemented { .. } => S3Error::with_message(S3ErrorCode::NotImplemented, message), + QueryError::Datafusion { .. } if looks_like_invalid_scan_range(&message) => { + S3Error::with_message(S3ErrorCode::InvalidRequestParameter, INVALID_SCAN_RANGE_MESSAGE.to_string()) + } + QueryError::Datafusion { .. } if looks_like_missing_binding(&message) => { + S3Error::with_message(S3ErrorCode::EvaluatorBindingDoesNotExist, message) + } + QueryError::Datafusion { .. } => S3Error::with_message(S3ErrorCode::UnsupportedSqlOperation, message), + QueryError::StoreError { .. } if looks_like_invalid_scan_range(&message) => { + S3Error::with_message(S3ErrorCode::InvalidRequestParameter, INVALID_SCAN_RANGE_MESSAGE.to_string()) + } + QueryError::StoreError { .. } if looks_like_bucket_not_found(&message) => { + S3Error::with_message(S3ErrorCode::NoSuchBucket, message) + } + QueryError::StoreError { .. } if looks_like_object_not_found(&message) => { + S3Error::with_message(S3ErrorCode::NoSuchKey, message) + } + QueryError::StoreError { .. } => S3Error::with_message(S3ErrorCode::InternalError, message), + QueryError::BuildQueryDispatcher { .. } + | QueryError::Cancel + | QueryError::FunctionNotExists { .. } + | QueryError::FunctionExists { .. } => S3Error::with_message(S3ErrorCode::InternalError, message), + } +} + +fn looks_like_bucket_not_found(message: &str) -> bool { + message.contains("NoSuchBucket") || message.contains("bucket not found") || message.contains("BucketNotFound") +} + +fn looks_like_object_not_found(message: &str) -> bool { + message.contains("NoSuchKey") + || message.contains("NoSuchVersion") + || message.contains("ObjectNotFound") + || message.contains("object not found") + || message.contains("NotFound") +} + +fn looks_like_missing_binding(message: &str) -> bool { + message.contains("No field named") + || message.contains("field not found") + || message.contains("Schema error") + || message.contains("No such column") +} + +fn looks_like_invalid_scan_range(message: &str) -> bool { + message.contains("ScanRange:") || message.contains(INVALID_SCAN_RANGE_MESSAGE) +} + +fn is_json_document(json: &JSONInput) -> bool { + json.type_ + .as_ref() + .is_some_and(|json_type| json_type.as_str() == JSONType::DOCUMENT) +} + +#[cfg(test)] +mod tests { + use super::*; + use datafusion::sql::sqlparser::parser::ParserError; + use http::HeaderMap; + + fn base_input() -> SelectObjectContentInput { + SelectObjectContentInput { + bucket: "bucket".to_string(), + expected_bucket_owner: None, + key: "object.csv".to_string(), + sse_customer_algorithm: None, + sse_customer_key: None, + sse_customer_key_md5: None, + request: SelectObjectContentRequest { + expression: "SELECT * FROM s3object".to_string(), + expression_type: ExpressionType::from_static(ExpressionType::SQL), + input_serialization: InputSerialization { + csv: Some(CSVInput::default()), + compression_type: None, + json: None, + parquet: None, + }, + output_serialization: OutputSerialization { + csv: Some(CSVOutput::default()), + json: None, + }, + request_progress: None, + scan_range: None, + }, + } + } + + #[test] + fn validate_rejects_http_range() { + let mut input = base_input(); + let mut headers = HeaderMap::new(); + headers.insert(RANGE, "bytes=0-1".parse().unwrap()); + let err = validate_select_request(&headers, &mut input).unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::UnsupportedRangeHeader); + } + + #[test] + fn validate_rejects_empty_select_expression_as_parse_failure() { + for expression in ["", " \t\n"] { + let mut input = base_input(); + input.request.expression = expression.to_string(); + + let err = validate_select_request(&HeaderMap::new(), &mut input).unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::Custom("ParseSelectFailure".into())); + assert_eq!(err.status_code(), Some(http::StatusCode::BAD_REQUEST)); + } + } + + #[test] + fn map_parser_error_to_parse_select_failure() { + let err = map_query_error_to_s3(QueryError::Parser { + source: ParserError::ParserError("syntax error".to_string()), + }); + + assert_eq!(err.code(), &S3ErrorCode::Custom("ParseSelectFailure".into())); + assert_eq!(err.status_code(), Some(http::StatusCode::BAD_REQUEST)); + assert_eq!(err.message(), Some("sql parser error: syntax error")); + } + + #[test] + fn validate_defaults_csv_header_and_compression() { + let mut input = base_input(); + let validation = validate_select_request(&HeaderMap::new(), &mut input).unwrap(); + assert!(matches!(validation.output_format, SelectOutputFormat::Csv(_))); + assert_eq!( + input + .request + .input_serialization + .csv + .as_ref() + .and_then(|csv| csv.file_header_info.as_ref()) + .map(|value| value.as_str()), + Some(FileHeaderInfo::NONE) + ); + assert_eq!( + input + .request + .input_serialization + .compression_type + .as_ref() + .map(|value| value.as_str()), + Some(CompressionType::NONE) + ); + } + + #[test] + fn json_encoder_outputs_line_delimited_records() { + let schema = + std::sync::Arc::new(datafusion::arrow::datatypes::Schema::new(vec![datafusion::arrow::datatypes::Field::new( + "name", + datafusion::arrow::datatypes::DataType::Utf8, + false, + )])); + let batch = RecordBatch::try_new( + schema, + vec![std::sync::Arc::new(datafusion::arrow::array::StringArray::from(vec![ + "a", "b", + ]))], + ) + .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"); + } + + #[test] + fn json_encoder_honors_custom_record_delimiter() { + let schema = + std::sync::Arc::new(datafusion::arrow::datatypes::Schema::new(vec![datafusion::arrow::datatypes::Field::new( + "name", + datafusion::arrow::datatypes::DataType::Utf8, + false, + )])); + let batch = RecordBatch::try_new( + schema, + vec![std::sync::Arc::new(datafusion::arrow::array::StringArray::from(vec![ + "a", "b", + ]))], + ) + .unwrap(); + + let bytes = encode_json_batch( + &batch, + &JSONOutput { + record_delimiter: Some("|".to_string()), + }, + ) + .unwrap(); + let output = String::from_utf8(bytes).unwrap(); + assert_eq!(output, "{\"name\":\"a\"}|{\"name\":\"b\"}|"); + } + + #[test] + fn csv_encoder_honors_output_delimiters() { + let schema = std::sync::Arc::new(datafusion::arrow::datatypes::Schema::new(vec![ + datafusion::arrow::datatypes::Field::new("name", datafusion::arrow::datatypes::DataType::Utf8, false), + datafusion::arrow::datatypes::Field::new("score", datafusion::arrow::datatypes::DataType::Int32, false), + ])); + let batch = RecordBatch::try_new( + schema, + vec![ + std::sync::Arc::new(datafusion::arrow::array::StringArray::from(vec!["a", "b"])), + std::sync::Arc::new(datafusion::arrow::array::Int32Array::from(vec![1, 2])), + ], + ) + .unwrap(); + + let bytes = encode_csv_batch( + &batch, + &CSVOutput { + field_delimiter: Some("|".to_string()), + record_delimiter: Some("\r\n".to_string()), + ..Default::default() + }, + ) + .unwrap(); + assert_eq!(String::from_utf8(bytes).unwrap(), "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]); + let mut progress = SelectProgress::default(); + for payload in &payloads { + progress.add_returned(payload.len()); + } + assert_eq!(progress.to_stats().bytes_returned, Some((RECORDS_CHUNK_TARGET + 7) as i64)); + assert!(payloads.len() > 1); + } + + #[test] + fn validate_rejects_scan_range_for_json_document_as_request_parameter() { + let mut input = base_input(); + input.request.input_serialization = InputSerialization { + csv: None, + json: Some(JSONInput { + type_: Some(JSONType::from_static(JSONType::DOCUMENT)), + }), + parquet: None, + compression_type: None, + }; + input.request.scan_range = Some(ScanRange { + start: Some(0), + end: Some(10), + }); + + let err = validate_select_request(&HeaderMap::new(), &mut input).unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::InvalidRequestParameter); + assert_eq!(err.message(), Some(INVALID_SCAN_RANGE_MESSAGE)); + } + + #[test] + fn validate_rejects_scan_range_start_after_object() { + let mut input = base_input(); + input.request.scan_range = Some(ScanRange { + start: Some(10), + end: None, + }); + + validate_select_request(&HeaderMap::new(), &mut input).unwrap(); + let err = validate_scan_range_for_object_size(&input.request, 10).unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::InvalidRequestParameter); + assert_eq!(err.message(), Some(INVALID_SCAN_RANGE_MESSAGE)); + } + + #[test] + fn validate_rejects_scan_range_start_after_end() { + let mut input = base_input(); + input.request.scan_range = Some(ScanRange { + start: Some(20), + end: Some(10), + }); + + let err = validate_select_request(&HeaderMap::new(), &mut input).unwrap_err(); + assert_eq!(err.code(), &S3ErrorCode::InvalidRequestParameter); + assert_eq!(err.message(), Some(INVALID_SCAN_RANGE_MESSAGE)); + } + + #[test] + fn validate_allows_scan_range_end_only_suffix_form() { + let mut input = base_input(); + input.request.scan_range = Some(ScanRange { + start: None, + end: Some(35), + }); + + validate_select_request(&HeaderMap::new(), &mut input).unwrap(); + validate_scan_range_for_object_size(&input.request, 10).unwrap(); + } + + #[test] + fn progress_does_not_report_unknown_input_bytes_as_zero() { + let mut progress = SelectProgress::default(); + progress.add_returned(12); + let stats = progress.to_stats(); + assert_eq!(stats.bytes_returned, Some(12)); + assert_eq!(stats.bytes_scanned, None); + assert_eq!(stats.bytes_processed, None); + } + + #[test] + fn map_store_error_not_found_to_no_such_key() { + let err = map_query_error_to_s3(QueryError::StoreError { + e: "ObjectStore NotFound: bucket/object.csv".to_string(), + }); + assert_eq!(err.code(), &S3ErrorCode::NoSuchKey); + } + + #[test] + fn map_store_error_bucket_not_found_to_no_such_bucket() { + let err = map_query_error_to_s3(QueryError::StoreError { + e: "bucket not found".to_string(), + }); + assert_eq!(err.code(), &S3ErrorCode::NoSuchBucket); + } + + #[test] + fn map_scan_range_store_error_to_invalid_request_parameter() { + let err = map_query_error_to_s3(QueryError::StoreError { + e: "ScanRange: Start after EOF".to_string(), + }); + assert_eq!(err.code(), &S3ErrorCode::InvalidRequestParameter); + assert_eq!(err.message(), Some(INVALID_SCAN_RANGE_MESSAGE)); + } +}