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


##########
spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala:
##########
@@ -904,16 +904,9 @@ class CometIcebergWriteActionSuite
     }
   }
 
-  test("native acceleration: ReplaceData (CoW MERGE) falls back (MergeRowsExec 
not Comet)") {
-    // TODO(comet-merge-rows): native MERGE engagement requires a Comet 
equivalent of Iceberg's
-    // `MergeRowsExec` (the per-row dispatch operator that assigns 
__row_operation codes from
-    // MATCHED/NOT MATCHED clauses). Without it, `MergeRowsExec` stays JVM, 
the upstream chain
-    // breaks Comet-native partway, and `requiresNativeChildren=true` declines 
the
-    // `IcebergWriteExec -> CometIcebergWriteExec` conversion. Until that 
lands, MERGE
-    // falls back to the JVM two-op path -- this test pins that contract. 
Native `MergeRowsExec`
-    // is being added in https://github.com/apache/datafusion-comet/pull/5318; 
when that lands
-    // this test will start failing and needs to flip to 
`assertNativeWriteEngages`.
+  test("native acceleration: ReplaceData (CoW MERGE) honors the versioned 
native contract") {
     assumeNativeAcceleration()
+    assume(isSpark35Plus, "MergeRowsExec requires Spark 3.5+")

Review Comment:
   The test this replaced also pinned the fallback on 3.4, and this `assume` 
now cancels it there. 3.4 registers nothing, so this MERGE stays on the JVM 
just as it does on 4.1+. I dropped the `assume` and took the fallback branch 
for `isSpark41Plus || !isSpark35Plus`, and it passes on 3.4. Could it do that 
instead of cancelling?



##########
spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala:
##########
@@ -933,6 +927,230 @@ class CometIcebergWriteActionSuite
           |WHEN NOT MATCHED THEN INSERT (id, region, amount) VALUES (s.id, 
s.region, s.amount)
           |""".stripMargin)
       }
+
+      withSQLConf(CometConf.COMET_EXEC_MERGE_ROWS_ENABLED.key -> "true") {
+        if (isSpark41Plus) {
+          assertNativeWriteDoesNotEngage("native_cow_merge", Seq(1, 2, 
3))(runMerge())
+        } else {
+          val snapshot = withNativeEnabled {
+            captureWrite("native_cow_merge")(runMerge())
+          }
+          assert(
+            snapshot.snapshotDelta == 1L,
+            s"expected exactly one Iceberg snapshot from native MERGE, got 
${snapshot.snapshotDelta}")
+          val mergeExecs = snapshot.plans.flatMap { plan =>
+            collectWithSubqueries(plan) { case e: CometMergeRowsExec => e }
+          }
+          assert(
+            mergeExecs.nonEmpty,
+            "expected Iceberg MERGE to execute through CometMergeRowsExec. 
Plans:\n" +
+              snapshot.plans.mkString("\n--\n"))
+          val nativeWrites = snapshot.plans.flatMap { plan =>
+            collectWithSubqueries(plan) { case e: CometIcebergWriteExec => e }
+          }
+          assert(
+            nativeWrites.nonEmpty,
+            "expected Iceberg MERGE to feed CometIcebergWriteExec. Plans:\n" +
+              snapshot.plans.mkString("\n--\n"))
+          assertRows("native_cow_merge", expectedIds = Seq(1, 2, 3))
+        }
+      }
+
+      val updated = spark
+        .sql(s"SELECT amount FROM $catalog.$ns.native_cow_merge WHERE id = 2")
+        .collect()
+      assert(updated.length == 1 && updated.head.getDouble(0) == 200.0)
+    }
+  }
+
+  test("native MergeRows matches Spark on partitioned Iceberg copy-on-write 
and merge-on-read") {
+    assumeNativeAcceleration()
+    assume(
+      isSpark35Plus && !isSpark41Plus,
+      "native MergeRows is registered only on Spark 3.5 and 4.0")
+    withIcebergCatalog { warehouseDir =>
+      spark
+        .range(0, 20000, 1, 8)
+        .selectExpr(
+          "CAST(id AS INT) AS id",
+          "concat('r', CAST(id % 8 AS STRING)) AS region",
+          "CAST(id AS DOUBLE) AS amount")
+        .createOrReplaceTempView("merge_parity_seed")
+      spark
+        .range(10000, 30000, 1, 8)
+        .selectExpr(
+          "CAST(id AS INT) AS id",
+          "concat('s', CAST(id % 8 AS STRING)) AS region",
+          "CASE WHEN id % 11 = 0 THEN CAST(NULL AS DOUBLE) ELSE CAST(id * 2 AS 
DOUBLE) END AS amount")
+        .createOrReplaceTempView("merge_parity_source")
+
+      def merge(table: String): Unit = {
+        spark.sql(s"""
+          |MERGE INTO $catalog.$ns.$table t
+          |USING merge_parity_source s
+          |ON t.id = s.id
+          |WHEN MATCHED AND CAST(NULL AS BOOLEAN) THEN UPDATE SET t.amount = 
-999.0
+          |WHEN MATCHED AND s.amount IS NULL THEN DELETE
+          |WHEN MATCHED AND s.id % 5 = 0 THEN DELETE
+          |WHEN MATCHED THEN UPDATE SET t.region = s.region, t.amount = 
s.amount + 0.5
+          |WHEN NOT MATCHED AND s.id IS NOT NULL THEN
+          |  INSERT (id, region, amount) VALUES (s.id, s.region, s.amount)
+          |WHEN NOT MATCHED BY SOURCE AND t.id % 7 = 0 THEN DELETE
+          |WHEN NOT MATCHED BY SOURCE THEN UPDATE SET t.amount = t.amount + 1.0
+          |""".stripMargin)
+      }
+
+      def rows(table: String): Seq[Row] =
+        spark
+          .sql(s"SELECT id, region, amount FROM $catalog.$ns.$table ORDER BY 
id")
+          .collect()
+          .toSeq
+
+      Seq("copy-on-write" -> "cow", "merge-on-read" -> "mor").foreach { case 
(mode, suffix) =>
+        val nativeTable = s"merge_parity_${suffix}_native"
+        val sparkTable = s"merge_parity_${suffix}_spark"
+        Seq(nativeTable, sparkTable).foreach { table =>
+          createTable(
+            warehouseDir,
+            table,
+            partitionSpec = "PARTITIONED BY (region)",
+            properties = Some(s"'format-version'='2', 
'write.merge.mode'='$mode'"))
+        }
+
+        withSQLConf(
+          CometConf.COMET_ENABLED.key -> "false",
+          "spark.sql.adaptive.coalescePartitions.enabled" -> "false",
+          "spark.sql.shuffle.partitions" -> "8") {
+          spark.sql(
+            s"INSERT INTO $catalog.$ns.$nativeTable SELECT id, region, amount 
FROM merge_parity_seed")
+          spark.sql(
+            s"INSERT INTO $catalog.$ns.$sparkTable SELECT id, region, amount 
FROM merge_parity_seed")
+        }
+
+        val inputFiles = spark
+          .sql(s"SELECT count(*) FROM $catalog.$ns.$nativeTable.data_files")
+          .collect()
+          .head
+          .getLong(0)
+        assert(inputFiles > 1L, s"expected a multi-file target, got 
$inputFiles file(s)")
+
+        var snapshot: Option[WriteSnapshot] = None
+        withSQLConf(
+          CometConf.COMET_EXEC_MERGE_ROWS_ENABLED.key -> "true",

Review Comment:
   Could this also run with `spark.sql.adaptive.enabled=false`? The existing 
AQE-off MERGE test runs with native MergeRows disabled, so nothing covers the 
native operator when the whole plan is converted at once. I ran this test with 
AQE off on 3.5 and it matches, so this is only about keeping that covered.



##########
spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala:
##########
@@ -933,6 +927,230 @@ class CometIcebergWriteActionSuite
           |WHEN NOT MATCHED THEN INSERT (id, region, amount) VALUES (s.id, 
s.region, s.amount)
           |""".stripMargin)
       }
