From cef0300e01bb838b7da9c3614428420c5b89286e Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 00:21:43 +0300 Subject: [PATCH 01/16] Preserve common CASE result field metadata without recursive planning --- datafusion/expr/src/expr_schema.rs | 246 +++++++++++++++++- .../optimizer/src/analyzer/type_coercion.rs | 28 ++ 2 files changed, 271 insertions(+), 3 deletions(-) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 8aaca2b72b2c6..79d321978289f 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -18,7 +18,7 @@ use super::{Between, Expr, Like, predicate_bounds}; use crate::ValueOrLambda; use crate::expr::{ - AggregateFunction, AggregateFunctionParams, Alias, BinaryExpr, Cast, InList, + AggregateFunction, AggregateFunctionParams, Alias, BinaryExpr, Case, Cast, InList, InSubquery, Lambda, Placeholder, ScalarFunction, TryCast, Unnest, WindowFunction, WindowFunctionParams, in_subquery_tuple_values, }; @@ -37,7 +37,7 @@ use arrow_schema::extension::{EXTENSION_TYPE_METADATA_KEY, EXTENSION_TYPE_NAME_K use datafusion_common::datatype::FieldExt; use datafusion_common::{ Column, DataFusionError, ExprSchema, Result, ScalarValue, Spans, TableReference, - not_impl_err, plan_datafusion_err, plan_err, + internal_err, not_impl_err, plan_datafusion_err, plan_err, }; use datafusion_expr_common::type_coercion::binary::BinaryTypeCoercer; use datafusion_functions_window_common::field::WindowUDFFieldArgs; @@ -138,6 +138,146 @@ fn scalar_argument_for_field(expr: &Expr, arg_field: &FieldRef) -> Option Result { + enum Work<'a> { + Visit(&'a Expr), + FinishCase(&'a Case), + FinishCast(&'a FieldRef, bool), + FinishAlias(Option<&'a FieldMetadata>), + } + + struct BranchField { + field: FieldRef, + certainly_null: bool, + } + + fn schedule_case<'a>(case: &'a Case, work: &mut Vec>) { + work.push(Work::FinishCase(case)); + work.extend(case.else_expr.iter().map(|expr| Work::Visit(expr))); + work.extend( + case.when_then_expr + .iter() + .rev() + .map(|(_, then_expr)| Work::Visit(then_expr)), + ); + } + + let mut work = Vec::new(); + let mut fields: Vec = Vec::new(); + schedule_case(case, &mut work); + while let Some(item) = work.pop() { + match item { + Work::Visit(expr) => match expr { + Expr::Case(nested) => schedule_case(nested, &mut work), + Expr::Cast(cast) => { + work.push(Work::FinishCast(&cast.field, false)); + work.push(Work::Visit(&cast.expr)); + } + Expr::TryCast(cast) => { + work.push(Work::FinishCast(&cast.field, true)); + work.push(Work::Visit(&cast.expr)); + } + Expr::Alias(alias) => { + work.push(Work::FinishAlias(alias.metadata.as_ref())); + work.push(Work::Visit(&alias.expr)); + } + Expr::Negative(inner) => work.push(Work::Visit(inner)), + _ => fields.push(BranchField { + field: expr.to_field(schema)?.1, + certainly_null: matches!( + unwrap_certainly_null_expr(expr), + Expr::Literal(value, _) if value.is_null() + ), + }), + }, + Work::FinishCast(target, force_nullable) => { + let Some(source) = fields.pop() else { + return internal_err!("Missing CASE cast input field"); + }; + fields.push(BranchField { + field: cast_output_field(&source.field, target, force_nullable), + certainly_null: source.certainly_null, + }); + } + Work::FinishAlias(metadata) => { + let Some(source) = fields.pop() else { + return internal_err!("Missing CASE alias input field"); + }; + let mut combined = source.field.metadata().clone(); + if let Some(metadata) = metadata { + combined.extend(metadata.to_hashmap()); + } + fields.push(BranchField { + field: Arc::new( + source.field.as_ref().clone().with_metadata(combined), + ), + certainly_null: source.certainly_null, + }); + } + Work::FinishCase(case) => { + let count = + case.when_then_expr.len() + usize::from(case.else_expr.is_some()); + if fields.len() < count { + return internal_err!("Missing CASE result fields"); + } + let start = fields.len() - count; + let mut then_type = DataType::Null; + let mut else_type = DataType::Null; + let mut branch_type = None; + let mut metadata = None; + let mut conflict = false; + let mut certainly_null = true; + for (index, branch) in fields.drain(start..).enumerate() { + let data_type = branch.field.data_type(); + if index < case.when_then_expr.len() { + if then_type.is_null() && !data_type.is_null() { + then_type = data_type.clone(); + } + } else { + else_type = data_type.clone(); + } + certainly_null &= branch.certainly_null; + if data_type.is_null() + || (branch.certainly_null && branch.field.metadata().is_empty()) + { + continue; + } + if branch_type.as_ref().is_some_and(|other| other != data_type) + || branch.field.metadata().is_empty() + || metadata + .as_ref() + .is_some_and(|other| other != branch.field.metadata()) + { + conflict = true; + } + branch_type.get_or_insert_with(|| data_type.clone()); + metadata.get_or_insert_with(|| branch.field.metadata().clone()); + } + let data_type = if then_type.is_null() { + else_type + } else { + then_type + }; + if conflict || branch_type.as_ref() != Some(&data_type) { + metadata = None; + } + fields.push(BranchField { + field: Arc::new( + Field::new("", data_type, true) + .with_metadata(metadata.unwrap_or_default()), + ), + certainly_null, + }); + } + } + } + let Some(result) = fields.pop() else { + return internal_err!("Missing CASE output field"); + }; + Ok(result.field) +} + impl ExprSchemable for Expr { /// Returns the [arrow::datatypes::DataType] of the expression /// based on [ExprSchema] @@ -678,11 +818,40 @@ impl ExprSchemable for Expr { Expr::LambdaVariable(LambdaVariable { field: Some(field), .. }) => Ok(Arc::clone(field).renamed(&schema_name)), + Expr::Case(case) => { + let data_type = self.get_type(schema)?; + let nullable = self.nullable(schema)?; + if let Some((_, first_result)) = case.when_then_expr.first() + && matches!( + first_result.as_ref(), + Expr::Column(_) | Expr::Literal(_, _) + ) + { + let field = first_result.to_field(schema)?.1; + if !field.data_type().is_null() + && !matches!(first_result.as_ref(), Expr::Literal(value, _) if value.is_null()) + && field.metadata().is_empty() + { + return Ok(( + relation, + Arc::new(Field::new(&schema_name, data_type, nullable)), + )); + } + } + let branch_field = case_field_metadata(case, schema)?; + let metadata = if branch_field.data_type() == &data_type { + branch_field.metadata().clone() + } else { + Default::default() + }; + Ok(Arc::new( + Field::new(&schema_name, data_type, nullable).with_metadata(metadata), + )) + } Expr::Like(_) | Expr::SimilarTo(_) | Expr::Not(_) | Expr::Between(_) - | Expr::Case(_) | Expr::InList(_) | Expr::InSubquery(_) | Expr::SetComparison(_) @@ -1065,6 +1234,13 @@ mod tests { assert_not_nullable(&e, &nullable_schema); assert_not_nullable(&e, ¬_nullable_schema); + let varchar_schema = MockExprSchema::new() + .with_data_type(DataType::Utf8) + .with_nullable(true); + let try_cast = Expr::TryCast(TryCast::new(Box::new(col("x")), DataType::Int32)); + let e = when(col("x").is_not_null(), try_cast).otherwise(lit(0))?; + assert_nullable(&e, &varchar_schema); + // CASE WHEN NOT x IS NULL THEN x ELSE 0 let e = when(not(col("x").is_null()), col("x")).otherwise(lit(0))?; assert_not_nullable(&e, &nullable_schema); @@ -1257,6 +1433,70 @@ mod tests { assert_eq!(meta, outer_ref.metadata(&schema).unwrap()); } + #[test] + fn test_case_field_metadata() -> Result<()> { + let shared = HashMap::from([("type".to_string(), "structured".to_string())]); + let different = HashMap::from([("type".to_string(), "other".to_string())]); + let schema = DFSchema::from_unqualified_fields( + vec![ + Field::new("a", DataType::Int32, false).with_metadata(shared.clone()), + Field::new("b", DataType::Int32, false).with_metadata(shared.clone()), + Field::new("c", DataType::Int32, false).with_metadata(different), + Field::new("d", DataType::Int32, false), + ] + .into(), + HashMap::new(), + )?; + + let same = when(lit(true), col("a")).otherwise(col("b"))?; + assert_eq!(same.to_field(&schema)?.1.metadata(), &shared); + + let null_else = when(lit(true), col("a")).otherwise(lit(ScalarValue::Null))?; + assert_eq!(null_else.to_field(&schema)?.1.metadata(), &shared); + + let coerced_null_else = when(lit(true), col("a")) + .otherwise(lit(ScalarValue::Null).cast_to(&DataType::Int32, &schema)?)?; + assert_eq!(coerced_null_else.to_field(&schema)?.1.metadata(), &shared); + + let try_cast_null_else = when(lit(true), col("a")).otherwise(Expr::TryCast( + TryCast::new(Box::new(lit(ScalarValue::Null)), DataType::Int32), + ))?; + assert_eq!(try_cast_null_else.to_field(&schema)?.1.metadata(), &shared); + + let typed_null = Expr::Cast(Cast::new_from_field( + Box::new(lit(ScalarValue::Null)), + Arc::new(Field::new("", DataType::Int32, true).with_metadata(shared.clone())), + )); + let all_typed_null = when(lit(true), typed_null.clone()).otherwise(typed_null)?; + assert_eq!(all_typed_null.to_field(&schema)?.1.metadata(), &shared); + + let mut nested = col("a"); + for _ in 0..128 { + nested = when(lit(true), col("a")).otherwise(nested)?; + } + let nested = when(lit(true), nested).otherwise(col("b"))?; + assert_eq!(nested.to_field(&schema)?.1.metadata(), &shared); + + let mut cast_nested = col("a"); + for _ in 0..128 { + let inner = when(lit(true), col("a")).otherwise(cast_nested)?; + cast_nested = Expr::Cast(Cast::new(Box::new(inner), DataType::Int32)); + } + let cast_nested = when(lit(true), cast_nested).otherwise(col("b"))?; + assert_eq!(cast_nested.to_field(&schema)?.1.metadata(), &shared); + + let mismatched = when(lit(true), col("a")).otherwise(col("c"))?; + assert!(mismatched.to_field(&schema)?.1.metadata().is_empty()); + + let unmarked = when(lit(true), col("a")).otherwise(col("d"))?; + assert!(unmarked.to_field(&schema)?.1.metadata().is_empty()); + + let first_unmarked = when(lit(true), col("d")).otherwise(col("a"))?; + assert!(first_unmarked.to_field(&schema)?.1.metadata().is_empty()); + + Ok(()) + } + #[test] fn test_alias_metadata_is_preserved_in_field_metadata() { let schema = MockExprSchema::new().with_data_type(DataType::Int32); diff --git a/datafusion/optimizer/src/analyzer/type_coercion.rs b/datafusion/optimizer/src/analyzer/type_coercion.rs index 39310eec1f74f..060c9c783ae0e 100644 --- a/datafusion/optimizer/src/analyzer/type_coercion.rs +++ b/datafusion/optimizer/src/analyzer/type_coercion.rs @@ -3053,6 +3053,34 @@ mod test { } } + #[test] + fn test_case_coercion_preserves_matching_field_metadata() -> Result<()> { + let metadata = std::collections::HashMap::from([( + "structured_type".to_string(), + "ARRAY(NUMBER)".to_string(), + )]); + let schema = DFSchema::from_unqualified_fields( + vec![ + Field::new("structured", DataType::Int32, false) + .with_metadata(metadata.clone()), + ] + .into(), + std::collections::HashMap::new(), + )?; + let case = Case { + expr: None, + when_then_expr: vec![(Box::new(lit(true)), Box::new(col("structured")))], + else_expr: Some(Box::new(lit(ScalarValue::Null))), + }; + let coerced = coerce_case_expression(case, &schema, None)?; + assert!(matches!(coerced.else_expr.as_deref(), Some(Expr::Cast(_)))); + assert_eq!( + Expr::Case(coerced).to_field(&schema)?.1.metadata(), + &metadata + ); + Ok(()) + } + #[test] fn test_case_expression_coercion() -> Result<()> { let schema = Arc::new(DFSchema::from_unqualified_fields( From 89d9117019520ac81ace15b508d25a2cb202ac24 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 01:52:43 +0300 Subject: [PATCH 02/16] Avoid repeated CASE metadata planning through binary expressions --- datafusion/expr/src/expr_schema.rs | 38 ++++++++++++++++++++++++++++++ 1 file changed, 38 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 79d321978289f..1a249d1767428 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -143,6 +143,7 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result enum Work<'a> { Visit(&'a Expr), FinishCase(&'a Case), + FinishBinary(&'a BinaryExpr), FinishCast(&'a FieldRef, bool), FinishAlias(Option<&'a FieldMetadata>), } @@ -170,6 +171,11 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result match item { Work::Visit(expr) => match expr { Expr::Case(nested) => schedule_case(nested, &mut work), + Expr::BinaryExpr(binary) => { + work.push(Work::FinishBinary(binary)); + work.push(Work::Visit(&binary.right)); + work.push(Work::Visit(&binary.left)); + } Expr::Cast(cast) => { work.push(Work::FinishCast(&cast.field, false)); work.push(Work::Visit(&cast.expr)); @@ -191,6 +197,29 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result ), }), }, + Work::FinishBinary(binary) => { + let Some(right) = fields.pop() else { + return internal_err!("Missing CASE binary right field"); + }; + let Some(left) = fields.pop() else { + return internal_err!("Missing CASE binary left field"); + }; + let mut coercer = BinaryTypeCoercer::new( + left.field.data_type(), + &binary.op, + right.field.data_type(), + ); + coercer.set_lhs_spans(binary.left.spans().cloned().unwrap_or_default()); + coercer.set_rhs_spans(binary.right.spans().cloned().unwrap_or_default()); + let nullable = match binary.op { + Operator::IsDistinctFrom | Operator::IsNotDistinctFrom => false, + _ => left.field.is_nullable() || right.field.is_nullable(), + }; + fields.push(BranchField { + field: Arc::new(Field::new("", coercer.get_result_type()?, nullable)), + certainly_null: false, + }); + } Work::FinishCast(target, force_nullable) => { let Some(source) = fields.pop() else { return internal_err!("Missing CASE cast input field"); @@ -1485,6 +1514,15 @@ mod tests { let cast_nested = when(lit(true), cast_nested).otherwise(col("b"))?; assert_eq!(cast_nested.to_field(&schema)?.1.metadata(), &shared); + let mut binary_nested = col("a"); + for _ in 0..128 { + binary_nested = when(lit(true), binary_nested + lit(1)).otherwise(lit(0))?; + } + let binary_nested = when(lit(true), binary_nested).otherwise(col("b"))?; + let binary_field = binary_nested.to_field(&schema)?.1; + assert_eq!(binary_field.data_type(), &DataType::Int32); + assert!(binary_field.metadata().is_empty()); + let mismatched = when(lit(true), col("a")).otherwise(col("c"))?; assert!(mismatched.to_field(&schema)?.1.metadata().is_empty()); From 78383134faa13a3e79174b22137c1119b8caf92b Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 01:58:17 +0300 Subject: [PATCH 03/16] Resolve scalar UDF CASE metadata iteratively --- datafusion/expr/src/expr_schema.rs | 65 ++++++++++++++++++++++++++++++ 1 file changed, 65 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 1a249d1767428..d428aaeb9851c 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -144,6 +144,7 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result Visit(&'a Expr), FinishCase(&'a Case), FinishBinary(&'a BinaryExpr), + FinishScalar(&'a ScalarFunction), FinishCast(&'a FieldRef, bool), FinishAlias(Option<&'a FieldMetadata>), } @@ -176,6 +177,10 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result work.push(Work::Visit(&binary.right)); work.push(Work::Visit(&binary.left)); } + Expr::ScalarFunction(function) => { + work.push(Work::FinishScalar(function)); + work.extend(function.args.iter().rev().map(Work::Visit)); + } Expr::Cast(cast) => { work.push(Work::FinishCast(&cast.field, false)); work.push(Work::Visit(&cast.expr)); @@ -220,6 +225,34 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result certainly_null: false, }); } + Work::FinishScalar(function) => { + if fields.len() < function.args.len() { + return internal_err!("Missing CASE scalar function arguments"); + } + let start = fields.len() - function.args.len(); + let input_fields = fields + .drain(start..) + .map(|arg| arg.field) + .collect::>(); + let coerced_fields = + verify_function_arguments(function.func.as_ref(), &input_fields)?; + let scalar_arguments = function + .args + .iter() + .map(|arg| match arg { + Expr::Literal(value, _) => Some(value), + _ => None, + }) + .collect::>(); + let field = function.func.return_field_from_args(ReturnFieldArgs { + arg_fields: &coerced_fields, + scalar_arguments: &scalar_arguments, + })?; + fields.push(BranchField { + field, + certainly_null: false, + }); + } Work::FinishCast(target, force_nullable) => { let Some(source) = fields.pop() else { return internal_err!("Missing CASE cast input field"); @@ -1523,6 +1556,38 @@ mod tests { assert_eq!(binary_field.data_type(), &DataType::Int32); assert!(binary_field.metadata().is_empty()); + let identity = Arc::new(crate::expr_fn::create_udf( + "identity", + vec![DataType::Int32], + DataType::Int32, + crate::Volatility::Immutable, + Arc::new(|_| { + Ok( + datafusion_expr_common::columnar_value::ColumnarValue::Scalar( + ScalarValue::Int32(Some(0)), + ), + ) + }), + )); + let mut scalar_nested = col("a"); + for _ in 0..128 { + let nested_case = when(lit(true), scalar_nested).otherwise(col("b"))?; + scalar_nested = Expr::ScalarFunction(ScalarFunction::new_udf( + Arc::clone(&identity), + vec![nested_case], + )); + } + let Expr::Case(scalar_case) = + when(lit(true), scalar_nested).otherwise(col("b"))? + else { + unreachable!(); + }; + assert!( + case_field_metadata(&scalar_case, &schema)? + .metadata() + .is_empty() + ); + let mismatched = when(lit(true), col("a")).otherwise(col("c"))?; assert!(mismatched.to_field(&schema)?.1.metadata().is_empty()); From e225c31b2b2cf60d7feb6f43ba5f746d0a402398 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 02:06:05 +0300 Subject: [PATCH 04/16] Limit CASE metadata propagation to supported branch forms --- datafusion/expr/src/expr_schema.rs | 94 ++++-------------------------- 1 file changed, 12 insertions(+), 82 deletions(-) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index d428aaeb9851c..a3fbe677d6567 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -138,13 +138,11 @@ fn scalar_argument_for_field(expr: &Expr, arg_field: &FieldRef) -> Option Result { enum Work<'a> { Visit(&'a Expr), FinishCase(&'a Case), - FinishBinary(&'a BinaryExpr), - FinishScalar(&'a ScalarFunction), FinishCast(&'a FieldRef, bool), FinishAlias(Option<&'a FieldMetadata>), } @@ -172,15 +170,6 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result match item { Work::Visit(expr) => match expr { Expr::Case(nested) => schedule_case(nested, &mut work), - Expr::BinaryExpr(binary) => { - work.push(Work::FinishBinary(binary)); - work.push(Work::Visit(&binary.right)); - work.push(Work::Visit(&binary.left)); - } - Expr::ScalarFunction(function) => { - work.push(Work::FinishScalar(function)); - work.extend(function.args.iter().rev().map(Work::Visit)); - } Expr::Cast(cast) => { work.push(Work::FinishCast(&cast.field, false)); work.push(Work::Visit(&cast.expr)); @@ -194,65 +183,20 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result work.push(Work::Visit(&alias.expr)); } Expr::Negative(inner) => work.push(Work::Visit(inner)), - _ => fields.push(BranchField { + Expr::Column(_) + | Expr::Literal(_, _) + | Expr::OuterReferenceColumn(_, _) + | Expr::ScalarVariable(_, _) + | Expr::Placeholder(_) + | Expr::LambdaVariable(_) => fields.push(BranchField { field: expr.to_field(schema)?.1, certainly_null: matches!( unwrap_certainly_null_expr(expr), Expr::Literal(value, _) if value.is_null() ), }), + _ => return Ok(Arc::new(Field::new("", DataType::Null, true))), }, - Work::FinishBinary(binary) => { - let Some(right) = fields.pop() else { - return internal_err!("Missing CASE binary right field"); - }; - let Some(left) = fields.pop() else { - return internal_err!("Missing CASE binary left field"); - }; - let mut coercer = BinaryTypeCoercer::new( - left.field.data_type(), - &binary.op, - right.field.data_type(), - ); - coercer.set_lhs_spans(binary.left.spans().cloned().unwrap_or_default()); - coercer.set_rhs_spans(binary.right.spans().cloned().unwrap_or_default()); - let nullable = match binary.op { - Operator::IsDistinctFrom | Operator::IsNotDistinctFrom => false, - _ => left.field.is_nullable() || right.field.is_nullable(), - }; - fields.push(BranchField { - field: Arc::new(Field::new("", coercer.get_result_type()?, nullable)), - certainly_null: false, - }); - } - Work::FinishScalar(function) => { - if fields.len() < function.args.len() { - return internal_err!("Missing CASE scalar function arguments"); - } - let start = fields.len() - function.args.len(); - let input_fields = fields - .drain(start..) - .map(|arg| arg.field) - .collect::>(); - let coerced_fields = - verify_function_arguments(function.func.as_ref(), &input_fields)?; - let scalar_arguments = function - .args - .iter() - .map(|arg| match arg { - Expr::Literal(value, _) => Some(value), - _ => None, - }) - .collect::>(); - let field = function.func.return_field_from_args(ReturnFieldArgs { - arg_fields: &coerced_fields, - scalar_arguments: &scalar_arguments, - })?; - fields.push(BranchField { - field, - certainly_null: false, - }); - } Work::FinishCast(target, force_nullable) => { let Some(source) = fields.pop() else { return internal_err!("Missing CASE cast input field"); @@ -1569,24 +1513,10 @@ mod tests { ) }), )); - let mut scalar_nested = col("a"); - for _ in 0..128 { - let nested_case = when(lit(true), scalar_nested).otherwise(col("b"))?; - scalar_nested = Expr::ScalarFunction(ScalarFunction::new_udf( - Arc::clone(&identity), - vec![nested_case], - )); - } - let Expr::Case(scalar_case) = - when(lit(true), scalar_nested).otherwise(col("b"))? - else { - unreachable!(); - }; - assert!( - case_field_metadata(&scalar_case, &schema)? - .metadata() - .is_empty() - ); + let scalar_result = + Expr::ScalarFunction(ScalarFunction::new_udf(identity, vec![col("a")])); + let scalar_case = when(lit(true), scalar_result).otherwise(col("b"))?; + assert!(scalar_case.to_field(&schema)?.1.metadata().is_empty()); let mismatched = when(lit(true), col("a")).otherwise(col("c"))?; assert!(mismatched.to_field(&schema)?.1.metadata().is_empty()); From eac3c681d5fc97e5828c1a7ef4cd1968740d0998 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 10:53:54 +0300 Subject: [PATCH 05/16] Cover CASE metadata through bounded function branches --- datafusion/expr/src/expr_schema.rs | 156 +++++++++++++++++++++++++++++ 1 file changed, 156 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index a3fbe677d6567..08857fd3c3f9c 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -35,6 +35,7 @@ use arrow::datatypes::FieldRef; use arrow::datatypes::{DataType, Field}; use arrow_schema::extension::{EXTENSION_TYPE_METADATA_KEY, EXTENSION_TYPE_NAME_KEY}; use datafusion_common::datatype::FieldExt; +use datafusion_common::tree_node::{TreeNode, TreeNodeRecursion}; use datafusion_common::{ Column, DataFusionError, ExprSchema, Result, ScalarValue, Spans, TableReference, internal_err, not_impl_err, plan_datafusion_err, plan_err, @@ -140,6 +141,23 @@ fn scalar_argument_for_field(expr: &Expr, arg_field: &FieldRef) -> Option Result { + // `to_field` can revisit nested aliases; delegate only shallow branches. + const MAX_EXACT_BRANCH_DEPTH: usize = 8; + + fn can_infer_exact_branch(expr: &Expr) -> Result { + let mut work = vec![(expr, 1)]; + while let Some((node, depth)) = work.pop() { + if matches!(node, Expr::Case(_)) || depth > MAX_EXACT_BRANCH_DEPTH { + return Ok(false); + } + node.apply_children(|child| { + work.push((child, depth + 1)); + Ok(TreeNodeRecursion::Continue) + })?; + } + Ok(true) + } + enum Work<'a> { Visit(&'a Expr), FinishCase(&'a Case), @@ -170,19 +188,53 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result match item { Work::Visit(expr) => match expr { Expr::Case(nested) => schedule_case(nested, &mut work), + Expr::Cast(cast) + if !cast.field.metadata().is_empty() + && can_infer_exact_branch(expr)? => + { + fields.push(BranchField { + field: expr.to_field(schema)?.1, + certainly_null: false, + }); + } Expr::Cast(cast) => { work.push(Work::FinishCast(&cast.field, false)); work.push(Work::Visit(&cast.expr)); } + Expr::TryCast(cast) + if !cast.field.metadata().is_empty() + && can_infer_exact_branch(expr)? => + { + fields.push(BranchField { + field: expr.to_field(schema)?.1, + certainly_null: false, + }); + } Expr::TryCast(cast) => { work.push(Work::FinishCast(&cast.field, true)); work.push(Work::Visit(&cast.expr)); } + Expr::Alias(alias) + if alias.metadata.is_some() && can_infer_exact_branch(expr)? => + { + fields.push(BranchField { + field: expr.to_field(schema)?.1, + certainly_null: false, + }); + } Expr::Alias(alias) => { work.push(Work::FinishAlias(alias.metadata.as_ref())); work.push(Work::Visit(&alias.expr)); } Expr::Negative(inner) => work.push(Work::Visit(inner)), + Expr::ScalarFunction(_) | Expr::HigherOrderFunction(_) + if can_infer_exact_branch(expr)? => + { + fields.push(BranchField { + field: expr.to_field(schema)?.1, + certainly_null: false, + }); + } Expr::Column(_) | Expr::Literal(_, _) | Expr::OuterReferenceColumn(_, _) @@ -1441,6 +1493,40 @@ mod tests { #[test] fn test_case_field_metadata() -> Result<()> { + #[derive(Debug, PartialEq, Eq, Hash)] + struct MarkedIdentity { + signature: crate::Signature, + } + + impl crate::ScalarUDFImpl for MarkedIdentity { + fn name(&self) -> &str { + "marked_identity" + } + + fn signature(&self) -> &crate::Signature { + &self.signature + } + + fn return_type(&self, _arg_types: &[DataType]) -> Result { + Ok(DataType::Int32) + } + + fn return_field_from_args(&self, args: ReturnFieldArgs) -> Result { + let input = &args.arg_fields[0]; + Ok(Arc::new( + Field::new("marked_identity", DataType::Int32, input.is_nullable()) + .with_metadata(input.metadata().clone()), + )) + } + + fn invoke_with_args( + &self, + _args: crate::ScalarFunctionArgs, + ) -> Result { + Ok(crate::ColumnarValue::Scalar(ScalarValue::Int32(Some(0)))) + } + } + let shared = HashMap::from([("type".to_string(), "structured".to_string())]); let different = HashMap::from([("type".to_string(), "other".to_string())]); let schema = DFSchema::from_unqualified_fields( @@ -1457,6 +1543,19 @@ mod tests { let same = when(lit(true), col("a")).otherwise(col("b"))?; assert_eq!(same.to_field(&schema)?.1.metadata(), &shared); + let no_else = when(lit(true), col("a")).end()?; + assert_eq!(no_else.to_field(&schema)?.1.metadata(), &shared); + + let multiple_when = when(lit(false), col("a")) + .when(lit(true), col("b")) + .otherwise(lit(ScalarValue::Null))?; + assert_eq!(multiple_when.to_field(&schema)?.1.metadata(), &shared); + + let conflicting_when = when(lit(false), col("a")) + .when(lit(true), col("c")) + .otherwise(col("b"))?; + assert!(conflicting_when.to_field(&schema)?.1.metadata().is_empty()); + let null_else = when(lit(true), col("a")).otherwise(lit(ScalarValue::Null))?; assert_eq!(null_else.to_field(&schema)?.1.metadata(), &shared); @@ -1518,6 +1617,63 @@ mod tests { let scalar_case = when(lit(true), scalar_result).otherwise(col("b"))?; assert!(scalar_case.to_field(&schema)?.1.metadata().is_empty()); + let marked = Arc::new(crate::ScalarUDF::from(MarkedIdentity { + signature: crate::Signature::uniform( + 1, + vec![DataType::Int32], + crate::Volatility::Immutable, + ), + })); + let marked_call = |arg| { + Expr::ScalarFunction(ScalarFunction::new_udf(Arc::clone(&marked), vec![arg])) + }; + let marked_case = + when(lit(true), marked_call(col("a"))).otherwise(marked_call(col("b")))?; + assert_eq!(marked_case.to_field(&schema)?.1.metadata(), &shared); + + let nested_arg = when(lit(true), col("a")).otherwise(col("b"))?; + let nested_call = when(lit(true), marked_call(nested_arg)).otherwise(col("b"))?; + assert!(nested_call.to_field(&schema)?.1.metadata().is_empty()); + + let binary = col("d") + lit(1); + let Expr::Alias(alias) = binary.clone().alias("marked") else { + unreachable!(); + }; + let marked_alias = + Expr::Alias(alias.with_metadata(Some(FieldMetadata::from(shared.clone())))); + let aliased_case = when(lit(true), marked_alias).otherwise(col("b"))?; + assert_eq!(aliased_case.to_field(&schema)?.1.metadata(), &shared); + + let marked_cast = Expr::Cast(Cast::new_from_field( + Box::new(binary.clone()), + Arc::new(Field::new("", DataType::Int32, true).with_metadata(shared.clone())), + )); + let cast_case = when(lit(true), marked_cast).otherwise(col("b"))?; + assert_eq!(cast_case.to_field(&schema)?.1.metadata(), &shared); + + let marked_try_cast = Expr::TryCast(TryCast::new_from_field( + Box::new(binary), + Arc::new(Field::new("", DataType::Int32, true).with_metadata(shared.clone())), + )); + let try_cast_case = when(lit(true), marked_try_cast).otherwise(col("b"))?; + assert_eq!(try_cast_case.to_field(&schema)?.1.metadata(), &shared); + + let mut deep_alias = col("a"); + for _ in 0..10 { + deep_alias = deep_alias.alias("nested"); + } + let deep_udf_case = + when(lit(true), marked_call(deep_alias.clone())).otherwise(col("b"))?; + assert!(deep_udf_case.to_field(&schema)?.1.metadata().is_empty()); + + let Expr::Alias(alias) = deep_alias.alias("marked") else { + unreachable!(); + }; + let deep_marked_alias = + Expr::Alias(alias.with_metadata(Some(FieldMetadata::from(shared.clone())))); + let deep_alias_case = when(lit(true), deep_marked_alias).otherwise(col("b"))?; + assert_eq!(deep_alias_case.to_field(&schema)?.1.metadata(), &shared); + let mismatched = when(lit(true), col("a")).otherwise(col("c"))?; assert!(mismatched.to_field(&schema)?.1.metadata().is_empty()); From 0f8c16537bdcf0600b094e417502f12d086e42ca Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 11:47:08 +0300 Subject: [PATCH 06/16] Preserve shallow nested CASE metadata through functions --- datafusion/expr/src/expr_schema.rs | 25 ++++++++++++++++++------- 1 file changed, 18 insertions(+), 7 deletions(-) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 08857fd3c3f9c..6ff6212a8c336 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -144,10 +144,12 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result // `to_field` can revisit nested aliases; delegate only shallow branches. const MAX_EXACT_BRANCH_DEPTH: usize = 8; - fn can_infer_exact_branch(expr: &Expr) -> Result { + fn can_infer_exact_branch(expr: &Expr, allow_nested_case: bool) -> Result { let mut work = vec![(expr, 1)]; while let Some((node, depth)) = work.pop() { - if matches!(node, Expr::Case(_)) || depth > MAX_EXACT_BRANCH_DEPTH { + if (!allow_nested_case && matches!(node, Expr::Case(_))) + || depth > MAX_EXACT_BRANCH_DEPTH + { return Ok(false); } node.apply_children(|child| { @@ -190,7 +192,7 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result Expr::Case(nested) => schedule_case(nested, &mut work), Expr::Cast(cast) if !cast.field.metadata().is_empty() - && can_infer_exact_branch(expr)? => + && can_infer_exact_branch(expr, false)? => { fields.push(BranchField { field: expr.to_field(schema)?.1, @@ -203,7 +205,7 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result } Expr::TryCast(cast) if !cast.field.metadata().is_empty() - && can_infer_exact_branch(expr)? => + && can_infer_exact_branch(expr, false)? => { fields.push(BranchField { field: expr.to_field(schema)?.1, @@ -215,7 +217,8 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result work.push(Work::Visit(&cast.expr)); } Expr::Alias(alias) - if alias.metadata.is_some() && can_infer_exact_branch(expr)? => + if alias.metadata.is_some() + && can_infer_exact_branch(expr, false)? => { fields.push(BranchField { field: expr.to_field(schema)?.1, @@ -228,7 +231,7 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result } Expr::Negative(inner) => work.push(Work::Visit(inner)), Expr::ScalarFunction(_) | Expr::HigherOrderFunction(_) - if can_infer_exact_branch(expr)? => + if can_infer_exact_branch(expr, true)? => { fields.push(BranchField { field: expr.to_field(schema)?.1, @@ -1633,7 +1636,15 @@ mod tests { let nested_arg = when(lit(true), col("a")).otherwise(col("b"))?; let nested_call = when(lit(true), marked_call(nested_arg)).otherwise(col("b"))?; - assert!(nested_call.to_field(&schema)?.1.metadata().is_empty()); + assert_eq!(nested_call.to_field(&schema)?.1.metadata(), &shared); + + let mut deep_nested_arg = when(lit(true), col("a")).otherwise(col("b"))?; + for _ in 0..8 { + deep_nested_arg = deep_nested_arg.alias("nested"); + } + let deep_nested_call = + when(lit(true), marked_call(deep_nested_arg)).otherwise(col("b"))?; + assert!(deep_nested_call.to_field(&schema)?.1.metadata().is_empty()); let binary = col("d") + lit(1); let Expr::Alias(alias) = binary.clone().alias("marked") else { From 1ca200b089237d5dffb7350cda3507515e42d4bc Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 12:03:26 +0300 Subject: [PATCH 07/16] Test additional CASE metadata branch forms --- datafusion/expr/src/expr_schema.rs | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 6ff6212a8c336..9f3cf4fac959e 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -1538,6 +1538,8 @@ mod tests { Field::new("b", DataType::Int32, false).with_metadata(shared.clone()), Field::new("c", DataType::Int32, false).with_metadata(different), Field::new("d", DataType::Int32, false), + Field::new("e", DataType::Boolean, false).with_metadata(shared.clone()), + Field::new("f", DataType::Boolean, false).with_metadata(shared.clone()), ] .into(), HashMap::new(), @@ -1559,6 +1561,24 @@ mod tests { .otherwise(col("b"))?; assert!(conflicting_when.to_field(&schema)?.1.metadata().is_empty()); + let marked_literal = Expr::Literal( + ScalarValue::Int32(Some(1)), + Some(FieldMetadata::from(shared.clone())), + ); + let literal_case = when(lit(true), marked_literal).otherwise(col("b"))?; + assert_eq!(literal_case.to_field(&schema)?.1.metadata(), &shared); + + let unmarked_literal = when(lit(true), col("a")).otherwise(lit(1_i32))?; + assert!(unmarked_literal.to_field(&schema)?.1.metadata().is_empty()); + + let negative_case = when(lit(true), Expr::Negative(Box::new(col("a")))) + .otherwise(Expr::Negative(Box::new(col("b"))))?; + assert_eq!(negative_case.to_field(&schema)?.1.metadata(), &shared); + + let unsupported_case = + when(lit(true), Expr::Not(Box::new(col("e")))).otherwise(col("f"))?; + assert!(unsupported_case.to_field(&schema)?.1.metadata().is_empty()); + let null_else = when(lit(true), col("a")).otherwise(lit(ScalarValue::Null))?; assert_eq!(null_else.to_field(&schema)?.1.metadata(), &shared); From 3a22022fa59bf6de3cea69cd41b3552855da99d9 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 13:49:32 +0300 Subject: [PATCH 08/16] Test CASE metadata across supported leaf expressions --- datafusion/expr/src/expr_schema.rs | 41 ++++++++++++++++++++++++++++++ 1 file changed, 41 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 9f3cf4fac959e..4b7244af0d7df 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -1717,6 +1717,47 @@ mod tests { Ok(()) } + #[test] + fn test_case_field_metadata_leaf_branches() -> Result<()> { + let metadata = HashMap::from([("type".to_string(), "structured".to_string())]); + let field = Arc::new( + Field::new("value", DataType::Int32, false).with_metadata(metadata.clone()), + ); + let schema = DFSchema::from_unqualified_fields( + vec![ + Field::new("value", DataType::Int32, false) + .with_metadata(metadata.clone()), + ] + .into(), + HashMap::new(), + )?; + + let branches = [ + Expr::OuterReferenceColumn(Arc::clone(&field), Column::from_name("outer")), + Expr::ScalarVariable(Arc::clone(&field), vec!["value".to_string()]), + Expr::Placeholder(Placeholder::new_with_field( + "$1".to_string(), + Some(Arc::clone(&field)), + )), + Expr::LambdaVariable(LambdaVariable::new("arg".into(), Some(field))), + ]; + for branch in branches { + let case = when(lit(true), col("value")).otherwise(branch)?; + assert_eq!(case.to_field(&schema)?.1.metadata(), &metadata); + } + + let null_first = + when(lit(true), lit(ScalarValue::Null)).otherwise(col("value"))?; + assert_eq!(null_first.to_field(&schema)?.1.metadata(), &metadata); + + let wrong_type = when(lit(true), col("value")).otherwise(Expr::Literal( + ScalarValue::Boolean(Some(false)), + Some(FieldMetadata::from(metadata)), + ))?; + assert!(wrong_type.to_field(&schema)?.1.metadata().is_empty()); + Ok(()) + } + #[test] fn test_alias_metadata_is_preserved_in_field_metadata() { let schema = MockExprSchema::new().with_data_type(DataType::Int32); From 74ac5de5bfd09326fcc33a4c47159ace5e1f80c3 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 14:00:49 +0300 Subject: [PATCH 09/16] Test CASE field metadata through SQL planning --- datafusion/core/tests/sql/select.rs | 73 +++++++++++++++++++++++++++++ 1 file changed, 73 insertions(+) diff --git a/datafusion/core/tests/sql/select.rs b/datafusion/core/tests/sql/select.rs index fae2fb5bc5e76..ecd59492c39d6 100644 --- a/datafusion/core/tests/sql/select.rs +++ b/datafusion/core/tests/sql/select.rs @@ -387,6 +387,79 @@ async fn test_query_parameters_with_metadata() -> Result<()> { Ok(()) } +#[tokio::test] +async fn test_sql_case_field_metadata() -> Result<()> { + let ctx = SessionContext::new(); + let metadata = + HashMap::from([("structured_type".to_string(), "ARRAY(NUMBER)".to_string())]); + let other_metadata = + HashMap::from([("structured_type".to_string(), "OBJECT".to_string())]); + let schema = Arc::new(Schema::new(vec![ + Field::new("flag", DataType::Boolean, false), + Field::new("a", DataType::Int32, false).with_metadata(metadata.clone()), + Field::new("b", DataType::Int32, false).with_metadata(metadata.clone()), + Field::new("c", DataType::Int32, false).with_metadata(other_metadata), + ])); + ctx.register_batch( + "t", + RecordBatch::try_new( + schema, + vec![ + Arc::new(BooleanArray::from(vec![true, false])), + Arc::new(Int32Array::from(vec![1, 2])), + Arc::new(Int32Array::from(vec![3, 4])), + Arc::new(Int32Array::from(vec![5, 6])), + ], + )?, + )?; + + let empty_metadata = HashMap::new(); + for (sql, preserve_metadata) in [ + ( + "SELECT CASE WHEN flag THEN a ELSE b END AS value FROM t", + true, + ), + ( + "SELECT CASE WHEN flag THEN a ELSE NULL END AS value FROM t", + true, + ), + ( + "SELECT CASE WHEN flag THEN NULL ELSE a END AS value FROM t", + true, + ), + ( + "SELECT CASE WHEN flag THEN a ELSE c END AS value FROM t", + false, + ), + ] { + let df = ctx.sql(sql).await?; + let optimized = ctx.state().optimize(df.logical_plan())?; + let expected_metadata = if preserve_metadata { + &metadata + } else { + &empty_metadata + }; + for field in [df.schema().field(0), optimized.schema().field(0)] { + assert_eq!(field.metadata(), expected_metadata); + } + if sql == "SELECT CASE WHEN flag THEN a ELSE b END AS value FROM t" { + let batches = df.collect().await?; + datafusion::assert_batches_eq!( + [ + "+-------+", + "| value |", + "+-------+", + "| 1 |", + "| 4 |", + "+-------+" + ], + &batches + ); + } + } + Ok(()) +} + #[tokio::test] async fn test_limit_offset_parameters_named() -> Result<()> { let ctx = SessionContext::new(); From cc221b5fdc0840612f7a2b5a56b049f579f2d7be Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 16:55:32 +0300 Subject: [PATCH 10/16] Test logical CASE field metadata for marked NULL casts (cherry picked from commit 816119e839317cca1b31ef02ad0e8e302e0c60d0) --- datafusion/expr/src/expr_schema.rs | 33 ++++++++++++++++++++++++++++++ 1 file changed, 33 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 4b7244af0d7df..507936c510e2c 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -1598,6 +1598,39 @@ mod tests { let all_typed_null = when(lit(true), typed_null.clone()).otherwise(typed_null)?; assert_eq!(all_typed_null.to_field(&schema)?.1.metadata(), &shared); + let binary_metadata = HashMap::from([( + EXTENSION_TYPE_NAME_KEY.to_string(), + "geoarrow.wkb".to_string(), + )]); + let binary_schema = DFSchema::from_unqualified_fields( + vec![ + Field::new("binary", DataType::LargeBinary, false) + .with_metadata(binary_metadata.clone()), + ] + .into(), + HashMap::new(), + )?; + let target_field: FieldRef = Arc::new( + Field::new("", DataType::Binary, true).with_metadata(binary_metadata), + ); + for marked_null in [ + Expr::Cast(Cast::new_from_field( + Box::new(lit(ScalarValue::Null)), + Arc::clone(&target_field), + )), + Expr::TryCast(TryCast::new_from_field( + Box::new(lit(ScalarValue::Null)), + Arc::clone(&target_field), + )), + ] { + let nested_all_null = + when(lit(true), marked_null).otherwise(lit(ScalarValue::Null))?; + let type_only_cast = + Expr::Cast(Cast::new(Box::new(nested_all_null), DataType::LargeBinary)); + let outer = when(lit(true), type_only_cast).otherwise(col("binary"))?; + assert!(outer.to_field(&binary_schema)?.1.metadata().is_empty()); + } + let mut nested = col("a"); for _ in 0..128 { nested = when(lit(true), col("a")).otherwise(nested)?; From d033855761de7315b4a58b98710f59937e8e0ecf Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 17:03:29 +0300 Subject: [PATCH 11/16] Test logical CASE metadata depth for marked scalar functions (cherry picked from commit 185ae25f63f622ece719cc6a95e3d8db3d6135ed) --- datafusion/expr/src/expr_schema.rs | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 507936c510e2c..56741d28c8d6e 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -1687,6 +1687,19 @@ mod tests { when(lit(true), marked_call(col("a"))).otherwise(marked_call(col("b")))?; assert_eq!(marked_case.to_field(&schema)?.1.metadata(), &shared); + let mut deep_function = col("a"); + for _ in 0..9 { + deep_function = marked_call(deep_function); + } + let deep_function_case = when(lit(true), deep_function).otherwise(col("b"))?; + assert!( + deep_function_case + .to_field(&schema)? + .1 + .metadata() + .is_empty() + ); + let nested_arg = when(lit(true), col("a")).otherwise(col("b"))?; let nested_call = when(lit(true), marked_call(nested_arg)).otherwise(col("b"))?; assert_eq!(nested_call.to_field(&schema)?.1.metadata(), &shared); From 3a8631b7fd7a2d9f2188158733e1d9e1655881e5 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 18:13:44 +0300 Subject: [PATCH 12/16] Test simple CASE field metadata through SQL --- datafusion/core/tests/sql/select.rs | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/datafusion/core/tests/sql/select.rs b/datafusion/core/tests/sql/select.rs index ecd59492c39d6..428948486f6e8 100644 --- a/datafusion/core/tests/sql/select.rs +++ b/datafusion/core/tests/sql/select.rs @@ -419,6 +419,10 @@ async fn test_sql_case_field_metadata() -> Result<()> { "SELECT CASE WHEN flag THEN a ELSE b END AS value FROM t", true, ), + ( + "SELECT CASE flag WHEN TRUE THEN a ELSE b END AS value FROM t", + true, + ), ( "SELECT CASE WHEN flag THEN a ELSE NULL END AS value FROM t", true, From c9dbd638f01c380b9a0c91f058af8e198a659c7a Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 21:34:35 +0300 Subject: [PATCH 13/16] test: cover CASE higher-order and extension metadata --- datafusion/core/tests/sql/select.rs | 19 ++++++-- datafusion/expr/src/expr_schema.rs | 73 +++++++++++++++++++++++++++++ 2 files changed, 88 insertions(+), 4 deletions(-) diff --git a/datafusion/core/tests/sql/select.rs b/datafusion/core/tests/sql/select.rs index 428948486f6e8..397115526814d 100644 --- a/datafusion/core/tests/sql/select.rs +++ b/datafusion/core/tests/sql/select.rs @@ -18,6 +18,7 @@ use std::collections::HashMap; use super::*; +use arrow_schema::extension::EXTENSION_TYPE_NAME_KEY; use datafusion_common::{ParamValues, ScalarValue, metadata::ScalarAndMetadata}; use insta::assert_snapshot; @@ -390,10 +391,20 @@ async fn test_query_parameters_with_metadata() -> Result<()> { #[tokio::test] async fn test_sql_case_field_metadata() -> Result<()> { let ctx = SessionContext::new(); - let metadata = - HashMap::from([("structured_type".to_string(), "ARRAY(NUMBER)".to_string())]); - let other_metadata = - HashMap::from([("structured_type".to_string(), "OBJECT".to_string())]); + let metadata = HashMap::from([ + ("structured_type".to_string(), "ARRAY(NUMBER)".to_string()), + ( + EXTENSION_TYPE_NAME_KEY.to_string(), + "example.case".to_string(), + ), + ]); + let other_metadata = HashMap::from([ + ("structured_type".to_string(), "OBJECT".to_string()), + ( + EXTENSION_TYPE_NAME_KEY.to_string(), + "example.other".to_string(), + ), + ]); let schema = Arc::new(Schema::new(vec![ Field::new("flag", DataType::Boolean, false), Field::new("a", DataType::Int32, false).with_metadata(metadata.clone()), diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 56741d28c8d6e..9a703e5f7f7bd 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -1804,6 +1804,79 @@ mod tests { Ok(()) } + #[test] + fn test_case_field_metadata_higher_order_function() -> Result<()> { + #[derive(Debug, PartialEq, Eq, Hash)] + struct MetadataPassthrough { + signature: crate::HigherOrderSignature, + } + + impl crate::HigherOrderUDFImpl for MetadataPassthrough { + fn name(&self) -> &str { + "metadata_passthrough" + } + + fn signature(&self) -> &crate::HigherOrderSignature { + &self.signature + } + + fn lambda_parameters( + &self, + _step: usize, + _fields: &[ValueOrLambda>], + ) -> Result { + Ok(crate::LambdaParametersProgress::Complete(vec![])) + } + + fn return_field_from_args( + &self, + args: HigherOrderReturnFieldArgs, + ) -> Result { + let ValueOrLambda::Value(field) = &args.arg_fields[0] else { + unreachable!(); + }; + Ok(Arc::clone(field)) + } + + fn invoke_with_args( + &self, + _args: crate::HigherOrderFunctionArgs, + ) -> Result { + unreachable!() + } + } + + let metadata = HashMap::from([("kind".to_string(), "marked".to_string())]); + let schema = DFSchema::from_unqualified_fields( + vec![ + Field::new("a", DataType::Int32, false).with_metadata(metadata.clone()), + Field::new("b", DataType::Int32, false).with_metadata(metadata.clone()), + ] + .into(), + HashMap::new(), + )?; + let udf = Arc::new(crate::HigherOrderUDF::new_from_impl(MetadataPassthrough { + signature: crate::HigherOrderSignature::any(1, crate::Volatility::Immutable), + })); + let call = |arg| { + Expr::HigherOrderFunction(crate::expr::HigherOrderFunction::new( + Arc::clone(&udf), + vec![arg], + )) + }; + + let shallow = when(lit(true), call(col("a"))).otherwise(col("b"))?; + assert_eq!(shallow.to_field(&schema)?.1.metadata(), &metadata); + + let mut deep = col("a"); + for _ in 0..9 { + deep = call(deep); + } + let deep_case = when(lit(true), deep).otherwise(col("b"))?; + assert!(deep_case.to_field(&schema)?.1.metadata().is_empty()); + Ok(()) + } + #[test] fn test_alias_metadata_is_preserved_in_field_metadata() { let schema = MockExprSchema::new().with_data_type(DataType::Int32); From 279401f993457945e3fdb6aeb2fa9dad5a59b323 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Fri, 9 Oct 2026 23:38:02 +0300 Subject: [PATCH 14/16] test: preserve CASE metadata through empty aliases of NULL --- datafusion/expr/src/expr_schema.rs | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 9a703e5f7f7bd..3fb18efc17149 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -217,7 +217,7 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result work.push(Work::Visit(&cast.expr)); } Expr::Alias(alias) - if alias.metadata.is_some() + if alias.metadata.as_ref().is_some_and(|meta| !meta.is_empty()) && can_infer_exact_branch(expr, false)? => { fields.push(BranchField { @@ -1586,6 +1586,19 @@ mod tests { .otherwise(lit(ScalarValue::Null).cast_to(&DataType::Int32, &schema)?)?; assert_eq!(coerced_null_else.to_field(&schema)?.1.metadata(), &shared); + let Expr::Alias(empty_metadata_alias) = lit(ScalarValue::Null) + .cast_to(&DataType::Int32, &schema)? + .alias("empty_metadata") + else { + unreachable!(); + }; + let empty_metadata_alias = Expr::Alias( + empty_metadata_alias.with_metadata(Some(FieldMetadata::default())), + ); + let aliased_null_else = + when(lit(true), col("a")).otherwise(empty_metadata_alias)?; + assert_eq!(aliased_null_else.to_field(&schema)?.1.metadata(), &shared); + let try_cast_null_else = when(lit(true), col("a")).otherwise(Expr::TryCast( TryCast::new(Box::new(lit(ScalarValue::Null)), DataType::Int32), ))?; From 3b33f8c409bc7c2e0178eec865307cfcd6199366 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Sat, 10 Oct 2026 02:59:07 +0300 Subject: [PATCH 15/16] Keep annotated untyped NULL neutral during CASE coercion --- datafusion/expr/src/expr_schema.rs | 17 ++++++- .../optimizer/src/analyzer/type_coercion.rs | 49 +++++++++++++++++++ 2 files changed, 65 insertions(+), 1 deletion(-) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 3fb18efc17149..633546b273f71 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -256,8 +256,23 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result let Some(source) = fields.pop() else { return internal_err!("Missing CASE cast input field"); }; + let untyped_null = + source.certainly_null && source.field.data_type().is_null(); + let mut field = cast_output_field(&source.field, target, force_nullable); + // A type-only coercion must not make an untyped NULL constrain CASE metadata. + if untyped_null + && target.metadata().is_empty() + && !field.metadata().is_empty() + { + field = Arc::new( + field + .as_ref() + .clone() + .with_metadata(arrow_schema::Metadata::default()), + ); + } fields.push(BranchField { - field: cast_output_field(&source.field, target, force_nullable), + field, certainly_null: source.certainly_null, }); } diff --git a/datafusion/optimizer/src/analyzer/type_coercion.rs b/datafusion/optimizer/src/analyzer/type_coercion.rs index 060c9c783ae0e..bf27245477051 100644 --- a/datafusion/optimizer/src/analyzer/type_coercion.rs +++ b/datafusion/optimizer/src/analyzer/type_coercion.rs @@ -3081,6 +3081,55 @@ mod test { Ok(()) } + #[test] + fn test_case_coercion_preserves_metadata_with_annotated_null() -> Result<()> { + let metadata = std::collections::HashMap::from([ + ( + "ARROW:extension:name".to_string(), + "example.case".to_string(), + ), + ("structured_type".to_string(), "ARRAY(NUMBER)".to_string()), + ]); + let schema = DFSchema::from_unqualified_fields( + vec![ + Field::new("structured", DataType::Int32, false) + .with_metadata(metadata.clone()), + ] + .into(), + std::collections::HashMap::new(), + )?; + let conflicting = std::collections::HashMap::from([ + ( + "ARROW:extension:name".to_string(), + "example.other".to_string(), + ), + ("structured_type".to_string(), "OBJECT".to_string()), + ]); + for null_metadata in [metadata.clone(), conflicting] { + let annotated_null = Expr::Literal( + ScalarValue::Null, + Some(expr::FieldMetadata::from(null_metadata)), + ); + let case = Case { + expr: None, + when_then_expr: vec![(Box::new(lit(true)), Box::new(col("structured")))], + else_expr: Some(Box::new(annotated_null)), + }; + assert_eq!( + Expr::Case(case.clone()).to_field(&schema)?.1.metadata(), + &metadata + ); + + let coerced = coerce_case_expression(case, &schema, None)?; + assert!(matches!(coerced.else_expr.as_deref(), Some(Expr::Cast(_)))); + assert_eq!( + Expr::Case(coerced).to_field(&schema)?.1.metadata(), + &metadata + ); + } + Ok(()) + } + #[test] fn test_case_expression_coercion() -> Result<()> { let schema = Arc::new(DFSchema::from_unqualified_fields( From 1b89bff614dee980c9e6f931fe9c82b090667ee0 Mon Sep 17 00:00:00 2001 From: osipovartem Date: Sat, 10 Oct 2026 03:01:47 +0300 Subject: [PATCH 16/16] Use Arrow-version-independent empty CASE metadata --- datafusion/expr/src/expr_schema.rs | 9 +++------ 1 file changed, 3 insertions(+), 6 deletions(-) diff --git a/datafusion/expr/src/expr_schema.rs b/datafusion/expr/src/expr_schema.rs index 633546b273f71..7ac02eced6548 100644 --- a/datafusion/expr/src/expr_schema.rs +++ b/datafusion/expr/src/expr_schema.rs @@ -264,12 +264,9 @@ fn case_field_metadata(case: &Case, schema: &dyn ExprSchema) -> Result && target.metadata().is_empty() && !field.metadata().is_empty() { - field = Arc::new( - field - .as_ref() - .clone() - .with_metadata(arrow_schema::Metadata::default()), - ); + field = Arc::new(field.as_ref().clone().with_metadata( + std::collections::HashMap::::new(), + )); } fields.push(BranchField { field,