From 768b3e90f261c7aea58bdb98dc698b90deeeae34 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sun, 14 Dec 2025 16:24:01 +0400 Subject: [PATCH 1/6] impl map_from_entries --- native/core/src/execution/jni_api.rs | 2 + .../apache/comet/serde/QueryPlanSerde.scala | 3 +- .../scala/org/apache/comet/serde/maps.scala | 29 +++++++++++- .../comet/CometMapExpressionSuite.scala | 45 +++++++++++++++++++ 4 files changed, 77 insertions(+), 2 deletions(-) diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index a24d993059..4f53cea3e6 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -46,6 +46,7 @@ use datafusion_spark::function::datetime::date_add::SparkDateAdd; use datafusion_spark::function::datetime::date_sub::SparkDateSub; use datafusion_spark::function::hash::sha1::SparkSha1; use datafusion_spark::function::hash::sha2::SparkSha2; +use datafusion_spark::function::map::map_from_entries::MapFromEntries; use datafusion_spark::function::math::expm1::SparkExpm1; use datafusion_spark::function::string::char::CharFunc; use datafusion_spark::function::string::concat::SparkConcat; @@ -337,6 +338,7 @@ fn register_datafusion_spark_function(session_ctx: &SessionContext) { session_ctx.register_udf(ScalarUDF::new_from_impl(SparkSha1::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkConcat::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkBitwiseNot::default())); + session_ctx.register_udf(ScalarUDF::new_from_impl(MapFromEntries::default())); } /// Prepares arrow arrays for output. diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index 54df2f1688..a99cf3824b 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -125,7 +125,8 @@ object QueryPlanSerde extends Logging with CometExprShim { classOf[MapKeys] -> CometMapKeys, classOf[MapEntries] -> CometMapEntries, classOf[MapValues] -> CometMapValues, - classOf[MapFromArrays] -> CometMapFromArrays) + classOf[MapFromArrays] -> CometMapFromArrays, + classOf[MapFromEntries] -> CometMapFromEntries) private val structExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( classOf[CreateNamedStruct] -> CometCreateNamedStruct, diff --git a/spark/src/main/scala/org/apache/comet/serde/maps.scala b/spark/src/main/scala/org/apache/comet/serde/maps.scala index 2e217f6af0..498aa3594c 100644 --- a/spark/src/main/scala/org/apache/comet/serde/maps.scala +++ b/spark/src/main/scala/org/apache/comet/serde/maps.scala @@ -19,9 +19,12 @@ package org.apache.comet.serde +import scala.annotation.tailrec + import org.apache.spark.sql.catalyst.expressions._ -import org.apache.spark.sql.types.{ArrayType, MapType} +import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, MapType, StructType} +import org.apache.comet.serde.CometArrayReverse.containsBinary import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithInfo, scalarFunctionExprToProto, scalarFunctionExprToProtoWithReturnType} object CometMapKeys extends CometExpressionSerde[MapKeys] { @@ -89,3 +92,27 @@ object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { optExprWithInfo(mapFromArraysExpr, expr, expr.children: _*) } } + +object CometMapFromEntries extends CometScalarFunction[MapFromEntries]("map_from_entries") { + val keyUnsupportedReason = "Using BinaryType as Map keys is not allowed in map_from_entries" + val valueUnsupportedReason = "Using BinaryType as Map values is not allowed in map_from_entries" + + private def containsBinary(dataType: DataType): Boolean = { + dataType match { + case BinaryType => true + case StructType(fields) => fields.exists(field => containsBinary(field.dataType)) + case ArrayType(elementType, _) => containsBinary(elementType) + case _ => false + } + } + + override def getSupportLevel(expr: MapFromEntries): SupportLevel = { + if (containsBinary(expr.dataType.keyType)) { + return Incompatible(Some(keyUnsupportedReason)) + } + if (containsBinary(expr.dataType.valueType)) { + return Incompatible(Some(valueUnsupportedReason)) + } + Compatible(None) + } +} diff --git a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala index 88c13391a6..01b9744ed6 100644 --- a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala @@ -25,7 +25,9 @@ import org.apache.hadoop.fs.Path import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.BinaryType +import org.apache.comet.serde.CometMapFromEntries import org.apache.comet.testing.{DataGenOptions, ParquetGenerator, SchemaGenOptions} class CometMapExpressionSuite extends CometTestBase { @@ -125,4 +127,47 @@ class CometMapExpressionSuite extends CometTestBase { } } + test("map_from_entries") { + withTempDir { dir => + val path = new Path(dir.toURI.toString, "test.parquet") + val filename = path.toString + val random = new Random(42) + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + val schemaGenOptions = + SchemaGenOptions( + generateArray = true, + generateStruct = true, + primitiveTypes = SchemaGenOptions.defaultPrimitiveTypes.filterNot(_ == BinaryType)) + val dataGenOptions = DataGenOptions(allowNull = false, generateNegativeZero = false) + ParquetGenerator.makeParquetFile( + random, + spark, + filename, + 100, + schemaGenOptions, + dataGenOptions) + } + val df = spark.read.parquet(filename) + df.createOrReplaceTempView("t1") + for (field <- df.schema.fieldNames) { + checkSparkAnswerAndOperator( + spark.sql(s"SELECT map_from_entries(array(struct($field as a, $field as b))) FROM t1")) + } + } + } + + test("map_from_entries - fallback for binary type") { + val table = "t2" + withTable(table) { + sql( + s"create table $table using parquet as select cast(array() as array) as c1 from range(10)") + checkSparkAnswerAndFallbackReason( + sql(s"select map_from_entries(array(struct(c1, 0))) from $table"), + CometMapFromEntries.keyUnsupportedReason) + checkSparkAnswerAndFallbackReason( + sql(s"select map_from_entries(array(struct(0, c1))) from $table"), + CometMapFromEntries.valueUnsupportedReason) + } + } + } From c68c3428676b5d991e7ba9e13464bf2ce1ec84e8 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Tue, 16 Dec 2025 16:10:43 +0400 Subject: [PATCH 2/6] Revert "impl map_from_entries" This reverts commit 768b3e90f261c7aea58bdb98dc698b90deeeae34. --- native/core/src/execution/jni_api.rs | 2 - .../apache/comet/serde/QueryPlanSerde.scala | 3 +- .../scala/org/apache/comet/serde/maps.scala | 29 +----------- .../comet/CometMapExpressionSuite.scala | 45 ------------------- 4 files changed, 2 insertions(+), 77 deletions(-) diff --git a/native/core/src/execution/jni_api.rs b/native/core/src/execution/jni_api.rs index 4f53cea3e6..a24d993059 100644 --- a/native/core/src/execution/jni_api.rs +++ b/native/core/src/execution/jni_api.rs @@ -46,7 +46,6 @@ use datafusion_spark::function::datetime::date_add::SparkDateAdd; use datafusion_spark::function::datetime::date_sub::SparkDateSub; use datafusion_spark::function::hash::sha1::SparkSha1; use datafusion_spark::function::hash::sha2::SparkSha2; -use datafusion_spark::function::map::map_from_entries::MapFromEntries; use datafusion_spark::function::math::expm1::SparkExpm1; use datafusion_spark::function::string::char::CharFunc; use datafusion_spark::function::string::concat::SparkConcat; @@ -338,7 +337,6 @@ fn register_datafusion_spark_function(session_ctx: &SessionContext) { session_ctx.register_udf(ScalarUDF::new_from_impl(SparkSha1::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkConcat::default())); session_ctx.register_udf(ScalarUDF::new_from_impl(SparkBitwiseNot::default())); - session_ctx.register_udf(ScalarUDF::new_from_impl(MapFromEntries::default())); } /// Prepares arrow arrays for output. diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala index a99cf3824b..54df2f1688 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -125,8 +125,7 @@ object QueryPlanSerde extends Logging with CometExprShim { classOf[MapKeys] -> CometMapKeys, classOf[MapEntries] -> CometMapEntries, classOf[MapValues] -> CometMapValues, - classOf[MapFromArrays] -> CometMapFromArrays, - classOf[MapFromEntries] -> CometMapFromEntries) + classOf[MapFromArrays] -> CometMapFromArrays) private val structExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( classOf[CreateNamedStruct] -> CometCreateNamedStruct, diff --git a/spark/src/main/scala/org/apache/comet/serde/maps.scala b/spark/src/main/scala/org/apache/comet/serde/maps.scala index 498aa3594c..2e217f6af0 100644 --- a/spark/src/main/scala/org/apache/comet/serde/maps.scala +++ b/spark/src/main/scala/org/apache/comet/serde/maps.scala @@ -19,12 +19,9 @@ package org.apache.comet.serde -import scala.annotation.tailrec - import org.apache.spark.sql.catalyst.expressions._ -import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, MapType, StructType} +import org.apache.spark.sql.types.{ArrayType, MapType} -import org.apache.comet.serde.CometArrayReverse.containsBinary import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal, optExprWithInfo, scalarFunctionExprToProto, scalarFunctionExprToProtoWithReturnType} object CometMapKeys extends CometExpressionSerde[MapKeys] { @@ -92,27 +89,3 @@ object CometMapFromArrays extends CometExpressionSerde[MapFromArrays] { optExprWithInfo(mapFromArraysExpr, expr, expr.children: _*) } } - -object CometMapFromEntries extends CometScalarFunction[MapFromEntries]("map_from_entries") { - val keyUnsupportedReason = "Using BinaryType as Map keys is not allowed in map_from_entries" - val valueUnsupportedReason = "Using BinaryType as Map values is not allowed in map_from_entries" - - private def containsBinary(dataType: DataType): Boolean = { - dataType match { - case BinaryType => true - case StructType(fields) => fields.exists(field => containsBinary(field.dataType)) - case ArrayType(elementType, _) => containsBinary(elementType) - case _ => false - } - } - - override def getSupportLevel(expr: MapFromEntries): SupportLevel = { - if (containsBinary(expr.dataType.keyType)) { - return Incompatible(Some(keyUnsupportedReason)) - } - if (containsBinary(expr.dataType.valueType)) { - return Incompatible(Some(valueUnsupportedReason)) - } - Compatible(None) - } -} diff --git a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala index 01b9744ed6..88c13391a6 100644 --- a/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometMapExpressionSuite.scala @@ -25,9 +25,7 @@ import org.apache.hadoop.fs.Path import org.apache.spark.sql.CometTestBase import org.apache.spark.sql.functions._ import org.apache.spark.sql.internal.SQLConf -import org.apache.spark.sql.types.BinaryType -import org.apache.comet.serde.CometMapFromEntries import org.apache.comet.testing.{DataGenOptions, ParquetGenerator, SchemaGenOptions} class CometMapExpressionSuite extends CometTestBase { @@ -127,47 +125,4 @@ class CometMapExpressionSuite extends CometTestBase { } } - test("map_from_entries") { - withTempDir { dir => - val path = new Path(dir.toURI.toString, "test.parquet") - val filename = path.toString - val random = new Random(42) - withSQLConf(CometConf.COMET_ENABLED.key -> "false") { - val schemaGenOptions = - SchemaGenOptions( - generateArray = true, - generateStruct = true, - primitiveTypes = SchemaGenOptions.defaultPrimitiveTypes.filterNot(_ == BinaryType)) - val dataGenOptions = DataGenOptions(allowNull = false, generateNegativeZero = false) - ParquetGenerator.makeParquetFile( - random, - spark, - filename, - 100, - schemaGenOptions, - dataGenOptions) - } - val df = spark.read.parquet(filename) - df.createOrReplaceTempView("t1") - for (field <- df.schema.fieldNames) { - checkSparkAnswerAndOperator( - spark.sql(s"SELECT map_from_entries(array(struct($field as a, $field as b))) FROM t1")) - } - } - } - - test("map_from_entries - fallback for binary type") { - val table = "t2" - withTable(table) { - sql( - s"create table $table using parquet as select cast(array() as array) as c1 from range(10)") - checkSparkAnswerAndFallbackReason( - sql(s"select map_from_entries(array(struct(c1, 0))) from $table"), - CometMapFromEntries.keyUnsupportedReason) - checkSparkAnswerAndFallbackReason( - sql(s"select map_from_entries(array(struct(0, c1))) from $table"), - CometMapFromEntries.valueUnsupportedReason) - } - } - } From 68eb53df8832b75ce0e31a9c3fb3fd86526e3004 Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Thu, 6 Aug 2026 10:50:55 +0400 Subject: [PATCH 3/6] chore: delegate ANSI integer arithmetic to arrow checked kernels --- .../src/math_funcs/checked_arithmetic.rs | 117 ++++++++++++------ 1 file changed, 76 insertions(+), 41 deletions(-) diff --git a/native/spark-expr/src/math_funcs/checked_arithmetic.rs b/native/spark-expr/src/math_funcs/checked_arithmetic.rs index a76574e16d..3a2a8bd5e3 100644 --- a/native/spark-expr/src/math_funcs/checked_arithmetic.rs +++ b/native/spark-expr/src/math_funcs/checked_arithmetic.rs @@ -20,6 +20,7 @@ use arrow::array::{ArrayRef, AsArray}; use crate::{divide_by_zero_error, EvalMode, SparkError}; use arrow::buffer::NullBuffer; +use arrow::compute::kernels::{arity, numeric}; use arrow::datatypes::{ ArrowPrimitiveType, DataType, Float16Type, Float32Type, Float64Type, Int16Type, Int32Type, Int64Type, Int8Type, @@ -29,32 +30,58 @@ use datafusion::common::DataFusionError; use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; -pub fn try_arithmetic_kernel( +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum MathOp { + Add, + Sub, + Mul, + Div, +} + +fn try_arithmetic_kernel( left: &PrimitiveArray, right: &PrimitiveArray, - op: &str, is_ansi_mode: bool, + op: MathOp, ) -> Result where T: ArrowPrimitiveType, { match op { - "checked_add" => checked_binary(left, right, is_ansi_mode, false, |l, r| l.add_checked(r)), - "checked_sub" => checked_binary(left, right, is_ansi_mode, false, |l, r| l.sub_checked(r)), - "checked_mul" => checked_binary(left, right, is_ansi_mode, false, |l, r| l.mul_checked(r)), - "checked_div" => checked_binary(left, right, is_ansi_mode, true, |l, r| l.div_checked(r)), - _ => Err(DataFusionError::Internal(format!( - "Unsupported operation: {:?}", - op - ))), + MathOp::Add => checked_binary(left, right, is_ansi_mode, |l, r| l.add_checked(r)), + MathOp::Sub => checked_binary(left, right, is_ansi_mode, |l, r| l.sub_checked(r)), + MathOp::Mul => checked_binary(left, right, is_ansi_mode, |l, r| l.mul_checked(r)), + MathOp::Div => checked_binary(left, right, is_ansi_mode, |l, r| l.div_checked(r)), } } +fn ansi_arithmetic_kernel( + left: &PrimitiveArray, + right: &PrimitiveArray, + op: MathOp, +) -> Result +where + T: ArrowPrimitiveType, +{ + let result_array = match op { + MathOp::Add => numeric::add(left, right), + MathOp::Sub => numeric::sub(left, right), + MathOp::Mul => numeric::mul(left, right), + MathOp::Div => numeric::div(left, right), + }; + + result_array.map_err(|e| match e { + ArrowError::DivideByZero => divide_by_zero_error().into(), + _ => DataFusionError::from(SparkError::ArithmeticOverflow { + from_type: String::from("integer"), + }), + }) +} + fn checked_binary( left: &PrimitiveArray, right: &PrimitiveArray, is_ansi_mode: bool, - is_div: bool, op: F, ) -> Result where @@ -62,7 +89,7 @@ where F: Fn(T::Native, T::Native) -> Result, { if is_ansi_mode { - return arrow::compute::kernels::arity::try_binary::<_, _, _, T>(left, right, op) + return arity::try_binary::<_, _, _, T>(left, right, op) .map(|array| Arc::new(array) as ArrayRef) .map_err(|e| match e { ArrowError::DivideByZero => divide_by_zero_error().into(), @@ -84,18 +111,7 @@ where match op(l, r) { Ok(v) => *out = v, Err(_) => { - if !is_ansi_mode { - overflowed.push(i); - } else if nulls.as_ref().is_none_or(|n| n.is_valid(i)) { - return if is_div && r.is_zero() { - Err(divide_by_zero_error().into()) - } else { - Err(SparkError::ArithmeticOverflow { - from_type: String::from("integer"), - } - .into()) - }; - } + overflowed.push(i); } } } @@ -126,12 +142,13 @@ where Ok(Arc::new(PrimitiveArray::::new(values.into(), nulls)) as ArrayRef) } + pub fn checked_add( args: &[ColumnarValue], data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_add", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Add, eval_mode) } pub fn checked_sub( @@ -139,7 +156,7 @@ pub fn checked_sub( data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_sub", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Sub, eval_mode) } pub fn checked_mul( @@ -147,7 +164,7 @@ pub fn checked_mul( data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_mul", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Mul, eval_mode) } pub fn checked_div( @@ -155,13 +172,13 @@ pub fn checked_div( data_type: &DataType, eval_mode: EvalMode, ) -> Result { - checked_arithmetic_internal(args, data_type, "checked_div", eval_mode) + checked_arithmetic_internal(args, data_type, MathOp::Div, eval_mode) } fn checked_arithmetic_internal( args: &[ColumnarValue], data_type: &DataType, - op: &str, + op: MathOp, eval_mode: EvalMode, ) -> Result { let left = &args[0]; @@ -189,50 +206,68 @@ fn checked_arithmetic_internal( (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => (l.to_array()?, r.to_array()?), }; - // Rust only supports checked_arithmetic on numeric types let result_array = match data_type { - DataType::Int8 => try_arithmetic_kernel::( + DataType::Int8 if is_ansi_mode => ansi_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, + ), + DataType::Int8 => try_arithmetic_kernel::( + left_arr.as_primitive::(), + right_arr.as_primitive::(), is_ansi_mode, + op, ), - DataType::Int16 => try_arithmetic_kernel::( + DataType::Int16 if is_ansi_mode => ansi_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, + ), + DataType::Int16 => try_arithmetic_kernel::( + left_arr.as_primitive::(), + right_arr.as_primitive::(), is_ansi_mode, + op, ), - DataType::Int32 => try_arithmetic_kernel::( + DataType::Int32 if is_ansi_mode => ansi_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, + ), + DataType::Int32 => try_arithmetic_kernel::( + left_arr.as_primitive::(), + right_arr.as_primitive::(), is_ansi_mode, + op, ), - DataType::Int64 => try_arithmetic_kernel::( + DataType::Int64 if is_ansi_mode => ansi_arithmetic_kernel( left_arr.as_primitive::(), right_arr.as_primitive::(), op, + ), + DataType::Int64 => try_arithmetic_kernel::( + left_arr.as_primitive::(), + right_arr.as_primitive::(), is_ansi_mode, + op, ), - // Spark always casts division operands to floats - DataType::Float16 if (op == "checked_div") => try_arithmetic_kernel::( + DataType::Float16 if op == MathOp::Div => try_arithmetic_kernel::( left_arr.as_primitive::(), right_arr.as_primitive::(), - op, is_ansi_mode, + op, ), - DataType::Float32 if (op == "checked_div") => try_arithmetic_kernel::( + DataType::Float32 if op == MathOp::Div => try_arithmetic_kernel::( left_arr.as_primitive::(), right_arr.as_primitive::(), - op, is_ansi_mode, + op, ), - DataType::Float64 if (op == "checked_div") => try_arithmetic_kernel::( + DataType::Float64 if op == MathOp::Div => try_arithmetic_kernel::( left_arr.as_primitive::(), right_arr.as_primitive::(), - op, is_ansi_mode, + op, ), _ => Err(DataFusionError::Internal(format!( "Unsupported data type: {:?}", From c20948298aa3b5153a52cefbce96b95165e928aa Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Fri, 7 Aug 2026 21:47:59 +0400 Subject: [PATCH 4/6] address comments --- .../src/math_funcs/checked_arithmetic.rs | 291 +++++++++++------- 1 file changed, 180 insertions(+), 111 deletions(-) diff --git a/native/spark-expr/src/math_funcs/checked_arithmetic.rs b/native/spark-expr/src/math_funcs/checked_arithmetic.rs index 3a2a8bd5e3..a04ae76878 100644 --- a/native/spark-expr/src/math_funcs/checked_arithmetic.rs +++ b/native/spark-expr/src/math_funcs/checked_arithmetic.rs @@ -29,6 +29,8 @@ use arrow::error::ArrowError; use datafusion::common::DataFusionError; use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; +use common::ScalarValue; +use datafusion::common; #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum MathOp { @@ -41,64 +43,94 @@ enum MathOp { fn try_arithmetic_kernel( left: &PrimitiveArray, right: &PrimitiveArray, - is_ansi_mode: bool, op: MathOp, ) -> Result where T: ArrowPrimitiveType, { match op { - MathOp::Add => checked_binary(left, right, is_ansi_mode, |l, r| l.add_checked(r)), - MathOp::Sub => checked_binary(left, right, is_ansi_mode, |l, r| l.sub_checked(r)), - MathOp::Mul => checked_binary(left, right, is_ansi_mode, |l, r| l.mul_checked(r)), - MathOp::Div => checked_binary(left, right, is_ansi_mode, |l, r| l.div_checked(r)), + MathOp::Add => checked_binary(left, right, |l, r| l.add_checked(r)), + MathOp::Sub => checked_binary(left, right, |l, r| l.sub_checked(r)), + MathOp::Mul => checked_binary(left, right, |l, r| l.mul_checked(r)), + MathOp::Div => checked_binary(left, right, |l, r| l.div_checked(r)), } } -fn ansi_arithmetic_kernel( - left: &PrimitiveArray, - right: &PrimitiveArray, +fn ansi_arithmetic_kernel( + left: &ColumnarValue, + right: &ColumnarValue, op: MathOp, -) -> Result -where - T: ArrowPrimitiveType, -{ - let result_array = match op { - MathOp::Add => numeric::add(left, right), - MathOp::Sub => numeric::sub(left, right), - MathOp::Mul => numeric::mul(left, right), - MathOp::Div => numeric::div(left, right), +) -> Result { + let to_array = |cv: &ColumnarValue| -> Result<(ArrayRef, bool), DataFusionError> { + match cv { + ColumnarValue::Array(arr) => Ok((Arc::clone(arr), false)), + ColumnarValue::Scalar(scalar) => Ok((scalar.to_array()?, true)), + } }; - result_array.map_err(|e| match e { + let (left_arr, left_is_scalar) = to_array(left)?; + let (right_arr, right_is_scalar) = to_array(right)?; + + let run_kernel = |l: &dyn arrow::array::Datum, r: &dyn arrow::array::Datum| { + match op { + MathOp::Add => numeric::add(l, r), + MathOp::Sub => numeric::sub(l, r), + MathOp::Mul => numeric::mul(l, r), + MathOp::Div => numeric::div(l, r), + } + }; + + let result_array = match (left_is_scalar, right_is_scalar) { + (false, false) => run_kernel(&left_arr, &right_arr), + (false, true) => run_kernel(&left_arr, &arrow::array::Scalar::new(right_arr)), + (true, false) => run_kernel(&arrow::array::Scalar::new(left_arr), &right_arr), + (true, true) => run_kernel( + &arrow::array::Scalar::new(left_arr), + &arrow::array::Scalar::new(right_arr), + ), + }; + + let array = result_array.map_err(|e| match e { ArrowError::DivideByZero => divide_by_zero_error().into(), _ => DataFusionError::from(SparkError::ArithmeticOverflow { from_type: String::from("integer"), }), - }) + })?; + + if left_is_scalar && right_is_scalar { + let scalar_val = ScalarValue::try_from_array(array.as_ref(), 0)?; + Ok(ColumnarValue::Scalar(scalar_val)) + } else { + Ok(ColumnarValue::Array(array)) + } +} + +fn ansi_float_div( + left: &PrimitiveArray, + right: &PrimitiveArray, +) -> Result +where + T: ArrowPrimitiveType, +{ + arity::try_binary::<_, _, _, T>(left, right, |l, r| l.div_checked(r)) + .map(|array| Arc::new(array) as ArrayRef) + .map_err(|e| match e { + ArrowError::DivideByZero => divide_by_zero_error().into(), + _ => DataFusionError::from(SparkError::ArithmeticOverflow { + from_type: String::from("integer"), + }), + }) } fn checked_binary( left: &PrimitiveArray, right: &PrimitiveArray, - is_ansi_mode: bool, op: F, ) -> Result where T: ArrowPrimitiveType, F: Fn(T::Native, T::Native) -> Result, { - if is_ansi_mode { - return arity::try_binary::<_, _, _, T>(left, right, op) - .map(|array| Arc::new(array) as ArrayRef) - .map_err(|e| match e { - ArrowError::DivideByZero => divide_by_zero_error().into(), - _ => DataFusionError::from(SparkError::ArithmeticOverflow { - from_type: String::from("integer"), - }), - }); - } - let len = left.len(); let lhs = &left.values()[..len]; let rhs = &right.values()[..len]; @@ -195,87 +227,95 @@ fn checked_arithmetic_internal( } }; - let (left_arr, right_arr): (ArrayRef, ArrayRef) = match (left, right) { - (ColumnarValue::Array(l), ColumnarValue::Array(r)) => (Arc::clone(l), Arc::clone(r)), - (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { - (l.to_array_of_size(r.len())?, Arc::clone(r)) - } - (ColumnarValue::Array(l), ColumnarValue::Scalar(r)) => { - (Arc::clone(l), r.to_array_of_size(l.len())?) + match data_type { + DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::Int64 + if is_ansi_mode => + { + ansi_arithmetic_kernel(left, right, op) + } + _ => { + let (left_arr, right_arr): (ArrayRef, ArrayRef) = match (left, right) { + (ColumnarValue::Array(l), ColumnarValue::Array(r)) => (Arc::clone(l), Arc::clone(r)), + (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { + (l.to_array_of_size(r.len())?, Arc::clone(r)) + } + (ColumnarValue::Array(l), ColumnarValue::Scalar(r)) => { + (Arc::clone(l), r.to_array_of_size(l.len())?) + } + (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => (l.to_array()?, r.to_array()?), + }; + + let result_array = match data_type { + DataType::Int8 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Int16 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Int32 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Int64 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Float16 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ) + } + } + DataType::Float32 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ) + } + } + DataType::Float64 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ) + } + } + _ => Err(DataFusionError::Internal(format!( + "Unsupported data type: {:?}", + data_type + ))), + }; + Ok(ColumnarValue::Array(result_array?)) } - (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => (l.to_array()?, r.to_array()?), - }; - - let result_array = match data_type { - DataType::Int8 if is_ansi_mode => ansi_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ), - DataType::Int8 => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - DataType::Int16 if is_ansi_mode => ansi_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ), - DataType::Int16 => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - DataType::Int32 if is_ansi_mode => ansi_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ), - DataType::Int32 => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - DataType::Int64 if is_ansi_mode => ansi_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ), - DataType::Int64 => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - DataType::Float16 if op == MathOp::Div => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - DataType::Float32 if op == MathOp::Div => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - DataType::Float64 if op == MathOp::Div => try_arithmetic_kernel::( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - is_ansi_mode, - op, - ), - _ => Err(DataFusionError::Internal(format!( - "Unsupported data type: {:?}", - data_type - ))), - }; - - Ok(ColumnarValue::Array(result_array?)) + } } #[cfg(test)] @@ -367,4 +407,33 @@ mod tests { assert_eq!(result, Int32Array::from(vec![None, Some(2)])); assert_eq!(result.values()[0], 0); } + + #[test] + fn test_ansi_integer_div_by_zero() { + let args = int32_args(vec![Some(10)], vec![Some(0)]); + let result = checked_div(&args, &DataType::Int32, EvalMode::Ansi); + assert!(result.is_err()); + } + + #[test] + fn test_ansi_scalar_operand() { + let args = vec![ + ColumnarValue::Array(Arc::new(Int32Array::from(vec![Some(10), Some(20)]))), + ColumnarValue::Scalar(ScalarValue::Int32(Some(5))), + ]; + let result = as_int32(checked_add(&args, &DataType::Int32, EvalMode::Ansi).unwrap()); + assert_eq!(result, Int32Array::from(vec![Some(15), Some(25)])); + } + + #[test] + fn test_ansi_int64_overflow() { + let args = vec![ + ColumnarValue::Array(Arc::new(arrow::array::Int64Array::from(vec![Some( + i64::MAX, + )]))), + ColumnarValue::Array(Arc::new(arrow::array::Int64Array::from(vec![Some(1)]))), + ]; + let result = checked_add(&args, &DataType::Int64, EvalMode::Ansi); + assert!(result.is_err()); + } } From 9efc3a2d542678c6aebf910433bfbd4b3202218f Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sat, 8 Aug 2026 20:10:31 +0400 Subject: [PATCH 5/6] address comments --- .../src/math_funcs/checked_arithmetic.rs | 56 ++++++++++--------- 1 file changed, 29 insertions(+), 27 deletions(-) diff --git a/native/spark-expr/src/math_funcs/checked_arithmetic.rs b/native/spark-expr/src/math_funcs/checked_arithmetic.rs index a04ae76878..0594125e2d 100644 --- a/native/spark-expr/src/math_funcs/checked_arithmetic.rs +++ b/native/spark-expr/src/math_funcs/checked_arithmetic.rs @@ -29,6 +29,8 @@ use arrow::error::ArrowError; use datafusion::common::DataFusionError; use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; +use array::{Datum, Scalar}; +use arrow::array; use common::ScalarValue; use datafusion::common; @@ -61,17 +63,7 @@ fn ansi_arithmetic_kernel( right: &ColumnarValue, op: MathOp, ) -> Result { - let to_array = |cv: &ColumnarValue| -> Result<(ArrayRef, bool), DataFusionError> { - match cv { - ColumnarValue::Array(arr) => Ok((Arc::clone(arr), false)), - ColumnarValue::Scalar(scalar) => Ok((scalar.to_array()?, true)), - } - }; - - let (left_arr, left_is_scalar) = to_array(left)?; - let (right_arr, right_is_scalar) = to_array(right)?; - - let run_kernel = |l: &dyn arrow::array::Datum, r: &dyn arrow::array::Datum| { + let run_kernel = |l: &dyn Datum, r: &dyn Datum| { match op { MathOp::Add => numeric::add(l, r), MathOp::Sub => numeric::sub(l, r), @@ -80,14 +72,29 @@ fn ansi_arithmetic_kernel( } }; - let result_array = match (left_is_scalar, right_is_scalar) { - (false, false) => run_kernel(&left_arr, &right_arr), - (false, true) => run_kernel(&left_arr, &arrow::array::Scalar::new(right_arr)), - (true, false) => run_kernel(&arrow::array::Scalar::new(left_arr), &right_arr), - (true, true) => run_kernel( - &arrow::array::Scalar::new(left_arr), - &arrow::array::Scalar::new(right_arr), - ), + let result_array = match (left, right) { + (ColumnarValue::Array(l), ColumnarValue::Array(r)) => { + run_kernel(&l, &r) + } + (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { + let l_arr = l.to_array()?; + let l_scalar = Scalar::new(l_arr); + run_kernel(&l_scalar, &r) + } + (ColumnarValue::Array(l), ColumnarValue::Scalar(r)) => { + let r_arr = r.to_array()?; + let r_scalar = Scalar::new(r_arr); + run_kernel(&l, &r_scalar) + } + (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => { + let l_arr = l.to_array()?; + let r_arr = r.to_array()?; + let l_scalar = Scalar::new(l_arr); + let r_scalar = Scalar::new(r_arr); + let res = run_kernel(&l_scalar, &r_scalar)?; + let scalar_val = ScalarValue::try_from_array(res.as_ref(), 0)?; + return Ok(ColumnarValue::Scalar(scalar_val)); + } }; let array = result_array.map_err(|e| match e { @@ -97,12 +104,7 @@ fn ansi_arithmetic_kernel( }), })?; - if left_is_scalar && right_is_scalar { - let scalar_val = ScalarValue::try_from_array(array.as_ref(), 0)?; - Ok(ColumnarValue::Scalar(scalar_val)) - } else { - Ok(ColumnarValue::Array(array)) - } + Ok(ColumnarValue::Array(array)) } fn ansi_float_div( @@ -428,10 +430,10 @@ mod tests { #[test] fn test_ansi_int64_overflow() { let args = vec![ - ColumnarValue::Array(Arc::new(arrow::array::Int64Array::from(vec![Some( + ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some( i64::MAX, )]))), - ColumnarValue::Array(Arc::new(arrow::array::Int64Array::from(vec![Some(1)]))), + ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some(1)]))), ]; let result = checked_add(&args, &DataType::Int64, EvalMode::Ansi); assert!(result.is_err()); From 10d81fc5712da1d81c8f938779d4ac9c21a768ed Mon Sep 17 00:00:00 2001 From: Kazantsev Maksim Date: Sat, 8 Aug 2026 20:49:03 +0400 Subject: [PATCH 6/6] address comments --- .../spark-expr/benches/checked_arithmetic.rs | 46 ++++ .../src/math_funcs/checked_arithmetic.rs | 229 ++++++++++-------- 2 files changed, 174 insertions(+), 101 deletions(-) diff --git a/native/spark-expr/benches/checked_arithmetic.rs b/native/spark-expr/benches/checked_arithmetic.rs index 280852171c..ed03f3c616 100644 --- a/native/spark-expr/benches/checked_arithmetic.rs +++ b/native/spark-expr/benches/checked_arithmetic.rs @@ -18,6 +18,7 @@ use arrow::array::{ArrayRef, Float64Array, Int32Array, Int64Array}; use arrow::datatypes::DataType; use criterion::{criterion_group, criterion_main, Criterion}; +use datafusion::common::ScalarValue; use datafusion::physical_plan::ColumnarValue; use datafusion_comet_spark_expr::{checked_add, checked_div, checked_mul, checked_sub, EvalMode}; use std::hint::black_box; @@ -95,6 +96,21 @@ fn criterion_benchmark(c: &mut Criterion) { ColumnarValue::Array(f64_array(0x9e37_79b9, nulls)), ]; + let i32_scalar_right_args = [ + ColumnarValue::Array(i32_array(0x1234_5678, nulls)), + ColumnarValue::Scalar(ScalarValue::Int32(Some(42))), + ]; + + let i32_scalar_left_args = [ + ColumnarValue::Scalar(ScalarValue::Int32(Some(42))), + ColumnarValue::Array(i32_array(0x1234_5678, nulls)), + ]; + + let i64_scalar_args = [ + ColumnarValue::Array(i64_array(0x1234_5678, nulls)), + ColumnarValue::Scalar(ScalarValue::Int64(Some(100))), + ]; + group.bench_function(format!("checked_add_i32_ansi_{label}"), |b| { b.iter(|| { black_box(checked_add( @@ -144,6 +160,36 @@ fn criterion_benchmark(c: &mut Criterion) { )) }) }); + + group.bench_function(format!("checked_add_i32_scalar_right_ansi_{label}"), |b| { + b.iter(|| { + black_box(checked_add( + black_box(&i32_scalar_right_args), + &DataType::Int32, + EvalMode::Ansi, + )) + }) + }); + + group.bench_function(format!("checked_add_i32_scalar_left_ansi_{label}"), |b| { + b.iter(|| { + black_box(checked_add( + black_box(&i32_scalar_left_args), + &DataType::Int32, + EvalMode::Ansi, + )) + }) + }); + + group.bench_function(format!("checked_mul_i64_scalar_ansi_{label}"), |b| { + b.iter(|| { + black_box(checked_mul( + black_box(&i64_scalar_args), + &DataType::Int64, + EvalMode::Ansi, + )) + }) + }); } group.finish(); diff --git a/native/spark-expr/src/math_funcs/checked_arithmetic.rs b/native/spark-expr/src/math_funcs/checked_arithmetic.rs index 0594125e2d..8c0dc0fc33 100644 --- a/native/spark-expr/src/math_funcs/checked_arithmetic.rs +++ b/native/spark-expr/src/math_funcs/checked_arithmetic.rs @@ -19,6 +19,8 @@ use arrow::array::{Array, ArrowNativeTypeOp, BooleanBufferBuilder, PrimitiveArra use arrow::array::{ArrayRef, AsArray}; use crate::{divide_by_zero_error, EvalMode, SparkError}; +use array::{Datum, Scalar}; +use arrow::array; use arrow::buffer::NullBuffer; use arrow::compute::kernels::{arity, numeric}; use arrow::datatypes::{ @@ -27,12 +29,9 @@ use arrow::datatypes::{ }; use arrow::error::ArrowError; use datafusion::common::DataFusionError; +use datafusion::common::ScalarValue; use datafusion::physical_plan::ColumnarValue; use std::sync::Arc; -use array::{Datum, Scalar}; -use arrow::array; -use common::ScalarValue; -use datafusion::common; #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum MathOp { @@ -63,19 +62,15 @@ fn ansi_arithmetic_kernel( right: &ColumnarValue, op: MathOp, ) -> Result { - let run_kernel = |l: &dyn Datum, r: &dyn Datum| { - match op { - MathOp::Add => numeric::add(l, r), - MathOp::Sub => numeric::sub(l, r), - MathOp::Mul => numeric::mul(l, r), - MathOp::Div => numeric::div(l, r), - } + let run_kernel = |l: &dyn Datum, r: &dyn Datum| match op { + MathOp::Add => numeric::add(l, r), + MathOp::Sub => numeric::sub(l, r), + MathOp::Mul => numeric::mul(l, r), + MathOp::Div => numeric::div(l, r), }; let result_array = match (left, right) { - (ColumnarValue::Array(l), ColumnarValue::Array(r)) => { - run_kernel(&l, &r) - } + (ColumnarValue::Array(l), ColumnarValue::Array(r)) => run_kernel(&l, &r), (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { let l_arr = l.to_array()?; let l_scalar = Scalar::new(l_arr); @@ -209,6 +204,14 @@ pub fn checked_div( checked_arithmetic_internal(args, data_type, MathOp::Div, eval_mode) } +#[inline] +fn is_integer_type(data_type: &DataType) -> bool { + matches!( + data_type, + DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::Int64 + ) +} + fn checked_arithmetic_internal( args: &[ColumnarValue], data_type: &DataType, @@ -229,95 +232,93 @@ fn checked_arithmetic_internal( } }; - match data_type { - DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::Int64 - if is_ansi_mode => - { - ansi_arithmetic_kernel(left, right, op) - } - _ => { - let (left_arr, right_arr): (ArrayRef, ArrayRef) = match (left, right) { - (ColumnarValue::Array(l), ColumnarValue::Array(r)) => (Arc::clone(l), Arc::clone(r)), - (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { - (l.to_array_of_size(r.len())?, Arc::clone(r)) - } - (ColumnarValue::Array(l), ColumnarValue::Scalar(r)) => { - (Arc::clone(l), r.to_array_of_size(l.len())?) - } - (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => (l.to_array()?, r.to_array()?), - }; + // Early return for integer types in ANSI mode using the fast Datum/Scalar path + if is_ansi_mode && is_integer_type(data_type) { + return ansi_arithmetic_kernel(left, right, op); + } - let result_array = match data_type { - DataType::Int8 => try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ), - DataType::Int16 => try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), + // Materialize operands for Try-mode and float division + let (left_arr, right_arr): (ArrayRef, ArrayRef) = match (left, right) { + (ColumnarValue::Array(l), ColumnarValue::Array(r)) => (Arc::clone(l), Arc::clone(r)), + (ColumnarValue::Scalar(l), ColumnarValue::Array(r)) => { + (l.to_array_of_size(r.len())?, Arc::clone(r)) + } + (ColumnarValue::Array(l), ColumnarValue::Scalar(r)) => { + (Arc::clone(l), r.to_array_of_size(l.len())?) + } + (ColumnarValue::Scalar(l), ColumnarValue::Scalar(r)) => (l.to_array()?, r.to_array()?), + }; + + let result_array = match data_type { + DataType::Int8 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Int16 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Int32 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Int64 => try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + op, + ), + DataType::Float16 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), op, - ), - DataType::Int32 => try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), + ) + } + } + DataType::Float32 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), op, - ), - DataType::Int64 => try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), + ) + } + } + DataType::Float64 if op == MathOp::Div => { + if is_ansi_mode { + ansi_float_div( + left_arr.as_primitive::(), + right_arr.as_primitive::(), + ) + } else { + try_arithmetic_kernel( + left_arr.as_primitive::(), + right_arr.as_primitive::(), op, - ), - DataType::Float16 if op == MathOp::Div => { - if is_ansi_mode { - ansi_float_div( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - ) - } else { - try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ) - } - } - DataType::Float32 if op == MathOp::Div => { - if is_ansi_mode { - ansi_float_div( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - ) - } else { - try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ) - } - } - DataType::Float64 if op == MathOp::Div => { - if is_ansi_mode { - ansi_float_div( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - ) - } else { - try_arithmetic_kernel( - left_arr.as_primitive::(), - right_arr.as_primitive::(), - op, - ) - } - } - _ => Err(DataFusionError::Internal(format!( - "Unsupported data type: {:?}", - data_type - ))), - }; - Ok(ColumnarValue::Array(result_array?)) + ) + } } - } + _ => Err(DataFusionError::Internal(format!( + "Unsupported data type: {:?}", + data_type + ))), + }; + + Ok(ColumnarValue::Array(result_array?)) } #[cfg(test)] @@ -430,12 +431,38 @@ mod tests { #[test] fn test_ansi_int64_overflow() { let args = vec![ - ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some( - i64::MAX, - )]))), + ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some(i64::MAX)]))), ColumnarValue::Array(Arc::new(array::Int64Array::from(vec![Some(1)]))), ]; let result = checked_add(&args, &DataType::Int64, EvalMode::Ansi); assert!(result.is_err()); } + + #[test] + fn test_ansi_scalar_operands() { + let args_right = vec![ + ColumnarValue::Array(Arc::new(Int32Array::from(vec![Some(10), Some(20)]))), + ColumnarValue::Scalar(ScalarValue::Int32(Some(5))), + ]; + let res_right = + as_int32(checked_add(&args_right, &DataType::Int32, EvalMode::Ansi).unwrap()); + assert_eq!(res_right, Int32Array::from(vec![Some(15), Some(25)])); + + let args_left = vec![ + ColumnarValue::Scalar(ScalarValue::Int32(Some(5))), + ColumnarValue::Array(Arc::new(Int32Array::from(vec![Some(10), Some(20)]))), + ]; + let res_left = as_int32(checked_add(&args_left, &DataType::Int32, EvalMode::Ansi).unwrap()); + assert_eq!(res_left, Int32Array::from(vec![Some(15), Some(25)])); + + let args_scalar = vec![ + ColumnarValue::Scalar(ScalarValue::Int32(Some(10))), + ColumnarValue::Scalar(ScalarValue::Int32(Some(20))), + ]; + let res_scalar = checked_add(&args_scalar, &DataType::Int32, EvalMode::Ansi).unwrap(); + match res_scalar { + ColumnarValue::Scalar(ScalarValue::Int32(v)) => assert_eq!(v, Some(30)), + _ => panic!("Expected scalar result"), + } + } }