Skip to main content

datui_lib/query/
sql_group.rs

1//! What a SQL `GROUP BY` result was grouped from, read from the statement itself, so
2//! Enter on a row of the result can show the rows behind it.
3//!
4//! Only a plain grouping of the one table qualifies: `SELECT keys…, aggregates… FROM df
5//! [WHERE …] GROUP BY keys [HAVING …] [ORDER BY …] [LIMIT …]`, each key selected. Joins,
6//! subqueries, CTEs, set operations, window functions, `DISTINCT ON`, `GROUP BY ALL` and
7//! keys that are not in the select list give no plan, so drill-down stays off rather
8//! than guess.
9
10use 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/// How a grouped statement relates its result to its source rows.
21#[derive(Debug, PartialEq)]
22pub struct GroupPlan {
23    /// A statement over the same table and `WHERE` that returns every source row, each
24    /// computed key added as a column named by [`KeySource::Computed`].
25    pub source_sql: String,
26    pub keys: Vec<PlanKey>,
27    /// Whether the statement says which groups come back in which order: an ORDER BY,
28    /// or a LIMIT that picks some.
29    pub ordered: bool,
30}
31
32/// One grouping key: where it sits in the result and how the source computes it.
33#[derive(Debug, PartialEq)]
34pub struct PlanKey {
35    /// Position of the key's column in the result.
36    pub result_index: usize,
37    pub source: KeySource,
38}
39
40#[derive(Debug, PartialEq)]
41pub enum KeySource {
42    /// A column of the source as it stands.
43    Column(String),
44    /// An expression, added to the source rows under this name.
45    Computed(String),
46}
47
48/// Prefix of the columns `source_sql` adds for computed keys; the drill drops them.
49const KEY_PREFIX: &str = "__datui_group_key_";
50
51/// The plan for `sql`, run against a table whose columns are `columns` and giving a
52/// result `result_width` columns wide, or `None` when the statement is not a grouping
53/// whose rows can be traced back reliably.
54pub fn plan(sql: &str, columns: &[&str], result_width: usize) -> Option<GroupPlan> {
55    // Parsed as polars-sql parses it, so the two read the same statement.
56    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    // Every select item is one result column, in order; a wildcard or a
90    // multi-column alias breaks that.
91    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    // `df.dept` and `t.dept` name the column `dept`.
104    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
151/// The columns of `sql`'s result, named in `result`, that are columns of the table it
152/// reads (named in `columns`) unchanged, renamed or not: each result name with the
153/// column's. Only a plain statement's select list is read: `*`, a column, or a column
154/// `AS` a name. Anything else, or a statement that is not plain, gives none.
155pub 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    // A name given twice is an error in Polars, so a computed column never shares a
205    // name with one kept here.
206    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
225/// The select item a `GROUP BY` entry names, as Polars resolves it: an ordinal, then a
226/// select alias that is not also a source column, then the same expression selected.
227fn 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
251/// `e` spelled one way for comparing keys, keeping only what polars-sql reads: no
252/// parentheses, no quotes around names (`"dept"` is `dept`, and case always counts),
253/// no table name in front of a column, and function names in lowercase.
254fn 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
285/// No clause beyond ORDER BY and LIMIT around the one SELECT.
286fn 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
296/// One table, no join, and nothing that reshapes rows before or beside the grouping.
297fn 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/// Walks the whole statement for what the checks on its clauses cannot see: a second
336/// query anywhere (a subquery), a window function, or `UNNEST`, which changes the rows
337/// before they are grouped.
338#[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;