diff --git a/datafusion/expr/src/type_coercion/functions.rs b/datafusion/expr/src/type_coercion/functions.rs index 8e86cb3685e90..ad7b086f39da8 100644 --- a/datafusion/expr/src/type_coercion/functions.rs +++ b/datafusion/expr/src/type_coercion/functions.rs @@ -816,13 +816,21 @@ fn get_valid_types( TypeSignature::Numeric(number) => { function_length_check(function_name, current_types.len(), *number)?; + let non_nulls = current_types + .iter() + .filter(|&t| NativeType::from(t) != NativeType::Null) + .collect::>(); + let mut valid_type = non_nulls + .first() + .copied() + .cloned() + // Fallback to default type if we don't know which type to coerced to + // f64 is chosen since most of the math functions utilize Signature::numeric, + // and their default type is double precision + .unwrap_or(DataType::Float64); // Find common numeric type among given types except string - let mut valid_type = current_types.first().unwrap().to_owned(); - for t in current_types.iter().skip(1) { + for &t in non_nulls.iter().skip(1) { let logical_data_type: NativeType = t.into(); - if logical_data_type == NativeType::Null { - continue; - } if !logical_data_type.is_numeric() { return plan_err!( @@ -840,12 +848,7 @@ fn get_valid_types( } let logical_data_type: NativeType = valid_type.clone().into(); - // Fallback to default type if we don't know which type to coerced to - // f64 is chosen since most of the math functions utilize Signature::numeric, - // and their default type is double precision - if logical_data_type == NativeType::Null { - valid_type = DataType::Float64; - } else if !logical_data_type.is_numeric() { + if !logical_data_type.is_numeric() { return plan_err!( "Function '{function_name}' expects Numeric but received {logical_data_type}" ); @@ -1451,6 +1454,13 @@ mod tests { ); assert_eq!(got, [DataType::Float64]); + let got = get_valid_types_flatten( + "test", + &TypeSignature::Numeric(2), + &[DataType::Null, DataType::Null], + ); + assert_eq!(got, [DataType::Float64, DataType::Float64]); + // Rejects non-numeric arg. let got = get_valid_types( "test", @@ -1463,6 +1473,21 @@ mod tests { "Function 'test' expects Numeric but received Timestamp(s)" ); + // Nulls should get ignored among other valid types + let got = get_valid_types_flatten( + "test", + &TypeSignature::Numeric(2), + &[DataType::Null, DataType::Int32], + ); + assert_eq!(got, [DataType::Int32, DataType::Int32]); + + let got = get_valid_types_flatten( + "test", + &TypeSignature::Numeric(2), + &[DataType::Int32, DataType::Null], + ); + assert_eq!(got, [DataType::Int32, DataType::Int32]); + Ok(()) } diff --git a/datafusion/sqllogictest/test_files/spark/math/mod.slt b/datafusion/sqllogictest/test_files/spark/math/mod.slt index 7f8ac6bad8579..ca3b745df9309 100644 --- a/datafusion/sqllogictest/test_files/spark/math/mod.slt +++ b/datafusion/sqllogictest/test_files/spark/math/mod.slt @@ -102,6 +102,21 @@ SELECT MOD(NULL::int, NULL::int) as mod_null_3; ---- NULL +query I +SELECT MOD(NULL, 3); +---- +NULL + +query I +SELECT MOD(10, NULL); +---- +NULL + +query R +SELECT MOD(NULL, NULL); +---- +NULL + # Special values: NaN and Infinity query R SELECT MOD(5.0::float8, 'NaN'::float8) as mod_nan_1; diff --git a/datafusion/sqllogictest/test_files/spark/math/pmod.slt b/datafusion/sqllogictest/test_files/spark/math/pmod.slt index 5b9c84bfecd45..9ac8606d715cc 100644 --- a/datafusion/sqllogictest/test_files/spark/math/pmod.slt +++ b/datafusion/sqllogictest/test_files/spark/math/pmod.slt @@ -142,10 +142,10 @@ SELECT arrow_typeof(pmod(2.5::decimal(3,1), NULL)); ---- Decimal128(3, 1) -# An untyped NULL beside a typed non-decimal argument takes the Numeric path, -# which cannot coerce the pair. `mod` rejects it the same way. -statement error DataFusion error: Error during planning: Internal error: Function 'pmod' failed to match any signature +query I SELECT pmod(NULL, 3::int); +---- +NULL # PMOD tests with large integers query I