Repository navigation
fix: match Spark's ANSI integral SUM overflow error #6764
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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::<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() { | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. These six tests reach three |
||
| 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); | ||
| } | ||
| } | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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") | ||
|
Comment on lines
+3301
to
+3305
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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. |
||
| } | ||
|
|
||
| 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) | ||
| } | ||
|
|
||
There was a problem hiding this comment.
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::evaluatereturnsarithmetic_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 reportinginteger overflowwith notry_add, because Spark recomputes each sliding frame through the sameAdd. Could that site uselong_add_overflow_error()too, in whichever PR lands second? I have not built the merge.