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 <killme2008@gmail.com>

* fix: keep grouping sets on the frontend for partitioned tables

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>

---------

Signed-off-by: Dennis Zhuang <killme2008@gmail.com>
This commit is contained in:
dennis zhuang
2026-09-24 06:49:52 +00:00
committed by GitHub
parent b9f991502c
commit def0c2ec5c
4 changed files with 125 additions and 11 deletions
+35
View File
@@ -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
+26 -11
View File
@@ -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::<HashSet<_>>();
for all_alias in partition_cols.values() {
let all_alias = all_alias
.iter()
.map(|c| c.name.clone())
.collect::<HashSet<_>>();
// 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;
}
}
@@ -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
@@ -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;