1use std::ops::ControlFlow;
11
12use sqlparser::ast::{
13 Distinct, Expr, FunctionArguments, GroupByExpr, Ident, ObjectNamePart, Query, Select,
14 SelectItem, SetExpr, Statement, TableFactor, Visit, Visitor, WildcardAdditionalOptions,
15 visit_expressions_mut,
16};
17use sqlparser::dialect::GenericDialect;
18use sqlparser::parser::{Parser, ParserOptions};
19
20#[derive(Debug, PartialEq)]
22pub struct GroupPlan {
23 pub source_sql: String,
26 pub keys: Vec<PlanKey>,
27 pub ordered: bool,
30}
31
32#[derive(Debug, PartialEq)]
34pub struct PlanKey {
35 pub result_index: usize,
37 pub source: KeySource,
38}
39
40#[derive(Debug, PartialEq)]
41pub enum KeySource {
42 Column(String),
44 Computed(String),
46}
47
48const KEY_PREFIX: &str = "__datui_group_key_";
50
51pub fn plan(sql: &str, columns: &[&str], result_width: usize) -> Option<GroupPlan> {
55 let statements = Parser::new(&GenericDialect)
57 .with_options(ParserOptions {
58 trailing_commas: true,
59 ..Default::default()
60 })
61 .try_with_sql(sql)
62 .ok()?
63 .parse_statements()
64 .ok()?;
65 let [Statement::Query(query)] = statements.as_slice() else {
66 return None;
67 };
68 if !plain_query(query) || !Plain::check(query) {
69 return None;
70 }
71 let SetExpr::Select(select) = query.body.as_ref() else {
72 return None;
73 };
74 if !plain_select(select) {
75 return None;
76 }
77 let GroupByExpr::Expressions(group_by, modifiers) = &select.group_by else {
78 return None;
79 };
80 if group_by.is_empty() || !modifiers.is_empty() {
81 return None;
82 }
83 let [from] = select.from.as_slice() else {
84 return None;
85 };
86 let TableFactor::Table { name, alias, .. } = &from.relation else {
87 return None;
88 };
89 let items: Vec<(&Expr, Option<&Ident>)> = select
92 .projection
93 .iter()
94 .map(|item| match item {
95 SelectItem::UnnamedExpr(e) => Some((e, None)),
96 SelectItem::ExprWithAlias { expr, alias } => Some((expr, Some(alias))),
97 _ => None,
98 })
99 .collect::<Option<_>>()?;
100 if items.len() != result_width {
101 return None;
102 }
103 let qualifiers: Vec<&str> = name
105 .0
106 .last()
107 .and_then(|p| p.as_ident())
108 .into_iter()
109 .chain(alias.as_ref().map(|a| &a.name))
110 .map(|ident| ident.value.as_str())
111 .collect();
112 let normalize = |e: &Expr| normalized(e, &qualifiers);
113
114 let mut keys = Vec::with_capacity(group_by.len());
115 let mut computed = Vec::new();
116 for key in group_by {
117 let (index, expr) = resolve_key(key, &items, columns, &normalize)?;
118 let source = match normalize(expr) {
119 Expr::Identifier(ident) if columns.contains(&ident.value.as_str()) => {
120 KeySource::Column(ident.value)
121 }
122 _ => {
123 let name = format!("{KEY_PREFIX}{}", computed.len());
124 computed.push(format!("{expr} AS \"{name}\""));
125 KeySource::Computed(name)
126 }
127 };
128 keys.push(PlanKey {
129 result_index: index,
130 source,
131 });
132 }
133
134 let mut source_sql = String::from("SELECT *");
135 for column in &computed {
136 source_sql.push_str(", ");
137 source_sql.push_str(column);
138 }
139 source_sql.push_str(&format!(" FROM {from}"));
140 if let Some(selection) = &select.selection {
141 source_sql.push_str(&format!(" WHERE {selection}"));
142 }
143 let ordered = query.order_by.is_some() || query.limit_clause.is_some() || query.fetch.is_some();
144 Some(GroupPlan {
145 source_sql,
146 keys,
147 ordered,
148 })
149}
150
151pub fn passed_through(sql: &str, columns: &[&str], result: &[&str]) -> Vec<(String, String)> {
156 let Some(statements) = Parser::new(&GenericDialect)
157 .with_options(ParserOptions {
158 trailing_commas: true,
159 ..Default::default()
160 })
161 .try_with_sql(sql)
162 .and_then(|mut parser| parser.parse_statements())
163 .ok()
164 else {
165 return Vec::new();
166 };
167 let [Statement::Query(query)] = statements.as_slice() else {
168 return Vec::new();
169 };
170 if !plain_query(query) || !Plain::check(query) {
171 return Vec::new();
172 }
173 let SetExpr::Select(select) = query.body.as_ref() else {
174 return Vec::new();
175 };
176 if !plain_select(select) {
177 return Vec::new();
178 }
179 let [from] = select.from.as_slice() else {
180 return Vec::new();
181 };
182 let TableFactor::Table { name, alias, .. } = &from.relation else {
183 return Vec::new();
184 };
185 let qualifiers: Vec<&str> = name
186 .0
187 .last()
188 .and_then(|p| p.as_ident())
189 .into_iter()
190 .chain(alias.as_ref().map(|a| &a.name))
191 .map(|ident| ident.value.as_str())
192 .collect();
193 let column = |e: &Expr| match normalized(e, &qualifiers) {
194 Expr::Identifier(ident) if columns.contains(&ident.value.as_str()) => Some(ident.value),
195 _ => None,
196 };
197 let plain = |options: &WildcardAdditionalOptions| {
198 options.opt_ilike.is_none()
199 && options.opt_exclude.is_none()
200 && options.opt_except.is_none()
201 && options.opt_replace.is_none()
202 && options.opt_rename.is_none()
203 };
204 let mut kept: Vec<(String, String)> = Vec::new();
207 for item in &select.projection {
208 match item {
209 SelectItem::UnnamedExpr(e) => kept.extend(column(e).map(|c| (c.clone(), c))),
210 SelectItem::ExprWithAlias { expr, alias } => {
211 kept.extend(column(expr).map(|c| (alias.value.clone(), c)));
212 }
213 SelectItem::Wildcard(options) | SelectItem::QualifiedWildcard(_, options)
214 if plain(options) =>
215 {
216 kept.extend(columns.iter().map(|c| (c.to_string(), c.to_string())));
217 }
218 _ => {}
219 }
220 }
221 kept.retain(|(shown, _)| result.contains(&shown.as_str()));
222 kept
223}
224
225fn resolve_key<'a>(
228 key: &'a Expr,
229 items: &[(&'a Expr, Option<&Ident>)],
230 columns: &[&str],
231 normalize: &impl Fn(&Expr) -> Expr,
232) -> Option<(usize, &'a Expr)> {
233 if let Expr::Value(value) = key {
234 let ordinal: usize = value.to_string().parse().ok()?;
235 let (expr, _) = items.get(ordinal.checked_sub(1)?)?;
236 return Some((ordinal - 1, expr));
237 }
238 if let Expr::Identifier(ident) = key
239 && !columns.contains(&ident.value.as_str())
240 && let Some(index) = items
241 .iter()
242 .position(|(_, alias)| alias.is_some_and(|a| a.value == ident.value))
243 {
244 return Some((index, items[index].0));
245 }
246 let wanted = normalize(key);
247 let index = items.iter().position(|(e, _)| normalize(e) == wanted)?;
248 Some((index, key))
249}
250
251fn normalized(e: &Expr, qualifiers: &[&str]) -> Expr {
255 let mut e = e.clone();
256 let _ = visit_expressions_mut(&mut e, |e| {
257 match e {
258 Expr::Nested(inner) => *e = inner.as_ref().clone(),
259 Expr::Identifier(ident) => ident.quote_style = None,
260 Expr::CompoundIdentifier(parts) => {
261 for part in parts.iter_mut() {
262 part.quote_style = None;
263 }
264 if let [table, column] = parts.as_slice()
265 && qualifiers.contains(&table.value.as_str())
266 {
267 *e = Expr::Identifier(column.clone());
268 }
269 }
270 Expr::Function(f) => {
271 for part in f.name.0.iter_mut() {
272 if let ObjectNamePart::Identifier(ident) = part {
273 ident.value = ident.value.to_lowercase();
274 ident.quote_style = None;
275 }
276 }
277 }
278 _ => {}
279 }
280 ControlFlow::<()>::Continue(())
281 });
282 e
283}
284
285fn plain_query(query: &Query) -> bool {
287 query.with.is_none()
288 && matches!(query.body.as_ref(), SetExpr::Select(_))
289 && query.locks.is_empty()
290 && query.for_clause.is_none()
291 && query.settings.is_none()
292 && query.format_clause.is_none()
293 && query.pipe_operators.is_empty()
294}
295
296fn plain_select(select: &Select) -> bool {
298 let one_table = match select.from.as_slice() {
299 [from] => {
300 from.joins.is_empty()
301 && matches!(
302 &from.relation,
303 TableFactor::Table {
304 alias,
305 args: None,
306 sample: None,
307 version: None,
308 with_ordinality: false,
309 json_path: None,
310 ..
311 } if alias.as_ref().is_none_or(|a| a.columns.is_empty())
312 )
313 }
314 _ => false,
315 };
316 one_table
317 && matches!(
318 select.distinct,
319 None | Some(Distinct::Distinct | Distinct::All)
320 )
321 && select.top.is_none()
322 && select.exclude.is_none()
323 && select.into.is_none()
324 && select.lateral_views.is_empty()
325 && select.prewhere.is_none()
326 && select.connect_by.is_empty()
327 && select.cluster_by.is_empty()
328 && select.distribute_by.is_empty()
329 && select.sort_by.is_empty()
330 && select.named_window.is_empty()
331 && select.qualify.is_none()
332 && select.value_table_mode.is_none()
333}
334
335#[derive(Default)]
339struct Plain {
340 queries: usize,
341}
342
343impl Plain {
344 fn check(query: &Query) -> bool {
345 let mut plain = Plain::default();
346 query.visit(&mut plain).is_continue() && plain.queries == 1
347 }
348}
349
350impl Visitor for Plain {
351 type Break = ();
352
353 fn pre_visit_query(&mut self, _query: &Query) -> ControlFlow<()> {
354 self.queries += 1;
355 ControlFlow::Continue(())
356 }
357
358 fn pre_visit_expr(&mut self, expr: &Expr) -> ControlFlow<()> {
359 match expr {
360 Expr::Subquery(_) | Expr::InSubquery { .. } | Expr::Exists { .. } => {
361 ControlFlow::Break(())
362 }
363 Expr::Function(f)
364 if f.over.is_some()
365 || matches!(f.args, FunctionArguments::Subquery(_))
366 || f.name.to_string().eq_ignore_ascii_case("unnest") =>
367 {
368 ControlFlow::Break(())
369 }
370 _ => ControlFlow::Continue(()),
371 }
372 }
373}
374
375#[cfg(test)]
376mod tests {
377 use super::*;
378
379 const COLUMNS: &[&str] = &["dept", "salary", "ts", "name"];
380
381 fn keys(sql: &str, width: usize) -> Option<Vec<(usize, KeySource)>> {
382 plan(sql, COLUMNS, width).map(|p| {
383 p.keys
384 .into_iter()
385 .map(|k| (k.result_index, k.source))
386 .collect()
387 })
388 }
389
390 fn column(name: &str) -> KeySource {
391 KeySource::Column(name.to_string())
392 }
393
394 fn computed(i: usize) -> KeySource {
395 KeySource::Computed(format!("{KEY_PREFIX}{i}"))
396 }
397
398 #[test]
399 fn plain_grouping_keeps_where_and_drops_the_rest() {
400 let plan = plan(
401 "SELECT dept, AVG(salary) AS avg_salary FROM df WHERE salary > 100000 \
402 GROUP BY dept HAVING COUNT(*) > 1 ORDER BY avg_salary DESC LIMIT 3",
403 COLUMNS,
404 2,
405 )
406 .expect("a plan");
407 assert_eq!(plan.source_sql, "SELECT * FROM df WHERE salary > 100000");
408 assert_eq!(
409 plan.keys,
410 vec![PlanKey {
411 result_index: 0,
412 source: column("dept")
413 }]
414 );
415 }
416
417 #[test]
418 fn keys_resolve_by_alias_ordinal_and_expression() {
419 assert_eq!(
421 keys("SELECT COUNT(*) AS n, dept AS d FROM df GROUP BY dept", 2),
422 Some(vec![(1, column("dept"))])
423 );
424 assert_eq!(
425 keys("SELECT dept, COUNT(*) FROM df GROUP BY 1", 2),
426 Some(vec![(0, column("dept"))])
427 );
428 let p = plan(
430 "SELECT EXTRACT(HOUR FROM ts) AS h, COUNT(*) AS n FROM df GROUP BY h",
431 COLUMNS,
432 2,
433 )
434 .expect("a plan");
435 assert_eq!(p.keys[0].source, computed(0));
436 assert_eq!(
437 p.source_sql,
438 format!("SELECT *, EXTRACT(HOUR FROM ts) AS \"{KEY_PREFIX}0\" FROM df")
439 );
440 assert_eq!(
442 keys(
443 "SELECT dept, EXTRACT(HOUR FROM ts) AS h, SUM(salary) FROM df \
444 GROUP BY dept, EXTRACT(HOUR FROM ts)",
445 3
446 ),
447 Some(vec![(0, column("dept")), (1, computed(0))])
448 );
449 assert_eq!(
451 keys("SELECT t.dept, COUNT(*) FROM df AS t GROUP BY dept", 2),
452 Some(vec![(0, column("dept"))])
453 );
454 }
455
456 #[test]
459 fn keys_match_however_they_are_spelled() {
460 assert_eq!(
461 keys("SELECT \"dept\", COUNT(*) FROM df GROUP BY dept", 2),
462 Some(vec![(0, column("dept"))])
463 );
464 assert_eq!(
465 keys(
466 "SELECT t.dept, COUNT(*) FROM df t GROUP BY \"t\".\"dept\"",
467 2
468 ),
469 Some(vec![(0, column("dept"))])
470 );
471 assert_eq!(
472 keys(
473 "SELECT upper(dept) AS u, COUNT(*) FROM df GROUP BY UPPER((df.\"dept\"))",
474 2
475 ),
476 Some(vec![(0, computed(0))])
477 );
478 assert_eq!(keys("SELECT Dept, COUNT(*) FROM df GROUP BY dept", 2), None);
480 }
481
482 #[test]
483 fn a_source_column_wins_over_an_alias_of_the_same_name() {
484 assert_eq!(
486 keys(
487 "SELECT salary, dept AS salary2, COUNT(*) FROM df GROUP BY salary, dept",
488 3
489 ),
490 Some(vec![(0, column("salary")), (1, column("dept"))])
491 );
492 }
493
494 #[test]
495 fn shapes_without_a_reliable_source_give_no_plan() {
496 for sql in [
497 "SELECT dept, salary FROM df",
498 "SELECT AVG(salary) FROM df GROUP BY dept",
499 "SELECT * FROM df GROUP BY dept",
500 "SELECT a.dept, COUNT(*) FROM df a JOIN df b ON a.name = b.name GROUP BY a.dept",
501 "SELECT dept, COUNT(*) FROM df, df AS b GROUP BY dept",
502 "SELECT dept, COUNT(*) FROM (SELECT * FROM df) GROUP BY dept",
503 "WITH t AS (SELECT * FROM df) SELECT dept, COUNT(*) FROM t GROUP BY dept",
504 "SELECT dept, COUNT(*) FROM df WHERE salary > (SELECT AVG(salary) FROM df) \
505 GROUP BY dept",
506 "SELECT dept, COUNT(*) FROM df WHERE dept IN (SELECT dept FROM df) GROUP BY dept",
507 "SELECT dept, COUNT(*) FROM df GROUP BY dept UNION SELECT dept, 1 FROM df",
508 "SELECT dept, RANK() OVER (ORDER BY COUNT(*)) FROM df GROUP BY dept",
509 "SELECT DISTINCT ON (dept) dept, COUNT(*) FROM df GROUP BY dept",
510 "SELECT dept, COUNT(*) FROM df GROUP BY ALL",
511 "SELECT dept, COUNT(*) FROM df GROUP BY 3",
512 "SELECT UPPER(dept), COUNT(*) FROM df GROUP BY LOWER(dept)",
513 "SELECT dept, COUNT(*) FROM df GROUP BY dept; SELECT 1",
514 "SELECT dept, UNNEST(name), COUNT(*) FROM df GROUP BY dept",
515 "not sql at all",
516 ] {
517 assert_eq!(plan(sql, COLUMNS, 2), None, "{sql}");
518 }
519 assert_eq!(
521 plan("SELECT dept, COUNT(*) FROM df GROUP BY dept", COLUMNS, 3),
522 None
523 );
524 }
525
526 #[test]
527 fn passed_through_names_the_columns_a_statement_keeps() {
528 let columns = ["a", "b", "c d"];
529 let kept = |sql: &str, result: &[&str]| super::passed_through(sql, &columns, result);
530 let pairs = |p: &[(&str, &str)]| -> Vec<(String, String)> {
531 p.iter()
532 .map(|(s, f)| (s.to_string(), f.to_string()))
533 .collect()
534 };
535 assert_eq!(
536 kept("SELECT * FROM df WHERE a > 1", &["a", "b", "c d"]),
537 pairs(&[("a", "a"), ("b", "b"), ("c d", "c d")])
538 );
539 assert_eq!(
540 kept(
541 r#"SELECT b * 2 AS b, df.a AS x, "c d", COUNT(*) AS a FROM df GROUP BY 1, 2, 3"#,
542 &["b", "x", "c d", "a"]
543 ),
544 pairs(&[("x", "a"), ("c d", "c d")])
545 );
546 assert!(kept("SELECT * EXCLUDE (a) FROM df", &["b", "c d"]).is_empty());
547 assert!(kept("SELECT a FROM df JOIN df AS e ON true", &["a"]).is_empty());
548 assert!(kept("not sql", &["a"]).is_empty());
549 }
550}