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))()