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