From 6a008e6f72b6c768cd0bfb90951570856831261b Mon Sep 17 00:00:00 2001 From: Kevin-Li-2025 <2242139@qq.com> Date: Mon, 29 Jun 2026 14:59:24 +0800 Subject: [PATCH 1/4] Align scalar UDF return-field literal args Signed-off-by: Kevin-Li-2025 <2242139@qq.com> --- .../user_defined_scalar_functions.rs | 72 ++++++++++++++++++- datafusion/expr/src/expr_schema.rs | 36 +++++++--- 2 files changed, 99 insertions(+), 9 deletions(-) diff --git a/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs b/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs index b758aeb5209e8..6d4f7948bd4d9 100644 --- a/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs +++ b/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs @@ -40,13 +40,15 @@ use datafusion_common::{ DFSchema, DataFusionError, Result, ScalarValue, assert_batches_eq, assert_batches_sorted_eq, assert_contains, exec_datafusion_err, exec_err, not_impl_err, plan_err, + types::{NativeType, logical_int16}, }; use datafusion_expr::simplify::{ExprSimplifyResult, SimplifyContext}; use datafusion_expr::{ Accumulator, ColumnarValue, CreateFunction, CreateFunctionBody, LogicalPlanBuilder, OperateFunctionArg, ReturnFieldArgs, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, - Signature, Volatility, lit_with_metadata, + Signature, TypeSignatureClass, Volatility, lit_with_metadata, }; +use datafusion_expr_common::signature::Coercion; use datafusion_expr_common::signature::TypeSignature; use datafusion_functions_nested::range::range_udf; use parking_lot::Mutex; @@ -2078,6 +2080,74 @@ AS t(string, extension) Ok(()) } +/// https://github.com/apache/datafusion/issues/19982 +#[tokio::test] +async fn test_return_field_args_scalar_argument_types_match_arg_fields() -> Result<()> { + #[derive(Debug, PartialEq, Eq, Hash)] + struct TestUdf { + signature: Signature, + } + + impl Default for TestUdf { + fn default() -> Self { + Self { + signature: Signature::coercible( + vec![Coercion::new_implicit( + TypeSignatureClass::Native(logical_int16()), + vec![TypeSignatureClass::Numeric], + NativeType::Int16, + )], + Volatility::Immutable, + ), + } + } + } + + impl ScalarUDFImpl for TestUdf { + fn name(&self) -> &str { + "test_udf" + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, _arg_types: &[DataType]) -> Result { + unreachable!("return_field_from_args is implemented") + } + + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + assert_eq!(args.arg_fields.len(), 1); + assert_eq!(args.scalar_arguments.len(), 1); + assert_eq!( + args.arg_fields[0].data_type(), + &args.scalar_arguments[0] + .expect("literal argument") + .data_type() + ); + Ok( + Field::new(self.name(), args.arg_fields[0].data_type().clone(), true) + .into(), + ) + } + + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { + assert!(matches!( + args.args[0], + ColumnarValue::Scalar(ScalarValue::Int16(Some(_))) + )); + Ok(args.args[0].clone()) + } + } + + let ctx = SessionContext::new(); + ctx.register_udf(TestUdf::default().into()); + + ctx.sql("select test_udf(1)").await?.collect().await?; + + Ok(()) +} + /// https://github.com/apache/datafusion/issues/17422 #[tokio::test] async fn test_extension_metadata_preserve_in_subquery() -> Result<()> { diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 039bbad65a660..ac6b6d3728953 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -87,6 +87,30 @@ fn cast_output_field( Arc::new(f) } +fn scalar_arguments_for_fields( + args: &[Expr], + arg_fields: &[FieldRef], +) -> Result>> { + args.iter() + .zip(arg_fields) + .map(|(expr, field)| { + literal_scalar_value(expr) + .map(|sv| sv.cast_to(field.data_type())) + .transpose() + }) + .collect() +} + +fn literal_scalar_value(expr: &Expr) -> Option<&ScalarValue> { + match expr { + Expr::Literal(sv, _) => Some(sv), + Expr::Cast(Cast { expr, .. }) | Expr::TryCast(TryCast { expr, .. }) => { + literal_scalar_value(expr) + } + _ => None, + } +} + impl ExprSchemable for Expr { /// Returns the [arrow::datatypes::DataType] of the expression /// based on [ExprSchema] @@ -580,16 +604,12 @@ impl ExprSchemable for Expr { .collect::>>()?; let new_fields = verify_function_arguments(func.as_ref(), &fields)?; - let arguments = args - .iter() - .map(|e| match e { - Expr::Literal(sv, _) => Some(sv), - _ => None, - }) - .collect::>(); + let arguments = scalar_arguments_for_fields(args, &new_fields)?; + let argument_refs = + arguments.iter().map(Option::as_ref).collect::>(); let args = ReturnFieldArgs { arg_fields: &new_fields, - scalar_arguments: &arguments, + scalar_arguments: &argument_refs, }; func.return_field_from_args(args) From 053b1c0e8a10d37910d0a70b8d10064fc5ed7a88 Mon Sep 17 00:00:00 2001 From: Kevin-Li-2025 <2242139@qq.com> Date: Tue, 7 Jul 2026 09:52:55 +0800 Subject: [PATCH 2/4] Limit scalar literal coercion to implicit casts --- .../tests/dataframe/dataframe_functions.rs | 16 ++++++++ datafusion/expr/src/expr_schema.rs | 39 +++++++++++++++---- 2 files changed, 47 insertions(+), 8 deletions(-) diff --git a/datafusion/core/tests/dataframe/dataframe_functions.rs b/datafusion/core/tests/dataframe/dataframe_functions.rs index 2ada0411f4f8c..ad02c9e600c68 100644 --- a/datafusion/core/tests/dataframe/dataframe_functions.rs +++ b/datafusion/core/tests/dataframe/dataframe_functions.rs @@ -214,6 +214,22 @@ async fn test_fn_arrow_cast() -> Result<()> { Ok(()) } +#[tokio::test] +async fn test_fn_arrow_cast_requires_literal_type_arg() -> Result<()> { + let ctx = SessionContext::new(); + let df = ctx + .sql("select arrow_cast(1, cast('Utf8' as varchar))") + .await?; + let err = df.collect().await.unwrap_err(); + + assert!( + err.to_string() + .contains("arrow_cast requires its second argument") + ); + + Ok(()) +} + #[tokio::test] async fn test_nvl() -> Result<()> { let lit_null = lit(ScalarValue::Utf8(None)); diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index ac6b6d3728953..6648e80427764 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -93,19 +93,42 @@ fn scalar_arguments_for_fields( ) -> Result>> { args.iter() .zip(arg_fields) - .map(|(expr, field)| { - literal_scalar_value(expr) - .map(|sv| sv.cast_to(field.data_type())) - .transpose() - }) + .map(|(expr, field)| scalar_argument_for_field(expr, field)) .collect() } -fn literal_scalar_value(expr: &Expr) -> Option<&ScalarValue> { +fn scalar_argument_for_field( + expr: &Expr, + arg_field: &FieldRef, +) -> Result> { + let Some(sv) = literal_scalar_value(expr, arg_field) else { + return Ok(None); + }; + + // Preserve existing error behavior for functions that validate literal + // values themselves. This helper only normalizes scalar argument types when + // the planning-time cast is valid. + Ok(Some( + sv.cast_to(arg_field.data_type()) + .unwrap_or_else(|_| sv.clone()), + )) +} + +fn literal_scalar_value<'a>( + expr: &'a Expr, + arg_field: &FieldRef, +) -> Option<&'a ScalarValue> { match expr { Expr::Literal(sv, _) => Some(sv), - Expr::Cast(Cast { expr, .. }) | Expr::TryCast(TryCast { expr, .. }) => { - literal_scalar_value(expr) + Expr::Cast(Cast { expr, field }) + if field.data_type() == arg_field.data_type() => + { + match expr.as_ref() { + Expr::Literal(sv, _) if sv.data_type() != *arg_field.data_type() => { + Some(sv) + } + _ => None, + } } _ => None, } From 955c4e09906142559bbe7bf0d578e3f13467da9a Mon Sep 17 00:00:00 2001 From: Kevin-Li-2025 Date: Thu, 23 Jul 2026 23:44:26 +0800 Subject: [PATCH 3/4] fix: align scalar UDF literals with coerced fields --- .../tests/dataframe/dataframe_functions.rs | 9 +- .../user_defined_scalar_functions.rs | 48 +++++++---- datafusion/expr/src/expr_schema.rs | 86 +++++++++++-------- datafusion/expr/src/udf.rs | 4 + datafusion/functions/src/math/round.rs | 4 +- .../optimizer/src/analyzer/type_coercion.rs | 82 +++++++++++++++--- 6 files changed, 168 insertions(+), 65 deletions(-) diff --git a/datafusion/core/tests/dataframe/dataframe_functions.rs b/datafusion/core/tests/dataframe/dataframe_functions.rs index ad02c9e600c68..284acb361642a 100644 --- a/datafusion/core/tests/dataframe/dataframe_functions.rs +++ b/datafusion/core/tests/dataframe/dataframe_functions.rs @@ -217,10 +217,13 @@ async fn test_fn_arrow_cast() -> Result<()> { #[tokio::test] async fn test_fn_arrow_cast_requires_literal_type_arg() -> Result<()> { let ctx = SessionContext::new(); - let df = ctx + let err = match ctx .sql("select arrow_cast(1, cast('Utf8' as varchar))") - .await?; - let err = df.collect().await.unwrap_err(); + .await + { + Ok(df) => df.collect().await.unwrap_err(), + Err(err) => err, + }; assert!( err.to_string() diff --git a/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs b/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs index 6d4f7948bd4d9..b045d1df90d9d 100644 --- a/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs +++ b/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs @@ -20,8 +20,8 @@ use std::hash::{Hash, Hasher}; use std::sync::Arc; use arrow::array::{ - Array, ArrayRef, Float32Array, Float64Array, Int32Array, RecordBatch, StringArray, - builder::BooleanBuilder, cast::AsArray, + Array, ArrayRef, Float32Array, Float64Array, Int16Array, Int32Array, RecordBatch, + StringArray, builder::BooleanBuilder, cast::AsArray, }; use arrow::array::{Int8Array, UInt64Array, as_string_array, create_array, record_batch}; use arrow::compute::kernels::numeric::add; @@ -2085,12 +2085,16 @@ AS t(string, extension) async fn test_return_field_args_scalar_argument_types_match_arg_fields() -> Result<()> { #[derive(Debug, PartialEq, Eq, Hash)] struct TestUdf { + name: &'static str, + expect_literal: bool, signature: Signature, } - impl Default for TestUdf { - fn default() -> Self { + impl TestUdf { + fn new(name: &'static str, expect_literal: bool) -> Self { Self { + name, + expect_literal, signature: Signature::coercible( vec![Coercion::new_implicit( TypeSignatureClass::Native(logical_int16()), @@ -2105,7 +2109,7 @@ async fn test_return_field_args_scalar_argument_types_match_arg_fields() -> Resu impl ScalarUDFImpl for TestUdf { fn name(&self) -> &str { - "test_udf" + self.name } fn signature(&self) -> &Signature { @@ -2120,11 +2124,14 @@ async fn test_return_field_args_scalar_argument_types_match_arg_fields() -> Resu assert_eq!(args.arg_fields.len(), 1); assert_eq!(args.scalar_arguments.len(), 1); assert_eq!( - args.arg_fields[0].data_type(), - &args.scalar_arguments[0] - .expect("literal argument") - .data_type() + args.scalar_arguments[0].is_some(), + self.expect_literal, + "unexpected scalar argument for {}", + self.name ); + if let Some(scalar) = args.scalar_arguments[0] { + assert_eq!(args.arg_fields[0].data_type(), &scalar.data_type()); + } Ok( Field::new(self.name(), args.arg_fields[0].data_type().clone(), true) .into(), @@ -2132,18 +2139,29 @@ async fn test_return_field_args_scalar_argument_types_match_arg_fields() -> Resu } fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { - assert!(matches!( - args.args[0], - ColumnarValue::Scalar(ScalarValue::Int16(Some(_))) - )); + assert_eq!(args.args[0].data_type(), DataType::Int16); Ok(args.args[0].clone()) } } let ctx = SessionContext::new(); - ctx.register_udf(TestUdf::default().into()); + let coerced_literal_udf: ScalarUDF = TestUdf::new("coerced_literal_udf", true).into(); + let expression_arg_udf: ScalarUDF = TestUdf::new("expression_arg_udf", false).into(); + ctx.register_udf(coerced_literal_udf); + ctx.register_udf(expression_arg_udf.clone()); - ctx.sql("select test_udf(1)").await?.collect().await?; + ctx.sql("select coerced_literal_udf(1), coerced_literal_udf(NULL)") + .await? + .collect() + .await?; + let batch = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("a", DataType::Int16, false)])), + vec![Arc::new(Int16Array::from(vec![1]))], + )?; + ctx.read_batch(batch)? + .select(vec![expression_arg_udf.call(vec![col("a")])])? + .collect() + .await?; Ok(()) } diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 6648e80427764..25ac4f1f52093 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -90,46 +90,16 @@ fn cast_output_field( fn scalar_arguments_for_fields( args: &[Expr], arg_fields: &[FieldRef], -) -> Result>> { +) -> Vec> { args.iter() .zip(arg_fields) .map(|(expr, field)| scalar_argument_for_field(expr, field)) .collect() } -fn scalar_argument_for_field( - expr: &Expr, - arg_field: &FieldRef, -) -> Result> { - let Some(sv) = literal_scalar_value(expr, arg_field) else { - return Ok(None); - }; - - // Preserve existing error behavior for functions that validate literal - // values themselves. This helper only normalizes scalar argument types when - // the planning-time cast is valid. - Ok(Some( - sv.cast_to(arg_field.data_type()) - .unwrap_or_else(|_| sv.clone()), - )) -} - -fn literal_scalar_value<'a>( - expr: &'a Expr, - arg_field: &FieldRef, -) -> Option<&'a ScalarValue> { +fn scalar_argument_for_field(expr: &Expr, arg_field: &FieldRef) -> Option { match expr { - Expr::Literal(sv, _) => Some(sv), - Expr::Cast(Cast { expr, field }) - if field.data_type() == arg_field.data_type() => - { - match expr.as_ref() { - Expr::Literal(sv, _) if sv.data_type() != *arg_field.data_type() => { - Some(sv) - } - _ => None, - } - } + Expr::Literal(sv, _) => sv.cast_to(arg_field.data_type()).ok(), _ => None, } } @@ -627,7 +597,7 @@ impl ExprSchemable for Expr { .collect::>>()?; let new_fields = verify_function_arguments(func.as_ref(), &fields)?; - let arguments = scalar_arguments_for_fields(args, &new_fields)?; + let arguments = scalar_arguments_for_fields(args, &new_fields); let argument_refs = arguments.iter().map(Option::as_ref).collect::>(); let args = ReturnFieldArgs { @@ -850,6 +820,54 @@ mod tests { }}; } + #[test] + fn scalar_arguments_match_coerced_fields() { + let int16_field: FieldRef = Field::new("arg", DataType::Int16, true).into(); + + assert_eq!( + scalar_argument_for_field(&lit(1_i64), &int16_field), + Some(ScalarValue::Int16(Some(1))) + ); + assert_eq!( + scalar_argument_for_field(&lit(ScalarValue::Null), &int16_field), + Some(ScalarValue::Int16(None)) + ); + + let int32_list = ScalarValue::List(ScalarValue::new_list( + &[ScalarValue::Int32(Some(1))], + &DataType::Int32, + true, + )); + let int64_list_type = DataType::new_list(DataType::Int64, true); + let int64_list_field: FieldRef = + Field::new("arg", int64_list_type.clone(), true).into(); + assert_eq!( + scalar_argument_for_field(&lit(int32_list.clone()), &int64_list_field), + Some(int32_list.cast_to(&int64_list_type).unwrap()) + ); + } + + #[test] + fn scalar_arguments_exclude_expression_casts_and_invalid_values() { + let int16_field: FieldRef = Field::new("arg", DataType::Int16, true).into(); + let explicit_cast = Expr::Cast(Cast::new(Box::new(lit(1_i64)), DataType::Int16)); + let explicit_try_cast = + Expr::TryCast(TryCast::new(Box::new(lit(1_i64)), DataType::Int16)); + + assert_eq!( + scalar_argument_for_field(&explicit_cast, &int16_field), + None + ); + assert_eq!( + scalar_argument_for_field(&explicit_try_cast, &int16_field), + None + ); + assert_eq!( + scalar_argument_for_field(&lit("not an integer"), &int16_field), + None + ); + } + #[test] fn expr_schema_nullability() { let expr = col("foo").eq(lit(1)); diff --git a/datafusion/expr/src/udf.rs b/datafusion/expr/src/udf.rs index e206ce8b29108..616b227b88119 100644 --- a/datafusion/expr/src/udf.rs +++ b/datafusion/expr/src/udf.rs @@ -456,6 +456,10 @@ pub struct ReturnFieldArgs<'a> { /// /// If the argument `i` is not a scalar, it will be None /// + /// When present, the scalar value has the same data type as the corresponding + /// entry in [`Self::arg_fields`], including after implicit type coercion. + /// User-written `Cast` and `TryCast` expressions are not scalar arguments. + /// /// For example, if a function is called like `my_function(column_a, 5)` /// this field will be `[None, Some(ScalarValue::Int32(Some(5)))]` pub scalar_arguments: &'a [Option<&'a ScalarValue>], diff --git a/datafusion/functions/src/math/round.rs b/datafusion/functions/src/math/round.rs index 10500810a56b4..dec4856ba0a5a 100644 --- a/datafusion/functions/src/math/round.rs +++ b/datafusion/functions/src/math/round.rs @@ -242,8 +242,8 @@ impl ScalarUDFImpl for RoundFunc { // If decimal_places is a scalar literal, we can incorporate it into the output type // (scale reduction). Otherwise, keep the input scale as we can't pick a per-row scale. // - // Note: `scalar_arguments` contains the original literal values (pre-coercion), so - // integer literals may appear as Int64 even though the signature coerces them to Int32. + // `scalar_arguments` uses the coerced argument type, so an integer literal + // is Int32 here when the signature coerces it to Int32. let decimal_places: Option = match args.scalar_arguments.get(1) { None => Some(0), // No dp argument means default to 0 Some(None) => None, // dp is not a literal (e.g. column) diff --git a/datafusion/optimizer/src/analyzer/type_coercion.rs b/datafusion/optimizer/src/analyzer/type_coercion.rs index afd4e980b5424..8db406b811028 100644 --- a/datafusion/optimizer/src/analyzer/type_coercion.rs +++ b/datafusion/optimizer/src/analyzer/type_coercion.rs @@ -745,8 +745,11 @@ impl TreeNodeRewriter for TypeCoercionRewriter<'_> { Ok(Transformed::yes(Expr::Case(case))) } Expr::ScalarFunction(ScalarFunction { func, args }) => { - let new_expr = - coerce_arguments_for_signature(args, self.schema, func.as_ref())?; + let new_expr = coerce_scalar_function_arguments_for_signature( + args, + self.schema, + func.as_ref(), + )?; Ok(Transformed::yes(Expr::ScalarFunction( ScalarFunction::new_udf(func, new_expr), ))) @@ -1065,23 +1068,80 @@ fn coerce_arguments_for_signature( schema: &DFSchema, func: &F, ) -> Result> { - let current_fields = expressions - .iter() - .map(|e| e.to_field(schema).map(|(_, f)| f)) - .collect::>>()?; + let coerced_types = coerced_argument_types(&expressions, schema, func)?; - let coerced_types = fields_with_udf(¤t_fields, func)? + expressions .into_iter() - .map(|f| f.data_type().clone()) - .collect::>(); + .zip(coerced_types) + .map(|(expr, data_type)| expr.cast_to(&data_type, schema)) + .collect() +} + +/// Coerces scalar function arguments while materializing successful implicit +/// casts of literals. This preserves the literal for subsequent calls to +/// `return_field_from_args` without treating user-written `Cast` or `TryCast` +/// expressions as scalar arguments. +fn coerce_scalar_function_arguments_for_signature( + expressions: Vec, + schema: &DFSchema, + func: &F, +) -> Result> { + let coerced_types = coerced_argument_types(&expressions, schema, func)?; expressions .into_iter() - .enumerate() - .map(|(i, expr)| expr.cast_to(&coerced_types[i], schema)) + .zip(coerced_types) + .map(|(expr, data_type)| { + coerce_scalar_function_argument(expr, &data_type, schema) + }) .collect() } +fn coerced_argument_types( + expressions: &[Expr], + schema: &DFSchema, + func: &F, +) -> Result> { + let current_fields = expressions + .iter() + .map(|e| e.to_field(schema).map(|(_, f)| f)) + .collect::>>()?; + + fields_with_udf(¤t_fields, func).map(|fields| { + fields + .into_iter() + .map(|field| field.data_type().clone()) + .collect() + }) +} + +fn coerce_scalar_function_argument( + expr: Expr, + data_type: &DataType, + schema: &DFSchema, +) -> Result { + if matches!(&expr, Expr::Cast(_) | Expr::TryCast(_)) + && expr.get_type(schema)? == *data_type + { + return Ok(expr); + } + + let Expr::Literal(value, metadata) = expr else { + return expr.cast_to(data_type, schema); + }; + + if value.data_type() != *data_type + && let Ok(value) = value.cast_to(data_type) + { + return Ok(Expr::Literal(value, metadata)); + } + + // A failed value cast remains an expression cast so execution produces the + // same error as before. Since it is no longer a literal at the coerced type, + // it is reported as `None` in `ReturnFieldArgs::scalar_arguments`. + Expr::Literal(value, metadata).cast_to(data_type, schema) +} + fn coerce_case_expression(case: Case, schema: &DFSchema) -> Result { // Given expressions like: // From 6e56abb942d317fd6d84425a128c8f60adc915b0 Mon Sep 17 00:00:00 2001 From: Kevin-Li-2025 Date: Thu, 30 Jul 2026 02:39:24 +0800 Subject: [PATCH 4/4] test: address scalar argument review feedback Signed-off-by: Kevin-Li-2025 --- .../tests/dataframe/dataframe_functions.rs | 19 ------------------- .../user_defined_scalar_functions.rs | 3 ++- datafusion/expr/src/udf.rs | 4 ---- datafusion/functions/src/math/round.rs | 3 --- .../test_files/string_numeric_coercion.slt | 4 ++++ 5 files changed, 6 insertions(+), 27 deletions(-) diff --git a/datafusion/core/tests/dataframe/dataframe_functions.rs b/datafusion/core/tests/dataframe/dataframe_functions.rs index 284acb361642a..2ada0411f4f8c 100644 --- a/datafusion/core/tests/dataframe/dataframe_functions.rs +++ b/datafusion/core/tests/dataframe/dataframe_functions.rs @@ -214,25 +214,6 @@ async fn test_fn_arrow_cast() -> Result<()> { Ok(()) } -#[tokio::test] -async fn test_fn_arrow_cast_requires_literal_type_arg() -> Result<()> { - let ctx = SessionContext::new(); - let err = match ctx - .sql("select arrow_cast(1, cast('Utf8' as varchar))") - .await - { - Ok(df) => df.collect().await.unwrap_err(), - Err(err) => err, - }; - - assert!( - err.to_string() - .contains("arrow_cast requires its second argument") - ); - - Ok(()) -} - #[tokio::test] async fn test_nvl() -> Result<()> { let lit_null = lit(ScalarValue::Utf8(None)); diff --git a/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs b/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs index b045d1df90d9d..e39f75ccc74c0 100644 --- a/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs +++ b/datafusion/core/tests/user_defined/user_defined_scalar_functions.rs @@ -2129,7 +2129,8 @@ async fn test_return_field_args_scalar_argument_types_match_arg_fields() -> Resu "unexpected scalar argument for {}", self.name ); - if let Some(scalar) = args.scalar_arguments[0] { + if self.expect_literal { + let scalar = args.scalar_arguments[0].unwrap(); assert_eq!(args.arg_fields[0].data_type(), &scalar.data_type()); } Ok( diff --git a/datafusion/expr/src/udf.rs b/datafusion/expr/src/udf.rs index 616b227b88119..e206ce8b29108 100644 --- a/datafusion/expr/src/udf.rs +++ b/datafusion/expr/src/udf.rs @@ -456,10 +456,6 @@ pub struct ReturnFieldArgs<'a> { /// /// If the argument `i` is not a scalar, it will be None /// - /// When present, the scalar value has the same data type as the corresponding - /// entry in [`Self::arg_fields`], including after implicit type coercion. - /// User-written `Cast` and `TryCast` expressions are not scalar arguments. - /// /// For example, if a function is called like `my_function(column_a, 5)` /// this field will be `[None, Some(ScalarValue::Int32(Some(5)))]` pub scalar_arguments: &'a [Option<&'a ScalarValue>], diff --git a/datafusion/functions/src/math/round.rs b/datafusion/functions/src/math/round.rs index dec4856ba0a5a..120b07f4eb8ca 100644 --- a/datafusion/functions/src/math/round.rs +++ b/datafusion/functions/src/math/round.rs @@ -241,9 +241,6 @@ impl ScalarUDFImpl for RoundFunc { // If decimal_places is a scalar literal, we can incorporate it into the output type // (scale reduction). Otherwise, keep the input scale as we can't pick a per-row scale. - // - // `scalar_arguments` uses the coerced argument type, so an integer literal - // is Int32 here when the signature coerces it to Int32. let decimal_places: Option = match args.scalar_arguments.get(1) { None => Some(0), // No dp argument means default to 0 Some(None) => None, // dp is not a literal (e.g. column) diff --git a/datafusion/sqllogictest/test_files/string_numeric_coercion.slt b/datafusion/sqllogictest/test_files/string_numeric_coercion.slt index 1567a149bcdf4..885bb487529f0 100644 --- a/datafusion/sqllogictest/test_files/string_numeric_coercion.slt +++ b/datafusion/sqllogictest/test_files/string_numeric_coercion.slt @@ -553,6 +553,10 @@ SELECT arrow_cast('hello', 'RunEndEncoded("run_ends": non-null Int32, "values": ---- true +# The type argument to arrow_cast must remain a literal after planning. +statement error arrow_cast requires its second argument +SELECT arrow_cast(1, cast('Utf8' as varchar)); + # ------------------------------------------------- # Cleanup # -------------------------------------------------