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

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


The following commit(s) were added to refs/heads/master by this push:
     new c8d1cb5a1 [CALCITE-5000] Expand `AGGREGATE_REDUCE_FUNCTIONS`, when arg 
of agg-call exists in the aggregate's group
c8d1cb5a1 is described below

commit c8d1cb5a1e1b1dd98966c29308bc6b1aed245501
Author: xurenhe <[email protected]>
AuthorDate: Wed Feb 9 16:09:55 2022 +0800

    [CALCITE-5000] Expand `AGGREGATE_REDUCE_FUNCTIONS`, when arg of agg-call 
exists in the aggregate's group
---
 .../rel/rules/AggregateReduceFunctionsRule.java    | 62 +++++++++++++++++++++-
 .../org/apache/calcite/test/RelOptRulesTest.java   | 12 +++++
 .../org/apache/calcite/test/RelOptRulesTest.xml    | 32 +++++++++--
 3 files changed, 100 insertions(+), 6 deletions(-)

diff --git 
a/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
 
b/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
index 522835f8a..6cdf849be 100644
--- 
a/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
+++ 
b/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
@@ -20,12 +20,15 @@ import org.apache.calcite.plan.RelOptCluster;
 import org.apache.calcite.plan.RelOptRuleCall;
 import org.apache.calcite.plan.RelOptRuleOperand;
 import org.apache.calcite.plan.RelRule;
+import org.apache.calcite.rel.RelCollations;
+import org.apache.calcite.rel.RelNode;
 import org.apache.calcite.rel.core.Aggregate;
 import org.apache.calcite.rel.core.AggregateCall;
 import org.apache.calcite.rel.logical.LogicalAggregate;
 import org.apache.calcite.rel.type.RelDataType;
 import org.apache.calcite.rel.type.RelDataTypeFactory;
 import org.apache.calcite.rex.RexBuilder;
+import org.apache.calcite.rex.RexInputRef;
 import org.apache.calcite.rex.RexLiteral;
 import org.apache.calcite.rex.RexNode;
 import org.apache.calcite.sql.SqlAggFunction;
@@ -178,6 +181,37 @@ public class AggregateReduceFunctionsRule
         && config.extraCondition().test(call);
   }
 
