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]])
]]>