Skip to content
Open
Show file tree
Hide file tree
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 @@ -27,10 +27,11 @@ import org.apache.iceberg.expressions.{And => IcebergAnd, BoundPredicate, Expres
import org.apache.iceberg.spark.source.AuronIcebergSourceUtil
import org.apache.spark.internal.Logging
import org.apache.spark.sql.auron.{NativeConverters, Shims}
import org.apache.spark.sql.catalyst.expressions.{And => SparkAnd, AttributeReference, EqualTo, Expression => SparkExpression, GreaterThan, GreaterThanOrEqual, In, IsNaN, IsNotNull, IsNull, LessThan, LessThanOrEqual, Literal, Not => SparkNot, Or => SparkOr}
import org.apache.spark.sql.catalyst.expressions.{And => SparkAnd, AttributeReference, DynamicPruningExpression, EqualTo, Expression => SparkExpression, GreaterThan, GreaterThanOrEqual, In, IsNaN, IsNotNull, IsNull, LessThan, LessThanOrEqual, Literal, Not => SparkNot, Or => SparkOr}
import org.apache.spark.sql.catalyst.plans.physical.KeyGroupedPartitioning
import org.apache.spark.sql.catalyst.trees.TreeNodeTag
import org.apache.spark.sql.connector.read.{InputPartition, Scan}
import org.apache.spark.sql.execution.datasources.v2.BatchScanExec
import org.apache.spark.sql.execution.datasources.v2.{BatchScanExec, DataSourceRDDPartition}
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.types.{BinaryType, DataType, DecimalType, StringType, StructField, StructType}

Expand All @@ -53,7 +54,8 @@ final case class IcebergScanPlan(
fileSchema: StructType,
partitionSchema: StructType,
pruningPredicates: Seq[pb.PhysicalExprNode],
fieldIdsByName: Map[String, Int])
fieldIdsByName: Map[String, Int],
groupedScanTasks: Option[Seq[Seq[IcebergNativeScanTask]]] = None)

object IcebergScanSupport extends Logging {
private val scanPlanTag: TreeNodeTag[Option[IcebergScanPlan]] = TreeNodeTag(
Expand Down Expand Up @@ -154,7 +156,11 @@ object IcebergScanSupport extends Logging {
case None => return None
}

val partitions = inputPartitions(exec, useRuntimeFilters)
val (partitions, partitionGroups) =
plannedInputPartitions(exec, useRuntimeFilters) match {
case Some(plannedPartitions) => plannedPartitions
case None => return None
}
// Empty scan (e.g. empty table) should still build a plan to return no rows.
if (partitions.isEmpty) {
logWarning(s"Native Iceberg scan planned with empty partitions for $scanClassName.")
Expand All @@ -166,10 +172,14 @@ object IcebergScanSupport extends Logging {
fileSchema,
partitionSchema,
Seq.empty,
fieldIdsByName))
fieldIdsByName,
partitionGroups.map(_.map(_ => Seq.empty))))
}

val icebergPartitions = partitions.flatMap(icebergPartition)
val icebergPartitionGroups =
partitionGroups.map(_.map(_.flatMap(icebergPartition)))
val icebergPartitions =
icebergPartitionGroups.map(_.flatten).getOrElse(partitions.flatMap(icebergPartition))
// All partitions must be Iceberg SparkInputPartition with file scan tasks; otherwise fallback.
if (icebergPartitions.size != partitions.size) {
return None
Expand Down Expand Up @@ -203,7 +213,14 @@ object IcebergScanSupport extends Logging {
}

val pruningPredicates = collectPruningPredicates(scan.asInstanceOf[AnyRef], readSchema)
val nativeTasks = fileTasks.map(task => toNativeScanTask(task, partitionSchema))
val groupedNativeTasks = icebergPartitionGroups.map(_.map { group =>
group
.flatMap(_.tasks)
.collect { case task: FileScanTask => toNativeScanTask(task, partitionSchema) }
})
val nativeTasks = groupedNativeTasks
.map(_.flatten)
.getOrElse(fileTasks.map(task => toNativeScanTask(task, partitionSchema)))
Some(
IcebergScanPlan(
nativeTasks,
Expand All @@ -212,7 +229,8 @@ object IcebergScanSupport extends Logging {
fileSchema,
partitionSchema,
pruningPredicates,
fieldIdsByName))
fieldIdsByName,
groupedNativeTasks))
}

