Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion crates/frontend-sql/src/sql/clickhouse_ast.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,11 @@ pub(super) fn normalize(statement: &mut Statement) {
.value
.to_ascii_lowercase()
.as_str(),
"modulo" | "map" | "mapconcat" | "arrayelement"
"modulo"
| "map"
| "mapconcat"
| "arrayelement"
| "tupleelement"
);
}
ControlFlow::Continue(())
Expand Down
50 changes: 38 additions & 12 deletions crates/frontend-sql/src/sql/collection_planning.rs
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
//! DataFusion planning adapters. Types come from the canonical signature rules;
//! physical evaluation deliberately remains the query engine's responsibility.
use super::types::{arrow_to_dtype, dtype_to_arrow, scalar_value_to_asap};
use asap_types::pre_asap::scalar_signature::{element_access_type, MapScalarFunction};
use asap_types::pre_asap::scalar_signature::{
element_access_type, struct_field_type, MapScalarFunction,
};
use asap_types::pre_asap::{Column, QueryExpr, Schema};
use datafusion::arrow::datatypes::DataType;
use datafusion::common::{DataFusionError, ExprSchema, Result};
Expand All @@ -11,30 +13,45 @@ use datafusion::logical_expr::{
};
use datafusion::prelude::SessionContext;

#[derive(Debug, Clone, Copy)]
enum PlanningFunction {
Map(MapScalarFunction),
Element,
StructField,
}

pub(super) fn register(context: &SessionContext) {
for (name, function) in [
("map", MapScalarFunction::Construct),
("mapconcat", MapScalarFunction::Concat),
("arrayelement", MapScalarFunction::Access),
("map", PlanningFunction::Map(MapScalarFunction::Construct)),
(
"mapconcat",
PlanningFunction::Map(MapScalarFunction::Concat),
),
("arrayelement", PlanningFunction::Element),
("tupleelement", PlanningFunction::StructField),
] {
context.register_udf(ScalarUDF::from(CollectionPlanningFunction {
name,
function,
signature: match function {
MapScalarFunction::Construct => Signature::one_of(
PlanningFunction::Map(MapScalarFunction::Construct) => Signature::one_of(
vec![TypeSignature::Exact(vec![]), TypeSignature::VariadicAny],
Volatility::Immutable,
),
MapScalarFunction::Access => Signature::any(2, Volatility::Immutable),
MapScalarFunction::Concat => Signature::variadic_any(Volatility::Immutable),
PlanningFunction::Map(MapScalarFunction::Access)
| PlanningFunction::Element
| PlanningFunction::StructField => Signature::any(2, Volatility::Immutable),
PlanningFunction::Map(MapScalarFunction::Concat) => {
Signature::variadic_any(Volatility::Immutable)
}
},
}));
}
}
#[derive(Debug)]
struct CollectionPlanningFunction {
name: &'static str,
function: MapScalarFunction,
function: PlanningFunction,
signature: Signature,
}
impl CollectionPlanningFunction {
Expand All @@ -53,7 +70,10 @@ impl CollectionPlanningFunction {
.map_err(|e| DataFusionError::Plan(e.to_string()))
})
.collect::<Result<Vec<_>>>()?;
let (dtype, nullable) = if self.name == "arrayelement" {
let (dtype, nullable) = if matches!(
self.function,
PlanningFunction::Element | PlanningFunction::StructField
) {
// DataFusion asks for argument-dependent types before canonical
// expression binding. Reuse the shared resolver over typed argument
// slots; final canonical binding also validates literal selectors.
Expand All @@ -78,9 +98,15 @@ impl CollectionPlanningFunction {
}
})
.collect::<Result<Vec<_>>>()?;
element_access_type(&args, &schema)
match self.function {
PlanningFunction::Element => element_access_type(&args, &schema),
PlanningFunction::StructField => struct_field_type(&args, &schema),
PlanningFunction::Map(_) => unreachable!(),
}
} else if let PlanningFunction::Map(function) = self.function {
function.output_type(&inputs)
} else {
self.function.output_type(&inputs)
unreachable!()
}
.map_err(DataFusionError::Plan)?;
Ok((dtype_to_arrow(&dtype), nullable))
Expand Down Expand Up @@ -149,7 +175,7 @@ mod tests {
fn planning_adapter_explicitly_refuses_physical_execution() {
let adapter = CollectionPlanningFunction {
name: "map",
function: MapScalarFunction::Construct,
function: PlanningFunction::Map(MapScalarFunction::Construct),
signature: Signature::any(0, Volatility::Immutable),
};
assert!(matches!(
Expand Down
2 changes: 2 additions & 0 deletions crates/frontend-sql/src/sql/expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -204,6 +204,8 @@ pub(super) fn df_expr_to_unresolved(expr: &Expr) -> Result<Unresolved, LoweringE
Ok(Unresolved::FunctionCall {
name: if sf.func.name().eq_ignore_ascii_case("arrayelement") {
"asap_element_access".into()
} else if sf.func.name().eq_ignore_ascii_case("tupleelement") {
"asap_struct_field".into()
} else {
sf.func.name().to_string()
},
Expand Down
61 changes: 61 additions & 0 deletions crates/frontend-sql/tests/sql_lowering.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2393,3 +2393,64 @@ async fn clickhouse_list_element_uses_canonical_typed_access() {
);
}
}

#[tokio::test]
async fn clickhouse_tuple_element_preserves_declared_field_metadata() {
let catalog = SqlCatalog::new().with_table(
"t",
Schema::new(vec![
Column::new(
"sample",
DataType::Struct {
fields: vec![
Column::new("time", DataType::Int64, false),
Column::new("value", DataType::Float64, true),
],
},
false,
),
Column::new("index", DataType::Int64, false),
]),
);
for (sql, dtype, nullable) in [
(
"SELECT tupleElement(sample, 1) AS chosen FROM t",
DataType::Int64,
false,
),
(
"SELECT tupleElement(sample, 'value') AS chosen FROM t",
DataType::Float64,
true,
),
] {
let query = lower_sql_dialect(
sql,
&catalog,
SqlDialect::ClickhouseSQL,
AccuracyTarget::Exact,
)
.await
.unwrap();
let output = query.output_schema().unwrap();
assert_eq!(output.columns[0].dtype, dtype);
assert_eq!(output.columns[0].nullable, nullable);
assert!(serde_json::to_string(&query)
.unwrap()
.contains("asap_struct_field"));
}
for selector in ["0", "-1", "3", "'missing'", "index"] {
let sql = format!("SELECT tupleElement(sample, {selector}) FROM t");
assert!(
lower_sql_dialect(
&sql,
&catalog,
SqlDialect::ClickhouseSQL,
AccuracyTarget::Exact
)
.await
.is_err(),
"{sql}"
);
}
}
Loading