Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand All @@ -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}
Expand Down Expand Up @@ -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)
}
}

Expand All @@ -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)
}
Expand Down
Loading