Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/pr_build_linux.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions .github/workflows/pr_build_macos.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
72 changes: 42 additions & 30 deletions spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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(
Expand All @@ -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"),
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand All @@ -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.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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}

/**
Expand All @@ -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[_]] =
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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[_]] =
Expand Down
Original file line number Diff line number Diff line change
@@ -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")
}
}