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 a1ae1af1d1c..7add5f5f20f 100644 --- a/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala +++ b/spark/src/test/scala/org/apache/comet/parquet/CometParquetWriterSuite.scala @@ -20,6 +20,7 @@ package org.apache.comet.parquet import java.io.File +import java.util.concurrent.atomic.AtomicReference import scala.jdk.CollectionConverters._ import scala.util.{Random, Using} @@ -28,6 +29,7 @@ import org.apache.hadoop.fs.{FileSystem, Path} import org.apache.parquet.hadoop.ParquetFileReader import org.apache.parquet.hadoop.metadata.CompressionCodecName import org.apache.parquet.hadoop.util.HadoopInputFile +import org.apache.spark.CometListenerBusUtils import org.apache.spark.sql.{AnalysisException, CometTestBase, DataFrame, Row, SaveMode} import org.apache.spark.sql.comet.{CometBatchScanExec, CometNativeScanExec, CometNativeWriteExec, CometScanExec} import org.apache.spark.sql.execution.{FileSourceScanExec, QueryExecution, SparkPlan} @@ -710,12 +712,12 @@ class CometParquetWriterSuite extends CometTestBase { * The captured execution plan */ private def captureWritePlan(writeOp: String => Unit, outputPath: String): SparkPlan = { - var capturedPlan: Option[QueryExecution] = None + val capturedPlan = new AtomicReference[QueryExecution]() val listener = new org.apache.spark.sql.util.QueryExecutionListener { override def onSuccess(funcName: String, qe: QueryExecution, durationNs: Long): Unit = { if (funcName == "save" || funcName.contains("command")) { - capturedPlan = Some(qe) + capturedPlan.set(qe) } } @@ -725,27 +727,18 @@ class CometParquetWriterSuite extends CometTestBase { exception: Exception): Unit = {} } + // Listener events are delivered asynchronously, so drain the bus before registering: an + // earlier write's event still in flight would otherwise be captured in place of this one. + CometListenerBusUtils.waitUntilEmpty(spark.sparkContext) spark.listenerManager.register(listener) try { writeOp(outputPath) + CometListenerBusUtils.waitUntilEmpty(spark.sparkContext) - // Wait for listener to be called with timeout - val maxWaitTimeMs = 15000 - val checkIntervalMs = 100 - val maxIterations = maxWaitTimeMs / checkIntervalMs - var iterations = 0 - - while (capturedPlan.isEmpty && iterations < maxIterations) { - Thread.sleep(checkIntervalMs) - iterations += 1 - } - - assert( - capturedPlan.isDefined, - s"Listener was not called within ${maxWaitTimeMs}ms - no execution plan captured") - - stripAQEPlan(capturedPlan.get.executedPlan) + val plan = capturedPlan.get() + assert(plan != null, "Listener was not called - no execution plan captured") + stripAQEPlan(plan.executedPlan) } finally { spark.listenerManager.unregister(listener) }