This is an automated email from the ASF dual-hosted git repository.

ulysses-you pushed a commit to branch branch-4.x
in repository https://gitbox.apache.org/repos/asf/spark.git


The following commit(s) were added to refs/heads/branch-4.x by this push:
     new 7388a4cb84dc [SPARK-43317][SQL] Support combine adjacent aggregation
7388a4cb84dc is described below

commit 7388a4cb84dc442abbe637c2e901cc6c54b03495
Author: Xiduo You <[email protected]>
AuthorDate: Fri Jul 17 16:17:11 2026 +0800

    [SPARK-43317][SQL] Support combine adjacent aggregation
    
    ### What changes were proposed in this pull request?
    
    This is a rework for https://github.com/apache/spark/pull/40990.
    
    Add a new physical rule `CombineAdjacentAggregation` in the query stage 
preparation rules. When there is an adjacent aggregation pair with `Partial` 
and `Final` mode that groups by the same expressions and refers to the same 
logical aggregate, combine them into a single `Complete` mode aggregate so we 
do not need to merge the aggregation buffer.
    
    For example:
    ```
    HashAggregate (Final)         HashAggregate (Complete)
           |                             |
    HashAggregate (Partial)    =>    Exchange
           |
       Exchange
    ```
    
    It supports `HashAggregateExec`, `SortAggregateExec` and 
`ObjectHashAggregateExec`.
    
    The new rule runs before `ReplaceHashWithSortAgg`. The "merge adjacent 
partial and final aggregate" logic that used to live in 
`ReplaceHashWithSortAgg` is removed and consolidated into this new rule. 
`ReplaceHashWithSortAgg` is thus simplified to only replace a single hash 
aggregate with a sort aggregate when its child ordering is satisfied; it can 
still turn the combined `Complete` hash aggregate into a sort aggregate.
    
    The behavior is guarded by a new internal config 
`spark.sql.execution.combineAdjacentAggregation`, which falls back to 
`spark.sql.execution.replaceHashWithSortAgg` (disabled by default). So this 
does not change plans out of the box, and enabling `replaceHashWithSortAgg` 
also enables combining, preserving the previous merge behavior that was moved 
out of that rule.
    
    ### Why are the changes needed?
    
    Two reasons:
    
    1. **Performance.** Combining an adjacent `Partial` + `Final` aggregate 
into a single `Complete` aggregate skips the aggregation buffer merge, which is 
beneficial especially in the high cardinality case.
    
    2. **Correctness fix for the existing `ReplaceHashWithSortAgg` merge 
path.** When `spark.sql.execution.replaceHashWithSortAgg` is enabled, 
`ReplaceHashWithSortAgg` already merged an adjacent partial + final aggregate 
into a `Complete` sort aggregate, but it built the merged aggregate from the 
**final** aggregate's expressions:
    
       ```scala
       val sortAgg = hashAgg.toSortAggregate // final aggregate's expressions
       sortAgg.copy(
         aggregateExpressions = sortAgg.aggregateExpressions.map(_.copy(mode = 
Complete)),
         child = partialAgg.child)
       ```
    
       This is wrong, because after merging, the aggregate reads the original 
input (the partial aggregate's child) rather than the partial aggregation 
buffers, and some information only lives on the **partial** aggregate 
expressions. Concretely:
    
       - The `FILTER (WHERE ...)` clause of an aggregate function is only 
carried by the partial aggregate; the final aggregate merely merges buffers. 
Using the final expressions silently drops the filter and **produces wrong 
results**.
       - `initialInputBufferOffset` was not reset to `0`, even though the 
merged aggregate now reads the original input with no leading grouping-key 
columns.
       - The grouping expressions of the partial and final aggregate were not 
required to match.
    
       Repro (returns wrong counts on current master with 
`spark.sql.execution.replaceHashWithSortAgg=true`, the `FILTER` is dropped):
       ```sql
       SELECT key, count(*) FILTER (WHERE v > 10) FROM (
         SELECT /*+ SHUFFLE_MERGE(t1) */ t1.key AS key, t1.v AS v FROM t1 JOIN 
t2 ON t1.key = t2.key
       ) GROUP BY key
       ```
    
       The new `CombineAdjacentAggregation` fixes this by building the 
`Complete` aggregate from the partial aggregate's expressions (keeping the 
`FILTER` clause), resetting `initialInputBufferOffset` to `0`, and requiring 
the grouping expressions to match.
    
    ### Does this PR introduce _any_ user-facing change?
    
    No. The new config falls back to 
