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
17pub 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}