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