From e2755a341caa6690616eca17b4c8087780d62be7 Mon Sep 17 00:00:00 2001 From: ShreyeshArangath Date: Fri, 17 Jul 2026 16:50:15 -0700 Subject: [PATCH 1/6] Fix ANSI overflow behavior for unary minus Route integral UnaryMinus through Spark expression evaluation when ANSI mode is enabled so Int/Long overflow raises ARITHMETIC_OVERFLOW instead of wrapping. Update AuronExpressionSuite UnaryMinus test to validate overflow behavior for both vanilla Spark and Auron. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../apache/auron/AuronExpressionSuite.scala | 33 +++++++++++++++---- .../spark/sql/auron/NativeConverters.scala | 12 +++++++ 2 files changed, 38 insertions(+), 7 deletions(-) diff --git a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala index b0006d363..a94526cf4 100644 --- a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala +++ b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala @@ -33,16 +33,35 @@ class AuronExpressionSuite extends AuronQueryTest with BaseAuronSQLSuite { } test("UnaryMinus") { - // Negating Int.MinValue overflows. Under ANSI mode (default in Spark 4.x) vanilla Spark - // throws while the native engine wraps, so the comparison diverges. Disable ANSI so both - // engines wrap consistently and the boundary value can still be exercised. - withSQLConf("spark.sql.ansi.enabled" -> "false") { + withSQLConf("spark.sql.ansi.enabled" -> "true") { withTable("t1") { sql("create table t1(col1 int) using parquet") - sql( - "insert into t1 values(1), (2), (3), (3), (-1), (0), (null), (2147483647), (-2147483648)") - checkSparkAnswerAndOperator("SELECT negative(col1), -(col1) FROM t1") + sql("insert into t1 values(1), (0), (-2147483648)") + + withSQLConf("spark.auron.enable" -> "false") { + assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1")) + } + withSQLConf("spark.auron.enable" -> "true") { + assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1")) + } } } } + + private def assertArithmeticOverflow(df: => org.apache.spark.sql.DataFrame): Unit = { + val err = intercept[Exception] { + df.collect() + } + assert(allCauseMessages(err).contains("[ARITHMETIC_OVERFLOW]")) + } + + private def allCauseMessages(err: Throwable): String = { + val messages = scala.collection.mutable.ArrayBuffer.empty[String] + var current = err + while (current != null) { + Option(current.getMessage).foreach(messages += _) + current = current.getCause + } + messages.mkString(" | caused by: ") + } } diff --git a/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala b/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala index 378a8d662..6b031a8ab 100644 --- a/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala +++ b/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala @@ -114,6 +114,14 @@ object NativeConverters extends Logging { } } + private def shouldFallbackUnaryMinusToSpark(unaryMinus: UnaryMinus): Boolean = { + val ansiEnabled = SQLConf.get.getConfString("spark.sql.ansi.enabled", "false").toBoolean + ansiEnabled && unaryMinus.child.dataType match { + case ByteType | ShortType | IntegerType | LongType => true + case _ => false + } + } + def existTimestampType(dataType: DataType): Boolean = { dataType match { case TimestampType => @@ -563,6 +571,10 @@ object NativeConverters extends Logging { .setExpr(convertExprWithFallback(child, isPruningExpr, fallback)) .build()) } + // Spark ANSI mode requires overflow to raise, so use the Spark expression path for + // integral negation instead of the native wrapping implementation. + case unaryMinus: UnaryMinus if shouldFallbackUnaryMinusToSpark(unaryMinus) => + buildSparkUdfWrapperExpr(unaryMinus, fallback) case unaryMinus: UnaryMinus => buildExprNode { _.setNegative( From 954fbafb90916433ef16d7fafa1fbe7d5ca04719 Mon Sep 17 00:00:00 2001 From: ShreyeshArangath Date: Sat, 18 Jul 2026 08:34:50 -0700 Subject: [PATCH 2/6] Implement ANSI unary minus natively Add a checked native unary-minus expression for signed integral types, thread the ANSI flag through the plan proto, and update the Spark regression test to cover Int.MinValue and Long.MinValue with native execution preserved. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- native-engine/auron-planner/proto/auron.proto | 1 + native-engine/auron-planner/src/planner.rs | 20 +- native-engine/datafusion-ext-exprs/src/lib.rs | 1 + .../src/spark_negative.rs | 227 ++++++++++++++++++ .../apache/auron/AuronExpressionSuite.scala | 44 +++- .../spark/sql/auron/NativeConverters.scala | 14 +- 6 files changed, 291 insertions(+), 16 deletions(-) create mode 100644 native-engine/datafusion-ext-exprs/src/spark_negative.rs diff --git a/native-engine/auron-planner/proto/auron.proto b/native-engine/auron-planner/proto/auron.proto index a905c8a36..ef5c060a7 100644 --- a/native-engine/auron-planner/proto/auron.proto +++ b/native-engine/auron-planner/proto/auron.proto @@ -312,6 +312,7 @@ message PhysicalCastNode { message PhysicalNegativeNode { PhysicalExprNode expr = 1; + bool ansi_enabled = 2; } message PhysicalLikeExprNode { diff --git a/native-engine/auron-planner/src/planner.rs b/native-engine/auron-planner/src/planner.rs index b1ee15843..b3244ac01 100644 --- a/native-engine/auron-planner/src/planner.rs +++ b/native-engine/auron-planner/src/planner.rs @@ -55,6 +55,7 @@ use datafusion_ext_exprs::{ named_struct::NamedStructExpr, row_num::RowNumExpr, spark_monotonically_increasing_id::SparkMonotonicallyIncreasingIdExpr, spark_partition_id::SparkPartitionIdExpr, + spark_negative::SparkNegativeExpr, spark_scalar_subquery_wrapper::SparkScalarSubqueryWrapperExpr, spark_udf_wrapper::SparkUDFWrapperExpr, string_contains::StringContainsExpr, string_ends_with::StringEndsWithExpr, string_starts_with::StringStartsWithExpr, @@ -962,9 +963,22 @@ impl PhysicalPlanner { ExprType::NotExpr(e) => Arc::new(NotExpr::new( self.try_parse_physical_expr_box_required(&e.expr, input_schema)?, )), - ExprType::Negative(e) => Arc::new(NegativeExpr::new( - self.try_parse_physical_expr_box_required(&e.expr, input_schema)?, - )), + ExprType::Negative(e) => { + let expr = self.try_parse_physical_expr_box_required(&e.expr, input_schema)?; + if e.ansi_enabled { + match expr.data_type(input_schema)? { + datafusion::arrow::datatypes::DataType::Int8 + | datafusion::arrow::datatypes::DataType::Int16 + | datafusion::arrow::datatypes::DataType::Int32 + | datafusion::arrow::datatypes::DataType::Int64 => { + Arc::new(SparkNegativeExpr::new(expr)) + } + _ => Arc::new(NegativeExpr::new(expr)), + } + } else { + Arc::new(NegativeExpr::new(expr)) + } + } ExprType::InList(e) => { let expr = self.try_parse_physical_expr_box_required(&e.expr, input_schema)?; let dt = expr.data_type(input_schema)?; diff --git a/native-engine/datafusion-ext-exprs/src/lib.rs b/native-engine/datafusion-ext-exprs/src/lib.rs index 6400f7d21..4087e25d2 100644 --- a/native-engine/datafusion-ext-exprs/src/lib.rs +++ b/native-engine/datafusion-ext-exprs/src/lib.rs @@ -25,6 +25,7 @@ pub mod named_struct; pub mod row_num; pub mod spark_monotonically_increasing_id; pub mod spark_partition_id; +pub mod spark_negative; pub mod spark_scalar_subquery_wrapper; pub mod spark_udf_wrapper; pub mod string_contains; diff --git a/native-engine/datafusion-ext-exprs/src/spark_negative.rs b/native-engine/datafusion-ext-exprs/src/spark_negative.rs new file mode 100644 index 000000000..63809db2c --- /dev/null +++ b/native-engine/datafusion-ext-exprs/src/spark_negative.rs @@ -0,0 +1,227 @@ +// 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. + +use std::{ + any::Any, + fmt::{Debug, Display, Formatter}, + hash::{Hash, Hasher}, + sync::Arc, +}; + +use arrow::{ + array::{ArrayRef, Int16Array, Int32Array, Int64Array, Int8Array}, + datatypes::{DataType, Schema}, + record_batch::RecordBatch, +}; +use datafusion::{ + common::Result, + logical_expr::ColumnarValue, + physical_expr::{PhysicalExpr, PhysicalExprRef}, +}; +use datafusion_ext_commons::{df_execution_err, downcast_any}; + +pub struct SparkNegativeExpr { + expr: PhysicalExprRef, +} + +impl SparkNegativeExpr { + pub fn new(expr: PhysicalExprRef) -> Self { + Self { expr } + } +} + +impl Display for SparkNegativeExpr { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "negative({})", self.expr) + } +} + +impl Debug for SparkNegativeExpr { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "negative({})", self.expr) + } +} + +impl PartialEq for SparkNegativeExpr { + fn eq(&self, other: &Self) -> bool { + self.expr.eq(&other.expr) + } +} + +impl Eq for SparkNegativeExpr {} + +impl Hash for SparkNegativeExpr { + fn hash(&self, state: &mut H) { + self.expr.hash(state); + } +} + +impl PhysicalExpr for SparkNegativeExpr { + fn as_any(&self) -> &dyn Any { + self + } + + fn data_type(&self, input_schema: &Schema) -> Result { + self.expr.data_type(input_schema) + } + + fn nullable(&self, input_schema: &Schema) -> Result { + self.expr.nullable(input_schema) + } + + fn evaluate(&self, batch: &RecordBatch) -> Result { + let value = self.expr.evaluate(batch)?; + Ok(match value { + ColumnarValue::Scalar(scalar) => ColumnarValue::Scalar(negate_scalar(scalar)?), + ColumnarValue::Array(array) => ColumnarValue::Array(negate_array(array.as_ref())?), + }) + } + + fn children(&self) -> Vec<&PhysicalExprRef> { + vec![&self.expr] + } + + fn with_new_children( + self: Arc, + children: Vec, + ) -> Result { + Ok(Arc::new(Self::new(children[0].clone()))) + } + + fn fmt_sql(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "fmt_sql not used") + } +} + +fn negate_scalar(scalar: datafusion::common::ScalarValue) -> Result { + use datafusion::common::ScalarValue; + + Ok(match scalar { + ScalarValue::Int8(Some(v)) => ScalarValue::Int8(Some(negate_value(v)?)), + ScalarValue::Int16(Some(v)) => ScalarValue::Int16(Some(negate_value(v)?)), + ScalarValue::Int32(Some(v)) => ScalarValue::Int32(Some(negate_value(v)?)), + ScalarValue::Int64(Some(v)) => ScalarValue::Int64(Some(negate_value(v)?)), + ScalarValue::Int8(None) => ScalarValue::Int8(None), + ScalarValue::Int16(None) => ScalarValue::Int16(None), + ScalarValue::Int32(None) => ScalarValue::Int32(None), + ScalarValue::Int64(None) => ScalarValue::Int64(None), + other => return df_execution_err!("unsupported data type for SparkNegativeExpr: {other}"), + }) +} + +macro_rules! negate_primitive_array { + ($array:expr, $array_ty:ty) => {{ + let array = downcast_any!($array, $array_ty)?; + let mut values = Vec::with_capacity(array.len()); + for value in array.iter() { + values.push(match value { + Some(v) => Some(negate_value(v)?), + None => None, + }); + } + Ok(Arc::new(<$array_ty>::from(values)) as ArrayRef) + }}; +} + +fn negate_array(array: &dyn arrow::array::Array) -> Result { + match array.data_type() { + DataType::Int8 => negate_primitive_array!(array, Int8Array), + DataType::Int16 => negate_primitive_array!(array, Int16Array), + DataType::Int32 => negate_primitive_array!(array, Int32Array), + DataType::Int64 => negate_primitive_array!(array, Int64Array), + other => df_execution_err!("unsupported data type for SparkNegativeExpr: {other}"), + } +} + +fn negate_value(value: T) -> Result +where + T: CheckedNeg, +{ + value + .checked_neg() + .ok_or_else(|| datafusion::common::DataFusionError::Execution( + "[ARITHMETIC_OVERFLOW] arithmetic overflow in unary minus".to_string(), + )) +} + +trait CheckedNeg { + fn checked_neg(self) -> Option + where + Self: Sized; +} + +impl CheckedNeg for i8 { + fn checked_neg(self) -> Option { + i8::checked_neg(self) + } +} + +impl CheckedNeg for i16 { + fn checked_neg(self) -> Option { + i16::checked_neg(self) + } +} + +impl CheckedNeg for i32 { + fn checked_neg(self) -> Option { + i32::checked_neg(self) + } +} + +impl CheckedNeg for i64 { + fn checked_neg(self) -> Option { + i64::checked_neg(self) + } +} + +#[cfg(test)] +mod test { + use std::{error::Error, sync::Arc}; + + use arrow::{ + array::{ArrayRef, Int32Array, Int64Array}, + datatypes::{DataType, Field, Schema}, + record_batch::RecordBatch, + }; + use datafusion::physical_expr::{PhysicalExpr, expressions::Column}; + + use super::SparkNegativeExpr; + + #[test] + fn test_int32_array() -> Result<(), Box> { + let array: ArrayRef = Arc::new(Int32Array::from(vec![Some(1), Some(-2), None, Some(3)])); + let batch = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("c", DataType::Int32, true)])), + vec![array], + )?; + let expr = Arc::new(SparkNegativeExpr::new(Arc::new(Column::new("c", 0)))); + let output = expr.evaluate(&batch)?.into_array(batch.num_rows())?; + let expected: ArrayRef = Arc::new(Int32Array::from(vec![Some(-1), Some(2), None, Some(-3)])); + assert_eq!(&output, &expected); + Ok(()) + } + + #[test] + fn test_int64_scalar_overflow() -> Result<(), Box> { + let batch = RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("c", DataType::Int64, true)])), + vec![Arc::new(Int64Array::from(vec![Some(i64::MIN)])) as ArrayRef], + )?; + let expr = Arc::new(SparkNegativeExpr::new(Arc::new(Column::new("c", 0)))); + let err = expr.evaluate(&batch).expect_err("expected overflow"); + assert!(err.to_string().contains("[ARITHMETIC_OVERFLOW]")); + Ok(()) + } +} diff --git a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala index a94526cf4..dd2729565 100644 --- a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala +++ b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala @@ -36,14 +36,44 @@ class AuronExpressionSuite extends AuronQueryTest with BaseAuronSQLSuite { withSQLConf("spark.sql.ansi.enabled" -> "true") { withTable("t1") { sql("create table t1(col1 int) using parquet") - sql("insert into t1 values(1), (0), (-2147483648)") + sql(""" + |insert into t1 values + | (1), + | (0), + | (-2147483648) + |""".stripMargin) withSQLConf("spark.auron.enable" -> "false") { assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1")) } withSQLConf("spark.auron.enable" -> "true") { + val df = sql("SELECT negative(col1), -(col1) FROM t1") + assertArithmeticOverflow(df) + assertNativePlan(df) + } + } + } + } + + test("UnaryMinusLong") { + withSQLConf("spark.sql.ansi.enabled" -> "true") { + withTable("t1") { + sql("create table t1(col1 bigint) using parquet") + sql(""" + |insert into t1 values + | (1), + | (0), + | (cast(-9223372036854775808 as bigint)) + |""".stripMargin) + + withSQLConf("spark.auron.enable" -> "false") { assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1")) } + withSQLConf("spark.auron.enable" -> "true") { + val df = sql("SELECT negative(col1), -(col1) FROM t1") + assertArithmeticOverflow(df) + assertNativePlan(df) + } } } } @@ -64,4 +94,16 @@ class AuronExpressionSuite extends AuronQueryTest with BaseAuronSQLSuite { } messages.mkString(" | caused by: ") } + + private def assertNativePlan(df: org.apache.spark.sql.DataFrame): Unit = { + val plan = stripAQEPlan(df.queryExecution.executedPlan) + plan + .collectFirst { case op if !isNativeOrPassThrough(op) => op } + .foreach { op => + fail(s""" + |Found non-native operator: ${op.nodeName} + |plan: + |${plan}""".stripMargin) + } + } } diff --git a/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala b/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala index 6b031a8ab..2a24f48d7 100644 --- a/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala +++ b/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala @@ -114,14 +114,6 @@ object NativeConverters extends Logging { } } - private def shouldFallbackUnaryMinusToSpark(unaryMinus: UnaryMinus): Boolean = { - val ansiEnabled = SQLConf.get.getConfString("spark.sql.ansi.enabled", "false").toBoolean - ansiEnabled && unaryMinus.child.dataType match { - case ByteType | ShortType | IntegerType | LongType => true - case _ => false - } - } - def existTimestampType(dataType: DataType): Boolean = { dataType match { case TimestampType => @@ -571,16 +563,14 @@ object NativeConverters extends Logging { .setExpr(convertExprWithFallback(child, isPruningExpr, fallback)) .build()) } - // Spark ANSI mode requires overflow to raise, so use the Spark expression path for - // integral negation instead of the native wrapping implementation. - case unaryMinus: UnaryMinus if shouldFallbackUnaryMinusToSpark(unaryMinus) => - buildSparkUdfWrapperExpr(unaryMinus, fallback) case unaryMinus: UnaryMinus => buildExprNode { _.setNegative( pb.PhysicalNegativeNode .newBuilder() .setExpr(convertExprWithFallback(unaryMinus.child, isPruningExpr, fallback)) + .setAnsiEnabled( + SQLConf.get.getConfString("spark.sql.ansi.enabled", "false").toBoolean) .build()) } From 801ece9d31268fb7ba69d14abb253419941a0364 Mon Sep 17 00:00:00 2001 From: ShreyeshArangath Date: Sat, 18 Jul 2026 09:00:14 -0700 Subject: [PATCH 3/6] Consolidate native unary minus handling Use SparkNegativeExpr for all unary-minus plans. It delegates standard and non-ANSI behavior to DataFusion NegativeExpr and applies checked integral negation only when ANSI mode is enabled. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- native-engine/auron-planner/src/planner.rs | 27 +-- native-engine/datafusion-ext-exprs/src/lib.rs | 2 +- .../src/spark_negative.rs | 167 +++++++++++------- 3 files changed, 107 insertions(+), 89 deletions(-) diff --git a/native-engine/auron-planner/src/planner.rs b/native-engine/auron-planner/src/planner.rs index b3244ac01..72a89c86e 100644 --- a/native-engine/auron-planner/src/planner.rs +++ b/native-engine/auron-planner/src/planner.rs @@ -42,8 +42,8 @@ use datafusion::{ physical_plan::{ ColumnStatistics, ExecutionPlan, Statistics, expressions as phys_expr, expressions::{ - BinaryExpr, CaseExpr, CastExpr, Column, IsNotNullExpr, IsNullExpr, Literal, - NegativeExpr, NotExpr, PhysicalSortExpr, + BinaryExpr, CaseExpr, CastExpr, Column, IsNotNullExpr, IsNullExpr, Literal, NotExpr, + PhysicalSortExpr, }, metrics::ExecutionPlanMetricsSet, }, @@ -54,8 +54,7 @@ use datafusion_ext_exprs::{ get_indexed_field::GetIndexedFieldExpr, get_map_value::GetMapValueExpr, named_struct::NamedStructExpr, row_num::RowNumExpr, spark_monotonically_increasing_id::SparkMonotonicallyIncreasingIdExpr, - spark_partition_id::SparkPartitionIdExpr, - spark_negative::SparkNegativeExpr, + spark_negative::SparkNegativeExpr, spark_partition_id::SparkPartitionIdExpr, spark_scalar_subquery_wrapper::SparkScalarSubqueryWrapperExpr, spark_udf_wrapper::SparkUDFWrapperExpr, string_contains::StringContainsExpr, string_ends_with::StringEndsWithExpr, string_starts_with::StringStartsWithExpr, @@ -963,22 +962,10 @@ impl PhysicalPlanner { ExprType::NotExpr(e) => Arc::new(NotExpr::new( self.try_parse_physical_expr_box_required(&e.expr, input_schema)?, )), - ExprType::Negative(e) => { - let expr = self.try_parse_physical_expr_box_required(&e.expr, input_schema)?; - if e.ansi_enabled { - match expr.data_type(input_schema)? { - datafusion::arrow::datatypes::DataType::Int8 - | datafusion::arrow::datatypes::DataType::Int16 - | datafusion::arrow::datatypes::DataType::Int32 - | datafusion::arrow::datatypes::DataType::Int64 => { - Arc::new(SparkNegativeExpr::new(expr)) - } - _ => Arc::new(NegativeExpr::new(expr)), - } - } else { - Arc::new(NegativeExpr::new(expr)) - } - } + ExprType::Negative(e) => Arc::new(SparkNegativeExpr::new( + self.try_parse_physical_expr_box_required(&e.expr, input_schema)?, + e.ansi_enabled, + )), ExprType::InList(e) => { let expr = self.try_parse_physical_expr_box_required(&e.expr, input_schema)?; let dt = expr.data_type(input_schema)?; diff --git a/native-engine/datafusion-ext-exprs/src/lib.rs b/native-engine/datafusion-ext-exprs/src/lib.rs index 4087e25d2..36389ae46 100644 --- a/native-engine/datafusion-ext-exprs/src/lib.rs +++ b/native-engine/datafusion-ext-exprs/src/lib.rs @@ -24,8 +24,8 @@ pub mod get_map_value; pub mod named_struct; pub mod row_num; pub mod spark_monotonically_increasing_id; -pub mod spark_partition_id; pub mod spark_negative; +pub mod spark_partition_id; pub mod spark_scalar_subquery_wrapper; pub mod spark_udf_wrapper; pub mod string_contains; diff --git a/native-engine/datafusion-ext-exprs/src/spark_negative.rs b/native-engine/datafusion-ext-exprs/src/spark_negative.rs index 63809db2c..22644a66d 100644 --- a/native-engine/datafusion-ext-exprs/src/spark_negative.rs +++ b/native-engine/datafusion-ext-exprs/src/spark_negative.rs @@ -21,24 +21,26 @@ use std::{ }; use arrow::{ - array::{ArrayRef, Int16Array, Int32Array, Int64Array, Int8Array}, + array::{ArrayRef, Int8Array, Int16Array, Int32Array, Int64Array}, datatypes::{DataType, Schema}, record_batch::RecordBatch, }; use datafusion::{ - common::Result, + common::{Result, ScalarValue}, logical_expr::ColumnarValue, physical_expr::{PhysicalExpr, PhysicalExprRef}, + physical_plan::expressions::NegativeExpr, }; use datafusion_ext_commons::{df_execution_err, downcast_any}; pub struct SparkNegativeExpr { expr: PhysicalExprRef, + ansi_enabled: bool, } impl SparkNegativeExpr { - pub fn new(expr: PhysicalExprRef) -> Self { - Self { expr } + pub fn new(expr: PhysicalExprRef, ansi_enabled: bool) -> Self { + Self { expr, ansi_enabled } } } @@ -56,7 +58,7 @@ impl Debug for SparkNegativeExpr { impl PartialEq for SparkNegativeExpr { fn eq(&self, other: &Self) -> bool { - self.expr.eq(&other.expr) + self.expr.eq(&other.expr) && self.ansi_enabled == other.ansi_enabled } } @@ -65,6 +67,7 @@ impl Eq for SparkNegativeExpr {} impl Hash for SparkNegativeExpr { fn hash(&self, state: &mut H) { self.expr.hash(state); + self.ansi_enabled.hash(state); } } @@ -82,11 +85,16 @@ impl PhysicalExpr for SparkNegativeExpr { } fn evaluate(&self, batch: &RecordBatch) -> Result { - let value = self.expr.evaluate(batch)?; - Ok(match value { - ColumnarValue::Scalar(scalar) => ColumnarValue::Scalar(negate_scalar(scalar)?), - ColumnarValue::Array(array) => ColumnarValue::Array(negate_array(array.as_ref())?), - }) + if self.ansi_enabled + && matches!( + self.expr.data_type(batch.schema().as_ref())?, + DataType::Int8 | DataType::Int16 | DataType::Int32 | DataType::Int64 + ) + { + return checked_negate(self.expr.evaluate(batch)?); + } + + NegativeExpr::new(self.expr.clone()).evaluate(batch) } fn children(&self) -> Vec<&PhysicalExprRef> { @@ -97,7 +105,7 @@ impl PhysicalExpr for SparkNegativeExpr { self: Arc, children: Vec, ) -> Result { - Ok(Arc::new(Self::new(children[0].clone()))) + Ok(Arc::new(Self::new(children[0].clone(), self.ansi_enabled))) } fn fmt_sql(&self, f: &mut Formatter<'_>) -> std::fmt::Result { @@ -105,29 +113,34 @@ impl PhysicalExpr for SparkNegativeExpr { } } -fn negate_scalar(scalar: datafusion::common::ScalarValue) -> Result { - use datafusion::common::ScalarValue; +fn checked_negate(value: ColumnarValue) -> Result { + Ok(match value { + ColumnarValue::Scalar(scalar) => ColumnarValue::Scalar(checked_negate_scalar(scalar)?), + ColumnarValue::Array(array) => ColumnarValue::Array(checked_negate_array(array.as_ref())?), + }) +} +fn checked_negate_scalar(scalar: ScalarValue) -> Result { Ok(match scalar { - ScalarValue::Int8(Some(v)) => ScalarValue::Int8(Some(negate_value(v)?)), - ScalarValue::Int16(Some(v)) => ScalarValue::Int16(Some(negate_value(v)?)), - ScalarValue::Int32(Some(v)) => ScalarValue::Int32(Some(negate_value(v)?)), - ScalarValue::Int64(Some(v)) => ScalarValue::Int64(Some(negate_value(v)?)), + ScalarValue::Int8(Some(v)) => ScalarValue::Int8(Some(checked_negate_value(v)?)), + ScalarValue::Int16(Some(v)) => ScalarValue::Int16(Some(checked_negate_value(v)?)), + ScalarValue::Int32(Some(v)) => ScalarValue::Int32(Some(checked_negate_value(v)?)), + ScalarValue::Int64(Some(v)) => ScalarValue::Int64(Some(checked_negate_value(v)?)), ScalarValue::Int8(None) => ScalarValue::Int8(None), ScalarValue::Int16(None) => ScalarValue::Int16(None), ScalarValue::Int32(None) => ScalarValue::Int32(None), ScalarValue::Int64(None) => ScalarValue::Int64(None), - other => return df_execution_err!("unsupported data type for SparkNegativeExpr: {other}"), + other => return df_execution_err!("unsupported ANSI negative data type: {other}"), }) } -macro_rules! negate_primitive_array { +macro_rules! checked_negate_primitive_array { ($array:expr, $array_ty:ty) => {{ let array = downcast_any!($array, $array_ty)?; let mut values = Vec::with_capacity(array.len()); for value in array.iter() { values.push(match value { - Some(v) => Some(negate_value(v)?), + Some(v) => Some(checked_negate_value(v)?), None => None, }); } @@ -135,25 +148,25 @@ macro_rules! negate_primitive_array { }}; } -fn negate_array(array: &dyn arrow::array::Array) -> Result { +fn checked_negate_array(array: &dyn arrow::array::Array) -> Result { match array.data_type() { - DataType::Int8 => negate_primitive_array!(array, Int8Array), - DataType::Int16 => negate_primitive_array!(array, Int16Array), - DataType::Int32 => negate_primitive_array!(array, Int32Array), - DataType::Int64 => negate_primitive_array!(array, Int64Array), - other => df_execution_err!("unsupported data type for SparkNegativeExpr: {other}"), + DataType::Int8 => checked_negate_primitive_array!(array, Int8Array), + DataType::Int16 => checked_negate_primitive_array!(array, Int16Array), + DataType::Int32 => checked_negate_primitive_array!(array, Int32Array), + DataType::Int64 => checked_negate_primitive_array!(array, Int64Array), + other => df_execution_err!("unsupported ANSI negative data type: {other}"), } } -fn negate_value(value: T) -> Result +fn checked_negate_value(value: T) -> Result where T: CheckedNeg, { - value - .checked_neg() - .ok_or_else(|| datafusion::common::DataFusionError::Execution( + value.checked_neg().ok_or_else(|| { + datafusion::common::DataFusionError::Execution( "[ARITHMETIC_OVERFLOW] arithmetic overflow in unary minus".to_string(), - )) + ) + }) } trait CheckedNeg { @@ -162,36 +175,26 @@ trait CheckedNeg { Self: Sized; } -impl CheckedNeg for i8 { - fn checked_neg(self) -> Option { - i8::checked_neg(self) - } -} - -impl CheckedNeg for i16 { - fn checked_neg(self) -> Option { - i16::checked_neg(self) - } -} - -impl CheckedNeg for i32 { - fn checked_neg(self) -> Option { - i32::checked_neg(self) - } +macro_rules! impl_checked_neg { + ($($ty:ty),+) => { + $( + impl CheckedNeg for $ty { + fn checked_neg(self) -> Option { + <$ty>::checked_neg(self) + } + } + )+ + }; } -impl CheckedNeg for i64 { - fn checked_neg(self) -> Option { - i64::checked_neg(self) - } -} +impl_checked_neg!(i8, i16, i32, i64); #[cfg(test)] mod test { use std::{error::Error, sync::Arc}; use arrow::{ - array::{ArrayRef, Int32Array, Int64Array}, + array::{ArrayRef, Float64Array, Int32Array, Int64Array}, datatypes::{DataType, Field, Schema}, record_batch::RecordBatch, }; @@ -200,28 +203,56 @@ mod test { use super::SparkNegativeExpr; #[test] - fn test_int32_array() -> Result<(), Box> { - let array: ArrayRef = Arc::new(Int32Array::from(vec![Some(1), Some(-2), None, Some(3)])); - let batch = RecordBatch::try_new( - Arc::new(Schema::new(vec![Field::new("c", DataType::Int32, true)])), - vec![array], + fn test_ansi_checked_negation() -> Result<(), Box> { + let batch = batch( + DataType::Int64, + Arc::new(Int64Array::from(vec![Some(i64::MIN)])), )?; - let expr = Arc::new(SparkNegativeExpr::new(Arc::new(Column::new("c", 0)))); - let output = expr.evaluate(&batch)?.into_array(batch.num_rows())?; - let expected: ArrayRef = Arc::new(Int32Array::from(vec![Some(-1), Some(2), None, Some(-3)])); + let expr = expression(true); + let err = expr.evaluate(&batch).expect_err("expected overflow"); + assert!(err.to_string().contains("[ARITHMETIC_OVERFLOW]")); + Ok(()) + } + + #[test] + fn test_non_ansi_wrapping_negation() -> Result<(), Box> { + let batch = batch( + DataType::Int32, + Arc::new(Int32Array::from(vec![Some(i32::MIN), Some(1), None])), + )?; + let output = expression(false) + .evaluate(&batch)? + .into_array(batch.num_rows())?; + let expected: ArrayRef = Arc::new(Int32Array::from(vec![Some(i32::MIN), Some(-1), None])); assert_eq!(&output, &expected); Ok(()) } #[test] - fn test_int64_scalar_overflow() -> Result<(), Box> { - let batch = RecordBatch::try_new( - Arc::new(Schema::new(vec![Field::new("c", DataType::Int64, true)])), - vec![Arc::new(Int64Array::from(vec![Some(i64::MIN)])) as ArrayRef], + fn test_float_delegates_to_datafusion() -> Result<(), Box> { + let batch = batch( + DataType::Float64, + Arc::new(Float64Array::from(vec![Some(1.5), None])), )?; - let expr = Arc::new(SparkNegativeExpr::new(Arc::new(Column::new("c", 0)))); - let err = expr.evaluate(&batch).expect_err("expected overflow"); - assert!(err.to_string().contains("[ARITHMETIC_OVERFLOW]")); + let output = expression(true) + .evaluate(&batch)? + .into_array(batch.num_rows())?; + let expected: ArrayRef = Arc::new(Float64Array::from(vec![Some(-1.5), None])); + assert_eq!(&output, &expected); Ok(()) } + + fn expression(ansi_enabled: bool) -> Arc { + Arc::new(SparkNegativeExpr::new( + Arc::new(Column::new("c", 0)), + ansi_enabled, + )) + } + + fn batch(data_type: DataType, array: ArrayRef) -> Result> { + Ok(RecordBatch::try_new( + Arc::new(Schema::new(vec![Field::new("c", data_type, true)])), + vec![array], + )?) + } } From 118c529a216d358f073627a8d72acccc08b219b7 Mon Sep 17 00:00:00 2001 From: ShreyeshArangath Date: Sat, 18 Jul 2026 20:38:57 -0700 Subject: [PATCH 4/6] Fix unary minus test across Spark versions Accept Spark's version-specific overflow wording while still requiring the unary minus plan to execute natively. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../src/test/scala/org/apache/auron/AuronExpressionSuite.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala index dd2729565..b0b8d977e 100644 --- a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala +++ b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala @@ -82,7 +82,7 @@ class AuronExpressionSuite extends AuronQueryTest with BaseAuronSQLSuite { val err = intercept[Exception] { df.collect() } - assert(allCauseMessages(err).contains("[ARITHMETIC_OVERFLOW]")) + assert(allCauseMessages(err).toLowerCase.contains("overflow")) } private def allCauseMessages(err: Throwable): String = { From e657a5ac0e3c57439856627d37c844ee7f8d0034 Mon Sep 17 00:00:00 2001 From: ShreyeshArangath Date: Sun, 19 Jul 2026 20:26:55 -0700 Subject: [PATCH 5/6] Honor Spark default ANSI setting Use Spark's typed ANSI configuration and cover explicit ANSI overflow, non-ANSI wrapping, and Spark's version-specific default behavior. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../apache/auron/AuronExpressionSuite.scala | 40 ++++++++++++++++--- .../spark/sql/auron/NativeConverters.scala | 3 +- 2 files changed, 35 insertions(+), 8 deletions(-) diff --git a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala index b0b8d977e..29b04d897 100644 --- a/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala +++ b/spark-extension-shims-spark/src/test/scala/org/apache/auron/AuronExpressionSuite.scala @@ -44,11 +44,11 @@ class AuronExpressionSuite extends AuronQueryTest with BaseAuronSQLSuite { |""".stripMargin) withSQLConf("spark.auron.enable" -> "false") { - assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1")) + assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1"), "overflow") } withSQLConf("spark.auron.enable" -> "true") { val df = sql("SELECT negative(col1), -(col1) FROM t1") - assertArithmeticOverflow(df) + assertArithmeticOverflow(df, "[ARITHMETIC_OVERFLOW]") assertNativePlan(df) } } @@ -67,22 +67,50 @@ class AuronExpressionSuite extends AuronQueryTest with BaseAuronSQLSuite { |""".stripMargin) withSQLConf("spark.auron.enable" -> "false") { - assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1")) + assertArithmeticOverflow(sql("SELECT negative(col1), -(col1) FROM t1"), "overflow") } withSQLConf("spark.auron.enable" -> "true") { val df = sql("SELECT negative(col1), -(col1) FROM t1") - assertArithmeticOverflow(df) + assertArithmeticOverflow(df, "[ARITHMETIC_OVERFLOW]") assertNativePlan(df) } } } } - private def assertArithmeticOverflow(df: => org.apache.spark.sql.DataFrame): Unit = { + test("UnaryMinus without ANSI") { + withSQLConf("spark.sql.ansi.enabled" -> "false") { + withTable("t1") { + sql("create table t1(col1 int) using parquet") + sql( + "insert into t1 values(1), (2), (3), (3), (-1), (0), (null), (2147483647), (-2147483648)") + checkSparkAnswerAndOperator("SELECT negative(col1), -(col1) FROM t1") + } + } + } + + test("UnaryMinus honors Spark's default ANSI setting") { + withTable("t1") { + sql("create table t1(col1 int) using parquet") + sql("insert into t1 values(-2147483648)") + + if (spark.conf.get("spark.sql.ansi.enabled").toBoolean) { + val df = sql("SELECT negative(col1), -(col1) FROM t1") + assertArithmeticOverflow(df, "[ARITHMETIC_OVERFLOW]") + assertNativePlan(df) + } else { + checkSparkAnswerAndOperator("SELECT negative(col1), -(col1) FROM t1") + } + } + } + + private def assertArithmeticOverflow( + df: => org.apache.spark.sql.DataFrame, + expectedMessage: String): Unit = { val err = intercept[Exception] { df.collect() } - assert(allCauseMessages(err).toLowerCase.contains("overflow")) + assert(allCauseMessages(err).toLowerCase.contains(expectedMessage.toLowerCase)) } private def allCauseMessages(err: Throwable): String = { diff --git a/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala b/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala index 2a24f48d7..adecbc59e 100644 --- a/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala +++ b/spark-extension/src/main/scala/org/apache/spark/sql/auron/NativeConverters.scala @@ -569,8 +569,7 @@ object NativeConverters extends Logging { pb.PhysicalNegativeNode .newBuilder() .setExpr(convertExprWithFallback(unaryMinus.child, isPruningExpr, fallback)) - .setAnsiEnabled( - SQLConf.get.getConfString("spark.sql.ansi.enabled", "false").toBoolean) + .setAnsiEnabled(SQLConf.get.ansiEnabled) .build()) } From 728434ce178d429d66dc675bfefc728ef011f616 Mon Sep 17 00:00:00 2001 From: ShreyeshArangath Date: Mon, 20 Jul 2026 19:23:41 -0700 Subject: [PATCH 6/6] Test ANSI negation across integer widths Cover overflow behavior for Int8, Int16, Int32, and Int64 in the native checked-negation unit test. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> --- .../src/spark_negative.rs | 35 ++++++++++++++----- 1 file changed, 27 insertions(+), 8 deletions(-) diff --git a/native-engine/datafusion-ext-exprs/src/spark_negative.rs b/native-engine/datafusion-ext-exprs/src/spark_negative.rs index 22644a66d..15e395111 100644 --- a/native-engine/datafusion-ext-exprs/src/spark_negative.rs +++ b/native-engine/datafusion-ext-exprs/src/spark_negative.rs @@ -194,7 +194,7 @@ mod test { use std::{error::Error, sync::Arc}; use arrow::{ - array::{ArrayRef, Float64Array, Int32Array, Int64Array}, + array::{ArrayRef, Float64Array, Int8Array, Int16Array, Int32Array, Int64Array}, datatypes::{DataType, Field, Schema}, record_batch::RecordBatch, }; @@ -204,13 +204,32 @@ mod test { #[test] fn test_ansi_checked_negation() -> Result<(), Box> { - let batch = batch( - DataType::Int64, - Arc::new(Int64Array::from(vec![Some(i64::MIN)])), - )?; - let expr = expression(true); - let err = expr.evaluate(&batch).expect_err("expected overflow"); - assert!(err.to_string().contains("[ARITHMETIC_OVERFLOW]")); + let cases: Vec<(DataType, ArrayRef)> = vec![ + ( + DataType::Int8, + Arc::new(Int8Array::from(vec![Some(i8::MIN)])), + ), + ( + DataType::Int16, + Arc::new(Int16Array::from(vec![Some(i16::MIN)])), + ), + ( + DataType::Int32, + Arc::new(Int32Array::from(vec![Some(i32::MIN)])), + ), + ( + DataType::Int64, + Arc::new(Int64Array::from(vec![Some(i64::MIN)])), + ), + ]; + + for (data_type, array) in cases { + let batch = batch(data_type, array)?; + let err = expression(true) + .evaluate(&batch) + .expect_err("expected overflow"); + assert!(err.to_string().contains("[ARITHMETIC_OVERFLOW]")); + } Ok(()) }