Skip to content
Closed
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
100 changes: 93 additions & 7 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::{long_add_overflow_error, EvalMode};

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 long_add_overflow_error() too, in whichever PR lands second? I have not built the merge.

use arrow::array::{
as_primitive_array, cast::AsArray, Array, ArrayRef, ArrowNativeTypeOp, ArrowPrimitiveType,
BooleanArray, Int64Array, PrimitiveArray,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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(())
Expand Down Expand Up @@ -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()))?,
);
}
}
Expand Down Expand Up @@ -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};

Expand Down Expand Up @@ -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::<SparkError>() {
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<i64>) -> 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() {

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.

These six tests reach three map_err sites. Underflow takes the same add_checked branch as overflow, and SumIntegerAccumulatorAnsi::merge_batch only calls update_batch. So test_ansi_accumulator_update_batch_underflow, test_ansi_accumulator_merge_batch_overflow and test_ansi_groups_accumulator_merge_batch_underflow each repeat a site another test here already pins, and the Scala test drives both signs through all three sites. Could this be one test per site?

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);
}
}
9 changes: 9 additions & 0 deletions native/spark-expr/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(),
Expand Down
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 All @@ -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}

Expand Down Expand Up @@ -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")
Comment on lines +3301 to +3305

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.

#6661 asks for this test to compare the message parameters with Spark's. assertTryAddSuggestion only checks that Comet's alternative mentions try_add, so the Scala test would still pass if message came back as integer overflow. Only the Rust tests pin from_type. checkSparkError in CometTestBase already holds Spark's error as expected, and CometExpressionSuite and CometTemporalExpressionSuite each hand-roll the getMessageParameters comparison today. Could checkSparkError take a flag to also assert actual.getMessageParameters == expected.getMessageParameters? That checks message against Spark on every profile, including 4.2's overflow, without hard-coding the wording.

}

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)
Expand Down Expand Up @@ -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)
}
Expand All @@ -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)
}
Expand Down
Loading