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

morrySnow pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/doris.git


The following commit(s) were added to refs/heads/master by this push:
     new 749290cb041 [opt](repeat) deduplicate grouping scalar functions from 
projetion inrepeat (#65880)
749290cb041 is described below

commit 749290cb04107ae08e6b33980495f81b35874467
Author: morrySnow <[email protected]>
AuthorDate: Tue Jul 28 10:14:59 2026 +0800

    [opt](repeat) deduplicate grouping scalar functions from projetion inrepeat 
(#65880)
    
    ### What problem does this PR solve?
    
    Related PR: #34196
    
    Problem Summary:
    
    when normalize repeat, we maybe generate duplicate grouping functions
    because the same grouping functions appear in projection more than once.
---
 .../nereids/rules/analysis/NormalizeRepeat.java    | 34 +++++++++-------------
 .../doris/nereids/trees/plans/algebra/Repeat.java  |  2 +-
 .../apache/doris/nereids/util/ExpressionUtils.java |  4 +--
 .../rules/analysis/NormalizeRepeatTest.java        | 30 +++++++++++++++++++
 4 files changed, 47 insertions(+), 23 deletions(-)

diff --git 
a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
 
b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
index 1f00cc3e619..bbabf96f7f1 100644
--- 
a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
+++ 
b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeat.java
@@ -129,8 +129,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
     }
 
     private static void checkGroupingSetsSize(LogicalRepeat<Plan> repeat) {
-        Set<Expression> flattenGroupingSetExpr = ImmutableSet.copyOf(
-                ExpressionUtils.flatExpressions(repeat.getGroupingSets()));
+        Set<Expression> flattenGroupingSetExpr = 
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
         if (flattenGroupingSetExpr.size() > 
LogicalRepeat.MAX_GROUPING_SETS_NUM) {
             throw new AnalysisException(
                     "Too many sets in GROUP BY clause, the max grouping sets 
item is "
@@ -139,7 +138,9 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
     }
 
     private static LogicalAggregate<Plan> normalizeRepeat(LogicalRepeat<Plan> 
repeat) {
-        Set<Expression> needToSlotsGroupingExpr = 
collectNeedToSlotGroupingExpr(repeat);
+        // grouping sets should be pushed down, e.g. grouping sets((k + 1)),
+        // we should push down the `k + 1` to the bottom plan
+        Set<Expression> needToSlotsGroupingExpr = 
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
         NormalizeToSlotContext groupingExprContext = buildContext(repeat, 
needToSlotsGroupingExpr);
         Map<Expression, NormalizeToSlotTriplet> groupingExprMap = 
groupingExprContext.getNormalizeToSlotMap();
         Map<Expression, Alias> existsAlias = getExistsAlias(repeat, 
groupingExprMap);
@@ -159,7 +160,8 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
         // rewrite grouping scalar function to virtual slots
         // rewrite the arguments of agg function to slots
         List<NamedExpression> normalizedAggOutput = Lists.newArrayList();
-        List<NamedExpression> groupingFunctions = Lists.newArrayList();
+        // use a map to deduplicate grouping scalar functions in projection
+        Map<GroupingScalarFunction, NamedExpression> groupingFunctions = 
Maps.newHashMap();
         for (Expression expr : repeat.getOutputExpressions()) {
             Expression rewrittenExpr = expr.rewriteDownShortCircuit(
                     e -> 
normalizeAggFuncChildrenAndGroupingScalarFunc(argsContext, e, 
groupingFunctions));
@@ -172,17 +174,16 @@ public class NormalizeRepeat extends 
OneAnalysisRuleFactory {
         Set<SlotReference> aggUsedSlots = ExpressionUtils.collect(
                 normalizedAggOutput, expr -> 
expr.getClass().equals(SlotReference.class));
 
-        Set<Slot> groupingSetsUsedSlot = ImmutableSet.copyOf(
-                ExpressionUtils.flatExpressions(normalizedGroupingSets));
+        Set<Slot> groupingSetsUsedSlot = 
ExpressionUtils.flatExpressions(normalizedGroupingSets);
 
         SetView<SlotReference> aggUsedSlotNotInGroupBy
-                = Sets.difference(Sets.difference(aggUsedSlots, 
groupingFunctions.stream()
+                = Sets.difference(Sets.difference(aggUsedSlots, 
groupingFunctions.values().stream()
                 .map(NamedExpression::toSlot).collect(Collectors.toSet())), 
groupingSetsUsedSlot);
 
         List<NamedExpression> normalizedRepeatOutput = 
ImmutableList.<NamedExpression>builder()
                 .addAll(groupingSetsUsedSlot)
                 .addAll(aggUsedSlotNotInGroupBy)
-                .addAll(groupingFunctions)
+                .addAll(groupingFunctions.values())
                 .build();
 
         // 3 parts need push down:
@@ -218,7 +219,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
 
         List<Expression> normalizedAggGroupBy = 
ImmutableList.<Expression>builder()
                 .addAll(groupingSetsUsedSlot)
-                
.addAll(groupingFunctions.stream().map(NamedExpression::toSlot).collect(Collectors.toList()))
+                
.addAll(groupingFunctions.values().stream().map(NamedExpression::toSlot).collect(Collectors.toList()))
                 .add(groupingId)
                 .build();
 
@@ -227,13 +228,6 @@ public class NormalizeRepeat extends 
OneAnalysisRuleFactory {
                 Optional.of(normalizedRepeat), normalizedRepeat);
     }
 
-    private static Set<Expression> 
collectNeedToSlotGroupingExpr(LogicalRepeat<Plan> repeat) {
-        // grouping sets should be pushed down, e.g. grouping sets((k + 1)),
-        // we should push down the `k + 1` to the bottom plan
-        return ImmutableSet.copyOf(
-                ExpressionUtils.flatExpressions(repeat.getGroupingSets()));
-    }
-
     private static Set<Expression> 
collectNeedToSlotArgsOfGroupingScalarFuncAndAggFunc(LogicalRepeat<Plan> repeat) 
{
         Set<GroupingScalarFunction> groupingScalarFunctions = 
ExpressionUtils.collect(
                 repeat.getOutputExpressions(), 
GroupingScalarFunction.class::isInstance);
@@ -300,7 +294,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
 
     private static NormalizeToSlotContext buildContextWithAlias(Repeat<? 
extends Plan> repeat,
             Map<Expression, Alias> existsAliasMap, Collection<? extends 
Expression> sourceExpressions) {
-        List<Expression> groupingSetExpressions = 
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
+        Set<Expression> groupingSetExpressions = 
ExpressionUtils.flatExpressions(repeat.getGroupingSets());
         Map<Expression, NormalizeToSlotTriplet> normalizeToSlotMap = 
Maps.newLinkedHashMap();
         for (Expression expression : sourceExpressions) {
             Optional<NormalizeToSlotTriplet> pushDownTriplet;
@@ -325,7 +319,7 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
     }
 
     private static Expression 
normalizeAggFuncChildrenAndGroupingScalarFunc(NormalizeToSlotContext context,
-            Expression expr, List<NamedExpression> groupingSetExpressions) {
+            Expression expr, Map<GroupingScalarFunction, NamedExpression> 
groupingSetExpressions) {
         if (expr instanceof AggregateFunction) {
             AggregateFunction function = (AggregateFunction) expr;
             List<Expression> normalizedRealExpressions = 
context.normalizeToUseSlotRef(function.getArguments());
@@ -336,8 +330,8 @@ public class NormalizeRepeat extends OneAnalysisRuleFactory 
{
             List<Expression> normalizedRealExpressions = 
context.normalizeToUseSlotRef(function.getArguments());
             function = function.withChildren(normalizedRealExpressions);
             // eliminate GroupingScalarFunction and replace to 
VirtualSlotReference
-            Alias alias = new Alias(function, 
Repeat.generateVirtualSlotName(function));
-            groupingSetExpressions.add(alias);
+            NamedExpression alias = 
groupingSetExpressions.computeIfAbsent(function,
+                    f -> new Alias(f, Repeat.generateVirtualSlotName(f)));
             return alias.toSlot();
         } else {
             return expr;
diff --git 
a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
 
b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
index 7a7f1f1f3c2..2e0dde6c305 100644
--- 
a/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
+++ 
b/fe/fe-core/src/main/java/org/apache/doris/nereids/trees/plans/algebra/Repeat.java
@@ -51,7 +51,7 @@ public interface Repeat<CHILD_PLAN extends Plan> extends 
Aggregate<CHILD_PLAN> {
 
     @Override
     default List<Expression> getGroupByExpressions() {
-        return ExpressionUtils.flatExpressions(getGroupingSets());
+        return 
ImmutableList.copyOf(ExpressionUtils.flatExpressions(getGroupingSets()));
     }
 
     @Override
diff --git 
a/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java 
b/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
index 6d318cafd9c..fd7e761bc6e 100644
--- 
a/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
+++ 
b/fe/fe-core/src/main/java/org/apache/doris/nereids/util/ExpressionUtils.java
@@ -851,13 +851,13 @@ public class ExpressionUtils {
     }
 
     /** flatExpressions */
-    public static <E extends Expression> List<E> flatExpressions(List<List<E>> 
expressionLists) {
+    public static <E extends Expression> Set<E> flatExpressions(List<List<E>> 
expressionLists) {
         int num = 0;
         for (List<E> expressionList : expressionLists) {
             num += expressionList.size();
         }
 
-        ImmutableList.Builder<E> flatten = 
ImmutableList.builderWithExpectedSize(num);
+        ImmutableSet.Builder<E> flatten = 
ImmutableSet.builderWithExpectedSize(num);
         for (List<E> expressionList : expressionLists) {
             flatten.addAll(expressionList);
         }
diff --git 
a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
 
b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
index e2c874aad17..c1ba24c0e2d 100644
--- 
a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
+++ 
b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/analysis/NormalizeRepeatTest.java
@@ -21,8 +21,10 @@ import org.apache.doris.nereids.trees.expressions.Alias;
 import org.apache.doris.nereids.trees.expressions.Slot;
 import org.apache.doris.nereids.trees.expressions.functions.agg.Sum;
 import org.apache.doris.nereids.trees.expressions.functions.scalar.GroupingId;
+import 
org.apache.doris.nereids.trees.expressions.functions.scalar.GroupingScalarFunction;
 import org.apache.doris.nereids.trees.plans.Plan;
 import org.apache.doris.nereids.trees.plans.algebra.Repeat.RepeatType;
+import org.apache.doris.nereids.trees.plans.logical.LogicalAggregate;
 import org.apache.doris.nereids.trees.plans.logical.LogicalOlapScan;
 import org.apache.doris.nereids.trees.plans.logical.LogicalRepeat;
 import org.apache.doris.nereids.util.MemoPatternMatchSupported;
@@ -31,6 +33,7 @@ import org.apache.doris.nereids.util.PlanChecker;
 import org.apache.doris.nereids.util.PlanConstructor;
 
 import com.google.common.collect.ImmutableList;
+import org.junit.jupiter.api.Assertions;
 import org.junit.jupiter.api.Test;
 
 public class NormalizeRepeatTest implements MemoPatternMatchSupported {
@@ -92,4 +95,31 @@ public class NormalizeRepeatTest implements 
MemoPatternMatchSupported {
                         logicalAggregate(logicalRepeat(logicalOlapScan()))
                 );
     }
+
+    @Test
+    public void testDeduplicateGroupingScalarFunctions() {
+        Slot id = scan1.getOutput().get(0);
+        Slot name = scan1.getOutput().get(1);
+        Alias firstGroupingId = new Alias(new GroupingId(name), 
"first_grouping_id");
+        Alias secondGroupingId = new Alias(new GroupingId(name), 
"second_grouping_id");
+        LogicalRepeat<Plan> repeat = new LogicalRepeat<>(
+                ImmutableList.of(ImmutableList.of(id)),
+                ImmutableList.of(id, firstGroupingId, secondGroupingId),
+                RepeatType.GROUPING_SETS,
+                scan1
+        );
+
+        LogicalAggregate<Plan> normalizedAggregate = 
NormalizeRepeat.doNormalize(repeat);
+        LogicalRepeat<Plan> normalizedRepeat = (LogicalRepeat<Plan>) 
normalizedAggregate.child();
+        long groupingFunctionCount = 
normalizedRepeat.getOutputExpressions().stream()
+                .filter(expression -> expression instanceof Alias
+                        && ((Alias) expression).child() instanceof 
GroupingScalarFunction)
+                .count();
+
+        Assertions.assertEquals(1, groupingFunctionCount);
+        Assertions.assertEquals(3, 
normalizedAggregate.getOutputExpressions().size());
+        Assertions.assertEquals(
+                ((Alias) 
normalizedAggregate.getOutputExpressions().get(1)).child(),
+                ((Alias) 
normalizedAggregate.getOutputExpressions().get(2)).child());
+    }
 }


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

Reply via email to