private def planChangelogScan(
Expand All @@ -236,7 +254,11 @@ object IcebergScanSupport extends Logging {
case None => return None
}

val partitions = inputPartitions(exec, useRuntimeFilters)
val (partitions, partitionGroups) =
plannedInputPartitions(exec, useRuntimeFilters) match {
case Some(plannedPartitions) => plannedPartitions
case None => return None
}
if (partitions.isEmpty) {
return Some(
IcebergScanPlan(
Expand All @@ -246,10 +268,14 @@ object IcebergScanSupport extends Logging {
fileSchema,
partitionSchema,
Seq.empty,
fieldIdsByName))
fieldIdsByName,
partitionGroups.map(_.map(_ => Seq.empty))))
}

val icebergPartitions = partitions.flatMap(icebergPartition)
val icebergPartitionGroups =
partitionGroups.map(_.map(_.flatMap(icebergPartition)))
val icebergPartitions =
icebergPartitionGroups.map(_.flatten).getOrElse(partitions.flatMap(icebergPartition))
if (icebergPartitions.size != partitions.size) {
return None
}
Expand Down Expand Up @@ -291,7 +317,14 @@ object IcebergScanSupport extends Logging {
}

val pruningPredicates = collectPruningPredicates(scan.asInstanceOf[AnyRef], readSchema)
val nativeTasks = addedRowsTasks.map(task => toNativeScanTask(task, partitionSchema))
val groupedNativeTasks = icebergPartitionGroups.map(_.map { group =>
group
.flatMap(_.tasks)
.collect { case task: AddedRowsScanTask => toNativeScanTask(task, partitionSchema) }
})
val nativeTasks = groupedNativeTasks
.map(_.flatten)
.getOrElse(addedRowsTasks.map(task => toNativeScanTask(task, partitionSchema)))
Some(
IcebergScanPlan(
nativeTasks,
Expand All @@ -300,7 +333,8 @@ object IcebergScanSupport extends Logging {
fileSchema,
partitionSchema,
pruningPredicates,
fieldIdsByName))
fieldIdsByName,
groupedNativeTasks))
}

