jayzhan211 commented on code in PR #25781:
URL: https://github.com/apache/datafusion/pull/25781#discussion_r4225398313
##########
datafusion/optimizer/src/decorrelate.rs:
##########
@@ -828,6 +912,45 @@ fn can_pullup_over_aggregation(expr: &Expr) -> bool {
}
}
+/// Whether `expr` computes `GROUPING`/`GROUPING_ID`, which reads whether a
+/// row's grouping set groups by a particular column.
+fn is_grouping_call(expr: &Expr) -> bool {
+ matches!(expr, Expr::AggregateFunction(agg) if
agg.func.name().eq_ignore_ascii_case("grouping"))
+}
+
+/// Columns read by a node strictly above the first `Aggregate` found in
+/// `plan`, which is the subquery's original, not yet rewritten, plan.
+///
+/// [`PullUpCorrelatedExpr::f_up`] runs bottom-up, so by the time it visits an
+/// `Aggregate` it cannot yet tell whether a `HAVING` or `Projection` above it
+/// reads one of the columns a grouping set pull up would add. This walks the
+/// plan top-down instead, before the rewrite starts, and stops at the first
+/// `Aggregate` along each branch, collecting the columns every node above it
+/// reads in its own expressions. A nested `Subquery` is a different
+/// correlation scope and is skipped, the same way
+/// [`PullUpCorrelatedExpr::f_down`] skips it.
+fn columns_read_above_aggregate(plan: &LogicalPlan) -> BTreeSet<Column> {
+ fn walk(plan: &LogicalPlan, above: &mut BTreeSet<Column>) -> bool {
+ if matches!(plan, LogicalPlan::Subquery(_)) {
+ return false;
+ }
+ if matches!(plan, LogicalPlan::Aggregate(_)) {
Review Comment:
`walk` returns at the first `Aggregate` top-down, and `any` skips later
inputs once one holds an aggregate. So reads of the NULL-filled column above a
stacked aggregate, in a later join input, or between nested grouping sets are
never collected. The result is `SafeToExtend` and wrong rows; on main these
queries fail to plan instead.
Repro (tables from `subquery.slt`):
```sql
-- expected (1,1) (2,1) (4,0) (5,1) (NULL,0); returns 0 for every row
SELECT gs_outer.k, (SELECT count(*) FROM (SELECT gs_inner.k AS kk FROM
gs_inner WHERE gs_inner.k = gs_outer.k GROUP BY GROUPING SETS ((gs_inner.k),
(gs_inner.j))) t WHERE t.kk IS NULL) FROM gs_outer ORDER BY gs_outer.k;
-- expected true for 1, 2, 5; returns false for every row
SELECT gs_outer.k, EXISTS (SELECT 1 FROM (SELECT count(*) AS c FROM
gs_inner) a JOIN (SELECT gs_inner.k FROM gs_inner WHERE gs_inner.k = gs_outer.k
GROUP BY GROUPING SETS ((gs_inner.k), (gs_inner.j))) b ON a.c > 0 WHERE b.k IS
NULL) FROM gs_outer ORDER BY gs_outer.k;
-- expected true for 1, 2, 5; returns false for every row
SELECT gs_outer.k, EXISTS (SELECT 1 FROM (SELECT gs_inner.k AS kk,
gs_inner.j AS jj FROM gs_inner WHERE gs_inner.k = gs_outer.k GROUP BY GROUPING
SETS ((gs_inner.k), (gs_inner.j))) t GROUP BY GROUPING SETS ((t.kk), (t.kk,
t.jj)) HAVING t.kk IS NULL) FROM gs_outer ORDER BY gs_outer.k;
```
Fix: walk every input and through every aggregate, and count a node as
"above" when any grouping-set aggregate is below it. With this, the three
queries stay correlated, the `((k), (j))` EXISTS case still decorrelates, and
the full slt suite passes; please add the three as `statement error` cases.
```diff
-/// Columns read by a node strictly above the first `Aggregate` found in
-/// `plan`, which is the subquery's original, not yet rewritten, plan.
+/// Columns read by a node that has a grouping set `Aggregate` below it in
+/// `plan`, which is the subquery's original, not yet rewritten, plan.
@@
fn walk(plan: &LogicalPlan, above: &mut BTreeSet<Column>) -> bool {
if matches!(plan, LogicalPlan::Subquery(_)) {
return false;
}
- if matches!(plan, LogicalPlan::Aggregate(_)) {
- return true;
+ // Visit every input, and walk through every aggregate: a node above
+ // a grouping set may sit in a later input, or above another
aggregate.
+ let mut found_below = false;
+ for child in plan.inputs() {
+ found_below |= walk(child, above);
}
- let found_below = plan.inputs().into_iter().any(|child| walk(child,
above));
if found_below {
for expr in plan.expressions() {
above.extend(expr.column_refs().into_iter().cloned());
}
}
found_below
+ || matches!(plan, LogicalPlan::Aggregate(aggregate) if aggregate
+ .group_expr
+ .iter()
+ .any(|expr| matches!(expr, Expr::GroupingSet(_))))
}
```
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]