Skip to main content

datui_lib/query/
mod.rs

1pub(crate) mod query_prompt;
2pub(crate) mod sql_assist;
3// Public for the `sql_group_plan` fuzz target.
4#[cfg(feature = "sql")]
5pub mod sql_group;
6#[cfg(feature = "sql")]
7pub(crate) mod sql_plan;
8
9use polars::prelude::StrptimeOptions;
10use polars::prelude::*;
11use std::ops::{Add, Div, Mul, Rem, Sub};
12
13#[derive(Debug, Clone, PartialEq)]
14enum Token {
15    Identifier(String),
16    Number(f64),
17    String(String),
18    /// Date literal in YYYY.MM.DD format, stored as ISO "YYYY-MM-DD" for Polars
19    DateLiteral(String),
20    /// Timestamp literal YYYY.MM.DDTHH:MM:SS[.fff...]
21    TimestampLiteral {
22        iso: String,
23        format_str: String,
24        time_unit: TimeUnit,
25    },
26    Op(String),
27    LParen,
28    RParen,
29    LBracket,
30    RBracket,
31    Comma,
32    Colon,
33    Pipe,
34    Dot,
35    Select,
36    Where,
37    By,
38}
39
40/// Parse YYYY.MM.DDTHH:MM:SS[.fff...] timestamp. Consumes from chars. Returns (iso_string, format, time_unit) or None.
41fn parse_timestamp_literal(
42    date_part: &str,
43    chars: &mut std::iter::Peekable<std::str::Chars<'_>>,
44) -> Option<(String, String, TimeUnit)> {
45    if chars.peek() != Some(&'T') {
46        return None;
47    }
48    chars.next(); // consume 'T'
49    let mut time_part = String::new();
50    while let Some(&c) = chars.peek() {
51        if c.is_ascii_digit() || c == ':' || c == '.' {
52            time_part.push(c);
53            chars.next();
54        } else {
55            break;
56        }
57    }
58    let parts: Vec<&str> = time_part.split(':').collect();
59    if parts.len() != 3 {
60        return None;
61    }
62    let (h, m, s) = (parts[0], parts[1], parts[2]);
63    if h.len() != 2 || m.len() != 2 || s.len() < 2 {
64        return None;
65    }
66    let (sec_part, frac) = match s.split_once('.') {
67        Some((a, f)) => (a, f),
68        None => (s, ""),
69    };
70    let (time_unit, format_str) = match frac.len() {
71        0 => (TimeUnit::Microseconds, "%Y-%m-%dT%H:%M:%S".to_string()),
72        1..=3 => (TimeUnit::Milliseconds, "%Y-%m-%dT%H:%M:%S%.3f".to_string()),
73        4..=6 => (TimeUnit::Microseconds, "%Y-%m-%dT%H:%M:%S%.6f".to_string()),
74        7..=9 => (TimeUnit::Nanoseconds, "%Y-%m-%dT%H:%M:%S%.9f".to_string()),
75        _ => (TimeUnit::Nanoseconds, "%Y-%m-%dT%H:%M:%S%.9f".to_string()),
76    };
77    let iso_date = parse_date_literal(date_part)?;
78    let frac_padded = match time_unit {
79        TimeUnit::Milliseconds => format!("{:0<3}", frac),
80        TimeUnit::Microseconds => format!("{:0<6}", frac),
81        TimeUnit::Nanoseconds => format!("{:0<9}", frac),
82    };
83    let iso = if frac.is_empty() {
84        format!("{}T{}:{}:{}", iso_date, h, m, sec_part)
85    } else {
86        format!("{}T{}:{}:{}.{}", iso_date, h, m, sec_part, frac_padded)
87    };
88    Some((iso, format_str, time_unit))
89}
90
91/// Parse YYYY.MM.DD date literal (e.g. 2021.01.01). Returns ISO string "YYYY-MM-DD" or None.
92fn parse_date_literal(s: &str) -> Option<String> {
93    let parts: Vec<&str> = s.split('.').collect();
94    if parts.len() != 3 {
95        return None;
96    }
97    let year: u32 = parts[0].parse().ok()?;
98    let month: u32 = parts[1].parse().ok()?;
99    let day: u32 = parts[2].parse().ok()?;
100    if parts[0].len() != 4 || !(1000..=9999).contains(&year) {
101        return None;
102    }
103    if !(1..=12).contains(&month) || !(1..=31).contains(&day) {
104        return None;
105    }
106    Some(format!("{:04}-{:02}-{:02}", year, month, day))
107}
108
109fn tokenize(input: &str) -> Result<Vec<Token>, String> {
110    let mut tokens = Vec::new();
111    let mut chars = input.chars().peekable();
112
113    while let Some(&c) = chars.peek() {
114        match c {
115            ' ' | '\t' | '\n' | '\r' => {
116                chars.next();
117            }
118            ',' => {
119                tokens.push(Token::Comma);
120                chars.next();
121            }
122            ':' => {
123                tokens.push(Token::Colon);
124                chars.next();
125            }
126            '|' => {
127                tokens.push(Token::Pipe);
128                chars.next();
129            }
130            '(' => {
131                tokens.push(Token::LParen);
132                chars.next();
133            }
134            ')' => {
135                tokens.push(Token::RParen);
136                chars.next();
137            }
138            '[' => {
139                tokens.push(Token::LBracket);
140                chars.next();
141            }
142            ']' => {
143                tokens.push(Token::RBracket);
144                chars.next();
145            }
146            '"' => {
147                chars.next(); // consume opening quote
148                let mut string_val = String::new();
149                let mut found_closing_quote = false;
150                while let Some(&c) = chars.peek() {
151                    if c == '\\' {
152                        chars.next(); // consume backslash
153                        if let Some(&next_c) = chars.peek() {
154                            match next_c {
155                                'n' => {
156                                    string_val.push('\n');
157                                    chars.next();
158                                }
159                                't' => {
160                                    string_val.push('\t');
161                                    chars.next();
162                                }
163                                'r' => {
164                                    string_val.push('\r');
165                                    chars.next();
166                                }
167                                '\\' => {
168                                    string_val.push('\\');
169                                    chars.next();
170                                }
171                                '"' => {
172                                    string_val.push('"');
173                                    chars.next();
174                                }
175                                _ => {
176                                    // Unknown escape, just include the backslash and next char
177                                    string_val.push('\\');
178                                    string_val.push(next_c);
179                                    chars.next();
180                                }
181                            }
182                        } else {
183                            return Err("Unterminated escape sequence in string".to_string());
184                        }
185                    } else if c == '"' {
186                        chars.next(); // consume closing quote
187                        found_closing_quote = true;
188                        break;
189                    } else {
190                        string_val.push(c);
191                        chars.next();
192                    }
193                }
194                if !found_closing_quote {
195                    return Err("Unterminated string literal".to_string());
196                }
197                tokens.push(Token::String(string_val));
198            }
199            '^' => {
200                tokens.push(Token::Op("^".to_string()));
201                chars.next();
202            }
203            '+' | '-' | '*' | '%' | '/' | '=' | '<' | '>' | '!' => {
204                let mut op = c.to_string();
205                chars.next();
206                if let Some(&next_c) = chars.peek()
207                    && ((c == '<' && (next_c == '=' || next_c == '>'))
208                        || (c == '>' && next_c == '=')
209                        || (c == '!' && next_c == '='))
210                {
211                    op.push(next_c);
212                    chars.next();
213                }
214                tokens.push(Token::Op(op));
215            }
216            '.' => {
217                chars.next();
218                if chars.peek().is_some_and(|nc| nc.is_ascii_digit()) {
219                    let mut num_str = String::from('.');
220                    while let Some(&nc) = chars.peek() {
221                        if nc.is_ascii_digit() {
222                            num_str.push(nc);
223                            chars.next();
224                        } else {
225                            break;
226                        }
227                    }
228                    if let Ok(n) = num_str.parse::<f64>() {
229                        tokens.push(Token::Number(n));
230                    } else {
231                        return Err(format!("Invalid number: {}", num_str));
232                    }
233                } else {
234                    tokens.push(Token::Dot);
235                }
236            }
237            '0'..='9' => {
238                let mut num_str = String::new();
239                while let Some(&nc) = chars.peek() {
240                    if nc.is_ascii_digit() || nc == '.' {
241                        num_str.push(nc);
242                        chars.next();
243                    } else {
244                        break;
245                    }
246                }
247                // Check for YYYY.MM.DDTHH:MM:SS timestamp literal (peek for 'T' before consuming)
248                let is_timestamp =
249                    parse_date_literal(&num_str).is_some() && chars.peek() == Some(&'T');
250                if is_timestamp
251                    && let Some((iso, format_str, time_unit)) =
252                        parse_timestamp_literal(&num_str, &mut chars)
253                {
254                    tokens.push(Token::TimestampLiteral {
255                        iso,
256                        format_str,
257                        time_unit,
258                    });
259                    continue;
260                }
261                // Check for YYYY.MM.DD date literal
262                if let Some(iso) = parse_date_literal(&num_str) {
263                    tokens.push(Token::DateLiteral(iso));
264                } else if let Ok(n) = num_str.parse::<f64>() {
265                    tokens.push(Token::Number(n));
266                } else {
267                    return Err(format!("Invalid number: {}", num_str));
268                }
269            }
270            _ if c.is_alphabetic() || c == '_' => {
271                let mut ident = String::new();
272                while let Some(&nc) = chars.peek() {
273                    if nc.is_alphanumeric() || nc == '_' {
274                        ident.push(nc);
275                        chars.next();
276                    } else {
277                        break;
278                    }
279                }
280                match ident.as_str() {
281                    "select" => tokens.push(Token::Select),
282                    "where" => tokens.push(Token::Where),
283                    "by" => tokens.push(Token::By),
284                    _ => tokens.push(Token::Identifier(ident)),
285                }
286            }
287            _ => return Err(format!("Unexpected character: {}", c)),
288        }
289    }
290    Ok(tokens)
291}
292
293fn split_tokens(tokens: &[Token], delimiter: &Token) -> Vec<Vec<Token>> {
294    let mut result = Vec::new();
295    let mut current = Vec::new();
296    let mut depth = 0;
297    let mut bracket_depth = 0;
298
299    for token in tokens {
300        match token {
301            Token::LParen => depth += 1,
302            Token::RParen => depth -= 1,
303            Token::LBracket => bracket_depth += 1,
304            Token::RBracket => bracket_depth -= 1,
305            _ => {}
306        }
307
308        if depth == 0 && bracket_depth == 0 && token == delimiter {
309            result.push(current);
310            current = Vec::new();
311        } else {
312            current.push(token.clone());
313        }
314    }
315    result.push(current);
316    result
317}
318
319/// The token as the user typed it, for error messages.
320fn token_text(token: &Token) -> String {
321    match token {
322        Token::Identifier(s) => s.clone(),
323        Token::Number(n) => n.to_string(),
324        Token::String(s) => format!("\"{}\"", s),
325        Token::DateLiteral(iso) => iso.clone(),
326        Token::TimestampLiteral { iso, .. } => iso.clone(),
327        Token::Op(op) => op.clone(),
328        Token::LParen => "(".to_string(),
329        Token::RParen => ")".to_string(),
330        Token::LBracket => "[".to_string(),
331        Token::RBracket => "]".to_string(),
332        Token::Comma => ",".to_string(),
333        Token::Colon => ":".to_string(),
334        Token::Pipe => "|".to_string(),
335        Token::Dot => ".".to_string(),
336        Token::Select => "select".to_string(),
337        Token::Where => "where".to_string(),
338        Token::By => "by".to_string(),
339    }
340}
341
342/// Remedy shown when a clause keyword turns up out of place.
343const CLAUSE_ORDER: &str = "clause order is select [by group] [where conditions]";
344
345/// [`CLAUSE_ORDER`] with q's optional `from df`, for errors about it.
346const FROM_ORDER: &str = "clause order is select [by group] [from df] [where conditions]";
347
348/// The one table q reads: the one on screen, named as SQL names it.
349const TABLE: &str = "df";
350
351/// The table name after a `from` at `tokens[i]`: an identifier, or a dotted
352/// path like `data.csv`, that is not an operator word. Returns it and the index
353/// past it.
354fn from_table_at(tokens: &[Token], i: usize) -> Option<(String, usize)> {
355    if tokens.get(i) != Some(&Token::Identifier("from".to_string())) {
356        return None;
357    }
358    let mut name = match tokens.get(i + 1) {
359        Some(Token::Identifier(n)) if !WORD_OPS.contains(&n.as_str()) => n.clone(),
360        _ => return None,
361    };
362    let mut end = i + 2;
363    while let (Some(Token::Dot), Some(Token::Identifier(part))) =
364        (tokens.get(end), tokens.get(end + 1))
365    {
366        name.push('.');
367        name.push_str(part);
368        end += 2;
369    }
370    Some((name, end))
371}
372
373/// The query body without q's `from df` (after the select list and `by`, before
374/// `where`). `from` is the clause only where a column cannot be: followed by a table
375/// name, then `where`, `by` or the end, outside brackets.
376fn strip_from(body: &[Token]) -> Result<Vec<Token>, String> {
377    let mut depth = 0i32;
378    let mut found: Option<(usize, usize)> = None;
379    for (i, token) in body.iter().enumerate() {
380        match token {
381            Token::LParen | Token::LBracket => depth += 1,
382            Token::RParen | Token::RBracket => depth -= 1,
383            _ => {}
384        }
385        if depth != 0 {
386            continue;
387        }
388        let Some((name, end)) = from_table_at(body, i) else {
389            continue;
390        };
391        let after_where = body[..i].contains(&Token::Where);
392        match body.get(end) {
393            None | Some(Token::Where) | Some(Token::By) => {}
394            _ => continue,
395        }
396        if name != TABLE {
397            return Err("q reads the table on screen, named df: … from df …".to_string());
398        }
399        if after_where {
400            return Err(format!(
401                "Unexpected 'from df' after the where clause: {FROM_ORDER}"
402            ));
403        }
404        if body.get(end) == Some(&Token::By) {
405            return Err(format!("Unexpected 'by' after 'from df': {FROM_ORDER}"));
406        }
407        found = Some((i, end));
408    }
409    let mut body = body.to_vec();
410    if let Some((start, end)) = found {
411        body.drain(start..end);
412    }
413    Ok(body)
414}
415
416/// Infix operators spelled as words (q's); ordinary identifiers elsewhere, so a column
417/// named `in` or `mod` still works at an expression's start or after `.`.
418const WORD_OPS: [&str; 5] = ["in", "like", "xbar", "mod", "wavg"];
419
420/// The infix operator at `tokens[i]`, if there is one: a symbol, or an operator
421/// word that has an operand before it.
422fn infix_op_at(tokens: &[Token], i: usize) -> Option<&str> {
423    match tokens.get(i)? {
424        Token::Op(op) => Some(op.as_str()),
425        Token::Identifier(word)
426            if i > 0 && tokens[i - 1] != Token::Dot && WORD_OPS.contains(&word.as_str()) =>
427        {
428            Some(word.as_str())
429        }
430        _ => None,
431    }
432}
433
434/// A parsed expression, before becoming a Polars expression ([`Node::to_expr`]) or
435/// Python ([`Node::python`]): one parse, so "Copy as Python" is the query datui ran.
436#[derive(Debug, Clone, PartialEq)]
437pub(crate) enum Node {
438    Col(String),
439    /// A number as typed: a float literal.
440    Num(f64),
441    /// A whole number where the operator keeps integers whole (`mod`, `xbar`).
442    Int(i64),
443    Str(String),
444    Bool(bool),
445    Null,
446    /// A `YYYY.MM.DD` literal, as ISO `YYYY-MM-DD`.
447    Date(String),
448    /// A `YYYY.MM.DDTHH:MM:SS[.fff]` literal.
449    Timestamp {
450        iso: String,
451        format: String,
452        unit: TimeUnit,
453        /// The zone of the column it meets, read as a clock there; none for a column
454        /// without one. See [`Node::resolve_time_zones`].
455        zone: Option<String>,
456    },
457    Bin(BinOp, Box<Node>, Box<Node>),
458    Coalesce(Box<Node>, Box<Node>),
459    /// The values of the first where the second holds.
460    Filter(Box<Node>, Box<Node>),
461    /// `when(condition).then(value).otherwise(other)`.
462    When(Box<Node>, Box<Node>, Box<Node>),
463    Op(Box<Node>, Op),
464    Alias(Box<Node>, String),
465}
466
467#[derive(Debug, Clone, Copy, PartialEq, Eq)]
468pub(crate) enum BinOp {
469    Add,
470    Sub,
471    Mul,
472    /// Polars' `/` on two expressions.
473    Div,
474    TrueDiv,
475    FloorDiv,
476    Rem,
477    Eq,
478    Neq,
479    Lt,
480    Gt,
481    LtEq,
482    GtEq,
483    And,
484    Or,
485}
486
487/// A method applied to one expression.
488#[derive(Debug, Clone, PartialEq)]
489pub(crate) enum Op {
490    Mean,
491    Min,
492    Max,
493    Count,
494    Std,
495    Var,
496    Median,
497    Sum,
498    First,
499    Last,
500    NUnique,
501    Not,
502    IsNull,
503    IsNotNull,
504    LenChars,
505    Upper,
506    Lower,
507    Abs,
508    Floor,
509    Ceil,
510    Sqrt,
511    Ln,
512    Exp,
513    Date,
514    Time,
515    Year,
516    Quarter,
517    Month,
518    Week,
519    Day,
520    OrdinalDay,
521    Weekday,
522    Hour,
523    Minute,
524    Second,
525    MonthStart,
526    MonthEnd,
527    DtFormat(String),
528    StartsWith(String),
529    EndsWith(String),
530    ContainsLiteral(String),
531    /// A regex match, strict.
532    ContainsRegex(String),
533    /// Split on the text and take the piece at the index, null past the last.
534    Part(String, i64),
535    Slice(i64, Option<u64>),
536    ReplaceAll(String, String),
537    Strip,
538    ToDate(Option<String>),
539    ToDatetime(Option<String>),
540    Round(u32),
541    /// A non-strict cast.
542    Cast(CastTo),
543}
544
545#[derive(Debug, Clone, Copy, PartialEq, Eq)]
546pub(crate) enum CastTo {
547    Int64,
548    Float64,
549    String,
550}
551
552impl CastTo {
553    fn dtype(self) -> DataType {
554        match self {
555            CastTo::Int64 => DataType::Int64,
556            CastTo::Float64 => DataType::Float64,
557            CastTo::String => DataType::String,
558        }
559    }
560}
561
562/// The most nodes an expression may grow to: `wavg`, `xbar` and `in` repeat operands,
563/// so nesting multiplies (found by the `parse_query` fuzz target).
564const MAX_EXPR_NODES: usize = 10_000;
565
566/// Err when `copies` of `node` would pass [`MAX_EXPR_NODES`].
567fn check_copies(node: &Node, copies: usize) -> Result<(), String> {
568    if node.size().saturating_mul(copies) > MAX_EXPR_NODES {
569        return Err(
570            "Expression is too large: nested wavg, xbar or in repeat what they are \
571                    given. Simplify it or split it into steps."
572                .to_string(),
573        );
574    }
575    Ok(())
576}
577
578impl Node {
579    /// How many nodes the tree has.
580    fn size(&self) -> usize {
581        1 + match self {
582            Node::Col(_)
583            | Node::Num(_)
584            | Node::Int(_)
585            | Node::Str(_)
586            | Node::Bool(_)
587            | Node::Null
588            | Node::Date(_)
589            | Node::Timestamp { .. } => 0,
590            Node::Bin(_, a, b) | Node::Coalesce(a, b) | Node::Filter(a, b) => a.size() + b.size(),
591            Node::When(a, b, c) => a.size() + b.size() + c.size(),
592            Node::Op(a, _) | Node::Alias(a, _) => a.size(),
593        }
594    }
595
596    fn op(self, op: Op) -> Node {
597        Node::Op(Box::new(self), op)
598    }
599
600    fn bin(self, op: BinOp, right: Node) -> Node {
601        Node::Bin(op, Box::new(self), Box::new(right))
602    }
603
604    fn alias(self, name: impl Into<String>) -> Node {
605        Node::Alias(Box::new(self), name.into())
606    }
607
608    fn cast_text(self) -> Node {
609        self.op(Op::Cast(CastTo::String))
610    }
611
612    /// The Polars expression.
613    pub(crate) fn to_expr(&self) -> Expr {
614        match self {
615            Node::Col(name) => col(name),
616            Node::Num(n) => lit(*n),
617            Node::Int(n) => lit(*n),
618            Node::Str(s) => lit(s.as_str()),
619            Node::Bool(b) => lit(*b),
620            Node::Null => lit(NULL),
621            Node::Date(iso) => {
622                let opts = StrptimeOptions {
623                    format: Some("%Y-%m-%d".into()),
624                    ..Default::default()
625                };
626                lit(iso.as_str()).str().to_date(opts)
627            }
628            Node::Timestamp {
629                iso,
630                format,
631                unit,
632                zone,
633            } => {
634                let opts = StrptimeOptions {
635                    format: Some(format.as_str().into()),
636                    ..Default::default()
637                };
638                // Set only from a column's own dtype, so it parses.
639                let zone = TimeZone::opt_try_new(zone.as_deref()).ok().flatten();
640                // A clock time a fall back repeats is its first instant.
641                lit(iso.as_str())
642                    .str()
643                    .to_datetime(Some(*unit), zone, opts, lit("earliest"))
644            }
645            Node::Bin(op, left, right) => {
646                let (left, right) = (left.to_expr(), right.to_expr());
647                match op {
648                    BinOp::Add => left.add(right),
649                    BinOp::Sub => left.sub(right),
650                    BinOp::Mul => left.mul(right),
651                    BinOp::Div => left.div(right),
652                    BinOp::TrueDiv => left.true_div(right),
653                    BinOp::FloorDiv => left.floor_div(right),
654                    BinOp::Rem => left.rem(right),
655                    BinOp::Eq => left.eq(right),
656                    BinOp::Neq => left.neq(right),
657                    BinOp::Lt => left.lt(right),
658                    BinOp::Gt => left.gt(right),
659                    BinOp::LtEq => left.lt_eq(right),
660                    BinOp::GtEq => left.gt_eq(right),
661                    BinOp::And => left.and(right),
662                    BinOp::Or => left.or(right),
663                }
664            }
665            Node::Coalesce(left, right) => coalesce(&[left.to_expr(), right.to_expr()]),
666            Node::Filter(values, predicate) => values.to_expr().filter(predicate.to_expr()),
667            Node::When(condition, then, otherwise) => when(condition.to_expr())
668                .then(then.to_expr())
669                .otherwise(otherwise.to_expr()),
670            Node::Op(inner, op) => apply_op_expr(inner.to_expr(), op),
671            Node::Alias(inner, name) => inner.to_expr().alias(name.as_str()),
672        }
673    }
674
675    /// The same node with every alias inside it taken off, as Polars' `undo_aliases`.
676    pub(crate) fn without_aliases(&self) -> Node {
677        let strip = |n: &Node| Box::new(n.without_aliases());
678        match self {
679            Node::Alias(inner, _) => inner.without_aliases(),
680            Node::Bin(op, l, r) => Node::Bin(*op, strip(l), strip(r)),
681            Node::Coalesce(l, r) => Node::Coalesce(strip(l), strip(r)),
682            Node::Filter(v, p) => Node::Filter(strip(v), strip(p)),
683            Node::When(c, t, o) => Node::When(strip(c), strip(t), strip(o)),
684            Node::Op(inner, op) => Node::Op(strip(inner), op.clone()),
685            leaf => leaf.clone(),
686        }
687    }
688
689    /// Name each `/` as Polars runs it over `schema`: `Div` floor-divides integers and
690    /// divides otherwise, so Python needs `//` or `/` by the quotient's type.
691    pub(crate) fn resolve_division(&mut self, schema: &Schema) {
692        match self {
693            Node::Bin(_, left, right) | Node::Coalesce(left, right) | Node::Filter(left, right) => {
694                left.resolve_division(schema);
695                right.resolve_division(schema);
696            }
697            Node::When(c, t, o) => {
698                c.resolve_division(schema);
699                t.resolve_division(schema);
700                o.resolve_division(schema);
701            }
702            Node::Op(inner, _) | Node::Alias(inner, _) => inner.resolve_division(schema),
703            _ => {}
704        }
705        if let Node::Bin(BinOp::Div, ..) = self {
706            let quotient = DataFrame::empty_with_schema(schema)
707                .lazy()
708                .select([self.to_expr()])
709                .collect_schema()
710                .ok()
711                .and_then(|s| s.get_at_index(0).map(|(_, dtype)| dtype.is_integer()));
712            if let (Some(whole), Node::Bin(op, ..)) = (quotient, self) {
713                *op = if whole {
714                    BinOp::FloorDiv
715                } else {
716                    BinOp::TrueDiv
717                };
718            }
719        }
720    }
721
722    /// Each timestamp literal meeting a zoned column (comparison, arithmetic, `^`, `?`
723    /// branches) takes that zone; Polars refuses zoned-naive comparisons.
724    pub(crate) fn resolve_time_zones(&mut self, schema: &Schema) {
725        match self {
726            Node::Bin(_, left, right) | Node::Coalesce(left, right) => {
727                left.resolve_time_zones(schema);
728                right.resolve_time_zones(schema);
729                Self::share_zone(left, right, schema);
730            }
731            Node::Filter(values, predicate) => {
732                values.resolve_time_zones(schema);
733                predicate.resolve_time_zones(schema);
734            }
735            Node::When(c, t, o) => {
736                c.resolve_time_zones(schema);
737                t.resolve_time_zones(schema);
738                o.resolve_time_zones(schema);
739                Self::share_zone(t, o, schema);
740            }
741            Node::Op(inner, _) | Node::Alias(inner, _) => inner.resolve_time_zones(schema),
742            _ => {}
743        }
744    }
745
746    /// An error for the first comparison of a typed temporal column with quoted text (a
747    /// string in q, never a date), worded in q's terms.
748    fn check_quoted_temporal(&self, schema: &Schema) -> Result<(), String> {
749        match self {
750            Node::Bin(op, left, right) => {
751                left.check_quoted_temporal(schema)?;
752                right.check_quoted_temporal(schema)?;
753                let compares = matches!(
754                    op,
755                    BinOp::Eq | BinOp::Neq | BinOp::Lt | BinOp::Gt | BinOp::LtEq | BinOp::GtEq
756                );
757                if compares
758                    && let Some(err) = quoted_temporal(left, right, schema)
759                        .or_else(|| quoted_temporal(right, left, schema))
760                {
761                    return Err(err);
762                }
763            }
764            Node::Coalesce(left, right) | Node::Filter(left, right) => {
765                left.check_quoted_temporal(schema)?;
766                right.check_quoted_temporal(schema)?;
767            }
768            Node::When(c, t, o) => {
769                c.check_quoted_temporal(schema)?;
770                t.check_quoted_temporal(schema)?;
771                o.check_quoted_temporal(schema)?;
772            }
773            Node::Op(inner, _) | Node::Alias(inner, _) => inner.check_quoted_temporal(schema)?,
774            _ => {}
775        }
776        Ok(())
777    }
778
779    /// Give a zoneless timestamp literal on one side the zone of the other side's type.
780    fn share_zone(a: &mut Node, b: &mut Node, schema: &Schema) {
781        if !Self::take_zone(a, b, schema) {
782            Self::take_zone(b, a, schema);
783        }
784    }
785
786    /// Whether `literal`, a zoneless timestamp literal, took the zone of `other`'s type.
787    fn take_zone(literal: &mut Node, other: &Node, schema: &Schema) -> bool {
788        if let Node::Timestamp {
789            zone: zone @ None, ..
790        } = literal
791            && let Some(DataType::Datetime(_, Some(tz))) = other.dtype(schema)
792        {
793            *zone = Some(tz.to_string());
794            return true;
795        }
796        false
797    }
798
799    /// The type the expression has over `schema`, when Polars can say.
800    fn dtype(&self, schema: &Schema) -> Option<DataType> {
801        DataFrame::empty_with_schema(schema)
802            .lazy()
803            .select([self.to_expr()])
804            .collect_schema()
805            .ok()
806            .and_then(|s| s.get_at_index(0).map(|(_, dtype)| dtype.clone()))
807    }
808
809    /// The literal as Python, bare: `1.0`, `"a"`, `True`, `None`.
810    fn python_literal(&self) -> Option<String> {
811        Some(match self {
812            Node::Num(n) => crate::export::python_script::py_float(*n),
813            Node::Int(n) => n.to_string(),
814            Node::Str(s) => crate::export::python_script::py_str(s),
815            Node::Bool(b) => crate::export::python_script::py_bool(*b).to_string(),
816            Node::Null => "None".to_string(),
817            _ => return None,
818        })
819    }
820
821    /// Python Polars code for the expression.
822    pub(crate) fn python(&self) -> String {
823        use crate::export::python_script::py_str;
824        if let Some(literal) = self.python_literal() {
825            return format!("pl.lit({literal})");
826        }
827        match self {
828            Node::Col(name) => format!("pl.col({})", py_str(name)),
829            Node::Date(iso) => {
830                let parts: Vec<String> = iso
831                    .split('-')
832                    .map(|p| p.trim_start_matches('0').to_string())
833                    .map(|p| if p.is_empty() { "0".to_string() } else { p })
834                    .collect();
835                format!("pl.date({})", parts.join(", "))
836            }
837            Node::Timestamp {
838                iso,
839                format,
840                unit,
841                zone,
842            } => format!(
843                "pl.lit({}).str.to_datetime({}, time_unit={}{})",
844                py_str(iso),
845                py_str(format),
846                py_str(time_unit_name(*unit)),
847                zone.as_ref().map_or(String::new(), |zone| format!(
848                    ", time_zone={}, ambiguous=\"earliest\"",
849                    py_str(zone)
850                ))
851            ),
852            Node::Bin(op, left, right) => {
853                // A literal on the right stays bare (`pl.col("a") > 1.0`); Python's
854                // operators turn it into one.
855                let right = match right.python_literal() {
856                    Some(literal) => literal,
857                    None => right.python_operand(),
858                };
859                format!("{} {} {}", left.python_operand(), op.python(), right)
860            }
861            Node::Coalesce(left, right) => {
862                format!("pl.coalesce({}, {})", left.python(), right.python())
863            }
864            Node::Filter(values, predicate) => {
865                format!("{}.filter({})", values.python_operand(), predicate.python())
866            }
867            Node::When(condition, then, otherwise) => format!(
868                "pl.when({}).then({}).otherwise({})",
869                condition.python(),
870                then.python(),
871                otherwise.python()
872            ),
873            Node::Op(inner, op) => format!("{}{}", inner.python_operand(), op.python()),
874            Node::Alias(inner, name) => {
875                // Only the outer name survives an alias of an alias.
876                let mut inner = inner.as_ref();
877                while let Node::Alias(deeper, _) = inner {
878                    inner = deeper;
879                }
880                format!("{}.alias({})", inner.python_operand(), py_str(name))
881            }
882            _ => unreachable!("literals return above"),
883        }
884    }
885
886    /// As [`Self::python`], parenthesized where an operator or a method after it
887    /// would otherwise bind to part of it.
888    fn python_operand(&self) -> String {
889        match self {
890            Node::Bin(..) => format!("({})", self.python()),
891            _ => self.python(),
892        }
893    }
894}
895
896fn time_unit_name(unit: TimeUnit) -> &'static str {
897    match unit {
898        TimeUnit::Milliseconds => "ms",
899        TimeUnit::Microseconds => "us",
900        TimeUnit::Nanoseconds => "ns",
901    }
902}
903
904impl BinOp {
905    fn python(self) -> &'static str {
906        match self {
907            BinOp::Add => "+",
908            BinOp::Sub => "-",
909            BinOp::Mul => "*",
910            BinOp::Div | BinOp::TrueDiv => "/",
911            BinOp::FloorDiv => "//",
912            BinOp::Rem => "%",
913            BinOp::Eq => "==",
914            BinOp::Neq => "!=",
915            BinOp::Lt => "<",
916            BinOp::Gt => ">",
917            BinOp::LtEq => "<=",
918            BinOp::GtEq => ">=",
919            BinOp::And => "&",
920            BinOp::Or => "|",
921        }
922    }
923}
924
925impl Op {
926    /// The method call, from its dot: `.str.to_uppercase()`.
927    fn python(&self) -> String {
928        use crate::export::python_script::py_str;
929        let fixed = match self {
930            Op::Mean => ".mean()",
931            Op::Min => ".min()",
932            Op::Max => ".max()",
933            Op::Count => ".count()",
934            Op::Std => ".std()",
935            Op::Var => ".var()",
936            Op::Median => ".median()",
937            Op::Sum => ".sum()",
938            Op::First => ".first()",
939            Op::Last => ".last()",
940            Op::NUnique => ".n_unique()",
941            Op::Not => ".not_()",
942            Op::IsNull => ".is_null()",
943            Op::IsNotNull => ".is_not_null()",
944            Op::LenChars => ".str.len_chars()",
945            Op::Upper => ".str.to_uppercase()",
946            Op::Lower => ".str.to_lowercase()",
947            Op::Abs => ".abs()",
948            Op::Floor => ".floor()",
949            Op::Ceil => ".ceil()",
950            Op::Sqrt => ".sqrt()",
951            Op::Ln => ".log()",
952            Op::Exp => ".exp()",
953            Op::Date => ".dt.date()",
954            Op::Time => ".dt.time()",
955            Op::Year => ".dt.year()",
956            Op::Quarter => ".dt.quarter()",
957            Op::Month => ".dt.month()",
958            Op::Week => ".dt.week()",
959            Op::Day => ".dt.day()",
960            Op::OrdinalDay => ".dt.ordinal_day()",
961            Op::Weekday => ".dt.weekday()",
962            Op::Hour => ".dt.hour()",
963            Op::Minute => ".dt.minute()",
964            Op::Second => ".dt.second()",
965            Op::MonthStart => ".dt.month_start()",
966            Op::MonthEnd => ".dt.month_end()",
967            Op::Strip => ".str.strip_chars()",
968            Op::Cast(CastTo::Int64) => ".cast(pl.Int64, strict=False)",
969            Op::Cast(CastTo::Float64) => ".cast(pl.Float64, strict=False)",
970            Op::Cast(CastTo::String) => ".cast(pl.String)",
971            _ => "",
972        };
973        if !fixed.is_empty() {
974            return fixed.to_string();
975        }
976        let format_arg = |format: &Option<String>| match format {
977            Some(f) => format!("{}, strict=False", py_str(f)),
978            None => "strict=False".to_string(),
979        };
980        match self {
981            Op::DtFormat(f) => format!(".dt.to_string({})", py_str(f)),
982            Op::StartsWith(s) => format!(".str.starts_with({})", py_str(s)),
983            Op::EndsWith(s) => format!(".str.ends_with({})", py_str(s)),
984            Op::ContainsLiteral(s) => format!(".str.contains({}, literal=True)", py_str(s)),
985            Op::ContainsRegex(r) => format!(".str.contains({})", py_str(r)),
986            Op::Part(sep, i) => format!(
987                ".str.split({}).list.get({i}, null_on_oob=True)",
988                py_str(sep)
989            ),
990            Op::Slice(start, Some(len)) => format!(".str.slice({start}, {len})"),
991            Op::Slice(start, None) => format!(".str.slice({start})"),
992            Op::ReplaceAll(from, to) => format!(
993                ".str.replace_all({}, {}, literal=True)",
994                py_str(from),
995                py_str(to)
996            ),
997            Op::ToDate(format) => format!(".str.to_date({})", format_arg(format)),
998            Op::ToDatetime(format) => format!(".str.to_datetime({})", format_arg(format)),
999            Op::Round(d) => format!(".round({d}, mode=\"half_away_from_zero\")"),
1000            _ => unreachable!("fixed calls return above"),
1001        }
1002    }
1003}
1004
1005fn apply_op_expr(expr: Expr, op: &Op) -> Expr {
1006    let strptime = |format: &Option<String>| StrptimeOptions {
1007        format: format.as_deref().map(Into::into),
1008        // A value that does not match becomes null, as a failed parse does in q,
1009        // rather than one stray row failing the whole query.
1010        strict: false,
1011        ..Default::default()
1012    };
1013    match op {
1014        Op::Mean => expr.mean(),
1015        Op::Min => expr.min(),
1016        Op::Max => expr.max(),
1017        Op::Count => expr.count(),
1018        // Sample statistics (n - 1), like `std`; q's own var and dev divide by n.
1019        Op::Std => expr.std(1),
1020        Op::Var => expr.var(1),
1021        Op::Median => expr.median(),
1022        Op::Sum => expr.sum(),
1023        Op::First => expr.first(),
1024        Op::Last => expr.last(),
1025        Op::NUnique => expr.n_unique(),
1026        Op::Not => expr.not(),
1027        Op::IsNull => expr.is_null(),
1028        Op::IsNotNull => expr.is_not_null(),
1029        Op::LenChars => expr.str().len_chars(),
1030        Op::Upper => expr.str().to_uppercase(),
1031        Op::Lower => expr.str().to_lowercase(),
1032        Op::Abs => expr.abs(),
1033        Op::Floor => expr.floor(),
1034        Op::Ceil => expr.ceil(),
1035        Op::Sqrt => expr.sqrt(),
1036        Op::Ln => expr.log(lit(std::f64::consts::E)),
1037        Op::Exp => expr.exp(),
1038        Op::Date => expr.dt().date(),
1039        Op::Time => expr.dt().time(),
1040        Op::Year => expr.dt().year(),
1041        Op::Quarter => expr.dt().quarter(),
1042        Op::Month => expr.dt().month(),
1043        Op::Week => expr.dt().week(),
1044        Op::Day => expr.dt().day(),
1045        Op::OrdinalDay => expr.dt().ordinal_day(),
1046        Op::Weekday => expr.dt().weekday(),
1047        Op::Hour => expr.dt().hour(),
1048        Op::Minute => expr.dt().minute(),
1049        Op::Second => expr.dt().second(),
1050        Op::MonthStart => expr.dt().month_start(),
1051        Op::MonthEnd => expr.dt().month_end(),
1052        Op::DtFormat(f) => expr.dt().to_string(f),
1053        Op::StartsWith(s) => expr.str().starts_with(lit(s.as_str())),
1054        Op::EndsWith(s) => expr.str().ends_with(lit(s.as_str())),
1055        Op::ContainsLiteral(s) => expr.str().contains_literal(lit(s.as_str())),
1056        Op::ContainsRegex(r) => expr.str().contains(lit(r.as_str()), true),
1057        // Past the last piece is null, not an error; a negative index counts from the end.
1058        Op::Part(sep, i) => expr
1059            .str()
1060            .split(lit(sep.as_str()))
1061            .list()
1062            .get(lit(*i), true),
1063        Op::Slice(start, length) => {
1064            // No length: to the end of the string.
1065            let length = length.map_or_else(|| lit(NULL), lit);
1066            expr.str().slice(lit(*start), length)
1067        }
1068        Op::ReplaceAll(from, to) => {
1069            expr.str()
1070                .replace_all(lit(from.as_str()), lit(to.as_str()), true)
1071        }
1072        Op::Strip => expr.str().strip_chars(lit(NULL)),
1073        Op::ToDate(format) => expr.str().to_date(strptime(format)),
1074        Op::ToDatetime(format) => {
1075            expr.str()
1076                .to_datetime(None, None, strptime(format), lit("raise"))
1077        }
1078        // Half away from zero, the rounding people expect from a calculator or SQL.
1079        Op::Round(decimals) => expr.round(*decimals, RoundMode::HalfAwayFromZero),
1080        // Non-strict casts: a value that does not convert becomes null.
1081        Op::Cast(to) => expr.cast(to.dtype()),
1082    }
1083}
1084
1085/// An operand of `mod` or `xbar`: whole numbers as integer literals, so integer columns
1086/// keep their type (`5 xbar passenger_count` stays Int64).
1087fn int_or_node(tokens: &[Token]) -> Result<Node, String> {
1088    let whole = |n: f64| n.fract() == 0.0 && n.abs() < i64::MAX as f64;
1089    match tokens {
1090        [Token::Number(n)] if whole(*n) => Ok(Node::Int(*n as i64)),
1091        [Token::Op(minus), Token::Number(n)] if minus == "-" && whole(*n) => {
1092            Ok(Node::Int(-(*n as i64)))
1093        }
1094        _ => parse_node(tokens),
1095    }
1096}
1097
1098/// True when no `]` in `tokens` closes a `[` from outside them.
1099fn brackets_balanced(tokens: &[Token]) -> bool {
1100    let mut depth = 0usize;
1101    tokens.iter().all(|t| match t {
1102        Token::LBracket => {
1103            depth += 1;
1104            true
1105        }
1106        Token::RBracket => depth.checked_sub(1).map(|d| depth = d).is_some(),
1107        _ => true,
1108    })
1109}
1110
1111/// The error for `column` of a temporal type compared with the quoted `text`, naming
1112/// the literal the type takes: the text itself when it is one once unquoted.
1113fn quoted_temporal(column: &Node, text: &Node, schema: &Schema) -> Option<String> {
1114    let (Node::Col(name), Node::Str(s)) = (column, text) else {
1115        return None;
1116    };
1117    let shown = q_name(name);
1118    let unquoted = tokenize(s).ok();
1119    let literal = |is_kind: fn(&Token) -> bool, example: &str| match unquoted.as_deref() {
1120        Some([token]) if is_kind(token) => s.trim().to_string(),
1121        _ => example.to_string(),
1122    };
1123    let (kind, remedy) = match schema.get(name)? {
1124        DataType::Date => (
1125            "date",
1126            format!(
1127                "A date is {}",
1128                literal(|t| matches!(t, Token::DateLiteral(_)), "2024.01.01")
1129            ),
1130        ),
1131        DataType::Datetime(..) => (
1132            "timestamp",
1133            format!(
1134                "A timestamp is {}",
1135                literal(
1136                    |t| matches!(t, Token::TimestampLiteral { .. }),
1137                    "2024.01.01T05:00:00"
1138                )
1139            ),
1140        ),
1141        DataType::Time => (
1142            "time",
1143            format!(
1144                "A time has no literal; compare {shown}.hour, {shown}.minute or {shown}.second with a number"
1145            ),
1146        ),
1147        DataType::Duration(_) => ("duration", "A duration has no literal".to_string()),
1148        _ => return None,
1149    };
1150    Some(format!(
1151        "{shown} is a {kind}; \"{s}\" is a string. {remedy}"
1152    ))
1153}
1154
1155/// How q spells a column: bare when it can be, else `col["first name"]`.
1156pub(crate) fn q_name(name: &str) -> String {
1157    if is_plain_name(name) {
1158        name.to_string()
1159    } else {
1160        format!("col[\"{name}\"]")
1161    }
1162}
1163
1164/// Whether `name` reads as a column when typed bare, rather than needing `col["…"]`.
1165fn is_plain_name(name: &str) -> bool {
1166    let mut chars = name.chars();
1167    chars.next().is_some_and(|c| c.is_alphabetic() || c == '_')
1168        && chars.all(|c| c.is_alphanumeric() || c == '_')
1169        && !matches!(name, "select" | "where" | "by")
1170}
1171
1172/// OR of the conditions as a balanced tree, so a long `in` list nests
1173/// logarithmically rather than one level per element.
1174fn any_of(mut conditions: Vec<Node>) -> Node {
1175    if conditions.len() <= 1 {
1176        return conditions.pop().unwrap_or(Node::Bool(false));
1177    }
1178    let right = conditions.split_off(conditions.len() / 2);
1179    any_of(conditions).bin(BinOp::Or, any_of(right))
1180}
1181
1182/// A `like` pattern as an anchored regex: `*` is any run of characters, `?` any
1183/// one character, everything else literal.
1184fn like_regex(pattern: &str) -> String {
1185    let mut re = String::from("(?s)^");
1186    for c in pattern.chars() {
1187        match c {
1188            '*' => re.push_str(".*"),
1189            '?' => re.push('.'),
1190            _ => re.push_str(&regex::escape(c.encode_utf8(&mut [0; 4]))),
1191        }
1192    }
1193    re.push('$');
1194    re
1195}
1196
1197/// Operators whose right side is not an ordinary expression, or whose operands
1198/// need their tokens (literal checks, names). Everything else goes to `apply_op`.
1199fn apply_infix(left_tokens: &[Token], op: &str, right_tokens: &[Token]) -> Result<Node, String> {
1200    match op {
1201        "in" => {
1202            let list = match right_tokens {
1203                // One list: the first `[` closes at the last token, so `[1] + [2]` is not.
1204                [Token::LBracket, inner @ .., Token::RBracket] if brackets_balanced(inner) => inner,
1205                _ => {
1206                    return Err(
1207                        "in takes a list on its right, e.g. name in [\"Emma\", \"Olivia\"]"
1208                            .to_string(),
1209                    );
1210                }
1211            };
1212            let items = split_tokens(list, &Token::Comma);
1213            if items.iter().any(|item| item.is_empty()) {
1214                return Err(
1215                    "in needs a list of values, e.g. name in [\"Emma\", \"Olivia\"]".to_string(),
1216                );
1217            }
1218            let left = parse_node(left_tokens)?;
1219            // A column or literal repeated grows only as the list typed does; a larger
1220            // left side repeated per item is what multiplies.
1221            if left.size() > 1 {
1222                check_copies(&left, items.len())?;
1223            }
1224            // One `=` per value, so each value compares exactly as `x = value` would,
1225            // with the same literal casting (numbers, dates, timestamps).
1226            let conditions = items
1227                .iter()
1228                .map(|item| Ok(left.clone().bin(BinOp::Eq, parse_node(item)?)))
1229                .collect::<Result<Vec<_>, String>>()?;
1230            Ok(any_of(conditions))
1231        }
1232        "like" => {
1233            let [Token::String(pattern)] = right_tokens else {
1234                return Err(
1235                    "like takes a quoted pattern on its right, e.g. item like \"*Chicken*\""
1236                        .to_string(),
1237                );
1238            };
1239            let left = parse_node(left_tokens)?;
1240            // Cast first so numeric codes (zip, station ids read as numbers) match too.
1241            Ok(left.cast_text().op(Op::ContainsRegex(like_regex(pattern))))
1242        }
1243        "xbar" => {
1244            if let [Token::Number(n)] = left_tokens
1245                && *n <= 0.0
1246            {
1247                return Err(
1248                    "xbar needs a positive bucket size, e.g. 5 xbar fare_amount".to_string()
1249                );
1250            }
1251            let right = parse_node(right_tokens)?;
1252            let size = int_or_node(left_tokens)?;
1253            check_copies(&size, 2)?;
1254            // floor_div floors toward negative infinity for both ints and floats,
1255            // which is what makes every value land in the bucket at or below it.
1256            Ok(right
1257                .bin(BinOp::FloorDiv, size.clone())
1258                .bin(BinOp::Mul, size))
1259        }
1260        "mod" => {
1261            let right = int_or_node(right_tokens)?;
1262            let left = int_or_node(left_tokens)?;
1263            Ok(left.bin(BinOp::Rem, right))
1264        }
1265        "wavg" => {
1266            let values = parse_node(right_tokens)?;
1267            let weights = parse_node(left_tokens)?;
1268            // Below, weights appear five times and values three.
1269            check_copies(&weights, 5)?;
1270            check_copies(&values, 3)?;
1271            let weighted = weights.clone().bin(BinOp::Mul, values);
1272            // Only pairs with both a weight and a value count toward the total weight;
1273            // a null value would otherwise still pull the average toward zero.
1274            let total = Node::Filter(
1275                Box::new(weights),
1276                Box::new(weighted.clone().op(Op::IsNotNull)),
1277            )
1278            .op(Op::Sum);
1279            // No weight at all (every pair null) is no average, not 0/0 = NaN.
1280            let total = Node::When(
1281                Box::new(total.clone().bin(BinOp::Neq, Node::Int(0))),
1282                Box::new(total),
1283                Box::new(Node::Null),
1284            );
1285            let node = weighted.op(Op::Sum).bin(BinOp::TrueDiv, total);
1286            Ok(match simple_column_name(right_tokens) {
1287                Some(column) => node.alias(format!("wavg_{}", column)),
1288                None => node,
1289            })
1290        }
1291        _ => {
1292            // Parse right side first (right-to-left evaluation): it holds any
1293            // remaining operators, so c>c%n becomes c > (c%n).
1294            let right = parse_node(right_tokens)?;
1295            let left = parse_node(left_tokens)?;
1296            apply_op(left, op, right)
1297        }
1298    }
1299}
1300
1301fn apply_op(left: Node, op: &str, right: Node) -> Result<Node, String> {
1302    let op = match op {
1303        "+" => BinOp::Add,
1304        "-" => BinOp::Sub,
1305        "*" => BinOp::Mul,
1306        // `%` divides (q heritage); `/` is the alias everyone expects.
1307        "%" | "/" => BinOp::Div,
1308        "^" => return Ok(Node::Coalesce(Box::new(left), Box::new(right))),
1309        "=" => BinOp::Eq,
1310        "<" => BinOp::Lt,
1311        ">" => BinOp::Gt,
1312        "<=" => BinOp::LtEq,
1313        ">=" => BinOp::GtEq,
1314        "<>" | "!=" => BinOp::Neq,
1315        _ => return Err(format!("Unknown operator: {}", op)),
1316    };
1317    Ok(left.bin(op, right))
1318}
1319
1320/// The column name when the tokens are a bare column reference: `salary`, or
1321/// `col["first name"]` / `col[name]`. Anything more (literals, operators) is None.
1322fn simple_column_name(tokens: &[Token]) -> Option<String> {
1323    match tokens {
1324        [Token::Identifier(name)] => Some(name.clone()),
1325        [
1326            Token::Identifier(c),
1327            Token::LBracket,
1328            Token::String(name) | Token::Identifier(name),
1329            Token::RBracket,
1330        ] if c == "col" => Some(name.clone()),
1331        _ => None,
1332    }
1333}
1334
1335const WAVG_USAGE: &str = "wavg goes between weights and values, e.g. passengers wavg fare";
1336
1337/// Aggregation function names, lowercase.
1338const AGG_FUNCTIONS: [&str; 16] = [
1339    "avg", "mean", "min", "max", "count", "std", "stddev", "dev", "var", "med", "median", "sum",
1340    "first", "last", "nunique", "wavg",
1341];
1342
1343/// Scalar function names, lowercase.
1344const SCALAR_FUNCTIONS: [&str; 13] = [
1345    "len", "length", "not", "null", "upper", "lower", "abs", "floor", "ceil", "ceiling", "sqrt",
1346    "log", "exp",
1347];
1348
1349fn is_agg_function(name: &str) -> bool {
1350    AGG_FUNCTIONS.contains(&name.to_lowercase().as_str())
1351}
1352
1353fn is_function_name(name: &str) -> bool {
1354    // `wavg` is infix (`w wavg x`), so it never opens an expression.
1355    let name = name.to_lowercase();
1356    name != "wavg" && (is_agg_function(&name) || SCALAR_FUNCTIONS.contains(&name.as_str()))
1357}
1358
1359/// A call, `fn[args]` or `fn args`. The name is checked first so arguments parse once
1360/// (trying aggregates then scalars doubled work per nesting level).
1361fn parse_call(name: &str, args: &[Token]) -> Result<Node, String> {
1362    if is_agg_function(name) {
1363        parse_agg_function(name, args)
1364    } else {
1365        parse_function(name, args)
1366    }
1367}
1368
1369/// An aggregate call: `avg[a]`, `min[b]`.
1370fn parse_agg_function(name: &str, args: &[Token]) -> Result<Node, String> {
1371    if args.is_empty() {
1372        return Err(format!(
1373            "Aggregation function {} requires an argument",
1374            name
1375        ));
1376    }
1377    let fn_name = name.to_lowercase();
1378    if fn_name == "wavg" {
1379        return Err(WAVG_USAGE.to_string());
1380    }
1381    let node = parse_node(args)?;
1382    let op = match fn_name.as_str() {
1383        "avg" | "mean" => Op::Mean,
1384        "min" => Op::Min,
1385        "max" => Op::Max,
1386        "count" => Op::Count,
1387        "std" | "stddev" | "dev" => Op::Std,
1388        "var" => Op::Var,
1389        "med" | "median" => Op::Median,
1390        "sum" => Op::Sum,
1391        "first" => Op::First,
1392        "last" => Op::Last,
1393        "nunique" => Op::NUnique,
1394        _ => return Err(format!("Unknown aggregation function: {}", name)),
1395    };
1396    let node = node.op(op);
1397    // A bare-column aggregate is named `{fn}_{column}` so two aggregates of one column do
1398    // not collide; an explicit alias overrides it later.
1399    match simple_column_name(args) {
1400        Some(column) => Ok(node.alias(format!("{}_{}", fn_name, column))),
1401        None => Ok(node),
1402    }
1403}
1404
1405/// A scalar call: `not[a=b]`, `null[col]`, `len[x]`, `upper[x]`.
1406fn parse_function(name: &str, args: &[Token]) -> Result<Node, String> {
1407    if args.is_empty() {
1408        return Err(format!("Function {} requires an argument", name));
1409    }
1410    let name_lower = name.to_lowercase();
1411    if !SCALAR_FUNCTIONS.contains(&name_lower.as_str()) {
1412        return Err(format!("Unknown function: {}", name));
1413    }
1414    let node = parse_node(args)?;
1415    let op = match name_lower.as_str() {
1416        "not" => Op::Not,
1417        "null" => Op::IsNull,
1418        "len" | "length" => Op::LenChars,
1419        "upper" => Op::Upper,
1420        "lower" => Op::Lower,
1421        "abs" => Op::Abs,
1422        "floor" => Op::Floor,
1423        "ceil" | "ceiling" => Op::Ceil,
1424        "sqrt" => Op::Sqrt,
1425        "log" => Op::Ln,
1426        "exp" => Op::Exp,
1427        _ => return Err(format!("Unknown function: {}", name)),
1428    };
1429    Ok(node.op(op))
1430}
1431
1432/// An accessor's bracketed argument: `.part["-", 0]` has a string and a number.
1433#[derive(Debug, Clone, PartialEq)]
1434enum AccessorArg {
1435    Str(String),
1436    Num(f64),
1437}
1438
1439impl AccessorArg {
1440    /// As it goes into the result's auto-alias: `part_-_0`, not `part_-_0.0`.
1441    fn alias_text(&self) -> String {
1442        match self {
1443            AccessorArg::Str(s) => s.clone(),
1444            AccessorArg::Num(n) => n.to_string(),
1445        }
1446    }
1447}
1448
1449/// Every accessor with the argument counts it takes and an example for errors.
1450const ACCESSORS: &[(&str, usize, usize, &str)] = &[
1451    // Date and time parts.
1452    ("date", 0, 0, ".date"),
1453    ("time", 0, 0, ".time"),
1454    ("year", 0, 0, ".year"),
1455    ("quarter", 0, 0, ".quarter"),
1456    ("month", 0, 0, ".month"),
1457    ("week", 0, 0, ".week"),
1458    ("day", 0, 0, ".day"),
1459    ("doy", 0, 0, ".doy"),
1460    ("dow", 0, 0, ".dow"),
1461    ("weekday", 0, 0, ".weekday"),
1462    ("hour", 0, 0, ".hour"),
1463    ("minute", 0, 0, ".minute"),
1464    ("second", 0, 0, ".second"),
1465    ("month_start", 0, 0, ".month_start"),
1466    ("month_end", 0, 0, ".month_end"),
1467    ("format", 1, 1, ".format[\"%Y-%m\"]"),
1468    // Strings.
1469    ("len", 0, 0, ".len"),
1470    ("length", 0, 0, ".length"),
1471    ("upper", 0, 0, ".upper"),
1472    ("lower", 0, 0, ".lower"),
1473    ("starts_with", 1, 1, ".starts_with[\"x\"]"),
1474    ("ends_with", 1, 1, ".ends_with[\"x\"]"),
1475    ("contains", 1, 1, ".contains[\"x\"]"),
1476    ("part", 2, 2, ".part[\"-\", 0]"),
1477    ("slice", 1, 2, ".slice[0, 4]"),
1478    ("replace", 2, 2, ".replace[\"(P)\", \"\"]"),
1479    ("strip", 0, 0, ".strip"),
1480    ("to_date", 0, 1, ".to_date[\"%Y%m%d\"]"),
1481    ("to_datetime", 0, 1, ".to_datetime[\"%Y-%m-%d %H:%M\"]"),
1482    // Numbers and casts.
1483    ("round", 0, 1, ".round[1]"),
1484    ("int", 0, 0, ".int"),
1485    ("float", 0, 0, ".float"),
1486    ("str", 0, 0, ".str"),
1487];
1488
1489/// Names for the unknown-accessor error, by kind.
1490const ACCESSOR_HELP: &str = "Valid date/time: date, time, year, quarter, month, week, day, doy, dow, hour, minute, second, month_start, month_end, format. \
1491     Valid string: len, upper, lower, starts_with, ends_with, contains, part, slice, replace, strip, to_date, to_datetime. \
1492     Valid number: round, int, float, str";
1493
1494fn arg_count_text(min: usize, max: usize) -> String {
1495    match (min, max) {
1496        (0, 0) => "no arguments".to_string(),
1497        (1, 1) => "1 argument".to_string(),
1498        (a, b) if a == b => format!("{} arguments", a),
1499        (a, b) => format!("{} to {} arguments", a, b),
1500    }
1501}
1502
1503/// Apply accessor `name` with its bracketed arguments.
1504fn apply_accessor(node: Node, accessor: &str, args: &[AccessorArg]) -> Result<Node, String> {
1505    let name = accessor.to_lowercase();
1506    let Some(&(_, min, max, usage)) = ACCESSORS.iter().find(|(n, ..)| *n == name) else {
1507        return Err(format!(
1508            "Unknown accessor: '{}'. {}",
1509            accessor, ACCESSOR_HELP
1510        ));
1511    };
1512    if args.len() < min || args.len() > max {
1513        return Err(format!(
1514            "{} takes {}, e.g. {}; got {}",
1515            name,
1516            arg_count_text(min, max),
1517            usage,
1518            args.len()
1519        ));
1520    }
1521    let text = |i: usize| match args.get(i) {
1522        Some(AccessorArg::Str(s)) => Ok(s.clone()),
1523        _ => Err(format!(
1524            "{}: argument {} must be quoted text, e.g. {}",
1525            name,
1526            i + 1,
1527            usage
1528        )),
1529    };
1530    let int = |i: usize| match args.get(i) {
1531        Some(AccessorArg::Num(n)) if n.fract() == 0.0 && n.abs() <= u32::MAX as f64 => {
1532            Ok(*n as i64)
1533        }
1534        _ => Err(format!(
1535            "{}: argument {} must be a whole number, e.g. {}",
1536            name,
1537            i + 1,
1538            usage
1539        )),
1540    };
1541    // The string pieces cast first, so they also work on numbers and dates read
1542    // as such (NOAA's DATE, an integer zip code); a cast from String is a no-op.
1543    let as_str = || node.clone().cast_text();
1544    Ok(match name.as_str() {
1545        "date" => node.op(Op::Date),
1546        "time" => node.op(Op::Time),
1547        "year" => node.op(Op::Year),
1548        "quarter" => node.op(Op::Quarter),
1549        "month" => node.op(Op::Month),
1550        "week" => node.op(Op::Week),
1551        "day" => node.op(Op::Day),
1552        "doy" => node.op(Op::OrdinalDay),
1553        "dow" | "weekday" => node.op(Op::Weekday),
1554        "hour" => node.op(Op::Hour),
1555        "minute" => node.op(Op::Minute),
1556        "second" => node.op(Op::Second),
1557        "month_start" => node.op(Op::MonthStart),
1558        "month_end" => node.op(Op::MonthEnd),
1559        "format" => node.op(Op::DtFormat(text(0)?)),
1560        "len" | "length" => node.op(Op::LenChars),
1561        "upper" => node.op(Op::Upper),
1562        "lower" => node.op(Op::Lower),
1563        "starts_with" => node.op(Op::StartsWith(text(0)?)),
1564        "ends_with" => node.op(Op::EndsWith(text(0)?)),
1565        "contains" => node.op(Op::ContainsLiteral(text(0)?)),
1566        "part" => as_str().op(Op::Part(text(0)?, int(1)?)),
1567        "slice" => {
1568            let start = int(0)?;
1569            let length = match args.len() {
1570                2 => {
1571                    let n = int(1)?;
1572                    if n < 0 {
1573                        return Err(format!(
1574                            "slice: the length cannot be negative, e.g. {}",
1575                            usage
1576                        ));
1577                    }
1578                    Some(n as u64)
1579                }
1580                // No length: to the end of the string.
1581                _ => None,
1582            };
1583            as_str().op(Op::Slice(start, length))
1584        }
1585        "replace" => as_str().op(Op::ReplaceAll(text(0)?, text(1)?)),
1586        "strip" => as_str().op(Op::Strip),
1587        "to_date" => as_str().op(Op::ToDate(args.first().map(|_| text(0)).transpose()?)),
1588        "to_datetime" => as_str().op(Op::ToDatetime(args.first().map(|_| text(0)).transpose()?)),
1589        "round" => {
1590            let decimals = if args.is_empty() { 0 } else { int(0)? };
1591            let decimals = u32::try_from(decimals)
1592                .map_err(|_| format!("round: decimals cannot be negative, e.g. {}", usage))?;
1593            node.op(Op::Round(decimals))
1594        }
1595        "int" => node.op(Op::Cast(CastTo::Int64)),
1596        "float" => node.op(Op::Cast(CastTo::Float64)),
1597        "str" => node.op(Op::Cast(CastTo::String)),
1598        _ => {
1599            return Err(format!(
1600                "Unknown accessor: '{}'. {}",
1601                accessor, ACCESSOR_HELP
1602            ));
1603        }
1604    })
1605}
1606
1607/// The arguments inside an accessor's brackets: literals separated by commas.
1608fn parse_accessor_args(accessor: &str, tokens: &[Token]) -> Result<Vec<AccessorArg>, String> {
1609    if tokens.is_empty() {
1610        return Ok(Vec::new());
1611    }
1612    split_tokens(tokens, &Token::Comma)
1613        .iter()
1614        .map(|arg| match arg.as_slice() {
1615            [Token::String(s)] | [Token::Identifier(s)] => Ok(AccessorArg::Str(s.clone())),
1616            [Token::Number(n)] => Ok(AccessorArg::Num(*n)),
1617            [Token::Op(minus), Token::Number(n)] if minus == "-" => Ok(AccessorArg::Num(-n)),
1618            _ => Err(format!(
1619                "{} takes literal arguments, quoted text or numbers, e.g. .part[\"-\", 0]",
1620                accessor
1621            )),
1622        })
1623        .collect()
1624}
1625
1626/// Parse optional dot accessors from the remaining tokens, returning
1627/// (expr_with_accessors, remaining). With `base_name`, results are aliased
1628/// `{base}_{accessor}` (chained: `{base}_{acc1}_{acc2}`) to avoid duplicates.
1629fn parse_accessors<'a>(
1630    mut expr: Node,
1631    mut tokens: &'a [Token],
1632    base_name: Option<&str>,
1633) -> Result<(Node, &'a [Token]), String> {
1634    let mut alias_suffix = String::new();
1635    while let [Token::Dot, Token::Identifier(accessor), rest @ ..] = tokens {
1636        let (args, consumed) = if rest.first() == Some(&Token::LBracket) {
1637            let mut depth = 0;
1638            let close = rest
1639                .iter()
1640                .position(|t| {
1641                    match t {
1642                        Token::LBracket => depth += 1,
1643                        Token::RBracket => depth -= 1,
1644                        _ => {}
1645                    }
1646                    depth == 0
1647                })
1648                .ok_or_else(|| format!("Unmatched bracket after .{}", accessor))?;
1649            (parse_accessor_args(accessor, &rest[1..close])?, close + 3)
1650        } else {
1651            (Vec::new(), 2)
1652        };
1653        expr = apply_accessor(expr, accessor, &args)?;
1654        if !alias_suffix.is_empty() {
1655            alias_suffix.push('_');
1656        }
1657        alias_suffix.push_str(accessor);
1658        for arg in &args {
1659            alias_suffix.push('_');
1660            alias_suffix.push_str(&arg.alias_text());
1661        }
1662        tokens = &tokens[consumed..];
1663    }
1664    if !alias_suffix.is_empty() {
1665        let alias = match base_name {
1666            Some(name) => format!("{}_{}", name, alias_suffix),
1667            None => alias_suffix,
1668        };
1669        expr = expr.alias(alias);
1670    }
1671    Ok((expr, tokens))
1672}
1673
1674fn parse_term(tokens: &[Token]) -> Result<(Node, &[Token]), String> {
1675    if tokens.is_empty() {
1676        return Err("Unexpected end of expression".to_string());
1677    }
1678    match &tokens[0] {
1679        Token::Identifier(name) => {
1680            // Check if it's col[...] syntax for column names with spaces
1681            if name == "col" && tokens.len() > 1 && tokens[1] == Token::LBracket {
1682                let mut depth = 1;
1683                let mut i = 2;
1684                while i < tokens.len() && depth > 0 {
1685                    match tokens[i] {
1686                        Token::LBracket => depth += 1,
1687                        Token::RBracket => depth -= 1,
1688                        _ => {}
1689                    }
1690                    i += 1;
1691                }
1692                if depth > 0 {
1693                    return Err("Unmatched bracket in col[]".to_string());
1694                }
1695                let col_name_tokens = &tokens[2..i - 1];
1696                if col_name_tokens.len() != 1 {
1697                    return Err("col[] must contain a single string or identifier".to_string());
1698                }
1699                let col_name = match &col_name_tokens[0] {
1700                    Token::String(s) => s.clone(),
1701                    Token::Identifier(id) => id.clone(),
1702                    _ => return Err("col[] must contain a string or identifier".to_string()),
1703                };
1704                let expr = Node::Col(col_name.clone());
1705                let (expr, remaining) = parse_accessors(expr, &tokens[i..], Some(&col_name))?;
1706                Ok((expr, remaining))
1707            }
1708            // Check if it's a function call (using square brackets)
1709            else if tokens.len() > 1 && tokens[1] == Token::LBracket {
1710                let mut depth = 1;
1711                let mut i = 2;
1712                while i < tokens.len() && depth > 0 {
1713                    match tokens[i] {
1714                        Token::LBracket => depth += 1,
1715                        Token::RBracket => depth -= 1,
1716                        _ => {}
1717                    }
1718                    i += 1;
1719                }
1720                if depth > 0 {
1721                    return Err("Unmatched bracket in function call".to_string());
1722                }
1723                let expr = parse_call(name, &tokens[2..i - 1])?;
1724                parse_accessors(expr, &tokens[i..], None)
1725            } else {
1726                // Regular column reference
1727                // (Function calls without brackets are handled in parse_node)
1728                let expr = Node::Col(name.clone());
1729                let (expr, remaining) = parse_accessors(expr, &tokens[1..], Some(name))?;
1730                Ok((expr, remaining))
1731            }
1732        }
1733        Token::Number(n) => Ok((Node::Num(*n), &tokens[1..])), // Numbers don't support accessors
1734        Token::String(s) => Ok((Node::Str(s.clone()), &tokens[1..])), // Strings don't support accessors
1735        Token::DateLiteral(iso) => Ok((Node::Date(iso.clone()), &tokens[1..])),
1736        Token::TimestampLiteral {
1737            iso,
1738            format_str,
1739            time_unit,
1740        } => Ok((
1741            Node::Timestamp {
1742                iso: iso.clone(),
1743                format: format_str.clone(),
1744                unit: *time_unit,
1745                zone: None,
1746            },
1747            &tokens[1..],
1748        )),
1749        Token::LParen => {
1750            let mut depth = 1;
1751            let mut i = 1;
1752            while i < tokens.len() && depth > 0 {
1753                match tokens[i] {
1754                    Token::LParen => depth += 1,
1755                    Token::RParen => depth -= 1,
1756                    _ => {}
1757                }
1758                i += 1;
1759            }
1760            if depth > 0 {
1761                return Err("Unmatched parenthesis".to_string());
1762            }
1763            let inner = parse_node(&tokens[1..i - 1])?;
1764            let (expr, remaining) = parse_accessors(inner, &tokens[i..], None)?;
1765            Ok((expr, remaining))
1766        }
1767        // Square brackets are only for function calls, not grouping
1768        // Parentheses are used for grouping
1769        _ => Err(format!(
1770            "Unexpected '{}' where an expression was expected",
1771            token_text(&tokens[0])
1772        )),
1773    }
1774}
1775
1776/// Deepest nesting the recursive-descent parser follows, so a long chain (`------x`,
1777/// `((((x))))`) errors instead of overflowing the stack (found by the `parse_query`
1778/// fuzz target). Each level costs about 10 KiB of stack in a debug build, so 64 is safe
1779/// on a 2 MiB worker thread and far beyond handwritten nesting.
1780const MAX_EXPR_DEPTH: u32 = 64;
1781
1782thread_local! {
1783    static EXPR_DEPTH: std::cell::Cell<u32> = const { std::cell::Cell::new(0) };
1784}
1785
1786/// Holds the recursion counter up while alive: every recursive path passes through
1787/// `parse_node`, which returns from many places, so the decrement is scoped.
1788struct DepthGuard;
1789
1790impl DepthGuard {
1791    /// `None` once the limit is reached, leaving the counter untouched.
1792    fn enter() -> Option<Self> {
1793        EXPR_DEPTH.with(|depth| {
1794            let next = depth.get() + 1;
1795            if next > MAX_EXPR_DEPTH {
1796                return None;
1797            }
1798            depth.set(next);
1799            Some(DepthGuard)
1800        })
1801    }
1802}
1803
1804impl Drop for DepthGuard {
1805    fn drop(&mut self) {
1806        EXPR_DEPTH.with(|depth| depth.set(depth.get().saturating_sub(1)));
1807    }
1808}
1809
1810// Parse expression with right-to-left operator precedence
1811// This means operators are evaluated from right to left: a+b*c is parsed as a+(b*c)
1812fn parse_node(tokens: &[Token]) -> Result<Node, String> {
1813    let Some(_depth_guard) = DepthGuard::enter() else {
1814        return Err(
1815            "Expression is nested too deeply. Simplify it or split it into steps.".to_string(),
1816        );
1817    };
1818
1819    if tokens.is_empty() {
1820        return Err("Empty expression".to_string());
1821    }
1822
1823    // First check if this starts with a function call (without brackets)
1824    // This needs to be checked before operator parsing to ensure correct precedence
1825    if let Token::Identifier(name) = &tokens[0]
1826        && is_function_name(name)
1827        && tokens.len() > 1
1828        && tokens[1] != Token::LBracket
1829    {
1830        // A call without brackets: the rest is the argument, built as the bracketed form is.
1831        return parse_call(name, &tokens[1..]);
1832    }
1833
1834    let mut op_pos = None;
1835    let mut depth = 0;
1836    let mut bracket_depth = 0;
1837
1838    // Scan from left to right to find the leftmost operator
1839    for (i, token) in tokens.iter().enumerate() {
1840        match token {
1841            Token::LParen => depth += 1,
1842            Token::RParen => depth -= 1,
1843            Token::LBracket => bracket_depth += 1,
1844            Token::RBracket => bracket_depth -= 1,
1845            _ if depth == 0 && bracket_depth == 0 && infix_op_at(tokens, i).is_some() => {
1846                op_pos = Some(i);
1847                break;
1848            }
1849            _ => {}
1850        }
1851    }
1852
1853    if let Some(pos) = op_pos {
1854        let left_tokens = &tokens[..pos];
1855        let right_tokens = &tokens[pos + 1..];
1856
1857        if let Some(op) = infix_op_at(tokens, pos) {
1858            // Unary minus next to a literal with an operator on the other side: -0.1+discount → (-0.1)+discount
1859            if left_tokens.is_empty()
1860                && op == "-"
1861                && !right_tokens.is_empty()
1862                && matches!(right_tokens[0], Token::Number(_))
1863                && let Token::Number(n) = right_tokens[0]
1864            {
1865                if right_tokens.len() >= 3
1866                    && let Some(bin_op) = infix_op_at(right_tokens, 1)
1867                {
1868                    if WORD_OPS.contains(&bin_op) {
1869                        // The word operators read their operands as tokens, so hand
1870                        // them the negative number as one: -7 mod 3 is (-7) mod 3.
1871                        return apply_infix(&[Token::Number(-n)], bin_op, &right_tokens[2..]);
1872                    }
1873                    let right = parse_node(&right_tokens[2..])?;
1874                    return apply_op(Node::Int(0).bin(BinOp::Sub, Node::Num(n)), bin_op, right);
1875                }
1876                if right_tokens.len() == 1 {
1877                    return Ok(Node::Int(0).bin(BinOp::Sub, Node::Num(n)));
1878                }
1879            }
1880            // Unary plus/minus when there is no left operand (e.g. -x, +x, -(a+b))
1881            if left_tokens.is_empty() && (op == "+" || op == "-") {
1882                let inner = parse_node(right_tokens)?;
1883                return if op == "-" {
1884                    Ok(Node::Int(0).bin(BinOp::Sub, inner))
1885                } else {
1886                    Ok(inner)
1887                };
1888            }
1889            if left_tokens.is_empty() {
1890                return Err("Missing left operand".to_string());
1891            }
1892            apply_infix(left_tokens, op, right_tokens)
1893        } else {
1894            Err("Expected operator".to_string())
1895        }
1896    } else {
1897        // No operator: a term. Callers pass complete expressions, so leftover tokens are an
1898        // error (not silently dropped, as `where x > 1 by dept` once lost `by dept`).
1899        let (expr, remaining) = parse_term(tokens)?;
1900        if let Some(extra) = remaining.first() {
1901            if matches!(&tokens[0], Token::Identifier(w) if w == "wavg") {
1902                return Err(WAVG_USAGE.to_string());
1903            }
1904            return Err(format!(
1905                "Unexpected '{}' after the expression",
1906                token_text(extra)
1907            ));
1908        }
1909        Ok(expr)
1910    }
1911}
1912
1913/// A parsed q query, ready to apply to a LazyFrame.
1914#[derive(Debug, Default)]
1915pub struct ParsedQuery {
1916    /// The select list; empty means every column.
1917    pub cols: Vec<Expr>,
1918    /// The where clause, its terms ANDed.
1919    pub filter: Option<Expr>,
1920    /// The by expressions.
1921    pub group_by: Vec<Expr>,
1922    /// Names of the by columns that have one (a plain column or an alias).
1923    pub group_by_names: Vec<String>,
1924    /// `select distinct`: drop duplicate result rows.
1925    pub distinct: bool,
1926}
1927
1928impl ParsedQuery {
1929    /// The query with text casts and date parts guarded for dates past the calendar,
1930    /// where Polars panics ([`crate::past_calendar::guard_expr`]). With `schema`, only
1931    /// date and datetime operations change.
1932    pub fn past_calendar_safe(self, schema: Option<&Schema>) -> Self {
1933        let guard = |e: Expr| crate::past_calendar::guard_expr(e, schema);
1934        Self {
1935            cols: self.cols.into_iter().map(guard).collect(),
1936            filter: self.filter.map(guard),
1937            group_by: self.group_by.into_iter().map(guard).collect(),
1938            ..self
1939        }
1940    }
1941}
1942
1943/// Convert Polars-specific error messages to user-friendly query errors.
1944pub fn sanitize_query_error(msg: &str) -> String {
1945    let msg_lower = msg.to_lowercase();
1946    if msg_lower.contains("duplicate")
1947        && (msg_lower.contains("output name") || msg_lower.contains("projection"))
1948    {
1949        let name = msg
1950            .split('\'')
1951            .nth(1)
1952            .map(|s| s.to_string())
1953            .unwrap_or_else(|| "column".to_string());
1954        return format!(
1955            "Duplicate column name '{}' in result. Use aliases to rename columns, e.g. `select my_date: timestamp.date`",
1956            name
1957        );
1958    }
1959    if msg_lower.contains(".alias(") || msg_lower.contains("try renaming") {
1960        return "Duplicate column names in result. Use aliases to rename columns, e.g. `select my_date: timestamp.date`"
1961            .to_string();
1962    }
1963    msg.to_string()
1964}
1965
1966/// A q query as parsed, before it becomes Polars expressions: what
1967/// [`parse_query`] runs and "Copy as Python" writes out.
1968#[derive(Debug, Default)]
1969pub(crate) struct QueryNodes {
1970    pub cols: Vec<Node>,
1971    pub filter: Option<Node>,
1972    pub group_by: Vec<Node>,
1973    pub group_by_names: Vec<String>,
1974    pub distinct: bool,
1975}
1976
1977impl QueryNodes {
1978    fn into_parsed(self) -> ParsedQuery {
1979        let lower = |nodes: Vec<Node>| nodes.iter().map(Node::to_expr).collect();
1980        ParsedQuery {
1981            cols: lower(self.cols),
1982            filter: self.filter.as_ref().map(Node::to_expr),
1983            group_by: lower(self.group_by),
1984            group_by_names: self.group_by_names,
1985            distinct: self.distinct,
1986        }
1987    }
1988
1989    /// Each `/` named as Polars runs it over `schema`, the data the query reads; see
1990    /// [`Node::resolve_division`].
1991    pub(crate) fn resolve_division(&mut self, schema: &Schema) {
1992        let nodes = self
1993            .cols
1994            .iter_mut()
1995            .chain(self.filter.iter_mut())
1996            .chain(self.group_by.iter_mut());
1997        for node in nodes {
1998            node.resolve_division(schema);
1999        }
2000    }
2001
2002    /// Each timestamp literal read in the zone of the column it meets in `schema`; see
2003    /// [`Node::resolve_time_zones`].
2004    pub(crate) fn resolve_time_zones(&mut self, schema: &Schema) {
2005        let nodes = self
2006            .cols
2007            .iter_mut()
2008            .chain(self.filter.iter_mut())
2009            .chain(self.group_by.iter_mut());
2010        for node in nodes {
2011            node.resolve_time_zones(schema);
2012        }
2013    }
2014
2015    /// Fails on a temporal column compared with quoted text; see
2016    /// [`Node::check_quoted_temporal`].
2017    fn check_quoted_temporal(&self, schema: &Schema) -> Result<(), String> {
2018        self.cols
2019            .iter()
2020            .chain(self.filter.iter())
2021            .chain(self.group_by.iter())
2022            .try_for_each(|node| node.check_quoted_temporal(schema))
2023    }
2024
2025    /// The where clause as a Python `.filter(...)` call, if there is one.
2026    pub(crate) fn python_filter(&self) -> Option<String> {
2027        self.filter
2028            .as_ref()
2029            .map(|f| format!(".filter({})", f.python()))
2030    }
2031
2032    /// Python calls doing what `DataTableState::query` does: the where clause, then the
2033    /// grouping (ordered by keys named `key_names`) or the select list, then `distinct`.
2034    pub(crate) fn python_steps(&self, key_names: &[String]) -> Vec<String> {
2035        let mut steps: Vec<String> = self.python_filter().into_iter().collect();
2036        if !self.group_by.is_empty() {
2037            let keys = python_list(&self.group_by);
2038            let aggs = if !self.cols.is_empty() {
2039                python_list(&self.cols)
2040            } else if self.group_by_names.is_empty() {
2041                "pl.all()".to_string()
2042            } else {
2043                let names: Vec<String> = self
2044                    .group_by_names
2045                    .iter()
2046                    .map(|n| crate::export::python_script::py_str(n))
2047                    .collect();
2048                format!("pl.all().exclude({})", names.join(", "))
2049            };
2050            steps.push(format!(".group_by({keys})"));
2051            steps.push(format!(".agg({aggs})"));
2052            steps.push(crate::export::python_script::sort_call(
2053                key_names,
2054                &vec![false; key_names.len()],
2055            ));
2056        } else if !self.cols.is_empty() {
2057            steps.push(format!(".select({})", python_list(&self.cols)));
2058        }
2059        if self.distinct {
2060            steps.push(".unique(keep=\"first\", maintain_order=True)".to_string());
2061        }
2062        steps
2063    }
2064}
2065
2066/// Expressions as Python arguments: a plain column by its name, as Polars reads a
2067/// string there, anything else as an expression.
2068fn python_list(nodes: &[Node]) -> String {
2069    nodes
2070        .iter()
2071        .map(|n| match n {
2072            Node::Col(name) => crate::export::python_script::py_str(name),
2073            n => n.python(),
2074        })
2075        .collect::<Vec<_>>()
2076        .join(", ")
2077}
2078
2079pub fn parse_query(query: &str) -> Result<ParsedQuery, String> {
2080    parse_nodes(query).map(QueryNodes::into_parsed)
2081}
2082
2083/// [`parse_query`] for data of `schema`: timestamp literals against zoned columns read
2084/// in that zone, and temporal columns compared with quoted text are errors.
2085pub fn parse_query_over(query: &str, schema: Option<&Schema>) -> Result<ParsedQuery, String> {
2086    let mut nodes = parse_nodes(query)?;
2087    if let Some(schema) = schema {
2088        nodes.resolve_time_zones(schema);
2089        nodes.check_quoted_temporal(schema)?;
2090    }
2091    Ok(nodes.into_parsed())
2092}
2093
2094/// Parse a q query into nodes. An empty query selects every column.
2095pub(crate) fn parse_nodes(query: &str) -> Result<QueryNodes, String> {
2096    // Empty query is equivalent to "select" - return all columns with no filter or grouping
2097    let trimmed = query.trim();
2098    if trimmed.is_empty() {
2099        return Ok(QueryNodes::default());
2100    }
2101
2102    let tokens = tokenize(query)?;
2103    if tokens.is_empty() || tokens[0] != Token::Select {
2104        return Err("Query must start with 'select'".to_string());
2105    }
2106    // `distinct` after `select` is the keyword unless it is a column or alias
2107    // (`distinct: x`, `distinct, a`, `distinct + 1`); `col["distinct"]` always names it.
2108    let distinct = tokens.get(1) == Some(&Token::Identifier("distinct".to_string()))
2109        && !matches!(
2110            tokens.get(2),
2111            Some(Token::Colon | Token::Comma | Token::Dot | Token::Op(_))
2112        );
2113    let body = strip_from(&tokens[if distinct { 2 } else { 1 }..])?;
2114    let body = &body[..];
2115
2116    let mut parts = split_tokens(body, &Token::Where);
2117    let select_by_tokens = parts.remove(0);
2118    let where_tokens = if !parts.is_empty() {
2119        Some(parts.remove(0))
2120    } else {
2121        None
2122    };
2123    if !parts.is_empty() {
2124        return Err(
2125            "Unexpected second 'where': combine conditions with ',' (and) or '|' (or)".to_string(),
2126        );
2127    }
2128
2129    // `by` after `where` reads naturally but grouping comes first; catch it here,
2130    // outside parentheses and brackets, so the error can name the clause order.
2131    if let Some(ref wt) = where_tokens {
2132        let mut depth = 0;
2133        let mut bracket_depth = 0;
2134        for token in wt {
2135            match token {
2136                Token::LParen => depth += 1,
2137                Token::RParen => depth -= 1,
2138                Token::LBracket => bracket_depth += 1,
2139                Token::RBracket => bracket_depth -= 1,
2140                Token::By if depth == 0 && bracket_depth == 0 => {
2141                    return Err(format!(
2142                        "Unexpected 'by' after the where clause: {}",
2143                        CLAUSE_ORDER
2144                    ));
2145                }
2146                _ => {}
2147            }
2148        }
2149    }
2150
2151    let mut select_by_parts = split_tokens(&select_by_tokens, &Token::By);
2152    let cols_tokens = select_by_parts.remove(0);
2153    let by_tokens = if !select_by_parts.is_empty() {
2154        Some(select_by_parts.remove(0))
2155    } else {
2156        None
2157    };
2158    if !select_by_parts.is_empty() {
2159        return Err(format!("Unexpected second 'by': {}", CLAUSE_ORDER));
2160    }
2161
2162    let mut cols = Vec::new();
2163    if !cols_tokens.is_empty() {
2164        for chunk in split_tokens(&cols_tokens, &Token::Comma) {
2165            if chunk.is_empty() {
2166                continue;
2167            }
2168            // Find colon position (if any) - need to account for col[...] syntax
2169            let mut colon_pos = None;
2170            let mut depth = 0;
2171            for (i, token) in chunk.iter().enumerate() {
2172                match token {
2173                    Token::LBracket => depth += 1,
2174                    Token::RBracket => depth -= 1,
2175                    Token::Colon if depth == 0 => {
2176                        colon_pos = Some(i);
2177                        break;
2178                    }
2179                    _ => {}
2180                }
2181            }
2182            if let Some(pos) = colon_pos {
2183                // Has alias: parse left side for alias name, right side for expression
2184                let alias_tokens = &chunk[..pos];
2185                let expr_tokens = &chunk[pos + 1..];
2186
2187                // Parse alias - could be simple identifier or col[...]
2188                let alias_name = if alias_tokens.len() == 1 {
2189                    if let Token::Identifier(name) = &alias_tokens[0] {
2190                        name.clone()
2191                    } else {
2192                        return Err("Expected identifier or col[] for alias".to_string());
2193                    }
2194                } else if alias_tokens.len() == 4
2195                    && alias_tokens[0] == Token::Identifier("col".to_string())
2196                    && alias_tokens[1] == Token::LBracket
2197                    && alias_tokens[3] == Token::RBracket
2198                {
2199                    // col[...] syntax for alias
2200                    match &alias_tokens[2] {
2201                        Token::String(name) | Token::Identifier(name) => name.clone(),
2202                        _ => {
2203                            return Err(
2204                                "Expected string or identifier in col[] for alias".to_string()
2205                            );
2206                        }
2207                    }
2208                } else {
2209                    // Try to parse as expression and extract name (for simple cases)
2210                    // For now, require explicit identifier or col[]
2211                    return Err("Alias must be an identifier or col[]".to_string());
2212                };
2213
2214                let expr = parse_node(expr_tokens)?;
2215                cols.push(expr.alias(alias_name));
2216            } else {
2217                cols.push(parse_node(&chunk)?);
2218            }
2219        }
2220    }
2221
2222    let mut group_by_cols = Vec::new();
2223    let mut group_by_col_names = Vec::new();
2224    if let Some(bt) = by_tokens {
2225        for chunk in split_tokens(&bt, &Token::Comma) {
2226            if chunk.is_empty() {
2227                continue;
2228            }
2229            // Support column assignment in by clause (like select)
2230            // Find colon position (if any) - need to account for col[...] syntax
2231            let mut colon_pos = None;
2232            let mut depth = 0;
2233            for (i, token) in chunk.iter().enumerate() {
2234                match token {
2235                    Token::LBracket => depth += 1,
2236                    Token::RBracket => depth -= 1,
2237                    Token::Colon if depth == 0 => {
2238                        colon_pos = Some(i);
2239                        break;
2240                    }
2241                    _ => {}
2242                }
2243            }
2244            if let Some(pos) = colon_pos {
2245                // Has alias: parse left side for alias name, right side for expression
2246                let alias_tokens = &chunk[..pos];
2247                let expr_tokens = &chunk[pos + 1..];
2248
2249                // Parse alias - could be simple identifier or col[...]
2250                let alias_name = if alias_tokens.len() == 1 {
2251                    if let Token::Identifier(name) = &alias_tokens[0] {
2252                        name.clone()
2253                    } else {
2254                        return Err(
2255                            "Expected identifier or col[] for alias in by clause".to_string()
2256                        );
2257                    }
2258                } else if alias_tokens.len() == 4
2259                    && alias_tokens[0] == Token::Identifier("col".to_string())
2260                    && alias_tokens[1] == Token::LBracket
2261                    && alias_tokens[3] == Token::RBracket
2262                {
2263                    // col[...] syntax for alias
2264                    match &alias_tokens[2] {
2265                        Token::String(name) | Token::Identifier(name) => name.clone(),
2266                        _ => {
2267                            return Err(
2268                                "Expected string or identifier in col[] for alias in by clause"
2269                                    .to_string(),
2270                            );
2271                        }
2272                    }
2273                } else {
2274                    return Err("Alias must be an identifier or col[] in by clause".to_string());
2275                };
2276
2277                let expr = parse_node(expr_tokens)?;
2278                group_by_cols.push(expr.alias(alias_name.clone()));
2279                group_by_col_names.push(alias_name); // Use alias name
2280            } else {
2281                let expr = parse_node(&chunk)?;
2282                group_by_cols.push(expr.clone());
2283                // The column name of a simple expression: `name`, or `col[name]`.
2284                if chunk.len() == 1 {
2285                    if let Token::Identifier(name) = &chunk[0] {
2286                        group_by_col_names.push(name.clone());
2287                    }
2288                } else if chunk.len() == 4
2289                    && chunk[0] == Token::Identifier("col".to_string())
2290                    && chunk[1] == Token::LBracket
2291                    && chunk[3] == Token::RBracket
2292                {
2293                    // col[...] syntax
2294                    match &chunk[2] {
2295                        Token::String(name) | Token::Identifier(name) => {
2296                            group_by_col_names.push(name.clone());
2297                        }
2298                        _ => {}
2299                    }
2300                } else {
2301                    // Complex expressions have no simple name; sorting uses the expression itself.
2302                }
2303            }
2304        }
2305    }
2306
2307    let mut filter: Option<Node> = None;
2308    if let Some(wt) = where_tokens {
2309        for chunk in split_tokens(&wt, &Token::Comma) {
2310            if chunk.is_empty() {
2311                continue;
2312            }
2313            let mut or_expr: Option<Node> = None;
2314            for or_chunk in split_tokens(&chunk, &Token::Pipe) {
2315                if or_chunk.is_empty() {
2316                    continue;
2317                }
2318                let e = parse_node(&or_chunk)?;
2319                or_expr = match or_expr {
2320                    Some(curr) => Some(curr.bin(BinOp::Or, e)),
2321                    None => Some(e),
2322                };
2323            }
2324            if let Some(e) = or_expr {
2325                filter = match filter {
2326                    Some(curr) => Some(curr.bin(BinOp::And, e)),
2327                    None => Some(e),
2328                };
2329            }
2330        }
2331    }
2332
2333    Ok(QueryNodes {
2334        cols,
2335        filter,
2336        group_by: group_by_cols,
2337        group_by_names: group_by_col_names,
2338        distinct,
2339    })
2340}
2341
2342/// One expression as Polars runs it.
2343#[cfg(test)]
2344fn parse_expr(tokens: &[Token]) -> Result<Expr, String> {
2345    parse_node(tokens).map(|n| n.to_expr())
2346}
2347
2348#[cfg(test)]
2349mod tests;