From d279c6244fe84c9db63c27cc172fbc3e468b8aea Mon Sep 17 00:00:00 2001 From: osipovartem Date: Mon, 14 Sep 2026 07:14:40 +0300 Subject: [PATCH] Resolve scalar and aggregate UDF overloads by arity --- datafusion/expr-common/src/signature.rs | 48 ++++++++++++++++++ datafusion/sql/src/expr/function.rs | 14 +++++- datafusion/sql/tests/sql_integration.rs | 67 +++++++++++++++++++++++++ 3 files changed, 127 insertions(+), 2 deletions(-) diff --git a/datafusion/expr-common/src/signature.rs b/datafusion/expr-common/src/signature.rs index f0010f0a05014..d8d42d54a8afb 100644 --- a/datafusion/expr-common/src/signature.rs +++ b/datafusion/expr-common/src/signature.rs @@ -871,6 +871,35 @@ impl TypeSignature { } } + /// Returns whether this signature can accept `argument_count` arguments. + /// + /// User-defined signatures may implement arbitrary arity validation, so they + /// are conservatively treated as accepting any argument count here. + pub fn supports_argument_count(&self, argument_count: usize) -> bool { + match self { + TypeSignature::Variadic(_) | TypeSignature::VariadicAny => argument_count > 0, + TypeSignature::UserDefined => true, + TypeSignature::Uniform(count, _) + | TypeSignature::Comparable(count) + | TypeSignature::Any(count) + | TypeSignature::Numeric(count) + | TypeSignature::String(count) => *count == argument_count, + TypeSignature::Exact(types) => types.len() == argument_count, + TypeSignature::Coercible(coercions) => coercions.len() == argument_count, + TypeSignature::OneOf(signatures) => signatures + .iter() + .any(|signature| signature.supports_argument_count(argument_count)), + TypeSignature::ArraySignature(ArrayFunctionSignature::Array { + arguments, + .. + }) => arguments.len() == argument_count, + TypeSignature::ArraySignature( + ArrayFunctionSignature::RecursiveArray | ArrayFunctionSignature::MapArray, + ) => argument_count == 1, + TypeSignature::Nullary => argument_count == 0, + } + } + /// Returns true if the signature currently supports or used to supported 0 /// input arguments in a previous version of DataFusion. pub fn used_to_support_zero_arguments(&self) -> bool { @@ -1610,6 +1639,25 @@ mod tests { } } + #[test] + fn supports_argument_count_tests() { + let one_or_two = TypeSignature::OneOf(vec![ + TypeSignature::Any(1), + TypeSignature::Exact(vec![DataType::Utf8, DataType::Int64]), + ]); + assert!(one_or_two.supports_argument_count(1)); + assert!(one_or_two.supports_argument_count(2)); + assert!(!one_or_two.supports_argument_count(0)); + assert!(!one_or_two.supports_argument_count(3)); + + assert!(TypeSignature::VariadicAny.supports_argument_count(1)); + assert!(TypeSignature::VariadicAny.supports_argument_count(4)); + assert!(!TypeSignature::VariadicAny.supports_argument_count(0)); + assert!(TypeSignature::UserDefined.supports_argument_count(0)); + assert!(TypeSignature::Nullary.supports_argument_count(0)); + assert!(!TypeSignature::Nullary.supports_argument_count(1)); + } + #[test] fn type_signature_partial_ord() { // Test validates that partial ord is defined for TypeSignature and Signature. diff --git a/datafusion/sql/src/expr/function.rs b/datafusion/sql/src/expr/function.rs index 7b1b88424594e..b9a0a80fadc5a 100644 --- a/datafusion/sql/src/expr/function.rs +++ b/datafusion/sql/src/expr/function.rs @@ -345,8 +345,18 @@ impl SqlToRel<'_, S> { } } } - // User-defined function (UDF) should have precedence - if let Some(fm) = self.context_provider.get_function_meta(&name) { + // User-defined scalar functions take precedence unless their signature cannot + // accept this argument count and an aggregate overload exists. + let scalar_function = self.context_provider.get_function_meta(&name); + let prefer_aggregate_overload = + scalar_function.as_ref().is_some_and(|function| { + !function + .signature() + .type_signature + .supports_argument_count(args.len()) + && self.context_provider.get_aggregate_meta(&name).is_some() + }); + if let Some(fm) = scalar_function.filter(|_| !prefer_aggregate_overload) { let (args, arg_names) = self.function_args_to_expr_with_names(args, schema, planner_context)?; diff --git a/datafusion/sql/tests/sql_integration.rs b/datafusion/sql/tests/sql_integration.rs index fc1ef89331a84..0ac97600a9823 100644 --- a/datafusion/sql/tests/sql_integration.rs +++ b/datafusion/sql/tests/sql_integration.rs @@ -2345,6 +2345,73 @@ fn select_count_column() { ); } +#[test] +fn scalar_and_aggregate_udfs_can_share_a_name_when_arities_differ() { + let state = mock_session_state().with_scalar_function(Arc::new(make_udf( + "sum", + vec![DataType::Int32, DataType::Int32], + DataType::Int32, + ))); + let plan = logical_plan_from_state( + "SELECT sum(age) FROM person", + &GenericDialect {}, + ParserOptions::default(), + state, + ) + .unwrap(); + assert_snapshot!( + plan, + @r" + Projection: sum(person.age) + Aggregate: groupBy=[[]], aggr=[[sum(person.age)]] + TableScan: person + " + ); + + let state = mock_session_state().with_scalar_function(Arc::new(make_udf( + "sum", + vec![DataType::Int32, DataType::Int32], + DataType::Int32, + ))); + let plan = logical_plan_from_state( + "SELECT sum(age, age) FROM person", + &GenericDialect {}, + ParserOptions::default(), + state, + ) + .unwrap(); + assert_snapshot!( + plan, + @r" + Projection: sum(person.age, person.age) + TableScan: person + " + ); +} + +#[test] +fn scalar_udf_keeps_precedence_over_same_arity_aggregate_udf() { + let state = mock_session_state().with_scalar_function(Arc::new(make_udf( + "sum", + vec![DataType::Int32], + DataType::Int32, + ))); + let plan = logical_plan_from_state( + "SELECT sum(age) FROM person", + &GenericDialect {}, + ParserOptions::default(), + state, + ) + .unwrap(); + assert_snapshot!( + plan, + @r" + Projection: sum(person.age) + TableScan: person + " + ); +} + #[test] fn aggregate_expr_planner_can_resolve_qualified_wildcard_from_schema() { let state =