englefly commented on code in PR #64849:
URL: https://github.com/apache/doris/pull/64849#discussion_r3732912461
##########
fe/fe-core/src/main/java/org/apache/doris/nereids/rules/rewrite/EliminateGroupByKey.java:
##########
@@ -17,90 +17,230 @@
package org.apache.doris.nereids.rules.rewrite;
-import org.apache.doris.nereids.annotation.DependsRules;
+import org.apache.doris.nereids.jobs.JobContext;
import org.apache.doris.nereids.properties.DataTrait;
import org.apache.doris.nereids.properties.FuncDeps;
-import org.apache.doris.nereids.rules.Rule;
-import org.apache.doris.nereids.rules.RuleType;
+import org.apache.doris.nereids.trees.expressions.Alias;
+import org.apache.doris.nereids.trees.expressions.ExprId;
import org.apache.doris.nereids.trees.expressions.Expression;
import org.apache.doris.nereids.trees.expressions.NamedExpression;
import org.apache.doris.nereids.trees.expressions.Slot;
+import org.apache.doris.nereids.trees.expressions.functions.agg.AnyValue;
import org.apache.doris.nereids.trees.plans.Plan;
+import org.apache.doris.nereids.trees.plans.algebra.Aggregate;
import org.apache.doris.nereids.trees.plans.logical.LogicalAggregate;
+import org.apache.doris.nereids.trees.plans.logical.LogicalCTEConsumer;
+import org.apache.doris.nereids.trees.plans.logical.LogicalFilter;
+import org.apache.doris.nereids.trees.plans.logical.LogicalProject;
+import org.apache.doris.nereids.trees.plans.visitor.CustomRewriter;
+import org.apache.doris.nereids.trees.plans.visitor.DefaultPlanRewriter;
-import com.google.common.collect.ImmutableList;
+import com.google.common.collect.LinkedHashMultimap;
+import com.google.common.collect.Multimap;
import java.util.ArrayList;
import java.util.HashMap;
import java.util.HashSet;
+import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Map.Entry;
import java.util.Set;
-
/**
* Eliminate group by key based on fd item information.
* such as:
* for a -> b, we can get:
* group by a, b, c => group by a, c
+ *
+ * When a group-by key is FD-redundant but still needed in the output,
+ * it is wrapped with any_value() and assigned a fresh ExprId.
+ * Upper plan references are rewritten via ExprIdRewriter so that
+ * all ancestor nodes see the new ExprIds.
*/
-@DependsRules({EliminateGroupBy.class, ColumnPruning.class})
-public class EliminateGroupByKey implements RewriteRuleFactory {
+public class EliminateGroupByKey extends DefaultPlanRewriter<Map<ExprId,
ExprId>> implements CustomRewriter {
+ private ExprIdRewriter exprIdReplacer;
+
+ @Override
+ public Plan rewriteRoot(Plan plan, JobContext jobContext) {
+ if (!plan.containsType(Aggregate.class)) {
+ return plan;
+ }
+ Map<ExprId, ExprId> replaceMap = new HashMap<>();
+ ExprIdRewriter.ReplaceRule replaceRule = new
ExprIdRewriter.ReplaceRule(replaceMap, false);
+ exprIdReplacer = new ExprIdRewriter(replaceRule, jobContext);
+ return plan.accept(this, replaceMap);
+ }
+
+ @Override
+ public Plan visit(Plan plan, Map<ExprId, ExprId> replaceMap) {
+ plan = visitChildren(this, plan, replaceMap);
+ plan = exprIdReplacer.rewriteExpr(plan, replaceMap);
+ return plan;
+ }
+
+ @Override
+ public Plan visitLogicalProject(LogicalProject<? extends Plan> proj,
Map<ExprId, ExprId> replaceMap) {
+ proj = visitChildren(this, proj, replaceMap);
+
+ // Find the Aggregate child, possibly through a Filter
+ Plan child = proj.child(0);
+ LogicalAggregate<? extends Plan> agg;
+ boolean hasFilter = child instanceof LogicalFilter;
+ if (hasFilter && child.child(0) instanceof LogicalAggregate) {
+ agg = (LogicalAggregate<? extends Plan>) child.child(0);
+ } else if (child instanceof LogicalAggregate) {
+ agg = (LogicalAggregate<? extends Plan>) child;
+ } else {
+ return exprIdReplacer.rewriteExpr(proj, replaceMap);
+ }
+
+ // Don't transform if source repeat is present
+ if (agg.getSourceRepeat().isPresent()) {
+ return exprIdReplacer.rewriteExpr(proj, replaceMap);
+ }
+
+ // Rewrite proj and the filter (if present) through the replaceMap
accumulated
+ // by visitChildren, so that ExprId replacements from nested rewrites
+ // (e.g. inner aggregates) are reflected in the required-output slot
set.
+ proj = (LogicalProject<? extends Plan>)
exprIdReplacer.rewriteExpr(proj, replaceMap);
+ if (hasFilter) {
+ child = exprIdReplacer.rewriteExpr(child, replaceMap);
+ }
+
+ // Compute requireOutput: slots needed by the Project (and Filter, if
present)
+ Set<Slot> requireOutput = new HashSet<>(proj.getInputSlots());
+ if (hasFilter) {
+ requireOutput.addAll(child.getInputSlots());
+ }
+
+ // Transform the aggregate
+ EliminateResult result = eliminateGroupByKeyWithMap(agg,
requireOutput);
+ if (!result.changed) {
+ return proj;
+ }
+
+ // Merge into the global replaceMap so that all ancestor nodes get
rewritten
+ replaceMap.putAll(result.replaceMap);
+
+ // Rebuild the child chain with the new aggregate,
+ // and rewrite the Filter (if present) and Project expressions
+ Plan newChild;
+ if (hasFilter) {
+ Plan updatedFilter = child.withChildren(result.newAgg);
+ newChild = exprIdReplacer.rewriteExpr(updatedFilter, replaceMap);
+ } else {
+ newChild = result.newAgg;
+ }
+ Plan newProj = exprIdReplacer.rewriteExpr(proj.withChildren(newChild),
replaceMap);
+ return newProj;
+ }
@Override
- public List<Rule> buildRules() {
- return ImmutableList.of(
- RuleType.ELIMINATE_GROUP_BY_KEY.build(
- logicalProject(logicalAggregate().when(agg ->
!agg.getSourceRepeat().isPresent()))
- .then(proj -> {
- LogicalAggregate<? extends Plan> agg =
proj.child();
- LogicalAggregate<Plan> newAgg =
eliminateGroupByKey(agg, proj.getInputSlots());
- if (newAgg == null) {
- return null;
- }
- return proj.withChildren(newAgg);
- })),
- RuleType.ELIMINATE_FILTER_GROUP_BY_KEY.build(
- logicalProject(logicalFilter(logicalAggregate()
- .when(agg ->
!agg.getSourceRepeat().isPresent())))
- .then(proj -> {
- LogicalAggregate<? extends Plan> agg =
proj.child().child();
- Set<Slot> requireSlots = new
HashSet<>(proj.getInputSlots());
-
requireSlots.addAll(proj.child(0).getInputSlots());
- LogicalAggregate<Plan> newAgg =
eliminateGroupByKey(agg, requireSlots);
- if (newAgg == null) {
- return null;
- }
- return
proj.withChildren(proj.child().withChildren(newAgg));
- })
- )
- );
+ public Plan visitLogicalCTEConsumer(LogicalCTEConsumer cteConsumer,
Map<ExprId, ExprId> replaceMap) {
Review Comment:
不能,因为现在这个rule是whole tree rewrite rule.
当producer 的输出 slot id 因为这个rule 发生变化时 consumer 的slot map需要更新
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]