1use 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
25pub 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}