andygrove commented on code in PR #5318:
URL: https://github.com/apache/datafusion-comet/pull/5318#discussion_r4178476399


##########
spark/src/main/scala/org/apache/comet/CometConf.scala:
##########
@@ -277,6 +277,13 @@ object CometConf extends ShimCometConf {
     createExecEnabledConfig("emptyRelation", defaultValue = true)
   val COMET_EXEC_SAMPLE_ENABLED: ConfigEntry[Boolean] =
     createExecEnabledConfig("sample", defaultValue = true)
+  val COMET_EXEC_MERGE_ROWS_ENABLED: ConfigEntry[Boolean] =
+    createExecEnabledConfig(
+      "mergeRows",
+      defaultValue = false,
+      notes = Some(
+        "Ignored on Spark 4.1 and later, where MergeRowsExec remains on Spark 
so V2 writers " +

Review Comment:
   Since 4.1 is the default build, this flag does nothing there, and nothing 
tracks the work to change that. Could you open an issue for native 
`MergeRowsExec` on 4.1+ that covers the `MergeSummary` contract, and link it 
from this note and the compatibility page? Could the note also say the flag 
only takes effect on Spark 3.5 and 4.0, since 3.4 ignores it too? The comment 
at `CometExecRule.scala:118` has the same gap.



##########
native/core/src/execution/operators/merge_rows.rs:
##########
@@ -0,0 +1,1486 @@
+// 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.
+
+use arrow::array::{Array, ArrayRef, BooleanArray, Int64Array, RecordBatch};
+use arrow::compute::kernels::boolean::{and, and_not, not};
+use arrow::compute::{filter_record_batch, prep_null_mask_filter};
+use arrow::datatypes::{DataType, SchemaRef};
+use datafusion::common::tree_node::TreeNodeRecursion;
+use datafusion::common::utils::memory::estimate_memory_size;
+use datafusion::common::{DataFusionError, HashSet, ScalarValue};
+use datafusion::execution::memory_pool::{MemoryConsumer, MemoryReservation};
+use datafusion::logical_expr::ColumnarValue;
+use datafusion::physical_expr::{EquivalenceProperties, PhysicalExpr};
+use datafusion::physical_plan::execution_plan::{Boundedness, EmissionType};
+use datafusion::physical_plan::metrics::{BaselineMetrics, 
ExecutionPlanMetricsSet, MetricsSet};
+use datafusion::{
+    execution::TaskContext,
+    physical_plan::{
+        apply_expression_roots, ChildrenPropertiesMode, DisplayAs, 
DisplayFormatType,
+        ExecutionPlan, Partitioning, PlanProperties, RecordBatchStream, 
ReplaceChildrenOptions,
+        SendableRecordBatchStream,
+    },
+};
+use datafusion_comet_common::{cast_and_stamp_schema, SparkError};
+use futures::{Stream, StreamExt};
+use std::{
+    pin::Pin,
+    sync::Arc,
+    task::{Context, Poll},
+};
+
+/// A MergeRows instruction: condition plus zero (Discard), one (Keep), or two 
(Split)
+/// output row projections.
+#[derive(Debug, Clone)]
+pub struct MergeInstructionExec {
+    pub condition: Arc<dyn PhysicalExpr>,
+    pub outputs: Vec<Vec<Arc<dyn PhysicalExpr>>>,
+}
+
+#[derive(Debug)]
+struct MergeConfig {
+    is_source_row_present: Arc<dyn PhysicalExpr>,
+    is_target_row_present: Arc<dyn PhysicalExpr>,
+    matched_instructions: Vec<MergeInstructionExec>,
+    not_matched_instructions: Vec<MergeInstructionExec>,
+    not_matched_by_source_instructions: Vec<MergeInstructionExec>,
+    row_id_ordinal: Option<usize>,
+}
+
+impl MergeConfig {
+    fn validate(
+        &self,
+        child: &Arc<dyn ExecutionPlan>,
+        output_schema: &SchemaRef,
+    ) -> Result<(), DataFusionError> {
+        if let Some(ordinal) = self.row_id_ordinal {
+            let child_schema = child.schema();
+            let child_fields = child_schema.fields().len();
+            if ordinal >= child_fields {
+                return Err(DataFusionError::Internal(format!(
+                    "MergeRows: row id ordinal {ordinal} is out of range for a 
child with \
+                     {child_fields} columns"
+                )));
+            }
+            let data_type = child_schema.field(ordinal).data_type();
+            if data_type != &DataType::Int64 {
+                return Err(DataFusionError::Internal(format!(
+                    "MergeRows: row id column at ordinal {ordinal} must be 
Int64, got {data_type}"
+                )));
+            }
+        }
+
+        let output_width = output_schema.fields().len();
+        for (group, instructions) in [
+            ("matched", &self.matched_instructions),
+            ("not matched", &self.not_matched_instructions),
+            (
+                "not matched by source",
+                &self.not_matched_by_source_instructions,
+            ),
+        ] {
+            for (instruction_index, instruction) in 
instructions.iter().enumerate() {
+                if instruction.outputs.len() > 2 {
+                    return Err(DataFusionError::Internal(format!(
+                        "MergeRows: {group} instruction {instruction_index} 
has {} output rows; expected at most 2",
+                        instruction.outputs.len()
+                    )));
+                }
+                for (output_index, output) in 
instruction.outputs.iter().enumerate() {
+                    if output.len() != output_width {
+                        return Err(DataFusionError::Internal(format!(
+                            "MergeRows: {group} instruction 
{instruction_index} output {output_index} has {} expressions; expected 
{output_width}",
+                            output.len()
+                        )));
+                    }
+                }
+            }
+        }
+        Ok(())
+    }
+}
+
+#[derive(Debug)]
+pub struct MergeRowsExec {
+    config: Arc<MergeConfig>,
+    child: Arc<dyn ExecutionPlan>,
+    schema: SchemaRef,
+    cache: Arc<PlanProperties>,
+    metrics: ExecutionPlanMetricsSet,
+}
+
+impl MergeRowsExec {
+    #[allow(clippy::too_many_arguments)]
+    pub fn try_new(
+        is_source_row_present: Arc<dyn PhysicalExpr>,
+        is_target_row_present: Arc<dyn PhysicalExpr>,
+        matched_instructions: Vec<MergeInstructionExec>,
+        not_matched_instructions: Vec<MergeInstructionExec>,
+        not_matched_by_source_instructions: Vec<MergeInstructionExec>,
+        row_id_ordinal: Option<usize>,
+        child: Arc<dyn ExecutionPlan>,
+        schema: SchemaRef,
+    ) -> Result<Self, DataFusionError> {
+        let config = Arc::new(MergeConfig {
+            is_source_row_present,
+            is_target_row_present,
+            matched_instructions,
+            not_matched_instructions,
+            not_matched_by_source_instructions,
+            row_id_ordinal,
+        });
+        config.validate(&child, &schema)?;
+
+        let cache = Arc::new(PlanProperties::new(
+            EquivalenceProperties::new(Arc::clone(&schema)),
+            Partitioning::UnknownPartitioning(1),
+            EmissionType::Incremental,
+            Boundedness::Bounded,
+        ));
+
+        Ok(Self {
+            config,
+            child,
+            schema,
+            cache,
+            metrics: ExecutionPlanMetricsSet::new(),
+        })
+    }
+}
+
+impl DisplayAs for MergeRowsExec {
+    fn fmt_as(&self, t: DisplayFormatType, f: &mut std::fmt::Formatter) -> 
std::fmt::Result {
+        match t {
+            DisplayFormatType::Default | DisplayFormatType::Verbose => {
+                write!(f, "CometMergeRowsExec")
+            }
+            DisplayFormatType::TreeRender => unimplemented!(),
+        }
+    }
+}
+
+impl ExecutionPlan for MergeRowsExec {
+    fn schema(&self) -> SchemaRef {
+        Arc::clone(&self.schema)
+    }
+
+    fn children(&self) -> Vec<&Arc<dyn ExecutionPlan>> {
+        vec![&self.child]
+    }
+
+    fn replace_children(
+        self: Arc<Self>,
+        children: Vec<Arc<dyn ExecutionPlan>>,
+        _options: ReplaceChildrenOptions,
+    ) -> datafusion::common::Result<Arc<dyn ExecutionPlan>> {
+        let [child] = children.as_slice() else {
+            return Err(DataFusionError::Internal(format!(
+                "MergeRows expects exactly one child, got {}",
+                children.len()
+            )));
+        };
+        let child = Arc::clone(child);
+        self.config.validate(&child, &self.schema)?;
+        Ok(Arc::new(MergeRowsExec {
+            config: Arc::clone(&self.config),
+            child,
+            schema: Arc::clone(&self.schema),
+            cache: Arc::clone(&self.cache),
+            metrics: self.metrics.clone(),
+        }))
+    }
+
+    fn apply_expressions(
+        &self,
+        f: &mut dyn FnMut(&Arc<dyn PhysicalExpr>) -> 
datafusion::common::Result<TreeNodeRecursion>,
+    ) -> datafusion::common::Result<TreeNodeRecursion> {
+        let instructions = self
+            .config
+            .matched_instructions
+            .iter()
+            .chain(self.config.not_matched_instructions.iter())
+            .chain(self.config.not_matched_by_source_instructions.iter());
+        let instruction_expressions = instructions.flat_map(|instruction| {
+            std::iter::once(&instruction.condition)
+                .chain(instruction.outputs.iter().flat_map(|output| 
output.iter()))
+        });
+
+        apply_expression_roots(
+            [
+                &self.config.is_source_row_present,
+                &self.config.is_target_row_present,
+            ]
+            .into_iter()
+            .chain(instruction_expressions),
+            f,
+        )
+    }
+
+    fn with_new_children(
+        self: Arc<Self>,
+        children: Vec<Arc<dyn ExecutionPlan>>,
+    ) -> datafusion::common::Result<Arc<dyn ExecutionPlan>> {
+        self.replace_children(
+            children,
+            ReplaceChildrenOptions::new(ChildrenPropertiesMode::Recompute),
+        )
+    }
+
+    fn execute(
+        &self,
+        partition: usize,
+        context: Arc<TaskContext>,
+    ) -> datafusion::common::Result<SendableRecordBatchStream> {
+        let reservation = self.config.row_id_ordinal.map(|_| {
+            MemoryConsumer::new(format!("CometMergeRowsExec[{partition}]"))
+                .register(&context.runtime_env().memory_pool)
+        });
+        let child_stream = self.child.execute(partition, 
Arc::clone(&context))?;
+        Ok(Box::pin(MergeRowsStream {
+            config: Arc::clone(&self.config),
+            child_stream,
+            schema: Arc::clone(&self.schema),
+            seen: HashSet::new(),
+            reservation,
+            baseline: BaselineMetrics::new(&self.metrics, partition),
+        }))
+    }
+
+    fn properties(&self) -> &Arc<PlanProperties> {
+        &self.cache
+    }
+
+    fn metrics(&self) -> Option<MetricsSet> {
+        Some(self.metrics.clone_inner())
+    }
+
+    fn name(&self) -> &str {
+        "CometMergeRowsExec"
+    }
+}
+
+pub struct MergeRowsStream {
+    config: Arc<MergeConfig>,
+    child_stream: SendableRecordBatchStream,
+    schema: SchemaRef,
+    // Partition-scoped so duplicate matches across Arrow batches are still 
detected.
+    seen: HashSet<i64>,

Review Comment:
   Spark's `BitmapCardinalityValidator` keeps matched row ids in a 
`Roaring64Bitmap`, and this uses a `HashSet<i64>`. I measured both on 5M row 
ids. The set takes 75 MB, about 15 bytes per id. A roaring treemap takes 0.6 MB 
when the ids are dense, as in a co-partitioned join, and 5 to 17 MB when they 
are a shuffled share of 8 to 200 target partitions. Because this reservation 
can't spill, a MERGE with many matched rows per task can fail with 
`ResourcesExhausted` where Spark succeeds. `roaring` is already compiled in 
through iceberg-rust, and `RoaringBitmap::statistics()` reports capacity bytes 
that could size the reservation once per batch. Could you switch to 
`RoaringTreemap` here? If you'd rather keep that out of this PR, could you open 
an issue for it and add a line about the memory cost to the compatibility page?



##########
spark/src/test/scala/org/apache/comet/exec/CometMergeRowsNativeSuiteBase.scala:
##########
@@ -0,0 +1,499 @@
+/*
+ * 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.exec
+
+import org.apache.spark.{CometListenerBusUtils, SparkConf}
+import org.apache.spark.sql.CometTestBase
+import org.apache.spark.sql.catalyst.expressions.SubqueryExpression
+import org.apache.spark.sql.comet.CometMergeRowsExec
+import 
org.apache.spark.sql.connector.catalog.InMemoryRowLevelOperationTableCatalog
+import org.apache.spark.sql.execution.{QueryExecution, ScalarSubquery}
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
+import org.apache.spark.sql.util.QueryExecutionListener
+
+import org.apache.comet.CometConf
+import org.apache.comet.CometSparkSessionExtensions.isSpark35Plus
+
+/**
+ * `CometMergeRowsExec` converts Spark's `MergeRowsExec` and the `MergeRows` 
logical node that
+ * feeds it, both defined in Spark core (`execution.datasources.v2` / 
`catalyst.plans.logical`),
+ * not in any connector module. Spark plans a `MergeRowsExec` for MERGE INTO 
against any
+ * `SupportsRowLevelOperations` V2 table using group-based (copy-on-write) 
planning, independent
+ * of which connector implements the table.
+ *
+ * This suite pins that contract against Spark's own 
`InMemoryRowLevelOperationTableCatalog` test
+ * catalog rather than Iceberg. `InMemoryRowLevelOperationTable` selects its 
write shape via the
+ * `supports-deltas` table property: unset (default `false`) plans through 
group-based
+ * `MergeRows`; `supports-deltas=true` plans a JVM `WriteDelta` whose child is 
*also* a
+ * `MergeRowsExec` (and, with `split-updates=true`, emits `Split` for 
update-as-delete+insert).
+ * Comet's bottom-up conversion makes that child eligible for native execution 
under the JVM
+ * write, so both shapes are exercised here (`deltaMergeCase`) and checked 
against a pure-Spark
+ * baseline. See `CometIcebergWriteActionSuite` for MERGE INTO coverage 
against real Iceberg
+ * tables.
+ */
+abstract class CometMergeRowsNativeSuiteBase extends CometTestBase with 
AdaptiveSparkPlanHelper {
+
+  private val catalog = "generic_rowlevel"
+
+  override protected def sparkConf: SparkConf = {
+    super.sparkConf
+      .set(s"spark.sql.catalog.$catalog", 
classOf[InMemoryRowLevelOperationTableCatalog].getName)
+      // The test catalog's BatchScan is not a Comet-native scan. A broadcast 
join over these tiny
+      // fixtures therefore leaves MergeRowsExec on the JVM instead of 
exercising this suite's
+      // operator. Force the supported shuffle boundary while keeping AQE 
enabled.
+      .set("spark.sql.autoBroadcastJoinThreshold", "-1")
+      .set("spark.sql.adaptive.autoBroadcastJoinThreshold", "-1")
+      .set("spark.sql.shuffle.partitions", "4")
+  }
+
+  private def assumeMerge(): Unit = assume(isSpark35Plus, "MergeRowsExec 
requires Spark 3.5+")
+
+  test("MERGE all row-routing groups engage CometMergeRowsExec and match 
Spark") {
+    assumeMerge()
+    val target = s"$catalog.default.rowlevel_target"
+    val source = s"$catalog.default.rowlevel_source"
+
+    def resetTables(): Unit = {
+      sql(s"DROP TABLE IF EXISTS $target")
+      sql(s"DROP TABLE IF EXISTS $source")
+      sql(s"CREATE TABLE $target (id INT, region STRING, amount DOUBLE) USING 
parquet")
+      sql(s"CREATE TABLE $source (id INT, region STRING, amount DOUBLE) USING 
parquet")
+      sql(
+        s"INSERT INTO $target VALUES " +
+          (0 until 20).map(i => s"($i, 'r${i % 3}', ${i * 1.5})").mkString(", 
"))
+      sql(
+        s"INSERT INTO $source VALUES " +
+          (10 until 30).map(i => s"($i, 's${i % 3}', ${i * 2.0})").mkString(", 
"))
+    }
+
+    // Exercise all three MergeRows instruction groups in one end-to-end 
query: 10-19 are
+    // MATCHED, 20-29 are NOT MATCHED, and 0-9 are NOT MATCHED BY SOURCE. The 
predicate on the
+    // target-only clause also leaves 5-9 to Spark's generated catch-all Keep 
instruction.
+    val mergeSql =
+      s"""MERGE INTO $target t USING $source s ON t.id = s.id
+         |WHEN MATCHED THEN UPDATE SET t.amount = s.amount, t.region = s.region
+         |WHEN NOT MATCHED THEN INSERT (id, region, amount) VALUES (s.id, 
s.region, s.amount)
+         |WHEN NOT MATCHED BY SOURCE AND t.id < 5 THEN UPDATE SET t.amount = 
t.amount + 1000.0
+         |""".stripMargin
+
+    val captured = scala.collection.mutable.ArrayBuffer[QueryExecution]()
+    val listener = new QueryExecutionListener {
+      override def onSuccess(funcName: String, qe: QueryExecution, durationNs: 
Long): Unit =
+        captured += qe
+      override def onFailure(funcName: String, qe: QueryExecution, exception: 
Exception): Unit =
+        ()
+    }
+    spark.listenerManager.register(listener)
+    try {
+      resetTables()
+      captured.clear()
+      withSQLConf(
+        CometConf.COMET_ENABLED.key -> "true",

Review Comment:
   `CometTestBase` already enables Comet, so the `CometConf.COMET_ENABLED.key 
-> "true"` entries here, at lines 171 and 248, and in the 4.1+ 
`CometMergeRowsSuite` at line 74 can go. This was in my first review too.



-- 
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