+
+      withSQLConf(CometConf.COMET_EXEC_MERGE_ROWS_ENABLED.key -> "true") {
+        if (isSpark41Plus) {
+          assertNativeWriteDoesNotEngage("native_cow_merge", Seq(1, 2, 
3))(runMerge())
+        } else {
+          val snapshot = withNativeEnabled {
+            captureWrite("native_cow_merge")(runMerge())
+          }
+          assert(
+            snapshot.snapshotDelta == 1L,
+            s"expected exactly one Iceberg snapshot from native MERGE, got 
${snapshot.snapshotDelta}")
+          val mergeExecs = snapshot.plans.flatMap { plan =>
+            collectWithSubqueries(plan) { case e: CometMergeRowsExec => e }
+          }
+          assert(
+            mergeExecs.nonEmpty,
+            "expected Iceberg MERGE to execute through CometMergeRowsExec. 
Plans:\n" +
+              snapshot.plans.mkString("\n--\n"))
+          val nativeWrites = snapshot.plans.flatMap { plan =>
+            collectWithSubqueries(plan) { case e: CometIcebergWriteExec => e }
+          }
+          assert(
+            nativeWrites.nonEmpty,
+            "expected Iceberg MERGE to feed CometIcebergWriteExec. Plans:\n" +
+              snapshot.plans.mkString("\n--\n"))
+          assertRows("native_cow_merge", expectedIds = Seq(1, 2, 3))
+        }
+      }
+
+      val updated = spark
+        .sql(s"SELECT amount FROM $catalog.$ns.native_cow_merge WHERE id = 2")
+        .collect()
+      assert(updated.length == 1 && updated.head.getDouble(0) == 200.0)
+    }
+  }
+
+  test("native MergeRows matches Spark on partitioned Iceberg copy-on-write 
and merge-on-read") {
+    assumeNativeAcceleration()
+    assume(
+      isSpark35Plus && !isSpark41Plus,
+      "native MergeRows is registered only on Spark 3.5 and 4.0")
+    withIcebergCatalog { warehouseDir =>
+      spark
+        .range(0, 20000, 1, 8)
+        .selectExpr(
+          "CAST(id AS INT) AS id",
+          "concat('r', CAST(id % 8 AS STRING)) AS region",
+          "CAST(id AS DOUBLE) AS amount")
+        .createOrReplaceTempView("merge_parity_seed")
+      spark
+        .range(10000, 30000, 1, 8)
+        .selectExpr(
+          "CAST(id AS INT) AS id",
+          "concat('s', CAST(id % 8 AS STRING)) AS region",
+          "CASE WHEN id % 11 = 0 THEN CAST(NULL AS DOUBLE) ELSE CAST(id * 2 AS 
DOUBLE) END AS amount")
+        .createOrReplaceTempView("merge_parity_source")
+
+      def merge(table: String): Unit = {
+        spark.sql(s"""
+          |MERGE INTO $catalog.$ns.$table t
+          |USING merge_parity_source s
+          |ON t.id = s.id
+          |WHEN MATCHED AND CAST(NULL AS BOOLEAN) THEN UPDATE SET t.amount = 
-999.0
+          |WHEN MATCHED AND s.amount IS NULL THEN DELETE
+          |WHEN MATCHED AND s.id % 5 = 0 THEN DELETE
+          |WHEN MATCHED THEN UPDATE SET t.region = s.region, t.amount = 
s.amount + 0.5
+          |WHEN NOT MATCHED AND s.id IS NOT NULL THEN
+          |  INSERT (id, region, amount) VALUES (s.id, s.region, s.amount)
+          |WHEN NOT MATCHED BY SOURCE AND t.id % 7 = 0 THEN DELETE
+          |WHEN NOT MATCHED BY SOURCE THEN UPDATE SET t.amount = t.amount + 1.0
+          |""".stripMargin)
+      }
+
+      def rows(table: String): Seq[Row] =
+        spark
+          .sql(s"SELECT id, region, amount FROM $catalog.$ns.$table ORDER BY 
id")
+          .collect()
+          .toSeq
+
+      Seq("copy-on-write" -> "cow", "merge-on-read" -> "mor").foreach { case 
(mode, suffix) =>
+        val nativeTable = s"merge_parity_${suffix}_native"
+        val sparkTable = s"merge_parity_${suffix}_spark"
+        Seq(nativeTable, sparkTable).foreach { table =>
+          createTable(
+            warehouseDir,
+            table,
+            partitionSpec = "PARTITIONED BY (region)",
+            properties = Some(s"'format-version'='2', 
'write.merge.mode'='$mode'"))
+        }
+
+        withSQLConf(
+          CometConf.COMET_ENABLED.key -> "false",
+          "spark.sql.adaptive.coalescePartitions.enabled" -> "false",
+          "spark.sql.shuffle.partitions" -> "8") {
+          spark.sql(
+            s"INSERT INTO $catalog.$ns.$nativeTable SELECT id, region, amount 
FROM merge_parity_seed")
+          spark.sql(
+            s"INSERT INTO $catalog.$ns.$sparkTable SELECT id, region, amount 
FROM merge_parity_seed")
+        }
+
+        val inputFiles = spark
+          .sql(s"SELECT count(*) FROM $catalog.$ns.$nativeTable.data_files")
+          .collect()
+          .head
+          .getLong(0)
+        assert(inputFiles > 1L, s"expected a multi-file target, got 
$inputFiles file(s)")
+
+        var snapshot: Option[WriteSnapshot] = None
+        withSQLConf(
+          CometConf.COMET_EXEC_MERGE_ROWS_ENABLED.key -> "true",
+          "spark.sql.adaptive.coalescePartitions.enabled" -> "false",
+          "spark.sql.shuffle.partitions" -> "8") {
+          snapshot = Some(withNativeEnabled {
+            captureWrite(nativeTable)(merge(nativeTable))
+          })
+        }
+        val writeSnapshot =
+          snapshot.getOrElse(fail(s"$mode MERGE did not produce a write 
snapshot"))
+        assert(
+          writeSnapshot.snapshotDelta == 1L,
+          s"expected exactly 1 new Iceberg snapshot for $mode, got 
${writeSnapshot.snapshotDelta}. Plans:\n" +
+            writeSnapshot.plans.mkString("\n--\n"))
+        val commits = writeSnapshot.plans.flatMap { plan =>
+          collectWithSubqueries(plan) { case c: IcebergCommitExec => c }
+        }
+        if (mode == "copy-on-write") {
+          assert(
+            commits.nonEmpty,
+            s"expected >= 1 IcebergCommitExec for $mode, got 0. Plans:\n" +
+              writeSnapshot.plans.mkString("\n--\n"))
+        } else {
+          assert(
+            commits.isEmpty,
+            s"merge-on-read should stay on Iceberg WriteDelta, got 
${commits.size} IcebergCommitExec. Plans:\n" +
+              writeSnapshot.plans.mkString("\n--\n"))
+        }
+        val mergeExecs = writeSnapshot.plans.flatMap { plan =>
+          collectWithSubqueries(plan) { case e: CometMergeRowsExec => e }
+        }
+        assert(
+          mergeExecs.nonEmpty,
+          s"expected CometMergeRowsExec for $mode. Plans:\n" +
+            writeSnapshot.plans.mkString("\n--\n"))
+
+        val nativeWrites = writeSnapshot.plans.flatMap { plan =>
+          collectWithSubqueries(plan) { case e: CometIcebergWriteExec => e }
+        }
+        if (mode == "copy-on-write") {
+          assert(
+            nativeWrites.nonEmpty,
+            "copy-on-write MERGE should feed the native Iceberg writer on this 
partitioned plan")
+        } else {
+          assert(
+            nativeWrites.isEmpty,
+            "merge-on-read uses Iceberg WriteDelta and must not engage 
CometIcebergWriteExec")
+        }
+
+        withSQLConf(
+          CometConf.COMET_ENABLED.key -> "false",
+          "spark.sql.adaptive.coalescePartitions.enabled" -> "false",
+          "spark.sql.shuffle.partitions" -> "8") {
+          merge(sparkTable)
+        }
+        assert(rows(nativeTable) == rows(sparkTable), s"$mode MERGE result 
differs from Spark")
+      }
+    }
+  }
+
+  test("native MergeRows Iceberg cardinality violation matches Spark") {
+    assumeNativeAcceleration()
+    assume(
+      isSpark35Plus && !isSpark41Plus,
+      "native MergeRows is registered only on Spark 3.5 and 4.0")
+    withIcebergCatalog { warehouseDir =>
+      val nativeTable = "merge_cardinality_native"
+      val sparkTable = "merge_cardinality_spark"
+      Seq(nativeTable, sparkTable).foreach { table =>
+        createTable(
+          warehouseDir,
+          table,
+          partitionSpec = "PARTITIONED BY (region)",
+          properties = Some("'format-version'='2', 
'write.merge.mode'='copy-on-write'"))
+      }
+      withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+        Seq(nativeTable, sparkTable).foreach { table =>
+          spark.sql(
+            s"INSERT INTO $catalog.$ns.$table VALUES " +
+              "(42, 'r2', 42.0), (43, 'r3', 43.0)")
+        }
+      }
+
+      def duplicateMerge(table: String): Unit = {
+        spark.sql(s"""
+          |MERGE INTO $catalog.$ns.$table t
+          |USING (
+          |  SELECT 42 AS id, 'a' AS region, 420.0 AS amount
+          |  UNION ALL
+          |  SELECT 42 AS id, 'b' AS region, 421.0 AS amount
+          |) s
+          |ON t.id = s.id
+          |WHEN MATCHED THEN UPDATE SET t.amount = s.amount
+          |""".stripMargin)
+      }
+
+      val nativeError = intercept[Exception] {
+        withSQLConf(
+          CometConf.COMET_EXEC_MERGE_ROWS_ENABLED.key -> "true",
+          CometConf.COMET_ICEBERG_NATIVE_WRITE_ENABLED.key -> "false") {

Review Comment:
   This MERGE never reaches `CometMergeRowsExec`. The two-row source is 
broadcast, so the join stays on the JVM in the same stage as the copy-on-write 
target scan, and `MergeRowsExec` stays Spark's. I captured the failed plan with 
`captureFailedPlans`. It has Spark's `MergeRowsExec` and no 
`CometMergeRowsExec`, and the test still passes with 
`spark.comet.exec.mergeRows.enabled=false`. Setting 
`spark.sql.autoBroadcastJoinThreshold=-1` here, as 
`CometMergeRowsNativeSuiteBase` does, makes it engage, and the error still 
surfaces as `MERGE_CARDINALITY_VIOLATION`. Could it set that and assert on 
`CometMergeRowsExec` in the plans from `captureFailedPlans`, the way the parity 
test does with `captureWrite`?



##########
docs/source/user-guide/latest/compatibility/operators.md:
##########
@@ -113,6 +113,45 @@ runs natively; it is controlled by 
`spark.comet.exec.windowGroupLimit.enabled` (
   Scalar `FLOAT` and `DOUBLE` keys are normalized and match Spark; see
   [floating-point ordering](./floating-point.md).
 
+## MERGE INTO (MergeRowsExec)
+
+Spark `MergeRowsExec` appears as `CometMergeRows` when native execution is 
enabled.

Review Comment:
   I meant the operator reference in 
`docs/source/user-guide/latest/operators.md`, the page that calls itself the 
complete reference for each Spark physical operator. It still has no 
`MergeRowsExec` row. Could you add one under Writes, marked ⚠️ as opt-in on 3.5 
and 4.0, with a link to this section?



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