From a500c41fb7e03121c2dfb2c85f6c2902a602b973 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 30 Sep 2026 06:50:59 -0600 Subject: [PATCH 1/7] fix: rescale decimal results from dispatched functions to the declared type Since #5692 the codegen dispatcher runs Invoke and StaticInvoke calls, including DSv2 catalog functions with a magic method. Such a function can return a Decimal whose scale differs from its declared result type, and the dispatcher wrote it without rescaling: off by a power of ten for precision up to 18, and a task failure above that. Spark's row writer rescales the value with changePrecision and writes null when it does not fit. The dispatcher's decimal writer now does the same, for top-level and nested outputs. Closes #6425. (cherry picked from commit ee26712befff7bac4d581e4fdf358258af217c9b) --- .../CometBatchKernelCodegenOutput.scala | 30 ++++++- .../org/apache/comet/serde/statics.scala | 9 +- .../comet/CometCodegenSourceSuite.scala | 7 ++ .../org/apache/comet/CometCodegenSuite.scala | 87 ++++++++++++++++++- 4 files changed, 124 insertions(+), 9 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala index 33e6c0c0355..2c1c013cb92 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -231,15 +231,37 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { val set = if (nested) "setSafe" else "set" OutputEmit("", s"$targetVec.$set($idx, $source);") case dt: DecimalType => + // Rescale to the declared type, and write null when the value does not fit, as Spark's + // `UnsafeRowWriter` and `UnsafeArrayWriter` do in `write(ordinal, Decimal, precision, + // scale)`. A Spark expression already produces its declared precision and scale, but a + // DSv2 function called through `Invoke` / `StaticInvoke` can return a `Decimal` of any + // scale (#6425). Like Spark's writers, this rescales the value in place, and leaves it + // untouched when it does not fit. + // + // The precision and scale test repeats `changePrecision`'s own fast path. It keeps the call + // off the common path, so the JIT can still scalar-replace the `Decimal` that an input + // getter allocates. With the bare call, passing a `DECIMAL(18, 2)` column through took + // about half as long again per row. + // // DecimalOutputShortFastPath: precision <= 18 fits in a signed long, so pass the unscaled // value to `setSafe(int, long)` and skip the BigDecimal allocation. + val dec = ctx.freshName("dec") + val (precision, scale) = (dt.precision, dt.scale) val write = - if (dt.precision <= Decimal.MAX_LONG_DIGITS) { - s"$targetVec.setSafe($idx, $source.toUnscaledLong());" + if (precision <= Decimal.MAX_LONG_DIGITS) { + s"$targetVec.setSafe($idx, $dec.toUnscaledLong());" } else { - s"$targetVec.setSafe($idx, $source.toJavaBigDecimal());" + s"$targetVec.setSafe($idx, $dec.toJavaBigDecimal());" } - OutputEmit("", write) + OutputEmit( + "", + s"""org.apache.spark.sql.types.Decimal $dec = $source; + |if (($dec.precision() == $precision && $dec.scale() == $scale) || + | $dec.changePrecision($precision, $scale)) { + | $write + |} else { + | $targetVec.setNull($idx); + |}""".stripMargin) case _: StringType => // Utf8OutputOnHeapShortcut: when the UTF8String is on-heap (Spark's string functions // allocate results on-heap), pass its backing byte[] directly to `setSafe`, skipping the diff --git a/spark/src/main/scala/org/apache/comet/serde/statics.scala b/spark/src/main/scala/org/apache/comet/serde/statics.scala index 89e2d6e56ab..94b845b7348 100644 --- a/spark/src/main/scala/org/apache/comet/serde/statics.scala +++ b/spark/src/main/scala/org/apache/comet/serde/statics.scala @@ -81,10 +81,11 @@ object CometStaticInvoke extends CometExpressionSerde[StaticInvoke] { * dispatcher, and at least one of those is not dispatchable. `CometIcebergTruncate` declines a * decimal because Iceberg's `truncate` can return a value wider than the column's declared * precision, which Spark nulls only when the row is materialized; the dispatcher writes into an - * Arrow `Decimal128(precision, scale)` vector just like a native kernel does, so it produces - * the out-of-range value instead of a null. The mixin's contract ("the case must be something - * `doGenCode` can compile") does not cover a limit that lives at the Arrow output boundary, so - * enrollment stays with the individual handlers. + * Arrow `Decimal128(precision, scale)` vector just like a native kernel does, so it has to null + * the value at its own output, and an enclosing predicate or hash then sees a null where Spark + * sees the oversized value. The mixin's contract ("the case must be something `doGenCode` can + * compile") does not cover a limit that lives at the Arrow output boundary, so enrollment stays + * with the individual handlers. */ override def getSupportLevel(expr: StaticInvoke): SupportLevel = handlerFor(expr) diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala index 400ec6fcd2f..2830bca6265 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSourceSuite.scala @@ -401,6 +401,10 @@ class CometCodegenSourceSuite extends AnyFunSuite { !result.body.contains(".toJavaBigDecimal("), "expected no BigDecimal allocation for short-precision output; got:\n" + CodeFormatter.format(result.code)) + // Rescaled to the declared type first, as Spark's row writer does (#6425). + assert( + result.body.contains(".changePrecision(18, 2)"), + s"expected changePrecision before the write; got:\n${CodeFormatter.format(result.code)}") } test("DecimalVector setSafe uses BigDecimal slow path for long-precision output") { @@ -417,6 +421,9 @@ class CometCodegenSourceSuite extends AnyFunSuite { !result.body.contains(".toUnscaledLong()"), "expected no unscaled-long write for long-precision output; got:\n" + CodeFormatter.format(result.code)) + assert( + result.body.contains(".changePrecision(38, 10)"), + s"expected changePrecision before the write; got:\n${CodeFormatter.format(result.code)}") } test("VarCharVector setSafe uses on-heap UTF8String shortcut") { diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 8b691fedcec..d6c4eb4ebd6 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -25,14 +25,19 @@ import org.scalatest.exceptions.TestFailedException import org.apache.arrow.vector._ import org.apache.spark.{SparkConf, SparkEnv, TaskContext} -import org.apache.spark.sql.CometTestBase +import org.apache.spark.sql.{CometTestBase, Row} import org.apache.spark.sql.api.java.UDF1 +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.analysis.NoSuchFunctionException import org.apache.spark.sql.catalyst.expressions.{Add, Alias, AttributeReference, BoundReference, Cast, CreateArray, CreateMap, CreateNamedStruct, Expression, Hypot, Literal, MapConcat} import org.apache.spark.sql.catalyst.expressions.objects.Invoke import org.apache.spark.sql.comet.CometProjectExec +import org.apache.spark.sql.connector.catalog.{FunctionCatalog, Identifier} +import org.apache.spark.sql.connector.catalog.functions.{BoundFunction, ScalarFunction, UnboundFunction} import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.CaseInsensitiveStringMap import org.apache.spark.unsafe.types.UTF8String import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus @@ -2327,6 +2332,48 @@ class CometCodegenSuite Invoke(target, "twice", StringType, Seq(Literal(UTF8String.fromString("ab"), StringType))) assert(runKernel(folded, 1)(_.getUTF8String(0).toString) === "abab") } + + test("decimal results of a DSv2 function are rescaled to the declared type (#6425)") { + // Spark lowers a call to a DSv2 function with an instance `invoke` method to `Invoke`, which + // the dispatcher runs. Both functions return `Decimal(i)` at scale 0, one declaring + // `DECIMAL(10, 2)` and one `DECIMAL(20, 12)`, which covers both of the dispatcher's decimal + // writers. Spark's row writer rescales the value with `changePrecision` and writes null when + // it does not fit: 100000000 has nine integer digits and both types allow eight. Spark adds + // no overflow check around the call, so that null does not depend on ANSI mode. `map` is + // itself dispatched, so its value exercises the nested writer. + def dec(s: String) = new java.math.BigDecimal(s) + val expected = Seq( + Row(3, dec("3.00"), dec("3.000000000000"), Map("k" -> dec("3.00"))), + Row(-7, dec("-7.00"), dec("-7.000000000000"), Map("k" -> dec("-7.00"))), + Row(null, null, null, Map("k" -> null)), + Row( + 99999999, + dec("99999999.00"), + dec("99999999.000000000000"), + Map("k" -> dec("99999999.00"))), + Row(100000000, null, null, Map("k" -> null))) + withSQLConf( + "spark.sql.catalog.decfn" -> classOf[CometCodegenSuite.DecimalFunctionCatalog].getName) { + withTable("t") { + sql("CREATE TABLE t (i INT) USING parquet") + // One file, so the kernel sees every row in one batch. + sql( + "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM " + + "VALUES (3), (-7), (NULL), (99999999), (100000000) AS v(i)") + for (ansi <- Seq("true", "false")) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { + val df = sql( + "SELECT i, decfn.ns.as_money(i), decfn.ns.as_wide_money(i), " + + "map('k', decfn.ns.as_money(i)) FROM t") + assertCodegenRan { + checkSparkAnswerAndOperator(df) + } + checkAnswer(df, expected) + } + } + } + } + } } /** @@ -2346,6 +2393,44 @@ object CometCodegenSuite { class NotSerializableTarget { def twice(s: UTF8String): UTF8String = UTF8String.fromString(s.toString + s.toString) } + + /** + * DSv2 function catalog for the #6425 test. `as_money` declares `DECIMAL(10, 2)` and + * `as_wide_money` declares `DECIMAL(20, 12)`. + */ + class DecimalFunctionCatalog extends FunctionCatalog { + private val functions = Map( + "as_money" -> new ScaleZeroDecimalFunction(10, 2), + "as_wide_money" -> new ScaleZeroDecimalFunction(20, 12)) + private var catalogName: String = _ + + override def initialize(name: String, options: CaseInsensitiveStringMap): Unit = + catalogName = name + + override def name(): String = catalogName + + override def listFunctions(namespace: Array[String]): Array[Identifier] = + functions.keys.map(Identifier.of(namespace, _)).toArray + + override def loadFunction(ident: Identifier): UnboundFunction = + functions.getOrElse(ident.name(), throw new NoSuchFunctionException(ident)) + } + + /** + * Returns its `INT` argument as `Decimal(v)`, at scale 0, whatever scale it declares. `invoke` + * is an instance method, so Spark lowers a call to `Invoke`. The function binds to itself. + */ + class ScaleZeroDecimalFunction(precision: Int, scale: Int) + extends UnboundFunction + with ScalarFunction[Decimal] { + override def name(): String = "scale_zero_decimal" + override def description(): String = s"int -> decimal($precision, $scale), at scale 0" + override def bind(inputType: StructType): BoundFunction = this + override def inputTypes(): Array[DataType] = Array(IntegerType) + override def resultType(): DataType = DecimalType(precision, scale) + def invoke(v: Int): Decimal = Decimal(v) + override def produceResult(input: InternalRow): Decimal = invoke(input.getInt(0)) + } } /** From 27bc8d9408adcc751e0426ffa31c3438bd9787db Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 30 Sep 2026 14:00:54 -0600 Subject: [PATCH 2/7] test: cover rounding and the negative overflow boundary in the dispatcher decimal rescale Add a function that returns thousandths into DECIMAL(10, 2), so the rescale has to round half up, and rows at -99999999 and -100000000 for the negative side of the overflow boundary. (cherry picked from commit 4cf7c1c6a85f8e2c7ad42bacb67907651e4e5d5b) --- .../org/apache/comet/CometCodegenSuite.scala | 75 +++++++++++++------ 1 file changed, 52 insertions(+), 23 deletions(-) diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index d6c4eb4ebd6..eb55516f535 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -2335,36 +2335,60 @@ class CometCodegenSuite test("decimal results of a DSv2 function are rescaled to the declared type (#6425)") { // Spark lowers a call to a DSv2 function with an instance `invoke` method to `Invoke`, which - // the dispatcher runs. Both functions return `Decimal(i)` at scale 0, one declaring - // `DECIMAL(10, 2)` and one `DECIMAL(20, 12)`, which covers both of the dispatcher's decimal - // writers. Spark's row writer rescales the value with `changePrecision` and writes null when - // it does not fit: 100000000 has nine integer digits and both types allow eight. Spark adds - // no overflow check around the call, so that null does not depend on ANSI mode. `map` is - // itself dispatched, so its value exercises the nested writer. + // the dispatcher runs. `as_money` and `as_wide_money` return `Decimal(i)` at scale 0, one + // declaring `DECIMAL(10, 2)` and one `DECIMAL(20, 12)`, which covers both of the dispatcher's + // decimal writers. Spark's row writer rescales the value with `changePrecision` and writes + // null when it does not fit: 100000000 and -100000000 have nine integer digits and both types + // allow eight. Spark adds no overflow check around the call, so that null does not depend on + // ANSI mode. `map` is itself dispatched, so its value exercises the nested writer. + // `mills_as_money` returns `i` thousandths, at scale 3, into `DECIMAL(10, 2)`, so the rescale + // drops a digit and `changePrecision` rounds half up: -1.005 becomes -1.01 and 1.004 becomes + // 1.00. def dec(s: String) = new java.math.BigDecimal(s) val expected = Seq( - Row(3, dec("3.00"), dec("3.000000000000"), Map("k" -> dec("3.00"))), - Row(-7, dec("-7.00"), dec("-7.000000000000"), Map("k" -> dec("-7.00"))), - Row(null, null, null, Map("k" -> null)), + Row(3, dec("3.00"), dec("3.000000000000"), Map("k" -> dec("3.00")), dec("0.00")), + Row(-7, dec("-7.00"), dec("-7.000000000000"), Map("k" -> dec("-7.00")), dec("-0.01")), + Row(null, null, null, Map("k" -> null), null), Row( 99999999, dec("99999999.00"), dec("99999999.000000000000"), - Map("k" -> dec("99999999.00"))), - Row(100000000, null, null, Map("k" -> null))) + Map("k" -> dec("99999999.00")), + dec("100000.00")), + Row(100000000, null, null, Map("k" -> null), dec("100000.00")), + Row( + -99999999, + dec("-99999999.00"), + dec("-99999999.000000000000"), + Map("k" -> dec("-99999999.00")), + dec("-100000.00")), + Row(-100000000, null, null, Map("k" -> null), dec("-100000.00")), + Row( + -1005, + dec("-1005.00"), + dec("-1005.000000000000"), + Map("k" -> dec("-1005.00")), + dec("-1.01")), + Row( + 1004, + dec("1004.00"), + dec("1004.000000000000"), + Map("k" -> dec("1004.00")), + dec("1.00")), + Row(5, dec("5.00"), dec("5.000000000000"), Map("k" -> dec("5.00")), dec("0.01"))) withSQLConf( "spark.sql.catalog.decfn" -> classOf[CometCodegenSuite.DecimalFunctionCatalog].getName) { withTable("t") { sql("CREATE TABLE t (i INT) USING parquet") // One file, so the kernel sees every row in one batch. sql( - "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM " + - "VALUES (3), (-7), (NULL), (99999999), (100000000) AS v(i)") + "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM VALUES (3), (-7), (NULL), " + + "(99999999), (100000000), (-99999999), (-100000000), (-1005), (1004), (5) AS v(i)") for (ansi <- Seq("true", "false")) { withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { val df = sql( "SELECT i, decfn.ns.as_money(i), decfn.ns.as_wide_money(i), " + - "map('k', decfn.ns.as_money(i)) FROM t") + "map('k', decfn.ns.as_money(i)), decfn.ns.mills_as_money(i) FROM t") assertCodegenRan { checkSparkAnswerAndOperator(df) } @@ -2396,12 +2420,14 @@ object CometCodegenSuite { /** * DSv2 function catalog for the #6425 test. `as_money` declares `DECIMAL(10, 2)` and - * `as_wide_money` declares `DECIMAL(20, 12)`. + * `as_wide_money` declares `DECIMAL(20, 12)`, and both return their argument at scale 0. + * `mills_as_money` declares `DECIMAL(10, 2)` and returns its argument as thousandths. */ class DecimalFunctionCatalog extends FunctionCatalog { private val functions = Map( - "as_money" -> new ScaleZeroDecimalFunction(10, 2), - "as_wide_money" -> new ScaleZeroDecimalFunction(20, 12)) + "as_money" -> new IntAsDecimalFunction(10, 2, valueScale = 0), + "as_wide_money" -> new IntAsDecimalFunction(20, 12, valueScale = 0), + "mills_as_money" -> new IntAsDecimalFunction(10, 2, valueScale = 3)) private var catalogName: String = _ override def initialize(name: String, options: CaseInsensitiveStringMap): Unit = @@ -2417,18 +2443,21 @@ object CometCodegenSuite { } /** - * Returns its `INT` argument as `Decimal(v)`, at scale 0, whatever scale it declares. `invoke` - * is an instance method, so Spark lowers a call to `Invoke`. The function binds to itself. + * Returns its `INT` argument as the unscaled value of a `Decimal` at `valueScale`, whatever + * scale it declares. `invoke` is an instance method, so Spark lowers a call to `Invoke`. The + * function binds to itself. */ - class ScaleZeroDecimalFunction(precision: Int, scale: Int) + class IntAsDecimalFunction(precision: Int, scale: Int, valueScale: Int) extends UnboundFunction with ScalarFunction[Decimal] { - override def name(): String = "scale_zero_decimal" - override def description(): String = s"int -> decimal($precision, $scale), at scale 0" + override def name(): String = "int_as_decimal" + override def description(): String = + s"int -> decimal($precision, $scale), at scale $valueScale" override def bind(inputType: StructType): BoundFunction = this override def inputTypes(): Array[DataType] = Array(IntegerType) override def resultType(): DataType = DecimalType(precision, scale) - def invoke(v: Int): Decimal = Decimal(v) + // Ten digits hold any `INT`. + def invoke(v: Int): Decimal = Decimal(v.toLong, 10, valueScale) override def produceResult(input: InternalRow): Decimal = invoke(input.getInt(0)) } } From a700d6d354ff38181107c36ad5ecdd318ee8aee8 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 30 Sep 2026 15:55:59 -0600 Subject: [PATCH 3/7] test: cover StaticInvoke, nested writers and non-nullable overflow in the decimal rescale Add a StaticInvoke variant of as_money, DSv2 functions that return arrays and structs of decimals with nullable and non-nullable children, and a test that a non-nullable result that does not fit fails in both Spark and Comet. Declare mills_as_money as DECIMAL(7, 2) so that 99999.999 rounds up past the precision. Note in the writer why the decimal branch has no null test. (cherry picked from commit 04492d87bc4c38f829a8037cf9e3ee60b395bc8d) --- .../CometBatchKernelCodegenOutput.scala | 4 +- .../org/apache/comet/CometCodegenSuite.scala | 223 ++++++++++++------ 2 files changed, 154 insertions(+), 73 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala index 2c1c013cb92..89713827aa8 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -236,7 +236,9 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { // scale)`. A Spark expression already produces its declared precision and scale, but a // DSv2 function called through `Invoke` / `StaticInvoke` can return a `Decimal` of any // scale (#6425). Like Spark's writers, this rescales the value in place, and leaves it - // untouched when it does not fit. + // untouched when it does not fit. Unlike them, it does not test `source` for null: the + // callers write null values themselves, and skip that test only for a type that is not + // nullable. // // The precision and scale test repeats `changePrecision`'s own fast path. It keeps the call // off the common path, so the JIT can still scalar-replace the `Decimal` that an input diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index eb55516f535..9ab400c67a0 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -29,8 +29,9 @@ import org.apache.spark.sql.{CometTestBase, Row} import org.apache.spark.sql.api.java.UDF1 import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.analysis.NoSuchFunctionException -import org.apache.spark.sql.catalyst.expressions.{Add, Alias, AttributeReference, BoundReference, Cast, CreateArray, CreateMap, CreateNamedStruct, Expression, Hypot, Literal, MapConcat} +import org.apache.spark.sql.catalyst.expressions.{Add, Alias, AttributeReference, BoundReference, Cast, CreateArray, CreateMap, CreateNamedStruct, Expression, GenericInternalRow, Hypot, Literal, MapConcat} import org.apache.spark.sql.catalyst.expressions.objects.Invoke +import org.apache.spark.sql.catalyst.util.GenericArrayData import org.apache.spark.sql.comet.CometProjectExec import org.apache.spark.sql.connector.catalog.{FunctionCatalog, Identifier} import org.apache.spark.sql.connector.catalog.functions.{BoundFunction, ScalarFunction, UnboundFunction} @@ -2333,71 +2334,112 @@ class CometCodegenSuite assert(runKernel(folded, 1)(_.getUTF8String(0).toString) === "abab") } - test("decimal results of a DSv2 function are rescaled to the declared type (#6425)") { - // Spark lowers a call to a DSv2 function with an instance `invoke` method to `Invoke`, which - // the dispatcher runs. `as_money` and `as_wide_money` return `Decimal(i)` at scale 0, one - // declaring `DECIMAL(10, 2)` and one `DECIMAL(20, 12)`, which covers both of the dispatcher's - // decimal writers. Spark's row writer rescales the value with `changePrecision` and writes - // null when it does not fit: 100000000 and -100000000 have nine integer digits and both types - // allow eight. Spark adds no overflow check around the call, so that null does not depend on - // ANSI mode. `map` is itself dispatched, so its value exercises the nested writer. - // `mills_as_money` returns `i` thousandths, at scale 3, into `DECIMAL(10, 2)`, so the rescale - // drops a digit and `changePrecision` rounds half up: -1.005 becomes -1.01 and 1.004 becomes - // 1.00. - def dec(s: String) = new java.math.BigDecimal(s) - val expected = Seq( - Row(3, dec("3.00"), dec("3.000000000000"), Map("k" -> dec("3.00")), dec("0.00")), - Row(-7, dec("-7.00"), dec("-7.000000000000"), Map("k" -> dec("-7.00")), dec("-0.01")), - Row(null, null, null, Map("k" -> null), null), - Row( - 99999999, - dec("99999999.00"), - dec("99999999.000000000000"), - Map("k" -> dec("99999999.00")), - dec("100000.00")), - Row(100000000, null, null, Map("k" -> null), dec("100000.00")), - Row( - -99999999, - dec("-99999999.00"), - dec("-99999999.000000000000"), - Map("k" -> dec("-99999999.00")), - dec("-100000.00")), - Row(-100000000, null, null, Map("k" -> null), dec("-100000.00")), - Row( - -1005, - dec("-1005.00"), - dec("-1005.000000000000"), - Map("k" -> dec("-1005.00")), - dec("-1.01")), - Row( - 1004, - dec("1004.00"), - dec("1004.000000000000"), - Map("k" -> dec("1004.00")), - dec("1.00")), - Row(5, dec("5.00"), dec("5.000000000000"), Map("k" -> dec("5.00")), dec("0.01"))) + /** + * Runs `f` with [[CometCodegenSuite.DecimalFunctionCatalog]] registered as `decfn` and `values` + * in `t (i INT)`, for the #6425 tests. + */ + private def withDecimalFunctions(values: Any*)(f: => Unit): Unit = { withSQLConf( "spark.sql.catalog.decfn" -> classOf[CometCodegenSuite.DecimalFunctionCatalog].getName) { withTable("t") { sql("CREATE TABLE t (i INT) USING parquet") // One file, so the kernel sees every row in one batch. sql( - "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM VALUES (3), (-7), (NULL), " + - "(99999999), (100000000), (-99999999), (-100000000), (-1005), (1004), (5) AS v(i)") - for (ansi <- Seq("true", "false")) { - withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { - val df = sql( - "SELECT i, decfn.ns.as_money(i), decfn.ns.as_wide_money(i), " + - "map('k', decfn.ns.as_money(i)), decfn.ns.mills_as_money(i) FROM t") - assertCodegenRan { - checkSparkAnswerAndOperator(df) - } - checkAnswer(df, expected) + "INSERT INTO t SELECT /*+ REPARTITION(1) */ * FROM VALUES " + + values.map(v => s"($v)").mkString(", ") + " AS v(i)") + f + } + } + } + + private def dec(s: String) = if (s == null) null else new java.math.BigDecimal(s) + + test("decimal results of a DSv2 function are rescaled to the declared type (#6425)") { + // Spark lowers a call to a DSv2 function with an instance `invoke` method to `Invoke`, and one + // with a static `invoke` to `StaticInvoke`. The dispatcher runs both. `as_money` and + // `as_wide_money` return `Decimal(i)` at scale 0, one declaring `DECIMAL(10, 2)` and one + // `DECIMAL(20, 12)`, which covers both of the dispatcher's decimal writers. `static_as_money` + // is `as_money` with a static `invoke`. Spark's row writer rescales the value with + // `changePrecision` and writes null when it does not fit: 100000000 and -100000000 have nine + // integer digits and both types allow eight. Spark adds no overflow check around the call, so + // that null does not depend on ANSI mode. `map` is itself dispatched, so its value goes + // through the kernel's map writer. `mills_as_money` returns `i` thousandths, at scale 3, into + // `DECIMAL(7, 2)`, so the rescale drops a digit and `changePrecision` rounds half up: -1.005 + // becomes -1.01 and 1.004 becomes 1.00. 99999.999 has the five integer digits the type allows, + // but rounds up to 100000.00, which has six, so it is null. + // + // Each case is `(i, as_money, as_wide_money, mills_as_money)`. `static_as_money` and the `map` + // value match `as_money`. + val cases = Seq[(Any, String, String, String)]( + (3, "3.00", "3.000000000000", "0.00"), + (-7, "-7.00", "-7.000000000000", "-0.01"), + (null, null, null, null), + (99999999, "99999999.00", "99999999.000000000000", null), + (100000000, null, null, null), + (-99999999, "-99999999.00", "-99999999.000000000000", null), + (-100000000, null, null, null), + (-1005, "-1005.00", "-1005.000000000000", "-1.01"), + (1004, "1004.00", "1004.000000000000", "1.00"), + (5, "5.00", "5.000000000000", "0.01")) + val expected = cases.map { case (i, money, wide, mills) => + Row(i, dec(money), dec(money), dec(wide), Map("k" -> dec(money)), dec(mills)) + } + withDecimalFunctions(cases.map(_._1): _*) { + for (ansi <- Seq("true", "false")) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { + val df = sql( + "SELECT i, decfn.ns.as_money(i), decfn.ns.static_as_money(i), " + + "decfn.ns.as_wide_money(i), map('k', decfn.ns.as_money(i)), " + + "decfn.ns.mills_as_money(i) FROM t") + assertCodegenRan { + checkSparkAnswerAndImpl(df, dispatched = Seq("invoke", "staticinvoke")) } + checkAnswer(df, expected) } } } } + + test("decimals in a DSv2 function's array and struct results are rescaled (#6425)") { + // Each function returns `Decimal(i)`, at scale 0, in every decimal of its result, so the + // kernel's array and struct writers have to rescale them, as Spark's `UnsafeArrayWriter` and + // `UnsafeRowWriter` do. `money_array`'s element and `money_struct`'s `m` field are + // `DECIMAL(10, 2)`, so both are null at 100000000. The writers skip the null check for a + // non-nullable child, so `non_null_money_array`'s element and `money_struct`'s `non_null_m` + // field cover that path. They are `DECIMAL(12, 2)`, which holds any `INT`. + // + // `array(...)` or `named_struct(...)` around a scalar call would not reach these writers: + // Comet evaluates both natively, and dispatches only the call. + withDecimalFunctions(3, null, 100000000) { + val df = sql( + "SELECT i, decfn.ns.money_array(i), decfn.ns.non_null_money_array(i), " + + "decfn.ns.money_struct(i) FROM t") + assertCodegenRan { + checkSparkAnswerAndOperator(df) + } + checkAnswer( + df, + Seq( + Row(3, Seq(dec("3.00")), Seq(dec("3.00")), Row(dec("3.00"), dec("3.00"))), + Row(null, null, null, null), + Row(100000000, Seq(null), Seq(dec("100000000.00")), Row(null, dec("100000000.00"))))) + } + } + + test("a non-nullable DSv2 decimal result that does not fit fails, as in Spark (#6425)") { + // `non_null_money` declares a non-nullable `DECIMAL(10, 2)`, and `coalesce` makes its argument + // non-nullable too, so neither Spark nor the kernel checks the result for null. 100000000 does + // not fit, and both write null for it anyway. Spark then fails to decode the row, and Comet's + // native projection rejects the batch. The errors differ, but both engines fail. + withDecimalFunctions(3, 100000000) { + val (sparkError, cometError) = + checkSparkAnswerMaybeThrows(sql("SELECT decfn.ns.non_null_money(coalesce(i, 0)) FROM t")) + assert(sparkError.isDefined, "expected Spark to fail") + assert( + cometError.exists(_.getMessage.contains("is declared as non-nullable but contains null")), + s"expected Comet's native projection to reject the null, got $cometError") + } + } } /** @@ -2419,15 +2461,25 @@ object CometCodegenSuite { } /** - * DSv2 function catalog for the #6425 test. `as_money` declares `DECIMAL(10, 2)` and - * `as_wide_money` declares `DECIMAL(20, 12)`, and both return their argument at scale 0. - * `mills_as_money` declares `DECIMAL(10, 2)` and returns its argument as thousandths. + * DSv2 function catalog for the #6425 tests. Each function returns its argument as a `Decimal` + * whose scale need not match the type it declares. `mills_as_money` returns its argument as + * thousandths, and the rest at scale 0. */ class DecimalFunctionCatalog extends FunctionCatalog { - private val functions = Map( - "as_money" -> new IntAsDecimalFunction(10, 2, valueScale = 0), - "as_wide_money" -> new IntAsDecimalFunction(20, 12, valueScale = 0), - "mills_as_money" -> new IntAsDecimalFunction(10, 2, valueScale = 3)) + private val money = DecimalType(10, 2) + // Holds any `INT`. + private val intMoney = DecimalType(12, 2) + private val functions: Map[String, UnboundFunction] = Map( + "as_money" -> new IntAsDecimalFunction(money), + "static_as_money" -> new StaticAsMoneyFunction, + "as_wide_money" -> new IntAsDecimalFunction(DecimalType(20, 12)), + "mills_as_money" -> new IntAsDecimalFunction(DecimalType(7, 2), valueScale = 3), + "non_null_money" -> new IntAsDecimalFunction(money, nullable = false), + "money_array" -> new IntAsDecimalFunction(ArrayType(money, containsNull = true)), + "non_null_money_array" -> + new IntAsDecimalFunction(ArrayType(intMoney, containsNull = false)), + "money_struct" -> new IntAsDecimalFunction( + new StructType().add("m", money).add("non_null_m", intMoney, nullable = false))) private var catalogName: String = _ override def initialize(name: String, options: CaseInsensitiveStringMap): Unit = @@ -2444,24 +2496,51 @@ object CometCodegenSuite { /** * Returns its `INT` argument as the unscaled value of a `Decimal` at `valueScale`, whatever - * scale it declares. `invoke` is an instance method, so Spark lowers a call to `Invoke`. The - * function binds to itself. + * scale `declared` has. For an array or struct type, every decimal in the result holds that + * value: the array has one element, and each field of the struct has it. `invoke` is an + * instance method, so Spark lowers a call to `Invoke`. The function binds to itself. */ - class IntAsDecimalFunction(precision: Int, scale: Int, valueScale: Int) + class IntAsDecimalFunction(declared: DataType, valueScale: Int = 0, nullable: Boolean = true) extends UnboundFunction - with ScalarFunction[Decimal] { + with ScalarFunction[Any] { override def name(): String = "int_as_decimal" - override def description(): String = - s"int -> decimal($precision, $scale), at scale $valueScale" + override def description(): String = s"int -> ${declared.sql}, at scale $valueScale" override def bind(inputType: StructType): BoundFunction = this override def inputTypes(): Array[DataType] = Array(IntegerType) - override def resultType(): DataType = DecimalType(precision, scale) - // Ten digits hold any `INT`. - def invoke(v: Int): Decimal = Decimal(v.toLong, 10, valueScale) - override def produceResult(input: InternalRow): Decimal = invoke(input.getInt(0)) + override def resultType(): DataType = declared + override def isResultNullable(): Boolean = nullable + def invoke(v: Int): Any = valueOf(declared, v) + override def produceResult(input: InternalRow): Any = invoke(input.getInt(0)) + + private def valueOf(dataType: DataType, v: Int): Any = dataType match { + // Ten digits hold any `INT`. + case _: DecimalType => Decimal(v.toLong, 10, valueScale) + case ArrayType(elementType, _) => new GenericArrayData(Array(valueOf(elementType, v))) + case struct: StructType => + new GenericInternalRow(struct.fields.map(f => valueOf(f.dataType, v))) + } } } +/** + * `as_money` for the #6425 tests, with `invoke` on the companion object. Scala also compiles a + * top-level companion's methods to static methods on the class, so Spark finds a static `invoke` + * and lowers a call to `StaticInvoke`. + */ +class StaticAsMoneyFunction extends UnboundFunction with ScalarFunction[Decimal] { + override def name(): String = "static_as_money" + override def description(): String = "int -> decimal(10, 2), at scale 0" + override def bind(inputType: StructType): BoundFunction = this + override def inputTypes(): Array[DataType] = Array(IntegerType) + override def resultType(): DataType = DecimalType(10, 2) + override def produceResult(input: InternalRow): Decimal = + StaticAsMoneyFunction.invoke(input.getInt(0)) +} + +object StaticAsMoneyFunction { + def invoke(v: Int): Decimal = Decimal(v) +} + /** * Case class used by the struct-input / struct-output smoke tests. Must be declared at file scope * (not inside the test class) so Spark's TypeTag-based UDF encoder can resolve the Spark From 5c9482c9139d5ce2e716af00f83ec650ba37d57f Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 1 Oct 2026 12:34:26 -0600 Subject: [PATCH 4/7] fix: dispatch an expression with the DSv2 decimal call it reads, and fall back for aggregates Spark rescales a DSv2 function's decimal result to its declared type, and writes null when it does not fit, only when it writes a row. An expression around the call reads the Decimal the function returned. The dispatcher writes an Arrow vector of the declared type, so it nulls at its own output, and a native expression over that output read something else: IS NULL was true for a value that does not fit, a cast to string kept the declared scale, and hash read the unscaled value at the declared scale. An expression that takes such a call as an argument now runs in the same kernel as the call, so Spark's own code reads the value the function returned. An aggregate cannot run in the kernel, so count, max and sum over the call fall back to Spark. (cherry picked from commit 2949bda8a7270de490fdfe05097570934d5f139f) --- .../CometBatchKernelCodegenOutput.scala | 4 ++ .../apache/comet/serde/QueryPlanSerde.scala | 58 +++++++++++++++++ .../org/apache/comet/serde/statics.scala | 17 +++-- .../org/apache/comet/CometCodegenSuite.scala | 62 +++++++++++++++++++ 4 files changed, 135 insertions(+), 6 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala index 89713827aa8..c3fa9476c46 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -240,6 +240,10 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { // callers write null values themselves, and skip that test only for a type that is not // nullable. // + // Only the kernel's own output sees the null. Spark rescales such a value when it writes a + // row, and an expression around the call reads the value the function returned, so + // `QueryPlanSerde` dispatches that expression with the call, and falls an aggregate back. + // // The precision and scale test repeats `changePrecision`'s own fast path. It keeps the call // off the common path, so the JIT can still scalar-replace the `Decimal` that an input // getter allocates. With the bare call, passing a `DECIMAL(18, 2)` column through took 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 1d87294c20b..25a2256b348 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -31,6 +31,7 @@ import org.apache.spark.sql.catalyst.expressions.aggregate._ import org.apache.spark.sql.catalyst.expressions.objects.{Invoke, StaticInvoke} import org.apache.spark.sql.catalyst.expressions.xml.{XPathBoolean, XPathDouble, XPathFloat, XPathInt, XPathList, XPathLong, XPathShort, XPathString} import org.apache.spark.sql.comet.DecimalPrecision +import org.apache.spark.sql.connector.catalog.functions.ScalarFunction import org.apache.spark.sql.execution.{ScalarSubquery, SparkPlan} import org.apache.spark.sql.execution.datasources.parquet.ParquetUtils import org.apache.spark.sql.internal.SQLConf @@ -835,6 +836,17 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { } val fn = aggExpr.aggregateFunction + // An aggregate reads the `Decimal` a DSv2 function returned, whether or not it fits the + // declared type, and the kernel has nulled such a value at its own output: Spark counts it + // and takes it as the maximum, and the row writer nulls the result afterwards. An aggregate + // cannot run in the kernel, so it falls back. See [[readsDispatchedDsv2Decimal]]. + if (fn.children.exists(isDispatchedDsv2DecimalCall)) { + withFallbackReason( + aggExpr, + s"${fn.prettyName} aggregates the decimal result of a DSv2 function, which Spark " + + "rescales to the declared type only when it writes a row") + return None + } val cometExpr = aggrSerdeMap.get(fn.getClass) val protoAggExprOpt = cometExpr match { case Some(handler) => @@ -1017,6 +1029,12 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { s"${CometConf.getExprEnabledConfigKey(exprConfName)}=true to enable it.") return None } + if (readsDispatchedDsv2Decimal(expr)) { + // Returns None, tagged with the reason, when the dispatcher declines, which falls the + // operator back to Spark. Converting `expr` natively is not an alternative: it would read + // the call's output rather than the value the function returned. + return CometScalaUDF.emitJvmCodegenDispatch(expr, inputs, binding) + } handler.getSupportLevel(expr) match { case Unsupported(notes) => // `CodegenDispatchFallback` serdes have no native path for these cases either, but the @@ -1125,6 +1143,46 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { case _ => false } + /** + * Whether `expr` takes the decimal result of a DSv2 scalar function call, which the codegen + * dispatcher runs, as an argument. + * + * Spark does not rescale such a result to the type the function declares, or write null when it + * does not fit, until it writes a row. An expression around the call reads the `Decimal` the + * function returned. The dispatcher has to write an Arrow vector of the declared type, so it + * rescales and nulls at its own output, and a native expression over that output would read + * something else. `IS NULL` of a value that does not fit is false in Spark, a cast to string + * keeps the function's scale, and `hash` reads the unscaled value at that scale (#6425). So + * such an expression runs in the same kernel as the call, where Spark's own code reads the + * value the function returned. `Alias` is skipped because it computes nothing: the call under + * it is the root, and Spark writes a root as a row. + */ + private def readsDispatchedDsv2Decimal(expr: Expression): Boolean = + !isStructuralExpr(expr) && expr.children.exists(isDispatchedDsv2DecimalCall) + + private def isDispatchedDsv2DecimalCall(expr: Expression): Boolean = { + val dispatchedDsv2Call = expr match { + case i: Invoke => + i.targetObject match { + case Literal(_: ScalarFunction[_], _) => true + case _ => false + } + case s: StaticInvoke => + classOf[ScalarFunction[_]].isAssignableFrom(s.staticObject) && + CometStaticInvoke.runsInDispatcher(s) + case _ => false + } + dispatchedDsv2Call && containsDecimal(expr.dataType) + } + + private def containsDecimal(dataType: DataType): Boolean = dataType match { + case _: DecimalType => true + case ArrayType(elementType, _) => containsDecimal(elementType) + case MapType(keyType, valueType, _) => containsDecimal(keyType) || containsDecimal(valueType) + case StructType(fields) => fields.exists(f => containsDecimal(f.dataType)) + case _ => false + } + /** * Creates a UnaryExpr by calling exprToProtoInternal for the provided child expression and then * invokes the supplied function to wrap this UnaryExpr in a top-level Expr. diff --git a/spark/src/main/scala/org/apache/comet/serde/statics.scala b/spark/src/main/scala/org/apache/comet/serde/statics.scala index 94b845b7348..df42a5fa049 100644 --- a/spark/src/main/scala/org/apache/comet/serde/statics.scala +++ b/spark/src/main/scala/org/apache/comet/serde/statics.scala @@ -58,6 +58,10 @@ object CometStaticInvoke extends CometExpressionSerde[StaticInvoke] { private def handlerFor(expr: StaticInvoke): Option[CometExpressionSerde[StaticInvoke]] = staticInvokeExpressions.get((expr.functionName, expr.staticObject.getName)) + /** Whether [[convert]] hands `expr` to the codegen dispatcher, because no handler claims it. */ + def runsInDispatcher(expr: StaticInvoke): Boolean = + handlerFor(expr).isEmpty && icebergHandlerFor(expr).isEmpty + /** * Iceberg's system functions are keyed by class name only because both the `StaticInvoke` and * `ApplyFunctionExpression` lowerings share the same identity class. Consulted after the @@ -80,12 +84,13 @@ object CometStaticInvoke extends CometExpressionSerde[StaticInvoke] { * [[CodegenDispatchFallback]]: that would also route a *handler's* `Unsupported` through the * dispatcher, and at least one of those is not dispatchable. `CometIcebergTruncate` declines a * decimal because Iceberg's `truncate` can return a value wider than the column's declared - * precision, which Spark nulls only when the row is materialized; the dispatcher writes into an - * Arrow `Decimal128(precision, scale)` vector just like a native kernel does, so it has to null - * the value at its own output, and an enclosing predicate or hash then sees a null where Spark - * sees the oversized value. The mixin's contract ("the case must be something `doGenCode` can - * compile") does not cover a limit that lives at the Arrow output boundary, so enrollment stays - * with the individual handlers. + * precision, which Spark nulls only when the row is materialized. The dispatcher nulls such a + * value at its own output, as the row writer does, and an expression or aggregate that reads a + * dispatched call either runs in the same kernel or falls back (see + * `QueryPlanSerde.readsDispatchedDsv2Decimal`), but the decline predates that and stays. The + * mixin's contract ("the case must be something `doGenCode` can compile") does not cover a + * limit that lives at the Arrow output boundary, so enrollment stays with the individual + * handlers. */ override def getSupportLevel(expr: StaticInvoke): SupportLevel = handlerFor(expr) diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 9ab400c67a0..bb9aa867b4e 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -39,6 +39,7 @@ import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ import org.apache.spark.sql.util.CaseInsensitiveStringMap +import org.apache.spark.unsafe.hash.Murmur3_x86_32 import org.apache.spark.unsafe.types.UTF8String import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus @@ -2440,6 +2441,67 @@ class CometCodegenSuite s"expected Comet's native projection to reject the null, got $cometError") } } + + test("an expression around a DSv2 decimal result reads the value Spark reads (#6425)") { + // Spark rescales the call's `Decimal` to its declared type, and writes null when it does not + // fit, only when it writes a row. An expression around the call reads the `Decimal` the + // function returned, at the scale the function chose. `as_money` returns `Decimal(i)` at scale + // 0 for a `DECIMAL(10, 2)`: + // - `IS NULL` is false for 100000000, which needs nine integer digits where the type allows + // eight, so the row writer's null for it must not reach the expression; + // - a cast to a wider decimal keeps that value, and so does a comparison, which finds it + // greater than 5; + // - `CAST(... AS STRING)` prints "3", not "3.00"; + // - `hash` hashes the unscaled value, so `Decimal(3, 10, 0)` hashes as the long 3. + // The kernel nulls an oversized decimal at its own output, so each of these expressions has + // to run in the same kernel as the call. + val hashes = Seq(3L, null, 100000000L, -100000000L).map { + case null => 42 + case v: Long => Murmur3_x86_32.hashLong(v, 42) + } + withDecimalFunctions(3, null, 100000000, -100000000) { + for (ansi <- Seq("true", "false")) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { + val df = sql( + "SELECT i, decfn.ns.as_money(i) IS NULL, " + + "CAST(decfn.ns.as_money(i) AS DECIMAL(12, 2)), " + + "CAST(decfn.ns.as_money(i) AS STRING), hash(decfn.ns.as_money(i)) FROM t") + assertCodegenRan { + checkSparkAnswerAndImpl(df, dispatched = Seq("isnull", "cast", "hash")) + } + checkAnswer( + df, + Seq( + Row(3, false, dec("3.00"), "3", hashes(0)), + Row(null, true, null, null, hashes(1)), + Row(100000000, false, dec("100000000.00"), "100000000", hashes(2)), + Row(-100000000, false, dec("-100000000.00"), "-100000000", hashes(3)))) + } + } + val filtered = sql("SELECT i FROM t WHERE decfn.ns.as_money(i) > 5") + assertCodegenRan { + checkSparkAnswerAndImpl(filtered, dispatched = Seq("greaterthan")) + } + checkAnswer(filtered, Row(100000000)) + } + } + + test( + "an aggregate over a DSv2 decimal result falls back, as Spark aggregates the value (#6425)") { + // Spark aggregates the `Decimal` the function returned, and the row writer nulls the result + // afterwards. `count` counts 100000000, `max` is that value and so is null, and the two + // values that do not fit cancel in the sum of the even group. An aggregate cannot run in the + // kernel, which has nulled those values at its output. + withDecimalFunctions(3, null, 100000000, -100000000) { + val reason = "aggregates the decimal result of a DSv2 function" + val global = sql("SELECT count(decfn.ns.as_money(i)), max(decfn.ns.as_money(i)) FROM t") + checkSparkAnswerAndFallbackReason(global, reason) + checkAnswer(global, Row(3L, null)) + val grouped = sql("SELECT i % 2, sum(decfn.ns.as_money(i)) FROM t GROUP BY i % 2") + checkSparkAnswerAndFallbackReason(grouped, reason) + checkAnswer(grouped, Seq(Row(0, dec("0.00")), Row(1, dec("3.00")), Row(null, null))) + } + } } /** From 66ec6969645276d9d390a8bcfecc89d805b1a07c Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 1 Oct 2026 19:29:26 -0600 Subject: [PATCH 5/7] fix: preserve DSv2 decimal values through transitive consumers (cherry picked from commit 5d07da30e347909c38e288a951e839fdceaecb9e) --- .../apache/comet/serde/QueryPlanSerde.scala | 14 ++--- .../org/apache/comet/CometCodegenSuite.scala | 53 +++++++++++++++++++ 2 files changed, 61 insertions(+), 6 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 25a2256b348..dff3a158ed2 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -840,7 +840,7 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { // declared type, and the kernel has nulled such a value at its own output: Spark counts it // and takes it as the maximum, and the row writer nulls the result afterwards. An aggregate // cannot run in the kernel, so it falls back. See [[readsDispatchedDsv2Decimal]]. - if (fn.children.exists(isDispatchedDsv2DecimalCall)) { + if (fn.children.exists(_.exists(isDispatchedDsv2DecimalCall))) { withFallbackReason( aggExpr, s"${fn.prettyName} aggregates the decimal result of a DSv2 function, which Spark " + @@ -1144,8 +1144,8 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { } /** - * Whether `expr` takes the decimal result of a DSv2 scalar function call, which the codegen - * dispatcher runs, as an argument. + * Whether `expr` consumes a decimal result of a dispatched DSv2 scalar function anywhere in its + * argument trees, including through intermediate expressions and container access. * * Spark does not rescale such a result to the type the function declares, or write null when it * does not fit, until it writes a row. An expression around the call reads the `Decimal` the @@ -1154,11 +1154,13 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { * something else. `IS NULL` of a value that does not fit is false in Spark, a cast to string * keeps the function's scale, and `hash` reads the unscaled value at that scale (#6425). So * such an expression runs in the same kernel as the call, where Spark's own code reads the - * value the function returned. `Alias` is skipped because it computes nothing: the call under - * it is the root, and Spark writes a root as a row. + * value the function returned. Checking only immediate children would let an intermediate + * expression, such as `abs(call)` or `call[0]`, normalize the decimal before its parent reads + * it. `Alias` is skipped because it computes nothing: the call under it is the root, and Spark + * writes a root as a row. */ private def readsDispatchedDsv2Decimal(expr: Expression): Boolean = - !isStructuralExpr(expr) && expr.children.exists(isDispatchedDsv2DecimalCall) + !isStructuralExpr(expr) && expr.children.exists(_.exists(isDispatchedDsv2DecimalCall)) private def isDispatchedDsv2DecimalCall(expr: Expression): Boolean = { val dispatchedDsv2Call = expr match { diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index bb9aa867b4e..05924b7292c 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -2486,6 +2486,59 @@ class CometCodegenSuite } } + test("transitive consumers of DSv2 decimals preserve materialization timing (#6425)") { + withDecimalFunctions(3, null, 100000000, -100000000) { + for (ansi <- Seq("true", "false")) { + withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi) { + val projected = sql( + "SELECT i, abs(decfn.ns.as_money(i)) IS NULL, " + + "abs(decfn.ns.static_as_money(i)) IS NULL, " + + "decfn.ns.money_array(i)[0] IS NULL, " + + "decfn.ns.money_struct(i).m IS NULL, " + + "CAST(abs(decfn.ns.as_money(i)) AS STRING) FROM t") + assertCodegenRan { + checkSparkAnswerAndImpl(projected, dispatched = Seq("isnull", "cast")) + } + checkAnswer( + projected, + Seq( + Row(3, false, false, false, false, "3"), + Row(null, true, true, true, true, null), + Row(100000000, false, false, false, false, "100000000"), + Row(-100000000, false, false, false, false, "100000000"))) + val filtered = sql("SELECT i FROM t WHERE abs(decfn.ns.as_money(i)) IS NULL") + checkSparkAnswerAndImpl(filtered, dispatched = Seq("isnull")) + checkAnswer(filtered, Seq(Row(null))) + val aggregated = sql( + "SELECT count(abs(decfn.ns.as_money(i))), " + + "count(abs(decfn.ns.static_as_money(i))), " + + "count(decfn.ns.money_array(i)[0]), " + + "count(decfn.ns.money_struct(i).m), max(abs(decfn.ns.as_money(i))), " + + "sum(abs(decfn.ns.as_money(i))) FROM t") + checkSparkAnswerAndFallbackReason( + aggregated, + "aggregates the decimal result of a DSv2 function") + checkAnswer(aggregated, Row(3L, 3L, 3L, 3L, null, dec("200000003.00"))) + withSQLConf(CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false") { + checkSparkAnswerAndFallbackReason( + "SELECT abs(decfn.ns.as_money(i)) IS NULL FROM t", + "spark.comet.exec.scalaUDF.codegen.enabled=false") + } + // A physical write is a real boundary; a view would be collapsed by Catalyst. + withTable("decimal_materialized") { + sql( + "CREATE TABLE decimal_materialized USING parquet AS " + + "SELECT i, abs(decfn.ns.as_money(i)) AS d FROM t") + checkSparkAnswerAndOperator("SELECT i, d IS NULL FROM decimal_materialized") + checkAnswer( + sql("SELECT i FROM decimal_materialized WHERE d IS NULL"), + Seq(Row(null), Row(100000000), Row(-100000000))) + } + } + } + } + } + test( "an aggregate over a DSv2 decimal result falls back, as Spark aggregates the value (#6425)") { // Spark aggregates the `Decimal` the function returned, and the row writer nulls the result From 76a0e1c25d2f50dc7cf21e9f49bb7192acf502e0 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 1 Oct 2026 21:42:26 -0600 Subject: [PATCH 6/7] fix: preserve DSv2 decimal aliases across physical operators (cherry picked from commit c655d630696579c8d7ec3eac784afdc4a7fd1261) --- .../apache/comet/rules/CometExecRule.scala | 62 ++++++++++++++++++- .../apache/comet/serde/QueryPlanSerde.scala | 8 ++- .../org/apache/comet/CometCodegenSuite.scala | 54 ++++++++++++++++ 3 files changed, 121 insertions(+), 3 deletions(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 19c88a6a9d8..7c74fd174a5 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -22,7 +22,7 @@ package org.apache.comet.rules import scala.collection.mutable.ListBuffer import org.apache.spark.sql.SparkSession -import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder} +import org.apache.spark.sql.catalyst.expressions.{Divide, DoubleLiteral, EqualNullSafe, EqualTo, Expression, ExprId, FloatLiteral, GreaterThan, GreaterThanOrEqual, KnownFloatingPointNormalized, LessThan, LessThanOrEqual, NamedExpression, Remainder} import org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateMode, Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.optimizer.NormalizeNaNAndZero import org.apache.spark.sql.catalyst.rules.Rule @@ -70,6 +70,10 @@ object CometExecRule { val COMET_UNSAFE_PARTIAL: TreeNodeTag[String] = TreeNodeTag[String]("comet.unsafePartialAgg") + /** Keep a DSv2 decimal producer and its consumers within Spark's row-materialization path. */ + val COMET_UNMATERIALIZED_DECIMAL: TreeNodeTag[Unit] = + TreeNodeTag[Unit]("comet.unmaterializedDecimal") + /** * Fully native operators. */ @@ -339,6 +343,12 @@ case class CometExecRule(session: SparkSession) // spotless:on private def transform(plan: SparkPlan): SparkPlan = { def convertNode(op: SparkPlan): SparkPlan = op match { + case op if op.getTagValue(CometExecRule.COMET_UNMATERIALIZED_DECIMAL).isDefined => + withFallbackReason( + op, + "DSv2 decimal projection and its consumers must share Spark materialization") + op + // Scan marker produced by an optional, out-of-tree scan contrib (e.g. contrib/delta). // Matched by trait (no compile-time dependency on the contrib) and present only when that // contrib is on the classpath. The marker carries its own serde handler and typically wraps @@ -681,6 +691,55 @@ case class CometExecRule(session: SparkSession) } } + /** + * Spark can pass a Decimal between fused operators without writing a row. A retained Project + * therefore cannot write a dispatched DSv2 decimal to Arrow before its consumers run: that + * would rescale/null the value early. Follow projected attributes and keep both ends, including + * intervening operators, in Spark. Declining only the consumer would leave the early write. + * + * Paths are propagated conservatively through decimal outputs, including row boundaries. This + * may retain more Spark operators after a real materialization, but never introduces a new one + * before a consumer. Unrelated inputs, and trees that consume the call within one expression, + * remain eligible for native execution. + */ + private def tagUnmaterializedDecimals(plan: SparkPlan): Unit = { + type Paths = Map[ExprId, Set[SparkPlan]] + def visit(op: SparkPlan): Paths = { + val inputs = op.children.flatMap(visit).groupBy(_._1).map { case (id, paths) => + id -> paths.flatMap(_._2).toSet + } + val consumed = op.references.toSeq.flatMap(a => inputs.getOrElse(a.exprId, Set.empty)).toSet + if (consumed.nonEmpty) { + (consumed + op).foreach(_.setTagValue(CometExecRule.COMET_UNMATERIALIZED_DECIMAL, ())) + } + val sources = op match { + case p: ProjectExec => + p.projectList + .filter(QueryPlanSerde.producesUnmaterializedDsv2Decimal) + .map(_.exprId) + .toSet + case _ => Set.empty[ExprId] + } + op.output.flatMap { attr => + val inherited = inputs.getOrElse(attr.exprId, Set.empty) + // Computed decimal outputs can retain the raw value as well (e.g. abs(d), or max(d)). + val derived = if (containsDecimal(attr.dataType)) consumed else Set.empty[SparkPlan] + val path = inherited ++ derived + if (sources.contains(attr.exprId) || path.nonEmpty) Some(attr.exprId -> (path + op)) + else None + }.toMap + } + visit(plan) + } + + private def containsDecimal(dataType: DataType): Boolean = dataType match { + case _: DecimalType => true + case ArrayType(elementType, _) => containsDecimal(elementType) + case MapType(keyType, valueType, _) => containsDecimal(keyType) || containsDecimal(valueType) + case StructType(fields) => fields.exists(f => containsDecimal(f.dataType)) + case _ => false + } + private def normalizePlan(plan: SparkPlan): SparkPlan = { plan.transformUp { case p: ProjectExec => @@ -774,6 +833,7 @@ case class CometExecRule(session: SparkSession) // corresponding Final or PartialMerge cannot be converted and the intermediate buffer // formats are incompatible. This runs before transform() so the tags are checked // during the bottom-up conversion. Tags persist through AQE stage creation. + tagUnmaterializedDecimals(planWithJoinRewritten) tagUnsafePartialAggregates(planWithJoinRewritten) var newPlan = revertUnsafePartialAggregates(transform(planWithJoinRewritten)) 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 dff3a158ed2..d0e15e18a1f 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -1156,12 +1156,16 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { * such an expression runs in the same kernel as the call, where Spark's own code reads the * value the function returned. Checking only immediate children would let an intermediate * expression, such as `abs(call)` or `call[0]`, normalize the decimal before its parent reads - * it. `Alias` is skipped because it computes nothing: the call under it is the root, and Spark - * writes a root as a row. + * it. `Alias` is skipped because it computes nothing. A projected root is not necessarily a + * row-materialization boundary: CometExecRule keeps retained decimal aliases and their + * consumers in Spark until those consumers have read the original value. */ private def readsDispatchedDsv2Decimal(expr: Expression): Boolean = !isStructuralExpr(expr) && expr.children.exists(_.exists(isDispatchedDsv2DecimalCall)) + private[comet] def producesUnmaterializedDsv2Decimal(expr: Expression): Boolean = + containsDecimal(expr.dataType) && expr.exists(isDispatchedDsv2DecimalCall) + private def isDispatchedDsv2DecimalCall(expr: Expression): Boolean = { val dispatchedDsv2Call = expr match { case i: Invoke => diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index 05924b7292c..10ed793d260 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -2539,6 +2539,60 @@ class CometCodegenSuite } } + test("retained DSv2 decimal aliases preserve Spark materialization timing (#6425)") { + withDecimalFunctions(3, null, 100000000, -100000000) { + val reason = "DSv2 decimal projection and its consumers must share Spark materialization" + for (ansi <- Seq("true", "false"); aqe <- Seq("true", "false")) { + withSQLConf( + SQLConf.ANSI_ENABLED.key -> ansi, + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> aqe) { + for (expr <- Seq( + "decfn.ns.as_money(i)", + "decfn.ns.static_as_money(i)", + "abs(decfn.ns.as_money(i))")) { + val from = s"FROM (SELECT $expr AS d FROM t)" + val projected = sql(s"SELECT d, d IS NULL $from") + checkSparkAnswerAndFallbackReason(projected, reason) + checkAnswer( + projected, + Seq(Row(dec("3.00"), false), Row(null, true), Row(null, false), Row(null, false))) + // The retained inner projection must also fall back. A Spark consumer above a + // Comet producer would still read the prematurely normalized Arrow value. + assert( + !projected.queryExecution.executedPlan.exists(_.isInstanceOf[CometProjectExec])) + val aggregated = sql(s"SELECT count(d), max(d) $from") + checkSparkAnswerAndFallbackReason(aggregated, reason) + assert(aggregated.collect().head.getLong(0) == 3L) + checkSparkAnswerAndFallbackReason( + s"SELECT d, d IS NULL $from WHERE d IS NOT NULL", + reason) + checkSparkAnswerAndFallbackReason( + s"SELECT d, e, e IS NULL FROM (SELECT d, abs(d) AS e $from)", + reason) + } + for ((call, access) <- Seq("money_array" -> "d[0]", "money_struct" -> "d.m")) { + checkSparkAnswerAndFallbackReason( + s"SELECT d, $access IS NULL FROM (SELECT decfn.ns.$call(i) AS d FROM t)", + reason) + checkSparkAnswerAndFallbackReason( + s"SELECT count($access), max($access) FROM (SELECT decfn.ns.$call(i) AS d FROM t)", + reason) + } + checkSparkAnswerAndFallbackReason( + "SELECT d, CAST(d AS STRING) FROM (SELECT decfn.ns.mills_as_money(i) AS d FROM t)", + reason) + // Without fusion, Spark materializes each Project. Retaining the same operators also + // preserves that earlier boundary, rather than always exposing the raw Decimal. + withSQLConf(SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false") { + checkSparkAnswerAndFallbackReason( + "SELECT d, d IS NULL FROM (SELECT decfn.ns.as_money(i) AS d FROM t)", + reason) + } + } + } + } + } + test( "an aggregate over a DSv2 decimal result falls back, as Spark aggregates the value (#6425)") { // Spark aggregates the `Decimal` the function returned, and the row writer nulls the result From facb668c8af7e9eabc26d0118f5cf8a1fb11ce62 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Fri, 2 Oct 2026 03:45:08 -0600 Subject: [PATCH 7/7] fix: explicitly discard decimal provenance traversal result (cherry picked from commit e1972d144bc0ed18d121569dfe0d58716b5cd235) --- spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 7c74fd174a5..369cef6190e 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -729,7 +729,7 @@ case class CometExecRule(session: SparkSession) else None }.toMap } - visit(plan) + val _ = visit(plan) } private def containsDecimal(dataType: DataType): Boolean = dataType match {