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]

Reply via email to