Skip to content
Draft
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 @@ -62,9 +62,18 @@ case class EnsureRequirements(
shuffleOrigin: ShuffleOrigin): Seq[SparkPlan] = {
assert(requiredChildDistributions.length == originalChildren.length)
assert(requiredChildOrderings.length == originalChildren.length)
// Get the indexes of children which have specified distribution requirements and need to be
// co-partitioned.
val childrenIndexes = requiredChildDistributions.zipWithIndex.filter {
case (_: ClusteredDistribution, _) => true
case _ => false
}.map(_._2)
val isCoPartitioned = childrenIndexes.length > 1

// Ensure that the operator's children satisfy their output distribution requirements.
var children = originalChildren.zip(requiredChildDistributions).map {
case (child, distribution) if child.outputPartitioning.satisfies(distribution) =>
case (child, distribution) if satisfiesDistribution(
child.outputPartitioning, distribution, isCoPartitioned) =>
ensureOrdering(child, distribution)
case (child, BroadcastDistribution(mode)) =>
BroadcastExchangeExec(mode, child)
Expand All @@ -83,13 +92,6 @@ case class EnsureRequirements(
}
}

// Get the indexes of children which have specified distribution requirements and need to be
// co-partitioned.
val childrenIndexes = requiredChildDistributions.zipWithIndex.filter {
case (_: ClusteredDistribution, _) => true
case _ => false
}.map(_._2)

// Special case: if all sides of the join are single partition and it's physical size less than
// or equal spark.sql.maxSinglePartitionBytes.
val preferSinglePartition = childrenIndexes.forall { i =>
Expand Down Expand Up @@ -233,6 +235,31 @@ case class EnsureRequirements(
children
}

private def satisfiesDistribution(
partitioning: Partitioning,
distribution: Distribution,
isCoPartitioned: Boolean): Boolean = {
if (isCoPartitioned || !conf.v2BucketingAllowJoinKeysSubsetOfPartitionKeys) {
partitioning.satisfies(distribution)
} else {
(partitioning, distribution) match {
case (PartitioningCollection(partitionings), c: ClusteredDistribution) =>
partitionings.exists(satisfiesDistribution(_, c, isCoPartitioned = false))
case (k: KeyGroupedPartitioning, c: ClusteredDistribution) =>
// With subset keys enabled, satisfies() can be true even though rows sharing an
// operation key occupy different partitions. The multi-child path can regroup join
// scans, but a single-child operator needs its clustering before it runs. Without
// GroupPartitionsExec on this branch, use a shuffle when a partition expression is not
// covered by the operation's keys. Do not regroup a scan through an intervening operator.
k.satisfies(c) && k.expressions.forall { e =>
c.clustering.exists(_.semanticEquals(e)) ||
e.collectLeaves().forall(leaf => c.clustering.exists(_.semanticEquals(leaf)))
}
case _ => partitioning.satisfies(distribution)
}
}
}

private def reorder(
leftKeys: IndexedSeq[Expression],
rightKeys: IndexedSeq[Expression],
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ import org.apache.spark.sql.execution.datasources.v2.BatchScanExec
import org.apache.spark.sql.execution.datasources.v2.DataSourceV2ScanRelation
import org.apache.spark.sql.execution.exchange.ShuffleExchangeExec
import org.apache.spark.sql.execution.joins.SortMergeJoinExec
import org.apache.spark.sql.execution.window.WindowExec
import org.apache.spark.sql.internal.SQLConf
import org.apache.spark.sql.internal.SQLConf._
import org.apache.spark.sql.types._
Expand Down Expand Up @@ -1132,6 +1133,57 @@ class KeyGroupedPartitioningSuite extends DistributionAndOrderingSuiteBase {
}
}

for (joined <- Seq(false, true)) {
test(s"SPARK-58968: establish subset clustering below a window, joined=$joined") {
val windowColumns = Array("k", "discard", "v").map(name => Column.create(name, IntegerType))
createTable("window_subset", windowColumns, Array(identity("k"), identity("discard")))
sql("INSERT INTO testcat.ns.window_subset VALUES (1, 10, 10), (1, 20, 20), (2, 30, 30)")
createTable("window_right", Array(Column.create("k", IntegerType)), Array.empty)
sql("INSERT INTO testcat.ns.window_right VALUES (1), (2)")

for {
requireAll <- Seq(false, true)
fullKey <- Seq(false, true)
} {
withSQLConf(
SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
SQLConf.REQUIRE_ALL_CLUSTER_KEYS_FOR_DISTRIBUTION.key -> requireAll.toString,
SQLConf.REQUIRE_ALL_CLUSTER_KEYS_FOR_CO_PARTITION.key -> "false",
SQLConf.V2_BUCKETING_ALLOW_JOIN_KEYS_SUBSET_OF_PARTITION_KEYS.key -> "true",
SQLConf.V2_BUCKETING_SHUFFLE_ENABLED.key -> "true",
SQLConf.V2_BUCKETING_PARTIALLY_CLUSTERED_DISTRIBUTION_ENABLED.key -> "false") {
val windowKeys = if (fullKey) "k, discard" else "k"
val join = if (joined) "JOIN testcat.ns.window_right r ON w.k = r.k" else ""
val hint = if (joined) "/*+ MERGE(w, r) */" else ""
val df = sql(
s"""SELECT $hint w.k, w.discard, w.rn
|FROM (
| SELECT k, discard, ROW_NUMBER() OVER (
| PARTITION BY $windowKeys ORDER BY v) AS rn
| FROM testcat.ns.window_subset
|) w $join
|""".stripMargin)

withClue(s"requireAll=$requireAll, fullKey=$fullKey: ") {
checkAnswer(df, Seq(Row(1, 10, 1), Row(1, 20, if (fullKey) 1 else 2), Row(2, 30, 1)))
val plan = df.queryExecution.executedPlan
val Seq(window) = collect(plan) { case w: WindowExec => w }
val shuffles = collectAllShuffles(window.child)
if (fullKey) {
assert(shuffles.isEmpty, "full-key window should retain the scan's clustering")
} else {
assert(shuffles.size == 1, "the shuffle must be below the window")
val partitioning = shuffles.head.outputPartitioning
.asInstanceOf[physical.HashPartitioning]
assert(partitioning.expressions == window.partitionSpec)
}
assert(collect(plan) { case j: SortMergeJoinExec => j }.size == (if (joined) 1 else 0))
}
}
}
}
}

test("data source partitioning + dynamic partition filtering") {
withSQLConf(
SQLConf.AUTO_BROADCASTJOIN_THRESHOLD.key -> "-1",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1137,6 +1137,47 @@ class EnsureRequirementsSuite extends SharedSparkSession {
}
}

test("SPARK-58968: single-child clustering respects expressions, collections and counts") {
val key = AttributeReference("key", IntegerType)()
val extra = AttributeReference("extra", IntegerType)()
val transformedKey = bucket(4, key)
val subset = KeyGroupedPartitioning(Seq(key, extra), 3)
val full = KeyGroupedPartitioning(Seq(key), 3)
val required = ClusteredDistribution(Seq(key), requireAllClusterKeys = false,
requiredNumPartitions = Some(3))

val cases = Seq(
(subset, required, true),
(full, required, false),
(KeyGroupedPartitioning(Seq(transformedKey), 3), required, false),
(KeyGroupedPartitioning(Seq(transformedKey, extra), 3), required, true),
(KeyGroupedPartitioning(Seq(transformedKey), 3),
required.copy(clustering = Seq(transformedKey), requireAllClusterKeys = true), false),
(PartitioningCollection(Seq(subset, HashPartitioning(Seq(key), 3))), required, false),
(PartitioningCollection(Seq(subset, full)), required, false),
(PartitioningCollection(Seq(subset, subset)), required, true),
(full, required.copy(requiredNumPartitions = Some(4)), true))

withSQLConf(SQLConf.V2_BUCKETING_ALLOW_JOIN_KEYS_SUBSET_OF_PARTITION_KEYS.key -> "true") {
cases.foreach { case (partitioning, distribution, needsShuffle) =>
val child = DummySparkPlan(outputPartitioning = partitioning)
val parent = DummySparkPlan(children = Seq(child),
requiredChildDistribution = Seq(distribution), requiredChildOrdering = Seq(Nil))
val result = EnsureRequirements.apply(parent).children.head
withClue(s"$partitioning, $distribution: ") {
if (needsShuffle) {
val shuffle = result.asInstanceOf[ShuffleExchangeExec]
assert(shuffle.child == child)
assert(shuffle.outputPartitioning ==
HashPartitioning(distribution.clustering, distribution.requiredNumPartitions.get))
} else {
assert(result == child)
}
}
}
}
}

test("SPARK-42168: FlatMapCoGroupInPandas and Window function with differing key order") {
val lKey = AttributeReference("key", IntegerType)()
val lKey2 = AttributeReference("key2", IntegerType)()
Expand Down