From 5477a9bb1b7f9ea5ed20df9130c9f9eda61ca8f4 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 24 Sep 2026 11:58:55 -0600 Subject: [PATCH 1/2] fix: keep operators above a cached relation native after the AQE re-plan On Spark 3.5+, AQE wraps each cache scan in a TableCacheQueryStageExec and, once the stage materializes, plans the operators above it again. CometExecRule had no case for the stage, so those operators stayed on Spark in the final plan even though the initial plan was fully native. A stage over a Comet in-memory table scan is now a native input, as a Comet shuffle stage is. Spark's cache scan comes back from the re-plan inside the CometSparkToColumnarExec converted over it, which inherits the scan's logical link, so an existing CometSparkToColumnarExec is wrapped in a CometScanWrapper again to let the operators above it convert. --- .../apache/comet/rules/CometExecRule.scala | 14 ++++ .../apache/comet/exec/CometExecSuite.scala | 36 +++++++++ .../comet/exec/CometInMemoryCacheSuite.scala | 80 +++++++++++++++++-- 3 files changed, 123 insertions(+), 7 deletions(-) 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 286f5adea00..f04ada8e3ee 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -408,6 +408,20 @@ case class CometExecRule(session: SparkSession) } } + // On Spark 3.5+, AQE wraps each cache scan in a TableCacheQueryStageExec, a leaf, and once + // the stage materializes it plans the operators above it again. They can only convert over + // a native input. A Comet cache scan already produces Arrow batches, so its stage is a + // native input, as a Comet shuffle stage is. + case s: QueryStageExec if s.plan.isInstanceOf[CometInMemoryTableScanExec] => + convertToComet(s, CometExchangeSink).getOrElse(s) + + // A CometSparkToColumnarExec from an earlier pass. It inherits the logical link of the scan + // it converts, so when AQE plans the operators above a table-cache stage again, it reuses + // this node over the stage rather than the bare stage. Those operators need a native input, + // which the CometScanWrapper around this node gave them when it was converted. + case c: CometSparkToColumnarExec => + convertToComet(c, CometScanWrapper).getOrElse(c) + case op if shouldApplySparkToColumnar(conf, op) => convertToComet(op, CometSparkToColumnarExec).getOrElse(op) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala index 973260b69f3..20ec815b2a7 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala @@ -39,6 +39,7 @@ import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, Comet import org.apache.spark.sql.connector.catalog.InMemoryTableCatalog import org.apache.spark.sql.execution._ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, BroadcastQueryStageExec} +import org.apache.spark.sql.execution.aggregate.HashAggregateExec import org.apache.spark.sql.execution.columnar.CometInMemoryRelationHelper import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, BroadcastExchangeLike, ReusedExchangeExec, ShuffleExchangeExec} @@ -3949,6 +3950,41 @@ class CometExecSuite extends CometTestBase { } } + // https://github.com/apache/datafusion-comet/issues/6202 + test("SparkToColumnar over InMemoryTableScanExec survives the AQE re-plan") { + assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") + // AQE wraps the cache scan in a TableCacheQueryStageExec and, once the stage materializes, + // plans the aggregate above it again. Reset the serializer so the relation is cached in + // Spark's format, as in the test above. + CometInMemoryRelationHelper.clearSerializer() + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true") { + withTempView("table_cache_stage") { + spark + .range(0, 10000, 1, 4) + .selectExpr("id", "id % 10 AS k") + .createOrReplaceTempView("table_cache_stage") + spark.catalog.cacheTable("table_cache_stage") + try { + val df = spark.sql("SELECT k, count(*) FROM table_cache_stage GROUP BY k") + checkAnswer(df, (0L until 10L).map(k => Row(k, 1000L))) + val plan = df.queryExecution.executedPlan + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) + assert(collect(plan) { case a: HashAggregateExec => a }.isEmpty, plan) + assert(collect(plan) { case s: ShuffleExchangeExec => s }.isEmpty, plan) + assert( + collect(plan) { + case c: CometSparkToColumnarExec + if c.child.getClass.getSimpleName == "TableCacheQueryStageExec" => + c + }.size == 1, + plan) + } finally { + spark.catalog.uncacheTable("table_cache_stage") + } + } + } + } + test("SparkToColumnar eliminate redundant in AQE") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index bb9583e4a4d..35910f29c96 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -34,12 +34,12 @@ import org.apache.spark.SparkConf import org.apache.spark.sql.{CometTestBase, Row} import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, Expression, GreaterThanOrEqual, LessThan, Literal} import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} -import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} +import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometPlan, CometSortExec, CometSortMergeJoinExec, CometSparkToColumnarExec} import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} import org.apache.spark.sql.comet.util.Utils -import org.apache.spark.sql.execution.SortExec +import org.apache.spark.sql.execution.{InputAdapter, SortExec, SparkPlan, WholeStageCodegenExec} import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, QueryStageExec, ShuffleQueryStageExec} -import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation} +import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.exchange.{Exchange, ReusedExchangeExec, ShuffleExchangeLike} import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, SortMergeJoinExec} import org.apache.spark.sql.functions.max @@ -166,10 +166,7 @@ class CometInMemoryCacheSuite extends CometTestBase { checkAnswer(df, Seq(Row(1, 1), Row(2, 2))) val plan = df.queryExecution.executedPlan assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) - // Match by name because this suite must also compile on Spark 3.4. - assert(collect(plan) { - case s: QueryStageExec if s.getClass.getSimpleName == "TableCacheQueryStageExec" => s - }.size == 1) + assert(collect(plan) { case s: QueryStageExec if isTableCacheStage(s) => s }.size == 1) assert(collect(plan) { case s: CometInMemoryTableScanExec => s }.size == 1) assert(collect(plan) { case s: ShuffleQueryStageExec => s }.size == 1) assert(collect(plan) { case s: AQEShuffleReadExec => s }.isEmpty) @@ -230,6 +227,75 @@ class CometInMemoryCacheSuite extends CometTestBase { } } + // Match by name because this suite must also compile on Spark 3.4. + private def isTableCacheStage(plan: SparkPlan): Boolean = + plan.getClass.getSimpleName == "TableCacheQueryStageExec" + + // The operators of an executed adaptive plan that run on Spark, including those inside query + // stages. The scan under a CometSparkToColumnar is the Spark input that node converts, so it is + // not reported. + private def sparkOperators(plan: SparkPlan): Seq[String] = plan match { + case a: AdaptiveSparkPlanExec => sparkOperators(a.executedPlan) + case s: QueryStageExec => sparkOperators(s.plan) + case _: CometSparkToColumnarExec => Seq.empty + case _: CometPlan | _: AQEShuffleReadExec | _: WholeStageCodegenExec | _: InputAdapter => + plan.children.flatMap(sparkOperators) + case _ => plan.nodeName +: plan.children.flatMap(sparkOperators) + } + + // https://github.com/apache/datafusion-comet/issues/6202 + test("AQE keeps the operators above a table cache stage native once the stage materializes") { + assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") + // `t` holds 1000 ids for each k, and the ids for a given k sum to 1000 * k + 4995000. + val queries = Seq( + "SELECT k, count(*) FROM t GROUP BY k" -> (0L until 10L).map(k => Row(k, 1000L)), + "SELECT name, sum(id) FROM t JOIN d ON k = k2 GROUP BY name" -> + (0L until 10L).map(k => Row(s"n$k", 1000L * k + 4995000L)), + "SELECT sum(id) FROM t WHERE k = 3" -> Seq(Row(4998000L))) + for { + nativeCache <- Seq(true, false) + warm <- Seq(false, true) + (query, expected) <- queries + } { + withAQECache { + withSQLConf(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> nativeCache.toString) { + withTempView("t", "d") { + spark + .range(0, 10000, 1, 4) + .selectExpr("id", "id % 10 AS k") + .createOrReplaceTempView("t") + spark + .range(10) + .selectExpr("id AS k2", "concat('n', id) AS name") + .createOrReplaceTempView("d") + spark.catalog.cacheTable("t") + spark.catalog.cacheTable("d") + if (warm) { + spark.table("t").count() + spark.table("d").count() + } + + val df = sql(query) + checkAnswer(df, expected) + val plan = df.queryExecution.executedPlan + val clue = s"nativeCache=$nativeCache, warm=$warm: $query\n$plan" + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan, clue) + val cacheScans = collect(plan) { + case s: QueryStageExec if isTableCacheStage(s) => s.plan + } + assert(cacheScans.nonEmpty, clue) + if (nativeCache) { + assert(cacheScans.forall(_.isInstanceOf[CometInMemoryTableScanExec]), clue) + } else { + assert(cacheScans.forall(_.isInstanceOf[InMemoryTableScanExec]), clue) + } + assert(sparkOperators(plan).isEmpty, clue) + } + } + } + } + } + test("CometInMemoryTableScan over CometCachedBatch") { withSQLConf( SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", From a9b257b388e183196a77a61b667f5d15af96018c Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Thu, 24 Sep 2026 14:10:53 -0600 Subject: [PATCH 2/2] test: share the final-plan check and run the cold-cache case cold Move the final AQE plan check into CometTestBase as checkCometOperatorsInFinalPlan, which runs checkCometOperators over the final plan and the plan inside each query stage, and share isTableCacheStage between the two suites. checkAnswer runs the query once as an RDD before collecting it, which warmed the cache before the checked run. Run each query cold and then warm on one cache, with checkToRDD = false. --- .../apache/comet/rules/CometExecRule.scala | 12 +-- .../apache/comet/exec/CometExecSuite.scala | 32 +++---- .../comet/exec/CometInMemoryCacheSuite.scala | 89 ++++++++----------- .../org/apache/spark/sql/CometTestBase.scala | 21 ++++- 4 files changed, 70 insertions(+), 84 deletions(-) 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 f04ada8e3ee..6009906c843 100644 --- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala +++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala @@ -408,17 +408,13 @@ case class CometExecRule(session: SparkSession) } } - // On Spark 3.5+, AQE wraps each cache scan in a TableCacheQueryStageExec, a leaf, and once - // the stage materializes it plans the operators above it again. They can only convert over - // a native input. A Comet cache scan already produces Arrow batches, so its stage is a - // native input, as a Comet shuffle stage is. + // For AQE table-cache stage (Spark 3.5+) on a Comet cache scan. The operators above it are + // planned again once it materializes, and like a Comet shuffle stage it is a native input. case s: QueryStageExec if s.plan.isInstanceOf[CometInMemoryTableScanExec] => convertToComet(s, CometExchangeSink).getOrElse(s) - // A CometSparkToColumnarExec from an earlier pass. It inherits the logical link of the scan - // it converts, so when AQE plans the operators above a table-cache stage again, it reuses - // this node over the stage rather than the bare stage. Those operators need a native input, - // which the CometScanWrapper around this node gave them when it was converted. + // A CometSparkToColumnarExec from an earlier pass, which AQE reuses over a table-cache stage + // because it carries its scan's logical link. Wrap it again so re-planned parents convert. case c: CometSparkToColumnarExec => convertToComet(c, CometScanWrapper).getOrElse(c) diff --git a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala index 20ec815b2a7..f730c1f43df 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometExecSuite.scala @@ -39,8 +39,7 @@ import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, Comet import org.apache.spark.sql.connector.catalog.InMemoryTableCatalog import org.apache.spark.sql.execution._ import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, BroadcastQueryStageExec} -import org.apache.spark.sql.execution.aggregate.HashAggregateExec -import org.apache.spark.sql.execution.columnar.CometInMemoryRelationHelper +import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryTableScanExec} import org.apache.spark.sql.execution.datasources.parquet.ParquetFileFormat import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, BroadcastExchangeLike, ReusedExchangeExec, ShuffleExchangeExec} import org.apache.spark.sql.execution.joins.{BroadcastHashJoinExec, BroadcastNestedLoopJoinExec, CartesianProductExec, SortMergeJoinExec} @@ -3953,9 +3952,9 @@ class CometExecSuite extends CometTestBase { // https://github.com/apache/datafusion-comet/issues/6202 test("SparkToColumnar over InMemoryTableScanExec survives the AQE re-plan") { assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") - // AQE wraps the cache scan in a TableCacheQueryStageExec and, once the stage materializes, - // plans the aggregate above it again. Reset the serializer so the relation is cached in - // Spark's format, as in the test above. + // CometInMemoryCacheSuite covers this too, but registers Comet's rules twice, through its + // plugin and CometTestBase. This suite registers them once, as production does. Reset the + // serializer so the relation is cached in Spark's format, as in the test above. CometInMemoryRelationHelper.clearSerializer() withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true") { withTempView("table_cache_stage") { @@ -3964,23 +3963,14 @@ class CometExecSuite extends CometTestBase { .selectExpr("id", "id % 10 AS k") .createOrReplaceTempView("table_cache_stage") spark.catalog.cacheTable("table_cache_stage") - try { - val df = spark.sql("SELECT k, count(*) FROM table_cache_stage GROUP BY k") - checkAnswer(df, (0L until 10L).map(k => Row(k, 1000L))) - val plan = df.queryExecution.executedPlan - assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan) - assert(collect(plan) { case a: HashAggregateExec => a }.isEmpty, plan) - assert(collect(plan) { case s: ShuffleExchangeExec => s }.isEmpty, plan) - assert( - collect(plan) { - case c: CometSparkToColumnarExec - if c.child.getClass.getSimpleName == "TableCacheQueryStageExec" => - c - }.size == 1, - plan) - } finally { - spark.catalog.uncacheTable("table_cache_stage") + val df = spark.sql("SELECT k, count(*) FROM table_cache_stage GROUP BY k") + checkAnswer(df, (0L until 10L).map(k => Row(k, 1000L))) + val plan = df.queryExecution.executedPlan + checkCometOperatorsInFinalPlan(plan, classOf[InMemoryTableScanExec]) + val conversions = collect(plan) { + case c: CometSparkToColumnarExec if isTableCacheStage(c.child) => c } + assert(conversions.size == 1, plan) } } } diff --git a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala index 35910f29c96..b12ef02c5f7 100644 --- a/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala +++ b/spark/src/test/scala/org/apache/comet/exec/CometInMemoryCacheSuite.scala @@ -31,13 +31,13 @@ import org.apache.arrow.vector.compression.{CompressionCodec, CompressionUtil, N import org.apache.arrow.vector.types.pojo.ArrowType import org.apache.spark.CometDriverPlugin import org.apache.spark.SparkConf -import org.apache.spark.sql.{CometTestBase, Row} +import org.apache.spark.sql.{CometTestBase, QueryTest, Row} import org.apache.spark.sql.catalyst.expressions.{And, Attribute, AttributeReference, Expression, GreaterThanOrEqual, LessThan, Literal} import org.apache.spark.sql.columnar.{CachedBatch, SimpleMetricsCachedBatch} -import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometPlan, CometSortExec, CometSortMergeJoinExec, CometSparkToColumnarExec} +import org.apache.spark.sql.comet.{CometBroadcastHashJoinExec, CometInMemoryTableScanExec, CometSortExec, CometSortMergeJoinExec} import org.apache.spark.sql.comet.execution.arrow.{ArrowCachedBatchSerializer, CometCachedBatchHelper} import org.apache.spark.sql.comet.util.Utils -import org.apache.spark.sql.execution.{InputAdapter, SortExec, SparkPlan, WholeStageCodegenExec} +import org.apache.spark.sql.execution.SortExec import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AQEShuffleReadExec, QueryStageExec, ShuffleQueryStageExec} import org.apache.spark.sql.execution.columnar.{CometInMemoryRelationHelper, InMemoryRelation, InMemoryTableScanExec} import org.apache.spark.sql.execution.exchange.{Exchange, ReusedExchangeExec, ShuffleExchangeLike} @@ -227,22 +227,6 @@ class CometInMemoryCacheSuite extends CometTestBase { } } - // Match by name because this suite must also compile on Spark 3.4. - private def isTableCacheStage(plan: SparkPlan): Boolean = - plan.getClass.getSimpleName == "TableCacheQueryStageExec" - - // The operators of an executed adaptive plan that run on Spark, including those inside query - // stages. The scan under a CometSparkToColumnar is the Spark input that node converts, so it is - // not reported. - private def sparkOperators(plan: SparkPlan): Seq[String] = plan match { - case a: AdaptiveSparkPlanExec => sparkOperators(a.executedPlan) - case s: QueryStageExec => sparkOperators(s.plan) - case _: CometSparkToColumnarExec => Seq.empty - case _: CometPlan | _: AQEShuffleReadExec | _: WholeStageCodegenExec | _: InputAdapter => - plan.children.flatMap(sparkOperators) - case _ => plan.nodeName +: plan.children.flatMap(sparkOperators) - } - // https://github.com/apache/datafusion-comet/issues/6202 test("AQE keeps the operators above a table cache stage native once the stage materializes") { assume(isSpark35Plus, "Table-cache query stages require Spark 3.5+") @@ -252,44 +236,41 @@ class CometInMemoryCacheSuite extends CometTestBase { "SELECT name, sum(id) FROM t JOIN d ON k = k2 GROUP BY name" -> (0L until 10L).map(k => Row(s"n$k", 1000L * k + 4995000L)), "SELECT sum(id) FROM t WHERE k = 3" -> Seq(Row(4998000L))) - for { - nativeCache <- Seq(true, false) - warm <- Seq(false, true) - (query, expected) <- queries - } { - withAQECache { - withSQLConf(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> nativeCache.toString) { - withTempView("t", "d") { - spark - .range(0, 10000, 1, 4) - .selectExpr("id", "id % 10 AS k") - .createOrReplaceTempView("t") - spark - .range(10) - .selectExpr("id AS k2", "concat('n', id) AS name") - .createOrReplaceTempView("d") + withTempView("t", "d") { + spark + .range(0, 10000, 1, 4) + .selectExpr("id", "id % 10 AS k") + .createOrReplaceTempView("t") + spark + .range(10) + .selectExpr("id AS k2", "concat('n', id) AS name") + .createOrReplaceTempView("d") + for { + nativeCache <- Seq(true, false) + (query, expected) <- queries + } { + withAQECache { + withSQLConf(CometConf.COMET_EXEC_IN_MEMORY_CACHE_ENABLED.key -> nativeCache.toString) { spark.catalog.cacheTable("t") spark.catalog.cacheTable("d") - if (warm) { - spark.table("t").count() - spark.table("d").count() - } - - val df = sql(query) - checkAnswer(df, expected) - val plan = df.queryExecution.executedPlan - val clue = s"nativeCache=$nativeCache, warm=$warm: $query\n$plan" - assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan, clue) - val cacheScans = collect(plan) { - case s: QueryStageExec if isTableCacheStage(s) => s.plan - } - assert(cacheScans.nonEmpty, clue) - if (nativeCache) { - assert(cacheScans.forall(_.isInstanceOf[CometInMemoryTableScanExec]), clue) - } else { - assert(cacheScans.forall(_.isInstanceOf[InMemoryTableScanExec]), clue) + val scanClass: Class[_] = + if (nativeCache) classOf[CometInMemoryTableScanExec] + else classOf[InMemoryTableScanExec] + // The first run materializes the cache through the stage and the second reads it warm. + // checkToRDD = false keeps checkAnswer from warming the cache with a run of its own. + Seq("cold", "warm").foreach { cache => + val df = sql(query) + QueryTest.checkAnswer(df, expected, checkToRDD = false) + val plan = df.queryExecution.executedPlan + withClue(s"nativeCache=$nativeCache, $cache cache: $query\n$plan\n") { + val cacheScans = collect(plan) { + case s: QueryStageExec if isTableCacheStage(s) => s.plan + } + assert(cacheScans.nonEmpty) + assert(cacheScans.forall(scanClass.isInstance)) + checkCometOperatorsInFinalPlan(plan, classOf[InMemoryTableScanExec]) + } } - assert(sparkOperators(plan).isEmpty, clue) } } } diff --git a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala index 9740a4a5468..ff72204033f 100644 --- a/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala +++ b/spark/src/test/scala/org/apache/spark/sql/CometTestBase.scala @@ -42,7 +42,7 @@ import org.apache.spark.sql.catalyst.util.sideBySide import org.apache.spark.sql.comet.CometPlanChecker import org.apache.spark.sql.comet.execution.shuffle.{CometColumnarShuffle, CometNativeShuffle, CometShuffleExchangeExec} import org.apache.spark.sql.execution._ -import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper +import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, AdaptiveSparkPlanHelper, AQEShuffleReadExec, QueryStageExec} import org.apache.spark.sql.execution.exchange.ReusedExchangeExec import org.apache.spark.sql.internal._ import org.apache.spark.sql.test._ @@ -629,6 +629,25 @@ abstract class CometTestBase checkPlanNotMissingInput(plan) } + /** + * [[checkCometOperators]] for an adaptive query that has run. `checkSparkAnswerAndOperator` + * inspects a DataFrame that has not run, which for an adaptive query is its initial plan, and + * `checkCometOperators` treats query stages as leaves, so this checks the final plan and the + * plan inside each of its stages. + */ + protected def checkCometOperatorsInFinalPlan( + plan: SparkPlan, + excludedClasses: Class[_]*): Unit = { + assert(plan.asInstanceOf[AdaptiveSparkPlanExec].isFinalPlan, s"The query has not run:\n$plan") + val stagePlans = collect(plan) { case s: QueryStageExec => s.plan } + val excluded = excludedClasses :+ classOf[QueryStageExec] :+ classOf[AQEShuffleReadExec] + (stripAQEPlan(plan) +: stagePlans).foreach(checkCometOperators(_, excluded: _*)) + } + + // Matched by name because Spark 3.4 has no TableCacheQueryStageExec. + protected def isTableCacheStage(plan: SparkPlan): Boolean = + plan.getClass.getSimpleName == "TableCacheQueryStageExec" + // checks the plan node has no missing inputs // such nodes represented in plan with exclamation mark ! // example: !CometWindowExec