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
17/// `GROUP BY DISTINCT` operates on expanded grouping sets after parse analysis. At this point the input schema is known, so equivalent qualified and unqualified columns can share one identity and exact no-op casts can be removed without conflating expressions that `PostgreSQL` resolves to different operator inputs.
18pub fn prepare_distinct_grouping_sets(
19    engine: &dyn FunctionTypeResolver,
20    statement: &QueryBlockPlan,
21    schema: &RowSchema,
22    params: &[SQLParam],
23) -> Result<Option<QueryBlockPlan>, SQLError> {
24    if !statement.group_distinct {
25        return Ok(None);
26    }
27
28    let mut prepared = statement.clone();
29    prepared.group_distinct = false;
30    let mut seen = HashSet::with_capacity(prepared.grouping_sets.len());
31    let mut distinct = Vec::with_capacity(prepared.grouping_sets.len());
32    for grouping_set in std::mem::take(&mut prepared.grouping_sets) {
33        let identity = grouping_set_identity(engine, &grouping_set, schema, params)?;
34        if seen.insert(identity) {
35            distinct.push(grouping_set);
36        }
37    }
38    prepared.grouping_sets = distinct;
39    Ok(Some(prepared))
40}
41
42fn grouping_set_identity(
43    engine: &dyn FunctionTypeResolver,
44    grouping_set: &[ScalarExpr],
45    schema: &RowSchema,
46    params: &[SQLParam],
47) -> Result<Vec<Vec<u8>>, SQLError> {
48    let mut identity = grouping_set
49        .iter()
50        .map(|expression| expression_identity(engine, expression, schema, params))
51        .collect::<Result<Vec<_>, _>>()?;
52    identity.sort_unstable();
53    identity.dedup();
54    Ok(identity)
55}
56
57fn expression_identity(
58    engine: &dyn FunctionTypeResolver,
59    expression: &ScalarExpr,
60    schema: &RowSchema,
61    params: &[SQLParam],
62) -> Result<Vec<u8>, SQLError> {
63    let expression =
64        crate::bind_type_introspection_with_resolver(expression.clone(), schema, params, engine);
65    let expression = normalize_expression(engine, expression, schema, params)?;
66    serde_json::to_vec(&expression).map_err(|error| {
67        SQLError::Internal(format!(
68            "serialize GROUP BY DISTINCT expression identity: {error}"
69        ))
70    })
71}
72
73#[expect(
74    clippy::too_many_lines,
75    reason = "preserves SELECT schema and row identity"
76)]
77fn normalize_expression(
78    engine: &dyn FunctionTypeResolver,
79    expression: ScalarExpr,
80    schema: &RowSchema,
81    params: &[SQLParam],
82) -> Result<ScalarExpr, SQLError> {
83    Ok(match expression {
84        ScalarExpr::Column(column) => schema
85            .unqualified_position(&column)
86            .map_or(ScalarExpr::Column(column), ScalarExpr::Position),
87        ScalarExpr::QualifiedColumn { qualifier, column } => {
88            schema.qualified_position(&qualifier, &column).map_or(
89                ScalarExpr::QualifiedColumn { qualifier, column },
90                ScalarExpr::Position,
91            )
92        }
93        ScalarExpr::Func {
94            name,
95            binding,
96            args,
97            distinct,
98            order_by,
99            filter,
100        } => {
101            let name = canonical_function_name(name);
102            let argument_types = args
103                .iter()
104                .map(|argument| expression_type(engine, argument, schema, params))
105                .collect::<Result<Vec<_>, _>>()?;
106            let targets = crate::builtin_function_argument_targets(&name, &argument_types);
107            ScalarExpr::Func {
108                name,
109                binding,
110                args: args
111                    .into_iter()
112                    .zip(targets)
113                    .map(|(argument, target)| {
114                        normalize_unknown_literal(engine, argument, target.as_ref(), schema, params)
115                    })
116                    .collect::<Result<Vec<_>, _>>()?,
117                distinct,
118                order_by: order_by
119                    .into_iter()
120                    .map(|mut order| {
121                        order.expr = normalize_expression(engine, order.expr, schema, params)?;
122                        Ok(order)
123                    })
124                    .collect::<Result<Vec<_>, SQLError>>()?,
125                filter: filter
126                    .map(|expression| {
127                        normalize_expression(engine, *expression, schema, params).map(Box::new)
128                    })
129                    .transpose()?,
130            }
131        }
132        ScalarExpr::Array(items) => {
133            ScalarExpr::Array(normalize_items(engine, items, schema, params)?)
134        }
135        ScalarExpr::Row(items) => ScalarExpr::Row(normalize_items(engine, items, schema, params)?),
136        ScalarExpr::Binary { op, lhs, rhs } => {
137            let left_type = expression_type(engine, &lhs, schema, params)?;
138            let right_type = expression_type(engine, &rhs, schema, params)?;
139            ScalarExpr::Binary {
140                op,
141                lhs: Box::new(normalize_unknown_literal(
142                    engine,
143                    *lhs,
144                    left_type.is_none().then_some(right_type.as_ref()).flatten(),
145                    schema,
146                    params,
147                )?),
148                rhs: Box::new(normalize_unknown_literal(
149                    engine,
150                    *rhs,
151                    right_type.is_none().then_some(left_type.as_ref()).flatten(),
152                    schema,
153                    params,
154                )?),
155            }
156        }
157        ScalarExpr::UnaryMinus(expression) => ScalarExpr::UnaryMinus(Box::new(
158            normalize_expression(engine, *expression, schema, params)?,
159        )),
160        ScalarExpr::Not(expression) => ScalarExpr::Not(Box::new(normalize_expression(
161            engine,
162            *expression,
163            schema,
164            params,
165        )?)),
166        ScalarExpr::And(items) => ScalarExpr::And(normalize_items(engine, items, schema, params)?),
167        ScalarExpr::Or(items) => ScalarExpr::Or(normalize_items(engine, items, schema, params)?),
168        ScalarExpr::IsNull { expr, negated } => ScalarExpr::IsNull {
169            expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
170            negated,
171        },
172        ScalarExpr::Between { expr, low, high } => {
173            let target = expression_type(engine, &expr, schema, params)?;
174            ScalarExpr::Between {
175                expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
176                low: Box::new(normalize_unknown_literal(
177                    engine,
178                    *low,
179                    target.as_ref(),
180                    schema,
181                    params,
182                )?),
183                high: Box::new(normalize_unknown_literal(
184                    engine,
185                    *high,
186                    target.as_ref(),
187                    schema,
188                    params,
189                )?),
190            }
191        }
192        ScalarExpr::InList {
193            expr,
194            list,
195            negated,
196        } => {
197            let target = expression_type(engine, &expr, schema, params)?;
198            ScalarExpr::InList {
199                expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
200                list: list
201                    .into_iter()
202                    .map(|item| {
203                        normalize_unknown_literal(engine, item, target.as_ref(), schema, params)
204                    })
205                    .collect::<Result<Vec<_>, _>>()?,
206                negated,
207            }
208        }
209        ScalarExpr::WindowCall {
210            name,
211            args,
212            mut spec,
213        } => {
214            spec.partition_by = normalize_items(engine, spec.partition_by, schema, params)?;
215            for order in &mut spec.order_by {
216                order.expr = normalize_expression(engine, order.expr.clone(), schema, params)?;
217            }
218            if let Some(frame) = &mut spec.frame {
219                normalize_frame_bound(engine, &mut frame.start, schema, params)?;
220                normalize_frame_bound(engine, &mut frame.end, schema, params)?;
221            }
222            ScalarExpr::WindowCall {
223                name: canonical_function_name(name),
224                args: normalize_items(engine, args, schema, params)?,
225                spec,
226            }
227        }
228        ScalarExpr::Case {
229            base,
230            when,
231            else_branch,
232        } => ScalarExpr::Case {
233            base: base
234                .map(|expression| {
235                    normalize_expression(engine, *expression, schema, params).map(Box::new)
236                })
237                .transpose()?,
238            when: when
239                .into_iter()
240                .map(|(condition, result)| {
241                    Ok((
242                        normalize_expression(engine, condition, schema, params)?,
243                        normalize_expression(engine, result, schema, params)?,
244                    ))
245                })
246                .collect::<Result<Vec<_>, SQLError>>()?,
247            else_branch: else_branch
248                .map(|expression| {
249                    normalize_expression(engine, *expression, schema, params).map(Box::new)
250                })
251                .transpose()?,
252        },
253        ScalarExpr::Cast { expr, ty } => {
254            let source_type = expression_type(engine, &expr, schema, params)?;
255            let target_type = ColumnType::from_sql_name(&ty)?;
256            let expression = normalize_expression(engine, *expr, schema, params)?;
257            if source_type.as_ref() == Some(&target_type) {
258                expression
259            } else if let ScalarExpr::Literal(Value::Null) = expression {
260                ScalarExpr::TypedLiteral {
261                    value: Value::Null,
262                    ty: target_type.sql_name(),
263                    bound_type: Some(target_type),
264                    parameter_index: None,
265                }
266            } else {
267                ScalarExpr::Cast {
268                    expr: Box::new(expression),
269                    ty: target_type.sql_name(),
270                }
271            }
272        }
273        ScalarExpr::InSubquery {
274            expr,
275            subquery,
276            negated,
277        } => ScalarExpr::InSubquery {
278            expr: Box::new(normalize_expression(engine, *expr, schema, params)?),
279            subquery,
280            negated,
281        },
282        ScalarExpr::TypedLiteral {
283            value,
284            ty,
285            bound_type,
286            parameter_index,
287        } => {
288            let literal = ScalarExpr::Literal(value.clone());
289            let declared = match bound_type {
290                Some(ty) => ty,
291                None => crate::type_resolution::resolve_declared_column_type(
292                    engine,
293                    &ColumnType::Named(ty),
294                )?,
295            };
296            if parameter_index.is_none()
297                && expression_type(engine, &literal, schema, params)?.as_ref() == Some(&declared)
298            {
299                literal
300            } else {
301                ScalarExpr::TypedLiteral {
302                    value,
303                    ty: declared.sql_name(),
304                    bound_type: Some(declared),
305                    parameter_index,
306                }
307            }
308        }
309        expression @ (ScalarExpr::Star
310        | ScalarExpr::QualifiedStar(_)
311        | ScalarExpr::Default
312        | ScalarExpr::Position(_)
313        | ScalarExpr::InternalColumn(_)
314        | ScalarExpr::Literal(_)
315        | ScalarExpr::Param(_)
316        | ScalarExpr::ScalarSubquery(_)
317        | ScalarExpr::Exists { .. }) => expression,
318    })
319}
320
321fn normalize_items(
322    engine: &dyn FunctionTypeResolver,
323    items: Vec<ScalarExpr>,
324    schema: &RowSchema,
325    params: &[SQLParam],
326) -> Result<Vec<ScalarExpr>, SQLError> {
327    items
328        .into_iter()
329        .map(|item| normalize_expression(engine, item, schema, params))
330        .collect()
331}
332
333fn normalize_unknown_literal(
334    engine: &dyn FunctionTypeResolver,
335    expression: ScalarExpr,
336    target: Option<&ColumnType>,
337    schema: &RowSchema,
338    params: &[SQLParam],
339) -> Result<ScalarExpr, SQLError> {
340    if matches!(expression, ScalarExpr::Literal(Value::Null)) {
341        if let Some(target) = target {
342            return normalize_expression(
343                engine,
344                ScalarExpr::Cast {
345                    expr: Box::new(expression),
346                    ty: target.sql_name(),
347                },
348                schema,
349                params,
350            );
351        }
352    }
353    normalize_expression(engine, expression, schema, params)
354}
355
356fn normalize_frame_bound(
357    engine: &dyn FunctionTypeResolver,
358    bound: &mut ScalarFrameBound,
359    schema: &RowSchema,
360    params: &[SQLParam],
361) -> Result<(), SQLError> {
362    match bound {
363        ScalarFrameBound::Preceding(expression) | ScalarFrameBound::Following(expression) => {
364            **expression = normalize_expression(engine, (**expression).clone(), schema, params)?;
365        }
366        ScalarFrameBound::UnboundedPreceding
367        | ScalarFrameBound::UnboundedFollowing
368        | ScalarFrameBound::CurrentRow => {}
369    }
370    Ok(())
371}
372
373fn expression_type(
374    engine: &dyn FunctionTypeResolver,
375    expression: &ScalarExpr,
376    schema: &RowSchema,
377    params: &[SQLParam],
378) -> Result<Option<ColumnType>, SQLError> {
379    crate::scalar_type_with_resolver(expression, schema, params, engine)
380}
381
382fn canonical_function_name(name: String) -> String {
383    let lower = name.to_ascii_lowercase();
384    match lower.strip_prefix("pg_catalog.") {
385        Some(unqualified) => unqualified.to_owned(),
386        None => lower,
387    }
388}