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

zhouyuan pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/gluten.git


The following commit(s) were added to refs/heads/main by this push:
     new 5eca564841 [GLUTEN-12597][CORE] Migrate AggregateRel.Grouping to 
expression references (Substrait 0.98) (#12724)
5eca564841 is described below

commit 5eca56484134b91d4b74ae745a41bc2654504564
Author: Niels Pardon <[email protected]>
AuthorDate: Wed Aug 19 15:27:09 2026 +0200

    [GLUTEN-12597][CORE] Migrate AggregateRel.Grouping to expression references 
(Substrait 0.98) (#12724)
---
 .../covar_samp-covar_pop-final-agg-stage.json      | 25 +++----
 .../covar_samp-covar_pop-partial-agg-stage.json    | 25 +++----
 .../substrait-plans/tpch-q1-final-agg-stage.json   | 25 +++----
 .../substrait-plans/tpch-q2-in-one-wholestage.json | 26 +++----
 .../substrait-plans/tpch-q4-shj-stage.json         | 13 ++--
 .../Parser/RelParsers/AggregateRelParser.cpp       |  6 +-
 .../Parser/RelParsers/ExpandRelParser.cpp          |  2 +-
 .../tests/json/native_write_plan_1_spark33.json    | 19 ++---
 .../data/generic_q1/q1_first_stage_0.json          | 36 +++++-----
 cpp/velox/benchmarks/data/plan/q17_joins.json      | 19 ++---
 cpp/velox/substrait/SubstraitToVeloxPlan.cc        |  5 +-
 .../substrait/SubstraitToVeloxPlanValidator.cc     |  9 +--
 cpp/velox/substrait/VeloxToSubstraitPlan.cc        |  5 +-
 cpp/velox/tests/data/q1_first_stage.json           | 42 ++++++-----
 .../gluten/substrait/rel/AggregateRelNode.java     | 15 ++--
 .../substrait/proto/substrait/algebra.proto        | 17 ++++-
 .../gluten/utils/AggregateRelProtoSuite.scala      | 83 ++++++++++++++++++++++
 17 files changed, 251 insertions(+), 121 deletions(-)

diff --git 
a/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-final-agg-stage.json
 
b/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-final-agg-stage.json
index 3d8c5d5263..2c87e6a513 100644
--- 
a/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-final-agg-stage.json
+++ 
b/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-final-agg-stage.json
@@ -130,22 +130,23 @@
                 }
               },
               "groupings": [{
-                "groupingExpressions": [{
-                  "selection": {
-                    "directReference": {
-                      "structField": {
-                      }
+                "expressionReferences": [0, 1]
+              }],
+              "groupingExpressions": [{
+                "selection": {
+                  "directReference": {
+                    "structField": {
                     }
                   }
-                }, {
-                  "selection": {
-                    "directReference": {
-                      "structField": {
-                        "field": 1
-                      }
+                }
+              }, {
+                "selection": {
+                  "directReference": {
+                    "structField": {
+                      "field": 1
                     }
                   }
-                }]
+                }
               }],
               "measures": [{
                 "measure": {
diff --git 
a/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-partial-agg-stage.json
 
b/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-partial-agg-stage.json
index 98a7906974..f2f8b2dc59 100644
--- 
a/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-partial-agg-stage.json
+++ 
b/backends-clickhouse/src/test/resources/substrait-plans/covar_samp-covar_pop-partial-agg-stage.json
@@ -355,22 +355,23 @@
             }
           },
           "groupings": [{
-            "groupingExpressions": [{
-              "selection": {
-                "directReference": {
-                  "structField": {
-                  }
+            "expressionReferences": [0, 1]
+          }],
+          "groupingExpressions": [{
+            "selection": {
+              "directReference": {
+                "structField": {
                 }
               }
-            }, {
-              "selection": {
-                "directReference": {
-                  "structField": {
-                    "field": 1
-                  }
+            }
+          }, {
+            "selection": {
+              "directReference": {
+                "structField": {
+                  "field": 1
                 }
               }
-            }]
+            }
           }],
           "measures": [{
             "measure": {
diff --git 
a/backends-clickhouse/src/test/resources/substrait-plans/tpch-q1-final-agg-stage.json
 
b/backends-clickhouse/src/test/resources/substrait-plans/tpch-q1-final-agg-stage.json
index 335a84132c..cdfb2b7f2c 100644
--- 
a/backends-clickhouse/src/test/resources/substrait-plans/tpch-q1-final-agg-stage.json
+++ 
b/backends-clickhouse/src/test/resources/substrait-plans/tpch-q1-final-agg-stage.json
@@ -139,22 +139,23 @@
                 }
               },
               "groupings": [{
-                "groupingExpressions": [{
-                  "selection": {
-                    "directReference": {
-                      "structField": {
-                      }
+                "expressionReferences": [0, 1]
+              }],
+              "groupingExpressions": [{
+                "selection": {
+                  "directReference": {
+                    "structField": {
                     }
                   }
-                }, {
-                  "selection": {
-                    "directReference": {
-                      "structField": {
-                        "field": 1
-                      }
+                }
+              }, {
+                "selection": {
+                  "directReference": {
+                    "structField": {
+                      "field": 1
                     }
                   }
-                }]
+                }
               }],
               "measures": [{
                 "measure": {
diff --git 
a/backends-clickhouse/src/test/resources/substrait-plans/tpch-q2-in-one-wholestage.json
 
b/backends-clickhouse/src/test/resources/substrait-plans/tpch-q2-in-one-wholestage.json
index b910b1df0e..0843748475 100644
--- 
a/backends-clickhouse/src/test/resources/substrait-plans/tpch-q2-in-one-wholestage.json
+++ 
b/backends-clickhouse/src/test/resources/substrait-plans/tpch-q2-in-one-wholestage.json
@@ -750,14 +750,15 @@
                                                                                
     }
                                                                                
   },
                                                                                
   "groupings": [{
-                                                                               
     "groupingExpressions": [{
-                                                                               
       "selection": {
-                                                                               
         "directReference": {
-                                                                               
           "structField": {
-                                                                               
           }
+                                                                               
     "expressionReferences": [0]
+                                                                               
   }],
+                                                                               
   "groupingExpressions": [{
+                                                                               
     "selection": {
+                                                                               
       "directReference": {
+                                                                               
         "structField": {
                                                                                
         }
                                                                                
       }
-                                                                               
     }]
+                                                                               
     }
                                                                                
   }],
                                                                                
   "measures": [{
                                                                                
     "measure": {
@@ -784,14 +785,15 @@
                                                                                
 }
                                                                               
},
                                                                               
"groupings": [{
-                                                                               
 "groupingExpressions": [{
-                                                                               
   "selection": {
-                                                                               
     "directReference": {
-                                                                               
       "structField": {
-                                                                               
       }
+                                                                               
 "expressionReferences": [0]
+                                                                              
}],
+                                                                              
"groupingExpressions": [{
+                                                                               
 "selection": {
+                                                                               
   "directReference": {
+                                                                               
     "structField": {
                                                                                
     }
                                                                                
   }
-                                                                               
 }]
+                                                                               
 }
                                                                               
}],
                                                                               
"measures": [{
                                                                                
 "measure": {
diff --git 
a/backends-clickhouse/src/test/resources/substrait-plans/tpch-q4-shj-stage.json 
b/backends-clickhouse/src/test/resources/substrait-plans/tpch-q4-shj-stage.json
index 768856fd95..09fcce74ad 100644
--- 
a/backends-clickhouse/src/test/resources/substrait-plans/tpch-q4-shj-stage.json
+++ 
b/backends-clickhouse/src/test/resources/substrait-plans/tpch-q4-shj-stage.json
@@ -170,14 +170,15 @@
             }
           },
           "groupings": [{
-            "groupingExpressions": [{
-              "selection": {
-                "directReference": {
-                  "structField": {
-                  }
+            "expressionReferences": [0]
+          }],
+          "groupingExpressions": [{
+            "selection": {
+              "directReference": {
+                "structField": {
                 }
               }
-            }]
+            }
           }],
           "measures": [{
             "measure": {
diff --git a/cpp-ch/local-engine/Parser/RelParsers/AggregateRelParser.cpp 
b/cpp-ch/local-engine/Parser/RelParsers/AggregateRelParser.cpp
index 6880dc8d2c..5ddee9bfb5 100644
--- a/cpp-ch/local-engine/Parser/RelParsers/AggregateRelParser.cpp
+++ b/cpp-ch/local-engine/Parser/RelParsers/AggregateRelParser.cpp
@@ -100,7 +100,7 @@ AggregateRelParser::parse(DB::QueryPlanPtr query_plan, 
const substrait::Rel & re
     }
 
     /// If the groupings is empty, we still need to return one row with 
default values even if the input is empty.
-    if ((rel.aggregate().groupings().empty() || 
rel.aggregate().groupings()[0].grouping_expressions().empty())
+    if ((rel.aggregate().groupings().empty() || 
rel.aggregate().groupings()[0].expression_references().empty())
         && (has_final_stage || has_complete_stage || 
rel.aggregate().measures().empty()))
     {
         LOG_TRACE(&Poco::Logger::get("AggregateRelParser"), "default aggregate 
result step");
@@ -185,8 +185,10 @@ void AggregateRelParser::setup(DB::QueryPlanPtr 
query_plan, const substrait::Rel
 
     if (aggregate_rel->groupings_size() == 1)
     {
-        for (const auto & expr : 
aggregate_rel->groupings(0).grouping_expressions())
+        /// Grouping expressions live in the rel-level pool; the grouping 
references them by index.
+        for (const auto & ref : 
aggregate_rel->groupings(0).expression_references())
         {
+            const auto & expr = aggregate_rel->grouping_expressions(ref);
             auto field_index = SubstraitParserUtils::getStructFieldIndex(expr);
             if (field_index)
                 
grouping_keys.push_back(input_header.getByPosition(*field_index).name);
diff --git a/cpp-ch/local-engine/Parser/RelParsers/ExpandRelParser.cpp 
b/cpp-ch/local-engine/Parser/RelParsers/ExpandRelParser.cpp
index 495831ec1b..4bbdba6f90 100644
--- a/cpp-ch/local-engine/Parser/RelParsers/ExpandRelParser.cpp
+++ b/cpp-ch/local-engine/Parser/RelParsers/ExpandRelParser.cpp
@@ -168,7 +168,7 @@ DB::QueryPlanPtr ExpandRelParser::lazyAggregateExpandParse(
     auto aggregate_rel = rel.expand().input().aggregate();
     auto aggregate_descriptions = buildAggregations(*input_header, 
expand_field, aggregate_rel);
 
-    size_t grouping_keys = 
aggregate_rel.groupings(0).grouping_expressions_size();
+    size_t grouping_keys = 
aggregate_rel.groupings(0).expression_references_size();
 
     auto expand_step
         = std::make_unique<AdvancedExpandStep>(getContext(), input_header, 
grouping_keys, aggregate_descriptions, expand_field);
diff --git a/cpp-ch/local-engine/tests/json/native_write_plan_1_spark33.json 
b/cpp-ch/local-engine/tests/json/native_write_plan_1_spark33.json
index e053c99352..ad16d4c142 100644
--- a/cpp-ch/local-engine/tests/json/native_write_plan_1_spark33.json
+++ b/cpp-ch/local-engine/tests/json/native_write_plan_1_spark33.json
@@ -47,17 +47,20 @@
             },
             "groupings": [
               {
-                "groupingExpressions": [
-                  {
-                    "selection": {
-                      "directReference": {
-                        "structField": {}
-                      }
-                    }
-                  }
+                "expressionReferences": [
+                  0
                 ]
               }
             ],
+            "groupingExpressions": [
+              {
+                "selection": {
+                  "directReference": {
+                    "structField": {}
+                  }
+                }
+              }
+            ],
             "measures": [
               {
                 "measure": {
diff --git a/cpp/velox/benchmarks/data/generic_q1/q1_first_stage_0.json 
b/cpp/velox/benchmarks/data/generic_q1/q1_first_stage_0.json
index 5788baaf35..0fb4f4c2ed 100644
--- a/cpp/velox/benchmarks/data/generic_q1/q1_first_stage_0.json
+++ b/cpp/velox/benchmarks/data/generic_q1/q1_first_stage_0.json
@@ -190,24 +190,28 @@
                                 },
                                 "groupings": [
                                     {
-                                        "groupingExpressions": [
-                                            {
-                                                "selection": {
-                                                    "directReference": {
-                                                        "structField": {}
-                                                    }
-                                                }
-                                            },
-                                            {
-                                                "selection": {
-                                                    "directReference": {
-                                                        "structField": {
-                                                            "field": 1
-                                                        }
-                                                    }
+                                        "expressionReferences": [
+                                            0,
+                                            1
+                                        ]
+                                    }
+                                ],
+                                "groupingExpressions": [
+                                    {
+                                        "selection": {
+                                            "directReference": {
+                                                "structField": {}
+                                            }
+                                        }
+                                    },
+                                    {
+                                        "selection": {
+                                            "directReference": {
+                                                "structField": {
+                                                    "field": 1
                                                 }
                                             }
-                                        ]
+                                        }
                                     }
                                 ]
                             }
diff --git a/cpp/velox/benchmarks/data/plan/q17_joins.json 
b/cpp/velox/benchmarks/data/plan/q17_joins.json
index efb1c50239..b7cc137113 100644
--- a/cpp/velox/benchmarks/data/plan/q17_joins.json
+++ b/cpp/velox/benchmarks/data/plan/q17_joins.json
@@ -310,17 +310,20 @@
                                     },
                                     "groupings": [
                                       {
-                                        "groupingExpressions": [
-                                          {
-                                            "selection": {
-                                              "directReference": {
-                                                "structField": {}
-                                              }
-                                            }
-                                          }
+                                        "expressionReferences": [
+                                          0
                                         ]
                                       }
                                     ],
+                                    "groupingExpressions": [
+                                      {
+                                        "selection": {
+                                          "directReference": {
+                                            "structField": {}
+                                          }
+                                        }
+                                      }
+                                    ],
                                     "measures": [
                                       {
                                         "measure": {
diff --git a/cpp/velox/substrait/SubstraitToVeloxPlan.cc 
b/cpp/velox/substrait/SubstraitToVeloxPlan.cc
index 37a0a6ec40..888ab8d6ce 100644
--- a/cpp/velox/substrait/SubstraitToVeloxPlan.cc
+++ b/cpp/velox/substrait/SubstraitToVeloxPlan.cc
@@ -574,7 +574,10 @@ core::PlanNodePtr 
SubstraitToVeloxPlanConverter::toVeloxPlan(const ::substrait::
   VELOX_CHECK(
       aggRel.groupings().size() <= 1, "At most one grouping is supported, but 
got {}.", aggRel.groupings().size());
   if (aggRel.groupings().size() == 1) {
-    for (const auto& groupingExpr : 
aggRel.groupings()[0].grouping_expressions()) {
+    // Grouping expressions live in the rel-level pool; each grouping 
references
+    // them by index.
+    for (const auto& ref : aggRel.groupings()[0].expression_references()) {
+      const auto& groupingExpr = aggRel.grouping_expressions(ref);
       // Velox's groupings are limited to be Field.
       
veloxGroupingExprs.emplace_back(exprConverter_->toVeloxExpr(groupingExpr.selection(),
 inputType));
     }
diff --git a/cpp/velox/substrait/SubstraitToVeloxPlanValidator.cc 
b/cpp/velox/substrait/SubstraitToVeloxPlanValidator.cc
index 44313e423f..c041e580db 100644
--- a/cpp/velox/substrait/SubstraitToVeloxPlanValidator.cc
+++ b/cpp/velox/substrait/SubstraitToVeloxPlanValidator.cc
@@ -1256,10 +1256,11 @@ bool SubstraitToVeloxPlanValidator::validate(const 
::substrait::AggregateRel& ag
     }
   }
 
-  // Validate groupings.
+  // Validate groupings. Grouping expressions live in the rel-level pool; each
+  // grouping references them by index.
   for (const auto& grouping : aggRel.groupings()) {
-    for (const auto& groupingExpr : grouping.grouping_expressions()) {
-      const auto& typeCase = groupingExpr.rex_type_case();
+    for (const auto& ref : grouping.expression_references()) {
+      const auto& typeCase = aggRel.grouping_expressions(ref).rex_type_case();
       switch (typeCase) {
         case ::substrait::Expression::RexTypeCase::kSelection:
           break;
@@ -1369,7 +1370,7 @@ bool SubstraitToVeloxPlanValidator::validate(const 
::substrait::AggregateRel& ag
   if (aggRel.measures_size() == 0) {
     bool hasExpr = false;
     for (const auto& grouping : aggRel.groupings()) {
-      if (grouping.grouping_expressions().size() > 0) {
+      if (grouping.expression_references_size() > 0) {
         hasExpr = true;
         break;
       }
diff --git a/cpp/velox/substrait/VeloxToSubstraitPlan.cc 
b/cpp/velox/substrait/VeloxToSubstraitPlan.cc
index f21ba82dab..cdbea2ef32 100644
--- a/cpp/velox/substrait/VeloxToSubstraitPlan.cc
+++ b/cpp/velox/substrait/VeloxToSubstraitPlan.cc
@@ -247,9 +247,12 @@ void VeloxToSubstraitPlanConvertor::toSubstrait(
   int64_t groupingKeySize = groupingKeys.size();
   ::substrait::AggregateRel_Grouping* aggGroupings = 
aggregateRel->add_groupings();
 
+  // Populate the rel-level grouping expression pool in declaration order and
+  // have the single grouping reference every entry by index.
   for (int64_t i = 0; i < groupingKeySize; i++) {
-    aggGroupings->add_grouping_expressions()->mutable_selection()->MergeFrom(
+    aggregateRel->add_grouping_expressions()->mutable_selection()->MergeFrom(
         exprConvertor_->toSubstraitExpr(arena, groupingKeys.at(i), inputType));
+    aggGroupings->add_expression_references(i);
   }
 
   // AggregatesSize should be equal to or greater than the aggregateMasks Size.
diff --git a/cpp/velox/tests/data/q1_first_stage.json 
b/cpp/velox/tests/data/q1_first_stage.json
index 1413ffbd25..2ffdd9d19d 100644
--- a/cpp/velox/tests/data/q1_first_stage.json
+++ b/cpp/velox/tests/data/q1_first_stage.json
@@ -536,28 +536,32 @@
                         },
                         "groupings": [
                             {
-                                "grouping_expressions": [
-                                    {
-                                        "selection": {
-                                            "direct_reference": {
-                                                "struct_field": {
-                                                    "field": 0
-                                                }
-                                            },
-                                            "root_reference": {}
+                                "expression_references": [
+                                    0,
+                                    1
+                                ]
+                            }
+                        ],
+                        "grouping_expressions": [
+                            {
+                                "selection": {
+                                    "direct_reference": {
+                                        "struct_field": {
+                                            "field": 0
                                         }
                                     },
-                                    {
-                                        "selection": {
-                                            "direct_reference": {
-                                                "struct_field": {
-                                                    "field": 1
-                                                }
-                                            },
-                                            "root_reference": {}
+                                    "root_reference": {}
+                                }
+                            },
+                            {
+                                "selection": {
+                                    "direct_reference": {
+                                        "struct_field": {
+                                            "field": 1
                                         }
-                                    }
-                                ]
+                                    },
+                                    "root_reference": {}
+                                }
                             }
                         ],
                         "measures": [
diff --git 
a/gluten-substrait/src/main/java/org/apache/gluten/substrait/rel/AggregateRelNode.java
 
b/gluten-substrait/src/main/java/org/apache/gluten/substrait/rel/AggregateRelNode.java
index 31424e1774..c8b6090ac2 100644
--- 
a/gluten-substrait/src/main/java/org/apache/gluten/substrait/rel/AggregateRelNode.java
+++ 
b/gluten-substrait/src/main/java/org/apache/gluten/substrait/rel/AggregateRelNode.java
@@ -55,13 +55,18 @@ public class AggregateRelNode implements RelNode, 
Serializable {
     RelCommon.Builder relCommonBuilder = RelCommon.newBuilder();
     relCommonBuilder.setDirect(RelCommon.Direct.newBuilder());
 
-    AggregateRel.Grouping.Builder groupingBuilder = 
AggregateRel.Grouping.newBuilder();
-    for (ExpressionNode exprNode : groupings) {
-      groupingBuilder.addGroupingExpressions(exprNode.toProtobuf());
-    }
-
     AggregateRel.Builder aggBuilder = AggregateRel.newBuilder();
     aggBuilder.setCommon(relCommonBuilder.build());
+
+    // Gluten always emits a single grouping set with a flat list of grouping
+    // expressions (GROUPING SETS / CUBE / ROLLUP are expanded into an 
ExpandRel
+    // upstream). Populate the rel-level grouping expression pool in 
declaration
+    // order and have the single grouping reference every entry by index.
+    AggregateRel.Grouping.Builder groupingBuilder = 
AggregateRel.Grouping.newBuilder();
+    for (int i = 0; i < groupings.size(); i++) {
+      aggBuilder.addGroupingExpressions(groupings.get(i).toProtobuf());
+      groupingBuilder.addExpressionReferences(i);
+    }
     aggBuilder.addGroupings(groupingBuilder.build());
 
     for (int i = 0; i < aggregateFunctionNodes.size(); i++) {
diff --git 
a/gluten-substrait/src/main/resources/substrait/proto/substrait/algebra.proto 
b/gluten-substrait/src/main/resources/substrait/proto/substrait/algebra.proto
index aa3a5f1978..6647fe38e8 100644
--- 
a/gluten-substrait/src/main/resources/substrait/proto/substrait/algebra.proto
+++ 
b/gluten-substrait/src/main/resources/substrait/proto/substrait/algebra.proto
@@ -367,16 +367,29 @@ message AggregateRel {
   // Input of the aggregation
   Rel input = 2;
 
-  // A list of expression grouping that the aggregation measured should be 
calculated for.
+  // A list of zero or more grouping sets that the aggregation measures should
+  // be calculated for. There must be at least one grouping set if there are no
+  // measures (but it can be the empty grouping set).
   repeated Grouping groupings = 3;
 
   // A list of one or more aggregate expressions along with an optional filter.
+  // Required if there are no groupings.
   repeated Measure measures = 4;
 
+  // A list of zero or more grouping expressions that grouping sets (i.e.,
+  // `Grouping` messages in the `groupings` field) can reference. Each
+  // expression in this list must be referred to by at least one
+  // `Grouping.expression_references`.
+  repeated Expression grouping_expressions = 5;
+
   substrait.extensions.AdvancedExtension advanced_extension = 10;
 
   message Grouping {
-    repeated Expression grouping_expressions = 1;
+    reserved 1;
+
+    // A list of zero or more references to grouping expressions, i.e., indices
+    // into the `grouping_expression` list.
+    repeated uint32 expression_references = 2;
   }
 
   message Measure {
diff --git 
a/gluten-substrait/src/test/scala/org/apache/gluten/utils/AggregateRelProtoSuite.scala
 
b/gluten-substrait/src/test/scala/org/apache/gluten/utils/AggregateRelProtoSuite.scala
new file mode 100644
index 0000000000..e8e6f78e08
--- /dev/null
+++ 
b/gluten-substrait/src/test/scala/org/apache/gluten/utils/AggregateRelProtoSuite.scala
@@ -0,0 +1,83 @@
+/*
+ * 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.gluten.utils
+
+import org.apache.gluten.substrait.SubstraitContext
+import org.apache.gluten.substrait.expression.{AggregateFunctionNode, 
ExpressionBuilder, ExpressionNode}
+import org.apache.gluten.substrait.rel.RelBuilder
+
+import org.scalatest.funsuite.AnyFunSuite
+
+/**
+ * Locks the AggregateRel producer contract after the Substrait 0.98 
migration. 0.98 moved the
+ * grouping expressions out of the per-grouping 
`Grouping.grouping_expressions` (field 1) into a
+ * rel-level pool `AggregateRel.grouping_expressions` (field 5) that each 
grouping set references by
+ * index through `Grouping.expression_references` (field 2). Gluten only ever 
emits a single
+ * grouping set with a flat list of grouping expressions (GROUPING SETS / CUBE 
/ ROLLUP are expanded
+ * into an ExpandRel upstream), so the producer populates the pool in 
declaration order and has the
+ * single grouping reference every entry as `[0, 1, ..., n - 1]`. This suite 
pins that mapping.
+ */
+class AggregateRelProtoSuite extends AnyFunSuite {
+
+  test("makeAggregateRel emits a rel-level pool referenced by the single 
grouping") {
+    val context = new SubstraitContext
+    val groupings: java.util.List[ExpressionNode] =
+      java.util.Arrays.asList[ExpressionNode](
+        ExpressionBuilder.makeSelection(0),
+        ExpressionBuilder.makeSelection(1))
+    val rel = RelBuilder.makeAggregateRel(
+      null,
+      groupings,
+      java.util.Collections.emptyList[AggregateFunctionNode](),
+      java.util.Collections.emptyList[ExpressionNode](),
+      null,
+      context,
+      0L)
+    val aggRel = rel.toProtobuf.getAggregate
+
+    // The flat grouping list becomes the rel-level pool, in declaration order.
+    def poolField(i: Int): Int =
+      
aggRel.getGroupingExpressions(i).getSelection.getDirectReference.getStructField.getField
+    assert(aggRel.getGroupingExpressionsCount === 2)
+    assert(poolField(0) === 0)
+    assert(poolField(1) === 1)
+
+    // A single grouping references every pool entry by index.
+    assert(aggRel.getGroupingsCount === 1)
+    assert(aggRel.getGroupings(0).getExpressionReferencesCount === 2)
+    assert(aggRel.getGroupings(0).getExpressionReferences(0) === 0)
+    assert(aggRel.getGroupings(0).getExpressionReferences(1) === 1)
+  }
+
+  test("makeAggregateRel with no grouping keys emits one empty grouping 
(global aggregation)") {
+    val context = new SubstraitContext
+    val rel = RelBuilder.makeAggregateRel(
+      null,
+      java.util.Collections.emptyList[ExpressionNode](),
+      java.util.Collections.emptyList[AggregateFunctionNode](),
+      java.util.Collections.emptyList[ExpressionNode](),
+      null,
+      context,
+      0L
+    )
+    val aggRel = rel.toProtobuf.getAggregate
+
+    assert(aggRel.getGroupingExpressionsCount === 0)
+    assert(aggRel.getGroupingsCount === 1)
+    assert(aggRel.getGroupings(0).getExpressionReferencesCount === 0)
+  }
+}


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

Reply via email to