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

github-merge-queue[bot] pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/datafusion-comet.git


The following commit(s) were added to refs/heads/main by this push:
     new 6065705c16 fix: revert unsafe partial aggregates after final fallback 
(#5421)
6065705c16 is described below

commit 6065705c16340c0be293212a71decfd9df4daae4
Author: Chao Sun <[email protected]>
AuthorDate: Sun Sep 20 23:19:33 2026 +0000

    fix: revert unsafe partial aggregates after final fallback (#5421)
    
    * fix: revert unsafe partial aggregates after final fallback
    
    * fix: distinguish aggregate buffer compatibility by direction
    
    * fix: block unsafe native AVG partials before Spark final
    
    * test: cover aggregate fallback for unsupported array hash keys
    
    * test: align aggregate fallback checks across Spark versions
    
    Capture the Spark baseline inside withSQLConf so Spark 3 does not return
    Unit in place of the expected rows.
    
    Expect AVG to remain in Spark after Celeborn shuffle fallback, consistent
    with its unsafe empty partial buffer, and retain COUNT as the compatible
    native-partial control.
    
    * fix: clarify aggregate buffer compatibility and fallback
    
    * ci: group aggregate policy tests with aggregate suite
    
    * test: cover DISTINCT aggregate fallback boundaries
    
    Cover nondecimal AVG with SUM(DISTINCT) across global and grouped plans, 
shuffle modes, and selective lower-exchange fallback. Document why native-only 
fallback reaches the Final repair and qualify the auto/JVM boundary without 
changing the production predicate.
    
    * refactor: simplify aggregate fallback repair and regression tests
    
    * fix: address aggregate fallback review follow-ups
    
    Restore conditional diagnostics for known unrepaired buffer paths and 
focused coverage for direction defaults, immediate fallback reasons, and sticky 
shuffle refusal.
    
    Document the decimal SUM precision-overflow mismatch, preserve its fallback 
with a cancellation regression, and track remaining audits in #5975.
---
 .../org/apache/comet/rules/CometExecRule.scala     | 121 ++++++-
 .../RevertNativeForTransitionHeavyStages.scala     |   6 +-
 .../serde/CometAggregateExpressionSerde.scala      |  19 +-
 .../org/apache/comet/serde/QueryPlanSerde.scala    |  30 +-
 .../scala/org/apache/comet/serde/aggregates.scala  |  57 +++-
 .../org/apache/spark/sql/comet/operators.scala     |   2 +-
 .../apache/comet/exec/CometAggregateSuite.scala    | 367 ++++++++++++++++++++-
 .../apache/comet/rules/CometExecRuleSuite.scala    | 209 ++++++++++--
 .../CometCelebornShufflePlanningSuite.scala        |  21 +-
 9 files changed, 749 insertions(+), 83 deletions(-)

diff --git a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala 
b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
index 265c84a0df..54f8891624 100644
--- a/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
+++ b/spark/src/main/scala/org/apache/comet/rules/CometExecRule.scala
@@ -220,6 +220,16 @@ case class CometExecRule(session: SparkSession)
 
   private def isCometNative(op: SparkPlan): Boolean = 
op.isInstanceOf[CometNativeExec]
 
+  /**
+   * Restore a Spark Partial while retaining its current children. The tag 
prevents reconversion
+   * when AQE replans the exchange without its Final, and records why the 
Partial stays in Spark.
+   */
+  private def restoreSparkPartial(agg: CometHashAggregateExec, reason: 
String): SparkPlan = {
+    val partial = agg.originalPlan.withNewChildren(agg.children)
+    partial.setTagValue(CometExecRule.COMET_UNSAFE_PARTIAL, reason)
+    withFallbackReason(partial, reason)
+  }
+
   /**
    * A Celeborn exchange can fall back after its child has been converted, for 
example because of
    * the partition threshold or an unsupported hash key. Keep incompatible 
partial aggregate
@@ -237,10 +247,8 @@ case class CometExecRule(session: SparkSession)
       case _: QueryStageExec | _: ShuffleExchangeLike | _: 
BroadcastExchangeLike => plan
       case agg: CometHashAggregateExec
           if agg.modes == Seq(Partial) &&
-            
!QueryPlanSerde.allAggsSupportMixedExecution(agg.aggregateExpressions) =>
-        val sparkAggregate = agg.originalPlan.withNewChildren(agg.children)
-        sparkAggregate.setTagValue(CometExecRule.COMET_UNSAFE_PARTIAL, reason)
-        withFallbackReason(sparkAggregate, reason)
+            
!QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions)
 =>
+        restoreSparkPartial(agg, reason)
       // Final output is ordinary SQL data; any partial below it belongs to 
another aggregate.
       case agg: CometHashAggregateExec if agg.modes.contains(Final) => agg
       case agg: BaseAggregateExec if agg.aggregateExpressions.exists(_.mode == 
Final) => agg
@@ -758,7 +766,7 @@ case class CometExecRule(session: SparkSession)
       // during the bottom-up conversion. Tags persist through AQE stage 
creation.
       tagUnsafePartialAggregates(planWithJoinRewritten)
 
-      var newPlan = transform(planWithJoinRewritten)
+      var newPlan = 
revertUnsafePartialAggregates(transform(planWithJoinRewritten))
 
       // if the plan cannot be run fully natively then explain why (when 
appropriate
       // config is enabled)
@@ -1130,7 +1138,7 @@ case class CometExecRule(session: SparkSession)
         val consumerMode: AggregateMode =
           if (modes.contains(PartialMerge)) PartialMerge else Final
         if (consumesBuffers &&
-          
!QueryPlanSerde.allAggsSupportMixedExecution(agg.aggregateExpressions) &&
+          
!QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions)
 &&
           !canAggregateBeConverted(agg, consumerMode)) {
           findPartialAggInPlan(agg.child).foreach { partial =>
             // Only tag if the Partial would otherwise have been converted. If 
the Partial itself
@@ -1173,6 +1181,103 @@ case class CometExecRule(session: SparkSession)
     }
   }
 
+  /**
+   * Inspect a failed repair's buffer path without rewriting it or 
materializing any stage. Report
+   * only a native Partial/PartialMerge whose emitted state is not known to be 
Spark-compatible.
+   * Spark Partials and completed aggregates establish new buffers, so stop 
there rather than
+   * finding an unrelated native producer below them. Only known aggregate and 
exchange wrappers
+   * forward the same buffer path; an arbitrary operator is not evidence of a 
mixed boundary.
+   */
+  private def hasUnrepairedNativeBuffer(plan: SparkPlan): Boolean = plan match 
{
+    case agg: CometHashAggregateExec if agg.aggregateExpressions.isEmpty =>
+      hasUnrepairedNativeBuffer(agg.child)
+    case agg: CometHashAggregateExec =>
+      agg.modes.forall(m => m == Partial || m == PartialMerge) &&
+      
!QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions)
+    case agg: BaseAggregateExec
+        if agg.aggregateExpressions.nonEmpty &&
+          agg.aggregateExpressions.forall(_.mode == Partial) =>
+      false
+    case agg: BaseAggregateExec =>
+      agg.aggregateExpressions.forall(e => e.mode == Partial || e.mode == 
PartialMerge) &&
+      hasUnrepairedNativeBuffer(agg.child)
+    case placeholder: CometSinkPlaceHolder => 
hasUnrepairedNativeBuffer(placeholder.child)
+    case read: AQEShuffleReadExec => hasUnrepairedNativeBuffer(read.child)
+    case stage: ShuffleQueryStageExec => hasUnrepairedNativeBuffer(stage.plan)
+    case reused: ReusedExchangeExec => hasUnrepairedNativeBuffer(reused.child)
+    case shuffle: CometShuffleExchangeExec => 
hasUnrepairedNativeBuffer(shuffle.child)
+    case shuffle: ShuffleExchangeExec => 
hasUnrepairedNativeBuffer(shuffle.child)
+    case _ => false
+  }
+
+  /**
+   * The early tagging pass cannot know whether a Final's child will become 
native. Check the
+   * actual conversion result before serialization or AQE stage creation, 
restoring the feeding
+   * aggregate/exchange chain while keeping native work below its Partial. 
Return the repaired
+   * plan, or preserve an unrepairable path and record one warning on its 
Spark Final if an unsafe
+   * native producer remains. Existing stages and their buffers are never 
rewritten by this pass.
+   */
+  private[rules] def revertUnsafePartialAggregates(plan: SparkPlan): SparkPlan 
= {
+    def revertChain(node: SparkPlan): Option[SparkPlan] = node match {
+      case agg: CometHashAggregateExec if agg.modes == Seq(Partial) =>
+        Some(
+          restoreSparkPartial(
+            agg,
+            "Partial aggregate disabled: corresponding final aggregate " +
+              "cannot be converted to Comet and intermediate buffer formats 
are incompatible"))
+
+      case agg: CometHashAggregateExec
+          if agg.modes.forall(m => m == Partial || m == PartialMerge) =>
+        revertChain(agg.child).map(child => 
agg.originalPlan.withNewChildren(Seq(child)))
+
+      case agg: BaseAggregateExec
+          if agg.aggregateExpressions.nonEmpty &&
+            agg.aggregateExpressions.forall(_.mode == Partial) =>
+        // This producer already emits Spark buffers. Do not reach through it 
to an unrelated
+        // aggregate below it.
+        None
+
+      case agg: BaseAggregateExec
+          if agg.aggregateExpressions.forall(e => e.mode == Partial || e.mode 
== PartialMerge) =>
+        revertChain(agg.child).map(child => agg.withNewChildren(Seq(child)))
+
+      case CometSinkPlaceHolder(_, _, shuffle: CometShuffleExchangeExec) =>
+        revertChain(shuffle)
+      case shuffle: CometShuffleExchangeExec =>
+        revertChain(shuffle.child).map(child => 
shuffle.originalPlan.withNewChildren(Seq(child)))
+      case shuffle: ShuffleExchangeExec =>
+        revertChain(shuffle.child).map(child => 
shuffle.withNewChildren(Seq(child)))
+
+      // Stop at materialized stages and operators outside the feeding 
aggregate/exchange chain.
+      case _ => None
+    }
+
+    plan.transformUp {
+      case agg: BaseAggregateExec
+          if agg.aggregateExpressions.map(_.mode).distinct == Seq(Final) &&
+            
!QueryPlanSerde.allAggsSupportNativePartialToSparkFinal(agg.aggregateExpressions)
 =>
+        revertChain(agg.child)
+          // Rebuild native consumers and shuffles from their original Spark 
operators. Merely
+          // replacing their children would leave a native protobuf reading 
the old buffers.
+          .map(child => transform(agg.withNewChildren(Seq(child))))
+          .getOrElse {
+            if (hasUnrepairedNativeBuffer(agg.child)) {
+              val reason = "Comet could not restore a native intermediate 
buffer producer " +
+                "below Spark final aggregate; the remaining buffer may be 
incompatible"
+              // AQE can revisit the same consumer. Record the explanation and 
warn once,
+              // regardless of whether general fallback logging is enabled.
+              if (!agg
+                  .getTagValue(CometExplainInfo.FALLBACK_REASONS)
+                  .exists(_.contains(reason))) {
+                if (!CometConf.COMET_EXPLAIN_FALLBACK_LOG_ENABLED.get()) 
logWarning(reason)
+                withFallbackReason(agg, reason)
+              }
+            }
+            agg
+          }
+    }
+  }
+
   /**
    * Look for the bottom Partial-mode aggregate that feeds into the given plan 
(the child of a
    * Final). Walks through exchanges and AQE stages, and continues down 
through intermediate
@@ -1202,8 +1307,8 @@ case class CometExecRule(session: SparkSession)
   /**
    * Conservative check for whether an aggregate could be converted to Comet. 
Checks operator
    * enablement, grouping expressions, aggregate expressions, and result 
expressions.
-   * Intentionally skips the sparkFinalMode / child-native checks since those 
depend on
-   * transformation state.
+   * Intentionally skips the child-native checks since those depend on 
transformation state;
+   * [[revertUnsafePartialAggregates]] checks the actual conversion result 
before execution.
    *
    * WARNING: this intentionally mirrors the predicate checks in 
`CometBaseAggregate.doConvert`
    * (operators.scala). Any change to the convertibility rules there must be 
reflected here or
diff --git 
a/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala
 
b/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala
index 9a941c28ff..a56e477b03 100644
--- 
a/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala
+++ 
b/spark/src/main/scala/org/apache/comet/rules/RevertNativeForTransitionHeavyStages.scala
@@ -124,7 +124,11 @@ case class RevertNativeForTransitionHeavyStages(session: 
SparkSession)
     def visit(plan: SparkPlan): Boolean = plan match {
       case _ if isStageBoundary(plan) => false
       case aggregate: CometHashAggregateExec
-          if 
!QueryPlanSerde.allAggsSupportMixedExecution(aggregate.aggregateExpressions) =>
+          if !QueryPlanSerde
+            
.allAggsSupportNativePartialToSparkFinal(aggregate.aggregateExpressions) ||
+            QueryPlanSerde
+              
.aggsNotSupportingSparkPartialToNativeFinal(aggregate.aggregateExpressions)
+              .nonEmpty =>
         val producesBuffer =
           aggregate.modes.exists(mode => mode == Partial || mode == 
PartialMerge)
         val consumesAcrossBoundary =
diff --git 
a/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
 
b/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
index a52d600821..73630a4e8b 100644
--- 
a/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
+++ 
b/spark/src/main/scala/org/apache/comet/serde/CometAggregateExpressionSerde.scala
@@ -82,16 +82,21 @@ trait CometAggregateExpressionSerde[T <: AggregateFunction] 
{
   def getSupportLevel(expr: T): SupportLevel = Compatible(None)
 
   /**
-   * Whether this aggregate's intermediate buffer format is compatible between 
Spark and Comet for
-   * the given function instance, making it safe to run the Partial in one 
engine and the Final in
-   * the other. Aggregates with simple single-value buffers (MIN, MAX, 
bitwise) are always safe;
-   * SUM and non-decimal AVG match Spark's buffer and are safe except where 
noted per instance
-   * (e.g. TRY-mode SUM uses a Comet-internal flag column). COUNT is 
intentionally excluded
-   * despite a matching buffer: mixed COUNT partial/final regressed AQE's
+   * Whether a Comet aggregate can consume this function's Spark intermediate 
buffer. This covers
+   * Spark Partial to Comet Final, including intermediate PartialMerge stages. 
COUNT is excluded
+   * despite a matching buffer: a Comet Final above a Spark Partial regressed 
AQE's
    * PropagateEmptyRelationAfterAQE pattern (which matches BaseAggregateExec 
only) and the Spark
    * 4.0 count-bug decorrelation for correlated IN subqueries.
    */
-  def supportsMixedPartialFinal(fn: T): Boolean = false
+  def supportsSparkPartialToNativeFinal(fn: T): Boolean = false
+
+  /**
+   * Whether Spark can consume this function's Comet intermediate buffer. Opt 
in independently
+   * from the reverse direction: consuming Spark state does not establish that 
Comet emits state
+   * Spark can merge, especially from a never-updated or all-null partial 
accumulator. Remaining
+   * forward-compatibility audits are tracked in issue #5975.
+   */
+  def supportsNativePartialToSparkFinal(fn: T): Boolean = false
 
   /**
    * Convert a Spark expression into a protocol buffer representation that can 
be passed into
diff --git a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala 
b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
index eced2d7fdf..43d65fd418 100644
--- a/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/QueryPlanSerde.scala
@@ -443,31 +443,35 @@ object QueryPlanSerde extends Logging with CometExprShim 
with CometTypeShim {
     classOf[VarianceSamp] -> CometVarianceSamp)
 
   /**
-   * Returns true if all aggregate expressions in the list have intermediate 
buffer formats that
-   * are compatible between Spark and Comet, making it safe to run Partial in 
one engine and Final
-   * in the other.
+   * Returns true if Spark can consume all the intermediate buffers produced 
by Comet. Used when a
+   * Spark Final would otherwise consume a native Partial, including after 
shuffle fallback.
    */
-  def allAggsSupportMixedExecution(aggExprs: Seq[AggregateExpression]): 
Boolean = {
-    aggExprs.forall(aggExpr => 
supportsMixedExecution(aggExpr.aggregateFunction))
+  def allAggsSupportNativePartialToSparkFinal(aggExprs: 
Seq[AggregateExpression]): Boolean = {
+    aggExprs.forall { aggExpr =>
+      val fn = aggExpr.aggregateFunction
+      aggrSerdeMap.get(fn.getClass).exists { handler =>
+        handler
+          .asInstanceOf[CometAggregateExpressionSerde[AggregateFunction]]
+          .supportsNativePartialToSparkFinal(fn)
+      }
+    }
   }
 
   /**
-   * Returns the aggregate functions in the list whose intermediate buffer 
formats are not known
-   * to be compatible between Spark and Comet. These are the functions that 
prevent a Spark Final
-   * aggregate (without a Comet Partial) from running, since the buffer 
produced by one engine
-   * cannot be safely consumed by the other.
+   * Returns functions whose Spark intermediate buffers cannot safely be 
consumed by a Comet Final
+   * or PartialMerge. This is independent of native Partial to Spark Final 
compatibility.
    */
-  def aggsNotSupportingMixedExecution(
+  def aggsNotSupportingSparkPartialToNativeFinal(
       aggExprs: Seq[AggregateExpression]): Seq[AggregateFunction] = {
-    aggExprs.map(_.aggregateFunction).filterNot(supportsMixedExecution)
+    
aggExprs.map(_.aggregateFunction).filterNot(supportsSparkPartialToNativeFinal)
   }
 
-  private def supportsMixedExecution(fn: AggregateFunction): Boolean = {
+  private def supportsSparkPartialToNativeFinal(fn: AggregateFunction): 
Boolean = {
     aggrSerdeMap.get(fn.getClass) match {
       case Some(handler) =>
         handler
           .asInstanceOf[CometAggregateExpressionSerde[AggregateFunction]]
-          .supportsMixedPartialFinal(fn)
+          .supportsSparkPartialToNativeFinal(fn)
       case None => false
     }
   }
diff --git a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala 
b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala
index b3bc5e86b7..ad53cceb31 100644
--- a/spark/src/main/scala/org/apache/comet/serde/aggregates.scala
+++ b/spark/src/main/scala/org/apache/comet/serde/aggregates.scala
@@ -36,7 +36,10 @@ import org.apache.comet.shims.{CometCollectShim, 
CometEvalModeUtil, CometTypeShi
 
 object CometMin extends CometAggregateExpressionSerde[Min] {
 
-  override def supportsMixedPartialFinal(fn: Min): Boolean = true
+  override def supportsSparkPartialToNativeFinal(fn: Min): Boolean = true
+
+  // Native MIN emits one typed null for empty/all-null input; Spark's least 
merge ignores it.
+  override def supportsNativePartialToSparkFinal(fn: Min): Boolean = true
 
   override def getSupportLevel(expr: Min): SupportLevel =
     AggSerde.minMaxSupportLevel(expr.dataType)
@@ -72,7 +75,10 @@ object CometMin extends CometAggregateExpressionSerde[Min] {
 
 object CometMax extends CometAggregateExpressionSerde[Max] {
 
-  override def supportsMixedPartialFinal(fn: Max): Boolean = true
+  override def supportsSparkPartialToNativeFinal(fn: Max): Boolean = true
+
+  // Native MAX emits one typed null for empty/all-null input; Spark's 
greatest merge ignores it.
+  override def supportsNativePartialToSparkFinal(fn: Max): Boolean = true
 
   override def getSupportLevel(expr: Max): SupportLevel =
     AggSerde.minMaxSupportLevel(expr.dataType)
@@ -190,6 +196,10 @@ object CometMinBy extends CometMaxMinBy[MinBy] {
 }
 
 object CometCount extends CometAggregateExpressionSerde[Count] {
+  // Both buffers are a single non-null Long. The AQE/count-bug restrictions 
documented on the
+  // reverse direction concern a Comet Final; retaining Spark's Final 
preserves those rewrites.
+  override def supportsNativePartialToSparkFinal(fn: Count): Boolean = true
+
   override def convert(
       aggExpr: AggregateExpression,
       expr: Count,
@@ -213,7 +223,10 @@ object CometCount extends 
CometAggregateExpressionSerde[Count] {
 
 object CometAverage extends CometAggregateExpressionSerde[Average] {
 
-  override def supportsMixedPartialFinal(fn: Average): Boolean =
+  // Keep the default native-to-Spark restriction until #5420: an untouched 
native AVG emits
+  // (null, 0), but Spark's merge needs (0.0, 0).
+
+  override def supportsSparkPartialToNativeFinal(fn: Average): Boolean =
     // Non-decimal AVG has a (sum: double, count: long) buffer matching Spark. 
Decimal AVG is
     // deferred (overflow nulls count differently) and stays unsafe for mixed 
execution.
     !fn.child.dataType.isInstanceOf[DecimalType]
@@ -273,7 +286,17 @@ object CometAverage extends 
CometAggregateExpressionSerde[Average] {
 
 object CometSum extends CometAggregateExpressionSerde[Sum] {
 
-  override def supportsMixedPartialFinal(fn: Sum): Boolean =
+  // Non-decimal, non-TRY SUM emits one nullable sum, including null for 
empty/all-null input;
+  // Spark's coalesce-based merge accepts it. Decimal SUM has Spark's (sum, 
isEmpty) layout,
+  // but native updates make precision overflow sticky (or throw in ANSI 
mode). Spark's generated
+  // scalar SUM can recover before emitting its partial: decimal(38,38) inputs 
0.6, 0.6, -0.6
+  // sum to 0.6. Keep decimal partials in Spark until those update semantics 
match. Integer TRY
+  // SUM also remains excluded because its native state contains an extra 
has_all_nulls column.
+  override def supportsNativePartialToSparkFinal(fn: Sum): Boolean =
+    !fn.child.dataType.isInstanceOf[DecimalType] &&
+      CometEvalModeUtil.fromSparkEvalMode(CometEvalModeUtil.sumEvalMode(fn)) 
!= CometEvalMode.TRY
+
+  override def supportsSparkPartialToNativeFinal(fn: Sum): Boolean =
     // Decimal SUM is excluded: overflow detection (ANSI throw / Legacy null) 
does not survive a
     // Spark-partial / Comet-final split, so the required ArithmeticException 
is never raised.
     // TRY-mode integer SUM carries a Comet-internal has_all_nulls column that 
Spark cannot read.
@@ -390,7 +413,10 @@ object CometLast extends 
CometAggregateExpressionSerde[Last] {
 }
 
 object CometBitAndAgg extends CometAggregateExpressionSerde[BitAndAgg] {
-  override def supportsMixedPartialFinal(fn: BitAndAgg): Boolean = true
+  override def supportsSparkPartialToNativeFinal(fn: BitAndAgg): Boolean = true
+
+  // The single native buffer is null for empty/all-null input; Spark's merge 
skips nulls.
+  override def supportsNativePartialToSparkFinal(fn: BitAndAgg): Boolean = true
 
   override def getSupportLevel(expr: BitAndAgg): SupportLevel =
     if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) {
@@ -428,7 +454,10 @@ object CometBitAndAgg extends 
CometAggregateExpressionSerde[BitAndAgg] {
 }
 
 object CometBitOrAgg extends CometAggregateExpressionSerde[BitOrAgg] {
-  override def supportsMixedPartialFinal(fn: BitOrAgg): Boolean = true
+  override def supportsSparkPartialToNativeFinal(fn: BitOrAgg): Boolean = true
+
+  // The single native buffer is null for empty/all-null input; Spark's merge 
skips nulls.
+  override def supportsNativePartialToSparkFinal(fn: BitOrAgg): Boolean = true
 
   override def getSupportLevel(expr: BitOrAgg): SupportLevel =
     if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) {
@@ -466,7 +495,10 @@ object CometBitOrAgg extends 
CometAggregateExpressionSerde[BitOrAgg] {
 }
 
 object CometBitXOrAgg extends CometAggregateExpressionSerde[BitXorAgg] {
-  override def supportsMixedPartialFinal(fn: BitXorAgg): Boolean = true
+  override def supportsSparkPartialToNativeFinal(fn: BitXorAgg): Boolean = true
+
+  // The single native buffer is null for empty/all-null input; Spark's merge 
skips nulls.
+  override def supportsNativePartialToSparkFinal(fn: BitXorAgg): Boolean = true
 
   override def getSupportLevel(expr: BitXorAgg): SupportLevel =
     if (AggSerde.bitwiseAggTypeSupported(expr.dataType)) {
@@ -1015,7 +1047,11 @@ object CometRegrReplacement
 
 object CometBloomFilterAggregate extends 
CometAggregateExpressionSerde[BloomFilterAggregate] {
 
-  override def supportsMixedPartialFinal(fn: BloomFilterAggregate): Boolean = 
true
+  override def supportsSparkPartialToNativeFinal(fn: BloomFilterAggregate): 
Boolean = true
+
+  // Native state is Spark's serialized filter, non-null even for 
empty/all-null input; only
+  // the final result may be null, so Spark's deserialize always receives a 
valid filter.
+  override def supportsNativePartialToSparkFinal(fn: BloomFilterAggregate): 
Boolean = true
 
   override def getSupportLevel(expr: BloomFilterAggregate): SupportLevel =
     expr.child.dataType match {
@@ -1227,7 +1263,10 @@ object CometApproxCountDistinct extends 
CometAggregateExpressionSerde[HyperLogLo
   // The register buffer uses Spark's identical packed-`Long` layout 
(`numWords` `Long` columns),
   // matching Spark's `aggBufferSchema`, so a Comet partial and Spark final 
(or the reverse) can
   // be mixed in one plan.
-  override def supportsMixedPartialFinal(fn: HyperLogLogPlusPlus): Boolean = 
true
+  override def supportsSparkPartialToNativeFinal(fn: HyperLogLogPlusPlus): 
Boolean = true
+
+  // Native empty/all-null state contains non-null zero Long words, matching 
Spark's registers.
+  override def supportsNativePartialToSparkFinal(fn: HyperLogLogPlusPlus): 
Boolean = true
 
   // Types that Comet's native `xxhash64` hashes identically to Spark's 
`XxHash64Function`.
   // `StringType` here is the default UTF8_BINARY collation; a collated 
`StringType(collationId)`
diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala 
b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala
index fe1a2e637a..fa3b5da240 100644
--- a/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala
+++ b/spark/src/main/scala/org/apache/spark/sql/comet/operators.scala
@@ -1891,7 +1891,7 @@ trait CometBaseAggregate {
 
     if (missingCometProducer) {
       val incompatibleAggs =
-        
QueryPlanSerde.aggsNotSupportingMixedExecution(aggregate.aggregateExpressions)
+        
QueryPlanSerde.aggsNotSupportingSparkPartialToNativeFinal(aggregate.aggregateExpressions)
       if (incompatibleAggs.nonEmpty) {
         val names = 
incompatibleAggs.map(_.prettyName).distinct.sorted.mkString(", ")
         withFallbackReason(
diff --git 
a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala 
b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala
index 7cf5e69d6f..55c2ebdb57 100644
--- a/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala
+++ b/spark/src/test/scala/org/apache/comet/exec/CometAggregateSuite.scala
@@ -30,19 +30,21 @@ import org.apache.spark.sql.{CometTestBase, DataFrame, Row}
 import org.apache.spark.sql.catalyst.expressions.Cast
 import org.apache.spark.sql.catalyst.expressions.aggregate.{Final, Partial, 
PartialMerge}
 import org.apache.spark.sql.catalyst.optimizer.EliminateSorts
-import org.apache.spark.sql.catalyst.plans.physical.RangePartitioning
-import org.apache.spark.sql.comet.CometHashAggregateExec
+import org.apache.spark.sql.catalyst.plans.physical.{HashPartitioning, 
RangePartitioning}
+import org.apache.spark.sql.comet.{CometFilterExec, CometHashAggregateExec, 
CometNativeExec, CometProjectExec}
 import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec
 import org.apache.spark.sql.execution.SQLExecution
-import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, 
AdaptiveSparkPlanHelper}
-import org.apache.spark.sql.execution.exchange.ReusedExchangeExec
+import org.apache.spark.sql.execution.adaptive.{AdaptiveSparkPlanExec, 
AdaptiveSparkPlanHelper, ShuffleQueryStageExec}
+import org.apache.spark.sql.execution.aggregate.BaseAggregateExec
+import org.apache.spark.sql.execution.exchange.{ReusedExchangeExec, 
ShuffleExchangeExec}
 import org.apache.spark.sql.functions.{avg, col, count_distinct, expr, sum}
 import org.apache.spark.sql.internal.SQLConf
-import org.apache.spark.sql.types.{DataTypes, StructField, StructType}
+import org.apache.spark.sql.types.{ArrayType, DataTypes, StructField, 
StructType}
 
 import org.apache.comet.CometConf
 import org.apache.comet.CometConf.COMET_EXEC_STRICT_FLOATING_POINT
 import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus
+import org.apache.comet.rules.CometExecRule
 import org.apache.comet.serde.RegrSparkVersions
 import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator, 
ParquetGenerator, SchemaGenOptions}
 
@@ -356,6 +358,290 @@ class CometAggregateSuite extends CometTestBase with 
AdaptiveSparkPlanHelper {
     }
   }
 
+  for (adaptive <- Seq(false, true)) {
+    test(s"decimal AVG falls back across a Spark shuffle (AQE=$adaptive)") {
+      withTempDir { dir =>
+        val path = s"${dir.getAbsolutePath}/data"
+        withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+          spark
+            .range(0L, 8L, 1L, 4)
+            .selectExpr("id", "CAST(200 AS DECIMAL(20, 2)) AS v")
+            .write
+            .parquet(path)
+        }
+
+        withSQLConf(
+          SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString,
+          SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576",
+          SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false",
+          CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "false",
+          CometConf.COMET_CONVERT_FROM_PARQUET_ENABLED.key -> "true",
+          CometConf.COMET_SHUFFLE_ENABLED.key -> "false") {
+          withParquetTable(path, "decimal_avg_fallback") {
+            // The filter leaves three input partitions empty. Decimal AVG is 
not safe to mix
+            // between engines: a native empty partial can poison the Spark 
final's sum buffer.
+            val df = sql("SELECT AVG(v) FROM decimal_avg_fallback WHERE id = 
1")
+            val initialPlan = stripAQEPlan(df.queryExecution.executedPlan)
+            checkAnswer(df, Seq(Row(new java.math.BigDecimal("200.000000"))))
+            for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) {
+              assert(collect(plan) { case agg: CometHashAggregateExec => agg 
}.isEmpty)
+              val partials = collect(plan) {
+                case agg: BaseAggregateExec
+                    if agg.aggregateExpressions.forall(_.mode == Partial) =>
+                  agg
+              }
+              assert(partials.size == 1)
+              
assert(partials.forall(_.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined))
+              // Falling back the aggregate must not discard the native 
filter/scan conversion.
+              assert(collect(plan) { case filter: CometFilterExec => filter 
}.nonEmpty)
+            }
+            if (adaptive) {
+              val stages = collect(df.queryExecution.executedPlan) {
+                case stage: ShuffleQueryStageExec => stage
+              }
+              assert(stages.nonEmpty && stages.forall(_.isMaterialized))
+            }
+
+            // Compatible buffers may still use a native Partial and a Spark 
Final.
+            val safe = sql("SELECT MIN(v), MAX(v) FROM decimal_avg_fallback 
WHERE id = 1")
+            checkAnswer(
+              safe,
+              Seq(Row(new java.math.BigDecimal("200.00"), new 
java.math.BigDecimal("200.00"))))
+            assert(collect(safe.queryExecution.executedPlan) { case agg: 
CometHashAggregateExec =>
+              agg
+            }.size == 1)
+
+            // The same unsafe buffer is valid when both aggregate stages 
execute in Comet.
+            withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") {
+              val native = sql("SELECT AVG(v) FROM decimal_avg_fallback WHERE 
id = 1")
+              checkAnswer(native, Seq(Row(new 
java.math.BigDecimal("200.000000"))))
+              assert(collect(native.queryExecution.executedPlan) {
+                case agg: CometHashAggregateExec => agg
+              }.size == 2)
+            }
+          }
+        }
+      }
+    }
+  }
+
+  for (adaptive <- Seq(false, true)) {
+    test(s"COUNT and AVG fall back together across a Spark shuffle 
(AQE=$adaptive)") {
+      withTempDir { dir =>
+        val path = s"${dir.getAbsolutePath}/data"
+        withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+          spark
+            .range(0L, 8L, 1L, 4)
+            .selectExpr("id", "CAST(1 AS BIGINT) AS v")
+            .write
+            .parquet(path)
+        }
+
+        withSQLConf(
+          SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString,
+          SQLConf.FILES_MAX_PARTITION_BYTES.key -> "1048576",
+          SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false",
+          CometConf.COMET_NATIVE_SCAN_ENABLED.key -> "false",
+          CometConf.COMET_CONVERT_FROM_PARQUET_ENABLED.key -> "true",
+          CometConf.COMET_SHUFFLE_ENABLED.key -> "false",
+          CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "true") {
+          withParquetTable(path, "count_avg_fallback") {
+            assert(sql("SELECT * FROM 
count_avg_fallback").rdd.getNumPartitions == 4)
+            // A safe COUNT buffer must not admit an unsafe AVG buffer in the 
same Partial.
+            // Three partitions have no surviving rows, so AVG has no 
update_batch call and
+            // its native state is (null, 0), which poisons Spark Final's sum. 
Keeping Final
+            // enabled exercises repair after the shuffle falls back during 
conversion.
+            val df = sql("SELECT COUNT(*), AVG(v) FROM count_avg_fallback 
WHERE id = 1")
+            val initialPlan = stripAQEPlan(df.queryExecution.executedPlan)
+            checkAnswer(df, Seq(Row(1L, 1.0)))
+            for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) {
+              assert(collect(plan) { case agg: CometHashAggregateExec => agg 
}.isEmpty)
+              val partials = collect(plan) {
+                case agg: BaseAggregateExec
+                    if agg.aggregateExpressions.map(_.mode).distinct == 
Seq(Partial) =>
+                  agg
+              }
+              assert(partials.size == 1)
+              
assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)
+              assert(collect(plan) { case filter: CometFilterExec => filter 
}.nonEmpty)
+            }
+            withSQLConf(
+              CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
+              CometConf.COMET_SHUFFLE_MODE.key -> "native") {
+              // The native Final can consume its own empty AVG buffers; only 
the engine split
+              // is unsafe. Keep the fully native aggregate path enabled.
+              val native =
+                sql("SELECT COUNT(*), AVG(v) FROM count_avg_fallback WHERE id 
= 1")
+              val initialNativePlan = 
stripAQEPlan(native.queryExecution.executedPlan)
+              checkAnswer(native, Seq(Row(1L, 1.0)))
+              for (plan <- Seq(initialNativePlan, 
native.queryExecution.executedPlan)) {
+                assert(collect(plan) { case agg: CometHashAggregateExec => agg 
}.size == 2)
+              }
+            }
+          }
+        }
+      }
+    }
+
+    test(s"COUNT preserves safe native partials across a Spark shuffle 
(AQE=$adaptive)") {
+      val data = Seq((0, None), (0, None), (1, Some(3)), (1, None), (1, 
Some(4)))
+      withParquetTable(data, "count_fallback", false) {
+        withSQLConf(
+          SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString,
+          SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false",
+          CometConf.COMET_SHUFFLE_ENABLED.key -> "false",
+          CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "true") {
+          for (query <- Seq(
+              "SELECT _1, COUNT(_2), COUNT(*) FROM count_fallback GROUP BY _1",
+              "SELECT COUNT(_2), COUNT(*) FROM count_fallback WHERE _1 < 0")) {
+            val df = sql(query)
+            val initialPlan = stripAQEPlan(df.queryExecution.executedPlan)
+            assert(collect(initialPlan) {
+              case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => 
agg
+            }.size == 1)
+            assert(collect(initialPlan) {
+              case agg: BaseAggregateExec
+                  if agg.aggregateExpressions.map(_.mode).distinct == 
Seq(Final) =>
+                agg
+            }.size == 1)
+            checkSparkAnswer(df)
+          }
+        }
+      }
+    }
+
+    for (fn <- Seq("collect_list", "collect_set")) {
+      test(s"$fn falls back when enabled native shuffle is ineligible 
(AQE=$adaptive)") {
+        val data = (0 until 30).map(i => (i % 3, if (i % 7 == 0) None else 
Some(i % 5)))
+        withParquetTable(data, "collect_fallback", false) {
+          withSQLConf(
+            SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString,
+            SQLConf.PARQUET_FILTER_PUSHDOWN_ENABLED.key -> "false",
+            SQLConf.USE_OBJECT_HASH_AGG.key -> "true",
+            CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
+            CometConf.COMET_SHUFFLE_MODE.key -> "native",
+            CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> 
"false") {
+            // Integer keys isolate this from the wide-decimal shuffle 
restriction in #5420.
+            // The native Partial emits an Array buffer, but Spark's Final 
expects Binary.
+            val query = s"SELECT _1, sort_array($fn(_2)), COUNT(*) " +
+              "FROM collect_fallback WHERE _1 >= 0 GROUP BY _1"
+            val df = sql(query)
+            val initialPlan = stripAQEPlan(df.queryExecution.executedPlan)
+            checkSparkAnswer(df)
+            for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) {
+              assert(collect(plan) { case agg: CometHashAggregateExec => agg 
}.isEmpty)
+              val partials = collect(plan) {
+                case agg: BaseAggregateExec
+                    if agg.aggregateExpressions.map(_.mode).distinct == 
Seq(Partial) =>
+                  agg
+              }
+              assert(partials.size == 1)
+              
assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)
+              assert(collect(plan) { case filter: CometFilterExec => filter 
}.nonEmpty)
+            }
+            // A fully native producer/consumer pair can still use its native 
buffer format.
+            
withSQLConf(CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> 
"true") {
+              val native = sql(query)
+              checkSparkAnswer(native)
+              assert(getNumCometHashAggregate(native) == 2)
+            }
+          }
+        }
+      }
+    }
+
+    for (fn <- Seq("percentile", "collect_list", "sum")) {
+      test(
+        s"$fn preserves aggregate buffers with an unsupported array hash key 
(AQE=$adaptive)") {
+        withSQLConf(
+          SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> adaptive.toString,
+          SQLConf.SHUFFLE_PARTITIONS.key -> "4",
+          CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
+          CometConf.COMET_SHUFFLE_MODE.key -> "native",
+          CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key -> 
"true") {
+          withTempView("array_key_aggregate") {
+            // The array key itself makes native shuffle ineligible; no 
feature is disabled.
+            // 
https://github.com/apache/datafusion-comet/issues/5419#issuecomment-5464233245
+            spark
+              .range(0, 18, 1, 4)
+              .selectExpr("id % 3 AS k", "id % 5 AS v")
+              .createOrReplaceTempView("array_key_aggregate")
+            val aggregate = if (fn == "percentile") "percentile(v, 0.5)" else 
s"$fn(v)"
+            val query = s"SELECT array(k) AS ak, $aggregate " +
+              "FROM array_key_aggregate GROUP BY array(k)"
+
+            def normalizedRows(df: DataFrame): Seq[Row] = {
+              df.collect()
+                .toSeq
+                .map { row =>
+                  // Keep the reported collect_list SQL unchanged, normalizing 
its order only
+                  // after execution so another expression cannot cause an 
earlier fallback.
+                  if (fn == "collect_list") {
+                    Row(row.getSeq[Long](0), row.getSeq[Long](1).sorted)
+                  } else {
+                    row
+                  }
+                }
+                .sortBy(_.getSeq[Long](0).head)
+            }
+
+            // Spark 3's withSQLConf returns Unit, so capture the baseline 
inside its body.
+            var expected: Seq[Row] = Seq.empty
+            withSQLConf(CometConf.COMET_ENABLED.key -> "false") {
+              expected = normalizedRows(sql(query))
+            }
+            val df = sql(query)
+            val initialPlan = stripAQEPlan(df.queryExecution.executedPlan)
+            // Execute this same DataFrame before inspecting its materialized 
AQE plan.
+            assert(normalizedRows(df) == expected)
+            for (plan <- Seq(initialPlan, df.queryExecution.executedPlan)) {
+              val exchanges = collect(plan) { case exchange: 
ShuffleExchangeExec => exchange }
+              assert(exchanges.size == 1, s"$plan")
+              assert(exchanges.head.outputPartitioning match {
+                case HashPartitioning(Seq(key), 4) => 
key.dataType.isInstanceOf[ArrayType]
+                case _ => false
+              })
+              assert(collect(plan) { case exchange: CometShuffleExchangeExec =>
+                exchange
+              }.isEmpty)
+              val partials = collect(plan) {
+                case agg: BaseAggregateExec
+                    if agg.aggregateExpressions.map(_.mode).distinct == 
Seq(Partial) =>
+                  agg
+              }
+              val finals = collect(plan) {
+                case agg: BaseAggregateExec
+                    if agg.aggregateExpressions.map(_.mode).distinct == 
Seq(Final) =>
+                  agg
+              }
+              assert(finals.size == 1, s"$plan")
+              val nativeAggregates = collect(plan) { case agg: 
CometHashAggregateExec => agg }
+              if (fn == "sum") {
+                // SUM's Long buffer is safe for Spark's final, so retain its 
native partial.
+                assert(nativeAggregates.size == 1, s"$plan")
+                assert(nativeAggregates.head.modes == Seq(Partial))
+                assert(partials.isEmpty, s"$plan")
+              } else {
+                assert(nativeAggregates.isEmpty, s"$plan")
+                assert(partials.size == 1, s"$plan")
+                
assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)
+                assert(collect(partials.head.child) { case project: 
CometProjectExec =>
+                  project
+                }.nonEmpty)
+              }
+            }
+            if (adaptive) {
+              val stages = collect(df.queryExecution.executedPlan) {
+                case stage: ShuffleQueryStageExec => stage
+              }
+              assert(stages.nonEmpty && stages.forall(_.isMaterialized))
+            }
+          }
+        }
+      }
+    }
+  }
+
   test("stddev_pop should return NaN for some cases") {
     withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "true") {
       Seq(true, false).foreach { nullOnDivideByZero =>
@@ -431,15 +717,56 @@ class CometAggregateSuite extends CometTestBase with 
AdaptiveSparkPlanHelper {
     }
   }
 
-  test("mixed engine sum/avg: Comet partial + Spark final matches Spark") {
+  test("decimal SUM partial stays in Spark when a later input cancels 
precision overflow") {
+    // Keep all three values in one ordered input partition. Generated scalar 
Spark SUM can
+    // retain the temporary 1.2 and return 0.6 after cancellation; native 
decimal SUM instead
+    // makes that precision overflow sticky, or throws immediately in ANSI 
mode. A matching
+    // (sum, isEmpty) buffer schema therefore does not establish forward 
interoperability.
+    withSQLConf(
+      SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+      SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "true",
+      "spark.sql.files.minPartitionNum" -> "1",
+      CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
+      CometConf.COMET_SHUFFLE_ENABLED.key -> "false") {
+      withTempPath { path =>
+        spark
+          .range(0, 3, 1, 1)
+          .selectExpr("CAST(CASE WHEN id < 2 THEN '0.6' ELSE '-0.6' END AS 
DECIMAL(38,38)) AS v")
+          .write
+          .parquet(path.getCanonicalPath)
+        withParquetTable(path.getCanonicalPath, "decimal_sum_cancellation") {
+          for (ansi <- Seq(false, true)) {
+            withSQLConf(SQLConf.ANSI_ENABLED.key -> ansi.toString) {
+              val df = sql("SELECT SUM(v) FROM decimal_sum_cancellation")
+              val plan = df.queryExecution.executedPlan
+              assert(collect(plan) { case agg: CometHashAggregateExec => agg 
}.isEmpty)
+              val partials = collect(plan) {
+                case agg: BaseAggregateExec
+                    if agg.aggregateExpressions.map(_.mode).distinct == 
Seq(Partial) =>
+                  agg
+              }
+              assert(partials.size == 1)
+              
assert(partials.head.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).isDefined)
+              assert(collect(plan) { case native: CometNativeExec => native 
}.nonEmpty)
+              checkSparkAnswer(df)
+              checkAnswer(df, Seq(Row(new java.math.BigDecimal("0.6"))))
+            }
+          }
+        }
+      }
+    }
+  }
+
+  test("mixed engine sum/avg falls back when Spark Final would consume native 
AVG") {
     val data = (0 until 100).map(i => (i, i.toLong, i.toDouble, i % 7))
     withParquetTable(data, "tbl") {
       withSQLConf(
         CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
         CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
         CometConf.COMET_SHUFFLE_MODE.key -> "jvm") {
-        checkSparkAnswer(
-          "SELECT _4, SUM(_1), SUM(_2), SUM(_3), AVG(_1), AVG(_2), AVG(_3) 
FROM tbl GROUP BY _4")
+        checkSparkAnswerAndNumOfAggregates(
+          "SELECT _4, SUM(_1), SUM(_2), SUM(_3), AVG(_1), AVG(_2), AVG(_3) 
FROM tbl GROUP BY _4",
+          0)
       }
     }
   }
@@ -709,7 +1036,10 @@ class CometAggregateSuite extends CometTestBase with 
AdaptiveSparkPlanHelper {
             checkSparkAnswerAndNumOfAggregates("SELECT _2, COUNT(_1) FROM tbl 
GROUP BY _2", n)
             checkSparkAnswerAndNumOfAggregates("SELECT _2, MIN(_1) FROM tbl 
GROUP BY _2", n)
             checkSparkAnswerAndNumOfAggregates("SELECT _2, MAX(_1) FROM tbl 
GROUP BY _2", n)
-            checkSparkAnswerAndNumOfAggregates("SELECT _2, AVG(_1) FROM tbl 
GROUP BY _2", n)
+            val avgStages = if (nativeShuffleEnabled) 2 else 0
+            checkSparkAnswerAndNumOfAggregates(
+              "SELECT _2, AVG(_1) FROM tbl GROUP BY _2",
+              avgStages)
           }
         }
       }
@@ -914,26 +1244,29 @@ class CometAggregateSuite extends CometTestBase with 
AdaptiveSparkPlanHelper {
             val path = new Path(dir.toURI.toString, "test")
             makeParquetFile(path, 1000, 20, dictionaryEnabled)
             withParquetTable(path.toUri.toString, "tbl") {
+              // Spark rewrites _7's small decimal SUM to Long; _8 and _9 
remain decimal and
+              // cannot use a native Partial when the Final runs in Spark.
               val expectedNumOfCometAggregates = if (nativeShuffleEnabled) 2 
else 1
+              val expectedNumOfDecimalAggregates = if (nativeShuffleEnabled) 2 
else 0
 
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT _g2, SUM(_7) FROM tbl GROUP BY _g2",
                 expectedNumOfCometAggregates)
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT _g3, SUM(_8) FROM tbl GROUP BY _g3",
-                expectedNumOfCometAggregates)
+                expectedNumOfDecimalAggregates)
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT _g4, SUM(_9) FROM tbl GROUP BY _g4",
-                expectedNumOfCometAggregates)
+                expectedNumOfDecimalAggregates)
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT SUM(_7) FROM tbl",
                 expectedNumOfCometAggregates)
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT SUM(_8) FROM tbl",
-                expectedNumOfCometAggregates)
+                expectedNumOfDecimalAggregates)
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT SUM(_9) FROM tbl",
-                expectedNumOfCometAggregates)
+                expectedNumOfDecimalAggregates)
             }
           }
         }
@@ -1551,14 +1884,14 @@ class CometAggregateSuite extends CometTestBase with 
AdaptiveSparkPlanHelper {
     }
   }
 
-  test("test partial avg") {
+  test("AVG stays in Spark across a Spark shuffle") {
     Seq(true, false).foreach { dictionaryEnabled =>
       withParquetTable(
         (0 until 5).map(i => (i.toDouble, i.toDouble % 2)),
         "tbl",
         dictionaryEnabled) {
         withSQLConf(CometConf.COMET_SHUFFLE_ENABLED.key -> "false") {
-          checkSparkAnswerAndNumOfAggregates("SELECT _2 , AVG(_1) FROM tbl 
GROUP BY _2", 1)
+          checkSparkAnswerAndNumOfAggregates("SELECT _2 , AVG(_1) FROM tbl 
GROUP BY _2", 0)
         }
       }
     }
@@ -1595,7 +1928,9 @@ class CometAggregateSuite extends CometTestBase with 
AdaptiveSparkPlanHelper {
             val path = new Path(dir.toURI.toString, "test")
             makeParquetFile(path, 1000, 20, dictionaryEnabled)
             withParquetTable(path.toUri.toString, "tbl") {
-              val expectedNumOfCometAggregates = if (nativeShuffleEnabled) 2 
else 1
+              // Spark rewrites _7 to Long AVG, whose empty native buffer is 
also unsafe for a
+              // Spark Final until #5420. Keep all AVG partials in Spark 
across this boundary.
+              val expectedNumOfCometAggregates = if (nativeShuffleEnabled) 2 
else 0
 
               checkSparkAnswerAndNumOfAggregates(
                 "SELECT _g2, AVG(_7) FROM tbl GROUP BY _g2",
diff --git 
a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala 
b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
index 79b668444d..9d65dec834 100644
--- a/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
+++ b/spark/src/test/scala/org/apache/comet/rules/CometExecRuleSuite.scala
@@ -21,22 +21,23 @@ package org.apache.comet.rules
 
 import scala.util.Random
 
+import org.apache.logging.log4j.Level
 import org.apache.spark.sql._
 import org.apache.spark.sql.catalyst.FunctionIdentifier
-import org.apache.spark.sql.catalyst.expressions.{Expression, ExpressionInfo}
-import 
org.apache.spark.sql.catalyst.expressions.aggregate.{BloomFilterAggregate, 
Partial}
+import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, 
ExpressionInfo, Literal}
+import 
org.apache.spark.sql.catalyst.expressions.aggregate.{AggregateExpression, 
BloomFilterAggregate, Final, Min, Partial, PartialMerge}
 import org.apache.spark.sql.comet._
 import org.apache.spark.sql.comet.execution.shuffle.CometShuffleExchangeExec
 import org.apache.spark.sql.execution._
-import org.apache.spark.sql.execution.adaptive.QueryStageExec
+import org.apache.spark.sql.execution.adaptive.{QueryStageExec, 
ShuffleQueryStageExec}
 import org.apache.spark.sql.execution.aggregate.{HashAggregateExec, 
ObjectHashAggregateExec}
 import org.apache.spark.sql.execution.exchange.{BroadcastExchangeExec, 
ShuffleExchangeExec}
 import org.apache.spark.sql.internal.SQLConf
 import org.apache.spark.sql.types.{DataTypes, StructField, StructType}
 
 import org.apache.comet.{CometConf, CometExplainInfo, ExtendedExplainInfo}
-import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, 
isSpark40Plus, isSpark42Plus}
-import org.apache.comet.serde.{Compatible, Unsupported}
+import org.apache.comet.CometSparkSessionExtensions.{isSpark35Plus, 
isSpark40Plus, isSpark42Plus, withFallbackReason}
+import org.apache.comet.serde.{CometAggregateExpressionSerde, Compatible, 
ExprOuterClass, Unsupported}
 import org.apache.comet.testing.{DataGenOptions, FuzzDataGenerator}
 
 /**
@@ -353,8 +354,7 @@ class CometExecRuleSuite extends CometTestBase {
     }
   }
 
-  // Regression test for https://github.com/apache/datafusion-comet/issues/1389
-  test("CometExecRule should not allow Comet partial and Spark final hash 
aggregate") {
+  test("CometExecRule should allow COUNT Comet partial and Spark final hash 
aggregate") {
     withTempView("test_data") {
       createTestDataFrame.createOrReplaceTempView("test_data")
 
@@ -370,11 +370,10 @@ class CometExecRuleSuite extends CometTestBase {
         CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
         val transformedPlan = applyCometExecRule(sparkPlan)
 
-        // COUNT is intentionally excluded from mixed execution (AQE / 
count-bug reasons), so if
-        // the final aggregate cannot be converted to Comet, neither should 
the partial.
-        assert(
-          countOperators(transformedPlan, classOf[HashAggregateExec]) == 
originalHashAggCount)
-        assert(countOperators(transformedPlan, 
classOf[CometHashAggregateExec]) == 0)
+        // COUNT's buffer is compatible in this direction. Keeping the Final 
in Spark also keeps
+        // the AQE/count-bug rewrites that prevent the reverse direction from 
being admitted.
+        assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 
1)
+        assert(countOperators(transformedPlan, 
classOf[CometHashAggregateExec]) == 1)
       }
     }
   }
@@ -395,8 +394,8 @@ class CometExecRuleSuite extends CometTestBase {
         CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
         val transformedPlan = applyCometExecRule(sparkPlan)
 
-        // COUNT blocks mixed execution, so if the partial cannot be 
converted, neither should
-        // the final.
+        // COUNT still blocks Spark Partial to Comet Final, independently of 
the safe reverse
+        // direction, so if the partial cannot be converted, neither should 
the final.
         assert(
           countOperators(transformedPlan, classOf[HashAggregateExec]) == 
originalHashAggCount)
         assert(countOperators(transformedPlan, 
classOf[CometHashAggregateExec]) == 0)
@@ -483,7 +482,7 @@ class CometExecRuleSuite extends CometTestBase {
     }
   }
 
-  test("CometExecRule should allow AVG mixed Comet partial and Spark final") {
+  test("CometExecRule should not allow AVG Comet partial and Spark final 
before buffer repair") {
     withTempView("test_data") {
       createTestDataFrame.createOrReplaceTempView("test_data")
       val sparkPlan =
@@ -493,8 +492,9 @@ class CometExecRuleSuite extends CometTestBase {
         CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
         CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
         val transformedPlan = applyCometExecRule(sparkPlan)
-        assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 
1) // final
-        assert(countOperators(transformedPlan, 
classOf[CometHashAggregateExec]) == 1) // partial
+        // Matching field types do not make native AVG's empty (null, 0) state 
safe for Spark.
+        assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 
2)
+        assert(countOperators(transformedPlan, 
classOf[CometHashAggregateExec]) == 0)
       }
     }
   }
@@ -541,6 +541,175 @@ class CometExecRuleSuite extends CometTestBase {
     }
   }
 
+  for (distinct <- Seq(false, true)) {
+    test(
+      s"unsafe aggregate buffers fall back when native shuffle is ineligible 
(distinct=$distinct)") {
+      withTempView("test_data") {
+        createTestDataFrame.createOrReplaceTempView("test_data")
+        val aggregates = "AVG(id)" + (if (distinct) ", SUM(DISTINCT id)" else 
"")
+
+        for (fallback <- Seq("disabled hash partitioning", "prior shuffle 
fallback", "none")) {
+          withSQLConf(
+            CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true",
+            CometConf.COMET_SHUFFLE_ENABLED.key -> "true",
+            CometConf.COMET_SHUFFLE_MODE.key -> "native",
+            CometConf.COMET_SHUFFLE_NATIVE_HASH_PARTITIONING_ENABLED.key ->
+              (fallback != "disabled hash partitioning").toString) {
+            val sparkPlan =
+              createSparkPlan(spark, s"SELECT $aggregates FROM test_data GROUP 
BY (id % 3)")
+            val aggregateCount = countOperators(sparkPlan, 
classOf[HashAggregateExec])
+            assert(aggregateCount == (if (distinct) 4 else 2))
+            if (fallback == "prior shuffle fallback") {
+              // Tag only the lowest exchange. A DISTINCT plan's upper 
exchange must inherit
+              // the native-only refusal from its now-Spark merge inputs, not 
from another tag.
+              val lowerShuffle = stripAQEPlan(sparkPlan).collect {
+                case shuffle: ShuffleExchangeExec => shuffle
+              }.last
+              withFallbackReason(lowerShuffle, fallback)
+            }
+            val transformed = applyCometExecRule(sparkPlan)
+            val nativeExpected = fallback == "none"
+
+            // Shuffle is enabled, but a native-only shuffle can still fall 
back. The distinct
+            // rewrite also has intermediate PartialMerge and mixed 
Partial/PartialMerge stages.
+            for (plan <- Seq(transformed, applyCometExecRule(transformed))) {
+              assert(
+                countOperators(plan, classOf[CometHashAggregateExec]) ==
+                  (if (nativeExpected) aggregateCount else 0))
+              assert(
+                countOperators(plan, classOf[HashAggregateExec]) ==
+                  (if (nativeExpected) 0 else aggregateCount))
+            }
+            // AQE reapplies the rule to an exchange without its Final 
aggregate. The tagged
+            // Partial must remain in Spark in that stage-only pass too.
+            transformed.collect { case shuffle: ShuffleExchangeExec => shuffle 
}.foreach {
+              shuffle =>
+                val stage = applyCometExecRule(shuffle)
+                assert(countOperators(stage, classOf[CometHashAggregateExec]) 
== 0)
+            }
+          }
+        }
+      }
+    }
+  }
+
+  test("aggregate buffer direction opt-ins are independent") {
+    // A policy-only handler opts into consuming Spark state. Its inherited 
producer policy
+    // must stay false; serializing any expression is outside the scope of 
this fixture.
+    val reverseOnly = new CometAggregateExpressionSerde[Min] {
+      override def supportsSparkPartialToNativeFinal(fn: Min): Boolean = true
+
+      override def convert(
+          aggExpr: AggregateExpression,
+          expr: Min,
+          inputs: Seq[Attribute],
+          binding: Boolean,
+          conf: SQLConf): Option[ExprOuterClass.AggExpr] = None
+    }
+    val fn = Min(Literal(1L))
+    assert(reverseOnly.supportsSparkPartialToNativeFinal(fn))
+    assert(!reverseOnly.supportsNativePartialToSparkFinal(fn))
+  }
+
+  test("restored partial records its reason when its current child is not 
native") {
+    // Wrap a converted input to prevent a re-entrant serde call from 
supplying the reason.
+    // This planner-only fixture never executes its synthetic buffer boundary.
+    withSQLConf(
+      SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+      SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false",
+      CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true",
+      CometConf.COMET_SHUFFLE_MODE.key -> "native") {
+      withTempView("test_data") {
+        createTestDataFrame.createOrReplaceTempView("test_data")
+        val plan = applyCometExecRule(
+          createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id % 
3)"))
+        val partial = plan.collectFirst {
+          case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => agg
+        }.get
+        val sparkFinal = plan.collectFirst {
+          case agg: CometHashAggregateExec if agg.modes == Seq(Final) =>
+            agg.originalPlan.asInstanceOf[HashAggregateExec]
+        }.get
+        val nonNativeChild = InputAdapter(partial.child)
+        val restored = CometExecRule(spark).revertUnsafePartialAggregates(
+          sparkFinal.copy(child = partial.copy(child = nonNativeChild)))
+        val sparkPartial = restored.children.head
+        assert(sparkPartial.isInstanceOf[HashAggregateExec])
+        assert(sparkPartial.children.head.isInstanceOf[InputAdapter])
+        assert(sparkPartial.children.head.output == nonNativeChild.output)
+        val reason = 
sparkPartial.getTagValue(CometExecRule.COMET_UNSAFE_PARTIAL).get
+        
assert(sparkPartial.getTagValue(CometExplainInfo.FALLBACK_REASONS).get.contains(reason))
+        assert(new 
ExtendedExplainInfo().getFallbackReasons(sparkPartial).contains(reason))
+      }
+    }
+  }
+
+  test("unrepaired native aggregate buffers warn once without rewriting query 
stages") {
+    // Construct the stage placeholder emitted by CometExchangeSink, including 
a native merge
+    // above it. No SQL reproduction or materialization is assumed: this pins 
the diagnostic
+    // when repair stops at a stage, and the absence of warnings for unrelated 
inner producers.
+    withSQLConf(
+      SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false",
+      SQLConf.WHOLESTAGE_CODEGEN_ENABLED.key -> "false",
+      CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true",
+      CometConf.COMET_SHUFFLE_MODE.key -> "native") {
+      withTempView("test_data") {
+        createTestDataFrame.createOrReplaceTempView("test_data")
+        val plan = applyCometExecRule(
+          createSparkPlan(spark, "SELECT AVG(id) FROM test_data GROUP BY (id % 
3)"))
+        val partial = plan.collectFirst {
+          case agg: CometHashAggregateExec if agg.modes == Seq(Partial) => agg
+        }.get
+        val nativeFinal = plan.collectFirst {
+          case agg: CometHashAggregateExec if agg.modes == Seq(Final) => agg
+        }.get
+        val sparkFinal = 
nativeFinal.originalPlan.asInstanceOf[HashAggregateExec]
+        val sparkPartial = partial.originalPlan.asInstanceOf[HashAggregateExec]
+        val exchange = ShuffleExchangeExec(
+          org.apache.spark.sql.catalyst.plans.physical.SinglePartition,
+          partial)
+        val stage = ShuffleQueryStageExec(0, exchange, exchange.canonicalized)
+        val placeholder = CometSinkPlaceHolder(
+          
org.apache.comet.serde.OperatorOuterClass.Operator.getDefaultInstance,
+          stage,
+          stage)
+        val nativeMerge = partial.copy(
+          aggregateExpressions = partial.aggregateExpressions.map(_.copy(mode 
= PartialMerge)),
+          child = placeholder)
+        val warning = "Comet could not restore a native intermediate buffer 
producer"
+        val rule = CometExecRule(spark)
+        for {
+          (child, shouldWarn) <- Seq(
+            nativeMerge -> true,
+            placeholder -> true,
+            sparkPartial.copy(child = nativeFinal) -> false,
+            nativeFinal -> false)
+          logFallback <- Seq("false", "true")
+        } {
+          withSQLConf(CometConf.COMET_EXPLAIN_FALLBACK_LOG_ENABLED.key -> 
logFallback) {
+            val consumer = sparkFinal.copy(child = child)
+            val appender = new LogAppender("unrepaired aggregate buffers")
+            withLogAppender(appender, Seq("org.apache.comet"), 
Some(Level.WARN)) {
+              assert(rule.revertUnsafePartialAggregates(consumer) eq consumer)
+              assert(rule.revertUnsafePartialAggregates(consumer) eq consumer)
+            }
+            assert(consumer.child eq child)
+            val warnings =
+              
appender.loggingEvents.count(_.getMessage.getFormattedMessage.contains(warning))
+            assert(warnings == (if (shouldWarn) 1 else 0), s"$child: 
$warnings")
+            assert(
+              new ExtendedExplainInfo()
+                .getFallbackReasons(consumer)
+                .exists(_.contains(warning)) ==
+                shouldWarn)
+          }
+        }
+        assert(stage.plan eq exchange)
+        assert(exchange.child eq partial)
+      }
+    }
+  }
+
   test("CometExecRule should not allow decimal SUM mixed execution") {
     withTempView("test_data") {
       createTestDataFrame.createOrReplaceTempView("test_data")
@@ -556,9 +725,9 @@ class CometExecRuleSuite extends CometTestBase {
         CometConf.COMET_ENABLE_FINAL_HASH_AGGREGATE.key -> "false",
         CometConf.COMET_EXEC_LOCAL_TABLE_SCAN_ENABLED.key -> "true") {
         val transformedPlan = applyCometExecRule(sparkPlan)
-        // Decimal SUM overflow detection (ANSI throw / Legacy null) does not 
survive a
-        // Spark-partial / Comet-final split, so mixed execution is unsafe and 
the partial
-        // must also fall back to Spark.
+        // Native decimal SUM makes precision overflow sticky (or throws 
eagerly in ANSI),
+        // while Spark's generated scalar Partial can recover after a later 
cancelling input.
+        // Keep the Partial in Spark even though the emitted buffer field 
types match.
         assert(countOperators(transformedPlan, classOf[HashAggregateExec]) == 
2)
         assert(countOperators(transformedPlan, 
classOf[CometHashAggregateExec]) == 0)
       }
diff --git 
a/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala
 
b/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala
index a27f7c7b14..b2c6c0f8a3 100644
--- 
a/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala
+++ 
b/spark/src/test/scala/org/apache/spark/sql/comet/execution/shuffle/CometCelebornShufflePlanningSuite.scala
@@ -548,7 +548,7 @@ class CometCelebornShufflePlanningSuite extends 
CometTestBase {
 
     for {
       fallback <- Seq("partition threshold", "unsupported array hash key")
-      function <- Seq("collect_list", "collect_set", "avg")
+      function <- Seq("collect_list", "collect_set", "avg", "count")
     } {
       test(s"native $fallback preserves $function aggregate buffers with 
AQE=$adaptive") {
         val complexKey = fallback == "unsupported array hash key"
@@ -563,8 +563,11 @@ class CometCelebornShufflePlanningSuite extends 
CometTestBase {
             SQLConf.SHUFFLE_PARTITIONS.key -> "4",
             CometConf.COMET_SHUFFLE_MODE.key -> "native") {
             val grouping = if (complexKey) "array(id % 3)" else "id % 3"
-            val aggregate =
-              if (function == "avg") "avg(value)" else 
s"sort_array($function(value))"
+            val aggregate = if (function.startsWith("collect_")) {
+              s"sort_array($function(value))"
+            } else {
+              s"$function(value)"
+            }
             val query = spark
               .range(0, 18, 1, 4)
               .selectExpr(s"$grouping AS grouping_key", "id AS value")
@@ -579,13 +582,15 @@ class CometCelebornShufflePlanningSuite extends 
CometTestBase {
             val nativeAggregates = collect(executedPlan) {
               case aggregate: CometHashAggregateExec => aggregate
             }
-            if (function == "avg") {
-              // AVG's intermediate state is Spark-compatible; native partials 
remain safe.
-              assert(nativeAggregates.nonEmpty, s"$executedPlan")
+            if (function == "count") {
+              // COUNT's non-null Long buffer is safe for Spark Final to 
consume.
+              assert(nativeAggregates.size == 1, s"$executedPlan")
+              assert(nativeAggregates.head.modes == Seq(Partial), 
s"$executedPlan")
             } else {
               // A Spark final cannot deserialize Comet's ArrayType 
collect_list/collect_set
-              // state as its BinaryType buffer. Both halves must agree when 
the exchange falls
-              // back, not just when an aggregate operator itself is 
unsupported.
+              // state as BinaryType, or safely merge AVG's never-updated 
(null, 0) buffer.
+              // Both halves must agree when the exchange falls back, not just 
when an
+              // aggregate operator itself is unsupported.
               assert(nativeAggregates.isEmpty, s"$executedPlan")
             }
           }


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

Reply via email to