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]

Reply via email to