Skip to content
Draft
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
89 changes: 68 additions & 21 deletions spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala
Original file line number Diff line number Diff line change
Expand Up @@ -53,9 +53,61 @@ import org.apache.comet.udf.codegen.CometScalaUDFCodegen
*/
object CometScalaUDF extends CometExpressionSerde[ScalaUDF] {

override def getSupportLevel(expr: ScalaUDF): SupportLevel = dispatchSupportLevel(expr)

override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] =
emitJvmCodegenDispatch(expr, inputs, binding)

/**
* Bind `expr` the way [[emitJvmCodegenDispatch]] will.
*
* `RuntimeReplaceable` expressions (e.g. Spark 4's `StructsToJson`) have a `doGenCode` that
* always throws "Cannot generate code for expression". Catalyst's `ReplaceExpressions` rule
* normally rewrites them to their `replacement` form before codegen runs. Comet's serde
* sometimes works with the pre-rewrite form (via shim reconstruction) for matching purposes, so
* unwrap to the replacement here before binding so the kernel compiles.
*
* Binding is against only the `AttributeReference`s the tree actually reads, so ordinals align
* with the data args shipped alongside the closure. Those attributes are returned too, since
* [[emitJvmCodegenDispatch]] needs them in the same order to build the data args.
*/
private def bindForDispatch(expr: Expression): (Expression, Seq[AttributeReference]) = {
val target = expr match {
case rr: RuntimeReplaceable => rr.replacement
case other => other
}
val attrs = target.collect { case a: AttributeReference => a }.distinct
(BindReferences.bindReference(target, AttributeSeq(attrs)), attrs)
}

/**
* `SupportLevel` for a serde whose only path is the codegen dispatcher. `Compatible` when the
* dispatcher will accept the expression, `Unsupported` (with the same reason
* [[emitJvmCodegenDispatch]] would have tagged) when it will not.
*
* Reporting this from `getSupportLevel` rather than discovering it inside `convert` keeps the
* serde invariant intact and lets `exprToProtoInternal` handle the decline on its normal
* `Unsupported` path. Behaviour is unchanged: `CometCodegenDispatch` does not mix in
* `CodegenDispatchFallback`, so an `Unsupported` result still tags the reason and falls the
* operator back to Spark, exactly as the `convert`-side decline did.
*
* This is a pure predicate -- it records no fallback reason of its own, because the
* `Unsupported` arm of `exprToProtoInternal` already tags the notes returned here.
*/
def dispatchSupportLevel(expr: Expression): SupportLevel = {
val exprName = CometExplainInfo.exprDisplayName(expr)
if (!CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.get()) {
return Unsupported(
Some(
s"$exprName: ${CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key}=false; expression has " +
"no native path so the plan falls back to Spark"))
}
CometBatchKernelCodegen.canHandle(bindForDispatch(expr)._1) match {
case Some(reason) => Unsupported(Some(s"$exprName: $reason"))
case None => Compatible()
}
}

/**
* Bind `expr`, closure-serialize it, and emit a `JvmScalarUdf` proto routed through
* [[CometScalaUDFCodegen]] so that native execution evaluates the expression inside the
Expand All @@ -66,6 +118,10 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] {
* via [[CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED]] or when
* [[CometBatchKernelCodegen.canHandle]] refuses the expression tree. Callers should treat
* `None` as a clean Spark-fallback signal.
*
* Serdes that gate on [[dispatchSupportLevel]] have already screened both of those conditions,
* so for them the checks below are a cheap re-verification. They are kept because several
* serdes call this directly from `convert` without gating first.
*/
def emitJvmCodegenDispatch(
expr: Expression,
Expand All @@ -80,20 +136,7 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] {
return None
}

// `RuntimeReplaceable` expressions (e.g. Spark 4's `StructsToJson`) have a `doGenCode` that
// always throws "Cannot generate code for expression". Catalyst's `ReplaceExpressions` rule
// normally rewrites them to their `replacement` form before codegen runs. Comet's serde
// sometimes works with the pre-rewrite form (via shim reconstruction) for matching purposes,
// so unwrap to the replacement here before binding so the kernel compiles.
val target = expr match {
case rr: RuntimeReplaceable => rr.replacement
case other => other
}

// Bind against only the AttributeReferences the tree actually reads, so ordinals align with
// the data args we ship.
val attrs = target.collect { case a: AttributeReference => a }.distinct
val boundExpr = BindReferences.bindReference(target, AttributeSeq(attrs))
val (boundExpr, attrs) = bindForDispatch(expr)

