Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 11 additions & 6 deletions native/core/src/execution/planner.rs
Original file line number Diff line number Diff line change
Expand Up @@ -142,12 +142,13 @@ use datafusion_comet_proto::{
},
};
use datafusion_comet_spark_expr::{
create_case_when, create_if_expr, jvm_udf::JvmScalarUdfExpr, spark_in_list, ApproxPercentile,
ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow, Correlation, Covariance, CreateNamedStruct,
DecimalRescaleCheckOverflow, FloatOperands, GetArrayStructFields, GetStructField, HllPlusPlus,
HllSketchAgg, HllUnionAgg, IfExpr, ListExtract, MaxMinBy, Mode, NormalizeNaNAndZero,
NormalizeNestedFloats, Regr, RegrType, SparkCastOptions, SparkMinMax, Stddev, SumDecimal,
ToJson, UnboundColumn, Variance, WideDecimalBinaryExpr, WideDecimalOp,
coerce_to_common_type, create_case_when, create_if_expr, jvm_udf::JvmScalarUdfExpr,
spark_in_list, ApproxPercentile, ArrayInsert, Avg, AvgDecimal, Cast, CheckOverflow,
Correlation, Covariance, CreateNamedStruct, DecimalRescaleCheckOverflow, FloatOperands,
GetArrayStructFields, GetStructField, HllPlusPlus, HllSketchAgg, HllUnionAgg, IfExpr,
ListExtract, MaxMinBy, Mode, NormalizeNaNAndZero, NormalizeNestedFloats, Regr, RegrType,
SparkCastOptions, SparkMinMax, Stddev, SumDecimal, ToJson, UnboundColumn, Variance,
WideDecimalBinaryExpr, WideDecimalOp,
};
use itertools::Itertools;
use jni::objects::{Global, JObject};
Expand Down Expand Up @@ -3909,6 +3910,10 @@ impl PhysicalPlanner {
Self::coerce_child_to(arg, &input_schema, widen_map_entry_value_nullable)
})
.collect::<Result<Vec<_>, ExecutionError>>()?
} else if fun_name == "greatest" || fun_name == "least" {
// Spark compares struct arguments field by field by position, while DataFusion's
// coercion below would match fields by name. See `coerce_to_common_type`.
coerce_to_common_type(args, &input_schema)?
} else {
args
};
Expand Down
39 changes: 36 additions & 3 deletions native/spark-expr/src/conditional_funcs/case_when.rs
Original file line number Diff line number Diff line change
Expand Up @@ -96,9 +96,10 @@ pub fn create_if_expr(
)))
}

/// Reconciles Spark IF branches positionally, retaining THEN names and merging nullability.
/// Spark has already coerced the branches to the same SQL type. DataFusion's struct union may
/// instead match by name, pairing different positions when names differ only in case.
/// Reconciles Spark IF branches, or the arguments of `greatest` and `least`, positionally,
/// retaining the first one's names and merging nullability. Spark has already coerced them to the
/// same SQL type. DataFusion's struct union may instead match by name, pairing different positions
/// when names differ only in case.
fn if_common_type(then_type: &DataType, else_type: &DataType) -> Option<DataType> {
use arrow::datatypes::FieldRef;

Expand Down Expand Up @@ -162,6 +163,38 @@ fn coerce_branch(
))
}

/// Casts the arguments of a Spark `ComplexTypeMergingExpression` such as `greatest` or `least` to
/// their common type, which keeps the first argument's field names, as Spark's result type does.
///
/// Spark has already given the arguments the same SQL type up to nullability, and up to the case
/// of struct field names when the analysis is case-insensitive. It compares structs field by field
/// by position. DataFusion's struct coercion and Arrow's struct cast both match fields by name
/// when two structs hold the same set of names, which pairs different positions when the names
/// differ only in case, so the arguments are reconciled here positionally instead. The arguments
/// are returned unchanged when they have no positional common type.
pub fn coerce_to_common_type(

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

#6428 at its latest head (ba8f2a5cd4) renames if_common_type and coerce_branch to positional_common_type, which takes a PositionalTypeCoercion mode, and a public cast_to_common_type. Its new IN arm in the planner also folds a common type over the operands and casts each one, as this function does. The two PRs conflict in planner.rs, case_when.rs and mod.rs, so whichever lands second has to adapt. Could that one leave a single helper for IN, greatest and least? MetadataOnly looks like the right mode here, since #6428 documents it as the comparison mode where Catalyst has already coerced the leaf types.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed. #6428 is still open, so this PR keeps its own helper for now. Whichever lands second will fold positional_common_type(.., PositionalTypeCoercion::MetadataOnly) over the arguments and cast each one with cast_to_common_type, dropping coerce_to_common_type. That leaves one path for IN, greatest and least. If #6428 lands first, I'll rebase this PR onto it and make that change here.

exprs: Vec<Arc<dyn PhysicalExpr>>,
input_schema: &Schema,
) -> Result<Vec<Arc<dyn PhysicalExpr>>> {
let types = exprs
.iter()
.map(|e| e.data_type(input_schema))
.collect::<Result<Vec<_>>>()?;
let Some((first, rest)) = types.split_first() else {
return Ok(exprs);
};
let Some(common_type) = rest.iter().try_fold(first.clone(), |common, data_type| {
if_common_type(&common, data_type)
}) else {
return Ok(exprs);
};
Ok(exprs
.into_iter()
.zip(&types)
.map(|(e, data_type)| coerce_branch(e, data_type, &common_type))
.collect())
}

