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
48 changes: 48 additions & 0 deletions datafusion/expr-common/src/signature.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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.
Expand Down
14 changes: 12 additions & 2 deletions datafusion/sql/src/expr/function.rs
Original file line number Diff line number Diff line change
Expand Up @@ -345,8 +345,18 @@ impl<S: ContextProvider> 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)?;

Expand Down
67 changes: 67 additions & 0 deletions datafusion/sql/tests/sql_integration.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 =
Expand Down
Loading