From def0c2ec5ca0992083d955f6bcbc5ffaf914fe90 Mon Sep 17 00:00:00 2001 From: dennis zhuang Date: Thu, 24 Sep 2026 06:49:52 +0000 Subject: [PATCH] fix: only push down aggregates grouped by the partition columns themselves (#9337) * fix: only push down aggregates grouped by the partition columns themselves Signed-off-by: Dennis Zhuang * fix: keep grouping sets on the frontend for partitioned tables Signed-off-by: Dennis Zhuang --------- Signed-off-by: Dennis Zhuang --- src/query/src/dist_plan/analyzer/test.rs | 35 ++++++++++++++ src/query/src/dist_plan/commutativity.rs | 37 +++++++++----- .../common/aggregate/multi_regions.result | 48 +++++++++++++++++++ .../common/aggregate/multi_regions.sql | 16 +++++++ 4 files changed, 125 insertions(+), 11 deletions(-) diff --git a/src/query/src/dist_plan/analyzer/test.rs b/src/query/src/dist_plan/analyzer/test.rs index 5171688993c..694f2bcde59 100644 --- a/src/query/src/dist_plan/analyzer/test.rs +++ b/src/query/src/dist_plan/analyzer/test.rs @@ -40,6 +40,7 @@ use datafusion_expr::{ }; use datafusion_functions::datetime::date_bin; use datafusion_functions::datetime::expr_fn::now; +use datafusion_functions::unicode::expr_fn::substring; use datatypes::data_type::ConcreteDataType; use datatypes::schema::{ColumnSchema, SchemaBuilder, SchemaRef}; use futures::Stream; @@ -1001,6 +1002,40 @@ fn expand_proj_alias_aliased_part_col_aggr() { assert_eq!(expected, result.to_string()); } +/// `substr(pk1, 1, 1)` maps rows from different partitions to the same group, so the +/// aggregate has to be merged on the frontend even though it references `pk1`. +#[test] +fn expand_expr_over_part_col_aggr() { + init_default_ut_logging(); + let test_table = TestTable::table_with_name(0, "t".to_string()); + let table_source = Arc::new(DefaultTableSource::new(Arc::new( + DfTableProviderAdapter::new(test_table), + ))); + let plan = LogicalPlanBuilder::scan_with_filters("t", table_source, None, vec![]) + .unwrap() + .aggregate( + vec![substring(col("pk1"), lit(1i64), lit(1i64)), col("pk2")], + vec![min(col("number"))], + ) + .unwrap() + .build() + .unwrap(); + + let config = ConfigOptions::default(); + let result = DistPlannerAnalyzer {}.analyze(plan, &config).unwrap(); + + let expected = [ + "Projection: substr(t.pk1,Int64(1),Int64(1)), t.pk2, min(t.number)", + " Aggregate: groupBy=[[substr(t.pk1,Int64(1),Int64(1)), t.pk2]], aggr=[[__min_merge(__min_state(t.number)) AS min(t.number)]]", + " MergeScan [is_placeholder=false, remote_input=[", + "Aggregate: groupBy=[[substr(t.pk1, Int64(1), Int64(1)), t.pk2]], aggr=[[__min_state(t.number)]]", + " TableScan: t", + "]]", + ] + .join("\n"); + assert_eq!(expected, result.to_string()); +} + /// notice that step aggr then part col aggr seems impossible as the partition columns for part col aggr /// can't pass through the step aggr without making step aggr also a part col aggr /// so here only test part col aggr -> step aggr case diff --git a/src/query/src/dist_plan/commutativity.rs b/src/query/src/dist_plan/commutativity.rs index a236118822b..d02b746a8bf 100644 --- a/src/query/src/dist_plan/commutativity.rs +++ b/src/query/src/dist_plan/commutativity.rs @@ -117,7 +117,14 @@ impl Categorizer { LogicalPlan::Filter(filter) => Self::check_expr(&filter.predicate), LogicalPlan::Window(_) => Commutativity::Unimplemented, LogicalPlan::Aggregate(aggr) => { - let is_all_steppable = is_all_aggr_exprs_steppable(&aggr.aggr_expr); + // The state/merge split maps each group expression to one output column, + // which doesn't hold for grouping sets. + let has_grouping_set = aggr + .group_expr + .iter() + .any(|expr| matches!(expr, Expr::GroupingSet(_))); + let is_all_steppable = + !has_grouping_set && is_all_aggr_exprs_steppable(&aggr.aggr_expr); let matches_partition = Self::check_partition(&aggr.group_expr, &partition_cols); if !matches_partition && is_all_steppable { debug!("Plan is steppable: {plan}"); @@ -303,26 +310,34 @@ impl Categorizer { /// Return true if the given expr and partition cols satisfied the rule. /// In this case the plan can be treated as fully commutative. /// - /// So only if all partition columns show up in `exprs`, return true. + /// So only if every partition column is itself one of `exprs`, return true. /// Otherwise return false. /// + /// An expression that only references a partition column, like `substr(host, 3, 1)` + /// or `k % 2`, doesn't count: it can put rows from different partitions into the same + /// group. fn check_partition(exprs: &[Expr], partition_cols: &AliasMapping) -> bool { - let mut ref_cols = HashSet::new(); - for expr in exprs { - expr.add_column_refs(&mut ref_cols); - } - let ref_cols = ref_cols - .into_iter() - .map(|c| c.name.clone()) + let group_cols = exprs + .iter() + .filter_map(|expr| { + let mut expr = expr; + while let Expr::Alias(alias) = expr { + expr = &alias.expr; + } + match expr { + Expr::Column(column) => Some(column.name.clone()), + _ => None, + } + }) .collect::>(); for all_alias in partition_cols.values() { let all_alias = all_alias .iter() .map(|c| c.name.clone()) .collect::>(); - // check if ref columns intersect with all alias of partition columns + // check if group columns intersect with all alias of partition columns // is empty, if it's empty, not all partition columns show up in `exprs` - if ref_cols.intersection(&all_alias).count() == 0 { + if group_cols.intersection(&all_alias).count() == 0 { return false; } } diff --git a/tests/cases/standalone/common/aggregate/multi_regions.result b/tests/cases/standalone/common/aggregate/multi_regions.result index fa0bf929319..79ae85fbc04 100644 --- a/tests/cases/standalone/common/aggregate/multi_regions.result +++ b/tests/cases/standalone/common/aggregate/multi_regions.result @@ -113,6 +113,54 @@ select sum(val) from t group by idc; |_|_| Total rows: 0_| +-+-+-+ +insert into t values + (1000, 1, '1000-x', 'a'), + (1000, 2, '1000-y', 'a'), + (1000, 3, '2000-x', 'a'), + (1000, 4, '2000-y', 'a'); + +Affected Rows: 4 + +-- group keys derived from the partition column span regions +select substr(host, 6, 1) as g, count(*), sum(val) from t group by g order by g; + ++---+----------+------------+ +| g | count(*) | sum(t.val) | ++---+----------+------------+ +| x | 2 | 4.0 | +| y | 2 | 6.0 | ++---+----------+------------+ + +-- grouping sets are computed on the frontend, only the PostgreSQL dialect parses them +-- SQLNESS PROTOCOL POSTGRES +select host, idc, sum(val) from t group by grouping sets ((host, idc), (host)) order by host, idc; + ++--------+-----+------------+ +| host | idc | sum(t.val) | ++--------+-----+------------+ +| 1000-x | a | 1.0 | +| 1000-x | | 1.0 | +| 1000-y | a | 2.0 | +| 1000-y | | 2.0 | +| 2000-x | a | 3.0 | +| 2000-x | | 3.0 | +| 2000-y | a | 4.0 | +| 2000-y | | 4.0 | ++--------+-----+------------+ + +-- SQLNESS PROTOCOL POSTGRES +select host, sum(val) from t group by rollup(host) order by host; + ++--------+------------+ +| host | sum(t.val) | ++--------+------------+ +| 1000-x | 1.0 | +| 1000-y | 2.0 | +| 2000-x | 3.0 | +| 2000-y | 4.0 | +| | 10.0 | ++--------+------------+ + drop table t; Affected Rows: 0 diff --git a/tests/cases/standalone/common/aggregate/multi_regions.sql b/tests/cases/standalone/common/aggregate/multi_regions.sql index 3396bdddf19..4f99e1ed465 100644 --- a/tests/cases/standalone/common/aggregate/multi_regions.sql +++ b/tests/cases/standalone/common/aggregate/multi_regions.sql @@ -45,4 +45,20 @@ select sum(val) from t; explain analyze select sum(val) from t group by idc; +insert into t values + (1000, 1, '1000-x', 'a'), + (1000, 2, '1000-y', 'a'), + (1000, 3, '2000-x', 'a'), + (1000, 4, '2000-y', 'a'); + +-- group keys derived from the partition column span regions +select substr(host, 6, 1) as g, count(*), sum(val) from t group by g order by g; + +-- grouping sets are computed on the frontend, only the PostgreSQL dialect parses them +-- SQLNESS PROTOCOL POSTGRES +select host, idc, sum(val) from t group by grouping sets ((host, idc), (host)) order by host, idc; + +-- SQLNESS PROTOCOL POSTGRES +select host, sum(val) from t group by rollup(host) order by host; + drop table t;