Skip to main content

uqa_sql/semantics/
grouping_sets.rs

1//
2// Unified Query Algebra
3//
4// Copyright (c) 2023-2026 Cognica, Inc.
5//
6
7//! PostgreSQL-compatible grouping-set preparation after input schema binding.
8
9use std::collections::HashSet;
10
11use crate::ast::ColumnType;
12use crate::{RowSchema, ScalarExpr, ScalarFrameBound};
13use uqa_core::Value;
14
15use crate::{plan::QueryBlockPlan, FunctionTypeResolver, SQLError, SQLParam};
16
17mod expressions;
18mod names;
19pub(crate) use names::resolve_grouping_expression_reference;
20pub use names::{bind_grouping_names, resolve_grouping_expression};
21mod validation;
22pub use validation::validate_grouped_expressions;
23#[cfg(test)]
24mod tests;
25
26/// Resolve grouping names against the input before expanding output aliases and deduplicating `GROUP BY DISTINCT` sets. Equivalent qualified and unqualified columns share one identity without conflating expressions that `PostgreSQL` resolves to different operator inputs.
27pub fn prepare_grouping_sets(
28    engine: &dyn crate::routines::RoutineResolution,
29    statement: &QueryBlockPlan,
30    schema: &RowSchema,
31    params: &[SQLParam],
32) -> Result<Option<QueryBlockPlan>, SQLError> {
33    if statement.group_by.is_empty() && statement.grouping_sets.is_empty() {
34        return Ok(None);
35    }
36
37    let mut prepared = statement.clone();
38    // A block's operator holds the block's own columns; a lateral block binds its output names against the scope of its outer row before it builds the operator.
39    let mut changed = bind_grouping_names(engine, &mut prepared, schema, None, params)?;
40    changed |= expressions::bind_grouping_expressions(engine, &mut prepared, schema, params)?;
41    if !prepared.group_distinct {
42        return Ok(changed.then_some(prepared));
43    }
44    prepared.group_distinct = false;
45    let mut seen = HashSet::with_capacity(prepared.grouping_sets.len());
46    let mut distinct = Vec::with_capacity(prepared.grouping_sets.len());
47    for grouping_set in std::mem::take(&mut prepared.grouping_sets) {
48        let identity = grouping_set_identity(engine, &grouping_set, schema, params)?;
49        if seen.insert(identity) {
50            distinct.push(grouping_set);
51        }
52    }
53    prepared.grouping_sets = distinct;
54    Ok(Some(prepared))
55}
56
57fn grouping_set_identity(
58    engine: &dyn FunctionTypeResolver,
59    grouping_set: &[ScalarExpr],
60    schema: &RowSchema,
61    params: &[SQLParam],
62) -> Result<Vec<Vec<u8>>, SQLError> {
63    let mut identity = grouping_set
64        .iter()
65        .map(|expression| expression_identity(engine, expression, schema, params))
66        .collect::<Result<Vec<_>, _>>()?;
67    identity.sort_unstable();
68    identity.dedup();
69    Ok(identity)
70}
71
72fn expression_identity(
73    engine: &dyn FunctionTypeResolver,
74    expression: &ScalarExpr,
75    schema: &RowSchema,
76    params: &[SQLParam],
77) -> Result<Vec<u8>, SQLError> {
78    let expression =
79        crate::bind_type_introspection_with_resolver(expression.clone(), schema, params, engine);
80    let expression = normalize_expression(engine, expression, schema, params)?;
81    serde_json::to_vec(&expression).map_err(|error| {
82        SQLError::Internal(format!(
83            "serialize GROUP BY DISTINCT expression identity: {error}"
84        ))
85    })
86}
87
88#[expect(
89    clippy::too_many_lines,
90    reason = "preserves SELECT schema and row identity"
91)]
92fn normalize_expression(
93    engine: &dyn FunctionTypeResolver,
94    expression: ScalarExpr,
95    schema: &RowSchema,
96    params: &[SQLParam],
97) -> Result<ScalarExpr, SQLError> {
98    Ok(match expression {
99        ScalarExpr::Column(column) => schema
100            .unqualified_position(&column)
101            .map_or(ScalarExpr::Column(column), ScalarExpr::Position),
102        ScalarExpr::QualifiedColumn { qualifier, column } => {
103            schema.qualified_position(&qualifier, &column).map_or(
104                ScalarExpr::QualifiedColumn { qualifier, column },
105                ScalarExpr::Position,
106            )
107        }
108        ScalarExpr::Func {
109            order_syntax,
110            name,
111            binding,
112            args,
113            distinct,
114            order_by,
115            filter,
116        } => {
117            let name = canonical_function_name(name);
118            let argument_types = args
119                .iter()
120                .map(|argument| expression_type(engine, argument, schema, params))
121                .collect::<Result<Vec<_>, _>>()?;
122            let targets = crate::builtin_function_argument_targets(&name, &argument_types);
123            ScalarExpr::Func {
124                order_syntax,
125                name,
126                binding,
127                args: args
128                    .into_iter()
129                    .zip(targets)
130                    .map(|(argument, target)| {
131                        normalize_unknown_literal(engine, argument, target.as_ref(), schema, params)
132                    })
133                    .collect::<Result<Vec<_>, _>>()?,
134                distinct,
135                order_by: order_by
136                    .into_iter()
137                    .map(|mut order| {
138                        order.expr = normalize_expression(engine, order.expr, schema, params)?;
139                        Ok(order)
140                    })
141                    .collect::<Result<Vec<_>, SQLError>>()?,
142                filter: filter
143                    .map(|expression| {
144                        normalize_expression(engine, *expression, schema, params).map(Box::new)
145                    })
146                    .transpose()?,
147            }
148        }
149        ScalarExpr::Array(items) => {
150            ScalarExpr::Array(normalize_items(engine, items, schema, params)?)
151        }
152        ScalarExpr::Row(items) => ScalarExpr::Row(normalize_items(engine, items, schema, params)?),
153        ScalarExpr::CompositeRow {
154            items,
155            binding,
156            bound_type,
157        } => ScalarExpr::CompositeRow {
158            items: normalize_items(engine, items, schema, params)?,
159            binding: binding.clone(),
160            bound_type: bound_type.clone(),
161        },
162        ScalarExpr::Binary { op, lhs, rhs } => {
163            let left_type = expression_type(engine, &lhs, schema, params)?;
164            let right_type = expression_type(engine, &rhs, schema, params)?;
165            ScalarExpr::Binary {
166                op,
167                lhs: Box::new(normalize_unknown_literal(
168                    engine,
169                    *lhs,
170                    left_type.is_none().then_some(right_type.as_ref()).flatten(),
171                    schema,
172                    params,
173                )?),
174                rhs: Box::new(normalize_unknown_literal(
175                    engine,
176                    *rhs,
177                    right_type.is_none().then_some(left_type.as_ref()).flatten(),
178                    schema,
179                    params,
180                )?),
181            }
182        }
183        ScalarExpr::UnaryMinus(expression) => {
184            let expression = normalize_expression(engine, *expression, schema, params)?;
185            if let ScalarExpr::Literal(
186                value @ (Value::Int(_) | Value::Float(_) | Value::Decimal(_)),
187            ) = &expression
188            {
189                ScalarExpr::Literal(crate::expr::negate_value(value, None)?)
190            } else {
191                ScalarExpr::UnaryMinus(Box::new(expression))
192            }
193        }
194        ScalarExpr::Not(expression) => ScalarExpr::Not(Box::new(normalize_expression(
195            engine,
196            *expression,
197            schema,
198            params,
199        )?)),
200        ScalarExpr::And(items) => ScalarExpr::And(normalize_items(engine, items, schema, params)?),
201        ScalarExpr::Or(items) => ScalarExpr::Or(normalize_items(engine, items, schema, params)?),
202        ScalarExpr::IsNull { expr, negated } => ScalarExpr::IsNull {
203            expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
204            negated,
205        },
206        ScalarExpr::Between { expr, low, high } => {
207            let target = expression_type(engine, &expr, schema, params)?;
208            ScalarExpr::Between {
209                expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
210                low: Box::new(normalize_unknown_literal(
211                    engine,
212                    *low,
213                    target.as_ref(),
214                    schema,
215                    params,
216                )?),
217                high: Box::new(normalize_unknown_literal(
218                    engine,
219                    *high,
220                    target.as_ref(),
221                    schema,
222                    params,
223                )?),
224            }
225        }
226        ScalarExpr::InList {
227            expr,
228            list,
229            negated,
230        } => {
231            let target = expression_type(engine, &expr, schema, params)?;
232            ScalarExpr::InList {
233                expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
234                list: list
235                    .into_iter()
236                    .map(|item| {
237                        normalize_unknown_literal(engine, item, target.as_ref(), schema, params)
238                    })
239                    .collect::<Result<Vec<_>, _>>()?,
240                negated,
241            }
242        }
243        ScalarExpr::WindowCall {
244            name,
245            args,
246            mut spec,
247            filter,
248            modifiers,
249        } => {
250            spec.partition_by = normalize_items(engine, spec.partition_by, schema, params)?;
251            for order in &mut spec.order_by {
252                order.expr = normalize_expression(engine, order.expr.clone(), schema, params)?;
253            }
254            if let Some(frame) = &mut spec.frame {
255                normalize_frame_bound(engine, &mut frame.start, schema, params)?;
256                normalize_frame_bound(engine, &mut frame.end, schema, params)?;
257            }
258            ScalarExpr::WindowCall {
259                modifiers,
260                name: canonical_function_name(name),
261                args: normalize_items(engine, args, schema, params)?,
262                spec,
263                filter: filter
264                    .map(|expression| {
265                        normalize_expression(engine, *expression, schema, params).map(Box::new)
266                    })
267                    .transpose()?,
268            }
269        }
270        ScalarExpr::Case {
271            base,
272            when,
273            else_branch,
274        } => ScalarExpr::Case {
275            base: base
276                .map(|expression| {
277                    normalize_expression(engine, *expression, schema, params).map(Box::new)
278                })
279                .transpose()?,
280            when: when
281                .into_iter()
282                .map(|(condition, result)| {
283                    Ok((
284                        normalize_expression(engine, condition, schema, params)?,
285                        normalize_expression(engine, result, schema, params)?,
286                    ))
287                })
288                .collect::<Result<Vec<_>, SQLError>>()?,
289            else_branch: else_branch
290                .map(|expression| {
291                    normalize_expression(engine, *expression, schema, params).map(Box::new)
292                })
293                .transpose()?,
294        },
295        ScalarExpr::Cast { expr, ty, .. } => {
296            let source_type = expression_type(engine, &expr, schema, params)?;
297            let target_type = crate::type_resolution::resolve_declared_column_type(
298                engine,
299                &ColumnType::Named(ty),
300            )?;
301            let expression = normalize_expression(engine, *expr, schema, params)?;
302            if source_type.as_ref() == Some(&target_type) {
303                expression
304            } else if input_requires_catalog(&target_type) {
305                ScalarExpr::Cast {
306                    implicit: false,
307                    expr: Box::new(expression),
308                    ty: target_type.sql_name(),
309                }
310            } else if let ScalarExpr::Literal(value @ Value::Str(_)) = &expression {
311                let input_type = if matches!(
312                    target_type.without_temporal_modifiers(),
313                    ColumnType::Interval
314                ) {
315                    target_type.clone()
316                } else {
317                    target_type.without_type_modifiers()
318                };
319                let value = crate::expr::cast_value(value, &input_type.sql_name())?;
320                let input = normalize_expression(
321                    engine,
322                    ScalarExpr::TypedLiteral {
323                        value,
324                        ty: input_type.sql_name(),
325                        bound_type: Some(input_type.clone()),
326                        parameter_index: None,
327                    },
328                    schema,
329                    params,
330                )?;
331                if input_type == target_type {
332                    input
333                } else {
334                    ScalarExpr::Cast {
335                        implicit: false,
336                        expr: Box::new(input),
337                        ty: target_type.sql_name(),
338                    }
339                }
340            } else if let ScalarExpr::Literal(Value::Null) = expression {
341                ScalarExpr::TypedLiteral {
342                    value: Value::Null,
343                    ty: target_type.sql_name(),
344                    bound_type: Some(target_type),
345                    parameter_index: None,
346                }
347            } else {
348                ScalarExpr::Cast {
349                    implicit: false,
350                    expr: Box::new(expression),
351                    ty: target_type.sql_name(),
352                }
353            }
354        }
355        ScalarExpr::InSubquery {
356            expr,
357            subquery,
358            negated,
359        } => ScalarExpr::InSubquery {
360            expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
361            subquery,
362            negated,
363        },
364        ScalarExpr::TypedLiteral {
365            value,
366            ty,
367            bound_type,
368            parameter_index,
369        } => {
370            let literal = ScalarExpr::Literal(value.clone());
371            let declared = match bound_type {
372                Some(ty) => ty,
373                None => crate::type_resolution::resolve_declared_column_type(
374                    engine,
375                    &ColumnType::Named(ty),
376                )?,
377            };
378            if parameter_index.is_none()
379                && expression_type(engine, &literal, schema, params)?.as_ref() == Some(&declared)
380            {
381                literal
382            } else {
383                ScalarExpr::TypedLiteral {
384                    value,
385                    ty: declared.sql_name(),
386                    bound_type: Some(declared),
387                    parameter_index,
388                }
389            }
390        }
391        expression @ (ScalarExpr::Star
392        | ScalarExpr::QualifiedStar(_)
393        | ScalarExpr::Default
394        | ScalarExpr::Position(_)
395        | ScalarExpr::InternalColumn(_)
396        | ScalarExpr::Literal(_)
397        | ScalarExpr::Param(_)
398        | ScalarExpr::ScalarSubquery(_)
399        | ScalarExpr::Exists { .. }) => expression,
400    })
401}
402
403fn input_requires_catalog(ty: &ColumnType) -> bool {
404    match ty {
405        ColumnType::Named(_)
406        | ColumnType::Domain { .. }
407        | ColumnType::Regproc
408        | ColumnType::Regprocedure
409        | ColumnType::Regclass
410        | ColumnType::Regnamespace
411        | ColumnType::Regrole
412        | ColumnType::Regtype
413        | ColumnType::Record
414        | ColumnType::AnyArray => true,
415        ColumnType::Array(element) => input_requires_catalog(element),
416        _ => false,
417    }
418}
419
420fn normalize_items(
421    engine: &dyn FunctionTypeResolver,
422    items: Vec<ScalarExpr>,
423    schema: &RowSchema,
424    params: &[SQLParam],
425) -> Result<Vec<ScalarExpr>, SQLError> {
426    items
427        .into_iter()
428        .map(|item| normalize_expression(engine, item, schema, params))
429        .collect()
430}
431
432fn normalize_unknown_literal(
433    engine: &dyn FunctionTypeResolver,
434    expression: ScalarExpr,
435    target: Option<&ColumnType>,
436    schema: &RowSchema,
437    params: &[SQLParam],
438) -> Result<ScalarExpr, SQLError> {
439    if matches!(expression, ScalarExpr::Literal(Value::Null)) {
440        if let Some(target) = target {
441            return normalize_expression(
442                engine,
443                ScalarExpr::Cast {
444                    implicit: true,
445                    expr: Box::new(expression),
446                    ty: target.sql_name(),
447                },
448                schema,
449                params,
450            );
451        }
452    }
453    normalize_expression(engine, expression, schema, params)
454}
455
456fn normalize_frame_bound(
457    engine: &dyn FunctionTypeResolver,
458    bound: &mut ScalarFrameBound,
459    schema: &RowSchema,
460    params: &[SQLParam],
461) -> Result<(), SQLError> {
462    match bound {
463        ScalarFrameBound::Preceding(expression) | ScalarFrameBound::Following(expression) => {
464            **expression = normalize_expression(engine, (**expression).clone(), schema, params)?;
465        }
466        ScalarFrameBound::UnboundedPreceding
467        | ScalarFrameBound::UnboundedFollowing
468        | ScalarFrameBound::CurrentRow => {}
469    }
470    Ok(())
471}
472
473fn expression_type(
474    engine: &dyn FunctionTypeResolver,
475    expression: &ScalarExpr,
476    schema: &RowSchema,
477    params: &[SQLParam],
478) -> Result<Option<ColumnType>, SQLError> {
479    crate::scalar_type_with_resolver(expression, schema, params, engine)
480}
481
482fn canonical_function_name(name: String) -> String {
483    let lower = name.to_ascii_lowercase();
484    match lower.strip_prefix("pg_catalog.") {
485        Some(unqualified) => unqualified.to_owned(),
486        None => lower,
487    }
488}