+  /** Returns whether this rule can reduce some agg-call,
+   * which its arg exists in the aggregate's group. */
+  public boolean canReduceAggCallByGrouping(Aggregate oldAggRel, AggregateCall 
call) {
+    if (!Aggregate.isSimple(oldAggRel)) {
+      return false;
+    }
+    if (call.hasFilter() || call.distinctKeys != null || call.collation != 
RelCollations.EMPTY) {
+      return false;
+    }
+    final List<Integer> argList = call.getArgList();
+    if (argList.size() != 1) {
+      return false;
+    }
+    if (!oldAggRel.getGroupSet().asSet().contains(argList.get(0))) {
+      // arg doesn't exist in aggregate's group.
+      return false;
+    }
+    final SqlKind kind = call.getAggregation().getKind();
+    switch (kind) {
+    case AVG:
+    case MAX:
+    case MIN:
+    case ANY_VALUE:
+    case FIRST_VALUE:
+    case LAST_VALUE:
+      return true;
+    default:
+      return false;
+    }
+  }
+
   /**
    * Reduces calls to functions AVG, SUM, STDDEV_POP, STDDEV_SAMP, VAR_POP,
    * VAR_SAMP, COVAR_POP, COVAR_SAMP, REGR_SXX, REGR_SYY if the function is
@@ -228,7 +262,8 @@ public class AggregateReduceFunctionsRule
     }
     newAggregateRel(relBuilder, oldAggRel, newCalls);
     newCalcRel(relBuilder, oldAggRel.getRowType(), projList);
-    ruleCall.transformTo(relBuilder.build());
+    final RelNode build = relBuilder.build();
+    ruleCall.transformTo(build);
   }
 
   private RexNode reduceAgg(
@@ -237,7 +272,12 @@ public class AggregateReduceFunctionsRule
       List<AggregateCall> newCalls,
       Map<AggregateCall, RexNode> aggCallMapping,
       List<RexNode> inputExprs) {
-    if (canReduce(oldCall)) {
+    if (canReduceAggCallByGrouping(oldAggRel, oldCall)) {
+      // replace original MAX/MIN/AVG/ANY_VALUE/FIRST_VALUE/LAST_VALUE(x) with
+      // target field of x, when x exists in group
+      final RexNode reducedNode = reduceAggCallByGrouping(oldAggRel, oldCall);
+      return reducedNode;
+    } else if (canReduce(oldCall)) {
       final Integer y;
       final Integer x;
       final SqlKind kind = oldCall.getAggregation().getKind();
@@ -572,6 +612,24 @@ public class AggregateReduceFunctionsRule
         oldCall.getType(), result);
   }
 
+  private static RexNode reduceAggCallByGrouping(
+      Aggregate oldAggRel,
+      AggregateCall oldCall) {
+
+    final RexBuilder rexBuilder = oldAggRel.getCluster().getRexBuilder();
+    final List<Integer> oldGroups = oldAggRel.getGroupSet().asList();
+    final Integer firstArg = oldCall.getArgList().get(0);
+    final int index = oldGroups.lastIndexOf(firstArg);
+    assert index >= 0;
+
+    final RexInputRef refByGroup = RexInputRef.of(index, 
oldAggRel.getRowType().getFieldList());
+    if (refByGroup.getType().equals(oldCall.getType())) {
+      return refByGroup;
+    } else {
+      return rexBuilder.makeCast(oldCall.getType(), refByGroup);
+    }
+  }
+
   private static RexNode getSumAggregatedRexNode(Aggregate oldAggRel,
       AggregateCall oldCall,
       List<AggregateCall> newCalls,
diff --git a/core/src/test/java/org/apache/calcite/test/RelOptRulesTest.java 
b/core/src/test/java/org/apache/calcite/test/RelOptRulesTest.java
index 8aef5b85f..ad714d26d 100644
--- a/core/src/test/java/org/apache/calcite/test/RelOptRulesTest.java
+++ b/core/src/test/java/org/apache/calcite/test/RelOptRulesTest.java
@@ -6370,6 +6370,18 @@ class RelOptRulesTest extends RelOptTestBase {
     sql(sql).withRule(rule).check();
   }
 
+  /** Test case for
+   * <a 
href="https://issues.apache.org/jira/browse/CALCITE-5000";>[CALCITE-5000]
+   * Expand rule of `AGGREGATE_REDUCE_FUNCTIONS`, when arg of agg-call exist 
in the agg's group</a>.
+   */
+  @Test void testReduceAggregateFunctionsByGroup() {
+    final String sql = "select sal, max(sal) as sal_max, min(sal) as 
sal_min,\n"
+        + "avg(sal) sal_avg, any_value(sal) as sal_val, first_value(sal) as 
sal_first,\n"
+        + "last_value(sal) as sal_last\n"
+        + "from emp group by sal, deptno";
+    sql(sql).withRule(CoreRules.AGGREGATE_REDUCE_FUNCTIONS, 
CoreRules.PROJECT_MERGE).check();
+  }
+
   @Test void testReduceAllAggregateFunctions() {
     // configure rule to reduce all used functions
     final RelOptRule rule = AggregateReduceFunctionsRule.Config.DEFAULT
diff --git 
a/core/src/test/resources/org/apache/calcite/test/RelOptRulesTest.xml 
b/core/src/test/resources/org/apache/calcite/test/RelOptRulesTest.xml
index a748e544b..e2774111b 100644
--- a/core/src/test/resources/org/apache/calcite/test/RelOptRulesTest.xml
+++ b/core/src/test/resources/org/apache/calcite/test/RelOptRulesTest.xml
@@ -9922,6 +9922,30 @@ LogicalAggregate(group=[{0}], EXPR$1=[SUM($1)])
     LogicalAggregate(group=[{0}], EXPR$1=[SUM($1)])
       LogicalProject(ENAME=[$1], MGR=[$3])
         LogicalTableScan(table=[[CATALOG, SALES, EMP]])
+]]>
+    </Resource>
+  </TestCase>
+  <TestCase name="testReduceAggregateFunctionsByGroup">
+    <Resource name="sql">
+      <![CDATA[select sal, max(sal) as sal_max, min(sal) as sal_min,
+avg(sal) sal_avg, any_value(sal) as sal_val, first_value(sal) as sal_first,
+last_value(sal) as sal_last
+from emp group by sal, deptno]]>
+    </Resource>
+    <Resource name="planBefore">
+      <![CDATA[
+LogicalProject(SAL=[$0], SAL_MAX=[$2], SAL_MIN=[$3], SAL_AVG=[$4], 
SAL_VAL=[$5], SAL_FIRST=[$6], SAL_LAST=[$7])
+  LogicalAggregate(group=[{0, 1}], SAL_MAX=[MAX($0)], SAL_MIN=[MIN($0)], 
SAL_AVG=[AVG($0)], SAL_VAL=[ANY_VALUE($0)], SAL_FIRST=[FIRST_VALUE($0)], 
SAL_LAST=[LAST_VALUE($0)])
+    LogicalProject(SAL=[$5], DEPTNO=[$7])
+      LogicalTableScan(table=[[CATALOG, SALES, EMP]])
+]]>
+    </Resource>
+    <Resource name="planAfter">
+      <![CDATA[
+LogicalProject(SAL=[$0], SAL_MAX=[$0], SAL_MIN=[$0], SAL_AVG=[$0], 
SAL_VAL=[$0], SAL_FIRST=[$0], SAL_LAST=[$0])
+  LogicalAggregate(group=[{0, 1}])
+    LogicalProject(SAL=[$5], DEPTNO=[$7])
+      LogicalTableScan(table=[[CATALOG, SALES, EMP]])
 ]]>
     </Resource>
   </TestCase>
