Skip to content
Open
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
118 changes: 118 additions & 0 deletions datafusion/functions-nested/src/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,23 @@ impl ExprPlanner for NestedFunctionPlanner {
return Ok(PlannerResult::Planned(array_has_all(right, left)));
}
}
} else {
#[cfg(feature = "sql")]
let is_colon =
matches!(&op, BinaryOperator::Custom(operator) if operator == ":");
#[cfg(not(feature = "sql"))]
let is_colon = op == BinaryOperator::Colon;

if is_colon
&& matches!(left.get_type(schema)?, DataType::Struct(_))
&& let Expr::Literal(field_name, _) = &right
&& let Some(field_name) = field_name.try_as_str().flatten()
{
return Ok(PlannerResult::Planned(get_field(
left,
field_name.to_owned(),
)));
}
}

Ok(PlannerResult::Original(RawBinaryExpr { op, left, right }))
Expand Down Expand Up @@ -194,3 +211,104 @@ impl ExprPlanner for FieldAccessPlanner {
fn is_array_agg(func: &Arc<AggregateUDF>) -> bool {
func.name() == "array_agg"
}

#[cfg(all(test, feature = "sql"))]
mod tests {
use super::*;
use arrow::datatypes::{Field, Fields};
use datafusion_common::{Column, ScalarValue};
use datafusion_expr::planner::ExprPlanner;
use std::collections::HashMap;

fn nested_struct_schema() -> DFSchema {
let country = DataType::Struct(Fields::from(vec![Field::new(
"name",
DataType::Utf8,
true,
)]));
let payload =
DataType::Struct(Fields::from(vec![Field::new("country", country, true)]));
DFSchema::from_unqualified_fields(
vec![Field::new("payload", payload, true)].into(),
HashMap::new(),
)
.unwrap()
}

fn colon(left: Expr, field_name: &str) -> RawBinaryExpr {
RawBinaryExpr {
op: BinaryOperator::Custom(":".to_owned()),
left,
right: Expr::Literal(ScalarValue::from(field_name), None),
}
}

fn planned(expr: RawBinaryExpr, schema: &DFSchema) -> Expr {
let PlannerResult::Planned(expr) =
NestedFunctionPlanner.plan_binary_op(expr, schema).unwrap()
else {
panic!("expected colon access to be planned");
};
expr
}

#[test]
fn plans_nested_struct_colon_access_as_get_field() {
let schema = nested_struct_schema();
let payload = Expr::Column(Column::new_unqualified("payload"));
let country = planned(colon(payload.clone(), "country"), &schema);
let name = planned(colon(country.clone(), "name"), &schema);

assert_eq!(country, get_field(payload.clone(), "country"));
assert_eq!(name, get_field(get_field(payload, "country"), "name"));
}

#[test]
fn leaves_non_struct_colon_access_for_other_planners() {
let schema = DFSchema::from_unqualified_fields(
vec![Field::new("text", DataType::Utf8, true)].into(),
HashMap::new(),
)
.unwrap();
let original = colon(Expr::Column(Column::new_unqualified("text")), "field");

assert!(matches!(
NestedFunctionPlanner
.plan_binary_op(original, &schema)
.unwrap(),
PlannerResult::Original(_)
));
}

#[test]
fn leaves_unsupported_struct_access_for_other_planners() {
let schema = nested_struct_schema();
let payload = Expr::Column(Column::new_unqualified("payload"));
let unsupported = [
RawBinaryExpr {
op: BinaryOperator::Custom("other".to_owned()),
left: payload.clone(),
right: Expr::Literal(ScalarValue::from("country"), None),
},
RawBinaryExpr {
op: BinaryOperator::Custom(":".to_owned()),
left: payload.clone(),
right: Expr::Column(Column::new_unqualified("field_name")),
},
RawBinaryExpr {
op: BinaryOperator::Custom(":".to_owned()),
left: payload,
right: Expr::Literal(ScalarValue::Int64(Some(1)), None),
},
];

for expression in unsupported {
assert!(matches!(
NestedFunctionPlanner
.plan_binary_op(expression, &schema)
.unwrap(),
PlannerResult::Original(_)
));
}
}
}