This is an automated email from the ASF dual-hosted git repository.
andygrove pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datafusion-comet.git
The following commit(s) were added to refs/heads/main by this push:
new 4989b5ece0 perf: avoid repeated decimal promotion in expression
serialization (#5736)
4989b5ece0 is described below
commit 4989b5ece0c440211f0726edb8f5c86662036f43
Author: Peter Lee <[email protected]>
AuthorDate: Thu Sep 10 04:14:21 2026 +0800
perf: avoid repeated decimal promotion in expression serialization (#5736)
* perf: avoid repeated decimal promotion in expression serialization
* perf: finish decimal promotion audit and verify overflow results
---
.../org/apache/comet/serde/QueryPlanSerde.scala | 15 ++-
.../main/scala/org/apache/comet/serde/arrays.scala | 77 +++++------
.../scala/org/apache/comet/serde/bitwise.scala | 6 +-
.../sql/comet/CometDecimalPromotionSuite.scala | 143 ++++++++++++++++++++-
4 files changed, 194 insertions(+), 47 deletions(-)
diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
index d6c9a49604..ff6ecc471f 100644
--- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
@@ -741,6 +741,10 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
}
}
+ /**
+ * This method does not promote the aggregate tree. Its arguments and
filters are independent
+ * roots and must be serialized through [[exprToProto]] to receive decimal
promotion.
+ */
def aggExprToProto(
aggExpr: AggregateExpression,
inputs: Seq[Attribute],
@@ -845,7 +849,9 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
* expression.
*
* This method performs a transformation on the plan to handle decimal
promotion and then calls
- * into the recursive method [[exprToProtoInternal]].
+ * into the recursive method [[exprToProtoInternal]]. Use this entry point
for independent roots
+ * (including aggregate arguments and filters) and synthesized trees needing
decimal promotion.
+ * Serdes must use [[exprToProtoInternal]] for children of the
already-promoted tree.
*
* @param expr
* The input expression
@@ -914,6 +920,11 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
* Convert a Spark expression to a protocol-buffer representation of a
native Comet/DataFusion
* expression.
*
+ * The caller owns decimal promotion: this method serializes children of an
already-promoted
+ * root without traversing them again. Literals and wrappers that introduce
no decimal
+ * arithmetic can also use this path. Newly synthesized arithmetic must
enter through
+ * [[exprToProto]].
+ *
* @param expr
* The input expression
* @param inputs
@@ -1063,7 +1074,7 @@ object QueryPlanSerde extends Logging with CometExprShim
with CometTypeShim {
binding: Boolean,
f: (ExprOuterClass.Expr.Builder, ExprOuterClass.UnaryExpr) =>
ExprOuterClass.Expr.Builder)
: Option[ExprOuterClass.Expr] = {
- val childExpr = exprToProtoInternal(child, inputs, binding) // TODO review
+ val childExpr = exprToProtoInternal(child, inputs, binding)
if (childExpr.isDefined) {
// create the generic UnaryExpr message
val inner = ExprOuterClass.UnaryExpr
diff --git a/spark/src/main/scala/org/apache/comet/serde/arrays.scala
b/spark/src/main/scala/org/apache/comet/serde/arrays.scala
index 32f8d08e10..cca9f63f8b 100644
--- a/spark/src/main/scala/org/apache/comet/serde/arrays.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/arrays.scala
@@ -44,8 +44,8 @@ object CometArrayRemove
expr: ArrayRemove,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.left, inputs, binding)
- val keyExprProto = exprToProto(expr.right, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
+ val keyExprProto = exprToProtoInternal(expr.right, inputs, binding)
scalarFunctionExprToProto("array_remove_all", arrayExprProto, keyExprProto)
}
@@ -60,8 +60,8 @@ object CometArrayAppend extends
CometExpressionSerde[ArrayAppend] with ArraysBas
val (srcChild, itemChild) = widenElementInLockstep(expr.children.head,
expr.children(1))
val elementType = srcChild.dataType.asInstanceOf[ArrayType].elementType
- val arrayExprProto = exprToProto(srcChild, inputs, binding)
- val keyExprProto = exprToProto(itemChild, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(srcChild, inputs, binding)
+ val keyExprProto = exprToProtoInternal(itemChild, inputs, binding)
// DataFusion's array_append always returns a list with nullable elements,
// so we must promise ArrayType(elementType, containsNull = true) here
even if
@@ -83,7 +83,8 @@ object CometArrayAppend extends
CometExpressionSerde[ArrayAppend] with ArraysBas
binding,
(builder, unaryExpr) => builder.setIsNotNull(unaryExpr))
- val nullLiteralProto = exprToProto(Literal(null, elementType), Seq.empty)
+ val nullLiteralProto =
+ exprToProtoInternal(Literal(null, elementType), Seq.empty, binding =
true)
if (arrayAppendScalarExpr.isDefined && isNotNullExpr.isDefined &&
nullLiteralProto.isDefined) {
val caseWhenExpr = ExprOuterClass.CaseWhen
@@ -130,8 +131,8 @@ object CometArrayContains
expr: ArrayContains,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.children.head, inputs, binding)
- val keyExprProto = exprToProto(expr.children(1), inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.children.head, inputs,
binding)
+ val keyExprProto = exprToProtoInternal(expr.children(1), inputs, binding)
scalarFunctionExprToProto("array_contains", arrayExprProto, keyExprProto)
}
@@ -230,8 +231,8 @@ object CometArrayIntersect
expr: ArrayIntersect,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val leftArrayExprProto = exprToProto(expr.children.head, inputs, binding)
- val rightArrayExprProto = exprToProto(expr.children(1), inputs, binding)
+ val leftArrayExprProto = exprToProtoInternal(expr.children.head, inputs,
binding)
+ val rightArrayExprProto = exprToProtoInternal(expr.children(1), inputs,
binding)
val arraysIntersectScalarExpr =
scalarFunctionExprToProto("array_intersect", leftArrayExprProto,
rightArrayExprProto)
@@ -244,7 +245,7 @@ object CometArrayMax extends CometExpressionSerde[ArrayMax]
{
expr: ArrayMax,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.children.head, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.children.head, inputs,
binding)
val arrayMaxScalarExpr =
scalarFunctionExprToProto("array_max", arrayExprProto)
@@ -257,7 +258,7 @@ object CometArrayMin extends CometExpressionSerde[ArrayMin]
{
expr: ArrayMin,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.children.head, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.children.head, inputs,
binding)
val arrayMinScalarExpr = scalarFunctionExprToProto("array_min",
arrayExprProto)
arrayMinScalarExpr
@@ -269,8 +270,8 @@ object CometArraysOverlap extends
CometExpressionSerde[ArraysOverlap] {
expr: ArraysOverlap,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val leftArrayExprProto = exprToProto(expr.left, inputs, binding)
- val rightArrayExprProto = exprToProto(expr.right, inputs, binding)
+ val leftArrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
+ val rightArrayExprProto = exprToProtoInternal(expr.right, inputs, binding)
val arraysOverlapScalarExpr = scalarFunctionExprToProtoWithReturnType(
"spark_arrays_overlap",
@@ -289,7 +290,7 @@ object CometArrayCompact extends
CometExpressionSerde[Expression] {
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
val child = expr.children.head
- val arrayExprProto = exprToProto(child, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(child, inputs, binding)
val arrayCompactScalarExpr = scalarFunctionExprToProto("array_compact",
arrayExprProto)
arrayCompactScalarExpr
@@ -345,8 +346,8 @@ object CometArrayExcept
return None
case None =>
}
- val leftArrayExprProto = exprToProto(expr.left, inputs, binding)
- val rightArrayExprProto = exprToProto(expr.right, inputs, binding)
+ val leftArrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
+ val rightArrayExprProto = exprToProtoInternal(expr.right, inputs, binding)
val arrayExceptScalarExpr =
scalarFunctionExprToProto("array_except", leftArrayExprProto,
rightArrayExprProto)
@@ -402,8 +403,8 @@ object CometArrayJoin
expr: ArrayJoin,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.array, inputs, binding)
- val delimiterExprProto = exprToProto(expr.delimiter, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.array, inputs, binding)
+ val delimiterExprProto = exprToProtoInternal(expr.delimiter, inputs,
binding)
val joined = expr.nullReplacement match {
case Some(nullReplacementExpr) =>
@@ -411,7 +412,7 @@ object CometArrayJoin
"array_to_string",
arrayExprProto,
delimiterExprProto,
- exprToProto(nullReplacementExpr, inputs, binding))
+ exprToProtoInternal(nullReplacementExpr, inputs, binding))
case None =>
scalarFunctionExprToProto("array_to_string", arrayExprProto,
delimiterExprProto)
}
@@ -422,8 +423,8 @@ object CometArrayJoin
case Some(nullReplacementExpr) =>
for {
innerProto <- joined
- replacementIsNull <- exprToProto(IsNull(nullReplacementExpr),
inputs, binding)
- nullLiteral <- exprToProto(Literal(null, expr.dataType), inputs,
binding)
+ replacementIsNull <-
exprToProtoInternal(IsNull(nullReplacementExpr), inputs, binding)
+ nullLiteral <- exprToProtoInternal(Literal(null, expr.dataType),
inputs, binding)
} yield ExprOuterClass.Expr
.newBuilder()
.setIf(
@@ -478,9 +479,9 @@ object CometSlice extends CometExpressionSerde[Slice] {
expr: Slice,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.x, inputs, binding)
- val startExprProto = exprToProto(Cast(expr.start, LongType), inputs,
binding)
- val lengthExprProto = exprToProto(Cast(expr.length, LongType), inputs,
binding)
+ val arrayExprProto = exprToProtoInternal(expr.x, inputs, binding)
+ val startExprProto = exprToProtoInternal(Cast(expr.start, LongType),
inputs, binding)
+ val lengthExprProto = exprToProtoInternal(Cast(expr.length, LongType),
inputs, binding)
// No serialized return type: native `spark_array_slice` reuses its
input's list field for the
// output, so only `return_field_from_args` is guaranteed to match.
Spark's `expr.dataType` is
// not: `CometCreateArray` may have widened the input to a deeply-nullable
element type, and
@@ -498,8 +499,8 @@ object CometArrayUnion extends
CometExpressionSerde[ArrayUnion] {
expr: ArrayUnion,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val leftArrayExprProto = exprToProto(expr.children.head, inputs, binding)
- val rightArrayExprProto = exprToProto(expr.children(1), inputs, binding)
+ val leftArrayExprProto = exprToProtoInternal(expr.children.head, inputs,
binding)
+ val rightArrayExprProto = exprToProtoInternal(expr.children(1), inputs,
binding)
val arraysUnionScalarExpr =
scalarFunctionExprToProto("array_union", leftArrayExprProto,
rightArrayExprProto)
@@ -606,7 +607,7 @@ object CometArrayReverse extends
CometExpressionSerde[Reverse] with ArraysBase {
withFallbackReason(expr, s"child data type not supported:
${expr.child.dataType}")
return None
}
- val reverseExprProto = exprToProto(expr.child, inputs, binding)
+ val reverseExprProto = exprToProtoInternal(expr.child, inputs, binding)
val reverseScalarExpr = scalarFunctionExprToProto("array_reverse",
reverseExprProto)
reverseScalarExpr
}
@@ -694,7 +695,8 @@ object CometElementAt extends
CometExpressionSerde[ElementAt] {
inputs,
binding,
(builder, unaryExpr) => builder.setIsNotNull(unaryExpr))
- val nullLiteralProto = exprToProto(Literal(null, expr.dataType),
Seq.empty)
+ val nullLiteralProto =
+ exprToProtoInternal(Literal(null, expr.dataType), Seq.empty, binding =
true)
for {
base <- baseExpr
notNull <- isNotNullExpr
@@ -732,7 +734,7 @@ object CometFlatten extends CometExpressionSerde[Flatten]
with ArraysBase {
expr: Flatten,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val flattenExprProto = exprToProto(expr.child, inputs, binding)
+ val flattenExprProto = exprToProtoInternal(expr.child, inputs, binding)
val flattenScalarExpr = scalarFunctionExprToProto("flatten",
flattenExprProto)
flattenScalarExpr
}
@@ -777,7 +779,7 @@ object CometSize extends CometExpressionSerde[Size] {
expr: Size,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.child, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.child, inputs, binding)
for {
isNotNullExprProto <- createIsNotNullExprProto(expr, inputs, binding)
sizeScalarExprProto <- scalarFunctionExprToProto("size", arrayExprProto)
@@ -810,7 +812,7 @@ object CometSize extends CometExpressionSerde[Size] {
private def createLiteralExprProto(legacySizeOfNull: Boolean):
Option[ExprOuterClass.Expr] = {
val value = if (legacySizeOfNull) -1 else null
- exprToProto(Literal(value, IntegerType), Seq.empty)
+ exprToProtoInternal(Literal(value, IntegerType), Seq.empty, binding = true)
}
}
@@ -830,8 +832,8 @@ object CometArrayPosition extends
CometExpressionSerde[ArrayPosition] with Array
expr: ArrayPosition,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val arrayExprProto = exprToProto(expr.left, inputs, binding)
- val elementExprProto = exprToProto(expr.right, inputs, binding)
+ val arrayExprProto = exprToProtoInternal(expr.left, inputs, binding)
+ val elementExprProto = exprToProtoInternal(expr.right, inputs, binding)
// Use spark_array_position which returns Int64 and 0 when not found
// (matching Spark's behavior)
@@ -881,7 +883,8 @@ object CometArraysZip extends
CometExpressionSerde[ArraysZip] {
// mimic Spark's ArraysZip behavior: returns NULL if any argument is NULL
val combinedNullCheck = expr.children.map(child =>
IsNotNull(child)).reduce(And)
val isNotNullExpr = exprToProtoInternal(combinedNullCheck, inputs, binding)
- val nullLiteralProto = exprToProto(Literal(null, expr.dataType), Seq.empty)
+ val nullLiteralProto =
+ exprToProtoInternal(Literal(null, expr.dataType), Seq.empty, binding =
true)
if (exprChildren.forall(
_.isDefined) && isNotNullExpr.isDefined && nullLiteralProto.isDefined)
{
@@ -1017,12 +1020,12 @@ object CometSequence extends
CometExpressionSerde[Sequence] with CodegenDispatch
expr: Sequence,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val startExprProto = exprToProto(expr.start, inputs, binding)
- val stopExprProto = exprToProto(expr.stop, inputs, binding)
+ val startExprProto = exprToProtoInternal(expr.start, inputs, binding)
+ val stopExprProto = exprToProtoInternal(expr.stop, inputs, binding)
// With no step argument the native kernel computes Spark's per-row
default,
// `start <= stop ? 1 : -1`, which cannot be expressed as a plan-time
literal.
val argProtos = Seq(startExprProto, stopExprProto) ++
- expr.stepOpt.map(exprToProto(_, inputs, binding))
+ expr.stepOpt.map(exprToProtoInternal(_, inputs, binding))
scalarFunctionExprToProtoWithReturnType(
"spark_sequence",
expr.dataType,
diff --git a/spark/src/main/scala/org/apache/comet/serde/bitwise.scala
b/spark/src/main/scala/org/apache/comet/serde/bitwise.scala
index 115bd80422..ec06c71253 100644
--- a/spark/src/main/scala/org/apache/comet/serde/bitwise.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/bitwise.scala
@@ -50,7 +50,7 @@ object CometBitwiseNot extends
CometExpressionSerde[BitwiseNot] {
expr: BitwiseNot,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val childProto = exprToProto(expr.child, inputs, binding)
+ val childProto = exprToProtoInternal(expr.child, inputs, binding)
val bitNotScalarExpr =
scalarFunctionExprToProto("bitwise_not", childProto)
bitNotScalarExpr
@@ -144,8 +144,8 @@ object CometBitwiseGet extends
CometExpressionSerde[BitwiseGet] {
expr: BitwiseGet,
inputs: Seq[Attribute],
binding: Boolean): Option[ExprOuterClass.Expr] = {
- val argProto = exprToProto(expr.left, inputs, binding)
- val posProto = exprToProto(expr.right, inputs, binding)
+ val argProto = exprToProtoInternal(expr.left, inputs, binding)
+ val posProto = exprToProtoInternal(expr.right, inputs, binding)
val bitGetScalarExpr =
scalarFunctionExprToProtoWithReturnType("bit_get", ByteType, false,
argProto, posProto)
bitGetScalarExpr
diff --git
a/spark/src/test/scala/org/apache/spark/sql/comet/CometDecimalPromotionSuite.scala
b/spark/src/test/scala/org/apache/spark/sql/comet/CometDecimalPromotionSuite.scala
index 346e7bef27..8181628242 100644
---
a/spark/src/test/scala/org/apache/spark/sql/comet/CometDecimalPromotionSuite.scala
+++
b/spark/src/test/scala/org/apache/spark/sql/comet/CometDecimalPromotionSuite.scala
@@ -20,10 +20,12 @@
package org.apache.spark.sql.comet
import org.apache.spark.sql.CometTestBase
-import org.apache.spark.sql.catalyst.expressions.{ArrayContains,
AttributeReference, Divide, EvalMode}
+import org.apache.spark.sql.catalyst.expressions.{Add, ArrayContains,
AttributeReference, BitwiseNot, Cast, CreateArray, Divide, EvalMode, Multiply,
NamedExpression}
+import
org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression,
Partial, Sum}
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types.{DecimalType, IntegerType}
+import org.apache.comet.CometConf
import org.apache.comet.serde.{CometDivide, ExprOuterClass, QueryPlanSerde,
Unsupported}
class CometDecimalPromotionSuite extends CometTestBase {
@@ -62,11 +64,10 @@ class CometDecimalPromotionSuite extends CometTestBase {
DecimalPrecision.promote(promoted) == promoted,
s"$name promotion is not idempotent: $promoted")
- // This proto-shape check relies on CometArrayContains re-entering
exprToProto for its
- // children. If https://github.com/apache/datafusion-comet/issues/5248
changes that,
- // re-point it to another recursively serializing serde.
+ // Deliberately re-enter the public serializer with an
already-promoted tree. Recursive
+ // serdes no longer do this, but re-promotion must still preserve the
protobuf shape.
val arithmeticProto = QueryPlanSerde
- .exprToProto(expression, plan.children.head.output)
+ .exprToProto(promoted, plan.children.head.output)
.get
.getScalarFunc
.getArgs(1)
@@ -95,6 +96,138 @@ class CometDecimalPromotionSuite extends CometTestBase {
s"try_divide($left, $right)")
}
+ test("issue #5248: nested decimal children and aggregate roots retain
overflow wrappers") {
+ val left = AttributeReference("left", DecimalType(10, 0))()
+ val right = AttributeReference("right", DecimalType(10, 0))()
+ val inputs = Seq(left, right)
+
+ Seq(EvalMode.LEGACY, EvalMode.ANSI, EvalMode.TRY).foreach { mode =>
+ val multiply = Multiply(left, right, mode)
+ val arithmetic = Add(multiply, right, mode)
+ val contains = ArrayContains(CreateArray(Seq(arithmetic)), arithmetic)
+
+ def check(proto: ExprOuterClass.Expr, binding: Boolean = true): Unit = {
+ assert(proto.hasCheckOverflow, s"$mode: $proto")
+ val outer = proto.getCheckOverflow
+ assert(outer.getDatatype ===
QueryPlanSerde.serializeDataType(arithmetic.dataType).get)
+ assert(outer.getFailOnError === (mode == EvalMode.ANSI))
+ assert(outer.getChild.hasAdd, s"Duplicate outer CheckOverflow: $proto")
+ val inner = outer.getChild.getAdd.getLeft
+ assert(inner.hasCheckOverflow, s"Missing nested CheckOverflow: $proto")
+ assert(
+ inner.getCheckOverflow.getDatatype ===
+ QueryPlanSerde.serializeDataType(multiply.dataType).get)
+ assert(inner.getCheckOverflow.getFailOnError === (mode ==
EvalMode.ANSI))
+ assert(
+ inner.getCheckOverflow.getChild.hasMultiply,
+ s"Duplicate inner CheckOverflow: $proto")
+ val reference = inner.getCheckOverflow.getChild.getMultiply.getLeft
+ if (binding) {
+ assert(reference.hasBound && reference.getBound.getIndex == 0)
+ } else {
+ assert(reference.hasUnbound && reference.getUnbound.getName ==
"left")
+ }
+ }
+
+ // Exercise both array child paths, including CreateArray's recursive
serialization.
+ Seq(true, false).foreach { binding =>
+ val proto = QueryPlanSerde.exprToProto(contains, inputs,
binding).get.getScalarFunc
+ check(proto.getArgs(0).getScalarFunc.getArgs(0), binding)
+ check(proto.getArgs(1), binding)
+ val bitwise = BitwiseNot(Cast(arithmetic, IntegerType))
+ check(
+ QueryPlanSerde
+ .exprToProto(bitwise, inputs, binding)
+ .get
+ .getScalarFunc
+ .getArgs(0)
+ .getCast
+ .getChild,
+ binding)
+ }
+
+ // Aggregate serialization does not promote the aggregate tree before
visiting its inputs.
+ val aggregate = AggregateExpression(
+ Sum(arithmetic),
+ Partial,
+ false,
+ Some(contains),
+ NamedExpression.newExprId)
+ val proto = QueryPlanSerde.aggExprToProto(aggregate, inputs, true,
SQLConf.get).get
+ check(proto.getSum.getChild)
+ check(proto.getFilter.getScalarFunc.getArgs(1))
+ }
+ }
+
+ Seq(false, true).foreach { ansi =>
+ test(s"issue #5248: decimal overflow values under recursive serdes,
ANSI=$ansi") {
+ withSQLConf(
+ SQLConf.ANSI_ENABLED.key -> ansi.toString,
+ CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false",
+ CometConf.getExprAllowIncompatConfigKey("ArrayIntersect") -> "true",
+ CometConf.getExprAllowIncompatConfigKey("ArrayExcept") -> "true",
+ CometConf.getExprAllowIncompatConfigKey("ArrayJoin") -> "true") {
+ withTempPath { path =>
+ // Read actual decimal columns from Parquet so constant folding
cannot hide promotion.
+ withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+ sql(s"""SELECT CAST(a AS DECIMAL(38,0)) a, CAST(b AS
DECIMAL(38,0)) b,
+ |CAST(c AS DECIMAL(38,6)) c, CAST(d AS DECIMAL(38,6)) d
+ |FROM VALUES
+ |('${"9" * 38}', '2', '${"9" * 32}.999999', '0.000001'),
+ |('4', '2', '4', '2'), (NULL, NULL, NULL, NULL) AS
t(a,b,c,d)
+ |""".stripMargin).write.parquet(path.toString)
+ }
+ withParquetTable(path.toString, "decimal_overflow") {
+ val expressions = Seq(
+ "array_remove(array($e), $e)",
+ "array_append(array($e), $e)",
+ "array_contains(array($e), $e)",
+ "array_intersect(array($e), array($e))",
+ "array_max(array($e))",
+ "array_min(array($e))",
+ "arrays_overlap(array($e), array($e))",
+ "array_compact(array($e))",
+ "array_except(array($e), array($e))",
+ "array_join(array(CAST($e AS STRING)), ',')",
+ // Spark's ArrayJoin codegen needs a nullable array or delimiter
to clear isNull
+ // when the replacement is nullable. Use a column-based
delimiter for this case.
+ "array_join(array('x', NULL), CAST(a AS STRING), CAST($e AS
STRING))",
+ "slice(array($e), 1, 1)",
+ "slice(array(1, 2), CAST($e AS INT), 2)",
+ "slice(array(1, 2), 1, CAST($e AS INT))",
+ "array_union(array($e), array($e))",
+ "reverse(array($e))",
+ "flatten(array(array($e)))",
+ "size(array($e))",
+ "array_position(array($e), $e)",
+ "~CAST($e AS BIGINT)",
+ "bit_get(CAST($e AS BIGINT), 1)",
+ "element_at(array($e), 1)",
+ "arrays_zip(array($e), array($e))")
+ // Decimal division's overflow sentinel needs CheckOverflow to
become NULL in LEGACY.
+ for (arithmetic <- Seq("a * b", "c / d"); expression <-
expressions) {
+ val query = s"SELECT ${expression.replace("$e", arithmetic)}
FROM decimal_overflow"
+ withClue(query) {
+ if (ansi) {
+ val df = sql(query)
+ assert(df.queryExecution.executedPlan.collect { case _:
CometProjectExec =>
+ true
+ }.nonEmpty)
+ val (sparkError, cometError) =
checkSparkAnswerMaybeThrows(df)
+ assert(sparkError.isDefined == cometError.isDefined)
+ } else {
+ checkSparkAnswerAndOperator(
+ sql(query),
+ includeClasses = Seq(classOf[CometNativeScanExec]))
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+
test("decimal Divide with a non-decimal operand is unsupported") {
// This is only a sanity check; Spark's type coercion should prevent this
case.
val decimal = AttributeReference("decimal", DecimalType(10, 0))()
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]