comphead commented on code in PR #6547:
URL: https://github.com/apache/datafusion-comet/pull/6547#discussion_r4170344602


##########
spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala:
##########
@@ -608,6 +610,89 @@ case class CometExecRule(session: SparkSession)
     }
   }
 
+  /**
+   * AQE re-plans around a materialized stage by reusing the physical node 
linked to it, so a
+   * native operator that shares its logical node with a shuffle stage (the 
final aggregate of a
+   * two-phase aggregate) keeps the native plan it got while that input was a 
bare exchange, read
+   * through a plain `Scan`. Once the input is a sink that reads the shuffle 
directly, its
+   * `ShuffleScan` takes the place of the stale leaf. The leaf is patched in 
place because
+   * converting the node again from `originalPlan` would drop the stage's 
logical link that AQE
+   * relies on and re-run serde on a node that is already planned.
+   */
+  private def refreshStaleShuffleScans(op: SparkPlan): SparkPlan = op match {
+    // These build their own `Scan` over the child rather than embedding the 
child's plan.
+    case _: CometNativeWriteExec | _: CometIcebergWriteExec | _: 
CometWriteFilesExec => op
+    case native: CometNativeExec if native.children.nonEmpty =>
+      refreshedNativeOp(native) match {
+        case Some(newOp) =>
+          val refreshed = native.withRefreshedNativeOp(newOp)
+          // An operator that does not hold its native plan as a field cannot 
take a new one.
+          if (refreshed.nativeOp eq newOp) refreshed else op
+        case None => op
+      }
+    case _ => op
+  }
+
+  /**
+   * The native plan of `native` with each `Scan` leaf whose input is now a 
`ShuffleScan` of the
+   * same field types replaced by that `ShuffleScan`, or None if there is no 
such leaf or the plan
+   * children cannot be matched to the leaves.
+   */
+  private def refreshedNativeOp(native: CometNativeExec): Option[Operator] = {
+    val children = native.children.collect { case child: CometNativeExec => 
child }
+    // Only a sink that reads a shuffle directly, or a native child that may 
hold one, can feed
+    // a `ShuffleScan`.
+    val mayFeedShuffleScan = children.exists {
+      case sink: CometSinkPlaceHolder => sink.nativeOp.hasShuffleScan
+      case _ => true
+    }
+    if (children.length != native.children.length || !mayFeedShuffleScan) 
return None
+    val leaves = nativeLeaves(native.nativeOp)
+    if (!leaves.exists(_.hasScan)) return None
+
+    // Each plan child feeds a run of leaves, in order: a sink feeds one, and 
a native child
+    // feeds the leaves of its own native plan.
+    val current = children.flatMap {
+      case sink: CometSinkPlaceHolder => Seq(sink.nativeOp)
+      case child => nativeLeaves(child.nativeOp)
+    }
+    def isStale(leaf: Operator, input: Operator): Boolean = leaf.hasScan && 
input.hasShuffleScan
+    val stale = leaves.zip(current).filter { case (leaf, input) => 
isStale(leaf, input) }
+    val isRefreshable = current.length == leaves.length && stale.nonEmpty &&
+      stale.forall { case (leaf, input) =>
+        leaf.getScan.getFieldsList == input.getShuffleScan.getFieldsList
+      }
+    if (isRefreshable) {
+      val newLeaves = leaves.zip(current).map { case (leaf, input) =>
+        if (isStale(leaf, input)) input else leaf
+      }
+      Some(withLeaves(native.nativeOp, newLeaves))
+    } else {
+      None
+    }
+  }
+
+  private def nativeLeaves(op: Operator): Seq[Operator] =
+    if (op.getChildrenCount == 0) Seq(op)
+    else op.getChildrenList.asScala.toSeq.flatMap(nativeLeaves)

Review Comment:
   This leaf walk now exists in four places. 
`CometNativeExec.findShuffleScanIndices` in `operators.scala` already visits 
the same leaves in the same order. `CometExecRuleSuite.nativeLeaves` and the 
local `leaves` in `CometNativeShuffleSuite.shuffleReadingBlocks` copy this body.
   
   Could one shared `nativeLeaves` serve all four? `findShuffleScanIndices` 
