From d27cd8dad121caf25ace3b27b61f7ef443f07422 Mon Sep 17 00:00:00 2001 From: 0lai0 Date: Wed, 7 Oct 2026 23:59:44 +0300 Subject: [PATCH] fix: match Spark's ANSI integral SUM overflow error --- native/spark-expr/src/agg_funcs/sum_int.rs | 100 ++++++++++++++++-- native/spark-expr/src/lib.rs | 9 ++ .../comet/exec/CometAggregateSuite.scala | 89 ++++++++-------- 3 files changed, 149 insertions(+), 49 deletions(-) diff --git a/native/spark-expr/src/agg_funcs/sum_int.rs b/native/spark-expr/src/agg_funcs/sum_int.rs index 56a8c122c41..7e243580347 100644 --- a/native/spark-expr/src/agg_funcs/sum_int.rs +++ b/native/spark-expr/src/agg_funcs/sum_int.rs @@ -15,7 +15,7 @@ // specific language governing permissions and limitations // under the License. -use crate::{arithmetic_overflow_error, EvalMode}; +use crate::{long_add_overflow_error, EvalMode}; use arrow::array::{ as_primitive_array, cast::AsArray, Array, ArrayRef, ArrowNativeTypeOp, ArrowPrimitiveType, BooleanArray, Int64Array, PrimitiveArray, @@ -201,7 +201,7 @@ impl Accumulator for SumIntegerAccumulatorAnsi { })?; sum = v .add_checked(sum) - .map_err(|_| DataFusionError::from(arithmetic_overflow_error("integer")))?; + .map_err(|_| DataFusionError::from(long_add_overflow_error()))?; } } Ok(sum) @@ -574,10 +574,12 @@ impl GroupsAccumulator for SumIntGroupsAccumulatorAnsi { let v = int_array.value(i).to_i64().ok_or_else(|| { DataFusionError::Internal("Failed to convert value to i64".to_string()) })?; - sums[group_index] = - Some(sums[group_index].unwrap_or(0).add_checked(v).map_err(|_| { - DataFusionError::from(arithmetic_overflow_error("integer")) - })?); + sums[group_index] = Some( + sums[group_index] + .unwrap_or(0) + .add_checked(v) + .map_err(|_| DataFusionError::from(long_add_overflow_error()))?, + ); } } Ok(()) @@ -669,7 +671,7 @@ impl GroupsAccumulator for SumIntGroupsAccumulatorAnsi { self.sums[group_index] .unwrap() .add_checked(that_sum) - .map_err(|_| DataFusionError::from(arithmetic_overflow_error("integer")))?, + .map_err(|_| DataFusionError::from(long_add_overflow_error()))?, ); } } @@ -900,6 +902,7 @@ impl GroupsAccumulator for SumIntGroupsAccumulatorTry { #[cfg(test)] mod tests { use super::*; + use crate::SparkError; use arrow::array::Int64Array; use datafusion::logical_expr::{EmitTo, GroupsAccumulator}; @@ -1017,4 +1020,87 @@ mod tests { acc.merge_batch(&[states]).unwrap(); assert_eq!(acc.evaluate().unwrap(), ScalarValue::Int64(Some(60))); } + + /// Spark's integral `SUM` adds through `Add` on `LONG`, so every ANSI overflow path must + /// report a long overflow with the `try_add` suggestion. + fn assert_long_add_overflow(error: DataFusionError) { + let DataFusionError::External(error) = error else { + panic!("Expected structured Spark error, got {error:?}") + }; + match error.downcast_ref::() { + Some(SparkError::ArithmeticOverflow { + from_type, + function_name, + }) => { + assert_eq!(from_type, "long"); + assert_eq!(function_name, "try_add"); + } + other => panic!("Expected ArithmeticOverflow, got {other:?}"), + } + } + + fn int64_array(values: Vec) -> ArrayRef { + Arc::new(Int64Array::from(values)) + } + + #[test] + fn test_ansi_accumulator_update_batch_overflow() { + let mut acc = SumIntegerAccumulatorAnsi::new(); + let error = acc + .update_batch(&[int64_array(vec![i64::MAX, 1])]) + .unwrap_err(); + assert_long_add_overflow(error); + } + + #[test] + fn test_ansi_accumulator_merge_batch_overflow() { + let mut acc = SumIntegerAccumulatorAnsi::new(); + acc.merge_batch(&[int64_array(vec![i64::MAX])]).unwrap(); + let error = acc.merge_batch(&[int64_array(vec![1])]).unwrap_err(); + assert_long_add_overflow(error); + } + + #[test] + fn test_ansi_accumulator_update_batch_underflow() { + let mut acc = SumIntegerAccumulatorAnsi::new(); + let error = acc + .update_batch(&[int64_array(vec![i64::MIN, -1])]) + .unwrap_err(); + assert_long_add_overflow(error); + } + + #[test] + fn test_ansi_groups_accumulator_update_batch_overflow() { + let mut acc = SumIntGroupsAccumulatorAnsi::new(); + // Only group 1 overflows. + let error = acc + .update_batch( + &[int64_array(vec![1, i64::MAX, 2, 1])], + &[0, 1, 0, 1], + None, + 2, + ) + .unwrap_err(); + assert_long_add_overflow(error); + } + + #[test] + fn test_ansi_groups_accumulator_merge_batch_overflow() { + let mut acc = SumIntGroupsAccumulatorAnsi::new(); + acc.merge_batch(&[int64_array(vec![i64::MAX, 5])], &[0, 1], 2) + .unwrap(); + let error = acc + .merge_batch(&[int64_array(vec![1, 5])], &[0, 1], 2) + .unwrap_err(); + assert_long_add_overflow(error); + } + + #[test] + fn test_ansi_groups_accumulator_merge_batch_underflow() { + let mut acc = SumIntGroupsAccumulatorAnsi::new(); + let error = acc + .merge_batch(&[int64_array(vec![i64::MIN, -1])], &[0, 0], 1) + .unwrap_err(); + assert_long_add_overflow(error); + } } diff --git a/native/spark-expr/src/lib.rs b/native/spark-expr/src/lib.rs index 6d07ce13244..d7237c47fa3 100644 --- a/native/spark-expr/src/lib.rs +++ b/native/spark-expr/src/lib.rs @@ -140,6 +140,15 @@ pub(crate) fn arithmetic_overflow_error(from_type: &str) -> SparkError { } } +/// Spark adds integral `SUM` inputs as `LONG` through `Add`, so an overflow reports a long +/// overflow with the `try_add` suggestion whatever the input type. +pub(crate) fn long_add_overflow_error() -> SparkError { + SparkError::ArithmeticOverflow { + from_type: "long".to_string(), + function_name: "try_add".to_string(), + } +} + pub(crate) fn decimal_sum_overflow_error(function_name: &str) -> SparkError { SparkError::DecimalSumOverflow { function_name: function_name.to_string(), diff --git a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala index a8bd439c188..35dcd42a109 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala @@ -24,7 +24,7 @@ import java.util.concurrent.atomic.AtomicLong import scala.util.Random import org.apache.hadoop.fs.Path -import org.apache.spark.{CometListenerBusUtils, SparkConf} +import org.apache.spark.{CometListenerBusUtils, SparkConf, SparkThrowable} import org.apache.spark.scheduler.{SparkListener, SparkListenerTaskEnd} import org.apache.spark.sql.{Column, CometTestBase, DataFrame, QueryTest, Row} import org.apache.spark.sql.catalyst.expressions.Cast @@ -38,7 +38,7 @@ import org.apache.spark.sql.execution.SQLExecution import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, ShuffleQueryStageExec} import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, ObjectHashAggregateExec, SortAggregateExec} import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, ShuffleExchangeExec} -import org.apache.spark.sql.functions.{avg, col, collect_list, collect_set, count_distinct, expr, sort_array, sum} +import org.apache.spark.sql.functions.{avg, col, collect_list, collect_set, count_distinct, expr, lit, sort_array, sum, when} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{ArrayType, DataTypes, StructField, StructType} @@ -3297,36 +3297,49 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { (1 to 50).flatMap(_ => Seq((maxDec38_0, 1))) } + /** Spark's integral `SUM` overflow suggests `try_add`, so Comet's error must too. */ + private def assertTryAddSuggestion(error: SparkThrowable, clue: String): Unit = { + val alternative = error.getMessageParameters.get("alternative") + assert( + alternative != null && alternative.contains("'try_add'"), + s"$clue -> alternative=$alternative") + } + test("ANSI support - SUM function") { Seq(true, false).foreach { ansiEnabled => withSQLConf(SQLConf.ANSI_ENABLED.key -> ansiEnabled.toString) { - // Test long overflow - withParquetTable(Seq((Long.MaxValue, 1L), (100L, 1L)), "tbl") { - val res = sql("SELECT SUM(_1) FROM tbl") - if (ansiEnabled) { - checkSparkAnswerMaybeThrows(res) match { - case (Some(sparkExc), Some(cometExc)) => - // make sure that the error message throws overflow exception only - assert(sparkExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - assert(cometExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - case _ => fail("Exception should be thrown for Long overflow in ANSI mode") - } - } else { - checkSparkAnswerAndOperator(res) - } - } - // Test long underflow - withParquetTable(Seq((Long.MinValue, 1L), (-100L, 1L)), "tbl") { - val res = sql("SELECT SUM(_1) FROM tbl") - if (ansiEnabled) { - checkSparkAnswerMaybeThrows(res) match { - case (Some(sparkExc), Some(cometExc)) => - assert(sparkExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - assert(cometExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - case _ => fail("Exception should be thrown for Long underflow in ANSI mode") + // Spark's integral SUM adds through `Add` on LONG, so an overflow reports + // `long overflow` with the `try_add` suggestion. With one partition both rows meet in + // the partial aggregate. With two partitions each holds one row, so the overflow only + // happens when the final aggregate merges the partial sums. + for ((extreme, step) <- Seq((Long.MaxValue, 1L), (Long.MinValue, -1L)); + partitions <- Seq(1, 2); + grouped <- Seq(false, true)) { + withTempPath { dir => + val path = dir.getCanonicalPath + spark + .range(0, 2, 1, partitions) + .select(when(col("id") === 0, extreme).otherwise(step).as("l"), lit(1).as("g")) + .write + .parquet(path) + // An open cost as large as the split size keeps each file in its own partition. + val openCost = SQLConf.FILES_MAX_PARTITION_BYTES.defaultValue.get.toString + withSQLConf(SQLConf.FILES_OPEN_COST_IN_BYTES.key -> openCost) { + withTempView("tbl") { + val input = spark.read.parquet(path) + assert(input.rdd.getNumPartitions == partitions) + input.createOrReplaceTempView("tbl") + val query = + if (grouped) "SELECT g, SUM(l) FROM tbl GROUP BY g" + else "SELECT SUM(l) FROM tbl" + val res = sql(query) + if (ansiEnabled) { + assertTryAddSuggestion(checkSparkError(res, "ARITHMETIC_OVERFLOW"), query) + } else { + checkSparkAnswerAndOperator(res) + } + } } - } else { - checkSparkAnswerAndOperator(res) } } // Test Int SUM (should not overflow) @@ -3381,13 +3394,9 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { "tbl") { val res = sql("SELECT _2, SUM(_1) FROM tbl GROUP BY _2").repartition(2) if (ansiEnabled) { - checkSparkAnswerMaybeThrows(res) match { - case (Some(sparkExc), Some(cometExc)) => - assert(sparkExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - assert(cometExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - case _ => - fail("Exception should be thrown for Long overflow with GROUP BY in ANSI mode") - } + assertTryAddSuggestion( + checkSparkError(res, "ARITHMETIC_OVERFLOW"), + "SELECT _2, SUM(_1) FROM tbl GROUP BY _2") } else { checkSparkAnswerAndOperator(res) } @@ -3398,13 +3407,9 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper { "tbl") { val res = sql("SELECT _2, SUM(_1) FROM tbl GROUP BY _2") if (ansiEnabled) { - checkSparkAnswerMaybeThrows(res) match { - case (Some(sparkExc), Some(cometExc)) => - assert(sparkExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - assert(cometExc.getMessage.contains("ARITHMETIC_OVERFLOW")) - case _ => - fail("Exception should be thrown for Long underflow with GROUP BY in ANSI mode") - } + assertTryAddSuggestion( + checkSparkError(res, "ARITHMETIC_OVERFLOW"), + "SELECT _2, SUM(_1) FROM tbl GROUP BY _2") } else { checkSparkAnswerAndOperator(res) }