sunchao commented on code in PR #4744:
URL: https://github.com/apache/datafusion-comet/pull/4744#discussion_r4124859420
##########
native/core/src/execution/planner.rs:
##########
@@ -3707,6 +3727,137 @@ impl PhysicalPlanner {
}
}
+ fn create_high_order_function_expr(
+ &self,
+ expr: &HigherOrderFunc,
+ input_schema: SchemaRef,
+ ) -> Result<Arc<dyn PhysicalExpr>, ExecutionError> {
+ let udf = create_comet_hof_func(&expr.func_name,
&self.session_ctx.state())?;
+
+ // 1. Plan value args.
+ let value_args: Vec<Arc<dyn PhysicalExpr>> = expr
+ .value_args
+ .iter()
+ .map(|e| self.create_expr(e, Arc::clone(&input_schema)))
+ .collect::<Result<_, _>>()?;
+
+ // 2. Resolve lambda param field types via the UDF (mirrors runtime).
+ let param_fields = Self::resolve_lambda_param_fields(
+ &udf,
+ &expr.func_name,
+ &value_args,
+ expr.lambdas.len(),
+ input_schema.as_ref(),
+ )?;
+
+ // 3. Plan lambdas with resolved param fields.
+ let lambdas: Vec<Arc<dyn PhysicalExpr>> = expr
+ .lambdas
+ .iter()
+ .zip(¶m_fields)
+ .map(|(l, fields)| self.create_lambda_expr(l, &input_schema,
fields))
+ .collect::<Result<_, _>>()?;
+
+ // 4. NOTE: assumes value args precede lambdas (holds for
array_filter).
+ let mut args = value_args;
+ args.extend(lambdas);
+
+ Ok(Arc::new(HigherOrderFunctionExpr::try_new_with_schema(
Review Comment:
[P2] Avoid copying captured arrays once per lambda element
Could we prevent deep replication of captured arrays before selecting this
path by default? For `filter(a, x -> x >= 0 AND size(b) > 0)`, DataFusion's
`evaluate_single_list_lambda` broadcasts captured `b` with `take_arrays`,
copying its entire child buffer once for each element of `a`. With one row
containing 4,096 integers in each array, a focused native allocation probe
measured 67,705,058 bytes allocated from 33,200 bytes of input. The repeated
integer payload alone is 64 MiB, although this predicate only needs `b`'s
length. The cost grows with both array lengths and batch row count.
Spark retains the outer-row array reference, and the base implementation
dispatches this filter to Spark's evaluator. I confirmed that Catalyst keeps
`size(b)` inside the lambda and that the current Comet query takes the native
path with dispatch disabled. Please avoid expanding complex captures this way,
or retain the existing JVM dispatcher for affected shapes until native captures
can represent them efficiently. An allocation benchmark with a captured array
would cover this regression; the current capture benchmark uses a scalar
integer. These allocation numbers are component measurements, not
complete-query memory or timing results.
##########
native/core/src/execution/lambda.rs:
##########
@@ -0,0 +1,375 @@
+// 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.
+
+//! Helpers for planning DataFusion higher-order functions (HOFs) coming
+//! from Spark.
+//!
+//! The planner needs three things that don't belong in `planner.rs`:
+//! 1. A stack of *lambda scopes* so nested `NamedLambdaVariable`s resolve
+//! by Spark `exprId` (immune to name shadowing / column collisions).
+//! 2. A stack of scopes, popped on the `Ok`/`Err` paths of `with_scope`.
+//! Not unwind-safe: a panic unwinds through the whole planner and is
+//! caught at the JNI boundary (`try_unwrap_or_throw`), which tears down
+//! the planner and its scope stack together, so no stale scope survives.
+//! 3. A tiny `PhysicalExpr` wrapper that keeps *unused* lambda parameters
+//! visible in `children()` so `LambdaExpr::new`'s projection compaction
+//! stays consistent with the runtime batch layout.
+
+use std::cell::RefCell;
+use std::collections::HashMap;
+
+use arrow::datatypes::FieldRef;
+use datafusion::common::{DataFusionError, Result};
+
+use arrow::array::{Array, BooleanArray, BooleanBuilder};
+use std::fmt::{Display, Formatter};
+use std::hash::{Hash, Hasher};
+use std::sync::Arc;
+
+use arrow::datatypes::{DataType, Schema};
+use arrow::record_batch::RecordBatch;
+use datafusion::logical_expr::Operator;
+use datafusion::physical_expr::expressions::BinaryExpr;
+use datafusion::physical_expr::PhysicalExpr;
+use datafusion::physical_plan::ColumnarValue;
+
+/// Maps Spark `exprId` -> (column index in the extended body schema, field).
+pub(crate) type LambdaScope = HashMap<i64, (usize, FieldRef)>;
+
+/// A stack of lambda variable scopes, innermost last.
+/// Planning is single-threaded per planner, so `RefCell` is sufficient to
manage
+/// the stack of scopes during the recursive planning process.
+#[derive(Default)]
+pub(crate) struct LambdaScopes {
+ stack: RefCell<Vec<LambdaScope>>,
+}
+
+impl LambdaScopes {
+ /// Resolve a lambda variable by Spark `exprId`, searching innermost
+ /// scope first.
+ pub(crate) fn resolve_variable(&self, expr_id: i64) -> Option<(usize,
FieldRef)> {
+ self.stack
+ .borrow()
+ .iter()
+ .rev()
+ .find_map(|s| s.get(&expr_id).cloned())
+ }
+
+ /// Push `scope`, run `f`, pop unconditionally. The pop happens on both
+ /// the `Ok` and `Err` paths — this replaces the earlier RAII guard.
+ pub(crate) fn with_scope<T, E>(
+ &self,
+ scope: LambdaScope,
+ f: impl FnOnce() -> Result<T, E>,
+ ) -> Result<T, E> {
+ self.stack.borrow_mut().push(scope);
+ let out = f();
+ self.stack.borrow_mut().pop();
+ out
+ }
+}
+
+/// An expression adapter that short-circuits evaluation on empty batches (0
rows),
+/// returning an empty array without evaluating the inner expression.
+/// This prevents runtime errors (such as division by zero in ANSI mode) when
+/// Spark's short-circuit semantics guarantee the lambda predicate is never
invoked
+/// for empty arrays.
+#[derive(Debug)]
+pub struct EmptyBatchGuardExpr {
+ inner: Arc<dyn PhysicalExpr>,
+}
+
+impl PartialEq for EmptyBatchGuardExpr {
+ fn eq(&self, other: &Self) -> bool {
+ self.inner.eq(&other.inner)
+ }
+}
+
+impl Eq for EmptyBatchGuardExpr {}
+
+impl Hash for EmptyBatchGuardExpr {
+ fn hash<H: Hasher>(&self, state: &mut H) {
+ self.inner.dyn_hash(state);
+ }
+}
+
+impl EmptyBatchGuardExpr {
+ pub fn new(inner: Arc<dyn PhysicalExpr>) -> Self {
+ Self { inner }
+ }
+}
+
+impl Display for EmptyBatchGuardExpr {
+ fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
+ write!(f, "EmptyBatchGuard({})", self.inner)
+ }
+}
+
+impl PhysicalExpr for EmptyBatchGuardExpr {
+ fn data_type(&self, input_schema: &Schema) -> Result<DataType> {
+ self.inner.data_type(input_schema)
+ }
+
+ fn nullable(&self, input_schema: &Schema) -> Result<bool> {
+ self.inner.nullable(input_schema)
Review Comment:
[P2] Keep lambda nullability checks from consuming predicate state
Could the lambda wrapper report nullability without evaluating stateful
predicates? With one Parquet row containing `a=[1,2,3]`, this query returns
`[1]` in Spark and JVM dispatch, but `[]` through the native path at this head:
```sql
SELECT filter(a, x ->
(CASE WHEN monotonically_increasing_id() = 0 THEN x ELSE 0 END) > 0)
FROM t;
```
`HigherOrderFunctionExpr` asks the lambda body for its return field during
construction and evaluation. This delegation reaches DataFusion's
`CaseExpr::nullable()`, which evaluates the WHEN condition on a synthetic
one-row batch and advances the same counter later used for real elements. The
conditional guard admits the expression because its result branches are not
fallible. The previous whole-filter JVM dispatcher does not perform these
evaluations. Please keep metadata inspection free of these side effects and add
a native regression with a stateful CASE condition. I reproduced this through
full Spark 4.1.3/Comet execution with JVM dispatch disabled.
##########
spark/src/main/scala/org/apache/comet/serde/CometHighOrderFunction.scala:
##########
@@ -0,0 +1,219 @@
+/*
+ * 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 scala.jdk.CollectionConverters._
+import scala.util.control.NonFatal
+
+import org.apache.spark.sql.catalyst.expressions.{Abs, Add, AssertTrue,
Attribute, CaseWhen, Cast, Coalesce, Divide, ElementAt, Expression,
GetArrayItem, HigherOrderFunction, If, IntegralDivide, LambdaFunction =>
SparkLambdaFunction, Multiply, NamedLambdaVariable => SparkNamedLambdaVariable,
RaiseError, Remainder, Subtract, UnaryMinus}
+import org.apache.spark.sql.internal.SQLConf
+
+import org.apache.comet.CometConf
+import org.apache.comet.serde.CometHighOrderFunction.{containsJvmDispatch,
hasGuardedFallibleBranch, namedLambdaVariable2Proto}
+import org.apache.comet.serde.ExprOuterClass.{HigherOrderFunc, LambdaFunction,
NamedLambdaVariable}
+import org.apache.comet.serde.QueryPlanSerde.{exprToProtoInternal,
serializeDataType}
+
+/**
+ * Generic expression serializer for Spark higher-order functions (e.g.,
`filter`, `transform`).
+ *
+ * This class implements a three-tier execution hierarchy:
+ * 1. '''Native DataFusion execution:''' Produced when
+ * `COMET_EXEC_HIGHER_ORDER_FUNCTION_NATIVE_ENABLED` is enabled and the
lambda structure
+ * meets native runtime constraints. 2. '''JVM codegen dispatch:''' Emits
a `JvmScalarUdf`
+ * fallback via `CometScalaUDF.emitJvmCodegenDispatch` if the native path
cannot be taken,
+ * provided `COMET_SCALA_UDF_CODEGEN_ENABLED` is enabled. 3. '''Vanilla
Spark:''' Final
+ * fallback if neither native DataFusion nor codegen dispatch is
available.
+ *
+ * ===Short-Circuiting and Safety Guarantees===
+ * - '''Boolean short-circuiting (AND / OR):''' Handled natively in Rust via
+ * `ShortCircuitBinaryExpr`. It enforces strict SQL Three-Valued Logic
(3VL) per-element
+ * masking via `evaluate_selection`, ensuring that stateful functions
+ * (`monotonically_increasing_id`, `rand`) and fallible operations (`DIV`,
`abs`,
+ * `element_at`) are never evaluated on skipped elements.
+ * - '''Conditional expressions (CASE WHEN, IF, COALESCE):''' Guarded
branches containing
+ * fallible operations are checked via [[hasGuardedFallibleBranch]] and
safely routed to JVM
+ * codegen dispatch.
+ * - '''Speculative serialization:''' Lambda traversal is wrapped in a
`NonFatal` catch to
+ * decline the native path if eager expression evaluation (e.g.
`CometCast` evaluating literal
+ * arguments) fails during plan generation under ANSI mode.
+ */
+case class CometHighOrderFunction[T <: HigherOrderFunction](name: String)
+ extends CometExpressionSerde[T] {
+
+ def convert(expr: T, inputs: Seq[Attribute], binding: Boolean):
Option[ExprOuterClass.Expr] = {
+ if (!CometConf.COMET_EXEC_HIGHER_ORDER_FUNCTION_NATIVE_ENABLED.get()) {
+ return CometScalaUDF.emitJvmCodegenDispatch(expr, inputs, binding)
+ }
+ val hofProto =
+ try {
+ highOrderFunction2Proto(expr, inputs, binding)
+ } catch {
+ // Speculative serialization traverses the lambda body where certain
expressions
+ // eagerly evaluate literal arguments (e.g., CometCast calling
cast.eval()).
+ // In ANSI mode, guarded branches of conditional expressions (e.g.,
CASE WHEN)
+ // may throw during planning even though they are never reached at
runtime.
+ // Decline the native path cleanly and let execution fall back to JVM
codegen dispatch.
+ case NonFatal(_) => None
+ }
+ hofProto.orElse {
+ CometScalaUDF.emitJvmCodegenDispatch(expr, inputs, binding)
+ }
+ }
+
+ private def highOrderFunction2Proto(
+ expr: T,
+ inputs: Seq[Attribute],
+ binding: Boolean): Option[ExprOuterClass.Expr] = {
+ val argumentsProto = expr.arguments.map(exprToProtoInternal(_, inputs,
binding))
+ val functionsProto = expr.functions
+ .map {
+ case slf: SparkLambdaFunction =>
+ if (hasGuardedFallibleBranch(slf.function)) {
+ return None
+ }
+ exprToProtoInternal(slf.function, inputs, binding)
+ .flatMap { bodyProto =>
+ if (containsJvmDispatch(bodyProto)) {
+ return None
+ }
+ val namedLambdaVariablesProto = slf.arguments
+ .map {
+ case arg: SparkNamedLambdaVariable =>
+ namedLambdaVariable2Proto(arg)
+ case _ => None
+ }
+ if (namedLambdaVariablesProto.forall(_.isDefined)) {
+ Some(
+ LambdaFunction
+ .newBuilder()
+ .addAllArgs(namedLambdaVariablesProto.map(_.get).asJava)
+ .setBody(bodyProto)
+ .build())
+ } else {
+ None
+ }
+ }
+ case _ => None
+ }
+ if (functionsProto.forall(_.isDefined) &&
argumentsProto.forall(_.isDefined)) {
+ val hof = HigherOrderFunc
+ .newBuilder()
+ .setFuncName(name)
+ .addAllValueArgs(argumentsProto.map(_.get).asJava)
+ .addAllLambdas(functionsProto.map(_.get).asJava)
+ .build()
+ Some(ExprOuterClass.Expr.newBuilder().setHighOrderFunc(hof).build())
+ } else {
+ None
+ }
+ }
+}
+
+object CometHighOrderFunction {
+
+ def containsJvmDispatch(e: ExprOuterClass.Expr): Boolean =
+ containsJvmDispatch(e.asInstanceOf[com.google.protobuf.Message])
+
+ private def containsJvmDispatch(m: com.google.protobuf.Message): Boolean =
+ m match {
+ case e: ExprOuterClass.Expr if e.hasJvmScalarUdf => true
+ case _ =>
+ m.getAllFields.values().asScala.exists {
+ case v: com.google.protobuf.Message => containsJvmDispatch(v)
+ case vs: java.util.List[_] =>
+ vs.asScala.exists {
+ case v: com.google.protobuf.Message => containsJvmDispatch(v)
+ case _ => false
+ }
+ case _ => false
+ }
+ }
+
+ /**
+ * Checks whether an expression can throw a runtime exception during
evaluation.
+ */
+ private def isFallibleExpr(expr: Expression): Boolean = {
+ val ansi = SQLConf.get.ansiEnabled
+ expr.exists {
+ case _: Divide | _: IntegralDivide | _: Remainder => true
+ case c: Cast if ansi => !c.evalMode.toString.contains("TRY")
+ case _: Add | _: Subtract | _: Multiply | _: UnaryMinus | _: Abs if ansi
=> true
+ case _: GetArrayItem | _: ElementAt => true
+ case _: RaiseError | _: AssertTrue => true
+ case _ => false
+ }
+ }
+
+ /**
+ * Checks whether conditional expressions (CASE WHEN, IF, COALESCE) contain
guarded fallible
+ * branches that require JVM codegen fallback.
+ *
+ * Note: AND and OR are handled natively with strict per-element masking in
Rust (via
+ * StrictBooleanExpr) and do not require fallback.
+ */
+ def hasGuardedFallibleBranch(expr: Expression): Boolean = {
+ expr.exists {
+ // CASE WHEN: THEN and ELSE branches
+ case CaseWhen(branches, elseValue) =>
+ branches.map(_._2).exists(isFallibleExpr) ||
elseValue.exists(isFallibleExpr)
+
+ // IF: true and false branches
+ case If(_, trueValue, falseValue) =>
+ isFallibleExpr(trueValue) || isFallibleExpr(falseValue)
+
+ // COALESCE: tail arguments
+ case Coalesce(children) if children.length > 1 =>
+ children.tail.exists(isFallibleExpr)
Review Comment:
[P2] Preserve single evaluation of nondeterministic COALESCE children
Could native admission also account for nondeterministic nonfinal COALESCE
arguments? With one Parquet row `a=[1,2,3,4]`, this returns `[2,4]` in Spark
and JVM dispatch, but `[4]` with native HOF execution and dispatch disabled:
```sql
SELECT filter(a, x ->
coalesce(IF(rand(42L) < 0.6, x, CAST(NULL AS INT)), -1) = x)
FROM t;
```
`CometCoalesce` serializes the child separately as `WHEN IS NOT NULL(child)`
and `THEN child`. The second random-expression instance runs only on the
selected subset, so its draws correspond to different elements. This remains
wrong even when CASE metadata evaluation is suppressed, so it is separate from
the nullability issue. Plain IF with the same seeded predicate and
deterministic COALESCE controls both agree with Spark. Please ensure each child
is evaluated once per element, or retain whole-filter JVM dispatch for these
shapes until the native lowering preserves that contract. This mismatch is
reproduced through full current-head Spark 4.1.3/Comet execution.
--
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]