diff --git a/datafusion/functions/src/string/lower.rs b/datafusion/functions/src/string/lower.rs index 88f2c800e9e0c..18e718526f2b1 100644 --- a/datafusion/functions/src/string/lower.rs +++ b/datafusion/functions/src/string/lower.rs @@ -15,14 +15,14 @@ // specific language governing permissions and limitations // under the License. -use arrow::datatypes::DataType; +use arrow::datatypes::{DataType, Field, FieldRef}; use crate::string::common::to_lower; use datafusion_common::Result; use datafusion_common::types::logical_string; use datafusion_expr::{ - Coercion, ColumnarValue, Documentation, EncodingPreservation, ScalarFunctionArgs, - ScalarUDFImpl, Signature, TypeSignatureClass, Volatility, + Coercion, ColumnarValue, Documentation, EncodingPreservation, ReturnFieldArgs, + ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignatureClass, Volatility, }; use datafusion_macros::user_doc; @@ -80,6 +80,14 @@ impl ScalarUDFImpl for LowerFunc { Ok(arg_types[0].clone()) } + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + let input = &args.arg_fields[0]; + Ok( + Field::new(self.name(), input.data_type().clone(), input.is_nullable()) + .into(), + ) + } + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { to_lower(&args.args, "lower") } @@ -97,6 +105,21 @@ mod tests { use datafusion_common::config::ConfigOptions; use std::sync::Arc; + #[test] + fn preserves_input_nullability() -> Result<()> { + let func = LowerFunc::new(); + for nullable in [false, true] { + let input = Field::new("input", DataType::Utf8, nullable); + let result = func.return_field_from_args(ReturnFieldArgs { + arg_fields: &[input.into()], + scalar_arguments: &[None], + })?; + assert_eq!(result.data_type(), &DataType::Utf8); + assert_eq!(result.is_nullable(), nullable); + } + Ok(()) + } + fn invoke_lower(input: ArrayRef) -> Result { let func = LowerFunc::new(); let data_type = input.data_type().clone(); diff --git a/datafusion/functions/src/string/upper.rs b/datafusion/functions/src/string/upper.rs index 789ab2c046203..c7aa430d978fe 100644 --- a/datafusion/functions/src/string/upper.rs +++ b/datafusion/functions/src/string/upper.rs @@ -16,12 +16,12 @@ // under the License. use crate::string::common::to_upper; -use arrow::datatypes::DataType; +use arrow::datatypes::{DataType, Field, FieldRef}; use datafusion_common::Result; use datafusion_common::types::logical_string; use datafusion_expr::{ - Coercion, ColumnarValue, Documentation, EncodingPreservation, ScalarFunctionArgs, - ScalarUDFImpl, Signature, TypeSignatureClass, Volatility, + Coercion, ColumnarValue, Documentation, EncodingPreservation, ReturnFieldArgs, + ScalarFunctionArgs, ScalarUDFImpl, Signature, TypeSignatureClass, Volatility, }; use datafusion_macros::user_doc; @@ -79,6 +79,14 @@ impl ScalarUDFImpl for UpperFunc { Ok(arg_types[0].clone()) } + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + let input = &args.arg_fields[0]; + Ok( + Field::new(self.name(), input.data_type().clone(), input.is_nullable()) + .into(), + ) + } + fn invoke_with_args(&self, args: ScalarFunctionArgs) -> Result { to_upper(&args.args, "upper") } @@ -96,6 +104,21 @@ mod tests { use datafusion_common::config::ConfigOptions; use std::sync::Arc; + #[test] + fn preserves_input_nullability() -> Result<()> { + let func = UpperFunc::new(); + for nullable in [false, true] { + let input = Field::new("input", DataType::Utf8, nullable); + let result = func.return_field_from_args(ReturnFieldArgs { + arg_fields: &[input.into()], + scalar_arguments: &[None], + })?; + assert_eq!(result.data_type(), &DataType::Utf8); + assert_eq!(result.is_nullable(), nullable); + } + Ok(()) + } + fn invoke_upper(input: ArrayRef) -> Result { let func = UpperFunc::new(); let data_type = input.data_type().clone(); diff --git a/datafusion/sqllogictest/test_files/functions.slt b/datafusion/sqllogictest/test_files/functions.slt index 4baee8ce6fbe5..fbabbb07fee0a 100644 --- a/datafusion/sqllogictest/test_files/functions.slt +++ b/datafusion/sqllogictest/test_files/functions.slt @@ -486,6 +486,30 @@ BAR Dictionary(Int32, Utf8) statement ok DROP TABLE upper_dictionary_test +# upper and lower preserve input nullability in the planned schema +query TTT +DESCRIBE SELECT upper(a) AS ua, lower(a) AS la, upper(b) AS ub +FROM (VALUES ('Ab', CAST(NULL AS VARCHAR))) AS t(a, b) +---- +ua Utf8 NO +la Utf8 NO +ub Utf8View YES + +# An outer join widens a non-nullable input before it reaches upper +query TTT +DESCRIBE SELECT upper(t.a) AS ua +FROM (SELECT 1 AS k) x +LEFT JOIN (VALUES ('Ab')) AS t(a) ON false +---- +ua Utf8 YES + +query T +SELECT upper(t.a) +FROM (SELECT 1 AS k) x +LEFT JOIN (VALUES ('Ab')) AS t(a) ON false +---- +NULL + query T SELECT btrim(' foo ') ----