Skip to main content

uqa_sql/semantics/
aggregates.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Aggregate discovery and expression equivalence for SQL grouping.
8
9use crate::plan::{AggregateClassifier, ProjectionPlan, QueryBlockPlan};
10use crate::ScalarExpr;
11use uqa_core::Value;
12
13#[cfg(test)]
14mod tests;
15
16pub fn exprs_match(lhs: &ScalarExpr, rhs: &ScalarExpr) -> bool {
17    match (lhs, rhs) {
18        (ScalarExpr::Star, ScalarExpr::Star) => true,
19        (ScalarExpr::Column(a), ScalarExpr::Column(b)) => a == b,
20        (
21            ScalarExpr::QualifiedColumn {
22                qualifier: aq,
23                column: ac,
24                ..
25            },
26            ScalarExpr::QualifiedColumn {
27                qualifier: bq,
28                column: bc,
29                ..
30            },
31        ) => aq == bq && ac == bc,
32        (ScalarExpr::Column(c), ScalarExpr::QualifiedColumn { column, .. })
33        | (ScalarExpr::QualifiedColumn { column, .. }, ScalarExpr::Column(c)) => c == column,
34        (ScalarExpr::Literal(a), ScalarExpr::Literal(b)) => literals_equal(a, b),
35        (
36            ScalarExpr::TypedLiteral {
37                value: av,
38                ty: at,
39                bound_type: ab,
40                parameter_index: ap,
41            },
42            ScalarExpr::TypedLiteral {
43                value: bv,
44                ty: bt,
45                bound_type: bb,
46                parameter_index: bp,
47            },
48        ) => at == bt && ab == bb && ap == bp && literals_equal(av, bv),
49        (ScalarExpr::Param(a), ScalarExpr::Param(b)) => a == b,
50        (ScalarExpr::Position(a), ScalarExpr::Position(b)) => a == b,
51        (ScalarExpr::InternalColumn(a), ScalarExpr::InternalColumn(b)) => a == b,
52        (
53            ScalarExpr::Func {
54                order_syntax: aos,
55                name: an,
56                binding: ab,
57                args: aa,
58                distinct: ad,
59                order_by: ao,
60                filter: af,
61            },
62            ScalarExpr::Func {
63                order_syntax: bos,
64                name: bn,
65                binding: bb,
66                args: ba,
67                distinct: bd,
68                order_by: bo,
69                filter: bf,
70            },
71        ) => {
72            an.eq_ignore_ascii_case(bn)
73                && ab == bb
74                && ad == bd
75                && aos == bos
76                && aa.len() == ba.len()
77                && aa.iter().zip(ba.iter()).all(|(x, y)| exprs_match(x, y))
78                && ao.len() == bo.len()
79                && ao.iter().zip(bo.iter()).all(|(x, y)| {
80                    x.descending == y.descending
81                        && x.nulls == y.nulls
82                        && exprs_match(&x.expr, &y.expr)
83                })
84                && match (af.as_deref(), bf.as_deref()) {
85                    (None, None) => true,
86                    (Some(x), Some(y)) => exprs_match(x, y),
87                    _ => false,
88                }
89        }
90        (
91            ScalarExpr::Binary {
92                op: ao,
93                lhs: al,
94                rhs: ar,
95            },
96            ScalarExpr::Binary {
97                op: bo,
98                lhs: bl,
99                rhs: br,
100            },
101        ) => ao == bo && exprs_match(al, bl) && exprs_match(ar, br),
102        (ScalarExpr::And(a), ScalarExpr::And(b)) | (ScalarExpr::Or(a), ScalarExpr::Or(b)) => {
103            a.len() == b.len() && a.iter().zip(b.iter()).all(|(x, y)| exprs_match(x, y))
104        }
105        (ScalarExpr::Not(a), ScalarExpr::Not(b))
106        | (ScalarExpr::UnaryMinus(a), ScalarExpr::UnaryMinus(b)) => exprs_match(a, b),
107        (
108            ScalarExpr::Cast {
109                expr: a, ty: at, ..
110            },
111            ScalarExpr::Cast {
112                expr: b, ty: bt, ..
113            },
114        ) => at == bt && exprs_match(a, b),
115        _ => false,
116    }
117}
118
119pub fn literals_equal(a: &Value, b: &Value) -> bool {
120    a.has_same_representation(b)
121}
122
123pub fn has_aggregate(aggregates: &dyn AggregateClassifier, projections: &[ProjectionPlan]) -> bool {
124    projections
125        .iter()
126        .any(|p| contains_aggregate(aggregates, &p.expr))
127}
128
129pub fn is_aggregate(aggregates: &dyn AggregateClassifier, expr: &ScalarExpr) -> bool {
130    is_builtin_aggregate(expr)
131        || matches!(expr, ScalarExpr::Func { name, .. } if aggregates.is_registered_aggregate(name))
132}
133
134pub use super::is_builtin_aggregate;
135
136pub fn aggregate_exprs<'a>(
137    aggregates: &dyn AggregateClassifier,
138    projections: &'a [ProjectionPlan],
139) -> Vec<&'a ScalarExpr> {
140    let mut out = Vec::new();
141    for projection in projections {
142        collect_aggregate_exprs(aggregates, &projection.expr, &mut out);
143    }
144    out
145}
146
147/// Aggregate states needed by a query block, in accumulator order.
148///
149/// Projection aggregates remain first (including repeated expressions) because
150/// projection rewriting consumes them positionally. Aggregates referenced only
151/// by HAVING are appended once as hidden targets, matching `PostgreSQL`'s rule
152/// that HAVING need not expose an aggregate in the SELECT list.
153pub fn aggregate_targets<'a>(
154    aggregates: &dyn AggregateClassifier,
155    statement: &'a QueryBlockPlan,
156) -> Vec<&'a ScalarExpr> {
157    let mut targets = aggregate_exprs(aggregates, &statement.projections);
158    if let Some(having) = statement.having.as_ref() {
159        let mut hidden = Vec::new();
160        collect_aggregate_exprs(aggregates, having, &mut hidden);
161        for aggregate in hidden {
162            if !targets
163                .iter()
164                .any(|existing| exprs_match(existing, aggregate))
165            {
166                targets.push(aggregate);
167            }
168        }
169    }
170    targets
171}
172
173pub fn collect_aggregate_exprs<'a>(
174    aggregates: &dyn AggregateClassifier,
175    expr: &'a ScalarExpr,
176    out: &mut Vec<&'a ScalarExpr>,
177) {
178    if is_aggregate(aggregates, expr) {
179        out.push(expr);
180        return;
181    }
182    match expr {
183        ScalarExpr::Func { args, filter, .. } => {
184            for arg in args {
185                collect_aggregate_exprs(aggregates, arg, out);
186            }
187            if let Some(filter) = filter.as_deref() {
188                collect_aggregate_exprs(aggregates, filter, out);
189            }
190        }
191        ScalarExpr::Array(items)
192        | ScalarExpr::Row(items)
193        | ScalarExpr::CompositeRow { items, .. }
194        | ScalarExpr::And(items)
195        | ScalarExpr::Or(items) => {
196            for item in items {
197                collect_aggregate_exprs(aggregates, item, out);
198            }
199        }
200        ScalarExpr::Binary { lhs, rhs, .. } => {
201            collect_aggregate_exprs(aggregates, lhs, out);
202            collect_aggregate_exprs(aggregates, rhs, out);
203        }
204        ScalarExpr::Not(inner)
205        | ScalarExpr::UnaryMinus(inner)
206        | ScalarExpr::Cast { expr: inner, .. } => {
207            collect_aggregate_exprs(aggregates, inner, out);
208        }
209        ScalarExpr::IsNull { expr, .. } => collect_aggregate_exprs(aggregates, expr, out),
210        ScalarExpr::Between { expr, low, high } => {
211            collect_aggregate_exprs(aggregates, expr, out);
212            collect_aggregate_exprs(aggregates, low, out);
213            collect_aggregate_exprs(aggregates, high, out);
214        }
215        ScalarExpr::InList { expr, list, .. } => {
216            collect_aggregate_exprs(aggregates, expr, out);
217            for item in list {
218                collect_aggregate_exprs(aggregates, item, out);
219            }
220        }
221        ScalarExpr::Case {
222            base,
223            when,
224            else_branch,
225        } => {
226            if let Some(base) = base.as_deref() {
227                collect_aggregate_exprs(aggregates, base, out);
228            }
229            for (condition, result) in when {
230                collect_aggregate_exprs(aggregates, condition, out);
231                collect_aggregate_exprs(aggregates, result, out);
232            }
233            if let Some(else_branch) = else_branch.as_deref() {
234                collect_aggregate_exprs(aggregates, else_branch, out);
235            }
236        }
237        ScalarExpr::InSubquery { expr, .. } => collect_aggregate_exprs(aggregates, expr, out),
238        ScalarExpr::Default
239        | ScalarExpr::Star
240        | ScalarExpr::QualifiedStar(_)
241        | ScalarExpr::Column(_)
242        | ScalarExpr::Position(_)
243        | ScalarExpr::InternalColumn(_)
244        | ScalarExpr::QualifiedColumn { .. }
245        | ScalarExpr::Literal(_)
246        | ScalarExpr::TypedLiteral { .. }
247        | ScalarExpr::Param(_)
248        | ScalarExpr::WindowCall { .. }
249        | ScalarExpr::ScalarSubquery(_)
250        | ScalarExpr::Exists { .. } => {}
251    }
252}
253
254pub fn contains_aggregate(aggregates: &dyn AggregateClassifier, expr: &ScalarExpr) -> bool {
255    let mut found = Vec::new();
256    collect_aggregate_exprs(aggregates, expr, &mut found);
257    !found.is_empty()
258}
259
260/// Collect the top-level column names an expression reads. Returns
261/// `false` when the expression can reach arbitrary fields (`*`,
262/// subqueries, window calls), in which case callers must materialise
263/// whole documents.
264pub fn expr_references_columns(expr: &ScalarExpr) -> bool {
265    match expr {
266        ScalarExpr::Star
267        | ScalarExpr::QualifiedStar(_)
268        | ScalarExpr::Column(_)
269        | ScalarExpr::Position(_)
270        | ScalarExpr::InternalColumn(_)
271        | ScalarExpr::QualifiedColumn { .. } => true,
272        ScalarExpr::Func { args, filter, .. } => {
273            args.iter().any(expr_references_columns)
274                || filter.as_deref().is_some_and(expr_references_columns)
275        }
276        ScalarExpr::Array(items)
277        | ScalarExpr::Row(items)
278        | ScalarExpr::CompositeRow { items, .. }
279        | ScalarExpr::And(items)
280        | ScalarExpr::Or(items) => items.iter().any(expr_references_columns),
281        ScalarExpr::Binary { lhs, rhs, .. } => {
282            expr_references_columns(lhs) || expr_references_columns(rhs)
283        }
284        ScalarExpr::Not(inner)
285        | ScalarExpr::UnaryMinus(inner)
286        | ScalarExpr::Cast { expr: inner, .. } => expr_references_columns(inner),
287        ScalarExpr::IsNull { expr, .. } => expr_references_columns(expr),
288        ScalarExpr::Between { expr, low, high } => {
289            expr_references_columns(expr)
290                || expr_references_columns(low)
291                || expr_references_columns(high)
292        }
293        ScalarExpr::InList { expr, list, .. } => {
294            expr_references_columns(expr) || list.iter().any(expr_references_columns)
295        }
296        ScalarExpr::WindowCall { args, filter, .. } => {
297            args.iter().any(expr_references_columns)
298                || filter.as_deref().is_some_and(expr_references_columns)
299        }
300        ScalarExpr::Case {
301            base,
302            when,
303            else_branch,
304        } => {
305            base.as_deref().is_some_and(expr_references_columns)
306                || when.iter().any(|(condition, result)| {
307                    expr_references_columns(condition) || expr_references_columns(result)
308                })
309                || else_branch.as_deref().is_some_and(expr_references_columns)
310        }
311        ScalarExpr::InSubquery { expr, .. } => expr_references_columns(expr),
312        ScalarExpr::ScalarSubquery(_) | ScalarExpr::Exists { .. } => true,
313        ScalarExpr::Default
314        | ScalarExpr::Literal(_)
315        | ScalarExpr::TypedLiteral { .. }
316        | ScalarExpr::Param(_) => false,
317    }
318}
319
320mod slots;
321pub use slots::{
322    aggregate_slot_index, compile_group_slots, compile_having_aggregate_slots,
323    compile_projection_aggregate_slots, select_grouping_set,
324};