private def inspectFieldIdSupport(
Expand Down Expand Up @@ -391,6 +425,59 @@ object IcebergScanSupport extends Logging {
private def deletesEmpty(deletes: java.util.List[_]): Boolean =
deletes == null || deletes.isEmpty

private def plannedInputPartitions(exec: BatchScanExec, useRuntimeFilters: Boolean)
: Option[(Seq[InputPartition], Option[Seq[Seq[InputPartition]]])] = {
exec.outputPartitioning match {
case partitioning: KeyGroupedPartitioning =>
// Runtime filtering can change the final groups after static planning. Keep this
// combination on Spark until native execution can preserve those dynamic groups.
val hasEffectiveRuntimeFilters = exec.runtimeFilters
.exists(_ != DynamicPruningExpression(Literal.TrueLiteral))
if (hasEffectiveRuntimeFilters) {
None
} else {
keyGroupedInputPartitions(exec, partitioning)
.map(groups => groups.flatten -> Some(groups))
}
case _ =>
Some(inputPartitions(exec, useRuntimeFilters) -> None)
}
}

private def keyGroupedInputPartitions(
exec: BatchScanExec,
partitioning: KeyGroupedPartitioning): Option[Seq[Seq[InputPartition]]] = {
try {
// BatchScanExec.inputRDD contains Spark's final partition groups, including empty groups
// introduced while aligning both sides of a storage-partitioned join.
val rddPartitions = exec.inputRDD.partitions.toSeq
val dataSourcePartitions = rddPartitions.collect { case partition: DataSourceRDDPartition =>
partition
}
if (dataSourcePartitions.size != rddPartitions.size) {
logWarning(
s"Expected DataSourceRDDPartition for every key-grouped Iceberg partition in " +
s"${exec.getClass.getName}.")
return None
}

val groups = dataSourcePartitions.map(_.inputPartitions.toSeq)
if (groups.size != partitioning.numPartitions) {
logWarning(
s"Key-grouped Iceberg partition count mismatch: planned ${groups.size}, " +
s"declared ${partitioning.numPartitions}.")
return None
}
Some(groups)
} catch {
case NonFatal(t) =>
logWarning(
s"Failed to obtain final key-grouped input partitions for ${exec.getClass.getName}.",
t)
None
}
}

private def inputPartitions(
exec: BatchScanExec,
useRuntimeFilters: Boolean): Seq[InputPartition] = {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ import org.apache.spark.sql.auron.{EmptyNativeRDD, NativeConverters, NativeHelpe
import org.apache.spark.sql.auron.iceberg.{IcebergNativeScanTask, IcebergScanPlan, IcebergScanSupport}
import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.catalyst.expressions.{Expression, GenericInternalRow, Literal}
import org.apache.spark.sql.catalyst.plans.physical.SinglePartition
import org.apache.spark.sql.catalyst.plans.physical.{KeyGroupedPartitioning, SinglePartition}
import org.apache.spark.sql.execution.{LeafExecNode, SparkPlan, SQLExecution}
import org.apache.spark.sql.execution.datasources.{FilePartition, PartitionedFile}
import org.apache.spark.sql.execution.datasources.v2.BatchScanExec
Expand Down Expand Up @@ -278,6 +278,25 @@ case class NativeIcebergTableScanExec(
}

private def buildFilePartitions(): Array[FilePartition] = {
outputPartitioning match {
case partitioning: KeyGroupedPartitioning =>
val taskGroups = plan.groupedScanTasks.getOrElse {
throw new IllegalStateException(
"Missing grouped Iceberg scan tasks for KeyGroupedPartitioning.")
}
require(
taskGroups.size == partitioning.numPartitions,
s"Key-grouped Iceberg task count ${taskGroups.size} did not match declared " +
s"partition count ${partitioning.numPartitions}.")
require(
taskGroups.flatten == scanTasks,
"Key-grouped Iceberg tasks did not flatten to the planned scan tasks.")
return taskGroups.zipWithIndex.map { case (tasks, index) =>
FilePartition(index, tasks.map(partitionedFile).toArray)
}.toArray
case _ =>
}

// Convert Iceberg scan tasks into Spark FilePartition groups for execution.
if (scanTasks.isEmpty) {
return Array.empty
Expand All @@ -286,13 +305,7 @@ case class NativeIcebergTableScanExec(
val sparkSession = Shims.get.getSqlContext(basedScan).sparkSession
val maxSplitBytes = getMaxSplitBytes(sparkSession, scanTasks)
val partitionedFiles = scanTasks
.map { task =>
Shims.get.getPartitionedFile(
partitionValuesRow(task),
task.location,
task.start,
task.length)
}
.map(partitionedFile)
.sortBy(_.length)(Ordering[Long].reverse)
.toSeq

Expand All @@ -304,6 +317,10 @@ case class NativeIcebergTableScanExec(
}
}

private def partitionedFile(task: IcebergNativeScanTask): PartitionedFile = {
Shims.get.getPartitionedFile(partitionValuesRow(task), task.location, task.start, task.length)
}

private def partitionValuesRow(task: IcebergNativeScanTask): InternalRow = {
val values = partitionSchema.fields.zip(task.partitionValues).map { case (field, value) =>
Literal.create(value, field.dataType).eval()
Expand Down
Loading
Loading