From 7ac7f592521b8c7e5604de2b321e33dcc30f0564 Mon Sep 17 00:00:00 2001 From: naman Date: Fri, 11 Sep 2026 18:42:10 +0530 Subject: [PATCH 1/4] fix: Emit equality conditions for Substrait CASE base expressions Substrait's `IfThen` has no base expression. Every `IfClause` is a standalone boolean condition, and `then` is the value that clause yields. The producer instead encoded `CASE WHEN THEN ...` by pushing a leading `IfClause` that carries the base expression in `if` and leaves `then` unset, followed by one clause per WHEN whose `if` is the raw WHEN operand. For `CASE a WHEN 1 THEN 'x' WHEN 2 THEN 'y' ELSE 'z' END` that emits three clauses whose conditions are `a`, `1` and `2`, none of which is boolean, and a first clause with no result. The convention is private to DataFusion: the consumer reads a `then`-less first clause back as the base expression, so a DataFusion-to-DataFusion round trip is unaffected. Any other engine sees clauses it cannot evaluate. Emit ` = ` as each clause condition instead, the same desugaring `from_between` already applies to `BETWEEN`. DataFusion matches a base expression with `=` semantics, so the plan keeps its meaning, including a NULL WHEN operand never matching. A base `CASE` now round trips as the equivalent searched `CASE`, keeping its original projection name and schema. --- .../src/logical_plan/producer/expr/if_then.rs | 31 ++++---- .../tests/cases/roundtrip_logical_plan.rs | 8 ++- datafusion/substrait/tests/cases/serialize.rs | 72 ++++++++++++++++++- 3 files changed, 96 insertions(+), 15 deletions(-) diff --git a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs index 2c10b26436f50..2ee7510a0931a 100644 --- a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs +++ b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs @@ -17,7 +17,7 @@ use crate::logical_plan::producer::SubstraitProducer; use datafusion::common::DFSchemaRef; -use datafusion::logical_expr::Case; +use datafusion::logical_expr::{Case, Expr}; use substrait::proto::Expression; use substrait::proto::expression::if_then::IfClause; use substrait::proto::expression::{IfThen, RexType}; @@ -32,19 +32,24 @@ pub fn from_case( when_then_expr, else_expr, } = case; - let mut ifs: Vec = vec![]; - // Parse base - if let Some(e) = expr { - // Base expression exists - ifs.push(IfClause { - r#if: Some(producer.handle_expr(e, schema)?), - then: None, - }); - } - // Parse `when`s - for (r#if, then) in when_then_expr { + + // Substrait's `IfThen` has no notion of a base expression: every `IfClause` + // is a standalone boolean condition. A `CASE WHEN THEN ...` + // is therefore emitted as `IfClause`s over ` = `, the same + // desugaring `from_between` applies to `BETWEEN`. DataFusion matches a base + // expression with `=` semantics, so this preserves the plan's meaning, + // including a `NULL` `` never matching. + let mut ifs: Vec = Vec::with_capacity(when_then_expr.len()); + for (when, then) in when_then_expr { + let condition = match expr { + Some(base) => { + let eq = Expr::eq(*base.clone(), *when.clone()); + producer.handle_expr(&eq, schema)? + } + None => producer.handle_expr(when, schema)?, + }; ifs.push(IfClause { - r#if: Some(producer.handle_expr(r#if, schema)?), + r#if: Some(condition), then: Some(producer.handle_expr(then, schema)?), }); } diff --git a/datafusion/substrait/tests/cases/roundtrip_logical_plan.rs b/datafusion/substrait/tests/cases/roundtrip_logical_plan.rs index 813d0ed6c3489..2c3c989dd2d4d 100644 --- a/datafusion/substrait/tests/cases/roundtrip_logical_plan.rs +++ b/datafusion/substrait/tests/cases/roundtrip_logical_plan.rs @@ -656,12 +656,18 @@ async fn case_without_base_expression() -> Result<()> { #[tokio::test] async fn case_with_base_expression() -> Result<()> { - roundtrip( + // Substrait has no base expression in `IfThen`, so a base `CASE` is emitted + // as conditions over ` = ` and comes back in that form. The + // projection keeps its original name, so the schema is unchanged. + assert_expected_plan( "SELECT (CASE a WHEN 0 THEN 'zero' WHEN 1 THEN 'one' ELSE 'other' END) FROM data", + "Projection: CASE WHEN data.a = Int64(0) THEN Utf8(\"zero\") WHEN data.a = Int64(1) THEN Utf8(\"one\") ELSE Utf8(\"other\") END AS CASE data.a WHEN Int64(0) THEN Utf8(\"zero\") WHEN Int64(1) THEN Utf8(\"one\") ELSE Utf8(\"other\") END\ + \n TableScan: data projection=[a]", + true, ) .await } diff --git a/datafusion/substrait/tests/cases/serialize.rs b/datafusion/substrait/tests/cases/serialize.rs index 4a8413718edb9..d0aa2718964e9 100644 --- a/datafusion/substrait/tests/cases/serialize.rs +++ b/datafusion/substrait/tests/cases/serialize.rs @@ -30,7 +30,8 @@ mod tests { use std::{fs, sync::Arc}; use substrait::proto::expression::field_reference::{ReferenceType, RootType}; use substrait::proto::expression::reference_segment; - use substrait::proto::expression::{ReferenceSegment, RexType}; + use substrait::proto::expression::{IfThen, ReferenceSegment, RexType}; + use substrait::proto::extensions::simple_extension_declaration::MappingType; use substrait::proto::function_argument::ArgType; use substrait::proto::plan_rel::RelType; use substrait::proto::rel_common::{Emit, EmitKind}; @@ -321,6 +322,75 @@ mod tests { Ok(()) } + /// Substrait's `IfThen` has no base expression: every `IfClause` is a + /// standalone boolean condition and `then` is the value that clause yields. + /// A `CASE WHEN ...` must therefore be emitted as conditions + /// over ` = `. A round trip cannot catch a regression here, + /// because the consumer reads back whatever the producer writes. + #[tokio::test] + async fn case_with_base_expression_emits_equality_conditions() -> Result<()> { + let ctx = create_context().await?; + let sql = "SELECT CASE a WHEN 1 THEN 'x' WHEN 2 THEN 'y' ELSE 'z' END FROM data"; + + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let proto = to_substrait_plan(&plan, &ctx.state())?; + + let equal_anchors: Vec = proto + .extensions + .iter() + .filter_map(|e| match e.mapping_type.as_ref().unwrap() { + MappingType::ExtensionFunction(f) if f.name == "equal" => { + Some(f.function_anchor) + } + _ => None, + }) + .collect(); + assert!(!equal_anchors.is_empty(), "no `equal` function registered"); + + let root = match proto.relations.first().unwrap().rel_type.as_ref() { + Some(RelType::Root(root)) => root.input.as_ref().unwrap(), + _ => panic!("expected Root"), + }; + let Some(rel::RelType::Project(project)) = root.rel_type.as_ref() else { + panic!("expected Project") + }; + + let if_thens: Vec<&IfThen> = project + .expressions + .iter() + .filter_map(|expr| match expr.rex_type.as_ref() { + Some(RexType::IfThen(if_then)) => Some(if_then.as_ref()), + _ => None, + }) + .collect(); + assert_eq!(if_thens.len(), 1, "expected one IfThen for `{sql}`"); + let if_then = if_thens[0]; + + // One clause per WHEN, with no extra clause carrying the base expression. + assert_eq!(if_then.ifs.len(), 2); + assert!(if_then.r#else.is_some()); + + for (i, clause) in if_then.ifs.iter().enumerate() { + let condition = clause + .r#if + .as_ref() + .unwrap_or_else(|| panic!("clause {i} has no condition")); + assert!(clause.then.is_some(), "clause {i} has no `then`"); + + match condition.rex_type.as_ref().unwrap() { + RexType::ScalarFunction(f) => assert!( + equal_anchors.contains(&f.function_reference), + "clause {i} condition is not an `equal` call" + ), + other => { + panic!("clause {i} condition is not a scalar function: {other:?}") + } + } + } + + Ok(()) + } + fn assert_emit(rel_common: Option<&RelCommon>, output_mapping: Vec) { assert_eq!( rel_common.unwrap().emit_kind.clone(), From 2574122663b3b6f78694f68140060d934405486f Mon Sep 17 00:00:00 2001 From: naman Date: Wed, 16 Sep 2026 08:52:20 +0530 Subject: [PATCH 2/4] Reject a volatile CASE base expression The base is written once per WHEN, so a volatile base would be evaluated once per arm. CaseExpr evaluates it once and compares every WHEN against that one value, so such a plan has no faithful IfThen encoding. The regression uses a counter UDF rather than random(), so the repeated evaluation is deterministic: the base form returns 10 after one call, and the same CASE after this desugaring returns 99 after two. The serialization test now also asserts each condition is equal(base, literal) with the operands in that order. --- .../src/logical_plan/producer/expr/if_then.rs | 13 +- datafusion/substrait/tests/cases/serialize.rs | 166 +++++++++++++++++- 2 files changed, 170 insertions(+), 9 deletions(-) diff --git a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs index 2ee7510a0931a..015c1aa9c3ac4 100644 --- a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs +++ b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs @@ -16,7 +16,7 @@ // under the License. use crate::logical_plan::producer::SubstraitProducer; -use datafusion::common::DFSchemaRef; +use datafusion::common::{DFSchemaRef, not_impl_err}; use datafusion::logical_expr::{Case, Expr}; use substrait::proto::Expression; use substrait::proto::expression::if_then::IfClause; @@ -39,6 +39,17 @@ pub fn from_case( // desugaring `from_between` applies to `BETWEEN`. DataFusion matches a base // expression with `=` semantics, so this preserves the plan's meaning, // including a `NULL` `` never matching. + // + // The base is written once per WHEN, which a volatile base would then + // evaluate once per arm. `CaseExpr` evaluates it once and compares every + // WHEN against that one value, so such a plan has no faithful `IfThen` + // encoding and is rejected instead. + if let Some(base) = expr.as_ref().filter(|base| base.is_volatile()) { + return not_impl_err!( + "Substrait does not support a volatile CASE base expression: {base}" + ); + } + let mut ifs: Vec = Vec::with_capacity(when_then_expr.len()); for (when, then) in when_then_expr { let condition = match expr { diff --git a/datafusion/substrait/tests/cases/serialize.rs b/datafusion/substrait/tests/cases/serialize.rs index d0aa2718964e9..9f9e4f88adb72 100644 --- a/datafusion/substrait/tests/cases/serialize.rs +++ b/datafusion/substrait/tests/cases/serialize.rs @@ -23,12 +23,22 @@ mod tests { use datafusion_substrait::logical_plan::producer::to_substrait_plan; use datafusion_substrait::serializer; + use datafusion::arrow::array::Int64Array; + use datafusion::arrow::datatypes::DataType; + use datafusion::common::ScalarValue; use datafusion::error::Result; + use datafusion::logical_expr::{ + ColumnarValue, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, Signature, + Volatility, + }; use datafusion::prelude::*; use insta::assert_snapshot; + use std::hash::{Hash, Hasher}; + use std::sync::atomic::{AtomicUsize, Ordering}; use std::{fs, sync::Arc}; use substrait::proto::expression::field_reference::{ReferenceType, RootType}; + use substrait::proto::expression::literal::LiteralType; use substrait::proto::expression::reference_segment; use substrait::proto::expression::{IfThen, ReferenceSegment, RexType}; use substrait::proto::extensions::simple_extension_declaration::MappingType; @@ -377,20 +387,160 @@ mod tests { .unwrap_or_else(|| panic!("clause {i} has no condition")); assert!(clause.then.is_some(), "clause {i} has no `then`"); - match condition.rex_type.as_ref().unwrap() { - RexType::ScalarFunction(f) => assert!( - equal_anchors.contains(&f.function_reference), - "clause {i} condition is not an `equal` call" - ), - other => { - panic!("clause {i} condition is not a scalar function: {other:?}") - } + let RexType::ScalarFunction(f) = condition.rex_type.as_ref().unwrap() else { + panic!("clause {i} condition is not a scalar function: {condition:?}") + }; + assert!( + equal_anchors.contains(&f.function_reference), + "clause {i} condition is not an `equal` call" + ); + assert_eq!(f.arguments.len(), 2, "clause {i} condition arity"); + + // The condition must be ` = `, in that order: the base + // field reference on the left, the WHEN literal on the right. + let args: Vec<&Expression> = f + .arguments + .iter() + .map(|arg| match arg.arg_type.as_ref().unwrap() { + ArgType::Value(value) => value, + other => panic!("clause {i} argument is not a value: {other:?}"), + }) + .collect(); + + let Some(RexType::Selection(field)) = args[0].rex_type.as_ref() else { + panic!( + "clause {i} left operand is not a field reference: {:?}", + args[0] + ) + }; + assert!( + matches!(field.root_type, Some(RootType::RootReference(_))), + "clause {i} left operand is not rooted at the input" + ); + let Some(ReferenceType::DirectReference(ReferenceSegment { + reference_type: + Some(reference_segment::ReferenceType::StructField(struct_field)), + })) = field.reference_type.as_ref() + else { + panic!("clause {i} left operand is not a direct struct reference") + }; + // `data.a` is the first field of the scan. + assert_eq!(struct_field.field, 0, "clause {i} left operand field index"); + + let Some(RexType::Literal(literal)) = args[1].rex_type.as_ref() else { + panic!("clause {i} right operand is not a literal: {:?}", args[1]) + }; + assert_eq!( + literal.literal_type, + Some(LiteralType::I64(i as i64 + 1)), + "clause {i} right operand literal" + ); + } + + Ok(()) + } + + /// A nullary volatile function returning 1 on its first call, 2 on its + /// second, and so on, so that a repeated evaluation is visible in the + /// result rather than being random. + #[derive(Debug)] + struct CallCounter { + signature: Signature, + calls: Arc, + } + + impl CallCounter { + fn new(calls: Arc) -> Self { + Self { + signature: Signature::nullary(Volatility::Volatile), + calls, } } + } + + impl PartialEq for CallCounter { + fn eq(&self, other: &Self) -> bool { + self.signature == other.signature + } + } + + impl Eq for CallCounter {} + + impl Hash for CallCounter { + fn hash(&self, state: &mut H) { + self.signature.hash(state); + } + } + + impl ScalarUDFImpl for CallCounter { + fn name(&self) -> &str { + "call_counter" + } + + fn signature(&self) -> &Signature { + &self.signature + } + + fn return_type(&self, _arg_types: &[DataType]) -> Result { + Ok(DataType::Int64) + } + + fn invoke_with_args(&self, _args: ScalarFunctionArgs) -> Result { + let call = self.calls.fetch_add(1, Ordering::SeqCst) as i64 + 1; + Ok(ColumnarValue::Scalar(ScalarValue::Int64(Some(call)))) + } + } + + /// `CaseExpr` evaluates a base expression once and compares every WHEN + /// against that one value, so a volatile base cannot be emitted as + /// ` = ` conditions: each condition would evaluate it again. + /// The producer rejects such a plan instead of changing its meaning. + #[tokio::test] + async fn case_with_volatile_base_expression_is_rejected() -> Result<()> { + let ctx = create_context().await?; + let calls = Arc::new(AtomicUsize::new(0)); + ctx.register_udf(ScalarUDF::from(CallCounter::new(Arc::clone(&calls)))); + + // One row, so the difference below is only in how often the base runs. + let base_sql = "SELECT CASE call_counter() WHEN 2 THEN 20 WHEN 1 THEN 10 ELSE 99 END FROM data WHERE a = 1"; + // The same CASE after the desugaring this file applies to a base CASE. + let desugared_sql = "SELECT CASE WHEN call_counter() = 2 THEN 20 WHEN call_counter() = 1 THEN 10 ELSE 99 END FROM data WHERE a = 1"; + + // The base is evaluated once, returns 1, and matches the second WHEN. + assert_eq!(single_i64(&ctx, base_sql).await?, 10); + assert_eq!(calls.swap(0, Ordering::SeqCst), 1); + + // Desugared, it is evaluated once per condition: 1 does not equal 2, + // then 2 does not equal 1, so the row falls through to ELSE. + assert_eq!(single_i64(&ctx, desugared_sql).await?, 99); + assert_eq!(calls.swap(0, Ordering::SeqCst), 2); + + let plan = ctx.sql(base_sql).await?.into_optimized_plan()?; + let err = to_substrait_plan(&plan, &ctx.state()) + .expect_err("a volatile CASE base expression must be rejected") + .to_string(); + assert!( + err.contains("volatile CASE base expression"), + "unexpected error: {err}" + ); Ok(()) } + /// Runs `sql` and returns the single `Int64` value it produces. + async fn single_i64(ctx: &SessionContext, sql: &str) -> Result { + let batches = ctx.sql(sql).await?.collect().await?; + let rows: usize = batches.iter().map(|batch| batch.num_rows()).sum(); + assert_eq!(rows, 1, "expected one row from `{sql}`"); + let batch = batches.iter().find(|batch| batch.num_rows() == 1).unwrap(); + let values = batch + .column(0) + .as_any() + .downcast_ref::() + .expect("expected an Int64 column"); + Ok(values.value(0)) + } + fn assert_emit(rel_common: Option<&RelCommon>, output_mapping: Vec) { assert_eq!( rel_common.unwrap().emit_kind.clone(), From 879e22926d1d03e6f246eaa18e33e5123c08592d Mon Sep 17 00:00:00 2001 From: naman Date: Thu, 17 Sep 2026 23:43:38 +0530 Subject: [PATCH 3/4] Look inside subqueries when rejecting a volatile CASE base Expr::is_volatile walks the expression tree, where a subquery is a leaf, so a base such as (SELECT random()) passed the guard and was then duplicated into every equal(base, when) condition. The check now walks the plan inside a subquery as well, and calls back into itself so a subquery nested in one is covered. It reads the four expressions that carry a subquery, since a base is not restricted to a scalar one. A base with nothing volatile inside it still serializes. --- .../src/logical_plan/producer/expr/if_then.rs | 43 ++++++++++++++++++- datafusion/substrait/tests/cases/serialize.rs | 17 ++++++++ 2 files changed, 58 insertions(+), 2 deletions(-) diff --git a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs index 015c1aa9c3ac4..daf9c475cd6b3 100644 --- a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs +++ b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs @@ -16,8 +16,10 @@ // under the License. use crate::logical_plan::producer::SubstraitProducer; +use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion}; use datafusion::common::{DFSchemaRef, not_impl_err}; -use datafusion::logical_expr::{Case, Expr}; +use datafusion::logical_expr::expr::{Exists, InSubquery, SetComparison}; +use datafusion::logical_expr::{Case, Expr, LogicalPlan}; use substrait::proto::Expression; use substrait::proto::expression::if_then::IfClause; use substrait::proto::expression::{IfThen, RexType}; @@ -44,7 +46,9 @@ pub fn from_case( // evaluate once per arm. `CaseExpr` evaluates it once and compares every // WHEN against that one value, so such a plan has no faithful `IfThen` // encoding and is rejected instead. - if let Some(base) = expr.as_ref().filter(|base| base.is_volatile()) { + if let Some(base) = expr + && is_volatile_including_subqueries(base)? + { return not_impl_err!( "Substrait does not support a volatile CASE base expression: {base}" ); @@ -75,3 +79,38 @@ pub fn from_case( rex_type: Some(RexType::IfThen(Box::new(IfThen { ifs, r#else }))), }) } + +/// Whether evaluating `expr` twice can give two different values. +/// +/// [`Expr::is_volatile`] walks the expression tree, where a subquery is a leaf, +/// so it reports `(SELECT random())` as not volatile. The plan inside one has to +/// be walked as well, or a base holding it would be duplicated by the +/// desugaring above. +fn is_volatile_including_subqueries(expr: &Expr) -> datafusion::common::Result { + expr.exists(|expr| match expr { + Expr::ScalarSubquery(subquery) + | Expr::Exists(Exists { subquery, .. }) + | Expr::InSubquery(InSubquery { subquery, .. }) + | Expr::SetComparison(SetComparison { subquery, .. }) => { + plan_is_volatile(&subquery.subquery) + } + expr => Ok(expr.is_volatile_node()), + }) +} + +/// Whether any expression in `plan`, or in a plan nested in one of them, is +/// volatile. +fn plan_is_volatile(plan: &LogicalPlan) -> datafusion::common::Result { + plan.exists(|plan| { + let mut volatile = false; + plan.apply_expressions(|expr| { + volatile = is_volatile_including_subqueries(expr)?; + Ok(if volatile { + TreeNodeRecursion::Stop + } else { + TreeNodeRecursion::Continue + }) + })?; + Ok(volatile) + }) +} diff --git a/datafusion/substrait/tests/cases/serialize.rs b/datafusion/substrait/tests/cases/serialize.rs index 9f9e4f88adb72..e71fb660b8f65 100644 --- a/datafusion/substrait/tests/cases/serialize.rs +++ b/datafusion/substrait/tests/cases/serialize.rs @@ -524,6 +524,23 @@ mod tests { "unexpected error: {err}" ); + // `Expr::is_volatile` does not look inside a subquery's plan, but the + // desugaring duplicates the base all the same, so this is rejected too. + let subquery_sql = "SELECT CASE (SELECT call_counter()) WHEN 2 THEN 20 WHEN 1 THEN 10 ELSE 99 END FROM data WHERE a = 1"; + let plan = ctx.sql(subquery_sql).await?.into_optimized_plan()?; + let err = to_substrait_plan(&plan, &ctx.state()) + .expect_err("a volatile scalar subquery base must be rejected") + .to_string(); + assert!( + err.contains("volatile CASE base expression"), + "unexpected error: {err}" + ); + + // A subquery base with nothing volatile in it is still emitted. + let pure_sql = "SELECT CASE (SELECT max(a) FROM data) WHEN 2 THEN 20 ELSE 99 END FROM data WHERE a = 1"; + let plan = ctx.sql(pure_sql).await?.into_optimized_plan()?; + to_substrait_plan(&plan, &ctx.state())?; + Ok(()) } From 908ea39a6beee2b5bee73d2c810c5b764e38d8a4 Mon Sep 17 00:00:00 2001 From: naman Date: Fri, 18 Sep 2026 11:43:22 +0530 Subject: [PATCH 4/4] Skip the WHEN operands a NULL CASE base never evaluates `CaseExpr::case_when_with_expr` fills in the result for rows whose base is NULL and drops them from the batch before it evaluates the first WHEN, so a WHEN operand never runs on those rows. ` = ` evaluates both of its operands, so the desugaring did run it there, which turns a query that succeeds into one that fails when the operand errors on such a row. Emit that skip as a leading clause yielding the ELSE value. It is added for a nullable base whose WHEN operands are not all literals or columns: reading a literal or a column on those rows cannot fail and has no side effect, so the common `CASE WHEN ...` keeps the exact encoding it had. --- .../src/logical_plan/producer/expr/if_then.rs | 53 +++- datafusion/substrait/tests/cases/serialize.rs | 229 +++++++++++++++--- 2 files changed, 248 insertions(+), 34 deletions(-) diff --git a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs index daf9c475cd6b3..0afa361eaf2fb 100644 --- a/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs +++ b/datafusion/substrait/src/logical_plan/producer/expr/if_then.rs @@ -17,9 +17,9 @@ use crate::logical_plan::producer::SubstraitProducer; use datafusion::common::tree_node::{TreeNode, TreeNodeRecursion}; -use datafusion::common::{DFSchemaRef, not_impl_err}; +use datafusion::common::{DFSchemaRef, ScalarValue, not_impl_err}; use datafusion::logical_expr::expr::{Exists, InSubquery, SetComparison}; -use datafusion::logical_expr::{Case, Expr, LogicalPlan}; +use datafusion::logical_expr::{Case, Expr, ExprSchemable, LogicalPlan}; use substrait::proto::Expression; use substrait::proto::expression::if_then::IfClause; use substrait::proto::expression::{IfThen, RexType}; @@ -54,7 +54,54 @@ pub fn from_case( ); } - let mut ifs: Vec = Vec::with_capacity(when_then_expr.len()); + // A NULL base answers from ELSE without any WHEN being evaluated: + // `CaseExpr::case_when_with_expr` fills those rows in and drops them from + // the batch before it evaluates the first WHEN. The desugaring below would + // evaluate them, because ` = ` evaluates both of its operands, + // so a WHEN that errors or has a side effect would reach rows the plan + // never ran it on. Emitting that skip as a leading clause restores it. + // + // It is only needed when a WHEN operand can do something on those rows. + // Reading a literal or a column cannot fail and has no side effect, so the + // common `CASE WHEN ...` keeps the encoding it had. + let when_operand_is_inert = |(when, _): &(Box, Box)| { + matches!(when.as_ref(), Expr::Literal(..) | Expr::Column(_)) + }; + let null_base_guard = match expr { + Some(base) + if !when_then_expr.iter().all(when_operand_is_inert) + && base.nullable(schema.as_ref())? => + { + Some(base) + } + _ => None, + }; + + let mut ifs: Vec = + Vec::with_capacity(when_then_expr.len() + usize::from(null_base_guard.is_some())); + + if let Some(base) = null_base_guard { + let condition = producer.handle_expr(&base.clone().is_null(), schema)?; + // The value a NULL base yields: ELSE, or a NULL of the result type when + // the CASE has none. + let then = match else_expr { + Some(e) => producer.handle_expr(e, schema)?, + None => { + let result_type = match when_then_expr.first() { + Some((_, then)) => then.get_type(schema.as_ref())?, + None => { + return not_impl_err!("CASE with no WHEN clause"); + } + }; + let null = Expr::Literal(ScalarValue::try_from(&result_type)?, None); + producer.handle_expr(&null, schema)? + } + }; + ifs.push(IfClause { + r#if: Some(condition), + then: Some(then), + }); + } for (when, then) in when_then_expr { let condition = match expr { Some(base) => { diff --git a/datafusion/substrait/tests/cases/serialize.rs b/datafusion/substrait/tests/cases/serialize.rs index e71fb660b8f65..a69094ba895be 100644 --- a/datafusion/substrait/tests/cases/serialize.rs +++ b/datafusion/substrait/tests/cases/serialize.rs @@ -24,8 +24,11 @@ mod tests { use datafusion_substrait::serializer; use datafusion::arrow::array::Int64Array; - use datafusion::arrow::datatypes::DataType; + use datafusion::arrow::datatypes::{DataType, Field, Schema}; + use datafusion::arrow::record_batch::RecordBatch; + use datafusion::arrow::util::pretty; use datafusion::common::ScalarValue; + use datafusion::datasource::MemTable; use datafusion::error::Result; use datafusion::logical_expr::{ ColumnarValue, ScalarFunctionArgs, ScalarUDF, ScalarUDFImpl, Signature, @@ -46,7 +49,7 @@ mod tests { use substrait::proto::plan_rel::RelType; use substrait::proto::rel_common::{Emit, EmitKind}; use substrait::proto::r#type::{I64, Kind as TypeKind, List, Nullability, Struct}; - use substrait::proto::{Expression, RelCommon, Type, rel}; + use substrait::proto::{Expression, Plan, RelCommon, Type, rel}; use crate::cases::roundtrip_logical_plan::higher_order_function_ctx; @@ -345,40 +348,20 @@ mod tests { let plan = ctx.sql(sql).await?.into_optimized_plan()?; let proto = to_substrait_plan(&plan, &ctx.state())?; - let equal_anchors: Vec = proto - .extensions - .iter() - .filter_map(|e| match e.mapping_type.as_ref().unwrap() { - MappingType::ExtensionFunction(f) if f.name == "equal" => { - Some(f.function_anchor) - } - _ => None, - }) - .collect(); + let equal_anchors = function_anchors(&proto, "equal"); assert!(!equal_anchors.is_empty(), "no `equal` function registered"); - let root = match proto.relations.first().unwrap().rel_type.as_ref() { - Some(RelType::Root(root)) => root.input.as_ref().unwrap(), - _ => panic!("expected Root"), - }; - let Some(rel::RelType::Project(project)) = root.rel_type.as_ref() else { - panic!("expected Project") - }; + let if_then = single_if_then(&proto); - let if_thens: Vec<&IfThen> = project - .expressions - .iter() - .filter_map(|expr| match expr.rex_type.as_ref() { - Some(RexType::IfThen(if_then)) => Some(if_then.as_ref()), - _ => None, - }) - .collect(); - assert_eq!(if_thens.len(), 1, "expected one IfThen for `{sql}`"); - let if_then = if_thens[0]; - - // One clause per WHEN, with no extra clause carrying the base expression. + // One clause per WHEN, with no extra clause carrying the base + // expression. These WHEN operands are literals, so reading one on a + // NULL base row does nothing and no guard clause is needed. assert_eq!(if_then.ifs.len(), 2); assert!(if_then.r#else.is_some()); + assert!( + function_anchors(&proto, "is_null").is_empty(), + "literal WHEN operands should not need an `is_null` guard" + ); for (i, clause) in if_then.ifs.iter().enumerate() { let condition = clause @@ -440,6 +423,190 @@ mod tests { Ok(()) } + /// The function anchors registered under `name` in `proto`. + fn function_anchors(proto: &Plan, name: &str) -> Vec { + proto + .extensions + .iter() + .filter_map(|e| match e.mapping_type.as_ref().unwrap() { + MappingType::ExtensionFunction(f) if f.name == name => { + Some(f.function_anchor) + } + _ => None, + }) + .collect() + } + + /// The single `IfThen` in the plan's projection. + fn single_if_then(proto: &Plan) -> &IfThen { + let root = match proto.relations.first().unwrap().rel_type.as_ref() { + Some(RelType::Root(root)) => root.input.as_ref().unwrap(), + _ => panic!("expected Root"), + }; + let Some(rel::RelType::Project(project)) = root.rel_type.as_ref() else { + panic!("expected Project") + }; + let if_thens: Vec<&IfThen> = project + .expressions + .iter() + .filter_map(|expr| match expr.rex_type.as_ref() { + Some(RexType::IfThen(if_then)) => Some(if_then.as_ref()), + _ => None, + }) + .collect(); + assert_eq!(if_thens.len(), 1, "expected one IfThen"); + if_thens[0] + } + + /// A nullable base with a WHEN operand that is neither a literal nor a + /// column gets the guard clause, which yields the ELSE value. + #[tokio::test] + async fn case_with_null_base_emits_guard_clause() -> Result<()> { + let ctx = create_context().await?; + let sql = "SELECT CASE a WHEN 10 / a THEN 'x' ELSE 'z' END FROM data"; + + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let proto = to_substrait_plan(&plan, &ctx.state())?; + + let if_then = single_if_then(&proto); + assert_eq!(if_then.ifs.len(), 2, "one guard clause plus one WHEN"); + + let guard = &if_then.ifs[0]; + let guard_condition = guard.r#if.as_ref().expect("guard has no condition"); + let RexType::ScalarFunction(guard_fn) = + guard_condition.rex_type.as_ref().unwrap() + else { + panic!("guard condition is not a scalar function: {guard_condition:?}") + }; + assert!( + function_anchors(&proto, "is_null").contains(&guard_fn.function_reference), + "guard condition is not an `is_null` call" + ); + // The guard yields what a NULL base yields: the ELSE value. + let Some(RexType::Literal(literal)) = + guard.then.as_ref().and_then(|t| t.rex_type.as_ref()) + else { + panic!("guard `then` is not a literal: {:?}", guard.then) + }; + assert_eq!( + literal.literal_type, + Some(LiteralType::String("z".to_string())) + ); + + Ok(()) + } + + /// The guard is only needed when the base can be NULL. A base that cannot + /// be NULL keeps the conditions on their own. + #[tokio::test] + async fn case_with_non_nullable_base_emits_no_null_guard() -> Result<()> { + let schema = Arc::new(Schema::new(vec![Field::new("a", DataType::Int64, false)])); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(Int64Array::from(vec![1, 2]))], + )?; + let ctx = SessionContext::new(); + ctx.register_table( + "t", + Arc::new(MemTable::try_new(Arc::clone(&schema), vec![vec![batch]])?), + )?; + + // The same WHEN operand that earns a guard over a nullable base. + let sql = "SELECT CASE a WHEN 10 / a THEN 'x' ELSE 'z' END FROM t"; + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let proto = to_substrait_plan(&plan, &ctx.state())?; + + let if_then = single_if_then(&proto); + assert_eq!( + if_then.ifs.len(), + 1, + "a non-nullable base should not add a guard clause" + ); + assert!( + function_anchors(&proto, "is_null").is_empty(), + "no `is_null` should be registered for a non-nullable base" + ); + + Ok(()) + } + + /// `CaseExpr::case_when_with_expr` fills the result for rows whose base is + /// NULL and drops them before it evaluates the first WHEN, so a WHEN that + /// errors never runs on them. ` = ` evaluates both operands, so + /// without the guard clause the emitted plan fails on a query that succeeds. + #[tokio::test] + async fn case_with_null_base_does_not_evaluate_when_operands() -> Result<()> { + let ctx = SessionContext::new(); + // `10 / b` divides by zero on the second row, whose base is NULL. + let sql = "SELECT CASE a WHEN 10 / b THEN 'x' ELSE 'y' END AS r \ + FROM (VALUES (1, 1), (NULL, 0)) AS t(a, b)"; + + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let native = DataFrame::new(ctx.state(), plan.clone()).collect().await?; + + let proto = to_substrait_plan(&plan, &ctx.state())?; + let plan2 = from_substrait_plan(&ctx.state(), &proto).await?; + let roundtrip = DataFrame::new(ctx.state(), plan2).collect().await?; + + let native = pretty::pretty_format_batches(&native)?.to_string(); + assert_eq!( + native, + pretty::pretty_format_batches(&roundtrip)?.to_string() + ); + assert_snapshot!(native, @r" + +---+ + | r | + +---+ + | y | + | y | + +---+ + "); + + // With no ELSE, the guard yields a NULL of the result type, which is + // what the base CASE returns for those rows. + let sql = "SELECT CASE a WHEN 10 / b THEN 'x' END AS r \ + FROM (VALUES (1, 1), (NULL, 0)) AS t(a, b)"; + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + let native = DataFrame::new(ctx.state(), plan.clone()).collect().await?; + let proto = to_substrait_plan(&plan, &ctx.state())?; + let plan2 = from_substrait_plan(&ctx.state(), &proto).await?; + let roundtrip = DataFrame::new(ctx.state(), plan2).collect().await?; + let native = pretty::pretty_format_batches(&native)?.to_string(); + assert_eq!( + native, + pretty::pretty_format_batches(&roundtrip)?.to_string() + ); + assert_snapshot!(native, @r" + +---+ + | r | + +---+ + | | + | | + +---+ + "); + + // A genuine error is still reported: the same WHEN over a row whose + // base is not NULL fails on both sides. + let sql = "SELECT CASE a WHEN 10 / b THEN 'x' ELSE 'y' END AS r \ + FROM (VALUES (1, 1), (2, 0)) AS t(a, b)"; + let plan = ctx.sql(sql).await?.into_optimized_plan()?; + assert!( + DataFrame::new(ctx.state(), plan.clone()) + .collect() + .await + .is_err(), + "the base CASE should report the division by zero" + ); + let proto = to_substrait_plan(&plan, &ctx.state())?; + let plan2 = from_substrait_plan(&ctx.state(), &proto).await?; + assert!( + DataFrame::new(ctx.state(), plan2).collect().await.is_err(), + "the emitted plan should report the division by zero" + ); + + Ok(()) + } + /// A nullary volatile function returning 1 on its first call, 2 on its /// second, and so on, so that a repeated evaluation is visible in the /// result rather than being random.