// Gate at plan time. Surface the reason via withFallbackReason rather than crashing Janino
// at execute.
Expand Down Expand Up @@ -159,18 +202,22 @@ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] {

/**
* Convenience base for serdes that route a non-ScalaUDF Spark expression through the codegen
* dispatcher. Delegates `convert` to [[CometScalaUDF.emitJvmCodegenDispatch]] and marks the
* expression `Compatible()` because the dispatcher runs Spark's own `doGenCode` inside the
* kernel: behavior matches Spark exactly when [[CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED]] is
* enabled, and the operator falls back to Spark cleanly when it is not.
* dispatcher. Delegates `convert` to [[CometScalaUDF.emitJvmCodegenDispatch]], and reports
* [[CometScalaUDF.dispatchSupportLevel]] so that the two conditions the dispatcher can refuse on
* -- the global [[CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED]] flag being off, and
* `CometBatchKernelCodegen.canHandle` rejecting the bound tree -- are reported from
* `getSupportLevel` instead of surfacing as a `Compatible` serde that then declines in `convert`.
*
* When the dispatcher will run the expression this is `Compatible()`: behavior then matches Spark
* exactly, because the kernel runs Spark's own `doGenCode`.
*/
class CometCodegenDispatch[T <: Expression] extends CometExpressionSerde[T] {
override def getSupportLevel(expr: T): SupportLevel = Compatible()
override def getSupportLevel(expr: T): SupportLevel = CometScalaUDF.dispatchSupportLevel(expr)
// Intentionally no getCompatibleNotes override: the docs generator emits compat notes under
// a heading that promises "no additional configuration required". The dispatcher flag is a
// global concern documented elsewhere; tagging each expression here would contradict the
// heading. When the flag is off, `convert` returns None with a clear fallback reason that
// shows up in EXPLAIN, which is the right place for that signal.
// heading. When the flag is off, `getSupportLevel` reports Unsupported with a clear reason
// that shows up in EXPLAIN, which is the right place for that signal.
override def convert(expr: T, inputs: Seq[Attribute], binding: Boolean): Option[Expr] =
CometScalaUDF.emitJvmCodegenDispatch(expr, inputs, binding)
}
33 changes: 32 additions & 1 deletion spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@ import org.apache.arrow.vector._
import org.apache.spark.{SparkConf, SparkEnv, TaskContext}
import org.apache.spark.sql.CometTestBase
import org.apache.spark.sql.api.java.UDF1
import org.apache.spark.sql.catalyst.expressions.{BoundReference, CreateArray, CreateMap, CreateNamedStruct, Expression, Literal, MapConcat}
import org.apache.spark.sql.catalyst.expressions.{BoundReference, CreateArray, CreateMap, CreateNamedStruct, Expression, FindInSet, Literal, MapConcat}
import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types._
Expand All @@ -34,6 +34,7 @@ import org.apache.spark.unsafe.types.UTF8String
import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus
import org.apache.comet.codegen.CometBatchKernelCodegen
import org.apache.comet.codegen.CometBatchKernelCodegen.ArrowColumnSpec
import org.apache.comet.serde.{CometFindInSet, CometScalaUDF, Compatible, Unsupported}
import org.apache.comet.udf.codegen.CometScalaUDFCodegen
import org.apache.comet.vector.CometVector

Expand Down Expand Up @@ -1754,6 +1755,36 @@ class CometCodegenSuite
}
}
}

// The two conditions the dispatcher can refuse on are reported from `getSupportLevel` rather
// than discovered inside `convert`, so a dispatch-only serde never claims `Compatible` and then
// declines. `find_in_set` stands in for the ~62 plain `CometCodegenDispatch` serdes.
private def findInSet = FindInSet(Literal("b"), Literal("a,b,c"))

test("dispatch-only serdes report Compatible when the dispatcher will run the expression") {
assert(CometScalaUDF.dispatchSupportLevel(findInSet).isInstanceOf[Compatible])
assert(CometFindInSet.getSupportLevel(findInSet).isInstanceOf[Compatible])
}

test("dispatch-only serdes report Unsupported when the dispatcher is disabled") {
withSQLConf(CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key -> "false") {
Seq(
CometScalaUDF.dispatchSupportLevel(findInSet),
CometFindInSet.getSupportLevel(findInSet)).foreach {
case Unsupported(Some(reason)) =>
assert(reason.contains(CometConf.COMET_SCALA_UDF_CODEGEN_ENABLED.key))
case other => fail(s"expected Unsupported, got $other")
}
}
}

test("dispatch-only serdes report Unsupported when canHandle refuses the tree") {
// NullType is outside CometBatchKernelCodegen.isSupportedDataType.
CometScalaUDF.dispatchSupportLevel(Literal(null, NullType)) match {
case Unsupported(Some(reason)) => assert(reason.contains("unsupported output type"))
case other => fail(s"expected Unsupported, got $other")
}
}
}

/**
Expand Down
Loading