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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 54 additions & 9 deletions native/spark-expr/src/agg_funcs/sum_int.rs
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.

use crate::{arithmetic_overflow_error, EvalMode};
use crate::{EvalMode, SparkError};

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

#6069 is open and approved, and it adds a fourth ANSI site in this file: SlidingSumIntegerAccumulator::evaluate returns arithmetic_overflow_error("integer"). Git merges the two PRs without a conflict, but this line drops that import, so from reading the merged tree I expect whichever lands second not to compile. That sliding frame would also keep reporting integer overflow with no try_add, because Spark recomputes each sliding frame through the same Add. Could that site use sum_overflow_error() too, in whichever PR lands second? I have not built the merge.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

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

Agreed. I can't change that site here, because it only exists in #6069. Whichever PR lands second will handle it:

use arrow::array::{
as_primitive_array, cast::AsArray, Array, ArrayRef, ArrowNativeTypeOp, ArrowPrimitiveType,
BooleanArray, Int64Array, PrimitiveArray,
Expand All @@ -31,6 +31,17 @@ use datafusion::logical_expr::{
};
use std::sync::Arc;

/// Spark's integral `SUM` always returns `LONG` and adds through `Add` in both its update and
/// merge expressions, so an ANSI overflow is reported by `MathUtils.addExact(Long, Long)` as a
/// `long` overflow carrying the `try_add` suggestion.
fn sum_overflow_error() -> DataFusionError {
SparkError::ArithmeticOverflow {
from_type: "long".to_string(),
function_name: "try_add".to_string(),
}
.into()
}

#[derive(Debug, PartialEq, Eq, Hash)]
pub struct SumInteger {
signature: Signature,
Expand Down Expand Up @@ -199,9 +210,7 @@ impl Accumulator for SumIntegerAccumulatorAnsi {
int_array.value(i)
))
})?;
sum = v
.add_checked(sum)
.map_err(|_| DataFusionError::from(arithmetic_overflow_error("integer")))?;
sum = v.add_checked(sum).map_err(|_| sum_overflow_error())?;
}
}
Ok(sum)
Expand Down Expand Up @@ -574,10 +583,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(|_| sum_overflow_error())?,
);
}
}
Ok(())
Expand Down Expand Up @@ -669,7 +680,7 @@ impl GroupsAccumulator for SumIntGroupsAccumulatorAnsi {
self.sums[group_index]
.unwrap()
.add_checked(that_sum)
.map_err(|_| DataFusionError::from(arithmetic_overflow_error("integer")))?,
.map_err(|_| sum_overflow_error())?,
);
}
}
Expand Down Expand Up @@ -1017,4 +1028,38 @@ mod tests {
acc.merge_batch(&[states]).unwrap();
assert_eq!(acc.evaluate().unwrap(), ScalarValue::Int64(Some(60)));
}

/// Spark reports an integral SUM overflow as `long overflow` with the `try_add` suggestion,
/// on the update and merge paths of both the scalar and the grouped ANSI accumulators.
#[test]
fn test_ansi_overflow_matches_spark_long_add() {
fn assert_spark_long_add_overflow(error: DataFusionError) {
let DataFusionError::External(error) = error else {
panic!("Expected structured Spark error, got {error:?}")
};
let error = error.downcast_ref::<SparkError>().unwrap();
let json: serde_json::Value = serde_json::from_str(&error.to_json()).unwrap();
assert_eq!(json["errorClass"], "ARITHMETIC_OVERFLOW");
assert_eq!(
json["params"],
serde_json::json!({"fromType": "long", "functionName": "try_add"})
);
}

for (first, second) in [(i64::MAX, 1), (i64::MIN, -1)] {
let batch = || -> ArrayRef { Arc::new(Int64Array::from(vec![first, second])) };

let mut acc = SumIntegerAccumulatorAnsi::new();
assert_spark_long_add_overflow(acc.update_batch(&[batch()]).unwrap_err());
let mut acc = SumIntegerAccumulatorAnsi::new();
assert_spark_long_add_overflow(acc.merge_batch(&[batch()]).unwrap_err());

let mut acc = SumIntGroupsAccumulatorAnsi::new();
assert_spark_long_add_overflow(
acc.update_batch(&[batch()], &[0, 0], None, 1).unwrap_err(),
);
let mut acc = SumIntGroupsAccumulatorAnsi::new();
assert_spark_long_add_overflow(acc.merge_batch(&[batch()], &[0, 0], 1).unwrap_err());
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -3297,20 +3297,40 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper {
(1 to 50).flatMap(_ => Seq((maxDec38_0, 1)))
}

/**
* Spark's integral SUM adds through `Add`, so an ANSI overflow carries the `try_add`
* suggestion. Compare the structured error, not just its error class. The overflow wording is
* left to the message parameters: Spark 4.2 normalizes `long overflow` to `overflow`.
*/
private def assertAnsiSumOverflowMatchesSpark(df: DataFrame): Unit = {
// Without this, a SUM that fell back to Spark would compare Spark's error with itself.
checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan))
val (sparkError, cometError) = checkSparkAnswerMaybeThrows(df)
Comment on lines +3305 to +3308

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

In ANSI mode nothing here checks that Comet ran the aggregate natively. From reading the code, if SUM fell back to Spark, this would compare Spark's error with Spark's and pass. checkSparkError in CometTestBase starts with checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)), and the ANSI integral overflow fidelity tests in CometExpressionSuite make the same check before an identical comparison. Could the helper start with that line? Another option is a flag on checkSparkError that also compares getMessageParameters, which would make this helper unnecessary.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

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

