diff --git a/sql/core/src/main/scala/org/apache/spark/sql/execution/exchange/EnsureRequirements.scala b/sql/core/src/main/scala/org/apache/spark/sql/execution/exchange/EnsureRequirements.scala index 02827d39ec7e5..ebdc3769628ff 100644 --- a/sql/core/src/main/scala/org/apache/spark/sql/execution/exchange/EnsureRequirements.scala +++ b/sql/core/src/main/scala/org/apache/spark/sql/execution/exchange/EnsureRequirements.scala @@ -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) @@ -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 => @@ -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], diff --git a/sql/core/src/test/scala/org/apache/spark/sql/connector/KeyGroupedPartitioningSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/connector/KeyGroupedPartitioningSuite.scala index 5e035065b2ad2..644d85ebc4431 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/connector/KeyGroupedPartitioningSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/connector/KeyGroupedPartitioningSuite.scala @@ -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._ @@ -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", diff --git a/sql/core/src/test/scala/org/apache/spark/sql/execution/exchange/EnsureRequirementsSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/execution/exchange/EnsureRequirementsSuite.scala index a4ad8456c0263..9fd85151de51e 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/execution/exchange/EnsureRequirementsSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/execution/exchange/EnsureRequirementsSuite.scala @@ -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)()