/// Spark's `CASE WHEN`, which Comet also uses for `IF` and `COALESCE`.
///
/// Spark evaluates a WHEN only for the rows that no earlier WHEN matched, and a branch's value
Expand Down
2 changes: 1 addition & 1 deletion native/spark-expr/src/conditional_funcs/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -18,5 +18,5 @@
mod case_when;
mod if_expr;

pub use case_when::{create_case_when, create_if_expr, CaseWhenExpr};
pub use case_when::{coerce_to_common_type, create_case_when, create_if_expr, CaseWhenExpr};
pub use if_expr::IfExpr;
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
-- Licensed to the Apache Software Foundation (ASF) under one
-- or more contributor license agreements. See the NOTICE file
-- distributed with this work for additional information
-- regarding copyright ownership. The ASF licenses this file
-- to you under the Apache License, Version 2.0 (the
-- "License"); you may not use this file except in compliance
-- with the License. You may obtain a copy of the License at
--
-- http://www.apache.org/licenses/LICENSE-2.0
--
-- Unless required by applicable law or agreed to in writing,
-- software distributed under the License is distributed on an
-- "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
-- KIND, either express or implied. See the License for the
-- specific language governing permissions and limitations
-- under the License.

-- With case-insensitive analysis, Spark accepts greatest and least arguments whose struct field
-- names differ only in case, and compares the structs field by field by position. The result
-- takes the first argument's field names. Here the second argument holds the same names in the
-- other order, so matching the fields by name instead would pair x with x and give a different
-- row for id = 0 and id = 2.

-- Config: spark.comet.exec.range.enabled=true
-- Config: spark.sql.caseSensitive=false

-- Float fields, which take Comet's own greatest and least
query
SELECT id,
greatest(named_struct('x', CAST(id AS DOUBLE), 'X', 1D),
named_struct('X', 1D, 'x', CAST(id AS DOUBLE)))
FROM range(4)

query
SELECT id,
least(named_struct('x', CAST(id AS DOUBLE), 'X', 1D),
named_struct('X', 1D, 'x', CAST(id AS DOUBLE)))
FROM range(4)

-- Integer fields, which take DataFusion's greatest and least
query
SELECT id,
greatest(named_struct('x', CAST(id AS INT), 'X', 1),
named_struct('X', 1, 'x', CAST(id AS INT))),
least(named_struct('x', CAST(id AS INT), 'X', 1),
named_struct('X', 1, 'x', CAST(id AS INT)))
FROM range(4)

-- More than two arguments
query
SELECT id,
greatest(named_struct('x', CAST(id AS INT), 'X', 2),
named_struct('X', 1, 'x', CAST(id AS INT)),
named_struct('X', CAST(id AS INT), 'x', 1))
FROM range(4)

-- Structs nested in arrays
query
SELECT id,
greatest(array(named_struct('x', CAST(id AS DOUBLE), 'X', 1D)),
array(named_struct('X', 1D, 'x', CAST(id AS DOUBLE))))
FROM range(4)

query
SELECT id,
least(array(named_struct('x', CAST(id AS INT), 'X', 1)),
array(named_struct('X', 1, 'x', CAST(id AS INT))))
FROM range(4)

-- Arguments that also differ in whether a nested field can be NULL
query
SELECT id,
greatest(named_struct('x', CAST(id AS DOUBLE), 'X', 1D),
named_struct('X', IF(id = 3, NULL, 1D), 'x', CAST(id AS DOUBLE)))
FROM range(4)

query
SELECT id,
least(array(named_struct('x', CAST(id AS INT), 'X', 1)),
array(named_struct('X', 1, 'x', IF(id = 3, NULL, CAST(id AS INT)))))
FROM range(4)
Loading