`spark.sql.execution.replaceHashWithSortAgg`, which is disabled by default, so 
plans are unchanged out of the box. The correctness fix only affects the merge 
path that is enabled together with `replaceHashWithSortAgg`.
    
    ### How was this patch tested?
    
    Add `CombineAdjacentAggregationSuite` (covering hash / sort / object-hash 
aggregates, `FILTER` clauses, `count(distinct)`, the buffer-offset reset, the 
grouping / shuffle guards, and the config fallback) and extend 
`ReplaceHashWithSortAggSuite` to cover the combine-then-replace interaction.
    
    ### Was this patch authored or co-authored using generative AI tooling?
    
    Generated-by: Claude Code
    
    🤖 Generated with [Claude Code](https://claude.com/claude-code)
    
    Closes #57276 from ulysses-you/combine-agg.
    
    Authored-by: Xiduo You <[email protected]>
    Signed-off-by: Xiduo You <[email protected]>
    (cherry picked from commit c5a1880d48b7c5faf683b5b6467c20bf8288d4ce)
    Signed-off-by: Xiduo You <[email protected]>
---
 .../org/apache/spark/sql/internal/SQLConf.scala    |  11 +
 .../sql/execution/CombineAdjacentAggregation.scala |  93 +++++++
 .../spark/sql/execution/QueryExecution.scala       |   4 +
 .../sql/execution/ReplaceHashWithSortAgg.scala     |  65 +----
 .../execution/adaptive/AdaptiveSparkPlanExec.scala |   4 +
 .../CombineAdjacentAggregationSuite.scala          | 308 +++++++++++++++++++++
 .../execution/ReplaceHashWithSortAggSuite.scala    |  72 ++++-
 7 files changed, 504 insertions(+), 53 deletions(-)

diff --git 
a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala 
b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala
index 3d2c618732bc..0b3f4693853a 100644
--- a/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala
+++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/internal/SQLConf.scala
@@ -3008,6 +3008,17 @@ object SQLConf {
     .booleanConf
     .createWithDefault(false)
 
+  val COMBINE_ADJACENT_AGGREGATION_ENABLED =
+    buildConf("spark.sql.execution.combineAdjacentAggregation")
+      .internal()
+      .doc("When true, combine adjacent aggregation with `Partial` and `Final` 
to `Complete` " +
+        "mode. This defaults to the value of 
`spark.sql.execution.replaceHashWithSortAgg` since " +
+        "combining adjacent aggregation subsumes the partial-and-final merge 
that " +
+        "`replaceHashWithSortAgg` used to perform on its own.")
+      .version("4.3.0")
+      .withBindingPolicy(ConfigBindingPolicy.NOT_APPLICABLE)
+      .fallbackConf(REPLACE_HASH_WITH_SORT_AGG_ENABLED)
+
   val USE_PARTITION_EVALUATOR = 
buildConf("spark.sql.execution.usePartitionEvaluator")
     .internal()
     .doc("When true, use PartitionEvaluator to execute SQL operators.")
diff --git 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/CombineAdjacentAggregation.scala
 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/CombineAdjacentAggregation.scala
new file mode 100644
index 000000000000..00643e1638ff
--- /dev/null
+++ 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/CombineAdjacentAggregation.scala
@@ -0,0 +1,93 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.execution
+
+import org.apache.spark.sql.catalyst.expressions.aggregate.{Complete, Final, 
Partial}
+import org.apache.spark.sql.catalyst.rules.Rule
+import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, 
HashAggregateExec, ObjectHashAggregateExec, SortAggregateExec}
+import org.apache.spark.sql.internal.SQLConf
+
+/**
+ * This rule combines adjacent aggregation with `Partial` and `Final` to 
`Complete` mode.
+ * Example for hash aggregate:
+ *    HashAggregate (Final)         HashAggregate (Complete)
+ *          |                             |
+ *    HashAggregate (Partial)    =>    Exchange
+ *          |
+ *       Exchange
+ *
+ * Example for sort aggregate:
+ *    SortAggregateExec (Final)       SortAggregateExec (Complete)
+ *          |                               |
+ *    SortAggregateExec (Partial)    =>    Sort
+ *          |                               |
+ *         Sort                          Exchange
+ *          |
+ *       Exchange
+ *
+ * It supports [[HashAggregateExec]], [[SortAggregateExec]] and 
[[ObjectHashAggregateExec]].
+ */
+object CombineAdjacentAggregation extends Rule[SparkPlan] {
+  override def apply(plan: SparkPlan): SparkPlan = {
+    if (!conf.getConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED)) {
+      return plan
+    }
+
+    plan.transformDown {
+      case finalAgg @ HashAggregateExec(_, _, _, _, _, _, _, _, partialAgg: 
HashAggregateExec)
+          if isPartialAgg(partialAgg, finalAgg) =>
+        finalAgg.copy(
+          groupingExpressions = partialAgg.groupingExpressions,
+          aggregateExpressions = 
partialAgg.aggregateExpressions.map(_.copy(mode = Complete)),
+          initialInputBufferOffset = 0,
+          child = partialAgg.child)
+
+      case finalAgg @ SortAggregateExec(_, _, _, _, _, _, _, _, partialAgg: 
SortAggregateExec)
+          if isPartialAgg(partialAgg, finalAgg) =>
+        finalAgg.copy(
+          groupingExpressions = partialAgg.groupingExpressions,
+          aggregateExpressions = 
partialAgg.aggregateExpressions.map(_.copy(mode = Complete)),
+          initialInputBufferOffset = 0,
+          child = partialAgg.child)
+
+      case finalAgg @ ObjectHashAggregateExec(_, _, _, _, _, _, _, _,
+        partialAgg: ObjectHashAggregateExec)
+          if isPartialAgg(partialAgg, finalAgg) =>
+        finalAgg.copy(
+          groupingExpressions = partialAgg.groupingExpressions,
+          aggregateExpressions = 
partialAgg.aggregateExpressions.map(_.copy(mode = Complete)),
+          initialInputBufferOffset = 0,
+          child = partialAgg.child)
+    }
+  }
+
+  /**
+   * Check if `partialAgg` is the partial aggregate of `finalAgg`.
+   */
+  private def isPartialAgg(
+      partialAgg: BaseAggregateExec,
+      finalAgg: BaseAggregateExec): Boolean = {
+    partialAgg.aggregateExpressions.forall(_.mode == Partial) &&
+      finalAgg.aggregateExpressions.forall(_.mode == Final) &&
+      partialAgg.groupingExpressions.map(_.canonicalized) ==
+        finalAgg.groupingExpressions.map(_.canonicalized) &&
+      finalAgg.logicalLink.isDefined &&
+      partialAgg.logicalLink.isDefined &&
+      finalAgg.logicalLink.get.sameResult(partialAgg.logicalLink.get)
+  }
+}
diff --git 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala
index 2174c5899286..f70eafc70af3 100644
--- 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala
+++ 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/QueryExecution.scala
@@ -764,6 +764,10 @@ object QueryExecution {
       EnsureRequirements(),
       // This rule must be run after `EnsureRequirements`.
       InsertSortForLimitAndOffset,
+      // `CombineAdjacentAggregation` must run before 
`ReplaceHashWithSortAgg`: it combines a pair
+      // of adjacent partial and final aggregate into a single `Complete` mode 
aggregate, which
+      // `ReplaceHashWithSortAgg` can then replace with a sort aggregate when 
the ordering allows.
+      CombineAdjacentAggregation,
       // `ReplaceHashWithSortAgg` needs to be added after `EnsureRequirements` 
to guarantee the
       // sort order of each node is checked to be valid.
       ReplaceHashWithSortAgg,
diff --git 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAgg.scala
 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAgg.scala
index 4495bc9b6a58..ca913a19019d 100644
--- 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAgg.scala
+++ 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAgg.scala
@@ -18,31 +18,16 @@
 package org.apache.spark.sql.execution
 
 import org.apache.spark.sql.catalyst.expressions.SortOrder
-import org.apache.spark.sql.catalyst.expressions.aggregate.{Complete, Final, 
Partial}
 import org.apache.spark.sql.catalyst.rules.Rule
 import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, 
HashAggregateExec, ObjectHashAggregateExec}
 import org.apache.spark.sql.internal.SQLConf
 
 /**
- * Replace hash-based aggregate with sort aggregate in the spark plan if:
+ * Replace hash-based aggregate with sort aggregate in the spark plan if the 
plan is a
+ * [[HashAggregateExec]] or [[ObjectHashAggregateExec]], and the child 
satisfies the sort order
+ * of corresponding [[SortAggregateExec]].
  *
- * 1. The plan is a pair of partial and final [[HashAggregateExec]] or 
[[ObjectHashAggregateExec]],
- *    and the child of partial aggregate satisfies the sort order of 
corresponding
- *    [[SortAggregateExec]].
- * or
- * 2. The plan is a [[HashAggregateExec]] or [[ObjectHashAggregateExec]], and 
the child satisfies
- *    the sort order of corresponding [[SortAggregateExec]].
- *
- * Examples:
- * 1. aggregate after join:
- *
- *  HashAggregate(t1.i, SUM, final)
- *               |                         SortAggregate(t1.i, SUM, complete)
- * HashAggregate(t1.i, SUM, partial)   =>                |
- *               |                            SortMergeJoin(t1.i = t2.j)
- *    SortMergeJoin(t1.i = t2.j)
- *
- * 2. aggregate after sort:
+ * Example:
  *
  * HashAggregate(t1.i, SUM, partial)         SortAggregate(t1.i, SUM, partial)
  *               |                     =>                  |
@@ -51,6 +36,10 @@ import org.apache.spark.sql.internal.SQLConf
  * Hash-based aggregate can be replaced when its child satisfies the sort 
order of
  * corresponding sort aggregate. Sort aggregate is faster in the sense that
  * it does not have hashing overhead of hash aggregate.
+ *
+ * Note that [[CombineAdjacentAggregation]] runs before this rule, so a pair 
of adjacent partial
+ * and final aggregate has already been combined into a single `Complete` mode 
aggregate, which
+ * this rule can further replace with a sort aggregate when the ordering is 
satisfied.
  */
 object ReplaceHashWithSortAgg extends Rule[SparkPlan] {
   def apply(plan: SparkPlan): SparkPlan = {
@@ -68,44 +57,16 @@ object ReplaceHashWithSortAgg extends Rule[SparkPlan] {
     plan.transformDown {
       case hashAgg: BaseAggregateExec if isHashBasedAggWithKeys(hashAgg) =>
         val sortAgg = hashAgg.toSortAggregate
-        hashAgg.child match {
-          case partialAgg: BaseAggregateExec
-            if isHashBasedAggWithKeys(partialAgg) && isPartialAgg(partialAgg, 
hashAgg) =>
-            if (SortOrder.orderingSatisfies(
-                partialAgg.child.outputOrdering, 
sortAgg.requiredChildOrdering.head)) {
-              sortAgg.copy(
-                aggregateExpressions = 
sortAgg.aggregateExpressions.map(_.copy(mode = Complete)),
-                child = partialAgg.child)
-            } else {
-              hashAgg
-            }
-          case other =>
-            if (SortOrder.orderingSatisfies(
-                other.outputOrdering, sortAgg.requiredChildOrdering.head)) {
-              sortAgg
-            } else {
-              hashAgg
-            }
+        if (SortOrder.orderingSatisfies(
+            hashAgg.child.outputOrdering, sortAgg.requiredChildOrdering.head)) 
{
+          sortAgg
+        } else {
+          hashAgg
         }
       case other => other
     }
   }
 
-  /**
-   * Check if `partialAgg` to be partial aggregate of `finalAgg`.
-   */
-  private def isPartialAgg(partialAgg: BaseAggregateExec, finalAgg: 
BaseAggregateExec): Boolean = {
-    if (partialAgg.aggregateExpressions.forall(_.mode == Partial) &&
-        finalAgg.aggregateExpressions.forall(_.mode == Final)) {
-      (finalAgg.logicalLink, partialAgg.logicalLink) match {
-        case (Some(agg1), Some(agg2)) => agg1.sameResult(agg2)
-        case _ => false
-      }
-    } else {
-      false
-    }
-  }
-
   /**
    * Check if `agg` is [[HashAggregateExec]] or [[ObjectHashAggregateExec]],
    * and has grouping keys.
diff --git 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/adaptive/AdaptiveSparkPlanExec.scala
 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/adaptive/AdaptiveSparkPlanExec.scala
index f74115d98369..1905b31b49f9 100644
--- 
a/sql/core/src/main/scala/org/apache/spark/sql/execution/adaptive/AdaptiveSparkPlanExec.scala
+++ 
b/sql/core/src/main/scala/org/apache/spark/sql/execution/adaptive/AdaptiveSparkPlanExec.scala
@@ -131,6 +131,10 @@ case class AdaptiveSparkPlanExec(
       // join drops its child ordering, which `ReplaceHashWithSortAgg` would 
otherwise rely on to
       // turn a hash aggregate into a sort aggregate.
       ConvertSortMergeJoinToShuffledHashJoin(ensureRequirements),
+      // `CombineAdjacentAggregation` must run before 
`ReplaceHashWithSortAgg`: it combines a pair
+      // of adjacent partial and final aggregate into a single `Complete` mode 
aggregate, which
+      // `ReplaceHashWithSortAgg` can then replace with a sort aggregate when 
the ordering allows.
+      CombineAdjacentAggregation,
       ReplaceHashWithSortAgg,
       RemoveRedundantWindowGroupLimits,
       DisableUnnecessaryBucketedScan,
diff --git 
a/sql/core/src/test/scala/org/apache/spark/sql/execution/CombineAdjacentAggregationSuite.scala
 
b/sql/core/src/test/scala/org/apache/spark/sql/execution/CombineAdjacentAggregationSuite.scala
new file mode 100644
index 000000000000..e4ea41960da9
--- /dev/null
+++ 
b/sql/core/src/test/scala/org/apache/spark/sql/execution/CombineAdjacentAggregationSuite.scala
@@ -0,0 +1,308 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.execution
+
+import org.apache.spark.sql.{QueryTest, Row}
+import org.apache.spark.sql.catalyst.expressions.aggregate.{Complete, Final}
+import org.apache.spark.sql.execution.adaptive.AdaptiveSparkPlanHelper
+import org.apache.spark.sql.execution.aggregate.{BaseAggregateExec, 
HashAggregateExec, SortAggregateExec}
+import org.apache.spark.sql.internal.SQLConf
+import org.apache.spark.sql.test.SharedSparkSession
+
+class CombineAdjacentAggregationSuite extends QueryTest
+  with SharedSparkSession
+  with AdaptiveSparkPlanHelper {
+
+  private def numAggregates(query: String): Int = {
+    val df = sql(query)
+    df.collect()
+    collect(df.queryExecution.executedPlan) {
+      case agg: BaseAggregateExec => agg
+    }.size
+  }
+
+  private def checkNumAggregation(
+      query: String,
+      numAggWithDisabled: Int,
+      numAggWithEnabled: Int): Unit = {
+    var expectedResult: Array[Row] = null
+    withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false") {
+      val df = sql(query)
+      expectedResult = df.collect()
+      assert(collect(df.queryExecution.executedPlan) {
+        case agg: BaseAggregateExec => agg
+      }.size == numAggWithDisabled)
+    }
+
+    withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true") {
+      val df = sql(query)
+      checkAnswer(df, expectedResult)
+      assert(collect(df.queryExecution.executedPlan) {
+        case agg: BaseAggregateExec => agg
+      }.size == numAggWithEnabled)
+    }
+  }
+
+  test("Test combine adjacent aggregation") {
+    withTempView("t") {
+      spark.range(20).selectExpr(s"id % 3 as k", "id % 7 as v")
+        .createOrReplaceTempView("t")
+
+      // do not combine if no adjacent aggregation
+      checkNumAggregation(
+        "SELECT k, count(*) FROM t GROUP BY k",
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 2)
+
+      // combine adjacent hash aggregation
+      checkNumAggregation(
+        "SELECT k, count(*) FROM (SELECT /*+ repartition(k) */ * FROM t) GROUP 
BY k",
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 1)
+
+      // combine adjacent object hash aggregate
+      checkNumAggregation(
+        "SELECT k, collect_set(v) FROM (SELECT /*+ repartition(k) */ * FROM t) 
GROUP BY k",
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 1)
+
+      // do not combine adjacent hash aggregation
+      checkNumAggregation(
+        "SELECT k, count(distinct v) FROM t GROUP BY k",
+        numAggWithDisabled = 4,
+        numAggWithEnabled = 4)
+
+      // combine adjacent hash aggregation
+      checkNumAggregation(
+        "SELECT k, count(distinct v) FROM (SELECT /*+ repartition(k) */ * FROM 
t) GROUP BY k",
+        numAggWithDisabled = 4,
+        numAggWithEnabled = 2)
+    }
+  }
+
+  test("Combine adjacent sort aggregate") {
+    // `max`/`min` over a string column produces a non-mutable aggregation 
buffer, so it cannot be
+    // planned as a hash or object-hash aggregate and Spark falls back to 
`SortAggregateExec`. This
+    // exercises the sort-aggregate rewrite branch specifically (a numeric 
`max` would be planned as
+    // a `HashAggregateExec`, leaving that branch untested).
+    withTempView("t") {
+      spark.range(20).selectExpr("id % 3 as k", "cast(id % 7 as string) as v")
+        .createOrReplaceTempView("t")
+      val query = "SELECT k, max(v) FROM (SELECT /*+ repartition(k) */ * FROM 
t) GROUP BY k"
+
+      val expected = 
withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false") {
+        val df = sql(query)
+        val aggs = collect(df.queryExecution.executedPlan) { case agg: 
BaseAggregateExec => agg }
+        // Two sort aggregates (partial + final), confirming the 
sort-aggregate branch is exercised.
+        assert(aggs.size == 2)
+        assert(aggs.forall(_.isInstanceOf[SortAggregateExec]))
+        df.collect()
+      }
+
+      withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true") {
+        val df = sql(query)
+        checkAnswer(df, expected)
+        val aggs = collect(df.queryExecution.executedPlan) { case agg: 
BaseAggregateExec => agg }
+        // Combined into a single `Complete` mode sort aggregate.
+        assert(aggs.size == 1)
+        assert(aggs.head.isInstanceOf[SortAggregateExec])
+        assert(aggs.head.aggregateExpressions.forall(_.mode == Complete))
+      }
+    }
+  }
+
+  test("Do not combine adjacent aggregates with mismatched grouping 
expressions") {
+    // A grouping-expression mismatch normally forces `EnsureRequirements` to 
insert an Exchange
+    // between the partial and final aggregate, so the planner never produces 
an adjacent pair with
+    // differing grouping keys. To exercise the canonicalized 
grouping-equality guard in
+    // `isPartialAgg` directly, we take a genuinely adjacent partial/final 
pair (the input is
+    // already partitioned by `k`, so no Exchange is inserted between them) 
and rewrite the final
+    // aggregate's grouping expressions to mismatch. The rule must then refuse 
to combine the pair.
+    withTempView("t") {
+      spark.range(20).selectExpr("id % 3 as k", "id % 7 as 
v").createOrReplaceTempView("t")
+      val queryText = "SELECT k, count(*) FROM (SELECT /*+ repartition(k) */ * 
FROM t) GROUP BY k"
+      // Build the physical plan with the rule disabled so the raw adjacent 
partial/final pair
+      // survives, and with AQE disabled so `executedPlan` is a plain tree we 
can rewrite.
+      val finalAgg = withSQLConf(
+          SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false") {
+        collect(sql(queryText).queryExecution.executedPlan) {
+          case f: HashAggregateExec if f.child.isInstanceOf[HashAggregateExec] 
=> f
+        }.head
+      }
+
+      withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true") {
+        // Sanity check: with matching grouping expressions, the adjacent pair 
is combined into one.
+        assert(collect(CombineAdjacentAggregation(finalAgg)) {
+          case agg: BaseAggregateExec => agg
+        }.size == 1)
+
+        // Mismatched grouping expressions: the same adjacent pair must not be 
combined. We copy the
+        // tags (including the `logicalLink`) from the original node so that 
`isPartialAgg` reaches
+        // the grouping-equality check rather than short-circuiting on a 
missing `logicalLink`.
+        val mismatched = finalAgg.copy(groupingExpressions = Nil)
+        mismatched.copyTagsFrom(finalAgg)
+        val result = CombineAdjacentAggregation(mismatched)
+        val aggs = collect(result) { case agg: BaseAggregateExec => agg }
+        assert(aggs.size == 2)
+        assert(aggs.head.aggregateExpressions.forall(_.mode == Final))
+      }
+    }
+  }
+
+  test("Combined aggregate reads original input with a zero buffer offset") {
+    // When adjacent aggregates are combined into a single `Complete` mode 
aggregate, its child
+    // becomes the partial aggregate's child, so it reads the original input 
rather than a row of
+    // `[groupingKeys, aggregationBuffers]`. `initialInputBufferOffset` must 
therefore be reset to
+    // 0. Use multiple grouping keys and multiple aggregate functions so that 
a stale, non-zero
+    // offset would bind the aggregate functions against the wrong input 
columns.
+    withTempView("t") {
+      spark.range(60).selectExpr("id % 3 as k1", "id % 5 as k2", "id % 7 as v")
+        .createOrReplaceTempView("t")
+
+      // hash aggregate with declarative aggregate functions
+      checkNumAggregation(
+        """SELECT k1, k2, sum(v), count(v), avg(v), max(v), min(v)
+          |FROM (SELECT /*+ repartition(k1, k2) */ * FROM t) GROUP BY k1, 
k2""".stripMargin,
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 1)
+
+      // object hash aggregate with imperative aggregate functions
+      checkNumAggregation(
+        """SELECT k1, k2, sort_array(collect_list(v)), count(v)
+          |FROM (SELECT /*+ repartition(k1, k2) */ * FROM t) GROUP BY k1, 
k2""".stripMargin,
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 1)
+    }
+  }
+
+  test("Do not combine when a shuffle sits between the partial and final 
aggregate") {
+    withTempView("t") {
+      spark.range(20).selectExpr("id % 3 as k", "id % 7 as v")
+        .createOrReplaceTempView("t")
+
+      // The input is repartitioned by `k`, but the query groups by `k + 1`, 
so the partial
+      // aggregate's output is not partitioned the way the final aggregate 
requires.
+      // `EnsureRequirements` therefore inserts an Exchange between the two 
aggregates, leaving them
+      // non-adjacent, and the rule (which only matches a final aggregate 
whose child is the partial
+      // aggregate) must not combine them.
+      checkNumAggregation(
+        "SELECT k + 1, count(*) FROM (SELECT /*+ repartition(k) */ * FROM t) 
GROUP BY k + 1",
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 2)
+
+      // Same non-adjacency guard for an object-hash aggregate (`collect_set` 
is imperative).
+      checkNumAggregation(
+        "SELECT k + 1, collect_set(v) FROM (SELECT /*+ repartition(k) */ * 
FROM t) GROUP BY k + 1",
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 2)
+    }
+  }
+
+  test("Combine with a string grouping key and a HAVING clause") {
+    withTempView("t") {
+      spark.range(20).selectExpr("cast(id % 4 as string) as s", "id % 7 as v")
+        .createOrReplaceTempView("t")
+
+      // string grouping key
+      checkNumAggregation(
+        "SELECT s, count(*) FROM (SELECT /*+ repartition(s) */ * FROM t) GROUP 
BY s",
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 1)
+
+      // aggregate with a FILTER clause and a HAVING predicate
+      checkNumAggregation(
+        """SELECT s, count(*) FILTER (WHERE v > 2), sum(v)
+          |FROM (SELECT /*+ repartition(s) */ * FROM t) GROUP BY s HAVING 
sum(v) > 5""".stripMargin,
+        numAggWithDisabled = 2,
+        numAggWithEnabled = 1)
+    }
+  }
+
+  test("SPARK-43317: Combined aggregate keeps the FILTER clause of aggregate 
functions") {
+    // The FILTER (WHERE ...) clause only lives on the partial aggregate 
expressions; the final
+    // aggregate merely merges the partial aggregation buffers. Combining must 
therefore take the
+    // aggregate functions (and their filters) from the partial aggregate, 
otherwise the filter is
+    // silently dropped and the result is wrong.
+    withTempView("t") {
+      spark.range(20).selectExpr("id % 3 as k", "id as v")
+        .createOrReplaceTempView("t")
+      val query =
+        """SELECT k, count(*) FILTER (WHERE v > 5), sum(v) FILTER (WHERE v < 
15), count(*)
+          |FROM (SELECT /*+ repartition(k) */ * FROM t) GROUP BY 
k""".stripMargin
+
+      val expected = 
withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false") {
+        sql(query).collect()
+      }
+      withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true") {
+        val df = sql(query)
+        checkAnswer(df, expected)
+        assert(collect(df.queryExecution.executedPlan) {
+          case agg: BaseAggregateExec => agg
+        }.size == 1)
+      }
+
+      // A `count(distinct)` with a FILTER clause. The distinct-with-filter 
rewrite expands the
+      // input (adding a group id) and shuffles on the distinct key and again 
on the grouping key,
+      // so an Exchange sits between the partial and final aggregate. They are 
therefore not
+      // adjacent and the rule must not combine them; the FILTER must still 
produce correct results.
+      val distinctQuery =
+        """SELECT k, count(distinct v) FILTER (WHERE v > 5)
+          |FROM (SELECT /*+ repartition(k) */ * FROM t) GROUP BY 
k""".stripMargin
+      val distinctExpected =
+        withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> 
"false") {
+          val df = sql(distinctQuery)
+          assert(collect(df.queryExecution.executedPlan) {
+            case agg: BaseAggregateExec => agg
+          }.size == 4)
+          df.collect()
+        }
+      withSQLConf(SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true") {
+        val df = sql(distinctQuery)
+        checkAnswer(df, distinctExpected)
+        assert(collect(df.queryExecution.executedPlan) {
+          case agg: BaseAggregateExec => agg
+        }.size == 4)
+      }
+    }
+  }
+
+  test("Default value falls back to 
spark.sql.execution.replaceHashWithSortAgg") {
+    withTempView("t") {
+      spark.range(20).selectExpr("id % 3 as k", "id % 7 as v")
+        .createOrReplaceTempView("t")
+      val query = "SELECT k, count(*) FROM (SELECT /*+ repartition(k) */ * 
FROM t) GROUP BY k"
+
+      // `spark.sql.execution.combineAdjacentAggregation` is unset here, so it 
falls back to
+      // `spark.sql.execution.replaceHashWithSortAgg`.
+      withSQLConf(SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "false") {
+        assert(numAggregates(query) == 2)
+      }
+      withSQLConf(SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "true") {
+        assert(numAggregates(query) == 1)
+      }
+
+      // An explicit value overrides the fallback.
+      withSQLConf(
+          SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "true",
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false") {
+        assert(numAggregates(query) == 2)
+      }
+    }
+  }
+}
diff --git 
a/sql/core/src/test/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAggSuite.scala
 