would then be the positions of `ShuffleScan` among the `Scan` and `ShuffleScan` 
leaves. I compared the two on a few hand-built trees, including a `NativeScan` 
leaf that `nativeLeaves` counts and the index walk skips, and they agreed. That 
check was outside the repo and not run against the real suites.
   
   `shuffleReadingBlocks.inputs` is also a second version of 
`foreachUntilCometInput`, the walk that fixes the runtime input order. Using 
the real one would make that test check the pairing the runtime uses.



##########
spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala:
##########
@@ -608,6 +610,89 @@ case class CometExecRule(session: SparkSession)
     }
   }
 
+  /**
+   * AQE re-plans around a materialized stage by reusing the physical node 
linked to it, so a
+   * native operator that shares its logical node with a shuffle stage (the 
final aggregate of a
+   * two-phase aggregate) keeps the native plan it got while that input was a 
bare exchange, read
+   * through a plain `Scan`. Once the input is a sink that reads the shuffle 
directly, its
+   * `ShuffleScan` takes the place of the stale leaf. The leaf is patched in 
place because
+   * converting the node again from `originalPlan` would drop the stage's 
logical link that AQE
+   * relies on and re-run serde on a node that is already planned.
+   */
+  private def refreshStaleShuffleScans(op: SparkPlan): SparkPlan = op match {
+    // These build their own `Scan` over the child rather than embedding the 
child's plan.
+    case _: CometNativeWriteExec | _: CometIcebergWriteExec | _: 
CometWriteFilesExec => op
+    case native: CometNativeExec if native.children.nonEmpty =>
+      refreshedNativeOp(native) match {
+        case Some(newOp) =>
+          val refreshed = native.withRefreshedNativeOp(newOp)
+          // An operator that does not hold its native plan as a field cannot 
take a new one.
+          if (refreshed.nativeOp eq newOp) refreshed else op
+        case None => op
+      }
+    case _ => op
+  }
+
+  /**
+   * The native plan of `native` with each `Scan` leaf whose input is now a 
`ShuffleScan` of the
+   * same field types replaced by that `ShuffleScan`, or None if there is no 
such leaf or the plan
+   * children cannot be matched to the leaves.
+   */
+  private def refreshedNativeOp(native: CometNativeExec): Option[Operator] = {
+    val children = native.children.collect { case child: CometNativeExec => 
child }
+    // Only a sink that reads a shuffle directly, or a native child that may 
hold one, can feed
+    // a `ShuffleScan`.
+    val mayFeedShuffleScan = children.exists {
+      case sink: CometSinkPlaceHolder => sink.nativeOp.hasShuffleScan
+      case _ => true
+    }
+    if (children.length != native.children.length || !mayFeedShuffleScan) 
return None
+    val leaves = nativeLeaves(native.nativeOp)
+    if (!leaves.exists(_.hasScan)) return None
+
+    // Each plan child feeds a run of leaves, in order: a sink feeds one, and 
a native child
+    // feeds the leaves of its own native plan.
+    val current = children.flatMap {
+      case sink: CometSinkPlaceHolder => Seq(sink.nativeOp)
+      case child => nativeLeaves(child.nativeOp)
+    }
+    def isStale(leaf: Operator, input: Operator): Boolean = leaf.hasScan && 
input.hasShuffleScan
+    val stale = leaves.zip(current).filter { case (leaf, input) => 
isStale(leaf, input) }
+    val isRefreshable = current.length == leaves.length && stale.nonEmpty &&
+      stale.forall { case (leaf, input) =>
+        leaf.getScan.getFieldsList == input.getShuffleScan.getFieldsList
+      }
+    if (isRefreshable) {
+      val newLeaves = leaves.zip(current).map { case (leaf, input) =>
+        if (isStale(leaf, input)) input else leaf
+      }
+      Some(withLeaves(native.nativeOp, newLeaves))

Review Comment:
   `leaves.zip(current)` and `isStale` run twice here, once for `stale` and 
again for `newLeaves`. Rebuilding through a plain recursive `mapLeaves(op)(f)` 
removes the second `zip` and `newLeaves`. It also replaces `withLeaves`, which 
threads a `(Vector, List)` through a `foldLeft` and uses `head` and `tail`:
   
   ```scala
   val inputs = current.iterator
   Some(mapLeaves(native.nativeOp) { leaf =>
     val input = inputs.next()
     if (isStale(leaf, input)) input else leaf
   })
   ```
   
   `mapLeaves` applies `f` to each childless operator in walk order and 
rebuilds each parent with `toBuilder.clearChildren().addAllChildren(...)`. On a 
few hand-built trees and every keep or replace pattern it produced the same 
trees as `withLeaves`. I have not run it inside the rule.



##########
spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala:
##########
@@ -608,6 +610,89 @@ case class CometExecRule(session: SparkSession)
     }
   }
 
