diff --git a/docs/source/contributor-guide/adding_a_new_operator.md b/docs/source/contributor-guide/adding_a_new_operator.md index 15c4e9169b..09cac3a1c7 100644 --- a/docs/source/contributor-guide/adding_a_new_operator.md +++ b/docs/source/contributor-guide/adding_a_new_operator.md @@ -731,6 +731,34 @@ Use `QueryPlanSerde.exprToProto` to convert Spark expressions to protobuf: val protoExpr = exprToProto(sparkExpr, inputSchema) ``` +### Restoring the Spark operator (`sparkFallback`) + +`CometExec.originalPlan` is the Spark operator this node replaced. `CometExecRule` copies +`originalPlan.logicalLink` onto the Comet node, which is how AQE finds the node again when it +re-plans a stage. `RevertNativeForTransitionHeavyStages` calls `sparkFallback(newChildren)` to +rebuild that Spark operator with the children of the reverted stage. + +The default implementation is `originalPlan.withNewChildren(newChildren)`. It refuses a null +`originalPlan`, an `originalPlan` that is one of the node's own children, or a different number of +children than the Spark operator has. + +Override `sparkFallback` when conversion changes the plan shape, so the restored node is not that +Spark operator with the same children. `CometNativeWriteExec` replaces a `DataWritingCommandExec` +and drops the `WriteFilesExec` under it; its override puts that wrapper back around the restored +input. `CometIcebergWriteExec` keeps the same shape as `IcebergWriteExec`, so the default is +enough. + +Also override `sparkFallback` when the operator's live state differs from `originalPlan`. +`CometNativeScanExec` restores its current partition and data filters, and +`CometIcebergNativeScanExec` restores its current runtime filters, so AQE's executable DPP +subqueries survive reversion. `CometLocalTopKExec` returns the restored child directly: Comet +inserted that local candidate selection, and only the outer TopK restores Spark's offset and +projection. Rebuilding the original TopK at both nodes would apply it twice. + +Do not point `originalPlan` at a child. If that child is a shuffle or query stage, the copied +logical link puts this node inside the stage's `LogicalQueryStage`. AQE then re-plans a second +copy of the operator around the one that is already there. + ### Handling Fallback Use `withInfo` to tag operators with fallback reasons: diff --git a/docs/source/contributor-guide/iceberg-writes.md b/docs/source/contributor-guide/iceberg-writes.md index 763cb88565..b40062214b 100644 --- a/docs/source/contributor-guide/iceberg-writes.md +++ b/docs/source/contributor-guide/iceberg-writes.md @@ -107,9 +107,9 @@ serializer requires that tag on Spark 4.1+, leaving stock writers on Spark throu ## From `IcebergWrite` to `CometIcebergWrite` `CometExecRule` converts an `IcebergWriteExec` with the `CometIcebergNativeWrite` operator serde -when `spark.comet.write.iceberg.enabled` is on. Two arms in `CometExecRule` handle it: one unwraps -the double conversion AQE can produce when it re-fires write planning over a sub-tree that already -contains a `CometIcebergWriteExec`, and the other calls `convertToComet`. +when `spark.comet.write.iceberg.enabled` is on. A single arm in `CometExecRule` handles the +conversion by calling `convertToComet`. The converted node keeps the `IcebergWriteExec` as its +`originalPlan`, so AQE re-plans the write from that node. `CometIcebergNativeWrite.requiresNativeChildren` is `true`. The native writer consumes Arrow batches from its child over FFI, so the conversion is declined unless the child is already a Comet @@ -421,8 +421,11 @@ Each of these has caused a bug on this path: ([#5691](https://github.com/apache/datafusion-comet/issues/5691), [#5693](https://github.com/apache/datafusion-comet/issues/5693), [#6141](https://github.com/apache/datafusion-comet/issues/6141)). -- **Plan rewrites must keep the write node.** Rules that restore Spark operators from a Comet - node's `originalPlan` have to handle the write execs, whose `originalPlan` today is their child +- **Plan rewrites must keep the write node.** `CometIcebergWriteExec.originalPlan` is the + `IcebergWriteExec` it replaced. Restoring Spark execution goes through + [`CometExec.sparkFallback`](adding_a_new_operator.md#restoring-the-spark-operator-sparkfallback), + which rebuilds that node around the reverted children. Pointing `originalPlan` at the child + drops the write, and AQE then treats the write as part of that child's stage ([#5719](https://github.com/apache/datafusion-comet/issues/5719)). - **The kill switch must still work.** Code that runs for Iceberg writes has to respect `spark.comet.enabled`, so that disabling Comet restores Spark's own plan diff --git a/docs/source/contributor-guide/native_shuffle.md b/docs/source/contributor-guide/native_shuffle.md index 6108844788..e9cf777441 100644 --- a/docs/source/contributor-guide/native_shuffle.md +++ b/docs/source/contributor-guide/native_shuffle.md @@ -39,6 +39,12 @@ Compare this to JVM shuffle's data path: Comet Native (columnar) → ColumnarToRowExec → rows → JVM Shuffle → Arrow IPC → columnar ``` +When `RevertNativeForTransitionHeavyStages` restores the map stage to Spark execution, the +native exchange stays in place. Its input still needs Arrow-backed Comet vectors, even if the +restored Spark operator supports columnar output. The rule adds `CometSparkToColumnarExec` to +convert either Spark rows or Spark columnar batches to Arrow before the native shuffle consumes +them. Spark's `RowToColumnarExec` alone does not satisfy this input contract. + ## When Native Shuffle is Used Native shuffle (`CometExchange`) is selected when all of the following conditions are met: diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala index 736143baf6..121ab445e7 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -542,23 +542,15 @@ case class CometExecRule(session: SparkSession, queryStagePrep: Boolean = false) // DataWritingCommandExec and re-implement the write framework inside CometNativeWriteExec. // This path is retained only for 3.4/3.5 and goes away with them. // - // AQE reoptimization looks for `DataWritingCommandExec` or `WriteFilesExec` - // if there is none it would reinsert write nodes, and since Comet remap those nodes - // to Comet counterparties the write nodes are twice to the plan. - // Checking if AQE inserted another write Command on top of existing write command - case _ @DataWritingCommandExec(_, w: WriteFilesExec) - if !isSpark40Plus && w.child.isInstanceOf[CometNativeWriteExec] => - w.child - + // `originalPlan` is that command. This rule copies `originalPlan.logicalLink` onto the + // Comet node, so AQE re-plans the write with the command rather than with whatever child + // happened to sit under it. No second DataWritingCommandExec is inserted on top. case op: DataWritingCommandExec if !isSpark40Plus => convertToComet(op, CometDataWritingCommand).getOrElse(op) - // AQE re-fires the Iceberg write planning on every stage materialisation, so a - // partitioned write's physical sub-tree may already contain a `CometIcebergWriteExec`. - // Unwrap to avoid a double conversion. - case op: IcebergWriteExec if op.child.isInstanceOf[CometIcebergWriteExec] => - op.child - + // `originalPlan` is this IcebergWriteExec, so AQE re-plans the write as this node. + // A shuffle directly under the native write stays in the child stage and is not wrapped + // again. case op: IcebergWriteExec if CometConf.COMET_ICEBERG_NATIVE_WRITE_ENABLED.get(op.conf) => convertToComet(op, CometIcebergNativeWrite).getOrElse(op) diff --git a/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala b/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala index 1181a573bc..3f3c57f064 100644 --- a/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala +++ b/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala @@ -23,7 +23,8 @@ import org.apache.spark.internal.Logging import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, PartialMerge} import org.apache.spark.sql.catalyst.rules.Rule -import org.apache.spark.sql.comet.{CometBaseAggregateExec, CometColumnarToRowExec, CometExec, CometIcebergNativeScanExec, CometLocalTopKExec, CometNativeColumnarToRowExec, CometNativeScanExec, CometSparkToColumnarExec} +import org.apache.spark.sql.comet.{CometBaseAggregateExec, CometColumnarToRowExec, CometExec, CometNativeColumnarToRowExec, CometSparkToColumnarExec} +import org.apache.spark.sql.comet.execution.shuffle.{CometNativeShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution.{ColumnarToRowExec, ColumnarToRowTransition, RowToColumnarExec, SparkPlan} import org.apache.spark.sql.execution.adaptive.QueryStageExec import org.apache.spark.sql.execution.exchange.{BroadcastExchangeLike, ShuffleExchangeLike} @@ -62,19 +63,18 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession, wholePlan plan match { case _: BroadcastExchangeLike => plan case exchange: ShuffleExchangeLike => - revertStageIfNeeded(exchange.child, exchange.supportsColumnar) + revertShuffleStageIfNeeded(exchange) .map(reverted => exchange.withNewChildren(Seq(reverted))) .getOrElse(plan) case _ => - // Result stage: its output is collected as rows, so no consumer requires columnar input - // and the reverted stage needs no trailing R2C. + // Result stage: its output is collected as rows. revertStageIfNeeded(plan, outputColumnar = false).getOrElse(plan) } } private def applyForNonAQE(plan: SparkPlan): SparkPlan = { val withRevertedStages = plan.transformUp { case exchange: ShuffleExchangeLike => - revertStageIfNeeded(exchange.child, exchange.supportsColumnar) + revertShuffleStageIfNeeded(exchange) .map(reverted => exchange.withNewChildren(Seq(reverted))) .getOrElse(exchange) } @@ -82,12 +82,22 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession, wholePlan .getOrElse(withRevertedStages) } + private def revertShuffleStageIfNeeded(exchange: ShuffleExchangeLike): Option[SparkPlan] = { + val outputArrow = exchange match { + case comet: CometShuffleExchangeExec => comet.shuffleType == CometNativeShuffle + case _ => false + } + revertStageIfNeeded(exchange.child, exchange.supportsColumnar, outputArrow) + } + /** - * Reverts the stage if C2R count exceeds threshold. Wraps in R2C if exchange needs columnar. + * Reverts the stage if C2R count exceeds threshold, restoring the stage's output format when + * the reverted root does not satisfy it. */ private def revertStageIfNeeded( stagePlan: SparkPlan, - outputColumnar: Boolean): Option[SparkPlan] = { + outputColumnar: Boolean, + outputArrow: Boolean = false): Option[SparkPlan] = { val transitionCount = countTransitions(stagePlan) if (transitionCount <= maxTransitions) return None @@ -101,11 +111,27 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession, wholePlan val reason = s"Stage reverted: $transitionCount C2R transitions exceed threshold $maxTransitions" - val reverted = revertToSpark(stagePlan) - val result = if (outputColumnar && !reverted.supportsColumnar) { - RowToColumnarExec(withFallbackReason(reverted, reason)) + val reverted = + try { + revertToSpark(stagePlan) + } catch { + case e: CometExec.InvalidSparkFallbackException => + logWarning( + "Skipping transition-heavy stage reversion because a Comet operator could not " + + s"restore its Spark plan: ${e.getMessage}") + return None + } + val revertedWithReason = withFallbackReason(reverted, reason) + val result = if (outputArrow) { + // Native shuffle consumes Arrow-backed Comet vectors, not arbitrary Spark columnar + // batches. This bridge converts both row-based and vectorized Spark fallback roots. + CometSparkToColumnarExec(revertedWithReason) + } else if (outputColumnar && !reverted.supportsColumnar) { + RowToColumnarExec(revertedWithReason) + } else if (!outputColumnar && reverted.supportsColumnar) { + ColumnarToRowExec(revertedWithReason) } else { - withFallbackReason(reverted, reason) + revertedWithReason } Some(result) } @@ -146,16 +172,31 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession, wholePlan } /** - * Like `transformDown`, never descends stage-boundary children. + * Like `transformDown`, never descends stage-boundary children. If the rule rewrites the + * current node, re-apply it to the result so stacked transitions such as + * `CometSparkToColumnarExec(CometNativeColumnarToRowExec(x))` are fully unwrapped before + * children are visited. Spark's `transformDown` does not do this; leaving the inner C2R in + * place later calls `CometNativeColumnarToRowExec.withNewChildren` with a reverted row-based + * child, which asserts `child.supportsColumnar`. + * + * A rewrite can itself be the stage boundary. Unwrapping a transition that sits directly on a + * shuffle yields that shuffle, and descending into it strips transitions in the next stage. + * `transformStageUp` and `insertTransitions` do not cross the exchange, so those transitions + * would not be restored (#6152). Return the boundary unchanged. */ private def transformStageDown(plan: SparkPlan)( rule: PartialFunction[SparkPlan, SparkPlan]): SparkPlan = { val transformed = rule.applyOrElse(plan, identity[SparkPlan]) - val newChildren = transformed.children.map { child => - if (isStageBoundary(child)) child else transformStageDown(child)(rule) + if (transformed ne plan) { + if (isStageBoundary(transformed)) transformed + else transformStageDown(transformed)(rule) + } else { + val newChildren = transformed.children.map { child => + if (isStageBoundary(child)) child else transformStageDown(child)(rule) + } + if (newChildren == transformed.children) transformed + else transformed.withNewChildren(newChildren) } - if (newChildren == transformed.children) transformed - else transformed.withNewChildren(newChildren) } /** Like `transformUp`, never descends stage-boundary children. */ @@ -184,7 +225,27 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession, wholePlan count } + /** + * Checks for Comet operators whose original Spark plan is also their input. This must run + * before any bottom-up rewrite replaces children. Otherwise a stable `originalPlan` reference + * can keep pointing at the old Comet child after `withNewChildren`, hiding the alias and + * causing fallback to reconstruct that Comet child instead of a Spark operator. + */ + private def validateOriginalPlanAliases(plan: SparkPlan): Unit = plan match { + case _ if isStageBoundary(plan) => () + case cometExec: CometExec => + val sparkPlan = cometExec.originalPlan + if (sparkPlan != null && cometExec.children.exists(_ eq sparkPlan)) { + throw new CometExec.InvalidSparkFallbackException( + s"${cometExec.getClass.getSimpleName} aliases its original Spark plan with a child") + } + cometExec.children.foreach(validateOriginalPlanAliases) + case _ => + plan.children.foreach(validateOriginalPlanAliases) + } + private[rules] def revertToSpark(plan: SparkPlan): SparkPlan = { + validateOriginalPlanAliases(plan) val stripped = transformStageDown(plan) { case CometNativeColumnarToRowExec(child) => child case CometColumnarToRowExec(child) => child @@ -192,40 +253,12 @@ case class RevertNativeForTransitionHeavyStages(session: SparkSession, wholePlan case sparkToColumnar: CometSparkToColumnarExec => sparkToColumnar.child case RowToColumnarExec(child) => child } - val reverted = transformStageUp(stripped) { - // Local candidate selection was inserted by Comet. Only the outer TopK owns - // the original Spark operator's offset and projection. - case local: CometLocalTopKExec => local.child - case cometExec: CometExec => - if (cometExec.originalPlan.children.size == cometExec.children.size) { - val originalWithCurrentExpressions = cometExec match { - case scan: CometNativeScanExec => - // AQE's query-stage optimizer rewrites DPP placeholders in the live Comet scan - // before this post-columnar rule runs. The frozen FileSourceScanExec in - // originalPlan still contains SubqueryAdaptiveBroadcastExec, which cannot execute. - // Preserve the rewritten filters when reverting the scan to Spark. - val originalScan = scan.originalPlan.copy( - partitionFilters = scan.partitionFilters, - dataFilters = scan.dataFilters) - scan.originalPlan.logicalLink.foreach(originalScan.setLogicalLink) - originalScan - case scan: CometIcebergNativeScanExec => - // Iceberg's native scan has the same split between live and frozen filters. - // serializedPartitionData rebuilds originalPlan from runtimeFilters before - // execution, but transition reversion skips that path and executes the restored - // BatchScanExec directly. Carry the executable DPP filters across here as well. - val originalScan = scan.originalPlan.copy(runtimeFilters = scan.runtimeFilters) - scan.originalPlan.logicalLink.foreach(originalScan.setLogicalLink) - originalScan - case _ => cometExec.originalPlan - } - originalWithCurrentExpressions.withNewChildren(cometExec.children) - } else { - logWarning( - "Comet plan and original have different child count for " + - s"${cometExec.getClass.getSimpleName}, using originalPlan as-is.") - cometExec.originalPlan - } + if (isStageBoundary(stripped)) { + throw new CometExec.InvalidSparkFallbackException( + "Cannot revert a stage whose stripped root is a stage boundary") + } + val reverted = transformStageUp(stripped) { case cometExec: CometExec => + cometExec.sparkFallback(cometExec.children) } insertTransitions(reverted) } diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometDataWritingCommand.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometDataWritingCommand.scala index 6af2cc18ba..a56dea862c 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometDataWritingCommand.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometDataWritingCommand.scala @@ -211,7 +211,7 @@ object CometDataWritingCommand extends CometOperatorSerde[DataWritingCommandExec throw new SparkException(s"Could not instantiate FileCommitProtocol: ${e.getMessage}") } - CometNativeWriteExec(nativeOp, childPlan, outputPath, cmd.mode, committer, jobId) + CometNativeWriteExec(nativeOp, op, childPlan, outputPath, cmd.mode, committer, jobId) } } diff --git a/spark/src/main/scala/org/apache/comet/serde/operator/CometIcebergNativeWrite.scala b/spark/src/main/scala/org/apache/comet/serde/operator/CometIcebergNativeWrite.scala index 5d5510fb7a..69fbec350e 100644 --- a/spark/src/main/scala/org/apache/comet/serde/operator/CometIcebergNativeWrite.scala +++ b/spark/src/main/scala/org/apache/comet/serde/operator/CometIcebergNativeWrite.scala @@ -647,7 +647,9 @@ object CometIcebergNativeWrite extends CometOperatorSerde[IcebergWriteExec] { "Native Iceberg write conversion: SparkWrite.outputSpecId reflection failed")) CometIcebergWriteExec( nativeOp, + op, op.child, + op.output, op.batchWrite, table.asInstanceOf[AnyRef], outputSpecId) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala index c944d3af8c..e058a5c091 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergNativeScanExec.scala @@ -27,7 +27,7 @@ import org.apache.spark.sql.SparkSession import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, SortOrder} import org.apache.spark.sql.catalyst.plans.QueryPlan import org.apache.spark.sql.catalyst.plans.physical.{Partitioning, UnknownPartitioning} -import org.apache.spark.sql.execution.SQLExecution +import org.apache.spark.sql.execution.{SparkPlan, SQLExecution} import org.apache.spark.sql.execution.datasources.v2.BatchScanExec import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} import org.apache.spark.sql.vectorized.ColumnarBatch @@ -65,6 +65,14 @@ case class CometIcebergNativeScanExec( @transient nativeIcebergScanMetadata: CometIcebergNativeScanMetadata) extends CometLeafExec { + override def sparkFallback(newChildren: Seq[SparkPlan]): SparkPlan = { + // Native execution rebuilds originalPlan from the live runtimeFilters during partition + // serialization. Reversion skips that path, so carry the executable DPP filters across here. + val restoredScan = originalPlan.copy(runtimeFilters = runtimeFilters) + originalPlan.logicalLink.foreach(restoredScan.setLogicalLink) + restoredScan + } + override val supportsColumnar: Boolean = true override val nodeName: String = "CometIcebergNativeScan" diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala index 877eb84c1f..83149ced77 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometIcebergWriteExec.scala @@ -24,14 +24,13 @@ import java.util.concurrent.atomic.AtomicReference import org.apache.spark.TaskContext import org.apache.spark.rdd.RDD import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference} +import org.apache.spark.sql.catalyst.expressions.Attribute import org.apache.spark.sql.catalyst.expressions.UnsafeProjection import org.apache.spark.sql.comet.execution.arrow.CometArrowStream import org.apache.spark.sql.comet.util.{Utils => CometUtils} import org.apache.spark.sql.connector.write.{BatchWrite, WriterCommitMessage} import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} -import org.apache.spark.sql.types.BinaryType import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.spark.util.TaskFailureListener @@ -54,8 +53,12 @@ import org.apache.comet.serde.OperatorOuterClass.Operator * @param nativeOp * Template operator carrying the `IcebergWrite` proto. Per-task `partition_id` / * `task_attempt_id` get stamped on a copy at execution time. + * @param originalPlan + * The JVM Iceberg write operator restored when Comet reverts a transition-heavy stage. * @param child * Comet native child (must be a [[CometNativeExec]] so columnar batches flow through FFI). + * @param output + * Commit-message output attributes copied from the JVM Iceberg write plan. * @param batchWrite * Shared with the outer [[IcebergCommitExec]] -- the same instance the strategy materialised * via `write.toBatch`. Used here only to provide the `dataLocation` / partition spec needed by @@ -66,20 +69,15 @@ import org.apache.comet.serde.OperatorOuterClass.Operator */ case class CometIcebergWriteExec( nativeOp: Operator, + @transient override val originalPlan: IcebergWriteExec, child: SparkPlan, + override val output: Seq[Attribute], @transient batchWrite: BatchWrite, @transient table: AnyRef, partitionSpecId: Int) extends CometNativeExec with UnaryExecNode { - override def originalPlan: SparkPlan = child - - // Same output schema as IcebergWriteExec so the outer IcebergCommitExec consumes the - // commit messages identically regardless of which inner exec emitted them. - override def output: Seq[Attribute] = Seq( - AttributeReference(IcebergWriteExec.CommitMessageColumn, BinaryType, nullable = false)()) - // Native exec emits a single Binary column; the surrounding command framework expects rows, so // the outer commit exec calls executeCollect on us. supportsColumnar = false keeps Spark from // inserting a ColumnarToRow that would clash with our (Nil-output-like) row contract. diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTopKExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTopKExec.scala index ec3c85191a..16ecb20881 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTopKExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometLocalTopKExec.scala @@ -91,6 +91,15 @@ case class CometLocalTopKExec( override val serializedPlanOpt: SerializedPlan) extends CometUnaryExec { + override def sparkFallback(newChildren: Seq[SparkPlan]): SparkPlan = newChildren match { + // Local candidate selection was inserted by Comet. Only the outer TopK owns + // the original Spark operator's offset and projection. + case Seq(restoredChild) => restoredChild + case _ => + throw new CometExec.InvalidSparkFallbackException( + s"CometLocalTopKExec expected one restored child but received ${newChildren.size}") + } + override def outputPartitioning: Partitioning = child.outputPartitioning override def outputOrdering: Seq[SortOrder] = sortOrder diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala index 8ea48c1002..3769499221 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeScanExec.scala @@ -75,6 +75,15 @@ case class CometNativeScanExec( with ShimStreamSourceAwareSparkPlan with CometScanWithPlanData { + override def sparkFallback(newChildren: Seq[SparkPlan]): SparkPlan = { + // AQE rewrites DPP placeholders in the live scan before transition reversion. + // The frozen originalPlan can still contain unexecutable SubqueryAdaptiveBroadcastExecs. + val restoredScan = + originalPlan.copy(partitionFilters = partitionFilters, dataFilters = dataFilters) + originalPlan.logicalLink.foreach(restoredScan.setLogicalLink) + restoredScan + } + override lazy val metadata: Map[String, String] = if (originalPlan != null) originalPlan.metadata else Map.empty diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala index ced525c186..bb2c560ba6 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/CometNativeWriteExec.scala @@ -29,10 +29,13 @@ import org.apache.spark.internal.io.{FileCommitProtocol, FileNameSpec} import org.apache.spark.rdd.RDD import org.apache.spark.sql.SaveMode import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.Attribute import org.apache.spark.sql.comet.execution.arrow.CometArrowStream import org.apache.spark.sql.comet.util.{Utils => CometUtils} import org.apache.spark.sql.errors.{QueryCompilationErrors, QueryExecutionErrors} import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} +import org.apache.spark.sql.execution.command.DataWritingCommandExec +import org.apache.spark.sql.execution.datasources.WriteFilesExec import org.apache.spark.sql.execution.metric.{SQLMetric, SQLMetrics} import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.spark.util.Utils @@ -54,6 +57,8 @@ import org.apache.comet.serde.OperatorOuterClass.Operator * * @param nativeOp * The native operator representing the write operation (template, will be modified per task) + * @param originalPlan + * The JVM data-writing command restored when Comet reverts a transition-heavy stage * @param child * The child operator providing the data to write * @param outputPath @@ -69,6 +74,7 @@ import org.apache.comet.serde.OperatorOuterClass.Operator */ case class CometNativeWriteExec( nativeOp: Operator, + @transient override val originalPlan: DataWritingCommandExec, child: SparkPlan, outputPath: String, mode: SaveMode, @@ -77,7 +83,19 @@ case class CometNativeWriteExec( extends CometNativeExec with UnaryExecNode { - override def originalPlan: SparkPlan = child + override def output: Seq[Attribute] = child.output + + override def sparkFallback(newChildren: Seq[SparkPlan]): SparkPlan = newChildren match { + case Seq(newInput) => + val restoredCommandChild = originalPlan.child match { + case writeFiles: WriteFilesExec => writeFiles.withNewChildren(Seq(newInput)) + case _ => newInput + } + originalPlan.withNewChildren(Seq(restoredCommandChild)) + case _ => + throw new CometExec.InvalidSparkFallbackException( + s"${getClass.getSimpleName} expected one reverted input but received ${newChildren.size}") + } // Accumulator to collect TaskCommitMessages from all tasks // Must be eagerly initialized on driver, not lazy diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala index 909539bf2b..e1de72ac4f 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala @@ -661,6 +661,30 @@ abstract class CometExec extends CometPlan { /** The original Spark operator from which this Comet operator is converted from */ def originalPlan: SparkPlan + /** + * Rebuilds the Spark operator represented by this Comet operator with reverted children. + * + * Operators whose Comet representation changes the Spark plan shape or whose live state differs + * from the original Spark plan must override this method. + */ + def sparkFallback(newChildren: Seq[SparkPlan]): SparkPlan = { + val sparkPlan = originalPlan + if (sparkPlan == null) { + throw new CometExec.InvalidSparkFallbackException( + s"${getClass.getSimpleName} has no original Spark plan") + } + if (newChildren.exists(_ eq sparkPlan)) { + throw new CometExec.InvalidSparkFallbackException( + s"${getClass.getSimpleName} aliases its original Spark plan with a child") + } + if (sparkPlan.children.size != newChildren.size) { + throw new CometExec.InvalidSparkFallbackException( + s"${getClass.getSimpleName} cannot restore ${sparkPlan.getClass.getSimpleName}: " + + s"expected ${sparkPlan.children.size} children but received ${newChildren.size}") + } + sparkPlan.withNewChildren(newChildren) + } + /** Comet always support columnar execution */ override def supportsColumnar: Boolean = true @@ -711,6 +735,9 @@ abstract class CometExec extends CometPlan { } object CometExec { + final class InvalidSparkFallbackException(message: String) + extends IllegalArgumentException(message) + // An unique id for each CometExecIterator, used to identify the native query execution. private val curId = new java.util.concurrent.atomic.AtomicLong() diff --git a/spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala b/spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala index b9eeadbc69..7692e8ce96 100644 --- a/spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometIcebergWriteActionSuite.scala @@ -51,6 +51,7 @@ import org.apache.spark.sql.connector.catalog.InMemoryTableCatalog import org.apache.spark.sql.connector.write.{BatchWrite, DataWriterFactory, PhysicalWriteInfo, Write, WriterCommitMessage} import org.apache.spark.sql.execution.{ColumnarToRowTransition, LeafExecNode, SparkPlan} import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.exchange.ShuffleExchangeLike import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.streaming.Trigger import org.apache.spark.sql.types.{DoubleType, IntegerType, StringType, StructField, StructType} @@ -807,6 +808,114 @@ class CometIcebergWriteActionSuite } } + for (adaptive <- Seq(false, true)) { + test(s"transition-heavy fallback preserves Iceberg writes with AQE=$adaptive") { + assumeNativeAcceleration() + withIcebergCatalog { warehouseDir => + val suffix = if (adaptive) "aqe" else "no_aqe" + val nativeTable = s"transition_native_$suffix" + val fallbackTable = s"transition_fallback_$suffix" + createTable(warehouseDir, nativeTable, partitionSpec = "") + createTable(warehouseDir, fallbackTable, partitionSpec = "") + val values = "(1, 'us-east', 10.5), (2, 'us-west', 20.3), (3, 'eu', 30.7)" + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { + assertNativeWriteEngages(nativeTable, Seq(1, 2, 3)) { + spark.sql(s"INSERT INTO $catalog.$ns.$nativeTable VALUES $values") + } + } + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + assertNativeWriteDoesNotEngage(fallbackTable, Seq(1, 2, 3)) { + spark.sql(s"INSERT INTO $catalog.$ns.$fallbackTable VALUES $values") + } + } + } + } + } + + // An unpartitioned INSERT never puts an exchange under the write, so it misses #6152: with AQE + // off, unwrapping a transition that sits on a shuffle used to keep walking into the map stage. + // The IN-subquery on a partitioned copy-on-write DELETE is the plan that does, and the AQE-off + // case is the one that fails with `ColumnarBatch cannot be cast to InternalRow`. + // `us-west` keeps a row after deleting id 2. An emptied partition makes the rewrite emit + // nothing, and AQE then replaces the write input with an empty LocalTableScan, so the executed + // plan no longer contains the shuffle this test is checking for. + for (adaptive <- Seq(false, true)) { + test(s"transition-heavy fallback preserves partitioned CoW deletes with AQE=$adaptive") { + assumeNativeAcceleration() + withIcebergCatalog { warehouseDir => + val suffix = if (adaptive) "aqe" else "no_aqe" + val nativeTable = s"transition_cow_native_$suffix" + val fallbackTable = s"transition_cow_fallback_$suffix" + val props = Some("'write.delete.mode'='copy-on-write'") + val spec = "PARTITIONED BY (region)" + createTable(warehouseDir, nativeTable, spec, props) + createTable(warehouseDir, fallbackTable, spec, props) + val seed = + Seq( + (1, "us-east", 10.0), + (2, "us-west", 20.0), + (3, "eu", 30.0), + (4, "us-east", 40.0), + (5, "us-west", 50.0)) + withSQLConf(CometConf.COMET_ICEBERG_WRITE_SPLIT_OPERATOR_ENABLED.key -> "false") { + coalesceInsert(nativeTable, seed) + coalesceInsert(fallbackTable, seed) + } + def delete(table: String): Unit = + spark.sql( + s"DELETE FROM $catalog.$ns.$table WHERE id IN " + + "(SELECT col1 FROM VALUES (2) AS t(col1))") + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { + assertNativeWriteEngages(nativeTable, Seq(1, 3, 4, 5)) { + delete(nativeTable) + } + } + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + val snapshot = withNativeEnabled { + captureWrite(fallbackTable)(delete(fallbackTable)) + } + assertExactlyOneCommit(snapshot) + val nativeExecs = snapshot.plans.flatMap { plan => + collectWithSubqueries(plan) { case e: CometIcebergWriteExec => e } + } + assert( + nativeExecs.isEmpty, + "transition reversion should restore IcebergWriteExec. Plans:\n" + + snapshot.plans.mkString("\n--\n")) + // AdaptiveSparkPlanExec is a leaf, so only the AQE-aware collect reaches the exchange. + val hasExchange = snapshot.plans.exists { plan => + collect(plan) { case _: ShuffleExchangeLike => true }.nonEmpty + } + assert( + hasExchange, + "the delete must keep a shuffle under the write. Plans:\n" + + snapshot.plans.mkString("\n--\n")) + assertRows(fallbackTable, Seq(1, 3, 4, 5)) + } + + val nativeDirs = partitionDirs(warehouseDir, nativeTable) + val fallbackDirs = partitionDirs(warehouseDir, fallbackTable) + assert( + fallbackDirs == nativeDirs, + s"partition layout fallback=$fallbackDirs native=$nativeDirs") + } + } + } + // What Iceberg's Spark writer stamps for `sort_order_id` on appended files changed across // releases: through 1.10 `SparkWrite$WriterFactory` never wires the table sort order (files // get 0 even on a sorted table); 1.11 added `SparkWriteConf.outputSortOrderId` and stamps the diff --git a/spark/src/test/scala/org/apache/comet/CometIcebergWriteDetectionSuite.scala b/spark/src/test/scala/org/apache/comet/CometIcebergWriteDetectionSuite.scala index bed4e85d55..ba0537e1a7 100644 --- a/spark/src/test/scala/org/apache/comet/CometIcebergWriteDetectionSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometIcebergWriteDetectionSuite.scala @@ -38,7 +38,7 @@ import org.apache.spark.sql.catalyst.InternalRow import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeReference} import org.apache.spark.sql.comet.{CometIcebergWriteExec, CometSparkToColumnarExec, IcebergWriteExec} import org.apache.spark.sql.execution.{ApplyColumnarRulesAndInsertTransitions, ColumnarToRowExec, CommandExecutionMode, LeafExecNode, SparkPlan} -import org.apache.spark.sql.types.IntegerType +import org.apache.spark.sql.types.{BinaryType, IntegerType} import org.apache.spark.sql.vectorized.ColumnarBatch import org.apache.comet.CometSparkSessionExtensions.isSpark35Plus @@ -1252,9 +1252,15 @@ class CometIcebergWriteDetectionSuite extends CometTestBase with CometIcebergTes * directly. */ private def writeChildAfterTransitionRules(source: SparkPlan): SparkPlan = { + val child = CometSparkToColumnarExec(source) + val output = Seq( + AttributeReference(IcebergWriteExec.CommitMessageColumn, BinaryType, nullable = false)()) + val originalPlan = IcebergWriteExec(null, output, child) val write = CometIcebergWriteExec( Operator.newBuilder().build(), - CometSparkToColumnarExec(source), + originalPlan, + child, + output, batchWrite = null, table = null, partitionSpecId = 0) diff --git a/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala b/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala index babe86b3c7..adc00d33b2 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala @@ -35,15 +35,17 @@ import org.apache.parquet.schema.{MessageType, Type} import org.apache.spark.internal.io.FileCommitProtocol import org.apache.spark.sql.{AnalysisException, DataFrame, Row, SaveMode} import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.comet.{CometBatchScanExec, CometNativeScanExec, CometScanExec, CometWriteFilesExec} -import org.apache.spark.sql.execution.{FileSourceScanExec, SparkPlan} -import org.apache.spark.sql.execution.datasources.{BasicWriteTaskStats, SQLHadoopMapReduceCommitProtocol, WriteTaskStats, WriteTaskStatsTracker} +import org.apache.spark.sql.comet.{CometBatchScanExec, CometNativeColumnarToRowExec, CometNativeScanExec, CometNativeWriteExec, CometScanExec, CometSparkToColumnarExec, CometWriteFilesExec} +import org.apache.spark.sql.execution.{ColumnarToRowTransition, FileSourceScanExec, SparkPlan, SQLExecution} +import org.apache.spark.sql.execution.command.DataWritingCommandExec +import org.apache.spark.sql.execution.datasources.{BasicWriteTaskStats, SQLHadoopMapReduceCommitProtocol, WriteFilesExec, WriteTaskStats, WriteTaskStatsTracker} import org.apache.spark.sql.functions.{array, col, map, struct, when} import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types.{ArrayType, LongType, MapType, Metadata, MetadataBuilder, StringType, StructField, StructType} import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, isSpark40Plus} +import org.apache.comet.rules.RevertNativeForTransitionHeavyStages import org.apache.comet.serde.operator.NativeWriteUtils import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator, SchemaGenOptions} @@ -126,6 +128,161 @@ class CometParquetWriterSuite extends CometParquetWriterTestBase { } } + for (adaptive <- Seq(false, true)) { + test(s"transition-heavy fallback restores the Spark write plan with AQE=$adaptive") { + // Force a transition below the writer on both paths: Spark 3.x replaces the command, + // whereas Spark 4.x replaces only its WriteFilesExec child. + withTempPath { dir => + withTempPath { inputDir => + val inputPath = createTestData(inputDir) + val nativeOutput = new File(dir, "native-output").getAbsolutePath + val commonConf = Seq( + CometConf.COMET_NATIVE_PARQUET_WRITE_ENABLED.key -> "true", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "America/Halifax", + CometConf.COMET_OPERATOR_DATA_WRITING_COMMAND_ALLOW_INCOMPAT.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString) + + withSQLConf( + (commonConf :+ + (CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false")): _*) { + val df = spark.read.parquet(inputPath) + val nativePlan = captureWritePlan(path => df.write.parquet(path), nativeOutput) + assertHasCometNativeWriteExec(nativePlan) + verifyWrittenFile(nativeOutput) + val nativeWrite: SparkPlan = nativePlan + .collectFirst { + case write: CometNativeWriteExec => write + case write: CometWriteFilesExec => write + } + .getOrElse(fail(s"expected a native parquet writer:\n$nativePlan")) + val stageWithTransition = nativeWrite.withNewChildren(Seq( + CometSparkToColumnarExec(CometNativeColumnarToRowExec(nativeWrite.children.head)))) + assert( + stageWithTransition.collect { case _: ColumnarToRowTransition => true }.nonEmpty, + s"test requires a C2R transition:\n$stageWithTransition") + + var fallbackPlan: SparkPlan = null + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + val stagePlan = if (isSpark40Plus) { + val command = nativePlan + .collectFirst { case node: DataWritingCommandExec => node } + .getOrElse(fail(s"expected a Spark write command:\n$nativePlan")) + command.withNewChildren(Seq(stageWithTransition)) + } else { + stageWithTransition + } + fallbackPlan = RevertNativeForTransitionHeavyStages(spark)(stagePlan) + } + assertNoCometNativeWriteExec(fallbackPlan) + val commands = fallbackPlan.collect { case command: DataWritingCommandExec => + command + } + assert( + commands.exists(_.child.isInstanceOf[WriteFilesExec]), + s"expected DataWritingCommandExec -> WriteFilesExec after fallback:\n$fallbackPlan") + + deletePath(nativeOutput) + SQLExecution.withNewExecutionId(spark.range(0).queryExecution) { + fallbackPlan.executeCollect() + } + verifyWrittenFile(nativeOutput) + } + } + } + } + } + + for (adaptive <- Seq(false, true)) { + test( + s"transition-heavy fallback restores parquet writes through the query pipeline with AQE=$adaptive") { + withTempPath { dir => + val nativeOutput = new File(dir, "native-output").getAbsolutePath + val fallbackOutput = new File(dir, "fallback-output").getAbsolutePath + // Row source, matching the Iceberg VALUES e2e: native write engages via SparkToColumnar. + val commonConf = Seq( + CometConf.COMET_NATIVE_PARQUET_WRITE_ENABLED.key -> "true", + CometConf.COMET_SPARK_TO_ARROW_ENABLED.key -> "true", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "America/Halifax", + CometConf.COMET_OPERATOR_DATA_WRITING_COMMAND_ALLOW_INCOMPAT.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString) + + def rangeWrite(path: String): Unit = + spark.range(0, 1000, 1, numPartitions = 4).toDF("id").write.parquet(path) + + val sparkOutput = new File(dir, "spark-output").getAbsolutePath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + rangeWrite(sparkOutput) + } + + withSQLConf( + (commonConf :+ + (CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false")): _*) { + val nativePlan = captureWritePlan(rangeWrite, nativeOutput) + assertHasCometNativeWriteExec(nativePlan) + assertWrittenRows(nativeOutput, sparkOutput, 1000) + } + + withSQLConf( + (commonConf ++ Seq( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0")): _*) { + val fallbackPlan = captureWritePlan(rangeWrite, fallbackOutput) + assertTransitionHeavyParquetFallback(fallbackPlan) + assertWrittenRows(fallbackOutput, sparkOutput, 1000) + } + } + } + } + + for (adaptive <- Seq(false, true)) { + test(s"transition-heavy fallback restores a union of row sources with AQE=$adaptive") { + withTempPath { dir => + val nativeOutput = new File(dir, "native-union-output").getAbsolutePath + val fallbackOutput = new File(dir, "fallback-union-output").getAbsolutePath + val commonConf = Seq( + CometConf.COMET_NATIVE_PARQUET_WRITE_ENABLED.key -> "true", + CometConf.COMET_SPARK_TO_ARROW_ENABLED.key -> "true", + SQLConf.SESSION_LOCAL_TIMEZONE.key -> "America/Halifax", + CometConf.COMET_OPERATOR_DATA_WRITING_COMMAND_ALLOW_INCOMPAT.key -> "true", + CometConf.COMET_EXEC_ENABLED.key -> "true", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString) + + def unionedWrite(path: String): Unit = { + val left = spark.range(0, 1000, 1, numPartitions = 2).toDF("id") + val mid = spark.range(1000, 2000, 1, numPartitions = 2).toDF("id") + val right = spark.range(2000, 3000, 1, numPartitions = 2).toDF("id") + left.union(mid).union(right).write.parquet(path) + } + + val sparkOutput = new File(dir, "spark-output").getAbsolutePath + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + unionedWrite(sparkOutput) + } + + withSQLConf( + (commonConf :+ + (CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false")): _*) { + val nativePlan = captureWritePlan(unionedWrite, nativeOutput) + assertHasCometNativeWriteExec(nativePlan) + assertWrittenRows(nativeOutput, sparkOutput, 3000) + } + + withSQLConf( + (commonConf ++ Seq( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0")): _*) { + val fallbackPlan = captureWritePlan(unionedWrite, fallbackOutput) + assertTransitionHeavyParquetFallback(fallbackPlan) + assertWrittenRows(fallbackOutput, sparkOutput, 3000) + } + } + } + } + test("basic parquet write with repartition") { withTempPath { dir => // Create test data and write it to a temp parquet file first @@ -1514,6 +1671,58 @@ class CometParquetWriterSuite extends CometParquetWriterTestBase { } } + private def deletePath(path: String): Unit = { + def delete(file: File): Unit = { + if (file.isDirectory) { + Option(file.listFiles()).foreach(_.foreach(delete)) + } + file.delete() + } + delete(new File(path)) + } + + /** + * Spark 3.x replaces the whole command with `CometNativeWriteExec`, and Spark then inserts a + * C2R above that columnar node, so `maxTransitions = 0` reverts the stage back to + * `DataWritingCommandExec` -> `WriteFilesExec`. Spark 4.0+ leaves the command in place and does + * not insert a C2R above `CometWriteFilesExec`. These row sources only add `SparkToColumnar`, + * which is an R2C, so the stage is not transition-heavy and the native write stays. + */ + private def assertTransitionHeavyParquetFallback(plan: SparkPlan): Unit = { + if (isSpark40Plus) { + assertHasCometNativeWriteExec(plan) + } else { + assertNoCometNativeWriteExec(plan) + assertRestoredParquetWriteCommand(plan) + } + } + + private def assertRestoredParquetWriteCommand(plan: SparkPlan): Unit = { + assert( + plan + .collect { case command: DataWritingCommandExec => command } + .exists(_.child.isInstanceOf[WriteFilesExec]), + s"expected DataWritingCommandExec -> WriteFilesExec after fallback:\n$plan") + } + + // Compare every row with Spark-written data using the row-based parquet reader, so this + // check is independent of Comet scans and detects duplicates or missing union inputs. + private def assertWrittenRows( + outputPath: String, + sparkOutputPath: String, + expectedRows: Int): Unit = { + val outputDir = new File(outputPath) + val partFiles = outputDir.listFiles().filter(_.getName.startsWith("part-")) + assert(partFiles.length > 1, s"Expected multiple part files under $outputPath") + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + SQLConf.PARQUET_VECTORIZED_READER_ENABLED.key -> "false") { + val expected = spark.read.parquet(sparkOutputPath).collect().toSeq + assert(expected.size == expectedRows) + checkAnswer(spark.read.parquet(outputPath), expected) + } + } + private def writeWithCometNativeWriteExec( inputPath: String, outputPath: String, @@ -1529,10 +1738,10 @@ class CometParquetWriterSuite extends CometParquetWriterTestBase { Some(plan) } - private def verifyWrittenFile(outputPath: String): Unit = { + private def verifyWrittenFile(outputPath: String, expectedRows: Int = 1000): Unit = { // Verify the data was written correctly val resultDf = spark.read.parquet(outputPath) - assert(resultDf.count() == 1000, "Expected 1000 rows to be written") + assert(resultDf.count() == expectedRows, s"Expected $expectedRows rows to be written") // Verify multiple part files were created val outputDir = new File(outputPath) diff --git a/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala b/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala index e350badf8d..d0ca67c60d 100644 --- a/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala +++ b/spark/src/test/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStagesSuite.scala @@ -19,19 +19,122 @@ package org.apache.comet.rules -import org.apache.spark.sql.{CometTestBase, Row} +import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.AtomicReference + +import org.apache.spark.sql.{CometTestBase, Row, SaveMode} +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, Literal} import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial} import org.apache.spark.sql.comet._ import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.QueryStageExec +import org.apache.spark.sql.execution.adaptive.{AQEShuffleReadExec, QueryStageExec, ShuffleQueryStageExec} +import org.apache.spark.sql.execution.command.DataWritingCommandExec +import org.apache.spark.sql.execution.datasources.WriteFilesExec import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.BinaryType +import org.apache.spark.sql.util.QueryExecutionListener import org.apache.comet.CometConf import org.apache.comet.CometSparkSessionExtensions.isSpark35Plus +import org.apache.comet.serde.OperatorOuterClass.Operator + +private case class AliasingFallbackCometExec( + override val originalPlan: SparkPlan, + child: SparkPlan) + extends CometExec + with UnaryExecNode { + override protected def withNewChildInternal(newChild: SparkPlan): SparkPlan = + copy(child = newChild) +} class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { + private def cometIcebergWrite(child: SparkPlan): CometIcebergWriteExec = { + val output = Seq( + AttributeReference(IcebergWriteExec.CommitMessageColumn, BinaryType, nullable = false)()) + val originalPlan = IcebergWriteExec(null, output, child) + CometIcebergWriteExec( + Operator.newBuilder().build(), + originalPlan, + child, + output, + batchWrite = null, + table = null, + partitionSpecId = 0) + } + + private def cometFilter(child: SparkPlan): CometFilterExec = { + val condition = Literal.TrueLiteral + val sparkFilter = FilterExec(condition, child) + CometFilterExec( + Operator.newBuilder().build(), + sparkFilter, + sparkFilter.output, + condition, + child, + SerializedPlan(None)) + } + + private def captureDataWritingCommand(path: String): DataWritingCommandExec = { + val captured = new AtomicReference[SparkPlan]() + val callbackCompleted = new CountDownLatch(1) + val listener = new QueryExecutionListener { + override def onSuccess(funcName: String, qe: QueryExecution, durationNs: Long): Unit = { + if (funcName == "save" || funcName.contains("command")) { + captured.set(qe.executedPlan) + callbackCompleted.countDown() + } + } + override def onFailure( + funcName: String, + qe: QueryExecution, + exception: Exception): Unit = {} + } + spark.listenerManager.register(listener) + try { + withSQLConf(CometConf.COMET_ENABLED.key -> "false") { + spark.range(1).toDF("id").write.mode("overwrite").parquet(path) + } + assert( + callbackCompleted.await(10, TimeUnit.SECONDS), + "timed out waiting to capture the parquet write plan") + } finally { + spark.listenerManager.unregister(listener) + } + val plan = stripAQEPlan( + Option(captured.get()).getOrElse(fail("expected a captured parquet write plan"))) + plan + .collectFirst { case command: DataWritingCommandExec => command } + .getOrElse(fail(s"expected DataWritingCommandExec:\n$plan")) + } + + private def cometNativeWrite( + child: SparkPlan, + command: DataWritingCommandExec): CometNativeWriteExec = { + CometNativeWriteExec( + Operator.newBuilder().build(), + command, + child, + outputPath = "/tmp/unused-native-write", + mode = SaveMode.Overwrite) + } + + private def assertRestoredParquetWrite(reverted: SparkPlan): WriteFilesExec = { + val command = reverted match { + case node: DataWritingCommandExec => node + case other => fail(s"expected DataWritingCommandExec, got:\n$other") + } + val writeFiles = command.child match { + case node: WriteFilesExec => node + case other => fail(s"expected WriteFilesExec under DataWritingCommandExec, got:\n$other") + } + assert( + command.collect { case _: CometNativeWriteExec => true }.isEmpty, + s"native parquet write should be restored, not erased:\n$reverted") + writeFiles + } + private def createSparkPlan(sql: String): SparkPlan = { var plan: SparkPlan = null withSQLConf(CometConf.COMET_ENABLED.key -> "false") { @@ -59,6 +162,12 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { plan.collect { case _: ColumnarToRowTransition => true }.size } + private def unwrapCodegen(plan: SparkPlan): SparkPlan = plan match { + case wholeStage: WholeStageCodegenExec => unwrapCodegen(wholeStage.child) + case inputAdapter: InputAdapter => unwrapCodegen(inputAdapter.child) + case other => other + } + private def collectCometAggregates(plan: SparkPlan): Seq[CometHashAggregateExec] = { val current = plan match { case aggregate: CometHashAggregateExec => Seq(aggregate) @@ -142,6 +251,207 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { } } + test("revertToSpark preserves an Iceberg write with a leaf child") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + val write = cometIcebergWrite(leaf) + + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + val icebergWrite = reverted match { + case node: IcebergWriteExec => node + case other => fail(s"expected IcebergWriteExec, got:\n$other") + } + + assert(icebergWrite.child eq leaf) + assert(icebergWrite.output.map(_.name) == Seq(IcebergWriteExec.CommitMessageColumn)) + assert(icebergWrite.output.map(_.dataType) == Seq(BinaryType)) + } + + test("revertToSpark preserves an Iceberg write over SparkToColumnar of a row source") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + val write = cometIcebergWrite(CometSparkToColumnarExec(leaf)) + + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + val icebergWrite = reverted match { + case node: IcebergWriteExec => node + case other => fail(s"expected IcebergWriteExec, got:\n$other") + } + + assert( + icebergWrite.child eq leaf, + s"SparkToColumnar should unwrap to the row source:\n$reverted") + assert(icebergWrite.output.map(_.name) == Seq(IcebergWriteExec.CommitMessageColumn)) + assert(icebergWrite.output.map(_.dataType) == Seq(BinaryType)) + assert( + reverted.collect { case _: CometSparkToColumnarExec => true }.isEmpty, + s"SparkToColumnar should be fully unwrapped:\n$reverted") + } + + test("revertToSpark preserves an Iceberg write without duplicating its unary child") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1), (2) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + val write = cometIcebergWrite(cometFilter(leaf)) + + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + + assert(reverted.isInstanceOf[IcebergWriteExec], s"expected IcebergWriteExec:\n$reverted") + assert( + reverted.collect { case _: FilterExec => true }.size == 1, + s"expected exactly one Spark FilterExec:\n$reverted") + assert(countCometExecs(reverted) == 0, s"expected no Comet operators:\n$reverted") + } + + test("revertToSpark unwraps stacked SparkToColumnar(C2R) under a native write") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + val stacked = CometSparkToColumnarExec(CometNativeColumnarToRowExec(cometFilter(leaf))) + val write = cometIcebergWrite(stacked) + + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + + assert(reverted.isInstanceOf[IcebergWriteExec], s"expected IcebergWriteExec:\n$reverted") + assert( + reverted.collect { case _: FilterExec => true }.size == 1, + s"expected exactly one Spark FilterExec:\n$reverted") + assert( + reverted.collect { case _: CometNativeColumnarToRowExec | _: CometSparkToColumnarExec => + true + }.isEmpty, + s"stacked transitions should be fully unwrapped:\n$reverted") + assert(countCometExecs(reverted) == 0, s"expected no Comet operators:\n$reverted") + } + + test("revertToSpark preserves a native parquet write with a leaf child") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + withTempPath { dir => + val write = cometNativeWrite(leaf, captureDataWritingCommand(dir.getAbsolutePath)) + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + val writeFiles = assertRestoredParquetWrite(reverted) + assert(writeFiles.child eq leaf, s"expected the original leaf child:\n$reverted") + } + } + + test("revertToSpark preserves a native parquet write without duplicating its unary child") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1), (2) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + withTempPath { dir => + val write = + cometNativeWrite(cometFilter(leaf), captureDataWritingCommand(dir.getAbsolutePath)) + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + assertRestoredParquetWrite(reverted) + assert( + reverted.collect { case _: FilterExec => true }.size == 1, + s"expected exactly one Spark FilterExec:\n$reverted") + assert(countCometExecs(reverted) == 0, s"expected no Comet operators:\n$reverted") + } + } + + test("revertToSpark restores a native parquet write whose command has no WriteFilesExec") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + withTempPath { dir => + val command = captureDataWritingCommand(dir.getAbsolutePath) + val input = command.child match { + case writeFiles: WriteFilesExec => writeFiles.child + case other => other + } + val commandWithoutWriteFiles = + command.withNewChildren(Seq(input)).asInstanceOf[DataWritingCommandExec] + val write = cometNativeWrite(leaf, commandWithoutWriteFiles) + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + val restored = reverted match { + case node: DataWritingCommandExec => node + case other => fail(s"expected DataWritingCommandExec, got:\n$other") + } + assert(restored.child eq leaf, s"expected the original leaf child:\n$reverted") + assert( + restored.collect { case _: WriteFilesExec => true }.isEmpty, + s"WriteFilesExec should not be reinserted when the original command lacked it:\n$reverted") + assert( + restored.collect { case _: CometNativeWriteExec => true }.isEmpty, + s"native parquet write should be restored, not erased:\n$reverted") + } + } + + test("invalid original-plan alias skips the entire stage reversion") { + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + val aliasing = AliasingFallbackCometExec(leaf, leaf) + val stagePlan = CometNativeColumnarToRowExec(aliasing) + val rule = RevertNativeForTransitionHeavyStages(spark) + assert(rule.countTransitions(stagePlan) == 1) + + val result = rule(stagePlan) + + assert( + result eq stagePlan, + s"invalid fallback must leave the whole stage unchanged:\n$result") + } + } + + test("invalid original-plan alias to a Comet child skips the entire stage reversion") { + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + val sparkPlan = createSparkPlan("SELECT id FROM VALUES (1) AS t(id)") + val leaf = sparkPlan.collectFirst { case node: LeafExecNode => node }.getOrElse { + fail(s"expected a leaf node in test plan:\n$sparkPlan") + } + val cometChild = cometFilter(leaf) + val aliasing = AliasingFallbackCometExec(cometChild, cometChild) + val stagePlan = CometNativeColumnarToRowExec(aliasing) + val rule = RevertNativeForTransitionHeavyStages(spark) + assert(rule.countTransitions(stagePlan) == 1) + + val result = rule(stagePlan) + + assert( + result eq stagePlan, + s"invalid nested fallback must leave the whole stage unchanged:\n$result") + } + } + + test("local TopK sparkFallback returns the supplied restored child") { + val child = spark.range(20).queryExecution.sparkPlan + val originalPlan = TakeOrderedAndProjectExec(5, Seq.empty, child.output, child) + val local = CometLocalTopKExec( + Operator.newBuilder().build(), + originalPlan, + child.output, + 5, + Seq.empty, + dynamicFilterEnabled = false, + child, + SerializedPlan(None)) + val restoredChild = spark.range(10).queryExecution.sparkPlan + + assert(local.sparkFallback(Seq(restoredChild)) eq restoredChild) + intercept[CometExec.InvalidSparkFallbackException] { + local.sparkFallback(Seq.empty) + } + } + for (adaptive <- Seq(false, true)) { test(s"transition reversion preserves local TopK: AQE=$adaptive") { withSQLConf( @@ -393,6 +703,206 @@ class RevertNativeForTransitionHeavyStagesSuite extends CometTestBase { } } + test("revertToSpark leaves transitions below a shuffle that a stripped transition sat on") { + withSQLConf("spark.sql.adaptive.enabled" -> "false") { + withParquetTable((0 until 100).map(i => (i, i % 10)), "tbl") { + val df = sql("SELECT _2, count(*) FROM tbl GROUP BY _2") + df.collect() + val cometPlan = stripAQEPlan(df.queryExecution.executedPlan) + val shuffle = cometPlan + .collectFirst { case s: CometShuffleExchangeExec => s } + .getOrElse(fail(s"test requires a native shuffle:\n$cometPlan")) + // The transition that must survive lives in the map stage, under the exchange. + val preserved = + if (shuffle.child.supportsColumnar) CometNativeColumnarToRowExec(shuffle.child) + else CometSparkToColumnarExec(shuffle.child) + val shuffleWithTransition = shuffle.withNewChildren(Seq(preserved)) + // Stacked transitions sit directly on the shuffle, the shape #6152 strips through. + val onExchange = + CometSparkToColumnarExec(CometNativeColumnarToRowExec(shuffleWithTransition)) + val write = cometIcebergWrite(onExchange) + + val reverted = RevertNativeForTransitionHeavyStages(spark).revertToSpark(write) + val restoredWrite = reverted match { + case node: IcebergWriteExec => node + case other => fail(s"expected IcebergWriteExec, got:\n$other") + } + val restoredShuffle = restoredWrite.child match { + case ColumnarToRowExec(exchange: CometShuffleExchangeExec) => exchange + case exchange: CometShuffleExchangeExec => exchange + case other => + fail(s"expected the shuffle under the restored write, got:\n$other") + } + assert( + restoredShuffle.child eq preserved, + "unwrapping the transition on the shuffle must not strip the stage below it:\n" + + reverted.treeString) + } + } + } + + test("transition-heavy reversion rejects a stripped stage-boundary root") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false") { + withParquetTable((0 until 100).map(i => (i, i % 10)), "tbl") { + var cometPlan: SparkPlan = null + withSQLConf(CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { + val df = sql("SELECT _1, _2 FROM tbl DISTRIBUTE BY _2") + df.collect() + cometPlan = stripAQEPlan(df.queryExecution.executedPlan) + } + val exchange = cometPlan + .collectFirst { case node: CometShuffleExchangeExec => node } + .getOrElse(fail(s"test requires a native shuffle:\n$cometPlan")) + val stagePlan = CometColumnarToRowExec(exchange) + val rule = RevertNativeForTransitionHeavyStages(spark) + assert(rule.countTransitions(stagePlan) == 1) + + var result: SparkPlan = null + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + result = rule(stagePlan) + } + + assert( + result eq stagePlan, + "a transition-heavy stage whose stripped root is an exchange must stay unchanged:\n" + + result.treeString) + } + } + } + + test("transition-heavy revert restores row output for a columnar Spark scan") { + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true") { + withParquetTable((0 until 100).map(i => (i, i % 10)), "tbl") { + val (_, plan) = checkSparkAnswer("SELECT _1, _2 FROM tbl") + val executedPlan = stripAQEPlan(plan) + val resultStageRoot = unwrapCodegen(executedPlan) + + val scan = resultStageRoot match { + case transition: ColumnarToRowExec => + unwrapCodegen(transition.child) match { + case child: FileSourceScanExec => child + case other => + fail(s"expected a vectorized Spark scan under ColumnarToRow, got:\n$other") + } + case other => + fail( + "expected ColumnarToRow over a vectorized Spark scan, got " + + s"${other.getClass.getName}:\n$other") + } + assert(scan.supportsColumnar, s"the reverted scan must use its columnar path:\n$scan") + assert(countCometExecs(executedPlan) == 0, s"the stage must be reverted:\n$executedPlan") + } + } + } + + for (adaptive <- Seq(false, true); columnarRoot <- Seq(false, true)) { + test( + s"transition-heavy map-stage fallback supplies Arrow: AQE=$adaptive, columnar=$columnarRoot") { + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "true", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { + val rows = (0 until 100).map(i => (i, i % 10)) + withParquetTable(rows, "tbl") { + val df = sql("SELECT _1, _2 FROM tbl DISTRIBUTE BY _2") + df.collect() + val exchange = stripAQEPlan(df.queryExecution.executedPlan) + .collectFirst { case node: CometShuffleExchangeExec => node } + .getOrElse(fail("test requires a native shuffle")) + assert(exchange.child.isInstanceOf[CometNativeScanExec]) + val input = if (columnarRoot) exchange.child else cometFilter(exchange.child) + val stage = exchange.withNewChildren( + Seq(CometSparkToColumnarExec(CometNativeColumnarToRowExec(input)))) + val rule = RevertNativeForTransitionHeavyStages(spark) + assert(rule.countTransitions(stage.children.head) == 1) + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0") { + val reverted = rule(stage).asInstanceOf[CometShuffleExchangeExec] + val bridge = reverted.child match { + case node: CometSparkToColumnarExec => node + case other => fail(s"expected an Arrow bridge after map-stage fallback:\n$other") + } + assert(bridge.child.supportsColumnar == columnarRoot) + assert(bridge.child.collect { case _: FileSourceScanExec => true }.nonEmpty) + assert(countCometExecs(bridge.child) == 0) + // Execute the native shuffle, not just its Spark fallback child: Spark batches + // satisfy supportsColumnar but cannot be cast to CometVector by the Arrow stream. + SQLExecution.withNewExecutionId(df.queryExecution) { + val actual = ColumnarToRowExec(reverted) + .executeCollect() + .map(row => (row.getInt(0), row.getInt(1))) + .toSeq + assert(actual.sorted == rows.sorted) + } + } + } + } + } + } + + for (adaptive <- Seq(false, true)) { + test(s"transition-heavy revert preserves native exchange for DISTRIBUTE BY: AQE=$adaptive") { + withSQLConf( + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true", + CometConf.COMET_EXEC_TRANSITION_REVERT_MAX_TRANSITIONS.key -> "0", + CometConf.COMET_SHUFFLE_MODE.key -> "native", + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString, + SQLConf.SHUFFLE_PARTITIONS.key -> "8", + "spark.sql.adaptive.coalescePartitions.enabled" -> "true", + "spark.sql.adaptive.coalescePartitions.parallelismFirst" -> "false", + "spark.sql.adaptive.advisoryPartitionSizeInBytes" -> "67108864") { + withParquetTable((0 until 100).map(i => (i, i % 10)), "tbl") { + val query = "SELECT _1, _2 FROM tbl DISTRIBUTE BY _2" + var sparkAnswer: Seq[Row] = Seq.empty + withSQLConf( + CometConf.COMET_ENABLED.key -> "false", + CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "false") { + sparkAnswer = sql(query).collect().toSeq + } + val df = sql(query) + checkCometAnswer(df, sparkAnswer) + val executedPlan = stripAQEPlan(df.queryExecution.executedPlan) + val resultStageRoot = unwrapCodegen(executedPlan) + + assert( + resultStageRoot.isInstanceOf[ColumnarToRowTransition], + s"the result stage must end with a columnar-to-row transition:\n$executedPlan") + + if (adaptive) { + val read = executedPlan + .collectFirst { case node: AQEShuffleReadExec => node } + .getOrElse(fail(s"expected a coalesced AQE shuffle read:\n$executedPlan")) + assert(read.partitionSpecs.size < 8, s"shuffle must be coalesced:\n$read") + val stage = read.child match { + case node: ShuffleQueryStageExec => node + case other => fail(s"expected a shuffle query stage:\n$other") + } + assert(stage.plan.isInstanceOf[CometShuffleExchangeExec]) + } else { + val exchange = executedPlan + .collectFirst { case node: CometShuffleExchangeExec => node } + .getOrElse(fail(s"expected a native shuffle:\n$executedPlan")) + assert( + exchange.collect { case scan: CometNativeScanExec => scan }.nonEmpty, + s"the map stage must retain its native scan:\n$executedPlan") + assert( + exchange.collect { case scan: FileSourceScanExec => scan }.isEmpty, + s"fallback must not replace the scan below the exchange:\n$executedPlan") + } + } + } + } + } + test("non-AQE apply must not produce an invalid plan when the result stage reverts") { withSQLConf( CometConf.COMET_EXEC_TRANSITION_REVERT_ENABLED.key -> "true",