mirror of
https://github.com/rustfs/rustfs.git
synced 2026-08-31 01:09:23 +00:00
feat(s3select): expand typed JSON source paths (#6864)
This commit is contained in:
@@ -27,6 +27,9 @@ use std::time::Duration;
|
||||
const BUCKET: &str = "test-sql-bucket";
|
||||
const CSV_OBJECT: &str = "test-data.csv";
|
||||
const JSON_OBJECT: &str = "test-data.json";
|
||||
const JSON_DOCUMENT_OBJECT: &str = "nested-data.json";
|
||||
const JSON_ROOT_ARRAY_OBJECT: &str = "root-array.json";
|
||||
const JSON_ROOT_SCALAR_ARRAY_OBJECT: &str = "root-scalars.json";
|
||||
const SELECT_RESPONSE_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
type TestResult<T> = Result<T, Box<dyn Error + Send + Sync>>;
|
||||
@@ -74,6 +77,51 @@ async fn upload_test_json(client: &Client) -> TestResult<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn upload_nested_json_document(client: &Client) -> TestResult<()> {
|
||||
let json_data = r#"{"departments":[{"employees":[{"name":"Alice","active":true},{"name":"Bob","active":false}]},{"employees":[{"name":"Charlie","active":true}]}]}"#;
|
||||
|
||||
client
|
||||
.put_object()
|
||||
.bucket(BUCKET)
|
||||
.key(JSON_DOCUMENT_OBJECT)
|
||||
.body(Bytes::from_static(json_data.as_bytes()).into())
|
||||
.send()
|
||||
.await?;
|
||||
client
|
||||
.put_object()
|
||||
.bucket(BUCKET)
|
||||
.key(JSON_ROOT_ARRAY_OBJECT)
|
||||
.body(Bytes::from_static(br#"[{"name":"Alice"},{"name":"Bob"}]"#).into())
|
||||
.send()
|
||||
.await?;
|
||||
client
|
||||
.put_object()
|
||||
.bucket(BUCKET)
|
||||
.key(JSON_ROOT_SCALAR_ARRAY_OBJECT)
|
||||
.body(Bytes::from_static(b"[1,2]").into())
|
||||
.send()
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn select_json_document(client: &Client, key: &str, expression: &str) -> TestResult<String> {
|
||||
let response = client
|
||||
.select_object_content()
|
||||
.bucket(BUCKET)
|
||||
.key(key)
|
||||
.expression(expression)
|
||||
.expression_type(ExpressionType::Sql)
|
||||
.input_serialization(
|
||||
InputSerialization::builder()
|
||||
.json(JsonInput::builder().set_type(Some(JsonType::Document)).build())
|
||||
.build(),
|
||||
)
|
||||
.output_serialization(OutputSerialization::builder().json(JsonOutput::builder().build()).build())
|
||||
.send()
|
||||
.await?;
|
||||
process_select_response(response).await
|
||||
}
|
||||
|
||||
async fn process_select_response(
|
||||
mut event_stream: aws_sdk_s3::operation::select_object_content::SelectObjectContentOutput,
|
||||
) -> TestResult<String> {
|
||||
@@ -365,6 +413,107 @@ async fn test_select_object_content_json_basic() -> TestResult<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
||||
async fn test_select_object_content_nested_json_source_path() -> TestResult<()> {
|
||||
let (_env, client) = create_test_environment().await?;
|
||||
setup_test_bucket(&client).await?;
|
||||
upload_nested_json_document(&client).await?;
|
||||
|
||||
let result = select_json_document(
|
||||
&client,
|
||||
JSON_DOCUMENT_OBJECT,
|
||||
"SELECT e.name FROM S3Object[*].departments[*].employees[*] AS e WHERE e.active = true",
|
||||
)
|
||||
.await?;
|
||||
let names: Vec<String> = result
|
||||
.lines()
|
||||
.filter(|line| !line.trim().is_empty())
|
||||
.map(|line| -> TestResult<String> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["name"].as_str().ok_or("missing name field")?.to_string())
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
|
||||
assert_eq!(names, vec!["Alice", "Charlie"]);
|
||||
|
||||
let terminal_scalars = select_json_document(
|
||||
&client,
|
||||
JSON_DOCUMENT_OBJECT,
|
||||
"SELECT NAME FROM S3Object[*].DEPARTMENTS[*].employees[*].NAME",
|
||||
)
|
||||
.await?;
|
||||
let scalar_names: Vec<String> = terminal_scalars
|
||||
.lines()
|
||||
.map(|line| -> TestResult<String> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["name"].as_str().ok_or("missing scalar name field")?.to_string())
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
assert_eq!(scalar_names, vec!["Alice", "Bob", "Charlie"]);
|
||||
|
||||
let aliased_scalars = select_json_document(
|
||||
&client,
|
||||
JSON_DOCUMENT_OBJECT,
|
||||
"SELECT v FROM S3Object[*].departments[*].employees[*].name AS v",
|
||||
)
|
||||
.await?;
|
||||
let aliased_names: Vec<String> = aliased_scalars
|
||||
.lines()
|
||||
.map(|line| -> TestResult<String> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["v"].as_str().ok_or("missing aliased scalar field")?.to_string())
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
assert_eq!(aliased_names, vec!["Alice", "Bob", "Charlie"]);
|
||||
|
||||
let root_array = select_json_document(&client, JSON_ROOT_ARRAY_OBJECT, "SELECT c.name FROM S3Object[*][*] AS c").await?;
|
||||
let root_names: Vec<String> = root_array
|
||||
.lines()
|
||||
.map(|line| -> TestResult<String> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["name"].as_str().ok_or("missing root-array name field")?.to_string())
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
assert_eq!(root_names, vec!["Alice", "Bob"]);
|
||||
|
||||
let root_index = select_json_document(&client, JSON_ROOT_ARRAY_OBJECT, "SELECT c.name FROM S3Object[*][0] AS c").await?;
|
||||
let root_index_value: serde_json::Value = serde_json::from_str(root_index.trim())?;
|
||||
assert_eq!(root_index_value["name"], "Alice");
|
||||
|
||||
let root_scalars = select_json_document(&client, JSON_ROOT_SCALAR_ARRAY_OBJECT, "SELECT V FROM S3Object AS V").await?;
|
||||
let scalar_values: Vec<i64> = root_scalars
|
||||
.lines()
|
||||
.map(|line| -> TestResult<i64> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["v"].as_i64().ok_or("missing root scalar value")?)
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
assert_eq!(scalar_values, vec![1, 2]);
|
||||
|
||||
let implicit_root_scalars =
|
||||
select_json_document(&client, JSON_ROOT_SCALAR_ARRAY_OBJECT, "SELECT S3Object FROM S3Object").await?;
|
||||
let implicit_scalar_values: Vec<i64> = implicit_root_scalars
|
||||
.lines()
|
||||
.map(|line| -> TestResult<i64> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["s3object"].as_i64().ok_or("missing implicit root scalar value")?)
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
assert_eq!(implicit_scalar_values, vec![1, 2]);
|
||||
|
||||
let quoted_root_scalars =
|
||||
select_json_document(&client, JSON_ROOT_SCALAR_ARRAY_OBJECT, "SELECT \"S3Object\" FROM \"S3Object\"").await?;
|
||||
let quoted_scalar_values: Vec<i64> = quoted_root_scalars
|
||||
.lines()
|
||||
.map(|line| -> TestResult<i64> {
|
||||
let value: serde_json::Value = serde_json::from_str(line)?;
|
||||
Ok(value["S3Object"].as_i64().ok_or("missing quoted root scalar value")?)
|
||||
})
|
||||
.collect::<TestResult<_>>()?;
|
||||
assert_eq!(quoted_scalar_values, vec![1, 2]);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
|
||||
async fn test_select_object_content_csv_limit() -> TestResult<()> {
|
||||
let (_env, client) = create_test_environment().await?;
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -13,6 +13,43 @@
|
||||
// limitations under the License.
|
||||
|
||||
use datafusion::sql::sqlparser::ast::Statement;
|
||||
use std::sync::Arc;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum JsonPathSegment {
|
||||
Key { name: String, quoted: bool },
|
||||
Index(usize),
|
||||
ArrayWildcard,
|
||||
ObjectWildcard,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
pub struct JsonSource {
|
||||
path: Arc<[JsonPathSegment]>,
|
||||
scalar_column: Option<String>,
|
||||
}
|
||||
|
||||
impl JsonSource {
|
||||
pub fn new(path: Vec<JsonPathSegment>, scalar_column: Option<String>) -> Self {
|
||||
Self {
|
||||
path: path.into(),
|
||||
scalar_column,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn from_path(path: Vec<JsonPathSegment>) -> Self {
|
||||
Self::new(path, None)
|
||||
}
|
||||
|
||||
pub fn path(&self) -> &[JsonPathSegment] {
|
||||
&self.path
|
||||
}
|
||||
|
||||
pub fn scalar_column(&self) -> Option<&str> {
|
||||
self.scalar_column.as_deref()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum ExtStatement {
|
||||
|
||||
@@ -29,6 +29,7 @@ use tracing::debug;
|
||||
use crate::{QueryError, QueryResult};
|
||||
|
||||
use super::Query;
|
||||
use super::ast::ExtStatement;
|
||||
use super::logical_planner::Plan;
|
||||
use super::session::{QueryExecutionTracker, SessionCtx};
|
||||
|
||||
@@ -172,6 +173,7 @@ pub struct QueryStateMachine {
|
||||
pub session: SessionCtx,
|
||||
pub query: Query,
|
||||
|
||||
prepared_statement: Option<ExtStatement>,
|
||||
query_tracker: Option<QueryExecutionTracker>,
|
||||
state: RwLock<QueryState>,
|
||||
start: Instant,
|
||||
@@ -196,6 +198,7 @@ impl QueryStateMachine {
|
||||
Self {
|
||||
session,
|
||||
query,
|
||||
prepared_statement: None,
|
||||
query_tracker: None,
|
||||
state: RwLock::new(QueryState::ACCEPTING),
|
||||
start: Instant::now(),
|
||||
@@ -211,6 +214,21 @@ impl QueryStateMachine {
|
||||
Ok(state_machine)
|
||||
}
|
||||
|
||||
pub fn begin_tracked_prepared(
|
||||
query: Query,
|
||||
session: SessionCtx,
|
||||
query_tracker: QueryExecutionTracker,
|
||||
prepared_statement: ExtStatement,
|
||||
) -> QueryResult<Self> {
|
||||
let mut state_machine = Self::begin_tracked(query, session, query_tracker)?;
|
||||
state_machine.prepared_statement = Some(prepared_statement);
|
||||
Ok(state_machine)
|
||||
}
|
||||
|
||||
pub fn prepared_statement(&self) -> Option<&ExtStatement> {
|
||||
self.prepared_statement.as_ref()
|
||||
}
|
||||
|
||||
pub fn query_tracker(&self) -> Option<&QueryExecutionTracker> {
|
||||
self.query_tracker.as_ref()
|
||||
}
|
||||
|
||||
@@ -34,6 +34,10 @@ impl Dialect for RustFsDialect {
|
||||
fn supports_group_by_expr(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn supports_partiql(&self) -> bool {
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
pub trait Parser {
|
||||
|
||||
@@ -12,9 +12,11 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
use crate::query::{Context, Query};
|
||||
use crate::{QueryError, QueryResult, object_store::EcObjectStore};
|
||||
use crate::{SelectInputMetrics, SelectObjectSnapshot};
|
||||
use crate::query::{Context, Query, ast::JsonSource};
|
||||
use crate::{
|
||||
QueryError, QueryResult, SelectInputMetrics, SelectObjectSnapshot,
|
||||
object_store::{EcObjectStore, is_json_document_input, legacy_json_source_from_input},
|
||||
};
|
||||
use datafusion::{
|
||||
arrow::{
|
||||
array::{Int32Array, StringArray},
|
||||
@@ -314,8 +316,15 @@ impl SessionCtxFactory {
|
||||
}
|
||||
|
||||
pub async fn create_session_ctx(&self, context: &Context) -> QueryResult<SessionCtx> {
|
||||
self.create_session_ctx_inner(context, None, None, None, DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES)
|
||||
.await
|
||||
self.create_session_ctx_inner(
|
||||
context,
|
||||
None,
|
||||
legacy_json_source_from_input(&context.input),
|
||||
None,
|
||||
None,
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_session_ctx_with_tracker_and_memory_limit(
|
||||
@@ -324,8 +333,15 @@ impl SessionCtxFactory {
|
||||
query_tracker: QueryExecutionTracker,
|
||||
memory_limit_bytes: usize,
|
||||
) -> QueryResult<SessionCtx> {
|
||||
self.create_session_ctx_inner(context, None, Some(query_tracker), None, memory_limit_bytes)
|
||||
.await
|
||||
self.create_session_ctx_inner(
|
||||
context,
|
||||
None,
|
||||
legacy_json_source_from_input(&context.input),
|
||||
Some(query_tracker),
|
||||
None,
|
||||
memory_limit_bytes,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
|
||||
@@ -335,8 +351,33 @@ impl SessionCtxFactory {
|
||||
query_tracker: QueryExecutionTracker,
|
||||
memory_limit_bytes: usize,
|
||||
) -> QueryResult<SessionCtx> {
|
||||
self.create_session_ctx_inner(context, Some(snapshot), Some(query_tracker), None, memory_limit_bytes)
|
||||
.await
|
||||
self.create_session_ctx_inner(
|
||||
context,
|
||||
Some(snapshot),
|
||||
legacy_json_source_from_input(&context.input),
|
||||
Some(query_tracker),
|
||||
None,
|
||||
memory_limit_bytes,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_session_ctx_for_query_with_source_and_tracker_and_memory_limit(
|
||||
&self,
|
||||
query: &Query,
|
||||
source: JsonSource,
|
||||
query_tracker: QueryExecutionTracker,
|
||||
memory_limit_bytes: usize,
|
||||
) -> QueryResult<SessionCtx> {
|
||||
self.create_session_ctx_inner(
|
||||
query.context(),
|
||||
query.snapshot().cloned(),
|
||||
source,
|
||||
Some(query_tracker),
|
||||
Some(Arc::clone(query.input_metrics())),
|
||||
memory_limit_bytes,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn create_session_ctx_for_query_with_tracker_and_memory_limit(
|
||||
@@ -348,6 +389,7 @@ impl SessionCtxFactory {
|
||||
self.create_session_ctx_inner(
|
||||
query.context(),
|
||||
query.snapshot().cloned(),
|
||||
legacy_json_source_from_input(&query.context().input),
|
||||
Some(query_tracker),
|
||||
Some(Arc::clone(query.input_metrics())),
|
||||
memory_limit_bytes,
|
||||
@@ -359,12 +401,13 @@ impl SessionCtxFactory {
|
||||
&self,
|
||||
context: &Context,
|
||||
snapshot: Option<Arc<SelectObjectSnapshot>>,
|
||||
source: JsonSource,
|
||||
query_tracker: Option<QueryExecutionTracker>,
|
||||
input_metrics: Option<Arc<SelectInputMetrics>>,
|
||||
memory_limit_bytes: usize,
|
||||
) -> QueryResult<SessionCtx> {
|
||||
let df_session_ctx = self
|
||||
.build_df_session_context(context, snapshot, query_tracker.clone(), input_metrics, memory_limit_bytes)
|
||||
.build_df_session_context(context, snapshot, source, query_tracker.clone(), input_metrics, memory_limit_bytes)
|
||||
.await?;
|
||||
|
||||
Ok(SessionCtx {
|
||||
@@ -378,6 +421,7 @@ impl SessionCtxFactory {
|
||||
&self,
|
||||
context: &Context,
|
||||
snapshot: Option<Arc<SelectObjectSnapshot>>,
|
||||
source: JsonSource,
|
||||
query_tracker: Option<QueryExecutionTracker>,
|
||||
input_metrics: Option<Arc<SelectInputMetrics>>,
|
||||
memory_limit_bytes: usize,
|
||||
@@ -401,10 +445,12 @@ impl SessionCtxFactory {
|
||||
.is_some_and(|delimiter| delimiter.len() == 2 && delimiter.as_bytes() != b"\r\n");
|
||||
let scan_range_requires_single_file_scan =
|
||||
context.input.request.scan_range.is_some() && context.input.request.input_serialization.parquet.is_none();
|
||||
let json_document_requires_single_file_scan = is_json_document_input(&context.input);
|
||||
let metered_input_requires_single_file_scan =
|
||||
input_metrics.is_some() && context.input.request.input_serialization.parquet.is_none();
|
||||
let config = if custom_two_byte_record_delimiter
|
||||
|| scan_range_requires_single_file_scan
|
||||
|| json_document_requires_single_file_scan
|
||||
|| metered_input_requires_single_file_scan
|
||||
{
|
||||
config.with_repartition_file_scans(false)
|
||||
@@ -463,14 +509,21 @@ impl SessionCtxFactory {
|
||||
} else {
|
||||
let input_metrics = input_metrics.unwrap_or_else(|| Arc::new(SelectInputMetrics::default()));
|
||||
let store: EcObjectStore = match query_tracker {
|
||||
Some(query_tracker) => EcObjectStore::new_with_query_tracker(
|
||||
Some(query_tracker) => EcObjectStore::new_with_query_tracker_and_source(
|
||||
context.input.clone(),
|
||||
memory_pool,
|
||||
query_tracker,
|
||||
input_metrics,
|
||||
snapshot,
|
||||
source,
|
||||
),
|
||||
None => EcObjectStore::new_with_memory_pool_and_source(
|
||||
context.input.clone(),
|
||||
memory_pool,
|
||||
input_metrics,
|
||||
snapshot,
|
||||
source,
|
||||
),
|
||||
None => EcObjectStore::new_with_memory_pool(context.input.clone(), memory_pool, input_metrics, snapshot),
|
||||
}
|
||||
.map_err(|err| QueryError::Datafusion {
|
||||
source: Box::new(DataFusionError::External(Box::new(err))),
|
||||
@@ -543,15 +596,15 @@ mod tests {
|
||||
use crate::storage_api::object_store::ObjectIO as _;
|
||||
use datafusion::{
|
||||
datasource::{
|
||||
file_format::csv::CsvFormat,
|
||||
file_format::{csv::CsvFormat, json::JsonFormat},
|
||||
listing::{ListingOptions, ListingTable, ListingTableConfig, ListingTableUrl},
|
||||
},
|
||||
execution::memory_pool::MemoryLimit,
|
||||
};
|
||||
use http::HeaderMap;
|
||||
use s3s::dto::{
|
||||
CSVInput, CSVOutput, ExpressionType, InputSerialization, JSONInput, OutputSerialization, ParquetInput, ScanRange,
|
||||
SelectObjectContentInput, SelectObjectContentRequest,
|
||||
CSVInput, CSVOutput, ExpressionType, InputSerialization, JSONInput, JSONType, OutputSerialization, ParquetInput,
|
||||
ScanRange, SelectObjectContentInput, SelectObjectContentRequest,
|
||||
};
|
||||
use std::io::Write as _;
|
||||
|
||||
@@ -592,6 +645,103 @@ mod tests {
|
||||
)
|
||||
}
|
||||
|
||||
async fn test_query_tracker() -> QueryExecutionTracker {
|
||||
let permit = Arc::new(tokio::sync::Semaphore::new(1))
|
||||
.acquire_owned()
|
||||
.await
|
||||
.expect("query permit should be available");
|
||||
QueryExecutionTracker::new(
|
||||
&QueryExecutionOwner::new(),
|
||||
Arc::new(permit),
|
||||
Instant::now() + std::time::Duration::from_secs(300),
|
||||
300,
|
||||
)
|
||||
}
|
||||
|
||||
async fn assert_legacy_json_column(bucket: &str, expression: &str, document: &[u8], column_name: &str, expected: &[&str]) {
|
||||
const OBJECT: &str = "input.json";
|
||||
|
||||
let env = crate::storage_api::select_test_ecstore_env().await;
|
||||
let mut context = test_context();
|
||||
{
|
||||
let input = Arc::make_mut(&mut context.input);
|
||||
input.bucket = bucket.to_string();
|
||||
input.key = OBJECT.to_string();
|
||||
input.request.expression = expression.to_string();
|
||||
input.request.input_serialization.csv = None;
|
||||
input.request.input_serialization.json = Some(JSONInput {
|
||||
type_: Some(JSONType::from_static(JSONType::DOCUMENT)),
|
||||
});
|
||||
}
|
||||
env.make_bucket(bucket, false).await;
|
||||
env.put_object_bytes(bucket, OBJECT, document.to_vec()).await;
|
||||
let factory = SessionCtxFactory::new(false);
|
||||
let lazy_session = factory
|
||||
.create_session_ctx(&context)
|
||||
.await
|
||||
.expect("legacy lazy session should preserve the JSON source");
|
||||
let tracked_lazy_session = factory
|
||||
.create_session_ctx_with_tracker_and_memory_limit(
|
||||
&context,
|
||||
test_query_tracker().await,
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
|
||||
)
|
||||
.await
|
||||
.expect("legacy tracked lazy session should preserve the JSON source");
|
||||
let snapshot = prepare_test_snapshot(&context).await;
|
||||
let snapshot_session = factory
|
||||
.create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
|
||||
&context,
|
||||
snapshot,
|
||||
test_query_tracker().await,
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
|
||||
)
|
||||
.await
|
||||
.expect("legacy snapshot session should preserve the JSON source");
|
||||
|
||||
for (kind, session) in [
|
||||
("lazy", lazy_session),
|
||||
("tracked lazy", tracked_lazy_session),
|
||||
("tracked snapshot", snapshot_session),
|
||||
] {
|
||||
let table_path = ListingTableUrl::parse(format!("s3://{bucket}/{OBJECT}")).expect("parse JSON table URL");
|
||||
let listing_options = ListingOptions::new(Arc::new(JsonFormat::default())).with_file_extension(".json");
|
||||
let schema = listing_options
|
||||
.infer_schema(session.inner(), &table_path)
|
||||
.await
|
||||
.expect("infer expanded JSON schema");
|
||||
let table = ListingTable::try_new(
|
||||
ListingTableConfig::new(table_path)
|
||||
.with_listing_options(listing_options)
|
||||
.with_schema(schema),
|
||||
)
|
||||
.expect("build expanded JSON table");
|
||||
let query_context = SessionContext::new_with_state(session.inner().clone());
|
||||
query_context
|
||||
.register_table("legacy_input", Arc::new(table))
|
||||
.expect("register expanded JSON table");
|
||||
let batches = query_context
|
||||
.sql(&format!("SELECT {column_name} FROM legacy_input"))
|
||||
.await
|
||||
.expect("plan expanded JSON query")
|
||||
.collect()
|
||||
.await
|
||||
.expect("execute expanded JSON query");
|
||||
let mut values = Vec::new();
|
||||
for batch in batches {
|
||||
let column = batch
|
||||
.column(0)
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.expect("expanded column should be Utf8");
|
||||
for row in 0..batch.num_rows() {
|
||||
values.push(column.value(row).to_string());
|
||||
}
|
||||
}
|
||||
assert_eq!(values, expected, "{kind} legacy constructor");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn session_factory_fields_remain_source_compatible() {
|
||||
let factory = SessionCtxFactory {
|
||||
@@ -629,6 +779,7 @@ mod tests {
|
||||
.create_session_ctx_inner(
|
||||
context,
|
||||
None,
|
||||
JsonSource::default(),
|
||||
None,
|
||||
Some(Arc::new(SelectInputMetrics::default())),
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
|
||||
@@ -680,6 +831,40 @@ mod tests {
|
||||
assert!(!session.inner().config().options().optimizer.repartition_file_scans);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn json_lines_without_scan_range_keeps_file_repartitioning() {
|
||||
let mut context = test_context();
|
||||
let request = &mut Arc::make_mut(&mut context.input).request;
|
||||
request.input_serialization.csv = None;
|
||||
request.input_serialization.json = Some(JSONInput::default());
|
||||
|
||||
let session = SessionCtxFactory::new(true)
|
||||
.with_target_partitions(2)
|
||||
.create_session_ctx(&context)
|
||||
.await
|
||||
.expect("JSON LINES session should be created");
|
||||
|
||||
assert!(session.inner().config().options().optimizer.repartition_file_scans);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn json_document_disables_file_repartitioning() {
|
||||
let mut context = test_context();
|
||||
let request = &mut Arc::make_mut(&mut context.input).request;
|
||||
request.input_serialization.csv = None;
|
||||
request.input_serialization.json = Some(JSONInput {
|
||||
type_: Some(JSONType::from_static(JSONType::DOCUMENT)),
|
||||
});
|
||||
|
||||
let session = SessionCtxFactory::new(true)
|
||||
.with_target_partitions(2)
|
||||
.create_session_ctx(&context)
|
||||
.await
|
||||
.expect("JSON DOCUMENT session should be created");
|
||||
|
||||
assert!(!session.inner().config().options().optimizer.repartition_file_scans);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn csv_scan_range_disables_file_repartitioning() {
|
||||
let mut context = test_context();
|
||||
@@ -755,7 +940,7 @@ mod tests {
|
||||
async fn session_factory_applies_memory_limit() {
|
||||
let factory = SessionCtxFactory::new(true);
|
||||
let session = factory
|
||||
.create_session_ctx_inner(&test_context(), None, None, None, 1024)
|
||||
.create_session_ctx_inner(&test_context(), None, JsonSource::default(), None, None, 1024)
|
||||
.await
|
||||
.expect("session should be created with a bounded memory pool");
|
||||
|
||||
@@ -803,6 +988,45 @@ mod tests {
|
||||
assert!(session.is_bound_to(&tracker));
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial_test::serial]
|
||||
async fn legacy_session_factory_preserves_single_key_json_source() {
|
||||
assert_legacy_json_column(
|
||||
"s3select-legacy-session-json-source",
|
||||
"SELECT e.name FROM S3Object.employees AS e",
|
||||
br#"{"employees":[{"name":"Alice"}]}"#,
|
||||
"name",
|
||||
&["Alice"],
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial_test::serial]
|
||||
async fn legacy_session_factory_preserves_implicit_root_alias() {
|
||||
assert_legacy_json_column(
|
||||
"s3select-legacy-session-root-alias",
|
||||
"SELECT S3Object FROM S3Object",
|
||||
br#"["one","two"]"#,
|
||||
"s3object",
|
||||
&["one", "two"],
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial_test::serial]
|
||||
async fn legacy_session_factory_preserves_quoted_root_alias() {
|
||||
assert_legacy_json_column(
|
||||
"s3select-legacy-session-quoted-root-alias",
|
||||
"SELECT \"V\" FROM S3Object AS \"V\"",
|
||||
br#"["one","two"]"#,
|
||||
"\"V\"",
|
||||
&["one", "two"],
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
#[serial_test::serial]
|
||||
async fn session_factory_propagates_query_guard_to_ec_store() {
|
||||
|
||||
@@ -41,11 +41,11 @@ use rustfs_s3select_api::{
|
||||
QueryError, QueryResult, SelectError,
|
||||
query::{
|
||||
Query,
|
||||
ast::ExtStatement,
|
||||
ast::{ExtStatement, JsonPathSegment, JsonSource},
|
||||
dispatcher::{DispatchedQuery, QueryDispatcher},
|
||||
execution::{Output, QueryStateMachine},
|
||||
function::FuncMetaManagerRef,
|
||||
logical_planner::{LogicalPlanner, Plan},
|
||||
logical_planner::Plan,
|
||||
parser::Parser,
|
||||
session::{
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES, QueryAdmission, QueryExecutionOwner, QueryExecutionStatus,
|
||||
@@ -53,7 +53,7 @@ use rustfs_s3select_api::{
|
||||
},
|
||||
},
|
||||
};
|
||||
use s3s::dto::{FileHeaderInfo, SelectObjectContentInput};
|
||||
use s3s::dto::{FileHeaderInfo, JSONType, SelectObjectContentInput};
|
||||
use std::sync::LazyLock;
|
||||
use tokio::{
|
||||
sync::Semaphore,
|
||||
@@ -66,6 +66,7 @@ use crate::{
|
||||
instance::{DEFAULT_MAX_CONCURRENT_QUERIES, DEFAULT_QUERY_TIMEOUT_SECS},
|
||||
metadata::{ContextProviderExtension, MetadataProvider, TableHandleProviderRef, base_table::BaseTableProvider},
|
||||
sql::logical::planner::DefaultLogicalPlanner,
|
||||
sql::planner::prepare_s3_select_statement,
|
||||
};
|
||||
|
||||
static IGNORE: LazyLock<FileHeaderInfo> = LazyLock::new(|| FileHeaderInfo::from_static(FileHeaderInfo::IGNORE));
|
||||
@@ -160,25 +161,18 @@ impl QueryDispatcher for SimpleQueryDispatcher {
|
||||
let session = &query_state_machine.session;
|
||||
let query = &query_state_machine.query;
|
||||
|
||||
let scheme_provider = self.build_scheme_provider(session).await?;
|
||||
let logical_planner = DefaultLogicalPlanner::new(&scheme_provider);
|
||||
let statements = self.parser.parse(query.content())?;
|
||||
|
||||
if statements.len() > 1 {
|
||||
return Err(QueryError::MultiStatement {
|
||||
num: statements.len(),
|
||||
sql: query_state_machine.query.content().to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
let stmt = match statements.front() {
|
||||
Some(stmt) => stmt.clone(),
|
||||
let stmt = match query_state_machine.prepared_statement() {
|
||||
Some(statement) => statement.clone(),
|
||||
None => {
|
||||
return Err(QueryError::Parser {
|
||||
source: ParserError::ParserError("empty SQL expression".to_string()),
|
||||
});
|
||||
let (statement, source) = self.prepare_query_statement(query.content())?;
|
||||
if source_path_requires_expansion(source.path()) {
|
||||
return Err(SelectError::DataSourcePathUnsupported.into());
|
||||
}
|
||||
statement
|
||||
}
|
||||
};
|
||||
let scheme_provider = self.build_scheme_provider(session).await?;
|
||||
let logical_planner = DefaultLogicalPlanner::new(&scheme_provider);
|
||||
|
||||
let logical_plan = self
|
||||
.statement_to_logical_plan(stmt, &logical_planner, Arc::clone(&query_state_machine))
|
||||
@@ -269,12 +263,20 @@ impl SimpleQueryDispatcher {
|
||||
self.query_timeout.as_secs(),
|
||||
);
|
||||
let phase_guard = QueryPhaseGuard::new(&query_tracker, &self.query_execution_owner);
|
||||
// Keep parser and analyzer errors in the planning phase. Successful
|
||||
// preparation is cached here because the source path configures the
|
||||
// object store before schema inference starts.
|
||||
let (prepared_statement, source) = match self.prepare_query_statement(query.content()) {
|
||||
Ok((statement, source)) => (Some(statement), source),
|
||||
Err(_) => (None, JsonSource::default()),
|
||||
};
|
||||
let session = self
|
||||
.run_with_query_deadline(
|
||||
&query_tracker,
|
||||
self.session_factory
|
||||
.create_session_ctx_for_query_with_tracker_and_memory_limit(
|
||||
.create_session_ctx_for_query_with_source_and_tracker_and_memory_limit(
|
||||
&query,
|
||||
source,
|
||||
query_tracker.clone(),
|
||||
self.memory_limit_bytes,
|
||||
),
|
||||
@@ -285,7 +287,28 @@ impl SimpleQueryDispatcher {
|
||||
return Err(self.query_tracker_error(&query_tracker));
|
||||
}
|
||||
phase_guard.disarm();
|
||||
Ok(Arc::new(QueryStateMachine::begin_tracked(query, session, query_tracker)?))
|
||||
let state_machine = match prepared_statement {
|
||||
Some(statement) => QueryStateMachine::begin_tracked_prepared(query, session, query_tracker, statement)?,
|
||||
None => QueryStateMachine::begin_tracked(query, session, query_tracker)?,
|
||||
};
|
||||
Ok(Arc::new(state_machine))
|
||||
}
|
||||
|
||||
fn prepare_query_statement(&self, sql: &str) -> QueryResult<(ExtStatement, JsonSource)> {
|
||||
let mut statements = self.parser.parse(sql)?;
|
||||
if statements.len() > 1 {
|
||||
return Err(QueryError::MultiStatement {
|
||||
num: statements.len(),
|
||||
sql: sql.to_string(),
|
||||
});
|
||||
}
|
||||
let mut statement = statements.pop_front().ok_or_else(|| QueryError::Parser {
|
||||
source: ParserError::ParserError("empty SQL expression".to_string()),
|
||||
})?;
|
||||
let ExtStatement::SqlStatement(sql_statement) = &mut statement;
|
||||
let source = prepare_s3_select_statement(sql_statement)?;
|
||||
validate_json_source_path_input(&self.input, source.path())?;
|
||||
Ok((statement, source))
|
||||
}
|
||||
async fn run_with_query_deadline<T>(
|
||||
&self,
|
||||
@@ -365,7 +388,7 @@ impl SimpleQueryDispatcher {
|
||||
// begin analyze
|
||||
query_state_machine.begin_analyze();
|
||||
let logical_plan = logical_planner
|
||||
.create_logical_plan(stmt, &query_state_machine.session)
|
||||
.prepared_statement_to_plan(stmt, &query_state_machine.session)
|
||||
.await?;
|
||||
query_state_machine.end_analyze();
|
||||
|
||||
@@ -503,6 +526,31 @@ impl SimpleQueryDispatcher {
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_json_source_path_input(input: &SelectObjectContentInput, source_path: &[JsonPathSegment]) -> QueryResult<()> {
|
||||
if source_path.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let Some(json) = input.request.input_serialization.json.as_ref() else {
|
||||
return Err(SelectError::DataSourcePathUnsupported.into());
|
||||
};
|
||||
if !source_path_requires_expansion(source_path)
|
||||
|| json
|
||||
.type_
|
||||
.as_ref()
|
||||
.is_some_and(|json_type| json_type.as_str() == JSONType::DOCUMENT)
|
||||
{
|
||||
return Ok(());
|
||||
}
|
||||
Err(SelectError::DataSourcePathUnsupported.into())
|
||||
}
|
||||
|
||||
fn source_path_requires_expansion(source_path: &[JsonPathSegment]) -> bool {
|
||||
!source_path
|
||||
.strip_prefix(&[JsonPathSegment::ArrayWildcard])
|
||||
.unwrap_or(source_path)
|
||||
.is_empty()
|
||||
}
|
||||
|
||||
pub struct TrackedRecordBatchStream {
|
||||
state: Arc<TrackedRecordBatchState>,
|
||||
schema: SchemaRef,
|
||||
@@ -760,7 +808,10 @@ impl SimpleQueryDispatcherBuilder {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{QueryPhaseGuard, SimpleQueryDispatcher, SimpleQueryDispatcherBuilder, TrackedRecordBatchStream};
|
||||
use super::{
|
||||
QueryPhaseGuard, SimpleQueryDispatcher, SimpleQueryDispatcherBuilder, TrackedRecordBatchStream,
|
||||
validate_json_source_path_input,
|
||||
};
|
||||
use crate::{
|
||||
execution::{
|
||||
factory::{QueryExecutionFactoryRef, SqlQueryExecutionFactory},
|
||||
@@ -789,6 +840,7 @@ mod tests {
|
||||
QueryError, QueryResult, SelectError,
|
||||
query::{
|
||||
Context as QueryContext, Query,
|
||||
ast::JsonPathSegment,
|
||||
dispatcher::QueryDispatcher,
|
||||
execution::{
|
||||
Output, QueryExecution, QueryExecutionFactory, QueryExecutionRef, QueryStateMachine, QueryStateMachineRef,
|
||||
@@ -1031,7 +1083,17 @@ mod tests {
|
||||
query_execution_factory: QueryExecutionFactoryRef,
|
||||
) -> (Arc<SimpleQueryDispatcher>, Arc<SelectObjectContentInput>) {
|
||||
let input = Arc::new(test_input());
|
||||
let dispatcher = SimpleQueryDispatcherBuilder::default()
|
||||
let dispatcher = test_dispatcher_for_input(Arc::clone(&input), admission, query_timeout, query_execution_factory);
|
||||
(dispatcher, input)
|
||||
}
|
||||
|
||||
fn test_dispatcher_for_input(
|
||||
input: Arc<SelectObjectContentInput>,
|
||||
admission: Arc<Semaphore>,
|
||||
query_timeout: Duration,
|
||||
query_execution_factory: QueryExecutionFactoryRef,
|
||||
) -> Arc<SimpleQueryDispatcher> {
|
||||
SimpleQueryDispatcherBuilder::default()
|
||||
.with_input(Arc::clone(&input))
|
||||
.with_default_table_provider(Arc::new(BaseTableProvider::default()))
|
||||
.with_session_factory(Arc::new(SessionCtxFactory::new(true)))
|
||||
@@ -1041,8 +1103,7 @@ mod tests {
|
||||
.with_query_admission(admission)
|
||||
.with_query_timeout(query_timeout)
|
||||
.build()
|
||||
.expect("query dispatcher should build");
|
||||
(dispatcher, input)
|
||||
.expect("query dispatcher should build")
|
||||
}
|
||||
|
||||
async fn snapshot_test_env() -> &'static TestECStoreEnv {
|
||||
@@ -1171,6 +1232,68 @@ mod tests {
|
||||
})
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn nested_source_paths_require_json_document_input() {
|
||||
let lines_input = json_snapshot_input();
|
||||
let nested_path = [JsonPathSegment::Key {
|
||||
name: "employees".to_string(),
|
||||
quoted: false,
|
||||
}];
|
||||
|
||||
assert!(matches!(
|
||||
validate_json_source_path_input(&lines_input, &nested_path),
|
||||
Err(ref error) if matches!(error.s3_select_policy_error(), Some(SelectError::DataSourcePathUnsupported))
|
||||
));
|
||||
assert!(validate_json_source_path_input(&lines_input, &[JsonPathSegment::ArrayWildcard]).is_ok());
|
||||
|
||||
let csv_input = test_input();
|
||||
let parquet_input = parquet_snapshot_input();
|
||||
for input in [&csv_input, parquet_input.as_ref()] {
|
||||
assert!(matches!(
|
||||
validate_json_source_path_input(input, &nested_path),
|
||||
Err(ref error) if matches!(error.s3_select_policy_error(), Some(SelectError::DataSourcePathUnsupported))
|
||||
));
|
||||
}
|
||||
|
||||
let mut document_input = (*lines_input).clone();
|
||||
document_input.request.input_serialization.json.as_mut().unwrap().type_ = Some(JSONType::from_static(JSONType::DOCUMENT));
|
||||
assert!(validate_json_source_path_input(&document_input, &nested_path).is_ok());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn normal_planning_pipeline_rejects_json_lines_source_expansion() {
|
||||
let mut input = (*json_snapshot_input()).clone();
|
||||
input.request.expression = "SELECT * FROM S3Object.employees".to_string();
|
||||
let input = Arc::new(input);
|
||||
let admission = Arc::new(Semaphore::new(1));
|
||||
let optimizer = Arc::new(CascadeOptimizerBuilder::default().build());
|
||||
let scheduler = Arc::new(LocalScheduler {});
|
||||
let dispatcher = test_dispatcher_for_input(
|
||||
Arc::clone(&input),
|
||||
Arc::clone(&admission),
|
||||
Duration::from_secs(300),
|
||||
Arc::new(SqlQueryExecutionFactory::new(optimizer, scheduler)),
|
||||
);
|
||||
let query = Query::new(
|
||||
QueryContext {
|
||||
input: Arc::clone(&input),
|
||||
},
|
||||
input.request.expression.clone(),
|
||||
);
|
||||
let query_state_machine = dispatcher
|
||||
.build_query_state_machine(query)
|
||||
.await
|
||||
.expect("validation errors should remain in the planning phase");
|
||||
|
||||
let result = dispatcher.build_logical_plan(query_state_machine).await;
|
||||
|
||||
assert!(matches!(
|
||||
result,
|
||||
Err(ref error) if matches!(error.s3_select_policy_error(), Some(SelectError::DataSourcePathUnsupported))
|
||||
));
|
||||
assert_eq!(admission.available_permits(), 1);
|
||||
}
|
||||
|
||||
fn parquet_snapshot_input() -> Arc<SelectObjectContentInput> {
|
||||
Arc::new(SelectObjectContentInput {
|
||||
bucket: "s3select-parquet-snapshot-race".to_string(),
|
||||
@@ -1530,6 +1653,140 @@ mod tests {
|
||||
assert!(matches!(result, Err(QueryError::Cancel)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unprepared_tracked_session_rejects_source_path_expansion() {
|
||||
let mut input = test_input();
|
||||
input.key = "test.json".to_string();
|
||||
input.request.expression = "SELECT * FROM S3Object.employees".to_string();
|
||||
input.request.input_serialization = InputSerialization {
|
||||
json: Some(JSONInput {
|
||||
type_: Some(JSONType::from_static(JSONType::DOCUMENT)),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
input.request.output_serialization = OutputSerialization {
|
||||
json: Some(JSONOutput::default()),
|
||||
..Default::default()
|
||||
};
|
||||
let input = Arc::new(input);
|
||||
let admission = Arc::new(Semaphore::new(1));
|
||||
let optimizer = Arc::new(CascadeOptimizerBuilder::default().build());
|
||||
let scheduler = Arc::new(LocalScheduler {});
|
||||
let dispatcher = test_dispatcher_for_input(
|
||||
Arc::clone(&input),
|
||||
Arc::clone(&admission),
|
||||
Duration::from_secs(300),
|
||||
Arc::new(SqlQueryExecutionFactory::new(optimizer, scheduler)),
|
||||
);
|
||||
let query = Query::new(
|
||||
QueryContext {
|
||||
input: Arc::clone(&input),
|
||||
},
|
||||
input.request.expression.clone(),
|
||||
);
|
||||
let permit = Arc::clone(&admission).acquire_owned().await.expect("admission permit");
|
||||
let tracker = QueryExecutionTracker::new(
|
||||
&dispatcher.query_execution_owner,
|
||||
Arc::new(permit),
|
||||
Instant::now() + Duration::from_secs(300),
|
||||
300,
|
||||
);
|
||||
let session = SessionCtxFactory::new(true)
|
||||
.create_session_ctx_with_tracker_and_memory_limit(
|
||||
query.context(),
|
||||
tracker.clone(),
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
|
||||
)
|
||||
.await
|
||||
.expect("test session");
|
||||
assert!(tracker.mark_admitted(&dispatcher.query_execution_owner));
|
||||
let state_machine =
|
||||
Arc::new(QueryStateMachine::begin_tracked(query, session, tracker).expect("tracked state machine should be valid"));
|
||||
|
||||
let result = dispatcher.build_logical_plan(state_machine).await;
|
||||
|
||||
assert!(matches!(
|
||||
result,
|
||||
Err(ref error) if matches!(error.s3_select_policy_error(), Some(SelectError::DataSourcePathUnsupported))
|
||||
));
|
||||
assert_eq!(admission.available_permits(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn unprepared_tracked_session_preserves_root_scalar_bindings() {
|
||||
for (bucket, expression) in [
|
||||
("s3select-unprepared-root-scalar-alias", "SELECT V FROM S3Object AS V"),
|
||||
("s3select-unprepared-root-wildcard", "SELECT _1 FROM S3Object[*]"),
|
||||
] {
|
||||
let mut input = test_input();
|
||||
input.bucket = bucket.to_string();
|
||||
input.key = "input.json".to_string();
|
||||
input.request.expression = expression.to_string();
|
||||
input.request.input_serialization = InputSerialization {
|
||||
json: Some(JSONInput {
|
||||
type_: Some(JSONType::from_static(JSONType::DOCUMENT)),
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
input.request.output_serialization = OutputSerialization {
|
||||
json: Some(JSONOutput::default()),
|
||||
..Default::default()
|
||||
};
|
||||
let input = Arc::new(input);
|
||||
let env = snapshot_test_env().await;
|
||||
env.make_bucket(&input.bucket, false).await;
|
||||
env.put_object_bytes(&input.bucket, &input.key, br#"["one","two"]"#.to_vec())
|
||||
.await;
|
||||
let snapshot = env.prepare_select_object_snapshot(&input.bucket, &input.key).await;
|
||||
let dispatcher = production_dispatcher(Arc::clone(&input));
|
||||
let query = Query::new_with_snapshot(
|
||||
QueryContext {
|
||||
input: Arc::clone(&input),
|
||||
},
|
||||
input.request.expression.clone(),
|
||||
Arc::clone(&snapshot),
|
||||
);
|
||||
let permit = Arc::clone(&dispatcher.query_admission)
|
||||
.acquire_owned()
|
||||
.await
|
||||
.expect("query permit should be available");
|
||||
let tracker = QueryExecutionTracker::new(
|
||||
&dispatcher.query_execution_owner,
|
||||
Arc::new(permit),
|
||||
Instant::now() + Duration::from_secs(300),
|
||||
300,
|
||||
);
|
||||
let session = SessionCtxFactory::new(false)
|
||||
.create_session_ctx_with_snapshot_and_tracker_and_memory_limit(
|
||||
query.context(),
|
||||
snapshot,
|
||||
tracker.clone(),
|
||||
DEFAULT_S3SELECT_MEMORY_LIMIT_BYTES,
|
||||
)
|
||||
.await
|
||||
.expect("legacy session should preserve the root scalar binding");
|
||||
assert!(tracker.mark_admitted(&dispatcher.query_execution_owner));
|
||||
let state_machine = Arc::new(
|
||||
QueryStateMachine::begin_tracked(query, session, tracker).expect("tracked state machine should be valid"),
|
||||
);
|
||||
|
||||
let logical_plan = dispatcher
|
||||
.build_logical_plan(Arc::clone(&state_machine))
|
||||
.await
|
||||
.expect("unprepared query should plan")
|
||||
.expect("SELECT should produce a logical plan");
|
||||
let values = collect_utf8_output(
|
||||
dispatcher
|
||||
.execute_logical_plan(logical_plan, state_machine)
|
||||
.await
|
||||
.expect("unprepared query should execute"),
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(values, ["one", "two"]);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn staged_query_rejects_unbound_session() {
|
||||
let admission = Arc::new(Semaphore::new(1));
|
||||
@@ -1955,6 +2212,38 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_sql_precedes_malformed_json_snapshot_read() {
|
||||
let mut input = json_snapshot_input();
|
||||
let input_mut = Arc::make_mut(&mut input);
|
||||
input_mut.bucket = "s3select-invalid-sql-precedence".to_string();
|
||||
input_mut.key = "malformed.json".to_string();
|
||||
input_mut.request.expression = "SELECT * FROM".to_string();
|
||||
input_mut.request.input_serialization.json.as_mut().unwrap().type_ = Some(JSONType::from_static(JSONType::DOCUMENT));
|
||||
|
||||
let env = snapshot_test_env().await;
|
||||
env.make_bucket(&input.bucket, false).await;
|
||||
env.put_object_bytes(&input.bucket, &input.key, b"{bad".to_vec()).await;
|
||||
let snapshot = env.prepare_select_object_snapshot(&input.bucket, &input.key).await;
|
||||
let dispatcher = production_dispatcher(Arc::clone(&input));
|
||||
let query = Query::new_with_snapshot(
|
||||
QueryContext {
|
||||
input: Arc::clone(&input),
|
||||
},
|
||||
input.request.expression.clone(),
|
||||
snapshot,
|
||||
);
|
||||
let query_state_machine = dispatcher
|
||||
.build_query_state_machine(query)
|
||||
.await
|
||||
.expect("invalid SQL should remain a planning-phase error");
|
||||
|
||||
assert!(matches!(
|
||||
dispatcher.build_logical_plan(query_state_machine).await,
|
||||
Err(QueryError::Parser { .. })
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||
async fn cancelled_execution_start_drops_future_before_releasing_admission() {
|
||||
let admission = Arc::new(Semaphore::new(1));
|
||||
|
||||
@@ -184,6 +184,13 @@ mod tests {
|
||||
assert!(dialect.supports_group_by_expr(), "RustFsDialect should support GROUP BY expressions");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_supports_partiql_paths() {
|
||||
let dialect = RustFsDialect;
|
||||
|
||||
assert!(dialect.supports_partiql(), "RustFsDialect should support JSON source paths");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_identifier_validation_comprehensive() {
|
||||
let dialect = RustFsDialect;
|
||||
|
||||
@@ -16,6 +16,7 @@ use std::{collections::VecDeque, fmt::Display};
|
||||
|
||||
use datafusion::sql::sqlparser::{
|
||||
dialect::Dialect,
|
||||
keywords::{Keyword, RESERVED_FOR_TABLE_ALIAS},
|
||||
parser::{Parser, ParserError},
|
||||
tokenizer::{Token, Tokenizer},
|
||||
};
|
||||
@@ -53,7 +54,8 @@ impl<'a> ExtParser<'a> {
|
||||
/// Parse the specified tokens with dialect
|
||||
fn new_with_dialect(sql: &str, dialect: &'a dyn Dialect) -> Result<Self> {
|
||||
let mut tokenizer = Tokenizer::new(dialect, sql);
|
||||
let tokens = tokenizer.tokenize()?;
|
||||
let mut tokens = tokenizer.tokenize()?;
|
||||
rewrite_source_object_wildcards(&mut tokens);
|
||||
Ok(ExtParser {
|
||||
parser: Parser::new(dialect).with_tokens(tokens),
|
||||
})
|
||||
@@ -104,6 +106,41 @@ impl<'a> ExtParser<'a> {
|
||||
}
|
||||
}
|
||||
|
||||
fn rewrite_source_object_wildcards(tokens: &mut [Token]) {
|
||||
let mut paren_depth = 0usize;
|
||||
let mut in_from = false;
|
||||
let mut index = 0usize;
|
||||
|
||||
while index < tokens.len() {
|
||||
match &tokens[index] {
|
||||
Token::Word(word) if paren_depth == 0 && word.keyword == Keyword::FROM => {
|
||||
in_from = true;
|
||||
}
|
||||
Token::Word(word) if in_from && paren_depth == 0 && ends_from_source(word.keyword) => {
|
||||
in_from = false;
|
||||
}
|
||||
Token::SemiColon if paren_depth == 0 => in_from = false,
|
||||
Token::LParen => paren_depth = paren_depth.saturating_add(1),
|
||||
Token::RParen => paren_depth = paren_depth.saturating_sub(1),
|
||||
Token::Period if in_from => {
|
||||
if let Some(next) = tokens[index + 1..]
|
||||
.iter_mut()
|
||||
.find(|token| !matches!(token, Token::Whitespace(_)))
|
||||
&& matches!(next, Token::Mul)
|
||||
{
|
||||
*next = Token::make_word("*", None);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
index += 1;
|
||||
}
|
||||
}
|
||||
|
||||
fn ends_from_source(keyword: Keyword) -> bool {
|
||||
RESERVED_FOR_TABLE_ALIAS.contains(&keyword) || matches!(keyword, Keyword::PREWHERE | Keyword::SETTINGS | Keyword::FORMAT)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -172,6 +209,23 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_source_object_wildcard_without_rewriting_projection_wildcard() {
|
||||
let mut statements = ExtParser::parse_sql("SELECT e.* FROM S3Object[*].* AS e").expect("query should parse");
|
||||
let ExtStatement::SqlStatement(statement) = statements.pop_front().expect("one statement");
|
||||
|
||||
assert_eq!(statement.to_string(), "SELECT e.* FROM S3Object[*].* AS e");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn from_tokens_in_literals_and_comments_do_not_change_wildcard_scope() {
|
||||
let sql = "SELECT 'FROM x.*' AS marker /* FROM y.* */ FROM S3Object.*";
|
||||
let mut statements = ExtParser::parse_sql(sql).expect("query should parse");
|
||||
let ExtStatement::SqlStatement(statement) = statements.pop_front().expect("one statement");
|
||||
|
||||
assert_eq!(statement.to_string(), "SELECT 'FROM x.*' AS marker FROM S3Object.*");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_parser_multiple_statements() {
|
||||
let parser = DefaultParser::default();
|
||||
|
||||
@@ -12,20 +12,21 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
use std::ops::ControlFlow;
|
||||
use std::{convert::Infallible, ops::ControlFlow};
|
||||
|
||||
use async_recursion::async_recursion;
|
||||
use async_trait::async_trait;
|
||||
use datafusion::sql::{
|
||||
planner::SqlToRel,
|
||||
planner::{IdentNormalizer, SqlToRel},
|
||||
sqlparser::ast::{
|
||||
GroupByExpr, ObjectNamePart, OrderByKind, Query, Select, SelectFlavor, SetExpr, Statement, TableFactor, Visit, Visitor,
|
||||
AccessExpr, Expr, GroupByExpr, Ident, JsonPath, JsonPathElem, ObjectNamePart, OrderByKind, Query, Select, SelectFlavor,
|
||||
SetExpr, Statement, Subscript, TableAlias, TableFactor, Value, Visit, VisitMut, Visitor, VisitorMut,
|
||||
},
|
||||
};
|
||||
use rustfs_s3select_api::{
|
||||
QueryError, QueryResult, SelectError,
|
||||
query::{
|
||||
ast::ExtStatement,
|
||||
ast::{ExtStatement, JsonPathSegment, JsonSource},
|
||||
logical_planner::{LogicalPlanner, Plan, QueryPlan},
|
||||
session::SessionCtx,
|
||||
},
|
||||
@@ -64,21 +65,24 @@ impl<'a, S: ContextProviderExtension + Send + Sync + 'a> SqlPlanner<'a, S> {
|
||||
}
|
||||
}
|
||||
|
||||
async fn df_sql_to_plan(&self, stmt: Statement, _session: &SessionCtx) -> QueryResult<Plan> {
|
||||
match stmt {
|
||||
Statement::Query(_) => {
|
||||
validate_s3_select_statement(&stmt)?;
|
||||
let df_plan = self.df_planner.sql_statement_to_plan(stmt).map_err(classify_planner_error)?;
|
||||
let plan = Plan::Query(QueryPlan {
|
||||
df_plan,
|
||||
is_tag_scan: false,
|
||||
});
|
||||
|
||||
Ok(plan)
|
||||
}
|
||||
_ => Err(unsupported_structure("only SELECT queries are supported")),
|
||||
pub(crate) async fn prepared_statement_to_plan(&self, statement: ExtStatement, session: &SessionCtx) -> QueryResult<Plan> {
|
||||
match statement {
|
||||
ExtStatement::SqlStatement(stmt) => self.df_prepared_sql_to_plan(*stmt, session).await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn df_sql_to_plan(&self, mut stmt: Statement, session: &SessionCtx) -> QueryResult<Plan> {
|
||||
prepare_s3_select_statement(&mut stmt)?;
|
||||
self.df_prepared_sql_to_plan(stmt, session).await
|
||||
}
|
||||
|
||||
async fn df_prepared_sql_to_plan(&self, stmt: Statement, _session: &SessionCtx) -> QueryResult<Plan> {
|
||||
let df_plan = self.df_planner.sql_statement_to_plan(stmt).map_err(classify_planner_error)?;
|
||||
Ok(Plan::Query(QueryPlan {
|
||||
df_plan,
|
||||
is_tag_scan: false,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
fn classify_planner_error(error: datafusion::common::DataFusionError) -> QueryError {
|
||||
@@ -95,11 +99,10 @@ fn classify_planner_error(error: datafusion::common::DataFusionError) -> QueryEr
|
||||
error.into()
|
||||
}
|
||||
|
||||
fn validate_s3_select_statement(statement: &Statement) -> QueryResult<()> {
|
||||
pub(crate) fn prepare_s3_select_statement(statement: &mut Statement) -> QueryResult<JsonSource> {
|
||||
let Statement::Query(query) = statement else {
|
||||
return Err(unsupported_structure("only SELECT queries are supported"));
|
||||
};
|
||||
|
||||
if query.with.is_some()
|
||||
|| query.order_by.as_ref().is_some_and(|order_by| {
|
||||
order_by.interpolate.is_some()
|
||||
@@ -137,17 +140,81 @@ fn validate_s3_select_statement(statement: &Statement) -> QueryResult<()> {
|
||||
}
|
||||
|
||||
let mut detector = SubqueryDetector { visited_root: false };
|
||||
if query.visit(&mut detector).is_break() {
|
||||
if Visit::visit(&*query, &mut detector).is_break() {
|
||||
return Err(unsupported_structure("subqueries are not supported"));
|
||||
}
|
||||
|
||||
let SetExpr::Select(select) = query.body.as_ref() else {
|
||||
return Err(unsupported_structure("set operations and nested queries are not supported"));
|
||||
let source = {
|
||||
let SetExpr::Select(select) = query.body.as_mut() else {
|
||||
return Err(unsupported_structure("set operations and nested queries are not supported"));
|
||||
};
|
||||
prepare_select(select)?
|
||||
};
|
||||
validate_select(select)
|
||||
let mut normalizer = PartiQlSubscriptNormalizer;
|
||||
let _ = VisitMut::visit(query, &mut normalizer);
|
||||
Ok(source)
|
||||
}
|
||||
|
||||
fn validate_select(select: &Select) -> QueryResult<()> {
|
||||
struct PartiQlSubscriptNormalizer;
|
||||
|
||||
impl VisitorMut for PartiQlSubscriptNormalizer {
|
||||
type Break = Infallible;
|
||||
|
||||
fn post_visit_expr(&mut self, expr: &mut Expr) -> ControlFlow<Self::Break> {
|
||||
let Expr::JsonAccess { value, path } = expr else {
|
||||
return ControlFlow::Continue(());
|
||||
};
|
||||
if !matches!(path.path.first(), Some(JsonPathElem::Bracket { .. }))
|
||||
|| path
|
||||
.path
|
||||
.iter()
|
||||
.any(|element| matches!(element, JsonPathElem::ColonBracket { .. }))
|
||||
{
|
||||
return ControlFlow::Continue(());
|
||||
}
|
||||
|
||||
let mut appended_access = Vec::with_capacity(path.path.len());
|
||||
for element in std::mem::take(&mut path.path) {
|
||||
match element {
|
||||
JsonPathElem::Dot { key, quoted } => {
|
||||
let identifier = if quoted {
|
||||
Ident::with_quote('"', key)
|
||||
} else {
|
||||
Ident::new(key)
|
||||
};
|
||||
appended_access.push(AccessExpr::Dot(Expr::Identifier(identifier)));
|
||||
}
|
||||
JsonPathElem::Bracket { key } => {
|
||||
appended_access.push(AccessExpr::Subscript(Subscript::Index { index: key }));
|
||||
}
|
||||
JsonPathElem::ColonBracket { key } => {
|
||||
appended_access.push(AccessExpr::Subscript(Subscript::Index { index: key }));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let value = std::mem::replace(value, Box::new(Expr::Identifier(Ident::new(""))));
|
||||
let (root, mut access_chain) = match *value {
|
||||
Expr::CompoundFieldAccess { root, access_chain } => (root, access_chain),
|
||||
root => (Box::new(root), Vec::new()),
|
||||
};
|
||||
access_chain.extend(appended_access);
|
||||
*expr = Expr::CompoundFieldAccess { root, access_chain };
|
||||
ControlFlow::Continue(())
|
||||
}
|
||||
}
|
||||
|
||||
fn implicit_source_alias(source_path: &[JsonPathSegment]) -> Ident {
|
||||
match source_path.last() {
|
||||
Some(JsonPathSegment::Key { name, quoted: true }) => Ident::with_quote('"', name),
|
||||
Some(JsonPathSegment::Key { name, quoted: false }) => Ident::new(name),
|
||||
Some(JsonPathSegment::Index(_) | JsonPathSegment::ArrayWildcard | JsonPathSegment::ObjectWildcard) | None => {
|
||||
Ident::new("_1")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn prepare_select(select: &mut Select) -> QueryResult<JsonSource> {
|
||||
if !select.optimizer_hints.is_empty()
|
||||
|| select.distinct.is_some()
|
||||
|| select.select_modifiers.is_some()
|
||||
@@ -170,7 +237,7 @@ fn validate_select(select: &Select) -> QueryResult<()> {
|
||||
return Err(unsupported_structure("the SELECT contains an unsupported clause"));
|
||||
}
|
||||
|
||||
let [table] = select.from.as_slice() else {
|
||||
let [table] = select.from.as_mut_slice() else {
|
||||
return Err(unsupported_structure("exactly one S3Object source is required"));
|
||||
};
|
||||
if !table.joins.is_empty() {
|
||||
@@ -186,8 +253,8 @@ fn validate_select(select: &Select) -> QueryResult<()> {
|
||||
partitions,
|
||||
sample,
|
||||
index_hints,
|
||||
..
|
||||
} = &table.relation
|
||||
json_path,
|
||||
} = &mut table.relation
|
||||
else {
|
||||
return Err(unsupported_structure("subqueries and table functions are not supported"));
|
||||
};
|
||||
@@ -202,9 +269,7 @@ fn validate_select(select: &Select) -> QueryResult<()> {
|
||||
{
|
||||
return Err(unsupported_structure("the S3Object source contains unsupported modifiers"));
|
||||
}
|
||||
let ([ObjectNamePart::Identifier(table_name)] | [ObjectNamePart::Identifier(table_name), ObjectNamePart::Identifier(_)]) =
|
||||
name.0.as_slice()
|
||||
else {
|
||||
let Some(ObjectNamePart::Identifier(table_name)) = name.0.first() else {
|
||||
return Err(SelectError::DataSourcePathUnsupported.into());
|
||||
};
|
||||
let is_s3_object = if table_name.quote_style.is_some() {
|
||||
@@ -216,6 +281,72 @@ fn validate_select(select: &Select) -> QueryResult<()> {
|
||||
return Err(SelectError::DataSourcePathUnsupported.into());
|
||||
}
|
||||
|
||||
let mut source_path = Vec::new();
|
||||
for part in &name.0[1..] {
|
||||
let ObjectNamePart::Identifier(identifier) = part else {
|
||||
return Err(SelectError::DataSourcePathUnsupported.into());
|
||||
};
|
||||
if identifier.quote_style.is_none() && identifier.value == "*" {
|
||||
source_path.push(JsonPathSegment::ObjectWildcard);
|
||||
} else {
|
||||
source_path.push(JsonPathSegment::Key {
|
||||
name: identifier.value.clone(),
|
||||
quoted: identifier.quote_style.is_some(),
|
||||
});
|
||||
}
|
||||
}
|
||||
if let Some(json_path) = json_path.as_ref() {
|
||||
append_json_path_segments(&mut source_path, json_path)?;
|
||||
}
|
||||
if alias.is_none() && !source_path.is_empty() {
|
||||
*alias = Some(TableAlias {
|
||||
explicit: true,
|
||||
name: implicit_source_alias(&source_path),
|
||||
columns: Vec::new(),
|
||||
at: None,
|
||||
});
|
||||
}
|
||||
let scalar_column = alias
|
||||
.as_ref()
|
||||
.map(|alias| IdentNormalizer::default().normalize(alias.name.clone()))
|
||||
.or_else(|| {
|
||||
source_path
|
||||
.is_empty()
|
||||
.then(|| IdentNormalizer::default().normalize(table_name.clone()))
|
||||
});
|
||||
name.0.truncate(1);
|
||||
*json_path = None;
|
||||
Ok(JsonSource::new(source_path, scalar_column))
|
||||
}
|
||||
|
||||
fn append_json_path_segments(source_path: &mut Vec<JsonPathSegment>, json_path: &JsonPath) -> QueryResult<()> {
|
||||
for element in &json_path.path {
|
||||
let segment = match element {
|
||||
JsonPathElem::Dot { key, quoted } if key == "*" && !quoted => JsonPathSegment::ObjectWildcard,
|
||||
JsonPathElem::Dot { key, quoted } => JsonPathSegment::Key {
|
||||
name: key.clone(),
|
||||
quoted: *quoted,
|
||||
},
|
||||
JsonPathElem::Bracket { key: Expr::Wildcard(_) } => JsonPathSegment::ArrayWildcard,
|
||||
JsonPathElem::Bracket { key: Expr::Value(value) } => match &value.value {
|
||||
Value::Number(number, false) => JsonPathSegment::Index(
|
||||
number
|
||||
.to_string()
|
||||
.parse()
|
||||
.map_err(|_| QueryError::from(SelectError::DataSourcePathUnsupported))?,
|
||||
),
|
||||
Value::SingleQuotedString(key) => JsonPathSegment::Key {
|
||||
name: key.clone(),
|
||||
quoted: true,
|
||||
},
|
||||
_ => return Err(SelectError::DataSourcePathUnsupported.into()),
|
||||
},
|
||||
JsonPathElem::Bracket { .. } | JsonPathElem::ColonBracket { .. } => {
|
||||
return Err(SelectError::DataSourcePathUnsupported.into());
|
||||
}
|
||||
};
|
||||
source_path.push(segment);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -245,10 +376,14 @@ impl Visitor for SubqueryDetector {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::validate_s3_select_statement;
|
||||
use super::prepare_s3_select_statement;
|
||||
use crate::sql::parser::ExtParser;
|
||||
use datafusion::sql::sqlparser::ast::Statement;
|
||||
use rustfs_s3select_api::{SelectError, query::ast::ExtStatement};
|
||||
use datafusion::sql::sqlparser::ast::{AccessExpr, Expr, Statement, Visit, Visitor};
|
||||
use rustfs_s3select_api::{
|
||||
QueryResult, SelectError,
|
||||
query::ast::{ExtStatement, JsonPathSegment, JsonSource},
|
||||
};
|
||||
use std::ops::ControlFlow;
|
||||
|
||||
fn parse_statement(sql: &str) -> Statement {
|
||||
let mut statements = ExtParser::parse_sql(sql).expect("SQL should parse");
|
||||
@@ -256,6 +391,10 @@ mod tests {
|
||||
*statement
|
||||
}
|
||||
|
||||
fn validate_s3_select_statement(statement: &Statement) -> QueryResult<JsonSource> {
|
||||
prepare_s3_select_statement(&mut statement.clone())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_s3_select_query_shape() {
|
||||
let statement = parse_statement("SELECT s.id FROM S3Object AS s WHERE s.id = '1' LIMIT 10");
|
||||
@@ -270,6 +409,195 @@ mod tests {
|
||||
assert!(validate_s3_select_statement(&statement).is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepares_nested_json_source_path_and_normalizes_table() {
|
||||
let mut statement = parse_statement("SELECT e.name FROM S3Object[*].employees[*] AS e");
|
||||
|
||||
let source = prepare_s3_select_statement(&mut statement).expect("JSON source path should be supported");
|
||||
|
||||
assert_eq!(
|
||||
source.path(),
|
||||
&[
|
||||
JsonPathSegment::ArrayWildcard,
|
||||
JsonPathSegment::Key {
|
||||
name: "employees".to_string(),
|
||||
quoted: false,
|
||||
},
|
||||
JsonPathSegment::ArrayWildcard,
|
||||
]
|
||||
);
|
||||
assert_eq!(statement.to_string(), "SELECT e.name FROM S3Object AS e");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn partiql_source_support_preserves_projection_and_filter_subscripts() {
|
||||
let mut statement = parse_statement("SELECT s.tags[1] FROM S3Object AS s WHERE s.values[0] = 1");
|
||||
|
||||
prepare_s3_select_statement(&mut statement).expect("array expressions should remain supported");
|
||||
let mut counter = FieldAccessCounter::default();
|
||||
let _ = Visit::visit(&statement, &mut counter);
|
||||
|
||||
assert_eq!(counter.json_accesses, 0);
|
||||
assert_eq!(counter.subscripts, 2);
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct FieldAccessCounter {
|
||||
json_accesses: usize,
|
||||
subscripts: usize,
|
||||
}
|
||||
|
||||
impl Visitor for FieldAccessCounter {
|
||||
type Break = ();
|
||||
|
||||
fn pre_visit_expr(&mut self, expr: &Expr) -> ControlFlow<Self::Break> {
|
||||
match expr {
|
||||
Expr::JsonAccess { .. } => self.json_accesses += 1,
|
||||
Expr::CompoundFieldAccess { access_chain, .. } => {
|
||||
self.subscripts += access_chain
|
||||
.iter()
|
||||
.filter(|access| matches!(access, AccessExpr::Subscript(_)))
|
||||
.count();
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
ControlFlow::Continue(())
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prepares_array_index_and_object_wildcard_paths() {
|
||||
let mut index_statement = parse_statement("SELECT * FROM S3Object[0]");
|
||||
let mut wildcard_statement = parse_statement("SELECT * FROM S3Object[*].*");
|
||||
|
||||
assert_eq!(
|
||||
prepare_s3_select_statement(&mut index_statement)
|
||||
.expect("array index should be supported")
|
||||
.path(),
|
||||
&[JsonPathSegment::Index(0)]
|
||||
);
|
||||
assert_eq!(
|
||||
prepare_s3_select_statement(&mut wildcard_statement)
|
||||
.expect("object wildcard should be supported")
|
||||
.path(),
|
||||
&[JsonPathSegment::ArrayWildcard, JsonPathSegment::ObjectWildcard]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn quoted_star_remains_an_object_key() {
|
||||
let mut statement = parse_statement("SELECT * FROM S3Object.\"*\"");
|
||||
|
||||
let source = prepare_s3_select_statement(&mut statement).expect("quoted key should be supported");
|
||||
|
||||
assert_eq!(
|
||||
source.path(),
|
||||
&[JsonPathSegment::Key {
|
||||
name: "*".to_string(),
|
||||
quoted: true,
|
||||
}]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preserves_quoted_keys_and_adds_implicit_source_aliases() {
|
||||
let mut key_statement = parse_statement("SELECT employee.name FROM S3Object[*].department.employee");
|
||||
let mut wildcard_statement = parse_statement("SELECT _1.name FROM S3Object[*].employees[*]");
|
||||
|
||||
prepare_s3_select_statement(&mut key_statement).expect("named source path should be supported");
|
||||
prepare_s3_select_statement(&mut wildcard_statement).expect("wildcard source path should be supported");
|
||||
|
||||
assert_eq!(key_statement.to_string(), "SELECT employee.name FROM S3Object AS employee");
|
||||
assert_eq!(wildcard_statement.to_string(), "SELECT _1.name FROM S3Object AS _1");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn root_scalar_aliases_are_preserved_and_unquoted_aliases_are_normalized() {
|
||||
let mut implicit = parse_statement("SELECT S3Object FROM S3Object");
|
||||
let mut unquoted = parse_statement("SELECT V FROM S3Object AS V");
|
||||
let mut quoted = parse_statement("SELECT \"V\" FROM S3Object AS \"V\"");
|
||||
|
||||
let implicit_source = prepare_s3_select_statement(&mut implicit).expect("implicit root alias should be supported");
|
||||
let unquoted_source = prepare_s3_select_statement(&mut unquoted).expect("unquoted root alias should be supported");
|
||||
let quoted_source = prepare_s3_select_statement(&mut quoted).expect("quoted root alias should be supported");
|
||||
|
||||
assert!(implicit_source.path().is_empty());
|
||||
assert_eq!(implicit_source.scalar_column(), Some("s3object"));
|
||||
assert!(unquoted_source.path().is_empty());
|
||||
assert_eq!(unquoted_source.scalar_column(), Some("v"));
|
||||
assert!(quoted_source.path().is_empty());
|
||||
assert_eq!(quoted_source.scalar_column(), Some("V"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unquoted_terminal_scalar_alias_uses_datafusion_identifier_case() {
|
||||
let mut statement = parse_statement("SELECT NAME FROM S3Object[*].NAME");
|
||||
|
||||
let source = prepare_s3_select_statement(&mut statement).expect("terminal scalar source should be supported");
|
||||
|
||||
assert_eq!(source.scalar_column(), Some("name"));
|
||||
assert_eq!(statement.to_string(), "SELECT NAME FROM S3Object AS NAME");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_quoted_source_key_adds_a_quoted_implicit_alias() {
|
||||
let mut statement = parse_statement("SELECT \"Employee Data\".id FROM S3Object['Employee Data']");
|
||||
|
||||
let source = prepare_s3_select_statement(&mut statement).expect("single-quoted source key should be supported");
|
||||
|
||||
assert_eq!(
|
||||
source.path(),
|
||||
&[JsonPathSegment::Key {
|
||||
name: "Employee Data".to_string(),
|
||||
quoted: true,
|
||||
}]
|
||||
);
|
||||
assert_eq!(source.scalar_column(), Some("Employee Data"));
|
||||
assert_eq!(statement.to_string(), "SELECT \"Employee Data\".id FROM S3Object AS \"Employee Data\"");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_object_wildcard_continuation() {
|
||||
let mut statement = parse_statement("SELECT * FROM S3Object[*].groups.*.id");
|
||||
|
||||
assert_eq!(
|
||||
prepare_s3_select_statement(&mut statement)
|
||||
.expect("object wildcard continuation should be supported")
|
||||
.path(),
|
||||
&[
|
||||
JsonPathSegment::ArrayWildcard,
|
||||
JsonPathSegment::Key {
|
||||
name: "groups".to_string(),
|
||||
quoted: false,
|
||||
},
|
||||
JsonPathSegment::ObjectWildcard,
|
||||
JsonPathSegment::Key {
|
||||
name: "id".to_string(),
|
||||
quoted: false,
|
||||
},
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_non_literal_or_out_of_range_array_indexes() {
|
||||
for sql in [
|
||||
"SELECT * FROM S3Object[-1]",
|
||||
"SELECT * FROM S3Object[1 + 1]",
|
||||
"SELECT * FROM S3Object[999999999999999999999999999999999999]",
|
||||
] {
|
||||
let statement = parse_statement(sql);
|
||||
assert!(
|
||||
matches!(
|
||||
validate_s3_select_statement(&statement),
|
||||
Err(ref error)
|
||||
if matches!(error.s3_select_policy_error(), Some(SelectError::DataSourcePathUnsupported))
|
||||
),
|
||||
"query should reject an unsafe array index: {sql}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_group_by_and_order_by() {
|
||||
let statement = parse_statement("SELECT department, COUNT(*) FROM S3Object GROUP BY department ORDER BY department");
|
||||
|
||||
Reference in New Issue
Block a user