+  /**
+   * AQE re-plans around a materialized stage by reusing the physical node 
linked to it, so a
+   * native operator that shares its logical node with a shuffle stage (the 
final aggregate of a
+   * two-phase aggregate) keeps the native plan it got while that input was a 
bare exchange, read
+   * through a plain `Scan`. Once the input is a sink that reads the shuffle 
directly, its
+   * `ShuffleScan` takes the place of the stale leaf. The leaf is patched in 
place because
+   * converting the node again from `originalPlan` would drop the stage's 
logical link that AQE
+   * relies on and re-run serde on a node that is already planned.
+   */
+  private def refreshStaleShuffleScans(op: SparkPlan): SparkPlan = op match {
+    // These build their own `Scan` over the child rather than embedding the 
child's plan.
+    case _: CometNativeWriteExec | _: CometIcebergWriteExec | _: 
CometWriteFilesExec => op

Review Comment:
   This is the same three-class list as the `firstNativeOp` reset in `_apply` 
(the `isInstanceOf[CometNativeWriteExec] || ...` check around line 988). One 
small predicate named for what both places rely on, a native plan that builds 
its own `Scan` over a child that runs as a separate block, would keep the two 
in step if another writer is added.



##########
spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala:
##########
@@ -1391,6 +1394,178 @@ class CometNativeShuffleSuite extends CometTestBase 
with AdaptiveSparkPlanHelper
     }
   }
 
