andygrove commented on code in PR #4459: URL: https://github.com/apache/datafusion-comet/pull/4459#discussion_r4188642176
########## spark/src/main/scala/org/apache/comet/udf/CometRustUdfRegistry.scala: ########## @@ -0,0 +1,53 @@ +/* + * 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.udf + +import java.util.concurrent.ConcurrentHashMap + +import org.apache.spark.sql.types.DataType + +/** Metadata for a registered Rust UDF. */ +case class RustUdfMetadata( + libraryPath: String, + inputTypes: Seq[DataType], + returnType: DataType, + deterministic: Boolean) + +/** + * Driver-side registry of Rust UDFs. Looked up by `QueryPlanSerde` to recognize names that should + * be emitted as `RustUdfCall` instead of attempted as JVM-evaluated `ScalaUDF`s. + */ +class CometRustUdfRegistry { + private val byName = new ConcurrentHashMap[String, RustUdfMetadata]() + + /** Register or replace metadata for a name. */ + def register(name: String, meta: RustUdfMetadata): Unit = + byName.put(name, meta) + + /** Return metadata for a name, if registered. */ + def get(name: String): Option[RustUdfMetadata] = + Option(byName.get(name)) +} + +object CometRustUdfRegistry { + + /** Process-wide singleton. */ + lazy val instance: CometRustUdfRegistry = new CometRustUdfRegistry Review Comment: Fixed, and the singleton is gone rather than scoped. `register` now installs the function in the session's own function registry, as `spark.udf.register` does, and every call resolves to a `NativeUdfCall` expression that carries the registration, so planning needs no registry of its own (94197907e, which follows the registration in #6697). A test checks that another session neither sees the registration nor gets its own UDF of that name answered natively. ########## spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala: ########## @@ -53,8 +54,45 @@ import org.apache.comet.udf.codegen.CometScalaUDFCodegen */ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { - override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = - emitJvmCodegenDispatch(expr, inputs, binding) + override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = { + // First check if this udfName is a registered Rust UDF -- those get emitted as RustUdfCall + // and dispatched to the loaded cdylib rather than the JVM codegen dispatcher. + expr.udfName.flatMap(CometRustUdfRegistry.instance.get) match { Review Comment: Fixed. A native UDF's calls now resolve to a `NativeUdfCall` expression with a serde of its own, so nothing matches on `udfName` any more: an ordinary Scala UDF registered under the same name builds an ordinary `ScalaUDF` and stays on the JVM path. The `ignore`d test from this thread is enabled and passes. Registering a builder in the session's function registry directly, instead of going through `spark.udf.register`, is what made this possible, since `functions.udf` wraps whatever it is handed. The first version of the fix (7d5b36375) kept a `ScalaUDF` with a marker function, and 94197907e replaced it with the expression #6697 uses. ########## spark/src/main/scala/org/apache/comet/serde/CometScalaUDF.scala: ########## @@ -55,8 +57,74 @@ import org.apache.comet.udf.codegen.CometScalaUDFCodegen */ object CometScalaUDF extends CometExpressionSerde[ScalaUDF] { - override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = - emitJvmCodegenDispatch(expr, inputs, binding) + override def convert(expr: ScalaUDF, inputs: Seq[Attribute], binding: Boolean): Option[Expr] = { + // A registered native UDF is emitted as NativeScalarUdf and dispatched to the loaded shared + // library rather than to the JVM codegen dispatcher. + // + // The match is on the name alone, which is not enough to identify one: Spark sets `udfName` for + // every `spark.udf.register` call, and the registry is process-wide and keyed by bare name, so + // an ordinary Scala UDF sharing the name is currently answered out of the native library. The + // registration would have to be identified some other way to fix that, since the closure Spark + // holds for the catalog stub is one `functions.udf` wrapped rather than the one Comet passed + // in. See https://github.com/apache/datafusion-comet/issues/5295. + expr.udfName.flatMap(CometNativeUdfRegistry.get) match { + case Some(meta) => + emitNativeScalarUdf(expr, meta, inputs, binding) + case None => + emitJvmCodegenDispatch(expr, inputs, binding) + } + } + + private def emitNativeScalarUdf( + expr: ScalaUDF, + meta: NativeUdfMetadata, + inputs: Seq[Attribute], + binding: Boolean): Option[Expr] = { + val name = expr.udfName.get + checkArgumentTypes(name, expr, meta) + val argProtos = expr.children.map(c => exprToProtoInternal(c, inputs, binding)) + if (argProtos.exists(_.isEmpty)) { + withFallbackReason(expr, "one or more native UDF arguments are not supported") + return None + } + val returnTypeProto = serializeDataType(meta.returnType).getOrElse { + withFallbackReason(expr, s"return type ${meta.returnType} not serializable") + return None + } + val callBuilder = ExprOuterClass.NativeScalarUdf + .newBuilder() + .setName(name) + .setLibraryPath(meta.libraryPath) + .setReturnType(returnTypeProto) + .setDeterministic(expr.deterministic) + argProtos.foreach(a => callBuilder.addArgs(a.get)) + Some(ExprOuterClass.Expr.newBuilder().setNativeScalarUdf(callBuilder.build()).build()) + } + + /** + * Refuse a call whose argument types differ from the ones the UDF was registered with. + * + * The catalog stub Comet installs is untyped, so Spark inserts no casts for it and a call + * reaches this point with whatever types its arguments happen to have. Converting them here + * would be a semantic choice Spark never made, so the call is refused instead, naming both + * signatures. Nullability is disregarded because it does not change the values a UDF receives. + * + * This throws rather than falling back: the stub cannot evaluate the UDF on the JVM, so a + * fallback would only fail later with a less useful message. + */ + private def checkArgumentTypes(name: String, expr: ScalaUDF, meta: NativeUdfMetadata): Unit = { + val actual = expr.children.map(_.dataType) + val matches = actual.length == meta.inputTypes.length && + actual.zip(meta.inputTypes).forall { case (a, d) => deepNullable(a) == deepNullable(d) } Review Comment: Fixed. 74d7560d6 first made Comet's own check ignore metadata, and 94197907e then moved the check into Spark's analyzer: `NativeUdfCall` declares the registered types through `ExpectsInputTypes`, and Spark's comparison already ignores nullability and field metadata, so Comet's copy is gone. The test passes a struct whose field carries a comment, built with `struct(col.as(name, metadata))`. -- 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]
