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(&param_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]

Reply via email to