Skip to main content

uqa_sql/semantics/aggregates/
slots.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! Assign internal scalar slots to projection and HAVING aggregates.
8
9use super::{exprs_match, is_aggregate};
10use crate::plan::AggregateClassifier;
11use crate::{SQLError, ScalarExpr};
12
13pub fn compile_projection_aggregate_slots(
14    context: &dyn AggregateClassifier,
15    expr: &ScalarExpr,
16    relation: crate::ast::InternalRelationId,
17    cursor: &mut usize,
18) -> Result<ScalarExpr, SQLError> {
19    rewrite_selected(expr, &mut |expression| {
20        if !is_aggregate(context, expression) {
21            return Ok(None);
22        }
23        let slot = aggregate_slot(relation, *cursor);
24        *cursor += 1;
25        Ok(Some(slot))
26    })
27}
28
29pub fn compile_having_aggregate_slots(
30    context: &dyn AggregateClassifier,
31    expr: &ScalarExpr,
32    relation: crate::ast::InternalRelationId,
33    aggregate_targets: &[ScalarExpr],
34) -> Result<ScalarExpr, SQLError> {
35    rewrite_selected(expr, &mut |aggregate| {
36        if !is_aggregate(context, aggregate) {
37            return Ok(None);
38        }
39        aggregate_targets
40            .iter()
41            .position(|target| exprs_match(target, aggregate))
42            .map(|index| Some(aggregate_slot(relation, index)))
43            .ok_or_else(|| {
44                SQLError::Unsupported(
45                    "HAVING references an aggregate that is not in the aggregate plan".into(),
46                )
47            })
48    })
49}
50
51/// Group values follow aggregate finalizers in the retained scalar slot namespace. Match a complete grouped expression before descending into its inputs.
52pub fn compile_group_slots(
53    expression: &ScalarExpr,
54    groups: &[ScalarExpr],
55    relation: crate::ast::InternalRelationId,
56    first_slot: usize,
57) -> Result<ScalarExpr, SQLError> {
58    rewrite_selected(expression, &mut |expression| {
59        groups
60            .iter()
61            .position(|group| exprs_match(group, expression))
62            .map(|index| {
63                first_slot
64                    .checked_add(index)
65                    .map(|index| aggregate_slot(relation, index))
66                    .ok_or_else(|| SQLError::Internal("aggregate scalar slot overflow".into()))
67            })
68            .transpose()
69    })
70}
71
72/// An absent grouping-set key is NULL as a complete expression, even when another selected key exposes one of its inputs. Aggregate arguments still read the original input rows.
73pub fn select_grouping_set(
74    context: &dyn AggregateClassifier,
75    statement: &crate::plan::QueryBlockPlan,
76    selected: &[ScalarExpr],
77) -> Result<crate::plan::QueryBlockPlan, SQLError> {
78    let groups = statement
79        .group_by
80        .iter()
81        .chain(statement.grouping_sets.iter().flatten())
82        .collect::<Vec<_>>();
83    let mut active = statement.clone();
84    active.group_by = selected.to_vec();
85    active.grouping_sets.clear();
86    let mut rewrite = |expression: &ScalarExpr| {
87        if is_aggregate(context, expression) {
88            return Ok(Some(expression.clone()));
89        }
90        Ok(groups
91            .iter()
92            .any(|group| exprs_match(group, expression))
93            .then(|| {
94                if selected.iter().any(|group| exprs_match(group, expression)) {
95                    expression.clone()
96                } else {
97                    ScalarExpr::Literal(uqa_core::Value::Null)
98                }
99            }))
100    };
101    for projection in &mut active.projections {
102        projection.expr = rewrite_selected(&projection.expr, &mut rewrite)?;
103    }
104    if let Some(having) = &mut active.having {
105        *having = rewrite_selected(having, &mut rewrite)?;
106    }
107    Ok(active)
108}
109
110pub fn aggregate_slot_index(
111    column: crate::ast::InternalColumnRef,
112    relation: crate::ast::InternalRelationId,
113) -> Option<usize> {
114    (column.relation() == relation).then(|| column.attribute())
115}
116
117fn aggregate_slot(relation: crate::ast::InternalRelationId, index: usize) -> ScalarExpr {
118    ScalarExpr::InternalColumn(relation.column(index))
119}
120
121#[expect(
122    clippy::too_many_lines,
123    reason = "preserves aggregate NULL and type order"
124)]
125fn rewrite_selected(
126    expr: &ScalarExpr,
127    replace: &mut impl FnMut(&ScalarExpr) -> Result<Option<ScalarExpr>, SQLError>,
128) -> Result<ScalarExpr, SQLError> {
129    if let Some(expression) = replace(expr)? {
130        return Ok(expression);
131    }
132    match expr {
133        ScalarExpr::Func {
134            order_syntax,
135            name,
136            binding,
137            args,
138            distinct,
139            order_by,
140            filter,
141        } => Ok(ScalarExpr::Func {
142            order_syntax: *order_syntax,
143            name: name.clone(),
144            binding: binding.clone(),
145            args: args
146                .iter()
147                .map(|arg| rewrite_selected(arg, replace))
148                .collect::<Result<Vec<_>, _>>()?,
149            distinct: *distinct,
150            order_by: order_by.clone(),
151            filter: filter
152                .as_deref()
153                .map(|filter| rewrite_selected(filter, replace).map(Box::new))
154                .transpose()?,
155        }),
156        ScalarExpr::Array(items) => Ok(ScalarExpr::Array(
157            items
158                .iter()
159                .map(|item| rewrite_selected(item, replace))
160                .collect::<Result<Vec<_>, _>>()?,
161        )),
162        ScalarExpr::Row(items) => Ok(ScalarExpr::Row(
163            items
164                .iter()
165                .map(|item| rewrite_selected(item, replace))
166                .collect::<Result<Vec<_>, _>>()?,
167        )),
168        ScalarExpr::CompositeRow {
169            items,
170            binding,
171            bound_type,
172        } => Ok(ScalarExpr::CompositeRow {
173            items: items
174                .iter()
175                .map(|item| rewrite_selected(item, replace))
176                .collect::<Result<Vec<_>, _>>()?,
177            binding: binding.clone(),
178            bound_type: bound_type.clone(),
179        }),
180        ScalarExpr::Binary { op, lhs, rhs } => Ok(ScalarExpr::Binary {
181            op: *op,
182            lhs: Box::new(rewrite_selected(lhs, replace)?),
183            rhs: Box::new(rewrite_selected(rhs, replace)?),
184        }),
185        ScalarExpr::Not(inner) => Ok(ScalarExpr::Not(Box::new(rewrite_selected(inner, replace)?))),
186        ScalarExpr::UnaryMinus(inner) => Ok(ScalarExpr::UnaryMinus(Box::new(rewrite_selected(
187            inner, replace,
188        )?))),
189        ScalarExpr::And(parts) => Ok(ScalarExpr::And(
190            parts
191                .iter()
192                .map(|part| rewrite_selected(part, replace))
193                .collect::<Result<Vec<_>, _>>()?,
194        )),
195        ScalarExpr::Or(parts) => Ok(ScalarExpr::Or(
196            parts
197                .iter()
198                .map(|part| rewrite_selected(part, replace))
199                .collect::<Result<Vec<_>, _>>()?,
200        )),
201        ScalarExpr::IsNull { expr, negated } => Ok(ScalarExpr::IsNull {
202            expr: Box::new(rewrite_selected(expr, replace)?),
203            negated: *negated,
204        }),
205        ScalarExpr::Between { expr, low, high } => Ok(ScalarExpr::Between {
206            expr: Box::new(rewrite_selected(expr, replace)?),
207            low: Box::new(rewrite_selected(low, replace)?),
208            high: Box::new(rewrite_selected(high, replace)?),
209        }),
210        ScalarExpr::InList {
211            expr,
212            list,
213            negated,
214        } => Ok(ScalarExpr::InList {
215            expr: Box::new(rewrite_selected(expr, replace)?),
216            list: list
217                .iter()
218                .map(|item| rewrite_selected(item, replace))
219                .collect::<Result<Vec<_>, _>>()?,
220            negated: *negated,
221        }),
222        ScalarExpr::Case {
223            base,
224            when,
225            else_branch,
226        } => Ok(ScalarExpr::Case {
227            base: base
228                .as_deref()
229                .map(|base| rewrite_selected(base, replace).map(Box::new))
230                .transpose()?,
231            when: when
232                .iter()
233                .map(|(condition, result)| {
234                    Ok((
235                        rewrite_selected(condition, replace)?,
236                        rewrite_selected(result, replace)?,
237                    ))
238                })
239                .collect::<Result<Vec<_>, SQLError>>()?,
240            else_branch: else_branch
241                .as_deref()
242                .map(|branch| rewrite_selected(branch, replace).map(Box::new))
243                .transpose()?,
244        }),
245        ScalarExpr::Cast { expr, ty, implicit } => Ok(ScalarExpr::Cast {
246            implicit: *implicit,
247            expr: Box::new(rewrite_selected(expr, replace)?),
248            ty: ty.clone(),
249        }),
250        ScalarExpr::InSubquery {
251            expr,
252            subquery,
253            negated,
254        } => Ok(ScalarExpr::InSubquery {
255            expr: Box::new(rewrite_selected(expr, replace)?),
256            subquery: *subquery,
257            negated: *negated,
258        }),
259        other => Ok(other.clone()),
260    }
261}