@@ -9961,8 +9985,8 @@ LogicalAggregate(group=[{0}], EXPR$1=[MAX($0)], 
EXPR$2=[AVG($1)], EXPR$3=[MIN($0
     </Resource>
     <Resource name="planAfter">
       <![CDATA[
-LogicalProject(NAME=[$0], EXPR$1=[$1], EXPR$2=[CAST(/($2, $3)):INTEGER NOT 
NULL], EXPR$3=[$4])
-  LogicalAggregate(group=[{0}], EXPR$1=[MAX($0)], agg#1=[$SUM0($1)], 
agg#2=[COUNT()], EXPR$3=[MIN($0)])
+LogicalProject(NAME=[$0], EXPR$1=[$0], EXPR$2=[CAST(/($1, $2)):INTEGER NOT 
NULL], EXPR$3=[$0])
+  LogicalAggregate(group=[{0}], agg#0=[$SUM0($1)], agg#1=[COUNT()])
     LogicalProject(NAME=[$1], DEPTNO=[$0])
       LogicalTableScan(table=[[CATALOG, SALES, DEPT]])
 ]]>
@@ -10025,8 +10049,8 @@ LogicalAggregate(group=[{0}], EXPR$1=[MAX($0)], 
EXPR$2=[AVG($1)], EXPR$3=[MIN($0
     </Resource>
     <Resource name="planAfter">
       <![CDATA[
-LogicalProject(NAME=[$0], EXPR$1=[$1], EXPR$2=[CAST(/($2, $3)):INTEGER NOT 
NULL], EXPR$3=[$4])
-  LogicalAggregate(group=[{0}], EXPR$1=[MAX($0)], agg#1=[SUM($1)], 
agg#2=[COUNT()], EXPR$3=[MIN($0)])
+LogicalProject(NAME=[$0], EXPR$1=[$0], EXPR$2=[CAST(/($1, $2)):INTEGER NOT 
NULL], EXPR$3=[$0])
+  LogicalAggregate(group=[{0}], agg#0=[SUM($1)], agg#1=[COUNT()])
     LogicalProject(NAME=[$1], DEPTNO=[$0])
       LogicalTableScan(table=[[CATALOG, SALES, DEPT]])
 ]]>

Reply via email to