diff --git a/spark-extension-shims-spark/src/test/scala/org/apache/auron/exec/AuronExecSuite.scala b/spark-extension-shims-spark/src/test/scala/org/apache/auron/exec/AuronExecSuite.scala index 69de4d834..1251f7174 100644 --- a/spark-extension-shims-spark/src/test/scala/org/apache/auron/exec/AuronExecSuite.scala +++ b/spark-extension-shims-spark/src/test/scala/org/apache/auron/exec/AuronExecSuite.scala @@ -24,6 +24,7 @@ import org.apache.spark.sql.catalyst.expressions.WindowExpression import org.apache.spark.sql.catalyst.expressions.WindowSpecDefinition import org.apache.spark.sql.execution.auron.plan.{NativeCollectLimitExec, NativeGlobalLimitExec, NativeLocalLimitExec, NativeTakeOrderedExec} import org.apache.spark.sql.execution.auron.plan.NativeWindowExec +import org.apache.spark.sql.internal.SQLConf import org.apache.auron.BaseAuronSQLSuite import org.apache.auron.util.AuronTestUtils @@ -43,6 +44,35 @@ class AuronExecSuite extends AuronQueryTest with BaseAuronSQLSuite { } } + test("CollectLimit batches partition scans") { + withTempPath { path => + spark.range(0, 80, 1, 8).write.parquet(path.getCanonicalPath) + + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + "spark.sql.files.maxPartitionBytes" -> "4096", + "spark.sql.limit.initialNumPartitions" -> "1") { + val df = spark.read.parquet(path.getCanonicalPath).where("id < 0").limit(1) + val collectLimit = collectFirst(df.queryExecution.executedPlan) { + case exec: NativeCollectLimitExec => exec + }.get + val numPartitions = collectLimit.child.execute().getNumPartitions + assert(numPartitions > 1) + + val jobGroup = s"collect-limit-${System.nanoTime()}" + spark.sparkContext.setJobGroup(jobGroup, "test CollectLimit job count") + try { + assert(collectLimit.executeCollect().isEmpty) + waitUntilListenerBusEmpty() + val jobCount = spark.sparkContext.statusTracker.getJobIdsForGroup(jobGroup).length + assert(jobCount > 0 && jobCount < numPartitions) + } finally { + spark.sparkContext.clearJobGroup() + } + } + } + } + test("CollectLimit with offset") { if (AuronTestUtils.isSparkV34OrGreater) { withTempView("t1") { diff --git a/spark-extension-shims-spark/src/test/scala/org/apache/spark/sql/AuronQueryTest.scala b/spark-extension-shims-spark/src/test/scala/org/apache/spark/sql/AuronQueryTest.scala index 263f06231..5540b38e3 100644 --- a/spark-extension-shims-spark/src/test/scala/org/apache/spark/sql/AuronQueryTest.scala +++ b/spark-extension-shims-spark/src/test/scala/org/apache/spark/sql/AuronQueryTest.scala @@ -93,4 +93,8 @@ abstract class AuronQueryTest true case _ => false } + + protected def waitUntilListenerBusEmpty(): Unit = { + spark.sparkContext.listenerBus.waitUntilEmpty() + } } diff --git a/spark-extension/src/main/scala/org/apache/spark/sql/execution/auron/plan/NativeCollectLimitBase.scala b/spark-extension/src/main/scala/org/apache/spark/sql/execution/auron/plan/NativeCollectLimitBase.scala index d33661030..4a226fc65 100644 --- a/spark-extension/src/main/scala/org/apache/spark/sql/execution/auron/plan/NativeCollectLimitBase.scala +++ b/spark-extension/src/main/scala/org/apache/spark/sql/execution/auron/plan/NativeCollectLimitBase.scala @@ -17,7 +17,6 @@ package org.apache.spark.sql.execution.auron.plan import scala.collection.mutable -import scala.collection.mutable.ArrayBuffer import org.apache.spark.OneToOneDependency import org.apache.spark.sql.auron.{NativeHelper, NativeRDD, NativeSupports, Shims} @@ -43,15 +42,7 @@ abstract class NativeCollectLimitBase(limit: Int, offset: Int, override val chil override def executeCollect(): Array[InternalRow] = { val partial = Shims.get.createNativeLocalLimitExec(limit, child) - val buf = new ArrayBuffer[InternalRow] - - // collect rows partition-by-partition up to 'limit', avoiding full-partition collect. - val it = partial.execute().toLocalIterator - while (buf.size < limit && it.hasNext) { - val row = it.next().copy() - buf += row - } - val rows = buf.toArray + val rows = partial.executeTake(limit) if (offset > 0) rows.drop(offset) else rows }