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                        composite_source: None,
324                        value,
325                        ty: input_type.sql_name(),
326                        bound_type: Some(input_type.clone()),
327                        parameter_index: None,
328                    },
329                    schema,
330                    params,
331                )?;
332                if input_type == target_type {
333                    input
334                } else {
335                    ScalarExpr::Cast {
336                        implicit: false,
337                        expr: Box::new(input),
338                        ty: target_type.sql_name(),
339                    }
340                }
341            } else if let ScalarExpr::Literal(Value::Null) = expression {
342                ScalarExpr::TypedLiteral {
343                    composite_source: None,
344                    value: Value::Null,
345                    ty: target_type.sql_name(),
346                    bound_type: Some(target_type),
347                    parameter_index: None,
348                }
349            } else {
350                ScalarExpr::Cast {
351                    implicit: false,
352                    expr: Box::new(expression),
353                    ty: target_type.sql_name(),
354                }
355            }
356        }
357        ScalarExpr::InSubquery {
358            expr,
359            subquery,
360            negated,
361        } => ScalarExpr::InSubquery {
362            expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
363            subquery,
364            negated,
365        },
366        ScalarExpr::TypedLiteral {
367            value,
368            ty,
369            bound_type,
370            parameter_index,
371            composite_source,
372        } => {
373            let literal = ScalarExpr::Literal(value.clone());
374            let declared = match bound_type {
375                Some(ty) => ty,
376                None => crate::type_resolution::resolve_declared_column_type(
377                    engine,
378                    &ColumnType::Named(ty),
379                )?,
380            };
381            if parameter_index.is_none()
382                && composite_source.is_none()
383                && expression_type(engine, &literal, schema, params)?.as_ref() == Some(&declared)
384            {
385                literal
386            } else {
387                ScalarExpr::TypedLiteral {
388                    composite_source,
389                    value,
390                    ty: declared.sql_name(),
391                    bound_type: Some(declared),
392                    parameter_index,
393                }
394            }
395        }
396        expression @ (ScalarExpr::Star
397        | ScalarExpr::QualifiedStar(_)
398        | ScalarExpr::Default
399        | ScalarExpr::Position(_)
400        | ScalarExpr::InternalColumn(_)
401        | ScalarExpr::Literal(_)
402        | ScalarExpr::Param(_)
403        | ScalarExpr::ScalarSubquery(_)
404        | ScalarExpr::Exists { .. }) => expression,
405    })
406}
407
408fn input_requires_catalog(ty: &ColumnType) -> bool {
409    match ty {
410        ColumnType::Named(_)
411        | ColumnType::Domain { .. }
412        | ColumnType::Enum(_)
413        | ColumnType::Composite(_)
414        | ColumnType::Regproc
415        | ColumnType::Regprocedure
416        | ColumnType::Regclass
417        | ColumnType::Regcollation
418        | ColumnType::Regnamespace
419        | ColumnType::Regrole
420        | ColumnType::Regtype
421        | ColumnType::Record
422        | ColumnType::AnyArray => true,
423        ColumnType::Array(element) => input_requires_catalog(element),
424        _ => false,
425    }
426}
427
428fn normalize_items(
429    engine: &dyn FunctionTypeResolver,
430    items: Vec<ScalarExpr>,
431    schema: &RowSchema,
432    params: &[SQLParam],
433) -> Result<Vec<ScalarExpr>, SQLError> {
434    items
435        .into_iter()
436        .map(|item| normalize_expression(engine, item, schema, params))
437        .collect()
438}
439
440fn normalize_unknown_literal(
441    engine: &dyn FunctionTypeResolver,
442    expression: ScalarExpr,
443    target: Option<&ColumnType>,
444    schema: &RowSchema,
445    params: &[SQLParam],
446) -> Result<ScalarExpr, SQLError> {
447    if matches!(expression, ScalarExpr::Literal(Value::Null)) {
448        if let Some(target) = target {
449            return normalize_expression(
450                engine,
451                ScalarExpr::Cast {
452                    implicit: true,
453                    expr: Box::new(expression),
454                    ty: target.sql_name(),
455                },
456                schema,
457                params,
458            );
459        }
460    }
461    normalize_expression(engine, expression, schema, params)
462}
463
464fn normalize_frame_bound(
465    engine: &dyn FunctionTypeResolver,
466    bound: &mut ScalarFrameBound,
467    schema: &RowSchema,
468    params: &[SQLParam],
469) -> Result<(), SQLError> {
470    match bound {
471        ScalarFrameBound::Preceding(expression) | ScalarFrameBound::Following(expression) => {
472            **expression = normalize_expression(engine, (**expression).clone(), schema, params)?;
473        }
474        ScalarFrameBound::UnboundedPreceding
475        | ScalarFrameBound::UnboundedFollowing
476        | ScalarFrameBound::CurrentRow => {}
477    }
478    Ok(())
479}
480
481fn expression_type(
482    engine: &dyn FunctionTypeResolver,
483    expression: &ScalarExpr,
484    schema: &RowSchema,
485    params: &[SQLParam],
486) -> Result<Option<ColumnType>, SQLError> {
487    crate::scalar_type_with_resolver(expression, schema, params, engine)
488}
489
490fn canonical_function_name(name: String) -> String {
491    let lower = name.to_ascii_lowercase();
492    match lower.strip_prefix("pg_catalog.") {
493        Some(unqualified) => unqualified.to_owned(),
494        None => lower,
495    }
496}