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(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
26pub 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 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}