Good catch, thanks. The helper now starts with checkCometOperators(stripAQEPlan(df.queryExecution.executedPlan)), so a SUM that fell back to Spark fails the test instead of passing (61f490f).

def structured(error: Option[Throwable]): SparkThrowable with Throwable = {
val failure = error.getOrElse(fail("Expected SUM overflow in ANSI mode"))
causeChain(failure)
.collect { case e: SparkThrowable with Throwable => e }
.lastOption
.getOrElse(fail(s"Expected SparkThrowable: $failure"))
}
val expected = structured(sparkError)
val actual = structured(cometError)
assert(expected.getErrorClass == "ARITHMETIC_OVERFLOW")
assert(actual.getClass == expected.getClass)
assert(actual.getErrorClass == expected.getErrorClass)
assert(actual.getSqlState == expected.getSqlState)
assert(actual.getMessageParameters == expected.getMessageParameters)
assert(actual.getMessage.contains("try_add"))
}

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")
}
assertAnsiSumOverflowMatchesSpark(res)
} else {
checkSparkAnswerAndOperator(res)
}
Expand All @@ -3319,16 +3339,42 @@ class CometAggregateSuite extends CometTestBase with AdaptiveSparkPlanHelper {
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")
}
assertAnsiSumOverflowMatchesSpark(res)
} else {
checkSparkAnswerAndOperator(res)
}
}
// Overflow only when the partial sums of two scan partitions are merged, ungrouped and
// grouped. A large open cost keeps the two files in separate scan partitions.
withSQLConf(SQLConf.FILES_OPEN_COST_IN_BYTES.key -> (128L * 1024 * 1024).toString) {
withTempView("tbl") {
withTempPath { dir =>
Seq((Long.MaxValue, 1), (1L, 1)).foreach { row =>
spark
.createDataFrame(Seq(row))
.write
.mode("append")
.parquet(dir.getCanonicalPath)
}
spark.read.parquet(dir.getCanonicalPath).createOrReplaceTempView("tbl")
// One row per scan partition, so no partial update can overflow: the error
// below can only come from merging the two partial sums.
val rowsPerPartition =
spark.table("tbl").rdd.mapPartitions(rows => Iterator(rows.size)).collect()
assert(rowsPerPartition.toSeq == Seq(1, 1), rowsPerPartition.mkString(","))
for (query <- Seq(
"SELECT SUM(_1) FROM tbl",
"SELECT _2, SUM(_1) FROM tbl GROUP BY _2")) {
Comment on lines +3365 to +3367

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

The GROUP BY test further down (around line 3420) still only checks that both messages contain ARITHMETIC_OVERFLOW, so its grouped overflow and underflow cases pass with or without this fix. Could its two ANSI branches call assertAnsiSumOverflowMatchesSpark as well? That would also compare a grouped underflow with Spark, which this block does not cover because it only overflows upward.

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

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

Done in 61f490f. Both ANSI branches of the GROUP BY test now call assertAnsiSumOverflowMatchesSpark, so the grouped overflow and the grouped underflow are both compared with Spark's message parameters. All 12 CometAggregateSuite ANSI support tests pass locally on the default profile.

val res = sql(query)
if (ansiEnabled) {
assertAnsiSumOverflowMatchesSpark(res)
} else {
checkSparkAnswerAndOperator(res)
}
}
}
}
}
// Test Int SUM (should not overflow)
withParquetTable(Seq((Int.MaxValue, 1), (Int.MaxValue, 1), (100, 1)), "tbl") {
val res = sql("SELECT SUM(_1) FROM tbl")
Expand Down Expand Up @@ -3381,13 +3427,7 @@ 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")
}
assertAnsiSumOverflowMatchesSpark(res)
} else {
checkSparkAnswerAndOperator(res)
}
Expand All @@ -3398,13 +3438,7 @@ 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")
}
assertAnsiSumOverflowMatchesSpark(res)
} else {
checkSparkAnswerAndOperator(res)
}
Expand Down