uqa_sql/semantics/aggregates/
slots.rs1use 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
51pub 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
72pub 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}