diff --git a/datafusion/expr/src/planner.rs b/datafusion/expr/src/planner.rs index 7aaf3a98cbe5d..4009dfb3e9e19 100644 --- a/datafusion/expr/src/planner.rs +++ b/datafusion/expr/src/planner.rs @@ -263,6 +263,13 @@ pub trait ExprPlanner: Debug + Send + Sync { ) } + /// Plans scalar functions, such as `ABS()` + /// + /// Returns the original scalar function if not possible + fn plan_scalar(&self, expr: RawScalarExpr) -> Result> { + Ok(PlannerResult::Original(expr)) + } + /// Plans aggregate functions, such as `COUNT()` /// /// Returns original expression arguments if not possible @@ -318,6 +325,13 @@ pub struct RawDictionaryExpr { pub values: Vec, } +/// A scalar function to plan with [`ExprPlanner`]. +#[derive(Debug, Clone)] +pub struct RawScalarExpr { + pub func: Arc, + pub args: Vec, +} + /// This structure is used by `AggregateFunctionPlanner` to plan operators with /// custom expressions. #[derive(Debug, Clone)] diff --git a/datafusion/sql/src/expr/function.rs b/datafusion/sql/src/expr/function.rs index f1ab0db470008..a7f7979a67b36 100644 --- a/datafusion/sql/src/expr/function.rs +++ b/datafusion/sql/src/expr/function.rs @@ -30,7 +30,7 @@ use datafusion_expr::{ self, HigherOrderFunction, Lambda, NullTreatment, ScalarFunction, Unnest, WildcardOptions, WindowFunction, }, - planner::{PlannerResult, RawAggregateExpr, RawWindowExpr}, + planner::{PlannerResult, RawAggregateExpr, RawScalarExpr, RawWindowExpr}, type_coercion::functions::value_fields_with_higher_order_udf, }; use sqlparser::ast::{ @@ -341,7 +341,20 @@ impl SqlToRel<'_, S> { }; // After resolution, all arguments are positional - let inner = ScalarFunction::new_udf(fm, resolved_args); + let mut scalar_expr = RawScalarExpr { + func: fm, + args: resolved_args, + }; + + for planner in self.context_provider.get_expr_planners().iter() { + match planner.plan_scalar(scalar_expr)? { + PlannerResult::Planned(expr) => return Ok(expr), + PlannerResult::Original(expr) => scalar_expr = expr, + } + } + + let RawScalarExpr { func, args } = scalar_expr; + let inner = ScalarFunction::new_udf(func, args); if name.eq_ignore_ascii_case(inner.name()) { return Ok(Expr::ScalarFunction(inner)); diff --git a/datafusion/sql/tests/sql_integration.rs b/datafusion/sql/tests/sql_integration.rs index 00103bfd9f56a..9f57aaafb0686 100644 --- a/datafusion/sql/tests/sql_integration.rs +++ b/datafusion/sql/tests/sql_integration.rs @@ -35,9 +35,10 @@ use datafusion_expr::{ expr::{HigherOrderFunction, LambdaVariable, ScalarFunction}, lambda, logical_plan::LogicalPlan, + planner::{ExprPlanner, PlannerResult, RawScalarExpr}, test::function_stub::sum_udaf, }; -use datafusion_functions::{string, unicode}; +use datafusion_functions::{core as core_functions, string, unicode}; use datafusion_sql::{ parser::DFParser, planner::{NullOrdering, ParserOptions, PlannerContext, SqlToRel}, @@ -836,6 +837,42 @@ fn select_scalar_func_with_literal_no_relation() { ); } +#[derive(Debug)] +struct SingleArgumentCoalescePlanner; + +impl ExprPlanner for SingleArgumentCoalescePlanner { + fn plan_scalar(&self, expr: RawScalarExpr) -> Result> { + if expr.func.name() == "coalesce" + && let [arg] = expr.args.as_slice() + { + Ok(PlannerResult::Planned(arg.clone())) + } else { + Ok(PlannerResult::Original(expr)) + } + } +} + +#[test] +fn select_scalar_func_with_expr_planner() -> Result<()> { + let state = mock_session_state() + .with_scalar_function(core_functions::coalesce()) + .with_expr_planner(Arc::new(SingleArgumentCoalescePlanner)); + let plan = logical_plan_from_state( + "SELECT coalesce(42)", + &GenericDialect {}, + ParserOptions::default(), + state, + )?; + assert_snapshot!( + plan, + @r" + Projection: Int64(42) + EmptyRelation: rows=1 + " + ); + Ok(()) +} + #[test] fn select_simple_filter() { let sql = "SELECT id, first_name, last_name \