From 8311fccac27d6284fa0bef185221871c25a9123d Mon Sep 17 00:00:00 2001 From: Dustin Smith Date: Sat, 12 Sep 2026 10:52:00 +0700 Subject: [PATCH] chore: drop the redundant width_bucket shim registrations and guard serde uniqueness The Spark 3.5 and 4.x shims still registered WidthBucket over the same key the shared math map registers, a no-op that nothing would have caught had the two drifted. Remove both entries, hoist the shared maps out of their merge expressions, assemble the combined map from a named group list, and add a suite asserting that shims only add classes, that no class lives in two groups, and that every entry reaches the combined map unchanged. Closes #4485 --- .github/workflows/pr_build_linux.yml | 1 + .github/workflows/pr_build_macos.yml | 1 + .../apache/comet/serde/QueryPlanSerde.scala | 72 +++++++++++-------- .../apache/comet/shims/CometExprShim.scala | 4 +- .../comet/shims/Spark4xCometExprShim.scala | 4 +- .../comet/serde/SerdeRegistrationSuite.scala | 65 +++++++++++++++++ 6 files changed, 113 insertions(+), 34 deletions(-) create mode 100644 spark/src/test/scala/org/apache/comet/serde/SerdeRegistrationSuite.scala diff --git a/.github/workflows/pr_build_linux.yml b/.github/workflows/pr_build_linux.yml index a019e9a721f..4bdc78c0ae0 100644 --- a/.github/workflows/pr_build_linux.yml +++ b/.github/workflows/pr_build_linux.yml @@ -481,6 +481,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 66a1b62ba3d..e3559036b09 100644 --- a/.github/workflows/pr_build_macos.yml +++ b/.github/workflows/pr_build_macos.yml @@ -235,6 +235,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/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/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") + } +}