Skip to content

Commit 19da05d

Browse files
authored
[Enhancement] Push down non-group-by aggregate through union all (#73930)
Signed-off-by: JasonQuCode <88529388+JasonQuCode@users.noreply.github.com> Co-authored-by: JasonQuCode <88529388+JasonQuCode@users.noreply.github.com>
1 parent 5702ecb commit 19da05d

9 files changed

Lines changed: 1607 additions & 0 deletions

File tree

docs/en/administration/management/FE_parameters/user_query_loading.md

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -362,6 +362,15 @@ Starting from version 3.3.0, the system defaults to refreshing one partition at
362362
- Description: Whether to enable predicate columns collection. If disabled, predicate columns will not be recorded during query optimization.
363363
- Introduced in: -
364364

365+
### `push_down_non_grouped_aggregate_below_union`
366+
367+
- Default: false
368+
- Type: Boolean
369+
- Unit: -
370+
- Is mutable: Yes
371+
- Description: Whether to push down non-grouped aggregations below Union in the physical plan.
372+
- Introduced in: -
373+
365374
### `enable_query_queue_v2`
366375

367376
- Default: true

docs/ja/administration/management/FE_parameters/user_query_loading.md

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -353,6 +353,15 @@ ADMIN SET FRONTEND CONFIG ("key" = "value");
353353
- 説明:述語列収集を有効にするかどうか。無効にすると、クエリオプティマイゼーション中に述語列は記録されません。
354354
- 導入時期:-
355355

356+
### `push_down_non_grouped_aggregate_below_union`
357+
358+
- デフォルト:false
359+
- タイプ:Boolean
360+
- 単位:-
361+
- 変更可能:Yes
362+
- 説明:物理プランにおいて、GROUP BY を持たない集約を Union の下にプッシュダウンするかどうか。
363+
- 導入時期:-
364+
356365
### `enable_query_queue_v2`
357366

358367
- デフォルト:true

docs/zh/administration/management/FE_parameters/user_query_loading.md

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -362,6 +362,15 @@ ADMIN SET FRONTEND CONFIG ("key" = "value");
362362
- 描述: 是否启用谓词列收集。如果禁用,谓词列在查询优化期间将不会被记录。
363363
- 引入版本: -
364364

365+
### `push_down_non_grouped_aggregate_below_union`
366+
367+
- 默认值: false
368+
- 类型: Boolean
369+
- 单位: -
370+
- 是否可变: Yes
371+
- 描述: 是否在物理计划中将不带 GROUP BY 的聚合下推到 Union 下方。
372+
- 引入版本: -
373+
365374
### `enable_query_queue_v2`
366375

367376
- 默认值: true

fe/fe-core/src/main/java/com/starrocks/common/Config.java

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2496,6 +2496,10 @@ public class Config extends ConfigBase {
24962496
"generation and dump them to logs when plan generation fails (e.g. CBO timeout) for diagnosis.")
24972497
public static boolean enable_dump_optimizer_trace_on_error = false;
24982498

2499+
@ConfField(mutable = true, comment = "Whether to push down non-grouped aggregations below Union " +
2500+
"in the physical plan")
2501+
public static boolean push_down_non_grouped_aggregate_below_union = false;
2502+
24992503
/**
25002504
* Num of thread to handle statistic collect(analyze command)
25012505
*/

fe/fe-core/src/main/java/com/starrocks/sql/optimizer/QueryOptimizer.java

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -125,6 +125,7 @@
125125
import com.starrocks.sql.optimizer.rule.tree.PruneSubfieldsForComplexType;
126126
import com.starrocks.sql.optimizer.rule.tree.PushDownAggregateRule;
127127
import com.starrocks.sql.optimizer.rule.tree.PushDownDistinctAggregateRule;
128+
import com.starrocks.sql.optimizer.rule.tree.PushDownNonGroupedAggregateBelowUnion;
128129
import com.starrocks.sql.optimizer.rule.tree.RemoveUselessScanOutputPropertyRule;
129130
import com.starrocks.sql.optimizer.rule.tree.SemiJoinDeduplicateRule;
130131
import com.starrocks.sql.optimizer.rule.tree.SimplifyCaseWhenPredicateRule;
@@ -1038,6 +1039,8 @@ private OptExpression physicalRuleRewrite(ConnectContext connectContext, TaskCon
10381039
result = new ExtractAggregateColumn().rewrite(result, rootTaskContext);
10391040
result = new JoinLocalShuffleRule().rewrite(result, rootTaskContext);
10401041

1042+
result = new PushDownNonGroupedAggregateBelowUnion().rewrite(result, rootTaskContext);
1043+
10411044
// This must be put at last of the optimization. Because wrapping reused ColumnRefOperator with CloneOperator
10421045
// too early will prevent it from certain optimizations that depend on the equivalence of the ColumnRefOperator.
10431046
result = new CloneDuplicateColRefRule().rewrite(result, rootTaskContext);

fe/fe-core/src/main/java/com/starrocks/sql/optimizer/rule/tree/PushDownNonGroupedAggregateBelowUnion.java

Lines changed: 459 additions & 0 deletions
Large diffs are not rendered by default.
Lines changed: 286 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,286 @@
1+
// Copyright 2021-present StarRocks, Inc. All rights reserved.
2+
//
3+
// Licensed under the Apache License, Version 2.0 (the "License");
4+
// you may not use this file except in compliance with the License.
5+
// You may obtain a copy of the License at
6+
//
7+
// https://www.apache.org/licenses/LICENSE-2.0
8+
//
9+
// Unless required by applicable law or agreed to in writing, software
10+
// distributed under the License is distributed on an "AS IS" BASIS,
11+
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
// See the License for the specific language governing permissions and
13+
// limitations under the License.
14+
15+
package com.starrocks.sql.plan;
16+
17+
import com.starrocks.common.Config;
18+
import com.starrocks.qe.SessionVariableConstants;
19+
import org.apache.commons.lang3.StringUtils;
20+
import org.junit.jupiter.api.Assertions;
21+
import org.junit.jupiter.api.Test;
22+
23+
public class PushDownNonGroupedAggregateBelowUnionTest extends PlanTestBase {
24+
private static final long UNION_INPUT_ROW_COUNT = 100000000L;
25+
private static final String UNION_ALL_INPUT =
26+
"(SELECT v1 AS v FROM t0 UNION ALL SELECT v4 AS v FROM t1)";
27+
28+
private void setUnionInputStatistics(long rowCount) {
29+
setTableStatistics(getOlapTable("t0"), rowCount);
30+
setTableStatistics(getOlapTable("t1"), rowCount);
31+
setTableStatistics(getOlapTable("t2"), rowCount);
32+
}
33+
34+
private String getPlan(String sql, boolean enable) throws Exception {
35+
return getPlan(sql, enable, SessionVariableConstants.AggregationStage.TWO_STAGE.ordinal());
36+
}
37+
38+
private String getPlan(String sql, boolean enable, int aggStage) throws Exception {
39+
int oldAggStage = connectContext.getSessionVariable().getNewPlannerAggStage();
40+
boolean oldEnable = Config.push_down_non_grouped_aggregate_below_union;
41+
try {
42+
connectContext.getSessionVariable().setNewPlanerAggStage(aggStage);
43+
Config.push_down_non_grouped_aggregate_below_union = enable;
44+
setUnionInputStatistics(UNION_INPUT_ROW_COUNT);
45+
return getVerboseExplain(sql);
46+
} finally {
47+
Config.push_down_non_grouped_aggregate_below_union = oldEnable;
48+
setUnionInputStatistics(0);
49+
connectContext.getSessionVariable().setNewPlanerAggStage(oldAggStage);
50+
}
51+
}
52+
53+
private void assertPushDownPlanShape(String plan) {
54+
assertContains(plan,
55+
"AGGREGATE (merge finalize)\n" +
56+
" | aggregate:",
57+
"AGGREGATE (merge serialize)\n" +
58+
" | aggregate:",
59+
"0:UNION\n" +
60+
" | output exprs:\n" +
61+
" | [",
62+
"|----",
63+
"EXCHANGE\n" +
64+
" | distribution type: ROUND_ROBIN",
65+
"AGGREGATE (update serialize)\n" +
66+
" | aggregate:",
67+
"OlapScanNode\n" +
68+
" table:");
69+
Assertions.assertEquals(1, StringUtils.countMatches(plan, "AGGREGATE (merge serialize)"), plan);
70+
Assertions.assertEquals(2, StringUtils.countMatches(plan, "AGGREGATE (update serialize)"), plan);
71+
assertPlanOrder(plan, "AGGREGATE (merge finalize)", "AGGREGATE (merge serialize)",
72+
":UNION", ":EXCHANGE", "AGGREGATE (update serialize)", "OlapScanNode");
73+
74+
int exchangeBelowUnion = plan.indexOf(":EXCHANGE", plan.indexOf(":UNION"));
75+
Assertions.assertTrue(exchangeBelowUnion > 0 &&
76+
exchangeBelowUnion < plan.indexOf("AGGREGATE (update serialize)"), plan);
77+
}
78+
79+
private void assertPlanOrder(String plan, String... fragments) {
80+
int previous = -1;
81+
for (String fragment : fragments) {
82+
int current = plan.indexOf(fragment, previous + 1);
83+
Assertions.assertTrue(current > previous, plan);
84+
previous = current;
85+
}
86+
}
87+
88+
private String assertNotRewritten(String sql) throws Exception {
89+
return assertNotRewritten(sql, SessionVariableConstants.AggregationStage.TWO_STAGE.ordinal());
90+
}
91+
92+
private String assertNotRewritten(String sql, int aggStage) throws Exception {
93+
String disabledPlan = getPlan(sql, false, aggStage);
94+
String enabledPlan = getPlan(sql, true, aggStage);
95+
Assertions.assertEquals(disabledPlan, enabledPlan);
96+
return enabledPlan;
97+
}
98+
99+
@Test
100+
public void testPushDownTypicalStatisticAggregatesBelowUnionAll() throws Exception {
101+
String sql = "SELECT COUNT(*), COUNT(v), SUM(v), MIN(v), MAX(v), NDV(v), " +
102+
"HLL_CARDINALITY(HLL_RAW(v)), BITMAP_COUNT(BITMAP_UNION(TO_BITMAP(v))), " +
103+
"PERCENTILE_APPROX(v, 0.5, 1000) FROM " + UNION_ALL_INPUT + " u";
104+
String plan = getPlan(sql, true);
105+
assertPushDownPlanShape(plan);
106+
assertContains(plan,
107+
"AGGREGATE (merge serialize)\n" +
108+
" | aggregate: count",
109+
"AGGREGATE (update serialize)\n" +
110+
" | aggregate: count",
111+
"sum", "min", "max", "ndv", "hll_raw", "bitmap_agg", "percentile_approx", "0.5", "1000");
112+
}
113+
114+
@Test
115+
public void testPushDownAggregateArgumentsWithSimpleExpressions() throws Exception {
116+
String sql = "SELECT SUM(CHAR_LENGTH(CAST(v AS VARCHAR))), " +
117+
"MAX(LEFT(CAST(v AS VARCHAR), 3)), MIN(LEFT(CAST(v AS VARCHAR), 3)), " +
118+
"SUM(CAST(v AS DECIMAL(18, 2))) FROM " + UNION_ALL_INPUT + " u";
119+
String plan = getPlan(sql, true);
120+
assertPushDownPlanShape(plan);
121+
assertContains(plan,
122+
"AGGREGATE (merge serialize)\n" +
123+
" | aggregate: sum",
124+
"AGGREGATE (update serialize)\n" +
125+
" | aggregate: sum",
126+
"char_length", "left", "cast", "DECIMAL");
127+
}
128+
129+
@Test
130+
public void testPushDownAggregateWithBranchProjection() throws Exception {
131+
String sql = "SELECT SUM(v), AVG(v) " +
132+
"FROM (SELECT v1 + 1 AS v FROM t0 UNION ALL SELECT v4 + 1 AS v FROM t1) u";
133+
String plan = getPlan(sql, true);
134+
assertPushDownPlanShape(plan);
135+
assertContains(plan,
136+
"AGGREGATE (merge serialize)\n" +
137+
" | aggregate: sum",
138+
"AGGREGATE (update serialize)\n" +
139+
" | aggregate: sum",
140+
"avg", "+ 1");
141+
}
142+
143+
@Test
144+
public void testPushDownAggregateWithCommonSubExpression() throws Exception {
145+
String sql = "SELECT SUM((v + 1) * (v + 1)), MAX((v + 1) * (v + 1)) " +
146+
"FROM " + UNION_ALL_INPUT + " u";
147+
String plan = getPlan(sql, true);
148+
assertPushDownPlanShape(plan);
149+
assertContains(plan,
150+
"Project\n" +
151+
" | output columns:",
152+
"common expressions:",
153+
"AGGREGATE (update serialize)\n" +
154+
" | aggregate: sum",
155+
"*", "+ 1");
156+
}
157+
158+
@Test
159+
public void testPushDownAdjacentNonGroupedAggregateBelowUnionPatterns() throws Exception {
160+
String sql = "SELECT SUM(cnt), SUM(total) FROM (" +
161+
"SELECT COUNT(*) AS cnt, SUM(v) AS total FROM " +
162+
"(SELECT v1 AS v FROM t0 UNION ALL SELECT v4 AS v FROM t1) u1 " +
163+
"UNION ALL " +
164+
"SELECT COUNT(*) AS cnt, SUM(v) AS total FROM " +
165+
"(SELECT v7 AS v FROM t2 UNION ALL SELECT v1 AS v FROM t0) u2" +
166+
") outer_u";
167+
168+
String disabledPlan = getPlan(sql, false);
169+
Assertions.assertEquals(0, StringUtils.countMatches(disabledPlan, "AGGREGATE (merge serialize)"),
170+
disabledPlan);
171+
Assertions.assertEquals(3, StringUtils.countMatches(disabledPlan, "AGGREGATE (update serialize)"),
172+
disabledPlan);
173+
174+
String enabledPlan = getPlan(sql, true);
175+
Assertions.assertEquals(3, StringUtils.countMatches(enabledPlan, "AGGREGATE (merge serialize)"),
176+
enabledPlan);
177+
Assertions.assertEquals(6, StringUtils.countMatches(enabledPlan, "AGGREGATE (update serialize)"),
178+
enabledPlan);
179+
180+
assertPlanOrder(enabledPlan, ":UNION", "AGGREGATE (update serialize)");
181+
}
182+
183+
@Test
184+
public void testPushDownMultiArgumentAggregateBelowUnionAll() throws Exception {
185+
String sql = "SELECT INTERSECT_COUNT(TO_BITMAP(id), tag, 10) " +
186+
"FROM (SELECT v1 AS id, v1 AS tag FROM t0 UNION ALL SELECT v4 AS id, v4 AS tag FROM t1) u";
187+
String plan = getPlan(sql, true);
188+
assertPushDownPlanShape(plan);
189+
}
190+
191+
@Test
192+
public void testPushDownDataSketchAggregateFunctions() throws Exception {
193+
String[] sqls = {
194+
"SELECT APPROX_COUNT_DISTINCT(v) FROM " + UNION_ALL_INPUT + " u",
195+
"SELECT APPROX_COUNT_DISTINCT_HLL_SKETCH(v) FROM " + UNION_ALL_INPUT + " u",
196+
"SELECT DS_HLL_COUNT_DISTINCT(v) FROM " + UNION_ALL_INPUT + " u",
197+
"SELECT DS_THETA_COUNT_DISTINCT(v) FROM " + UNION_ALL_INPUT + " u"
198+
};
199+
for (String sql : sqls) {
200+
assertPushDownPlanShape(getPlan(sql, true));
201+
}
202+
}
203+
204+
@Test
205+
public void testDisablePushDownNonGroupedAggregateBelowUnion() throws Exception {
206+
String sql = "SELECT COUNT(*), SUM(CHAR_LENGTH(CAST(v AS VARCHAR))) FROM " + UNION_ALL_INPUT + " u";
207+
String plan = getPlan(sql, false);
208+
assertNotContains(plan, "AGGREGATE (merge serialize)");
209+
assertContains(plan,
210+
"AGGREGATE (merge finalize)\n" +
211+
" | aggregate:",
212+
"AGGREGATE (update serialize)\n" +
213+
" | aggregate:",
214+
"0:UNION\n" +
215+
" | output exprs:");
216+
assertPlanOrder(plan, "AGGREGATE (merge finalize)", "AGGREGATE (update serialize)", ":UNION");
217+
}
218+
219+
@Test
220+
public void testNotPushDownGroupByDistinctAndMultiDistinctAggregates() throws Exception {
221+
String[] sqls = {
222+
"SELECT v, COUNT(*) FROM " + UNION_ALL_INPUT + " u GROUP BY v",
223+
"SELECT COUNT(DISTINCT v) FROM " + UNION_ALL_INPUT + " u",
224+
"SELECT MULTI_DISTINCT_COUNT(v) FROM " + UNION_ALL_INPUT + " u",
225+
"SELECT MULTI_DISTINCT_SUM(v) FROM " + UNION_ALL_INPUT + " u",
226+
"SELECT ARRAY_AGG(DISTINCT v) FROM " + UNION_ALL_INPUT + " u"
227+
};
228+
for (String sql : sqls) {
229+
assertNotRewritten(sql);
230+
}
231+
}
232+
233+
@Test
234+
public void testNotPushDownUnsupportedAggregateFunctions() throws Exception {
235+
String[] sqls = {
236+
"SELECT GROUP_CONCAT(CAST(v AS VARCHAR)) FROM " + UNION_ALL_INPUT + " u",
237+
"SELECT ARRAY_AGG(v) FROM " + UNION_ALL_INPUT + " u",
238+
"SELECT ARRAY_AGG_DISTINCT(v) FROM " + UNION_ALL_INPUT + " u",
239+
"SELECT ARRAY_UNIQUE_AGG(a) FROM " +
240+
"(SELECT [v1] AS a FROM t0 UNION ALL SELECT [v4] AS a FROM t1) u"
241+
};
242+
for (String sql : sqls) {
243+
assertNotRewritten(sql);
244+
}
245+
}
246+
247+
@Test
248+
public void testNotPushDownUnionDistinctOrOneStageAggregate() throws Exception {
249+
assertNotRewritten("SELECT COUNT(*) FROM (SELECT v1 AS v FROM t0 UNION SELECT v4 AS v FROM t1) u");
250+
251+
String sql = "SELECT COUNT(*) FROM " + UNION_ALL_INPUT + " u";
252+
assertNotRewritten(sql, SessionVariableConstants.AggregationStage.ONE_STAGE.ordinal());
253+
}
254+
255+
@Test
256+
public void testNotPushDownWhenAggInputContainsJoinLimitOrWindow() throws Exception {
257+
assertNotRewritten("SELECT SUM(x) FROM (SELECT u.v + t2.v7 AS x FROM " + UNION_ALL_INPUT +
258+
" u JOIN t2 ON u.v = t2.v7) x");
259+
260+
assertNotRewritten("SELECT COUNT(*) FROM (SELECT v FROM " + UNION_ALL_INPUT + " u LIMIT 10) x");
261+
262+
assertNotRewritten("SELECT SUM(rn) FROM " +
263+
"(SELECT v, ROW_NUMBER() OVER (ORDER BY v) rn FROM " + UNION_ALL_INPUT + " u) x");
264+
}
265+
266+
@Test
267+
public void testPushDownWhenFilterHasBeenPushedBelowUnionAll() throws Exception {
268+
String plan = getPlan("SELECT COUNT(*) FROM (SELECT v FROM " + UNION_ALL_INPUT +
269+
" u WHERE rand() > 0.1) x", true);
270+
assertPushDownPlanShape(plan);
271+
assertContains(plan, "Predicates: rand", "AGGREGATE (update serialize)");
272+
}
273+
274+
@Test
275+
public void testNotPushDownCteReuse() throws Exception {
276+
double oldCteReuseRatio = connectContext.getSessionVariable().getCboCTERuseRatio();
277+
try {
278+
connectContext.getSessionVariable().setCboCTERuseRatio(0);
279+
String sql = "WITH c AS " + UNION_ALL_INPUT + " " +
280+
"SELECT COUNT(*) FROM c UNION ALL SELECT SUM(v) FROM c";
281+
assertNotRewritten(sql);
282+
} finally {
283+
connectContext.getSessionVariable().setCboCTERuseRatio(oldCteReuseRatio);
284+
}
285+
}
286+
}

0 commit comments

Comments
 (0)