+  /**
+   * For each native block in `plan`, query stages included, that reads a 
shuffle: each input of
+   * the block, in walk order, paired with the kind of the native leaf that 
reads it.
+   */
+  private def shuffleReadingBlocks(plan: SparkPlan): Seq[Seq[(String, 
String)]] = {
+    def inputs(op: SparkPlan): Seq[SparkPlan] = op.children.flatMap {
+      case native: CometNativeExec if native.children.nonEmpty => 
inputs(native)
+      case other => Seq(other)
+    }
+    def leaves(op: OperatorOuterClass.Operator): 
Seq[OperatorOuterClass.Operator] =
+      if (op.getChildrenCount == 0) Seq(op)
+      else op.getChildrenList.asScala.toSeq.flatMap(leaves)
+    def readsShuffle(input: SparkPlan): Boolean = input match {
+      case _: ShuffleQueryStageExec | _: AQEShuffleReadExec | _: 
CometShuffleExchangeExec => true
+      case _ => false
+    }
+    // A node AQE reuses can keep the serialized plan of a block it used to be 
the root of.
+    val natives = collect(plan) { case native: CometNativeExec => native }
+    val nested = natives.flatMap(_.children.collect { case child: 
CometNativeExec => child })
+    natives
+      .filter(root => root.serializedPlanOpt.isDefined && !nested.exists(_ eq 
root))
+      .filter(inputs(_).exists(readsShuffle))
+      .map { root =>
+        val nativePlan = 
OperatorOuterClass.Operator.parseFrom(root.serializedPlanOpt.plan.get)
+        val blockInputs = inputs(root).map(_.getClass.getSimpleName)
+        val leafKinds = leaves(nativePlan).map(_.getOpStructCase.name)
+        assert(blockInputs.length == leafKinds.length, s"$blockInputs vs 
$leafKinds")
+        blockInputs.zip(leafKinds)
+      }
+  }
+
+  private val shuffleStageRead = Seq("ShuffleQueryStageExec" -> "SHUFFLE_SCAN")
+
+  private val aqeWithoutCoalescing = Seq(
+    SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
+    SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "false")
+
+  /** A two-phase aggregate over many map partitions. */
+  private def groupedSum: DataFrame =
+    spark
+      .range(0, 10000, 1, 50)
+      .groupBy((col("id") % 997).as("k"))
+      .agg(sum("id"), count("id"))
+
+  test("AQE final aggregate reads its shuffle stage through ShuffleScan") {

Review Comment:
   These four tests (this one, and the ones at 1449, 1516 and 1528) run 
`groupedSum` under a different conf set and assert one expected block each. A 
single table-driven test over (confs, expected block) would keep all four cases 
and drop three copies of the setup. The direct-read-off case at 1449 is also 
covered at the rule level by `CometExecRuleSuite` (`keeps an AQE-reused 
aggregate's Scan with shuffle direct read off`), so one of the two can go.



##########
spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala:
##########
@@ -368,76 +439,148 @@ class CometExecRuleSuite extends CometTestBase {
 
           assert(
             
spark.sessionState.conf.getConf(SQLConf.ADAPTIVE_CUSTOM_COST_EVALUATOR_CLASS).isEmpty)
-          type Replan = (SparkPlan, LogicalPlan)
-          val pending = new ThreadLocal[List[Option[Replan]]] {
-            override def initialValue(): List[Option[Replan]] = Nil
-          }
-          val observed = new 
ConcurrentLinkedQueue[(CometBroadcastExchangeExec, LogicalPlan)]()
-          val tempTag = AdaptiveSparkPlanExec.TEMP_LOGICAL_PLAN_TAG
-          val costEvaluator = SimpleCostEvaluator(forceOptimizeSkewedJoin = 
false)
-          beforeCometPreparation = plan => {
-            val replan = plan match {
-              case broadcast: CometBroadcastExchangeExec =>
-                broadcast.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).collect {
-                  case stage: LogicalQueryStage =>
-                    assert(stage.physicalPlan eq broadcast)
-                    assert(broadcast.getTagValue(tempTag).exists(_ eq 
stage.logicalPlan))
-                    (broadcast.clone(), stage.logicalPlan)
-                }
-              case _ => None
-            }
-            pending.set(replan :: pending.get())
-          }
-          afterCometPreparation = plan => {
-            val replan = pending.get().head
-            val remaining = pending.get().tail
-            if (remaining.isEmpty) pending.remove() else pending.set(remaining)
-            replan.foreach { case (previous, logicalPlan) =>
-              val broadcast = plan.asInstanceOf[CometBroadcastExchangeExec]
-              assert(broadcast.logicalLink.exists(_ eq logicalPlan))
-              assert(broadcast.getTagValue(tempTag).exists(_ eq logicalPlan))
-              // Spark rejects an equal-cost candidate when its physical tree 
is unchanged.
-              // Pin both inputs to that decision, including Comet's retained 
temporary link.
-              assert(previous == broadcast)
-              assert(costEvaluator.evaluateCost(previous) == SimpleCost(0))
-              assert(costEvaluator.evaluateCost(broadcast) == SimpleCost(0))
-              observed.add(
-                (broadcast.clone().asInstanceOf[CometBroadcastExchangeExec], 
logicalPlan))
+          f
+        }
+      }
+    }
+  }
+
+  /**
+   * Runs the DPP broadcast link query, checks its answer and that the join 
and the pruning
+   * broadcast are native, and returns the executed plan.
+   */
+  private def checkDppLinkQuery(): SparkPlan = {
+    val df = sql("""
+        |SELECT /*+ BROADCAST(d) */ f.k, f.total, d.total
+        |FROM (SELECT k, SUM(v) AS total FROM dpp_link_fact GROUP BY k) f
+        |JOIN (SELECT k, SUM(v) AS total FROM dpp_link_dim
+        |      WHERE country = 'DE' GROUP BY k) d ON f.k = d.k
+        |""".stripMargin)
+    QueryTest.checkAnswer(
+      df,
+      (0 until 8 by 2).map(k => Row(k, 224L + 8L * k, 48L + 4L * k)),
+      checkToRDD = false)
+    val plan = df.queryExecution.executedPlan
+    assert(collect(plan) { case b: CometBroadcastHashJoinExec => b }.nonEmpty)
+    assert(collectWithSubqueries(plan) { case s: CometSubqueryBroadcastExec =>
+      s
+    }.nonEmpty)
+    plan
+  }
+
+  test("AQE DPP broadcast roots retain temporary logical links after an 
unchanged replan") {
+    assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+")
+    // With direct read, the replan that first sees the dim side's shuffle 
stage moves the
+    // aggregate under the broadcast onto a ShuffleScan, so that replan is not 
unchanged.
+    withDppLinkViews(directRead = false) {
+      type Replan = (SparkPlan, LogicalPlan)
+      val pending = new ThreadLocal[List[Option[Replan]]] {
+        override def initialValue(): List[Option[Replan]] = Nil
+      }
+      val observed = new ConcurrentLinkedQueue[(CometBroadcastExchangeExec, 
LogicalPlan)]()
+      val tempTag = AdaptiveSparkPlanExec.TEMP_LOGICAL_PLAN_TAG
+      val costEvaluator = SimpleCostEvaluator(forceOptimizeSkewedJoin = false)
+      beforeCometPreparation = plan => {
+        val replan = plan match {
+          case broadcast: CometBroadcastExchangeExec =>
+            broadcast.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).collect {
+              case stage: LogicalQueryStage =>
+                assert(stage.physicalPlan eq broadcast)
+                assert(broadcast.getTagValue(tempTag).exists(_ eq 
stage.logicalPlan))
+                (broadcast.clone(), stage.logicalPlan)
             }
-          }
-          try {
-            val df = sql("""
-                |SELECT /*+ BROADCAST(d) */ f.k, f.total, d.total
-                |FROM (SELECT k, SUM(v) AS total FROM dpp_link_fact GROUP BY 
k) f
-                |JOIN (SELECT k, SUM(v) AS total FROM dpp_link_dim
-                |      WHERE country = 'DE' GROUP BY k) d ON f.k = d.k
-                |""".stripMargin)
-            QueryTest.checkAnswer(
-              df,
-              (0 until 8 by 2).map(k => Row(k, 224L + 8L * k, 48L + 4L * k)),
-              checkToRDD = false)
-            val plan = df.queryExecution.executedPlan
-            assert(collect(plan) { case b: CometBroadcastHashJoinExec => b 
}.nonEmpty)
-            assert(collectWithSubqueries(plan) { case s: 
CometSubqueryBroadcastExec =>
-              s
-            }.nonEmpty)
-            assert(!observed.isEmpty, "Expected a DPP broadcast root with a 
direct logical stage")
-            observed.iterator().asScala.foreach { case (broadcast, 
logicalPlan) =>
-              // Give the isolated snapshot conflicting links to pin Spark's 
TEMP-over-direct
-              // precedence, which would otherwise be invisible after Comet 
repairs both.
-              broadcast.setLogicalLink(LogicalQueryStage(logicalPlan, 
broadcast))
-              val stage = BroadcastQueryStageExec(0, broadcast, 
broadcast.canonicalized)
-              val setStageLink = 
PrivateMethod[Unit](Symbol("setLogicalLinkForNewQueryStage"))
-              plan
-                .asInstanceOf[AdaptiveSparkPlanExec]
-                .invokePrivate(setStageLink(stage, broadcast))
-              assert(stage.logicalLink.exists(_ eq logicalPlan))
+          case _ => None
+        }
+        pending.set(replan :: pending.get())
+      }
+      afterCometPreparation = plan => {
+        val replan = pending.get().head
+        val remaining = pending.get().tail
+        if (remaining.isEmpty) pending.remove() else pending.set(remaining)
+        replan.foreach { case (previous, logicalPlan) =>
+          val broadcast = plan.asInstanceOf[CometBroadcastExchangeExec]
+          assert(broadcast.logicalLink.exists(_ eq logicalPlan))
+          assert(broadcast.getTagValue(tempTag).exists(_ eq logicalPlan))
+          // Spark rejects an equal-cost candidate when its physical tree is 
unchanged.
+          // Pin both inputs to that decision, including Comet's retained 
temporary link.
+          assert(previous == broadcast)
+          assert(costEvaluator.evaluateCost(previous) == SimpleCost(0))
+          assert(costEvaluator.evaluateCost(broadcast) == SimpleCost(0))
+          
observed.add((broadcast.clone().asInstanceOf[CometBroadcastExchangeExec], 
logicalPlan))
+        }
+      }
+      try {
+        val plan = checkDppLinkQuery()
+        assert(!observed.isEmpty, "Expected a DPP broadcast root with a direct 
logical stage")
+        observed.iterator().asScala.foreach { case (broadcast, logicalPlan) =>
+          // Give the isolated snapshot conflicting links to pin Spark's 
TEMP-over-direct
+          // precedence, which would otherwise be invisible after Comet 
repairs both.
+          broadcast.setLogicalLink(LogicalQueryStage(logicalPlan, broadcast))
+          val stage = BroadcastQueryStageExec(0, broadcast, 
broadcast.canonicalized)
+          val setStageLink = 
PrivateMethod[Unit](Symbol("setLogicalLinkForNewQueryStage"))
+          plan
+            .asInstanceOf[AdaptiveSparkPlanExec]
+            .invokePrivate(setStageLink(stage, broadcast))
+          assert(stage.logicalLink.exists(_ eq logicalPlan))
+        }
+      } finally {
+        beforeCometPreparation = _ => ()
+        afterCometPreparation = _ => ()
+      }
+    }
+  }
+
+  test("AQE DPP broadcast roots keep their logical links through a ShuffleScan 
replan") {
+    assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+")
+    withDppLinkViews(directRead = true) {
+      val pending = new ThreadLocal[List[Option[(SparkPlan, LogicalPlan)]]] {
+        override def initialValue(): List[Option[(SparkPlan, LogicalPlan)]] = 
Nil
+      }
+      // For each replanned broadcast root: whether Comet changed it, whether 
it kept the
+      // stage's direct and temporary links, and the leaf kinds of the native 
plan under it.
+      val observed = new ConcurrentLinkedQueue[(Boolean, Boolean, Boolean, 
Seq[String])]()
+      val tempTag = AdaptiveSparkPlanExec.TEMP_LOGICAL_PLAN_TAG
+      beforeCometPreparation = plan => {
+        val replan = plan match {
+          case broadcast: CometBroadcastExchangeExec =>
+            broadcast.getTagValue(SparkPlan.LOGICAL_PLAN_TAG).collect {
+              case stage: LogicalQueryStage => (broadcast.clone(), 
stage.logicalPlan)
             }
-          } finally {
-            beforeCometPreparation = _ => ()
-            afterCometPreparation = _ => ()
+          case _ => None
+        }
+        pending.set(replan :: pending.get())
+      }
+      afterCometPreparation = plan => {
+        val replan = pending.get().head
+        val remaining = pending.get().tail
+        if (remaining.isEmpty) pending.remove() else pending.set(remaining)
+        replan.foreach { case (previous, logicalPlan) =>

Review Comment:
   The `pending` `ThreadLocal` stack and the push and pop in 
`beforeCometPreparation` and `afterCometPreparation` are copied from the test 
above (lines 477 to 501). Only the captured value and the per-replan checks 
differ. A small `withReplanObserver` helper taking those two pieces would keep 
the two tests from drifting.



##########
spark/src/test/scala/org/apache/comet/exec/CometNativeShuffleSuite.scala:
##########
@@ -1391,6 +1394,178 @@ class CometNativeShuffleSuite extends CometTestBase 
with AdaptiveSparkPlanHelper
     }
   }
 
+  /**
+   * For each native block in `plan`, query stages included, that reads a 
shuffle: each input of
+   * the block, in walk order, paired with the kind of the native leaf that 
reads it.
+   */
+  private def shuffleReadingBlocks(plan: SparkPlan): Seq[Seq[(String, 
String)]] = {
+    def inputs(op: SparkPlan): Seq[SparkPlan] = op.children.flatMap {
+      case native: CometNativeExec if native.children.nonEmpty => 
inputs(native)
+      case other => Seq(other)
+    }
+    def leaves(op: OperatorOuterClass.Operator): 
Seq[OperatorOuterClass.Operator] =
+      if (op.getChildrenCount == 0) Seq(op)
+      else op.getChildrenList.asScala.toSeq.flatMap(leaves)
+    def readsShuffle(input: SparkPlan): Boolean = input match {
+      case _: ShuffleQueryStageExec | _: AQEShuffleReadExec | _: 
CometShuffleExchangeExec => true
+      case _ => false
+    }
+    // A node AQE reuses can keep the serialized plan of a block it used to be 
the root of.
+    val natives = collect(plan) { case native: CometNativeExec => native }
+    val nested = natives.flatMap(_.children.collect { case child: 
CometNativeExec => child })
+    natives
+      .filter(root => root.serializedPlanOpt.isDefined && !nested.exists(_ eq 
root))
+      .filter(inputs(_).exists(readsShuffle))
+      .map { root =>
+        val nativePlan = 
OperatorOuterClass.Operator.parseFrom(root.serializedPlanOpt.plan.get)
+        val blockInputs = inputs(root).map(_.getClass.getSimpleName)
+        val leafKinds = leaves(nativePlan).map(_.getOpStructCase.name)
+        assert(blockInputs.length == leafKinds.length, s"$blockInputs vs 
$leafKinds")
+        blockInputs.zip(leafKinds)
+      }
+  }
+
+  private val shuffleStageRead = Seq("ShuffleQueryStageExec" -> "SHUFFLE_SCAN")
+
+  private val aqeWithoutCoalescing = Seq(
+    SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
+    SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "false")
+
+  /** A two-phase aggregate over many map partitions. */
+  private def groupedSum: DataFrame =
+    spark
+      .range(0, 10000, 1, 50)
+      .groupBy((col("id") % 997).as("k"))
+      .agg(sum("id"), count("id"))
+
+  test("AQE final aggregate reads its shuffle stage through ShuffleScan") {
+    withSQLConf(aqeWithoutCoalescing: _*) {
+      val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+      checkCometOperatorsInFinalPlan(plan)
+      assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead), plan)
+    }
+  }
+
+  test("AQE final aggregate keeps a plain Scan with shuffle direct read 
disabled") {
+    withSQLConf(
+      aqeWithoutCoalescing :+ (CometConf.COMET_SHUFFLE_DIRECT_READ_ENABLED.key 
-> "false"): _*) {
+      val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+      assert(shuffleReadingBlocks(plan) == Seq(Seq("ShuffleQueryStageExec" -> 
"SCAN")), plan)
+    }
+  }
+
+  test("AQE native aggregate chain over a shuffle stage reads it through 
ShuffleScan") {
+    withSQLConf(aqeWithoutCoalescing: _*) {
+      // A distinct aggregate plans two shuffles, with two native aggregates 
between them.
+      val df = spark
+        .range(0, 10000, 1, 20)
+        .groupBy((col("id") % 97).as("k"))
+        .agg(countDistinct(col("id") % 13), sum("id"))
+      val (_, plan) = checkSparkAnswerAndOperator(df)
+      checkCometOperatorsInFinalPlan(plan)
+      assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead, 
shuffleStageRead), plan)
+    }
+  }
+
+  test("AQE join of final aggregates reads only its shuffle stages through 
ShuffleScan") {
+    val agg = spark
+      .range(0, 10000, 1, 20)
+      .groupBy((col("id") % 97).as("k"))
+      .agg(sum("id").as("s"))
+    withSQLConf(
+      aqeWithoutCoalescing ++ Seq(
+        SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
+        SQLConf.ADAPTIVE_AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1"): _*) {
+      // The second side reuses the first side's shuffle.
+      val df = agg.join(agg.select(col("k"), (col("s") + 1).as("s2")), "k")
+      val (_, plan) = checkSparkAnswerAndOperator(df)
+      checkCometOperatorsInFinalPlan(plan, classOf[ReusedExchangeExec])
+      assert(collect(plan) { case r: ReusedExchangeExec => r }.nonEmpty, plan)
+      assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead ++ 
shuffleStageRead), plan)
+    }
+    withSQLConf(aqeWithoutCoalescing: _*) {
+      // The broadcast side stays a plain Scan.
+      val dim = spark.range(0, 50).select(col("id").as("k"), (col("id") * 
2).as("w"))
+      val df = agg.join(broadcast(dim), "k")
+      val (_, plan) = checkSparkAnswerAndOperator(df)
+      checkCometOperatorsInFinalPlan(plan)
+      val joinBlock = shuffleStageRead :+ ("BroadcastQueryStageExec" -> "SCAN")
+      assert(shuffleReadingBlocks(plan) == Seq(joinBlock), plan)
+    }
+  }
+
+  test("AQE keeps exchange reuse above equivalent aggregates that read through 
ShuffleScan") {
+    withSQLConf(aqeWithoutCoalescing: _*) {
+      // Each side's aggregate reads its own shuffle stage, so their 
ShuffleScan sources differ.
+      val agg = spark
+        .range(0, 10000, 1, 20)
+        .groupBy((col("id") % 97).as("k"))
+        .agg(sum("id").as("s"))
+        .repartition(7, col("s"))
+      val (_, plan) = checkSparkAnswerAndOperator(agg.union(agg))
+      checkCometOperatorsInFinalPlan(plan, classOf[ReusedExchangeExec])
+      assertExchangeReuseOver(plan, "Expected the exchange above the aggregate 
to be reused") {
+        case a: CometHashAggregateExec if a.modes.contains(Final) => a
+      }
+      val blocks = shuffleReadingBlocks(plan)
+      val readsShuffleScan = blocks.forall(_.forall { case (_, leaf) => leaf 
== "SHUFFLE_SCAN" })
+      assert(blocks.nonEmpty && readsShuffleScan, plan)
+    }
+  }
+
+  test("AQE final aggregate over a coalesced shuffle read reads it through 
ShuffleScan") {
+    // The aggregate takes the stage's ShuffleScan before AQE puts the 
coalesced read between
+    // them, and the ShuffleScan then reads the partitions that the coalesced 
read specifies.
+    withSQLConf(
+      SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true",
+      SQLConf.COALESCE_PARTITIONS_ENABLED.key -> "true") {
+      val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+      checkCometOperatorsInFinalPlan(plan)
+      assert(shuffleReadingBlocks(plan) == Seq(Seq("AQEShuffleReadExec" -> 
"SHUFFLE_SCAN")), plan)
+    }
+  }
+
+  test("final aggregate without AQE reads its shuffle exchange through a plain 
Scan") {
+    withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") {
+      val (_, plan) = checkSparkAnswerAndOperator(groupedSum)
+      assert(shuffleReadingBlocks(plan) == Seq(Seq("CometShuffleExchangeExec" 
-> "SCAN")), plan)
+    }
+  }
+
+  test("AQE DPP query with a broadcast side reads its shuffle stage through 
ShuffleScan") {
+    assume(isSpark35Plus, "Native AQE DPP requires Spark 3.5+")
+    withTempDir { dir =>
+      withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+        spark
+          .range(0, 1000, 1, 10)
+          .selectExpr("id", "id % 10 AS p")
+          .write
+          .partitionBy("p")
+          .parquet(s"$dir/fact")
+        spark.range(0, 10).selectExpr("id AS p", "id % 3 AS 
x").write.parquet(s"$dir/dim")
+      }
+      withTempView("fact", "dim") {
+        spark.read.parquet(s"$dir/fact").createOrReplaceTempView("fact")
+        spark.read.parquet(s"$dir/dim").createOrReplaceTempView("dim")
+        withSQLConf(
+          aqeWithoutCoalescing :+
+            (SQLConf.DYNAMIC_PARTITION_PRUNING_ENABLED.key -> "true"): _*) {
+          val df = sql("""SELECT f.p, sum(f.id), count(*) FROM fact f JOIN dim 
d ON f.p = d.p
+              |WHERE d.x = 1 GROUP BY f.p""".stripMargin)
+          val (_, plan) = checkSparkAnswer(df)
+          assert(shuffleReadingBlocks(plan) == Seq(shuffleStageRead), plan)
+
+          // Aggregates on both sides, the broadcast one feeding the pruning 
filter. This checks
+          // the answer only.
+          checkSparkAnswer(sql("""SELECT /*+ BROADCAST(d) */ f.p, f.total, 
d.total
+              |FROM (SELECT p, sum(id) AS total FROM fact GROUP BY p) f
+              |JOIN (SELECT p, sum(x) AS total FROM dim WHERE x = 1 GROUP BY 
p) d
+              |ON f.p = d.p""".stripMargin))

Review Comment:
   This second query only checks the answer and has the same shape as 
`checkDppLinkQuery` in `CometExecRuleSuite` (an aggregate on each side and a 
`BROADCAST(d)` hint), which also asserts that the join and the pruning 
broadcast are native. Could it be dropped here?



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