andygrove opened a new pull request, #6455:
URL: https://github.com/apache/datafusion-comet/pull/6455

   ## Which issue does this PR close?
   
   Closes #6425.
   
   Found by the 1.1.0 regression audit (#6399) and tracked in #6402.
   
   ## Rationale for this change
   
   Since #5692, Comet runs `Invoke` and `StaticInvoke` calls that it has no 
native handler for through the JVM codegen dispatcher. That includes calls to 
DSv2 catalog functions, which Spark lowers to one of the two when the function 
has an `invoke` magic method. Such a function can return a `Decimal` whose 
scale differs from its declared result type.
   
   Spark rescales the value when it writes the row. 
`UnsafeRowWriter.write(ordinal, Decimal, precision, scale)` calls 
`changePrecision(precision, scale)` and writes null if that returns false. The 
dispatcher's writer skipped that step:
   
   - For precision up to 18 it wrote the unscaled long as is. A `DECIMAL(10, 
2)` result came back as 0.03 where Spark returns 3.00, and a value too wide for 
the type came back where Spark returns null.
   - Above precision 18, Arrow's `DecimalVector.setSafe(BigDecimal)` rejected 
the scale mismatch and the task failed.
   
   In 1.0.0 these calls fell back to Spark, so this is a 1.1.0 regression.
   
   I fixed the writer rather than declining dispatch for decimal-returning 
`Invoke` / `StaticInvoke`. Declining would not close the hole. A dispatched 
parent such as `CreateMap` serializes its whole subtree into one kernel, so 
`map('k', fn(i))` would still run the call in the JVM and write it through the 
same writer. For a value that already matches its declared type, as the values 
of Spark's own expressions do, the writer's output is unchanged.
   
   ## What changes are included in this PR?
   
   - `CometBatchKernelCodegenOutput.emitWrite` now rescales each decimal to the 
declared precision and scale before writing it, and writes null when it doesn't 
fit, as `UnsafeRowWriter` and `UnsafeArrayWriter` do. This covers top-level 
values and array, struct and map children.
   - Like Spark's writers, it rescales the `Decimal` in place rather than 
copying it. `changePrecision` leaves the value untouched when it returns false.
   - ANSI mode doesn't change the result. Spark adds no overflow check around 
the lowered call (`V2ExpressionUtils.resolveScalarFunction`), and its row 
writer nulls in both modes. The dispatcher now does the same.
   - The generated code compares precision and scale before it calls 
`changePrecision`, which repeats the method's own fast path. The method is 
about 830 bytes of bytecode, too big for C2 to inline. A bare call stops the 
JIT from scalar-replacing the `Decimal` that an input getter allocates, and 
cost about 50% more per row on a pass-through `DECIMAL(18, 2)` column (numbers 
below).
   - The `CometStaticInvoke.getSupportLevel` scaladoc described the old 
behavior, so I updated it. Iceberg's decimal `truncate` stays excluded: the 
dispatcher now nulls the oversized value at its output, but an enclosing 
predicate or hash then sees a null where Spark sees the value.
   
   One edge case: a function that declares a non-nullable result but returns a 
decimal that doesn't fit. Spark's `collect()` then fails with 
`EXPRESSION_DECODING_FAILED`, because its row writer put a null in a 
non-nullable column. With this change Comet also fails, with `Column ... is 
declared as non-nullable but contains null values` from the native projection. 
Before, it returned the out-of-range value.
   
   ## How are these changes tested?
   
   The new `CometCodegenSuite` test, "decimal results of a DSv2 function are 
rescaled to the declared type (#6425)", is built from the issue's reproducer. 
Its function catalog has `as_money(int)`, declared `DECIMAL(10, 2)`, and 
`as_wide_money(int)`, declared `DECIMAL(20, 12)`, which covers the writer for 
precision above 18. Both return `Decimal(i)` at scale 0.
   
   The test selects both functions, plus `map('k', as_money(i))` for the nested 
writer, over 3, -7, NULL, 99999999 and 100000000. The last value needs nine 
integer digits and both types allow eight, so Spark returns null for it. The 
test runs with ANSI on and off. It compares with Spark using 
`checkSparkAnswerAndOperator` inside `assertCodegenRan`, and also checks the 
expected rows explicitly.
   
   Without the fix, the test fails with `BigDecimal scale must equal that in 
the Arrow vector: 0 != 12`. With the `DECIMAL(20, 12)` column dropped, Comet 
returns 0.03, -0.07, 999999.99 and 1000000.00 where Spark returns 3.00, -7.00, 
99999999.00 and null, and the map value is wrong the same way. With the fix the 
test passes. I also added `changePrecision` assertions to the two decimal 
writer tests in `CometCodegenSourceSuite`.
   
   These runs all pass:
   
   - Spark 4.1 (default profile): `./mvnw test -Dtest=none 
-Dsuites="org.apache.comet.CometCodegenSuite,org.apache.comet.CometCodegenSourceSuite,org.apache.comet.CometCodegenFuzzSuite,org.apache.comet.CometCodegenHOFSuite,org.apache.comet.CometIcebergSystemFunctionSuite"`,
 210 tests. `CometCodegenFuzzSuite` sweeps ScalaUDF decimal outputs across the 
precision-18 boundary, so the ScalaUDF results are unchanged.
   - Spark 4.1: `./mvnw test -Dtest=none 
-Dsuites="org.apache.comet.CometSqlFileTestSuite to_number"`, which runs 
`to_number.sql` and `try_to_number.sql`. Their decimal results come from the 
dispatcher.
   - Spark 3.5: `CometCodegenSuite`, `CometCodegenSourceSuite` and 
`CometCodegenFuzzSuite`, 191 passed and 1 canceled (the TIME test, which needs 
Spark 4.1).
   - Spark 3.4: the new test.
   - `spotless:check`, scalastyle, and the scalafix check on Spark 3.5 and 
Scala 2.12.
   
   To measure the hot-path cost, I timed kernels for a pass-through decimal 
column with a scratch test that isn't committed. It used 8192-row batches, with 
all three variants interleaved in one JVM:
   
   | Writer                      | `DECIMAL(18, 2)` | `DECIMAL(38, 10)` |
   | --------------------------- | ---------------- | ----------------- |
   | Before this PR              | 4.42 ns/row      | 53.9 ns/row       |
   | Bare `changePrecision` call | 6.74 ns/row      | 54.0 ns/row       |
   | This PR                     | 4.40 ns/row      | 53.9 ns/row       |
   


-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to