b/sql/core/src/test/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAggSuite.scala
index 41edb53441f7..298e51bbb7c0 100644
--- 
a/sql/core/src/test/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAggSuite.scala
+++ 
b/sql/core/src/test/scala/org/apache/spark/sql/execution/ReplaceHashWithSortAggSuite.scala
@@ -113,7 +113,7 @@ abstract class ReplaceHashWithSortAggSuiteBase
              |)
              |GROUP BY key
            """.stripMargin
-        checkAggs(query, 2, 0, 2, 0)
+        checkAggs(query, 1, 0, 2, 0)
       }
     }
   }
@@ -131,6 +131,76 @@ abstract class ReplaceHashWithSortAggSuiteBase
       }
     }
   }
+
+  test("combine adjacent aggregate then replace it with sort aggregate") {
+    withTempView("t1", "t2") {
+      spark.range(100).selectExpr("id as key").createOrReplaceTempView("t1")
+      spark.range(50).selectExpr("id as key").createOrReplaceTempView("t2")
+      // The partial and final aggregate sit on top of a sort merge join, so 
the child of the
+      // combined `Complete` aggregate is already ordered on the grouping key.
+      val query =
+        """SELECT key, count(key) FROM (
+          |  SELECT /*+ SHUFFLE_MERGE(t1) */ t1.key AS key FROM t1 JOIN t2 ON 
t1.key = t2.key
+          |) GROUP BY key""".stripMargin
+
+      val expected = withSQLConf(
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false",
+          SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "false") {
+        sql(query).collect()
+      }
+
+      // Combine only: the partial and final hash aggregate are merged into a 
single hash aggregate.
+      withSQLConf(
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true",
+          SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "false") {
+        val df = sql(query)
+        checkNumAggs(df, hashAggCount = 1, sortAggCount = 0)
+        checkAnswer(df, expected)
+      }
+
+      // Combine + replace: the combined hash aggregate is further replaced 
with a sort aggregate.
+      withSQLConf(
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true",
+          SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "true") {
+        val df = sql(query)
+        checkNumAggs(df, hashAggCount = 0, sortAggCount = 1)
+        checkAnswer(df, expected)
+      }
+    }
+  }
+
+  test("SPARK-43317: Combined sort aggregate keeps the FILTER clause") {
+    // Regression test for the wrong-results bug this fix addresses: when the 
partial and final
+    // aggregate are merged into a single `Complete` mode aggregate that is 
then replaced with a
+    // `SortAggregateExec` (`replaceHashWithSortAgg=true`), the `FILTER (WHERE 
...)` clause must be
+    // preserved. The filter only lives on the partial aggregate expressions, 
so combining must take
+    // the aggregate functions from the partial aggregate; otherwise the 
filter is silently dropped
+    // and the result is wrong. This drives the PR's own repro query 
end-to-end through the sort
+    // aggregate path.
+    withTempView("t1", "t2") {
+      spark.range(100).selectExpr("id as key").createOrReplaceTempView("t1")
+      spark.range(50).selectExpr("id as key").createOrReplaceTempView("t2")
+      val query =
+        """SELECT key, count(*) FILTER (WHERE key > 10) FROM (
+          |  SELECT /*+ SHUFFLE_MERGE(t1) */ t1.key AS key FROM t1 JOIN t2 ON 
t1.key = t2.key
+          |) GROUP BY key""".stripMargin
+
+      val expected = withSQLConf(
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "false",
+          SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "false") {
+        sql(query).collect()
+      }
+
+      withSQLConf(
+          SQLConf.COMBINE_ADJACENT_AGGREGATION_ENABLED.key -> "true",
+          SQLConf.REPLACE_HASH_WITH_SORT_AGG_ENABLED.key -> "true") {
+        val df = sql(query)
+        // Combined into a single `Complete` mode sort aggregate.
+        checkNumAggs(df, hashAggCount = 0, sortAggCount = 1)
+        checkAnswer(df, expected)
+      }
+    }
+  }
 }
 
 class ReplaceHashWithSortAggSuite extends ReplaceHashWithSortAggSuiteBase


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]


Reply via email to