diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index ff864ea7510..6659440351b 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -489,6 +489,7 @@ jobs: org.apache.comet.CometJsonExpressionSuite org.apache.comet.CometJsonJvmSuite org.apache.comet.SparkErrorConverterSuite + org.apache.comet.serde.SerdeRegistrationSuite org.apache.comet.expressions.conditional.CometIfSuite org.apache.comet.expressions.conditional.CometCoalesceSuite org.apache.comet.expressions.conditional.CometCaseWhenSuite diff --git a/.github/workflows/pr_build_macos.yml b/.github/workflows/pr_build_macos.yml index d5480dca112..6a536232a67 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -234,6 +234,7 @@ jobs: org.apache.comet.CometJsonExpressionSuite org.apache.comet.CometJsonJvmSuite org.apache.comet.SparkErrorConverterSuite + org.apache.comet.serde.SerdeRegistrationSuite org.apache.comet.expressions.conditional.CometIfSuite org.apache.comet.expressions.conditional.CometCoalesceSuite org.apache.comet.expressions.conditional.CometCaseWhenSuite diff --git a/docs/source/contributor-guide/expression-audits/math_funcs.md b/docs/source/contributor-guide/expression-audits/math_funcs.md index 6a9d7868519..984a0468fd0 100644 --- a/docs/source/contributor-guide/expression-audits/math_funcs.md +++ b/docs/source/contributor-guide/expression-audits/math_funcs.md @@ -278,8 +278,9 @@ Internal fused expression that replaces the `CheckOverflow(Cast(expr, Decimal128 ## width_bucket -- Spark 3.5.8 (audited 2026-05-27): introduced; not available in 3.4.3. +- Spark 3.4.3 (audited 2026-09-14): present in catalyst and the function registry with the same semantics as 3.5.8. +- Spark 3.5.8 (audited 2026-05-27): baseline. - Spark 4.0.1, 4.1.1 (audited 2026-05-27): same semantics; `NullIntolerant` -> `nullIntolerant: Boolean` refactor. -- Known limitation: wired via per-version `CometExprShim` rather than a `CometExpressionSerde`, so it bypasses the support-level framework and the auto-generated compatibility doc ([#4485](https://github.com/apache/datafusion-comet/issues/4485)). Native path uses datafusion-spark `SparkWidthBucket`; interval input types are not exercised by Comet tests. +- Wiring (audited 2026-09-14): `CometWidthBucket`, a `CometCodegenDispatch` registered once in the shared math group, so it goes through the support-level framework and the compatibility doc on every Spark line. `width_bucket.sql` exercises double and both interval input types. [Spark Expression Support]: ../../user-guide/latest/expressions.md 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 be4bc9c3412..eced2d7fdf9 100644 --- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala +++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala @@ -110,10 +110,10 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[Not] -> CometNot, classOf[Or] -> CometOr) - private[comet] val mathExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = { - // Explicit type ascription on `base`: Scala 2.13 cannot infer the existential key type - // when `++` is applied directly to a `Map(...)` literal. - val base: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( + // The shared maps keep an explicit type: Scala 2.13 cannot infer the existential key type + // when `++` is applied to a `Map(...)` literal, and the version shims are merged over them. + private[comet] val baseMathExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + Map( classOf[Acos] -> CometScalarFunction("acos"), classOf[Acosh] -> CometScalarFunction("acosh"), classOf[Add] -> CometAdd, @@ -172,13 +172,11 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[Pmod] -> CometPmod, classOf[WidthBucket] -> CometWidthBucket, classOf[UnaryPositive] -> CometUnaryPositive) - base ++ sparkVersionSpecificMathExpressions - } + private[comet] val mathExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + baseMathExpressions ++ sparkVersionSpecificMathExpressions - private[comet] val mapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = { - // Explicit type ascription on `base`: Scala 2.13 cannot infer the existential key type - // when `++` is applied directly to a `Map(...)` literal. - val base: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( + private[comet] val baseMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + Map( classOf[GetMapValue] -> CometMapExtract, classOf[MapKeys] -> CometMapKeys, classOf[MapEntries] -> CometMapEntries, @@ -192,8 +190,8 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[TransformValues] -> CometTransformValues, classOf[MapZipWith] -> CometMapZipWith, classOf[CreateMap] -> CometCreateMap) - base ++ sparkVersionSpecificMapExpressions - } + private[comet] val mapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + baseMapExpressions ++ sparkVersionSpecificMapExpressions private[comet] val structExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( @@ -212,10 +210,8 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[XxHash64] -> CometXxHash64, classOf[Sha1] -> CometSha1) - private[comet] val stringExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = { - // Explicit type ascription on `base`: Scala 2.13 cannot infer the existential key type - // when `++` is applied directly to a `Map(...)` literal. - val base: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( + private[comet] val baseStringExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + Map( classOf[Ascii] -> CometScalarFunction("ascii"), classOf[BitLength] -> CometBitLength, classOf[Chr] -> CometScalarFunction("char"), @@ -267,8 +263,8 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[TryToNumber] -> CometTryToNumber, classOf[Mask] -> CometMask, classOf[Empty2Null] -> CometEmpty2Null) - base ++ sparkVersionSpecificStringExpressions - } + private[comet] val stringExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + baseStringExpressions ++ sparkVersionSpecificStringExpressions private val bitwiseExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( classOf[BitwiseAnd] -> CometBitwiseAnd, @@ -355,11 +351,9 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[XPathString] -> CometXPathString, classOf[XPathList] -> CometXPathList) - private[comet] val miscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = { - // TODO PromotePrecision - // Explicit type ascription on `base`: Scala 2.13 cannot infer the existential key type - // when `++` is applied directly to a `Map(...)` literal. - val base: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map( + // TODO PromotePrecision + private[comet] val baseMiscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + Map( classOf[Alias] -> CometAlias, classOf[ApplyFunctionExpression] -> CometApplyFunctionExpression, classOf[AttributeReference] -> CometAttributeReference, @@ -381,18 +375,36 @@ object QueryPlanSerde extends Logging with CometExprShim with CometTypeShim { classOf[TryEval] -> CometTryEval, classOf[UnscaledValue] -> CometUnscaledValue, classOf[Uuid] -> CometUuid) - base ++ sparkVersionSpecificMiscExpressions - } + private[comet] val miscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = + baseMiscExpressions ++ sparkVersionSpecificMiscExpressions /** * Mapping of Spark expression class to Comet expression handler. */ + // Every serde group, in merge order. A class must appear in one group only, which the + // registration suite checks, so the order never decides which serde wins. + private[comet] val serdeGroups + : Seq[(String, Map[Class[_ <: Expression], CometExpressionSerde[_]])] = + Seq( + "math" -> mathExpressions, + "hash" -> hashExpressions, + "string" -> stringExpressions, + "conditional" -> conditionalExpressions, + "map" -> mapExpressions, + "predicate" -> predicateExpressions, + "struct" -> structExpressions, + "bitwise" -> bitwiseExpressions, + "misc" -> miscExpressions, + "array" -> arrayExpressions, + "temporal" -> temporalExpressions, + "conversion" -> conversionExpressions, + "url" -> urlExpressions, + "json" -> jsonExpressions, + "csv" -> csvExpressions, + "xpath" -> xpathExpressions) + val exprSerdeMap: Map[Class[_ <: Expression], CometExpressionSerde[_]] = - mathExpressions ++ hashExpressions ++ stringExpressions ++ - conditionalExpressions ++ mapExpressions ++ predicateExpressions ++ - structExpressions ++ bitwiseExpressions ++ miscExpressions ++ arrayExpressions ++ - temporalExpressions ++ conversionExpressions ++ urlExpressions ++ jsonExpressions ++ - csvExpressions ++ xpathExpressions + serdeGroups.map(_._2).reduce(_ ++ _) /** * Mapping of Spark aggregate expression class to Comet expression handler. diff --git a/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala b/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala index 0bfea5cd6ec..5cd3ae549d9 100644 --- a/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala +++ b/spark/src/main/spark-3.5/org/apache/comet/shims/CometExprShim.scala @@ -23,7 +23,7 @@ import org.apache.spark.sql.catalyst.expressions._ import org.apache.spark.sql.catalyst.expressions.aggregate.Sum import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometEncode, CometExpressionSerde, CometStringDecode, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometEncode, CometExpressionSerde, CometStringDecode, CometToPrettyString} import org.apache.comet.serde.ExprOuterClass.{BinaryOutputStyle, Expr} /** @@ -39,7 +39,7 @@ trait CometExprShim { : Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[StringDecode] -> CometStringDecode, classOf[Encode] -> CometEncode) def sparkVersionSpecificMathExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = - Map(classOf[WidthBucket] -> CometWidthBucket) + Map.empty def sparkVersionSpecificMiscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[ToPrettyString] -> CometToPrettyString) def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = diff --git a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala index 0e6d3b4b4e6..89663800edc 100644 --- a/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala +++ b/spark/src/main/spark-4.x/org/apache/comet/shims/Spark4xCometExprShim.scala @@ -26,7 +26,7 @@ import org.apache.spark.sql.catalyst.expressions.url.ParseUrlEvaluator import org.apache.comet.CometExplainInfo import org.apache.comet.expressions.CometEvalMode -import org.apache.comet.serde.{CometExpressionSerde, CometMapSort, CometRandStr, CometToPrettyString, CometWidthBucket} +import org.apache.comet.serde.{CometExpressionSerde, CometMapSort, CometRandStr, CometToPrettyString} import org.apache.comet.serde.ExprOuterClass.Expr import org.apache.comet.serde.QueryPlanSerde.exprToProtoInternal @@ -43,7 +43,7 @@ trait Spark4xCometExprShim extends CometExprShim4x { : Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[RandStr] -> CometRandStr) def sparkVersionSpecificMathExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = - Map(classOf[WidthBucket] -> CometWidthBucket) + Map.empty def sparkVersionSpecificMiscExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = Map(classOf[ToPrettyString] -> CometToPrettyString) def sparkVersionSpecificMapExpressions: Map[Class[_ <: Expression], CometExpressionSerde[_]] = diff --git a/spark/src/test/resources/sql-tests/expressions/misc/width_bucket.sql b/spark/src/test/resources/sql-tests/expressions/misc/width_bucket.sql index 7e7375b5b7f..25dc2993450 100644 --- a/spark/src/test/resources/sql-tests/expressions/misc/width_bucket.sql +++ b/spark/src/test/resources/sql-tests/expressions/misc/width_bucket.sql @@ -15,8 +15,6 @@ -- specific language governing permissions and limitations -- under the License. --- MinSparkVersion: 3.5 - statement CREATE TABLE test_wb(v double) USING parquet @@ -35,3 +33,13 @@ SELECT v, width_bucket(v, 0, 10, 1) FROM test_wb -- literal arguments query SELECT width_bucket(5.0, 0, 10, 4), width_bucket(0.0, 0, 10, 4), width_bucket(NULL, 0, 10, 4) + +-- day-time and year-month interval inputs take the same dispatch as doubles +query +SELECT v, width_bucket(make_dt_interval(v), make_dt_interval(0), make_dt_interval(10), 4) FROM test_wb + +query +SELECT v, width_bucket(make_ym_interval(0, CAST(v AS INT)), make_ym_interval(0, 0), make_ym_interval(0, 10), 4) FROM test_wb + +query +SELECT width_bucket(INTERVAL '2' DAY, INTERVAL '0' DAY, INTERVAL '10' DAY, 5), width_bucket(INTERVAL '2' YEAR, INTERVAL '0' YEAR, INTERVAL '10' YEAR, 5), width_bucket(CAST(NULL AS INTERVAL DAY), INTERVAL '0' DAY, INTERVAL '10' DAY, 5) diff --git a/spark/src/test/scala/org/apache/comet/serde/SerdeRegistrationSuite.scala b/spark/src/test/scala/org/apache/comet/serde/SerdeRegistrationSuite.scala new file mode 100644 index 00000000000..6b82b2fcfbb --- /dev/null +++ b/spark/src/test/scala/org/apache/comet/serde/SerdeRegistrationSuite.scala @@ -0,0 +1,65 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one + * or more contributor license agreements. See the NOTICE file + * distributed with this work for additional information + * regarding copyright ownership. The ASF licenses this file + * to you under the Apache License, Version 2.0 (the + * "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, + * software distributed under the License is distributed on an + * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY + * KIND, either express or implied. See the License for the + * specific language governing permissions and limitations + * under the License. + */ + +package org.apache.comet.serde + +import org.scalatest.funsuite.AnyFunSuite + +class SerdeRegistrationSuite extends AnyFunSuite { + + // A version shim only adds serdes for classes the shared map cannot name on every Spark + // version. A key present on both sides is either a stale duplicate or a silent override of + // the shared serde; a version that needs a different serde for a shared class should move + // that class out of the shared map instead of shadowing it here. + test("version shims register only classes the shared serde maps do not") { + import QueryPlanSerde._ + val overlaps = Seq( + "math" -> (baseMathExpressions, sparkVersionSpecificMathExpressions), + "map" -> (baseMapExpressions, sparkVersionSpecificMapExpressions), + "string" -> (baseStringExpressions, sparkVersionSpecificStringExpressions), + "misc" -> (baseMiscExpressions, sparkVersionSpecificMiscExpressions)) + .flatMap { case (group, (base, shim)) => + base.keySet.intersect(shim.keySet).map(cls => s"$group: ${cls.getSimpleName}") + } + assert(overlaps.isEmpty, s"shim entries shadow shared serdes: ${overlaps.mkString(", ")}") + } + + // The combined map is built by merging the groups in order, so a class registered in two + // groups would silently take the later serde. Every group must own its classes alone. + test("no expression class is registered in more than one serde group") { + val owners = QueryPlanSerde.serdeGroups + .flatMap { case (name, group) => group.keys.map(cls => cls -> name) } + .groupBy(_._1) + .collect { + case (cls, entries) if entries.size > 1 => + s"${cls.getSimpleName}: ${entries.map(_._2).mkString(", ")}" + } + assert(owners.isEmpty, s"classes registered in several groups: ${owners.mkString("; ")}") + } + + test("every serde group entry reaches the combined map unchanged") { + for ((_, group) <- QueryPlanSerde.serdeGroups; (cls, serde) <- group) { + assert(QueryPlanSerde.exprSerdeMap.get(cls).exists(_ eq serde), cls.getSimpleName) + } + val total = QueryPlanSerde.serdeGroups.map(_._2.size).sum + assert( + QueryPlanSerde.exprSerdeMap.size == total, + s"${QueryPlanSerde.exprSerdeMap.size} != $total") + } +}