Skip to main content

datui_lib/
query.rs

1use polars::prelude::StrptimeOptions;
2use polars::prelude::*;
3use std::ops::{Add, Div, Mul, Rem, Sub};
4
5#[derive(Debug, Clone, PartialEq)]
6enum Token {
7    Identifier(String),
8    Number(f64),
9    String(String),
10    /// Date literal in YYYY.MM.DD format, stored as ISO "YYYY-MM-DD" for Polars
11    DateLiteral(String),
12    /// Timestamp literal YYYY.MM.DDTHH:MM:SS[.fff...]
13    TimestampLiteral {
14        iso: String,
15        format_str: String,
16        time_unit: TimeUnit,
17    },
18    Op(String),
19    LParen,
20    RParen,
21    LBracket,
22    RBracket,
23    Comma,
24    Colon,
25    Pipe,
26    Dot,
27    Select,
28    Where,
29    By,
30}
31
32/// Parse YYYY.MM.DDTHH:MM:SS[.fff...] timestamp. Consumes from chars. Returns (iso_string, format, time_unit) or None.
33fn parse_timestamp_literal(
34    date_part: &str,
35    chars: &mut std::iter::Peekable<std::str::Chars<'_>>,
36) -> Option<(String, String, TimeUnit)> {
37    if chars.peek() != Some(&'T') {
38        return None;
39    }
40    chars.next(); // consume 'T'
41    let mut time_part = String::new();
42    while let Some(&c) = chars.peek() {
43        if c.is_ascii_digit() || c == ':' || c == '.' {
44            time_part.push(c);
45            chars.next();
46        } else {
47            break;
48        }
49    }
50    let parts: Vec<&str> = time_part.split(':').collect();
51    if parts.len() != 3 {
52        return None;
53    }
54    let (h, m, s) = (parts[0], parts[1], parts[2]);
55    if h.len() != 2 || m.len() != 2 || s.len() < 2 {
56        return None;
57    }
58    let (sec_part, frac) = match s.split_once('.') {
59        Some((a, f)) => (a, f),
60        None => (s, ""),
61    };
62    let (time_unit, format_str) = match frac.len() {
63        0 => (TimeUnit::Microseconds, "%Y-%m-%dT%H:%M:%S".to_string()),
64        1..=3 => (TimeUnit::Milliseconds, "%Y-%m-%dT%H:%M:%S%.3f".to_string()),
65        4..=6 => (TimeUnit::Microseconds, "%Y-%m-%dT%H:%M:%S%.6f".to_string()),
66        7..=9 => (TimeUnit::Nanoseconds, "%Y-%m-%dT%H:%M:%S%.9f".to_string()),
67        _ => (TimeUnit::Nanoseconds, "%Y-%m-%dT%H:%M:%S%.9f".to_string()),
68    };
69    let iso_date = parse_date_literal(date_part)?;
70    let frac_padded = match time_unit {
71        TimeUnit::Milliseconds => format!("{:0<3}", frac),
72        TimeUnit::Microseconds => format!("{:0<6}", frac),
73        TimeUnit::Nanoseconds => format!("{:0<9}", frac),
74    };
75    let iso = if frac.is_empty() {
76        format!("{}T{}:{}:{}", iso_date, h, m, sec_part)
77    } else {
78        format!("{}T{}:{}:{}.{}", iso_date, h, m, sec_part, frac_padded)
79    };
80    Some((iso, format_str, time_unit))
81}
82
83/// Parse YYYY.MM.DD date literal (e.g. 2021.01.01). Returns ISO string "YYYY-MM-DD" or None.
84fn parse_date_literal(s: &str) -> Option<String> {
85    let parts: Vec<&str> = s.split('.').collect();
86    if parts.len() != 3 {
87        return None;
88    }
89    let year: u32 = parts[0].parse().ok()?;
90    let month: u32 = parts[1].parse().ok()?;
91    let day: u32 = parts[2].parse().ok()?;
92    if parts[0].len() != 4 || !(1000..=9999).contains(&year) {
93        return None;
94    }
95    if !(1..=12).contains(&month) || !(1..=31).contains(&day) {
96        return None;
97    }
98    Some(format!("{:04}-{:02}-{:02}", year, month, day))
99}
100
101fn tokenize(input: &str) -> Result<Vec<Token>, String> {
102    let mut tokens = Vec::new();
103    let mut chars = input.chars().peekable();
104
105    while let Some(&c) = chars.peek() {
106        match c {
107            ' ' | '\t' | '\n' | '\r' => {
108                chars.next();
109            }
110            ',' => {
111                tokens.push(Token::Comma);
112                chars.next();
113            }
114            ':' => {
115                tokens.push(Token::Colon);
116                chars.next();
117            }
118            '|' => {
119                tokens.push(Token::Pipe);
120                chars.next();
121            }
122            '(' => {
123                tokens.push(Token::LParen);
124                chars.next();
125            }
126            ')' => {
127                tokens.push(Token::RParen);
128                chars.next();
129            }
130            '[' => {
131                tokens.push(Token::LBracket);
132                chars.next();
133            }
134            ']' => {
135                tokens.push(Token::RBracket);
136                chars.next();
137            }
138            '"' => {
139                // Parse string literal with escape sequences
140                chars.next(); // consume opening quote
141                let mut string_val = String::new();
142                let mut found_closing_quote = false;
143                while let Some(&c) = chars.peek() {
144                    if c == '\\' {
145                        chars.next(); // consume backslash
146                        if let Some(&next_c) = chars.peek() {
147                            match next_c {
148                                'n' => {
149                                    string_val.push('\n');
150                                    chars.next();
151                                }
152                                't' => {
153                                    string_val.push('\t');
154                                    chars.next();
155                                }
156                                'r' => {
157                                    string_val.push('\r');
158                                    chars.next();
159                                }
160                                '\\' => {
161                                    string_val.push('\\');
162                                    chars.next();
163                                }
164                                '"' => {
165                                    string_val.push('"');
166                                    chars.next();
167                                }
168                                _ => {
169                                    // Unknown escape, just include the backslash and next char
170                                    string_val.push('\\');
171                                    string_val.push(next_c);
172                                    chars.next();
173                                }
174                            }
175                        } else {
176                            return Err("Unterminated escape sequence in string".to_string());
177                        }
178                    } else if c == '"' {
179                        chars.next(); // consume closing quote
180                        found_closing_quote = true;
181                        break;
182                    } else {
183                        string_val.push(c);
184                        chars.next();
185                    }
186                }
187                if !found_closing_quote {
188                    return Err("Unterminated string literal".to_string());
189                }
190                tokens.push(Token::String(string_val));
191            }
192            '^' => {
193                tokens.push(Token::Op("^".to_string()));
194                chars.next();
195            }
196            '+' | '-' | '*' | '%' | '/' | '=' | '<' | '>' | '!' => {
197                let mut op = c.to_string();
198                chars.next();
199                if let Some(&next_c) = chars.peek()
200                    && ((c == '<' && (next_c == '=' || next_c == '>'))
201                        || (c == '>' && next_c == '=')
202                        || (c == '!' && next_c == '='))
203                {
204                    op.push(next_c);
205                    chars.next();
206                }
207                tokens.push(Token::Op(op));
208            }
209            '.' => {
210                chars.next();
211                if chars.peek().is_some_and(|nc| nc.is_ascii_digit()) {
212                    let mut num_str = String::from('.');
213                    while let Some(&nc) = chars.peek() {
214                        if nc.is_ascii_digit() {
215                            num_str.push(nc);
216                            chars.next();
217                        } else {
218                            break;
219                        }
220                    }
221                    if let Ok(n) = num_str.parse::<f64>() {
222                        tokens.push(Token::Number(n));
223                    } else {
224                        return Err(format!("Invalid number: {}", num_str));
225                    }
226                } else {
227                    tokens.push(Token::Dot);
228                }
229            }
230            '0'..='9' => {
231                let mut num_str = String::new();
232                while let Some(&nc) = chars.peek() {
233                    if nc.is_ascii_digit() || nc == '.' {
234                        num_str.push(nc);
235                        chars.next();
236                    } else {
237                        break;
238                    }
239                }
240                // Check for YYYY.MM.DDTHH:MM:SS timestamp literal (peek for 'T' before consuming)
241                let is_timestamp =
242                    parse_date_literal(&num_str).is_some() && chars.peek() == Some(&'T');
243                if is_timestamp
244                    && let Some((iso, format_str, time_unit)) =
245                        parse_timestamp_literal(&num_str, &mut chars)
246                {
247                    tokens.push(Token::TimestampLiteral {
248                        iso,
249                        format_str,
250                        time_unit,
251                    });
252                    continue;
253                }
254                // Check for YYYY.MM.DD date literal
255                if let Some(iso) = parse_date_literal(&num_str) {
256                    tokens.push(Token::DateLiteral(iso));
257                } else if let Ok(n) = num_str.parse::<f64>() {
258                    tokens.push(Token::Number(n));
259                } else {
260                    return Err(format!("Invalid number: {}", num_str));
261                }
262            }
263            _ if c.is_alphabetic() || c == '_' => {
264                let mut ident = String::new();
265                while let Some(&nc) = chars.peek() {
266                    if nc.is_alphanumeric() || nc == '_' {
267                        ident.push(nc);
268                        chars.next();
269                    } else {
270                        break;
271                    }
272                }
273                match ident.as_str() {
274                    "select" => tokens.push(Token::Select),
275                    "where" => tokens.push(Token::Where),
276                    "by" => tokens.push(Token::By),
277                    _ => tokens.push(Token::Identifier(ident)),
278                }
279            }
280            _ => return Err(format!("Unexpected character: {}", c)),
281        }
282    }
283    Ok(tokens)
284}
285
286fn split_tokens(tokens: &[Token], delimiter: &Token) -> Vec<Vec<Token>> {
287    let mut result = Vec::new();
288    let mut current = Vec::new();
289    let mut depth = 0;
290    let mut bracket_depth = 0;
291
292    for token in tokens {
293        match token {
294            Token::LParen => depth += 1,
295            Token::RParen => depth -= 1,
296            Token::LBracket => bracket_depth += 1,
297            Token::RBracket => bracket_depth -= 1,
298            _ => {}
299        }
300
301        if depth == 0 && bracket_depth == 0 && token == delimiter {
302            result.push(current);
303            current = Vec::new();
304        } else {
305            current.push(token.clone());
306        }
307    }
308    result.push(current);
309    result
310}
311
312/// The token as the user typed it, for error messages.
313fn token_text(token: &Token) -> String {
314    match token {
315        Token::Identifier(s) => s.clone(),
316        Token::Number(n) => n.to_string(),
317        Token::String(s) => format!("\"{}\"", s),
318        Token::DateLiteral(iso) => iso.clone(),
319        Token::TimestampLiteral { iso, .. } => iso.clone(),
320        Token::Op(op) => op.clone(),
321        Token::LParen => "(".to_string(),
322        Token::RParen => ")".to_string(),
323        Token::LBracket => "[".to_string(),
324        Token::RBracket => "]".to_string(),
325        Token::Comma => ",".to_string(),
326        Token::Colon => ":".to_string(),
327        Token::Pipe => "|".to_string(),
328        Token::Dot => ".".to_string(),
329        Token::Select => "select".to_string(),
330        Token::Where => "where".to_string(),
331        Token::By => "by".to_string(),
332    }
333}
334
335/// Remedy shown when a clause keyword turns up out of place.
336const CLAUSE_ORDER: &str = "clause order is select [by group] [where conditions]";
337
338/// [`CLAUSE_ORDER`] with q's optional `from df`, for errors about it.
339const FROM_ORDER: &str = "clause order is select [by group] [from df] [where conditions]";
340
341/// The one table q reads: the one on screen, named as SQL names it.
342const TABLE: &str = "df";
343
344/// The table name after a `from` at `tokens[i]`: an identifier, or a dotted
345/// path like `data.csv`, that is not an operator word. Returns it and the index
346/// past it.
347fn from_table_at(tokens: &[Token], i: usize) -> Option<(String, usize)> {
348    if tokens.get(i) != Some(&Token::Identifier("from".to_string())) {
349        return None;
350    }
351    let mut name = match tokens.get(i + 1) {
352        Some(Token::Identifier(n)) if !WORD_OPS.contains(&n.as_str()) => n.clone(),
353        _ => return None,
354    };
355    let mut end = i + 2;
356    while let (Some(Token::Dot), Some(Token::Identifier(part))) =
357        (tokens.get(end), tokens.get(end + 1))
358    {
359        name.push('.');
360        name.push_str(part);
361        end += 2;
362    }
363    Some((name, end))
364}
365
366/// The query body without q's `from df`, which may sit after the select list and
367/// `by`, before `where`. `from` stays an identifier, so it is the clause only
368/// where a column could not be: followed by a table name, then `where`, `by` or
369/// the end, outside brackets. A column named `from` reads as one elsewhere.
370fn strip_from(body: &[Token]) -> Result<Vec<Token>, String> {
371    let mut depth = 0i32;
372    let mut found: Option<(usize, usize)> = None;
373    for (i, token) in body.iter().enumerate() {
374        match token {
375            Token::LParen | Token::LBracket => depth += 1,
376            Token::RParen | Token::RBracket => depth -= 1,
377            _ => {}
378        }
379        if depth != 0 {
380            continue;
381        }
382        let Some((name, end)) = from_table_at(body, i) else {
383            continue;
384        };
385        let after_where = body[..i].contains(&Token::Where);
386        match body.get(end) {
387            None | Some(Token::Where) | Some(Token::By) => {}
388            _ => continue,
389        }
390        if name != TABLE {
391            return Err("q reads the table on screen, named df: … from df …".to_string());
392        }
393        if after_where {
394            return Err(format!(
395                "Unexpected 'from df' after the where clause: {FROM_ORDER}"
396            ));
397        }
398        if body.get(end) == Some(&Token::By) {
399            return Err(format!("Unexpected 'by' after 'from df': {FROM_ORDER}"));
400        }
401        found = Some((i, end));
402    }
403    let mut body = body.to_vec();
404    if let Some((start, end)) = found {
405        body.drain(start..end);
406    }
407    Ok(body)
408}
409
410/// Infix operators spelled as words (q's names). They stay ordinary identifiers
411/// everywhere else, so a column called `in` or `mod` still reads as one when it
412/// opens an expression or follows a `.`.
413const WORD_OPS: [&str; 5] = ["in", "like", "xbar", "mod", "wavg"];
414
415/// The infix operator at `tokens[i]`, if there is one: a symbol, or an operator
416/// word that has an operand before it.
417fn infix_op_at(tokens: &[Token], i: usize) -> Option<&str> {
418    match tokens.get(i)? {
419        Token::Op(op) => Some(op.as_str()),
420        Token::Identifier(word)
421            if i > 0 && tokens[i - 1] != Token::Dot && WORD_OPS.contains(&word.as_str()) =>
422        {
423            Some(word.as_str())
424        }
425        _ => None,
426    }
427}
428
429/// A parsed expression, before it is a Polars expression ([`Node::to_expr`]) or
430/// Python Polars code ([`Node::python`]). One parse serves both, so the code
431/// "Copy as Python" writes is the query datui ran.
432#[derive(Debug, Clone, PartialEq)]
433pub(crate) enum Node {
434    Col(String),
435    /// A number as typed: a float literal.
436    Num(f64),
437    /// A whole number where the operator keeps integers whole (`mod`, `xbar`).
438    Int(i64),
439    Str(String),
440    Bool(bool),
441    Null,
442    /// A `YYYY.MM.DD` literal, as ISO `YYYY-MM-DD`.
443    Date(String),
444    /// A `YYYY.MM.DDTHH:MM:SS[.fff]` literal.
445    Timestamp {
446        iso: String,
447        format: String,
448        unit: TimeUnit,
449        /// The zone of the column it meets, read as a clock there; none for a column
450        /// without one. See [`Node::resolve_time_zones`].
451        zone: Option<String>,
452    },
453    Bin(BinOp, Box<Node>, Box<Node>),
454    Coalesce(Box<Node>, Box<Node>),
455    /// The values of the first where the second holds.
456    Filter(Box<Node>, Box<Node>),
457    /// `when(condition).then(value).otherwise(other)`.
458    When(Box<Node>, Box<Node>, Box<Node>),
459    Op(Box<Node>, Op),
460    Alias(Box<Node>, String),
461}
462
463#[derive(Debug, Clone, Copy, PartialEq, Eq)]
464pub(crate) enum BinOp {
465    Add,
466    Sub,
467    Mul,
468    /// Polars' `/` on two expressions.
469    Div,
470    TrueDiv,
471    FloorDiv,
472    Rem,
473    Eq,
474    Neq,
475    Lt,
476    Gt,
477    LtEq,
478    GtEq,
479    And,
480    Or,
481}
482
483/// A method applied to one expression.
484#[derive(Debug, Clone, PartialEq)]
485pub(crate) enum Op {
486    Mean,
487    Min,
488    Max,
489    Count,
490    Std,
491    Var,
492    Median,
493    Sum,
494    First,
495    Last,
496    NUnique,
497    Not,
498    IsNull,
499    IsNotNull,
500    LenChars,
501    Upper,
502    Lower,
503    Abs,
504    Floor,
505    Ceil,
506    Sqrt,
507    Ln,
508    Exp,
509    Date,
510    Time,
511    Year,
512    Quarter,
513    Month,
514    Week,
515    Day,
516    OrdinalDay,
517    Weekday,
518    Hour,
519    Minute,
520    Second,
521    MonthStart,
522    MonthEnd,
523    DtFormat(String),
524    StartsWith(String),
525    EndsWith(String),
526    ContainsLiteral(String),
527    /// A regex match, strict.
528    ContainsRegex(String),
529    /// Split on the text and take the piece at the index, null past the last.
530    Part(String, i64),
531    Slice(i64, Option<u64>),
532    ReplaceAll(String, String),
533    Strip,
534    ToDate(Option<String>),
535    ToDatetime(Option<String>),
536    Round(u32),
537    /// A non-strict cast.
538    Cast(CastTo),
539}
540
541#[derive(Debug, Clone, Copy, PartialEq, Eq)]
542pub(crate) enum CastTo {
543    Int64,
544    Float64,
545    String,
546}
547
548impl CastTo {
549    fn dtype(self) -> DataType {
550        match self {
551            CastTo::Int64 => DataType::Int64,
552            CastTo::Float64 => DataType::Float64,
553            CastTo::String => DataType::String,
554        }
555    }
556}
557
558/// The most nodes an expression may grow to. `wavg`, `xbar` and `in` repeat an
559/// operand, so nesting them multiplies its size at every level: two dozen nested
560/// `wavg` were billions of nodes. Found by the `parse_query` fuzz target.
561const MAX_EXPR_NODES: usize = 10_000;
562
563/// Err when `copies` of `node` would pass [`MAX_EXPR_NODES`].
564fn check_copies(node: &Node, copies: usize) -> Result<(), String> {
565    if node.size().saturating_mul(copies) > MAX_EXPR_NODES {
566        return Err(
567            "Expression is too large: nested wavg, xbar or in repeat what they are \
568                    given. Simplify it or split it into steps."
569                .to_string(),
570        );
571    }
572    Ok(())
573}
574
575impl Node {
576    /// How many nodes the tree has.
577    fn size(&self) -> usize {
578        1 + match self {
579            Node::Col(_)
580            | Node::Num(_)
581            | Node::Int(_)
582            | Node::Str(_)
583            | Node::Bool(_)
584            | Node::Null
585            | Node::Date(_)
586            | Node::Timestamp { .. } => 0,
587            Node::Bin(_, a, b) | Node::Coalesce(a, b) | Node::Filter(a, b) => a.size() + b.size(),
588            Node::When(a, b, c) => a.size() + b.size() + c.size(),
589            Node::Op(a, _) | Node::Alias(a, _) => a.size(),
590        }
591    }
592
593    fn op(self, op: Op) -> Node {
594        Node::Op(Box::new(self), op)
595    }
596
597    fn bin(self, op: BinOp, right: Node) -> Node {
598        Node::Bin(op, Box::new(self), Box::new(right))
599    }
600
601    fn alias(self, name: impl Into<String>) -> Node {
602        Node::Alias(Box::new(self), name.into())
603    }
604
605    fn cast_text(self) -> Node {
606        self.op(Op::Cast(CastTo::String))
607    }
608
609    /// The Polars expression.
610    pub(crate) fn to_expr(&self) -> Expr {
611        match self {
612            Node::Col(name) => col(name),
613            Node::Num(n) => lit(*n),
614            Node::Int(n) => lit(*n),
615            Node::Str(s) => lit(s.as_str()),
616            Node::Bool(b) => lit(*b),
617            Node::Null => lit(NULL),
618            Node::Date(iso) => {
619                let opts = StrptimeOptions {
620                    format: Some("%Y-%m-%d".into()),
621                    ..Default::default()
622                };
623                lit(iso.as_str()).str().to_date(opts)
624            }
625            Node::Timestamp {
626                iso,
627                format,
628                unit,
629                zone,
630            } => {
631                let opts = StrptimeOptions {
632                    format: Some(format.as_str().into()),
633                    ..Default::default()
634                };
635                // Set only from a column's own dtype, so it parses.
636                let zone = TimeZone::opt_try_new(zone.as_deref()).ok().flatten();
637                // A clock time a fall back repeats is its first instant.
638                lit(iso.as_str())
639                    .str()
640                    .to_datetime(Some(*unit), zone, opts, lit("earliest"))
641            }
642            Node::Bin(op, left, right) => {
643                let (left, right) = (left.to_expr(), right.to_expr());
644                match op {
645                    BinOp::Add => left.add(right),
646                    BinOp::Sub => left.sub(right),
647                    BinOp::Mul => left.mul(right),
648                    BinOp::Div => left.div(right),
649                    BinOp::TrueDiv => left.true_div(right),
650                    BinOp::FloorDiv => left.floor_div(right),
651                    BinOp::Rem => left.rem(right),
652                    BinOp::Eq => left.eq(right),
653                    BinOp::Neq => left.neq(right),
654                    BinOp::Lt => left.lt(right),
655                    BinOp::Gt => left.gt(right),
656                    BinOp::LtEq => left.lt_eq(right),
657                    BinOp::GtEq => left.gt_eq(right),
658                    BinOp::And => left.and(right),
659                    BinOp::Or => left.or(right),
660                }
661            }
662            Node::Coalesce(left, right) => coalesce(&[left.to_expr(), right.to_expr()]),
663            Node::Filter(values, predicate) => values.to_expr().filter(predicate.to_expr()),
664            Node::When(condition, then, otherwise) => when(condition.to_expr())
665                .then(then.to_expr())
666                .otherwise(otherwise.to_expr()),
667            Node::Op(inner, op) => apply_op_expr(inner.to_expr(), op),
668            Node::Alias(inner, name) => inner.to_expr().alias(name.as_str()),
669        }
670    }
671
672    /// The same node with every alias inside it taken off, as Polars' `undo_aliases`.
673    pub(crate) fn without_aliases(&self) -> Node {
674        let strip = |n: &Node| Box::new(n.without_aliases());
675        match self {
676            Node::Alias(inner, _) => inner.without_aliases(),
677            Node::Bin(op, l, r) => Node::Bin(*op, strip(l), strip(r)),
678            Node::Coalesce(l, r) => Node::Coalesce(strip(l), strip(r)),
679            Node::Filter(v, p) => Node::Filter(strip(v), strip(p)),
680            Node::When(c, t, o) => Node::When(strip(c), strip(t), strip(o)),
681            Node::Op(inner, op) => Node::Op(strip(inner), op.clone()),
682            leaf => leaf.clone(),
683        }
684    }
685
686    /// Each `/` named as Polars runs it over `schema`, the data the expression reads.
687    /// Polars' `/` on two expressions (`Div`) floor-divides whole numbers and divides
688    /// anything else; Python has no such operator, so the script needs `//` or `/`,
689    /// picked by the type of the quotient.
690    pub(crate) fn resolve_division(&mut self, schema: &Schema) {
691        match self {
692            Node::Bin(_, left, right) | Node::Coalesce(left, right) | Node::Filter(left, right) => {
693                left.resolve_division(schema);
694                right.resolve_division(schema);
695            }
696            Node::When(c, t, o) => {
697                c.resolve_division(schema);
698                t.resolve_division(schema);
699                o.resolve_division(schema);
700            }
701            Node::Op(inner, _) | Node::Alias(inner, _) => inner.resolve_division(schema),
702            _ => {}
703        }
704        if let Node::Bin(BinOp::Div, ..) = self {
705            let quotient = DataFrame::empty_with_schema(schema)
706                .lazy()
707                .select([self.to_expr()])
708                .collect_schema()
709                .ok()
710                .and_then(|s| s.get_at_index(0).map(|(_, dtype)| dtype.is_integer()));
711            if let (Some(whole), Node::Bin(op, ..)) = (quotient, self) {
712                *op = if whole {
713                    BinOp::FloorDiv
714                } else {
715                    BinOp::TrueDiv
716                };
717            }
718        }
719    }
720
721    /// Each timestamp literal that meets a column with a time zone, by a comparison,
722    /// arithmetic, `^` or the two sides of a `?`, takes that zone, so it reads as a clock
723    /// there; Polars refuses to compare a zoned datetime with a naive one.
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 (`=`, `<`, … or `in`) of a temporal column with
747    /// quoted text. Quoted text is a string in q, never a date, so it stays an error; this
748    /// one says so in q's words. Only a column `schema` types; the rest is Polars'.
749    fn check_quoted_temporal(&self, schema: &Schema) -> Result<(), String> {
750        match self {
751            Node::Bin(op, left, right) => {
752                left.check_quoted_temporal(schema)?;
753                right.check_quoted_temporal(schema)?;
754                let compares = matches!(
755                    op,
756                    BinOp::Eq | BinOp::Neq | BinOp::Lt | BinOp::Gt | BinOp::LtEq | BinOp::GtEq
757                );
758                if compares
759                    && let Some(err) = quoted_temporal(left, right, schema)
760                        .or_else(|| quoted_temporal(right, left, schema))
761                {
762                    return Err(err);
763                }
764            }
765            Node::Coalesce(left, right) | Node::Filter(left, right) => {
766                left.check_quoted_temporal(schema)?;
767                right.check_quoted_temporal(schema)?;
768            }
769            Node::When(c, t, o) => {
770                c.check_quoted_temporal(schema)?;
771                t.check_quoted_temporal(schema)?;
772                o.check_quoted_temporal(schema)?;
773            }
774            Node::Op(inner, _) | Node::Alias(inner, _) => inner.check_quoted_temporal(schema)?,
775            _ => {}
776        }
777        Ok(())
778    }
779
780    /// Give a zoneless timestamp literal on one side the zone of the other side's type.
781    fn share_zone(a: &mut Node, b: &mut Node, schema: &Schema) {
782        if !Self::take_zone(a, b, schema) {
783            Self::take_zone(b, a, schema);
784        }
785    }
786
787    /// Whether `literal`, a zoneless timestamp literal, took the zone of `other`'s type.
788    fn take_zone(literal: &mut Node, other: &Node, schema: &Schema) -> bool {
789        if let Node::Timestamp {
790            zone: zone @ None, ..
791        } = literal
792            && let Some(DataType::Datetime(_, Some(tz))) = other.dtype(schema)
793        {
794            *zone = Some(tz.to_string());
795            return true;
796        }
797        false
798    }
799
800    /// The type the expression has over `schema`, when Polars can say.
801    fn dtype(&self, schema: &Schema) -> Option<DataType> {
802        DataFrame::empty_with_schema(schema)
803            .lazy()
804            .select([self.to_expr()])
805            .collect_schema()
806            .ok()
807            .and_then(|s| s.get_at_index(0).map(|(_, dtype)| dtype.clone()))
808    }
809
810    /// The literal as Python, bare: `1.0`, `"a"`, `True`, `None`.
811    fn python_literal(&self) -> Option<String> {
812        Some(match self {
813            Node::Num(n) => crate::python_script::py_float(*n),
814            Node::Int(n) => n.to_string(),
815            Node::Str(s) => crate::python_script::py_str(s),
816            Node::Bool(b) => crate::python_script::py_bool(*b).to_string(),
817            Node::Null => "None".to_string(),
818            _ => return None,
819        })
820    }
821
822    /// Python Polars code for the expression.
823    pub(crate) fn python(&self) -> String {
824        use crate::python_script::py_str;
825        if let Some(literal) = self.python_literal() {
826            return format!("pl.lit({literal})");
827        }
828        match self {
829            Node::Col(name) => format!("pl.col({})", py_str(name)),
830            Node::Date(iso) => {
831                let parts: Vec<String> = iso
832                    .split('-')
833                    .map(|p| p.trim_start_matches('0').to_string())
834                    .map(|p| if p.is_empty() { "0".to_string() } else { p })
835                    .collect();
836                format!("pl.date({})", parts.join(", "))
837            }
838            Node::Timestamp {
839                iso,
840                format,
841                unit,
842                zone,
843            } => format!(
844                "pl.lit({}).str.to_datetime({}, time_unit={}{})",
845                py_str(iso),
846                py_str(format),
847                py_str(time_unit_name(*unit)),
848                zone.as_ref().map_or(String::new(), |zone| format!(
849                    ", time_zone={}, ambiguous=\"earliest\"",
850                    py_str(zone)
851                ))
852            ),
853            Node::Bin(op, left, right) => {
854                // A literal on the right stays bare (`pl.col("a") > 1.0`); Python's
855                // operators turn it into one.
856                let right = match right.python_literal() {
857                    Some(literal) => literal,
858                    None => right.python_operand(),
859                };
860                format!("{} {} {}", left.python_operand(), op.python(), right)
861            }
862            Node::Coalesce(left, right) => {
863                format!("pl.coalesce({}, {})", left.python(), right.python())
864            }
865            Node::Filter(values, predicate) => {
866                format!("{}.filter({})", values.python_operand(), predicate.python())
867            }
868            Node::When(condition, then, otherwise) => format!(
869                "pl.when({}).then({}).otherwise({})",
870                condition.python(),
871                then.python(),
872                otherwise.python()
873            ),
874            Node::Op(inner, op) => format!("{}{}", inner.python_operand(), op.python()),
875            Node::Alias(inner, name) => {
876                // Only the outer name survives an alias of an alias.
877                let mut inner = inner.as_ref();
878                while let Node::Alias(deeper, _) = inner {
879                    inner = deeper;
880                }
881                format!("{}.alias({})", inner.python_operand(), py_str(name))
882            }
883            _ => unreachable!("literals return above"),
884        }
885    }
886
887    /// As [`Self::python`], parenthesized where an operator or a method after it
888    /// would otherwise bind to part of it.
889    fn python_operand(&self) -> String {
890        match self {
891            Node::Bin(..) => format!("({})", self.python()),
892            _ => self.python(),
893        }
894    }
895}
896
897fn time_unit_name(unit: TimeUnit) -> &'static str {
898    match unit {
899        TimeUnit::Milliseconds => "ms",
900        TimeUnit::Microseconds => "us",
901        TimeUnit::Nanoseconds => "ns",
902    }
903}
904
905impl BinOp {
906    fn python(self) -> &'static str {
907        match self {
908            BinOp::Add => "+",
909            BinOp::Sub => "-",
910            BinOp::Mul => "*",
911            BinOp::Div | BinOp::TrueDiv => "/",
912            BinOp::FloorDiv => "//",
913            BinOp::Rem => "%",
914            BinOp::Eq => "==",
915            BinOp::Neq => "!=",
916            BinOp::Lt => "<",
917            BinOp::Gt => ">",
918            BinOp::LtEq => "<=",
919            BinOp::GtEq => ">=",
920            BinOp::And => "&",
921            BinOp::Or => "|",
922        }
923    }
924}
925
926impl Op {
927    /// The method call, from its dot: `.str.to_uppercase()`.
928    fn python(&self) -> String {
929        use crate::python_script::py_str;
930        let fixed = match self {
931            Op::Mean => ".mean()",
932            Op::Min => ".min()",
933            Op::Max => ".max()",
934            Op::Count => ".count()",
935            Op::Std => ".std()",
936            Op::Var => ".var()",
937            Op::Median => ".median()",
938            Op::Sum => ".sum()",
939            Op::First => ".first()",
940            Op::Last => ".last()",
941            Op::NUnique => ".n_unique()",
942            Op::Not => ".not_()",
943            Op::IsNull => ".is_null()",
944            Op::IsNotNull => ".is_not_null()",
945            Op::LenChars => ".str.len_chars()",
946            Op::Upper => ".str.to_uppercase()",
947            Op::Lower => ".str.to_lowercase()",
948            Op::Abs => ".abs()",
949            Op::Floor => ".floor()",
950            Op::Ceil => ".ceil()",
951            Op::Sqrt => ".sqrt()",
952            Op::Ln => ".log()",
953            Op::Exp => ".exp()",
954            Op::Date => ".dt.date()",
955            Op::Time => ".dt.time()",
956            Op::Year => ".dt.year()",
957            Op::Quarter => ".dt.quarter()",
958            Op::Month => ".dt.month()",
959            Op::Week => ".dt.week()",
960            Op::Day => ".dt.day()",
961            Op::OrdinalDay => ".dt.ordinal_day()",
962            Op::Weekday => ".dt.weekday()",
963            Op::Hour => ".dt.hour()",
964            Op::Minute => ".dt.minute()",
965            Op::Second => ".dt.second()",
966            Op::MonthStart => ".dt.month_start()",
967            Op::MonthEnd => ".dt.month_end()",
968            Op::Strip => ".str.strip_chars()",
969            Op::Cast(CastTo::Int64) => ".cast(pl.Int64, strict=False)",
970            Op::Cast(CastTo::Float64) => ".cast(pl.Float64, strict=False)",
971            Op::Cast(CastTo::String) => ".cast(pl.String)",
972            _ => "",
973        };
974        if !fixed.is_empty() {
975            return fixed.to_string();
976        }
977        let format_arg = |format: &Option<String>| match format {
978            Some(f) => format!("{}, strict=False", py_str(f)),
979            None => "strict=False".to_string(),
980        };
981        match self {
982            Op::DtFormat(f) => format!(".dt.to_string({})", py_str(f)),
983            Op::StartsWith(s) => format!(".str.starts_with({})", py_str(s)),
984            Op::EndsWith(s) => format!(".str.ends_with({})", py_str(s)),
985            Op::ContainsLiteral(s) => format!(".str.contains({}, literal=True)", py_str(s)),
986            Op::ContainsRegex(r) => format!(".str.contains({})", py_str(r)),
987            Op::Part(sep, i) => format!(
988                ".str.split({}).list.get({i}, null_on_oob=True)",
989                py_str(sep)
990            ),
991            Op::Slice(start, Some(len)) => format!(".str.slice({start}, {len})"),
992            Op::Slice(start, None) => format!(".str.slice({start})"),
993            Op::ReplaceAll(from, to) => format!(
994                ".str.replace_all({}, {}, literal=True)",
995                py_str(from),
996                py_str(to)
997            ),
998            Op::ToDate(format) => format!(".str.to_date({})", format_arg(format)),
999            Op::ToDatetime(format) => format!(".str.to_datetime({})", format_arg(format)),
1000            Op::Round(d) => format!(".round({d}, mode=\"half_away_from_zero\")"),
1001            _ => unreachable!("fixed calls return above"),
1002        }
1003    }
1004}
1005
1006fn apply_op_expr(expr: Expr, op: &Op) -> Expr {
1007    let strptime = |format: &Option<String>| StrptimeOptions {
1008        format: format.as_deref().map(Into::into),
1009        // A value that does not match becomes null, as a failed parse does in q,
1010        // rather than one stray row failing the whole query.
1011        strict: false,
1012        ..Default::default()
1013    };
1014    match op {
1015        Op::Mean => expr.mean(),
1016        Op::Min => expr.min(),
1017        Op::Max => expr.max(),
1018        Op::Count => expr.count(),
1019        // Sample statistics (n - 1), like `std`; q's own var and dev divide by n.
1020        Op::Std => expr.std(1),
1021        Op::Var => expr.var(1),
1022        Op::Median => expr.median(),
1023        Op::Sum => expr.sum(),
1024        Op::First => expr.first(),
1025        Op::Last => expr.last(),
1026        Op::NUnique => expr.n_unique(),
1027        Op::Not => expr.not(),
1028        Op::IsNull => expr.is_null(),
1029        Op::IsNotNull => expr.is_not_null(),
1030        Op::LenChars => expr.str().len_chars(),
1031        Op::Upper => expr.str().to_uppercase(),
1032        Op::Lower => expr.str().to_lowercase(),
1033        Op::Abs => expr.abs(),
1034        Op::Floor => expr.floor(),
1035        Op::Ceil => expr.ceil(),
1036        Op::Sqrt => expr.sqrt(),
1037        Op::Ln => expr.log(lit(std::f64::consts::E)),
1038        Op::Exp => expr.exp(),
1039        Op::Date => expr.dt().date(),
1040        Op::Time => expr.dt().time(),
1041        Op::Year => expr.dt().year(),
1042        Op::Quarter => expr.dt().quarter(),
1043        Op::Month => expr.dt().month(),
1044        Op::Week => expr.dt().week(),
1045        Op::Day => expr.dt().day(),
1046        Op::OrdinalDay => expr.dt().ordinal_day(),
1047        Op::Weekday => expr.dt().weekday(),
1048        Op::Hour => expr.dt().hour(),
1049        Op::Minute => expr.dt().minute(),
1050        Op::Second => expr.dt().second(),
1051        Op::MonthStart => expr.dt().month_start(),
1052        Op::MonthEnd => expr.dt().month_end(),
1053        Op::DtFormat(f) => expr.dt().to_string(f),
1054        Op::StartsWith(s) => expr.str().starts_with(lit(s.as_str())),
1055        Op::EndsWith(s) => expr.str().ends_with(lit(s.as_str())),
1056        Op::ContainsLiteral(s) => expr.str().contains_literal(lit(s.as_str())),
1057        Op::ContainsRegex(r) => expr.str().contains(lit(r.as_str()), true),
1058        // Past the last piece is null, not an error; a negative index counts from the end.
1059        Op::Part(sep, i) => expr
1060            .str()
1061            .split(lit(sep.as_str()))
1062            .list()
1063            .get(lit(*i), true),
1064        Op::Slice(start, length) => {
1065            // No length: to the end of the string.
1066            let length = length.map_or_else(|| lit(NULL), lit);
1067            expr.str().slice(lit(*start), length)
1068        }
1069        Op::ReplaceAll(from, to) => {
1070            expr.str()
1071                .replace_all(lit(from.as_str()), lit(to.as_str()), true)
1072        }
1073        Op::Strip => expr.str().strip_chars(lit(NULL)),
1074        Op::ToDate(format) => expr.str().to_date(strptime(format)),
1075        Op::ToDatetime(format) => {
1076            expr.str()
1077                .to_datetime(None, None, strptime(format), lit("raise"))
1078        }
1079        // Half away from zero, the rounding people expect from a calculator or SQL.
1080        Op::Round(decimals) => expr.round(*decimals, RoundMode::HalfAwayFromZero),
1081        // Non-strict casts: a value that does not convert becomes null.
1082        Op::Cast(to) => expr.cast(to.dtype()),
1083    }
1084}
1085
1086/// An operand of `mod` or `xbar`. A whole number is an integer literal, so an
1087/// integer column keeps its type: `5 xbar passenger_count` stays Int64 instead of
1088/// becoming 5.0, 10.0.
1089fn int_or_node(tokens: &[Token]) -> Result<Node, String> {
1090    let whole = |n: f64| n.fract() == 0.0 && n.abs() < i64::MAX as f64;
1091    match tokens {
1092        [Token::Number(n)] if whole(*n) => Ok(Node::Int(*n as i64)),
1093        [Token::Op(minus), Token::Number(n)] if minus == "-" && whole(*n) => {
1094            Ok(Node::Int(-(*n as i64)))
1095        }
1096        _ => parse_node(tokens),
1097    }
1098}
1099
1100/// True when no `]` in `tokens` closes a `[` from outside them.
1101fn brackets_balanced(tokens: &[Token]) -> bool {
1102    let mut depth = 0usize;
1103    tokens.iter().all(|t| match t {
1104        Token::LBracket => {
1105            depth += 1;
1106            true
1107        }
1108        Token::RBracket => depth.checked_sub(1).map(|d| depth = d).is_some(),
1109        _ => true,
1110    })
1111}
1112
1113/// The error for `column` of a temporal type compared with the quoted `text`, naming
1114/// the literal the type takes: the text itself when it is one once unquoted.
1115fn quoted_temporal(column: &Node, text: &Node, schema: &Schema) -> Option<String> {
1116    let (Node::Col(name), Node::Str(s)) = (column, text) else {
1117        return None;
1118    };
1119    let shown = q_name(name);
1120    let unquoted = tokenize(s).ok();
1121    let literal = |is_kind: fn(&Token) -> bool, example: &str| match unquoted.as_deref() {
1122        Some([token]) if is_kind(token) => s.trim().to_string(),
1123        _ => example.to_string(),
1124    };
1125    let (kind, remedy) = match schema.get(name)? {
1126        DataType::Date => (
1127            "date",
1128            format!(
1129                "A date is {}",
1130                literal(|t| matches!(t, Token::DateLiteral(_)), "2024.01.01")
1131            ),
1132        ),
1133        DataType::Datetime(..) => (
1134            "timestamp",
1135            format!(
1136                "A timestamp is {}",
1137                literal(
1138                    |t| matches!(t, Token::TimestampLiteral { .. }),
1139                    "2024.01.01T05:00:00"
1140                )
1141            ),
1142        ),
1143        DataType::Time => (
1144            "time",
1145            format!(
1146                "A time has no literal; compare {shown}.hour, {shown}.minute or {shown}.second with a number"
1147            ),
1148        ),
1149        DataType::Duration(_) => ("duration", "A duration has no literal".to_string()),
1150        _ => return None,
1151    };
1152    Some(format!(
1153        "{shown} is a {kind}; \"{s}\" is a string. {remedy}"
1154    ))
1155}
1156
1157/// How q spells a column: bare when it can be, else `col["first name"]`.
1158pub(crate) fn q_name(name: &str) -> String {
1159    if is_plain_name(name) {
1160        name.to_string()
1161    } else {
1162        format!("col[\"{name}\"]")
1163    }
1164}
1165
1166/// Whether `name` reads as a column when typed bare, rather than needing `col["…"]`.
1167fn is_plain_name(name: &str) -> bool {
1168    let mut chars = name.chars();
1169    chars.next().is_some_and(|c| c.is_alphabetic() || c == '_')
1170        && chars.all(|c| c.is_alphanumeric() || c == '_')
1171        && !matches!(name, "select" | "where" | "by")
1172}
1173
1174/// OR of the conditions as a balanced tree, so a long `in` list nests
1175/// logarithmically rather than one level per element.
1176fn any_of(mut conditions: Vec<Node>) -> Node {
1177    if conditions.len() <= 1 {
1178        return conditions.pop().unwrap_or(Node::Bool(false));
1179    }
1180    let right = conditions.split_off(conditions.len() / 2);
1181    any_of(conditions).bin(BinOp::Or, any_of(right))
1182}
1183
1184/// A `like` pattern as an anchored regex: `*` is any run of characters, `?` any
1185/// one character, everything else literal.
1186fn like_regex(pattern: &str) -> String {
1187    let mut re = String::from("(?s)^");
1188    for c in pattern.chars() {
1189        match c {
1190            '*' => re.push_str(".*"),
1191            '?' => re.push('.'),
1192            _ => re.push_str(&regex::escape(c.encode_utf8(&mut [0; 4]))),
1193        }
1194    }
1195    re.push('$');
1196    re
1197}
1198
1199/// Operators whose right side is not an ordinary expression, or whose operands
1200/// need their tokens (literal checks, names). Everything else goes to `apply_op`.
1201fn apply_infix(left_tokens: &[Token], op: &str, right_tokens: &[Token]) -> Result<Node, String> {
1202    match op {
1203        "in" => {
1204            let list = match right_tokens {
1205                // One list: the first `[` closes at the last token, so `[1] + [2]` is not.
1206                [Token::LBracket, inner @ .., Token::RBracket] if brackets_balanced(inner) => inner,
1207                _ => {
1208                    return Err(
1209                        "in takes a list on its right, e.g. name in [\"Emma\", \"Olivia\"]"
1210                            .to_string(),
1211                    );
1212                }
1213            };
1214            let items = split_tokens(list, &Token::Comma);
1215            if items.iter().any(|item| item.is_empty()) {
1216                return Err(
1217                    "in needs a list of values, e.g. name in [\"Emma\", \"Olivia\"]".to_string(),
1218                );
1219            }
1220            let left = parse_node(left_tokens)?;
1221            // A column or literal repeated grows only as the list typed does; a larger
1222            // left side repeated per item is what multiplies.
1223            if left.size() > 1 {
1224                check_copies(&left, items.len())?;
1225            }
1226            // One `=` per value, so each value compares exactly as `x = value` would,
1227            // with the same literal casting (numbers, dates, timestamps).
1228            let conditions = items
1229                .iter()
1230                .map(|item| Ok(left.clone().bin(BinOp::Eq, parse_node(item)?)))
1231                .collect::<Result<Vec<_>, String>>()?;
1232            Ok(any_of(conditions))
1233        }
1234        "like" => {
1235            let [Token::String(pattern)] = right_tokens else {
1236                return Err(
1237                    "like takes a quoted pattern on its right, e.g. item like \"*Chicken*\""
1238                        .to_string(),
1239                );
1240            };
1241            let left = parse_node(left_tokens)?;
1242            // Cast first so numeric codes (zip, station ids read as numbers) match too.
1243            Ok(left.cast_text().op(Op::ContainsRegex(like_regex(pattern))))
1244        }
1245        "xbar" => {
1246            if let [Token::Number(n)] = left_tokens
1247                && *n <= 0.0
1248            {
1249                return Err(
1250                    "xbar needs a positive bucket size, e.g. 5 xbar fare_amount".to_string()
1251                );
1252            }
1253            let right = parse_node(right_tokens)?;
1254            let size = int_or_node(left_tokens)?;
1255            check_copies(&size, 2)?;
1256            // floor_div floors toward negative infinity for both ints and floats,
1257            // which is what makes every value land in the bucket at or below it.
1258            Ok(right
1259                .bin(BinOp::FloorDiv, size.clone())
1260                .bin(BinOp::Mul, size))
1261        }
1262        "mod" => {
1263            let right = int_or_node(right_tokens)?;
1264            let left = int_or_node(left_tokens)?;
1265            Ok(left.bin(BinOp::Rem, right))
1266        }
1267        "wavg" => {
1268            let values = parse_node(right_tokens)?;
1269            let weights = parse_node(left_tokens)?;
1270            // Below, weights appear five times and values three.
1271            check_copies(&weights, 5)?;
1272            check_copies(&values, 3)?;
1273            let weighted = weights.clone().bin(BinOp::Mul, values);
1274            // Only pairs with both a weight and a value count toward the total weight;
1275            // a null value would otherwise still pull the average toward zero.
1276            let total = Node::Filter(
1277                Box::new(weights),
1278                Box::new(weighted.clone().op(Op::IsNotNull)),
1279            )
1280            .op(Op::Sum);
1281            // No weight at all (every pair null) is no average, not 0/0 = NaN.
1282            let total = Node::When(
1283                Box::new(total.clone().bin(BinOp::Neq, Node::Int(0))),
1284                Box::new(total),
1285                Box::new(Node::Null),
1286            );
1287            let node = weighted.op(Op::Sum).bin(BinOp::TrueDiv, total);
1288            Ok(match simple_column_name(right_tokens) {
1289                Some(column) => node.alias(format!("wavg_{}", column)),
1290                None => node,
1291            })
1292        }
1293        _ => {
1294            // Parse right side first (right-to-left evaluation): it holds any
1295            // remaining operators, so c>c%n becomes c > (c%n).
1296            let right = parse_node(right_tokens)?;
1297            let left = parse_node(left_tokens)?;
1298            apply_op(left, op, right)
1299        }
1300    }
1301}
1302
1303fn apply_op(left: Node, op: &str, right: Node) -> Result<Node, String> {
1304    let op = match op {
1305        "+" => BinOp::Add,
1306        "-" => BinOp::Sub,
1307        "*" => BinOp::Mul,
1308        // `%` divides (q heritage); `/` is the alias everyone expects.
1309        "%" | "/" => BinOp::Div,
1310        "^" => return Ok(Node::Coalesce(Box::new(left), Box::new(right))),
1311        "=" => BinOp::Eq,
1312        "<" => BinOp::Lt,
1313        ">" => BinOp::Gt,
1314        "<=" => BinOp::LtEq,
1315        ">=" => BinOp::GtEq,
1316        "<>" | "!=" => BinOp::Neq,
1317        _ => return Err(format!("Unknown operator: {}", op)),
1318    };
1319    Ok(left.bin(op, right))
1320}
1321
1322/// The column name when the tokens are a bare column reference: `salary`, or
1323/// `col["first name"]` / `col[name]`. Anything more (literals, operators) is None.
1324fn simple_column_name(tokens: &[Token]) -> Option<String> {
1325    match tokens {
1326        [Token::Identifier(name)] => Some(name.clone()),
1327        [
1328            Token::Identifier(c),
1329            Token::LBracket,
1330            Token::String(name) | Token::Identifier(name),
1331            Token::RBracket,
1332        ] if c == "col" => Some(name.clone()),
1333        _ => None,
1334    }
1335}
1336
1337const WAVG_USAGE: &str = "wavg goes between weights and values, e.g. passengers wavg fare";
1338
1339/// Aggregation function names, lowercase.
1340const AGG_FUNCTIONS: [&str; 16] = [
1341    "avg", "mean", "min", "max", "count", "std", "stddev", "dev", "var", "med", "median", "sum",
1342    "first", "last", "nunique", "wavg",
1343];
1344
1345/// Scalar function names, lowercase.
1346const SCALAR_FUNCTIONS: [&str; 13] = [
1347    "len", "length", "not", "null", "upper", "lower", "abs", "floor", "ceil", "ceiling", "sqrt",
1348    "log", "exp",
1349];
1350
1351fn is_agg_function(name: &str) -> bool {
1352    AGG_FUNCTIONS.contains(&name.to_lowercase().as_str())
1353}
1354
1355// Check if an identifier is a known function name
1356fn is_function_name(name: &str) -> bool {
1357    // `wavg` is infix (`w wavg x`), so it never opens an expression.
1358    let name = name.to_lowercase();
1359    name != "wavg" && (is_agg_function(&name) || SCALAR_FUNCTIONS.contains(&name.as_str()))
1360}
1361
1362/// A function call, `fn[args]` or `fn args`. The name is checked before the
1363/// arguments are parsed, so each argument is parsed once; trying aggregates and
1364/// then scalars on the same arguments doubled the work at every nesting level.
1365fn parse_call(name: &str, args: &[Token]) -> Result<Node, String> {
1366    if is_agg_function(name) {
1367        parse_agg_function(name, args)
1368    } else {
1369        parse_function(name, args)
1370    }
1371}
1372
1373// Parse aggregation function like avg[a], min[b], etc.
1374fn parse_agg_function(name: &str, args: &[Token]) -> Result<Node, String> {
1375    if args.is_empty() {
1376        return Err(format!(
1377            "Aggregation function {} requires an argument",
1378            name
1379        ));
1380    }
1381    let fn_name = name.to_lowercase();
1382    if fn_name == "wavg" {
1383        return Err(WAVG_USAGE.to_string());
1384    }
1385    let node = parse_node(args)?;
1386    let op = match fn_name.as_str() {
1387        "avg" | "mean" => Op::Mean,
1388        "min" => Op::Min,
1389        "max" => Op::Max,
1390        "count" => Op::Count,
1391        "std" | "stddev" | "dev" => Op::Std,
1392        "var" => Op::Var,
1393        "med" | "median" => Op::Median,
1394        "sum" => Op::Sum,
1395        "first" => Op::First,
1396        "last" => Op::Last,
1397        "nunique" => Op::NUnique,
1398        _ => return Err(format!("Unknown aggregation function: {}", name)),
1399    };
1400    let node = node.op(op);
1401    // Left unnamed, two aggregates of one column collide ("avg salary, max salary"),
1402    // so a bare-column aggregate is named {fn}_{column}, the convention the dot
1403    // accessors already use. An explicit alias is applied later and overrides this.
1404    match simple_column_name(args) {
1405        Some(column) => Ok(node.alias(format!("{}_{}", fn_name, column))),
1406        None => Ok(node),
1407    }
1408}
1409
1410// Parse function like not[a=b], null[col], len[x], upper[x], etc.
1411fn parse_function(name: &str, args: &[Token]) -> Result<Node, String> {
1412    if args.is_empty() {
1413        return Err(format!("Function {} requires an argument", name));
1414    }
1415    let name_lower = name.to_lowercase();
1416    if !SCALAR_FUNCTIONS.contains(&name_lower.as_str()) {
1417        return Err(format!("Unknown function: {}", name));
1418    }
1419    let node = parse_node(args)?;
1420    let op = match name_lower.as_str() {
1421        "not" => Op::Not,
1422        "null" => Op::IsNull,
1423        "len" | "length" => Op::LenChars,
1424        "upper" => Op::Upper,
1425        "lower" => Op::Lower,
1426        "abs" => Op::Abs,
1427        "floor" => Op::Floor,
1428        "ceil" | "ceiling" => Op::Ceil,
1429        "sqrt" => Op::Sqrt,
1430        "log" => Op::Ln,
1431        "exp" => Op::Exp,
1432        _ => return Err(format!("Unknown function: {}", name)),
1433    };
1434    Ok(node.op(op))
1435}
1436
1437/// An accessor's bracketed argument: `.part["-", 0]` has a string and a number.
1438#[derive(Debug, Clone, PartialEq)]
1439enum AccessorArg {
1440    Str(String),
1441    Num(f64),
1442}
1443
1444impl AccessorArg {
1445    /// As it goes into the result's auto-alias: `part_-_0`, not `part_-_0.0`.
1446    fn alias_text(&self) -> String {
1447        match self {
1448            AccessorArg::Str(s) => s.clone(),
1449            AccessorArg::Num(n) => n.to_string(),
1450        }
1451    }
1452}
1453
1454/// Every accessor with the argument counts it takes and an example for errors.
1455const ACCESSORS: &[(&str, usize, usize, &str)] = &[
1456    // Date and time parts.
1457    ("date", 0, 0, ".date"),
1458    ("time", 0, 0, ".time"),
1459    ("year", 0, 0, ".year"),
1460    ("quarter", 0, 0, ".quarter"),
1461    ("month", 0, 0, ".month"),
1462    ("week", 0, 0, ".week"),
1463    ("day", 0, 0, ".day"),
1464    ("doy", 0, 0, ".doy"),
1465    ("dow", 0, 0, ".dow"),
1466    ("weekday", 0, 0, ".weekday"),
1467    ("hour", 0, 0, ".hour"),
1468    ("minute", 0, 0, ".minute"),
1469    ("second", 0, 0, ".second"),
1470    ("month_start", 0, 0, ".month_start"),
1471    ("month_end", 0, 0, ".month_end"),
1472    ("format", 1, 1, ".format[\"%Y-%m\"]"),
1473    // Strings.
1474    ("len", 0, 0, ".len"),
1475    ("length", 0, 0, ".length"),
1476    ("upper", 0, 0, ".upper"),
1477    ("lower", 0, 0, ".lower"),
1478    ("starts_with", 1, 1, ".starts_with[\"x\"]"),
1479    ("ends_with", 1, 1, ".ends_with[\"x\"]"),
1480    ("contains", 1, 1, ".contains[\"x\"]"),
1481    ("part", 2, 2, ".part[\"-\", 0]"),
1482    ("slice", 1, 2, ".slice[0, 4]"),
1483    ("replace", 2, 2, ".replace[\"(P)\", \"\"]"),
1484    ("strip", 0, 0, ".strip"),
1485    ("to_date", 0, 1, ".to_date[\"%Y%m%d\"]"),
1486    ("to_datetime", 0, 1, ".to_datetime[\"%Y-%m-%d %H:%M\"]"),
1487    // Numbers and casts.
1488    ("round", 0, 1, ".round[1]"),
1489    ("int", 0, 0, ".int"),
1490    ("float", 0, 0, ".float"),
1491    ("str", 0, 0, ".str"),
1492];
1493
1494/// Names for the unknown-accessor error, by kind.
1495const ACCESSOR_HELP: &str = "Valid date/time: date, time, year, quarter, month, week, day, doy, dow, hour, minute, second, month_start, month_end, format. \
1496     Valid string: len, upper, lower, starts_with, ends_with, contains, part, slice, replace, strip, to_date, to_datetime. \
1497     Valid number: round, int, float, str";
1498
1499fn arg_count_text(min: usize, max: usize) -> String {
1500    match (min, max) {
1501        (0, 0) => "no arguments".to_string(),
1502        (1, 1) => "1 argument".to_string(),
1503        (a, b) if a == b => format!("{} arguments", a),
1504        (a, b) => format!("{} to {} arguments", a, b),
1505    }
1506}
1507
1508/// Apply accessor `name` with its bracketed arguments.
1509fn apply_accessor(node: Node, accessor: &str, args: &[AccessorArg]) -> Result<Node, String> {
1510    let name = accessor.to_lowercase();
1511    let Some(&(_, min, max, usage)) = ACCESSORS.iter().find(|(n, ..)| *n == name) else {
1512        return Err(format!(
1513            "Unknown accessor: '{}'. {}",
1514            accessor, ACCESSOR_HELP
1515        ));
1516    };
1517    if args.len() < min || args.len() > max {
1518        return Err(format!(
1519            "{} takes {}, e.g. {}; got {}",
1520            name,
1521            arg_count_text(min, max),
1522            usage,
1523            args.len()
1524        ));
1525    }
1526    let text = |i: usize| match args.get(i) {
1527        Some(AccessorArg::Str(s)) => Ok(s.clone()),
1528        _ => Err(format!(
1529            "{}: argument {} must be quoted text, e.g. {}",
1530            name,
1531            i + 1,
1532            usage
1533        )),
1534    };
1535    let int = |i: usize| match args.get(i) {
1536        Some(AccessorArg::Num(n)) if n.fract() == 0.0 && n.abs() <= u32::MAX as f64 => {
1537            Ok(*n as i64)
1538        }
1539        _ => Err(format!(
1540            "{}: argument {} must be a whole number, e.g. {}",
1541            name,
1542            i + 1,
1543            usage
1544        )),
1545    };
1546    // The string pieces cast first, so they also work on numbers and dates read
1547    // as such (NOAA's DATE, an integer zip code); a cast from String is a no-op.
1548    let as_str = || node.clone().cast_text();
1549    Ok(match name.as_str() {
1550        "date" => node.op(Op::Date),
1551        "time" => node.op(Op::Time),
1552        "year" => node.op(Op::Year),
1553        "quarter" => node.op(Op::Quarter),
1554        "month" => node.op(Op::Month),
1555        "week" => node.op(Op::Week),
1556        "day" => node.op(Op::Day),
1557        "doy" => node.op(Op::OrdinalDay),
1558        "dow" | "weekday" => node.op(Op::Weekday),
1559        "hour" => node.op(Op::Hour),
1560        "minute" => node.op(Op::Minute),
1561        "second" => node.op(Op::Second),
1562        "month_start" => node.op(Op::MonthStart),
1563        "month_end" => node.op(Op::MonthEnd),
1564        "format" => node.op(Op::DtFormat(text(0)?)),
1565        "len" | "length" => node.op(Op::LenChars),
1566        "upper" => node.op(Op::Upper),
1567        "lower" => node.op(Op::Lower),
1568        "starts_with" => node.op(Op::StartsWith(text(0)?)),
1569        "ends_with" => node.op(Op::EndsWith(text(0)?)),
1570        "contains" => node.op(Op::ContainsLiteral(text(0)?)),
1571        "part" => as_str().op(Op::Part(text(0)?, int(1)?)),
1572        "slice" => {
1573            let start = int(0)?;
1574            let length = match args.len() {
1575                2 => {
1576                    let n = int(1)?;
1577                    if n < 0 {
1578                        return Err(format!(
1579                            "slice: the length cannot be negative, e.g. {}",
1580                            usage
1581                        ));
1582                    }
1583                    Some(n as u64)
1584                }
1585                // No length: to the end of the string.
1586                _ => None,
1587            };
1588            as_str().op(Op::Slice(start, length))
1589        }
1590        "replace" => as_str().op(Op::ReplaceAll(text(0)?, text(1)?)),
1591        "strip" => as_str().op(Op::Strip),
1592        "to_date" => as_str().op(Op::ToDate(args.first().map(|_| text(0)).transpose()?)),
1593        "to_datetime" => as_str().op(Op::ToDatetime(args.first().map(|_| text(0)).transpose()?)),
1594        "round" => {
1595            let decimals = if args.is_empty() { 0 } else { int(0)? };
1596            let decimals = u32::try_from(decimals)
1597                .map_err(|_| format!("round: decimals cannot be negative, e.g. {}", usage))?;
1598            node.op(Op::Round(decimals))
1599        }
1600        "int" => node.op(Op::Cast(CastTo::Int64)),
1601        "float" => node.op(Op::Cast(CastTo::Float64)),
1602        "str" => node.op(Op::Cast(CastTo::String)),
1603        _ => {
1604            return Err(format!(
1605                "Unknown accessor: '{}'. {}",
1606                accessor, ACCESSOR_HELP
1607            ));
1608        }
1609    })
1610}
1611
1612/// The arguments inside an accessor's brackets: literals separated by commas.
1613fn parse_accessor_args(accessor: &str, tokens: &[Token]) -> Result<Vec<AccessorArg>, String> {
1614    if tokens.is_empty() {
1615        return Ok(Vec::new());
1616    }
1617    split_tokens(tokens, &Token::Comma)
1618        .iter()
1619        .map(|arg| match arg.as_slice() {
1620            [Token::String(s)] | [Token::Identifier(s)] => Ok(AccessorArg::Str(s.clone())),
1621            [Token::Number(n)] => Ok(AccessorArg::Num(*n)),
1622            [Token::Op(minus), Token::Number(n)] if minus == "-" => Ok(AccessorArg::Num(-n)),
1623            _ => Err(format!(
1624                "{} takes literal arguments, quoted text or numbers, e.g. .part[\"-\", 0]",
1625                accessor
1626            )),
1627        })
1628        .collect()
1629}
1630
1631/// Parse optional dot accessors from remaining tokens. Returns (expr_with_accessors, remaining).
1632/// When base_name is Some, each accessor result is aliased to {base}_{accessor} (or {base}_{acc1}_{acc2} for chained)
1633/// to avoid duplicate column names.
1634fn parse_accessors<'a>(
1635    mut expr: Node,
1636    mut tokens: &'a [Token],
1637    base_name: Option<&str>,
1638) -> Result<(Node, &'a [Token]), String> {
1639    let mut alias_suffix = String::new();
1640    while let [Token::Dot, Token::Identifier(accessor), rest @ ..] = tokens {
1641        let (args, consumed) = if rest.first() == Some(&Token::LBracket) {
1642            let mut depth = 0;
1643            let close = rest
1644                .iter()
1645                .position(|t| {
1646                    match t {
1647                        Token::LBracket => depth += 1,
1648                        Token::RBracket => depth -= 1,
1649                        _ => {}
1650                    }
1651                    depth == 0
1652                })
1653                .ok_or_else(|| format!("Unmatched bracket after .{}", accessor))?;
1654            (parse_accessor_args(accessor, &rest[1..close])?, close + 3)
1655        } else {
1656            (Vec::new(), 2)
1657        };
1658        expr = apply_accessor(expr, accessor, &args)?;
1659        if !alias_suffix.is_empty() {
1660            alias_suffix.push('_');
1661        }
1662        alias_suffix.push_str(accessor);
1663        for arg in &args {
1664            alias_suffix.push('_');
1665            alias_suffix.push_str(&arg.alias_text());
1666        }
1667        tokens = &tokens[consumed..];
1668    }
1669    if !alias_suffix.is_empty() {
1670        let alias = match base_name {
1671            Some(name) => format!("{}_{}", name, alias_suffix),
1672            None => alias_suffix,
1673        };
1674        expr = expr.alias(alias);
1675    }
1676    Ok((expr, tokens))
1677}
1678
1679fn parse_term(tokens: &[Token]) -> Result<(Node, &[Token]), String> {
1680    if tokens.is_empty() {
1681        return Err("Unexpected end of expression".to_string());
1682    }
1683    match &tokens[0] {
1684        Token::Identifier(name) => {
1685            // Check if it's col[...] syntax for column names with spaces
1686            if name == "col" && tokens.len() > 1 && tokens[1] == Token::LBracket {
1687                // Find matching closing bracket
1688                let mut depth = 1;
1689                let mut i = 2;
1690                while i < tokens.len() && depth > 0 {
1691                    match tokens[i] {
1692                        Token::LBracket => depth += 1,
1693                        Token::RBracket => depth -= 1,
1694                        _ => {}
1695                    }
1696                    i += 1;
1697                }
1698                if depth > 0 {
1699                    return Err("Unmatched bracket in col[]".to_string());
1700                }
1701                // Extract column name from inside brackets
1702                let col_name_tokens = &tokens[2..i - 1];
1703                if col_name_tokens.len() != 1 {
1704                    return Err("col[] must contain a single string or identifier".to_string());
1705                }
1706                let col_name = match &col_name_tokens[0] {
1707                    Token::String(s) => s.clone(),
1708                    Token::Identifier(id) => id.clone(),
1709                    _ => return Err("col[] must contain a string or identifier".to_string()),
1710                };
1711                let expr = Node::Col(col_name.clone());
1712                let (expr, remaining) = parse_accessors(expr, &tokens[i..], Some(&col_name))?;
1713                Ok((expr, remaining))
1714            }
1715            // Check if it's a function call (using square brackets)
1716            else if tokens.len() > 1 && tokens[1] == Token::LBracket {
1717                // Find matching closing bracket
1718                let mut depth = 1;
1719                let mut i = 2;
1720                while i < tokens.len() && depth > 0 {
1721                    match tokens[i] {
1722                        Token::LBracket => depth += 1,
1723                        Token::RBracket => depth -= 1,
1724                        _ => {}
1725                    }
1726                    i += 1;
1727                }
1728                if depth > 0 {
1729                    return Err("Unmatched bracket in function call".to_string());
1730                }
1731                let expr = parse_call(name, &tokens[2..i - 1])?;
1732                parse_accessors(expr, &tokens[i..], None)
1733            } else {
1734                // Regular column reference
1735                // (Function calls without brackets are handled in parse_node)
1736                let expr = Node::Col(name.clone());
1737                let (expr, remaining) = parse_accessors(expr, &tokens[1..], Some(name))?;
1738                Ok((expr, remaining))
1739            }
1740        }
1741        Token::Number(n) => Ok((Node::Num(*n), &tokens[1..])), // Numbers don't support accessors
1742        Token::String(s) => Ok((Node::Str(s.clone()), &tokens[1..])), // Strings don't support accessors
1743        Token::DateLiteral(iso) => Ok((Node::Date(iso.clone()), &tokens[1..])),
1744        Token::TimestampLiteral {
1745            iso,
1746            format_str,
1747            time_unit,
1748        } => Ok((
1749            Node::Timestamp {
1750                iso: iso.clone(),
1751                format: format_str.clone(),
1752                unit: *time_unit,
1753                zone: None,
1754            },
1755            &tokens[1..],
1756        )),
1757        Token::LParen => {
1758            let mut depth = 1;
1759            let mut i = 1;
1760            while i < tokens.len() && depth > 0 {
1761                match tokens[i] {
1762                    Token::LParen => depth += 1,
1763                    Token::RParen => depth -= 1,
1764                    _ => {}
1765                }
1766                i += 1;
1767            }
1768            if depth > 0 {
1769                return Err("Unmatched parenthesis".to_string());
1770            }
1771            let inner = parse_node(&tokens[1..i - 1])?;
1772            let (expr, remaining) = parse_accessors(inner, &tokens[i..], None)?;
1773            Ok((expr, remaining))
1774        }
1775        // Square brackets are only for function calls, not grouping
1776        // Parentheses are used for grouping
1777        _ => Err(format!(
1778            "Unexpected '{}' where an expression was expected",
1779            token_text(&tokens[0])
1780        )),
1781    }
1782}
1783
1784/// Deepest chain of nested subexpressions the parser will follow.
1785///
1786/// Parsing is recursive descent, so nesting in the query becomes nesting on the stack:
1787/// `select ------x` recurses once per sign and `select ((((x))))` once per parenthesis.
1788/// Without a ceiling a long enough chain overflows the stack and takes the process with
1789/// it, which is a crash rather than the error message a mistyped query deserves. Found
1790/// by the `parse_query` fuzz target.
1791///
1792/// The ceiling is set by the smallest stack this runs on, not by what is expressible.
1793/// One level of nesting costs a `parse_node` frame and a `parse_term` frame, and in an
1794/// unoptimised build those come to roughly 10 KiB together — enough that a 2 MiB worker
1795/// thread runs out somewhere around 200. 64 leaves a wide margin there and a far wider
1796/// one in a release build, while staying far past any expression written by hand:
1797/// commas and `where` are split off before this runs, so the count is nesting within a
1798/// single expression.
1799const MAX_EXPR_DEPTH: u32 = 64;
1800
1801thread_local! {
1802    static EXPR_DEPTH: std::cell::Cell<u32> = const { std::cell::Cell::new(0) };
1803}
1804
1805/// Holds the recursion counter up for as long as it is alive.
1806///
1807/// Every recursive path in this module passes back through `parse_node`, so counting
1808/// there alone bounds the whole cycle. `parse_node` returns from a dozen places, most
1809/// of them through `?`, so the decrement is tied to the scope rather than written out
1810/// at each exit.
1811struct DepthGuard;
1812
1813impl DepthGuard {
1814    /// `None` once the limit is reached, leaving the counter untouched.
1815    fn enter() -> Option<Self> {
1816        EXPR_DEPTH.with(|depth| {
1817            let next = depth.get() + 1;
1818            if next > MAX_EXPR_DEPTH {
1819                return None;
1820            }
1821            depth.set(next);
1822            Some(DepthGuard)
1823        })
1824    }
1825}
1826
1827impl Drop for DepthGuard {
1828    fn drop(&mut self) {
1829        EXPR_DEPTH.with(|depth| depth.set(depth.get().saturating_sub(1)));
1830    }
1831}
1832
1833// Parse expression with right-to-left operator precedence
1834// This means operators are evaluated from right to left: a+b*c is parsed as a+(b*c)
1835fn parse_node(tokens: &[Token]) -> Result<Node, String> {
1836    let Some(_depth_guard) = DepthGuard::enter() else {
1837        return Err(
1838            "Expression is nested too deeply. Simplify it or split it into steps.".to_string(),
1839        );
1840    };
1841
1842    if tokens.is_empty() {
1843        return Err("Empty expression".to_string());
1844    }
1845
1846    // First check if this starts with a function call (without brackets)
1847    // This needs to be checked before operator parsing to ensure correct precedence
1848    if let Token::Identifier(name) = &tokens[0]
1849        && is_function_name(name)
1850        && tokens.len() > 1
1851        && tokens[1] != Token::LBracket
1852    {
1853        // Function call without brackets - parse the rest as the argument, going
1854        // through the same builders as the bracketed form so both spellings get
1855        // the same expression and the same auto-alias.
1856        return parse_call(name, &tokens[1..]);
1857    }
1858
1859    // Find the leftmost operator for right-to-left evaluation
1860    let mut op_pos = None;
1861    let mut depth = 0;
1862    let mut bracket_depth = 0;
1863
1864    // Scan from left to right to find the leftmost operator
1865    for (i, token) in tokens.iter().enumerate() {
1866        match token {
1867            Token::LParen => depth += 1,
1868            Token::RParen => depth -= 1,
1869            Token::LBracket => bracket_depth += 1,
1870            Token::RBracket => bracket_depth -= 1,
1871            _ if depth == 0 && bracket_depth == 0 && infix_op_at(tokens, i).is_some() => {
1872                op_pos = Some(i);
1873                break;
1874            }
1875            _ => {}
1876        }
1877    }
1878
1879    if let Some(pos) = op_pos {
1880        // Split at the operator
1881        let left_tokens = &tokens[..pos];
1882        let right_tokens = &tokens[pos + 1..];
1883
1884        if let Some(op) = infix_op_at(tokens, pos) {
1885            // Unary minus next to a literal with an operator on the other side: -0.1+discount → (-0.1)+discount
1886            if left_tokens.is_empty()
1887                && op == "-"
1888                && !right_tokens.is_empty()
1889                && matches!(right_tokens[0], Token::Number(_))
1890                && let Token::Number(n) = right_tokens[0]
1891            {
1892                if right_tokens.len() >= 3
1893                    && let Some(bin_op) = infix_op_at(right_tokens, 1)
1894                {
1895                    if WORD_OPS.contains(&bin_op) {
1896                        // The word operators read their operands as tokens, so hand
1897                        // them the negative number as one: -7 mod 3 is (-7) mod 3.
1898                        return apply_infix(&[Token::Number(-n)], bin_op, &right_tokens[2..]);
1899                    }
1900                    let right = parse_node(&right_tokens[2..])?;
1901                    return apply_op(Node::Int(0).bin(BinOp::Sub, Node::Num(n)), bin_op, right);
1902                }
1903                if right_tokens.len() == 1 {
1904                    return Ok(Node::Int(0).bin(BinOp::Sub, Node::Num(n)));
1905                }
1906            }
1907            // Unary plus/minus when there is no left operand (e.g. -x, +x, -(a+b))
1908            if left_tokens.is_empty() && (op == "+" || op == "-") {
1909                let inner = parse_node(right_tokens)?;
1910                return if op == "-" {
1911                    Ok(Node::Int(0).bin(BinOp::Sub, inner))
1912                } else {
1913                    Ok(inner)
1914                };
1915            }
1916            if left_tokens.is_empty() {
1917                return Err("Missing left operand".to_string());
1918            }
1919            apply_infix(left_tokens, op, right_tokens)
1920        } else {
1921            Err("Expected operator".to_string())
1922        }
1923    } else {
1924        // No operator found, parse as term. Every caller hands this a complete
1925        // expression, so leftover tokens are a mistake in the query; dropping them
1926        // here used to make `where x > 1 by dept` silently ignore `by dept`.
1927        let (expr, remaining) = parse_term(tokens)?;
1928        if let Some(extra) = remaining.first() {
1929            if matches!(&tokens[0], Token::Identifier(w) if w == "wavg") {
1930                return Err(WAVG_USAGE.to_string());
1931            }
1932            return Err(format!(
1933                "Unexpected '{}' after the expression",
1934                token_text(extra)
1935            ));
1936        }
1937        Ok(expr)
1938    }
1939}
1940
1941/// A parsed q query, ready to apply to a LazyFrame.
1942#[derive(Debug, Default)]
1943pub struct ParsedQuery {
1944    /// The select list; empty means every column.
1945    pub cols: Vec<Expr>,
1946    /// The where clause, its terms ANDed.
1947    pub filter: Option<Expr>,
1948    /// The by expressions.
1949    pub group_by: Vec<Expr>,
1950    /// Names of the by columns that have one (a plain column or an alias).
1951    pub group_by_names: Vec<String>,
1952    /// `select distinct`: drop duplicate result rows.
1953    pub distinct: bool,
1954}
1955
1956impl ParsedQuery {
1957    /// The query with its casts to text and date parts safe on a date past the
1958    /// calendar, where Polars panics ([`crate::past_calendar::guard_expr`]).
1959    /// `schema` is the data the query runs against: with it, only operations on a
1960    /// date or datetime change.
1961    pub fn past_calendar_safe(self, schema: Option<&Schema>) -> Self {
1962        let guard = |e: Expr| crate::past_calendar::guard_expr(e, schema);
1963        Self {
1964            cols: self.cols.into_iter().map(guard).collect(),
1965            filter: self.filter.map(guard),
1966            group_by: self.group_by.into_iter().map(guard).collect(),
1967            ..self
1968        }
1969    }
1970}
1971
1972/// Convert Polars-specific error messages to user-friendly query errors.
1973pub fn sanitize_query_error(msg: &str) -> String {
1974    let msg_lower = msg.to_lowercase();
1975    if msg_lower.contains("duplicate")
1976        && (msg_lower.contains("output name") || msg_lower.contains("projection"))
1977    {
1978        let name = msg
1979            .split('\'')
1980            .nth(1)
1981            .map(|s| s.to_string())
1982            .unwrap_or_else(|| "column".to_string());
1983        return format!(
1984            "Duplicate column name '{}' in result. Use aliases to rename columns, e.g. `select my_date: timestamp.date`",
1985            name
1986        );
1987    }
1988    if msg_lower.contains(".alias(") || msg_lower.contains("try renaming") {
1989        return "Duplicate column names in result. Use aliases to rename columns, e.g. `select my_date: timestamp.date`"
1990            .to_string();
1991    }
1992    msg.to_string()
1993}
1994
1995/// A q query as parsed, before it becomes Polars expressions: what
1996/// [`parse_query`] runs and "Copy as Python" writes out.
1997#[derive(Debug, Default)]
1998pub(crate) struct QueryNodes {
1999    pub cols: Vec<Node>,
2000    pub filter: Option<Node>,
2001    pub group_by: Vec<Node>,
2002    pub group_by_names: Vec<String>,
2003    pub distinct: bool,
2004}
2005
2006impl QueryNodes {
2007    fn into_parsed(self) -> ParsedQuery {
2008        let lower = |nodes: Vec<Node>| nodes.iter().map(Node::to_expr).collect();
2009        ParsedQuery {
2010            cols: lower(self.cols),
2011            filter: self.filter.as_ref().map(Node::to_expr),
2012            group_by: lower(self.group_by),
2013            group_by_names: self.group_by_names,
2014            distinct: self.distinct,
2015        }
2016    }
2017
2018    /// Each `/` named as Polars runs it over `schema`, the data the query reads; see
2019    /// [`Node::resolve_division`].
2020    pub(crate) fn resolve_division(&mut self, schema: &Schema) {
2021        let nodes = self
2022            .cols
2023            .iter_mut()
2024            .chain(self.filter.iter_mut())
2025            .chain(self.group_by.iter_mut());
2026        for node in nodes {
2027            node.resolve_division(schema);
2028        }
2029    }
2030
2031    /// Each timestamp literal read in the zone of the column it meets in `schema`; see
2032    /// [`Node::resolve_time_zones`].
2033    pub(crate) fn resolve_time_zones(&mut self, schema: &Schema) {
2034        let nodes = self
2035            .cols
2036            .iter_mut()
2037            .chain(self.filter.iter_mut())
2038            .chain(self.group_by.iter_mut());
2039        for node in nodes {
2040            node.resolve_time_zones(schema);
2041        }
2042    }
2043
2044    /// Fails on a temporal column compared with quoted text; see
2045    /// [`Node::check_quoted_temporal`].
2046    fn check_quoted_temporal(&self, schema: &Schema) -> Result<(), String> {
2047        self.cols
2048            .iter()
2049            .chain(self.filter.iter())
2050            .chain(self.group_by.iter())
2051            .try_for_each(|node| node.check_quoted_temporal(schema))
2052    }
2053
2054    /// The where clause as a Python `.filter(...)` call, if there is one.
2055    pub(crate) fn python_filter(&self) -> Option<String> {
2056        self.filter
2057            .as_ref()
2058            .map(|f| format!(".filter({})", f.python()))
2059    }
2060
2061    /// Python method calls doing what `DataTableState::query` does with the query:
2062    /// the where clause, then the grouping (its rows ordered by the keys, which the
2063    /// result names `key_names`) or the select list, then `distinct`.
2064    pub(crate) fn python_steps(&self, key_names: &[String]) -> Vec<String> {
2065        let mut steps: Vec<String> = self.python_filter().into_iter().collect();
2066        if !self.group_by.is_empty() {
2067            let keys = python_list(&self.group_by);
2068            let aggs = if !self.cols.is_empty() {
2069                python_list(&self.cols)
2070            } else if self.group_by_names.is_empty() {
2071                "pl.all()".to_string()
2072            } else {
2073                let names: Vec<String> = self
2074                    .group_by_names
2075                    .iter()
2076                    .map(|n| crate::python_script::py_str(n))
2077                    .collect();
2078                format!("pl.all().exclude({})", names.join(", "))
2079            };
2080            steps.push(format!(".group_by({keys})"));
2081            steps.push(format!(".agg({aggs})"));
2082            steps.push(crate::python_script::sort_call(
2083                key_names,
2084                &vec![false; key_names.len()],
2085            ));
2086        } else if !self.cols.is_empty() {
2087            steps.push(format!(".select({})", python_list(&self.cols)));
2088        }
2089        if self.distinct {
2090            steps.push(".unique(keep=\"first\", maintain_order=True)".to_string());
2091        }
2092        steps
2093    }
2094}
2095
2096/// Expressions as Python arguments: a plain column by its name, as Polars reads a
2097/// string there, anything else as an expression.
2098fn python_list(nodes: &[Node]) -> String {
2099    nodes
2100        .iter()
2101        .map(|n| match n {
2102            Node::Col(name) => crate::python_script::py_str(name),
2103            n => n.python(),
2104        })
2105        .collect::<Vec<_>>()
2106        .join(", ")
2107}
2108
2109pub fn parse_query(query: &str) -> Result<ParsedQuery, String> {
2110    parse_nodes(query).map(QueryNodes::into_parsed)
2111}
2112
2113/// [`parse_query`] for data of `schema`: a timestamp literal compared with a column
2114/// that has a time zone reads as a clock in that zone, and a temporal column compared
2115/// with quoted text is an error.
2116pub fn parse_query_over(query: &str, schema: Option<&Schema>) -> Result<ParsedQuery, String> {
2117    let mut nodes = parse_nodes(query)?;
2118    if let Some(schema) = schema {
2119        nodes.resolve_time_zones(schema);
2120        nodes.check_quoted_temporal(schema)?;
2121    }
2122    Ok(nodes.into_parsed())
2123}
2124
2125/// Parse a q query into nodes. An empty query selects every column.
2126pub(crate) fn parse_nodes(query: &str) -> Result<QueryNodes, String> {
2127    // Empty query is equivalent to "select" - return all columns with no filter or grouping
2128    let trimmed = query.trim();
2129    if trimmed.is_empty() {
2130        return Ok(QueryNodes::default());
2131    }
2132
2133    let tokens = tokenize(query)?;
2134    if tokens.is_empty() || tokens[0] != Token::Select {
2135        return Err("Query must start with 'select'".to_string());
2136    }
2137    // `distinct` right after `select` is the keyword unless what follows makes it
2138    // a column or an alias (`select distinct: x`, `select distinct, a`, `distinct + 1`);
2139    // col["distinct"] always names the column.
2140    let distinct = tokens.get(1) == Some(&Token::Identifier("distinct".to_string()))
2141        && !matches!(
2142            tokens.get(2),
2143            Some(Token::Colon | Token::Comma | Token::Dot | Token::Op(_))
2144        );
2145    let body = strip_from(&tokens[if distinct { 2 } else { 1 }..])?;
2146    let body = &body[..];
2147
2148    // Split by "where" first
2149    let mut parts = split_tokens(body, &Token::Where);
2150    let select_by_tokens = parts.remove(0);
2151    let where_tokens = if !parts.is_empty() {
2152        Some(parts.remove(0))
2153    } else {
2154        None
2155    };
2156    if !parts.is_empty() {
2157        return Err(
2158            "Unexpected second 'where': combine conditions with ',' (and) or '|' (or)".to_string(),
2159        );
2160    }
2161
2162    // `by` after `where` reads naturally but grouping comes first; catch it here,
2163    // outside parentheses and brackets, so the error can name the clause order.
2164    if let Some(ref wt) = where_tokens {
2165        let mut depth = 0;
2166        let mut bracket_depth = 0;
2167        for token in wt {
2168            match token {
2169                Token::LParen => depth += 1,
2170                Token::RParen => depth -= 1,
2171                Token::LBracket => bracket_depth += 1,
2172                Token::RBracket => bracket_depth -= 1,
2173                Token::By if depth == 0 && bracket_depth == 0 => {
2174                    return Err(format!(
2175                        "Unexpected 'by' after the where clause: {}",
2176                        CLAUSE_ORDER
2177                    ));
2178                }
2179                _ => {}
2180            }
2181        }
2182    }
2183
2184    // Split select/by part
2185    let mut select_by_parts = split_tokens(&select_by_tokens, &Token::By);
2186    let cols_tokens = select_by_parts.remove(0);
2187    let by_tokens = if !select_by_parts.is_empty() {
2188        Some(select_by_parts.remove(0))
2189    } else {
2190        None
2191    };
2192    if !select_by_parts.is_empty() {
2193        return Err(format!("Unexpected second 'by': {}", CLAUSE_ORDER));
2194    }
2195
2196    let mut cols = Vec::new();
2197    if !cols_tokens.is_empty() {
2198        for chunk in split_tokens(&cols_tokens, &Token::Comma) {
2199            if chunk.is_empty() {
2200                continue;
2201            }
2202            // Find colon position (if any) - need to account for col[...] syntax
2203            let mut colon_pos = None;
2204            let mut depth = 0;
2205            for (i, token) in chunk.iter().enumerate() {
2206                match token {
2207                    Token::LBracket => depth += 1,
2208                    Token::RBracket => depth -= 1,
2209                    Token::Colon if depth == 0 => {
2210                        colon_pos = Some(i);
2211                        break;
2212                    }
2213                    _ => {}
2214                }
2215            }
2216            if let Some(pos) = colon_pos {
2217                // Has alias: parse left side for alias name, right side for expression
2218                let alias_tokens = &chunk[..pos];
2219                let expr_tokens = &chunk[pos + 1..];
2220
2221                // Parse alias - could be simple identifier or col[...]
2222                let alias_name = if alias_tokens.len() == 1 {
2223                    if let Token::Identifier(name) = &alias_tokens[0] {
2224                        name.clone()
2225                    } else {
2226                        return Err("Expected identifier or col[] for alias".to_string());
2227                    }
2228                } else if alias_tokens.len() == 4
2229                    && alias_tokens[0] == Token::Identifier("col".to_string())
2230                    && alias_tokens[1] == Token::LBracket
2231                    && alias_tokens[3] == Token::RBracket
2232                {
2233                    // col[...] syntax for alias
2234                    match &alias_tokens[2] {
2235                        Token::String(name) | Token::Identifier(name) => name.clone(),
2236                        _ => {
2237                            return Err(
2238                                "Expected string or identifier in col[] for alias".to_string()
2239                            );
2240                        }
2241                    }
2242                } else {
2243                    // Try to parse as expression and extract name (for simple cases)
2244                    // For now, require explicit identifier or col[]
2245                    return Err("Alias must be an identifier or col[]".to_string());
2246                };
2247
2248                let expr = parse_node(expr_tokens)?;
2249                cols.push(expr.alias(alias_name));
2250            } else {
2251                cols.push(parse_node(&chunk)?);
2252            }
2253        }
2254    }
2255
2256    let mut group_by_cols = Vec::new();
2257    let mut group_by_col_names = Vec::new();
2258    if let Some(bt) = by_tokens {
2259        for chunk in split_tokens(&bt, &Token::Comma) {
2260            if chunk.is_empty() {
2261                continue;
2262            }
2263            // Support column assignment in by clause (like select)
2264            // Find colon position (if any) - need to account for col[...] syntax
2265            let mut colon_pos = None;
2266            let mut depth = 0;
2267            for (i, token) in chunk.iter().enumerate() {
2268                match token {
2269                    Token::LBracket => depth += 1,
2270                    Token::RBracket => depth -= 1,
2271                    Token::Colon if depth == 0 => {
2272                        colon_pos = Some(i);
2273                        break;
2274                    }
2275                    _ => {}
2276                }
2277            }
2278            if let Some(pos) = colon_pos {
2279                // Has alias: parse left side for alias name, right side for expression
2280                let alias_tokens = &chunk[..pos];
2281                let expr_tokens = &chunk[pos + 1..];
2282
2283                // Parse alias - could be simple identifier or col[...]
2284                let alias_name = if alias_tokens.len() == 1 {
2285                    if let Token::Identifier(name) = &alias_tokens[0] {
2286                        name.clone()
2287                    } else {
2288                        return Err(
2289                            "Expected identifier or col[] for alias in by clause".to_string()
2290                        );
2291                    }
2292                } else if alias_tokens.len() == 4
2293                    && alias_tokens[0] == Token::Identifier("col".to_string())
2294                    && alias_tokens[1] == Token::LBracket
2295                    && alias_tokens[3] == Token::RBracket
2296                {
2297                    // col[...] syntax for alias
2298                    match &alias_tokens[2] {
2299                        Token::String(name) | Token::Identifier(name) => name.clone(),
2300                        _ => {
2301                            return Err(
2302                                "Expected string or identifier in col[] for alias in by clause"
2303                                    .to_string(),
2304                            );
2305                        }
2306                    }
2307                } else {
2308                    return Err("Alias must be an identifier or col[] in by clause".to_string());
2309                };
2310
2311                let expr = parse_node(expr_tokens)?;
2312                group_by_cols.push(expr.alias(alias_name.clone()));
2313                group_by_col_names.push(alias_name); // Use alias name
2314            } else {
2315                let expr = parse_node(&chunk)?;
2316                group_by_cols.push(expr.clone());
2317                // Try to extract column name from simple Expr
2318                // For simple identifiers: [Token::Identifier(name)]
2319                // For col[] syntax: [Token::Identifier("col"), Token::LBracket, Token::String/Identifier(name), Token::RBracket]
2320                if chunk.len() == 1 {
2321                    if let Token::Identifier(name) = &chunk[0] {
2322                        group_by_col_names.push(name.clone());
2323                    }
2324                } else if chunk.len() == 4
2325                    && chunk[0] == Token::Identifier("col".to_string())
2326                    && chunk[1] == Token::LBracket
2327                    && chunk[3] == Token::RBracket
2328                {
2329                    // col[...] syntax
2330                    match &chunk[2] {
2331                        Token::String(name) | Token::Identifier(name) => {
2332                            group_by_col_names.push(name.clone());
2333                        }
2334                        _ => {}
2335                    }
2336                } else {
2337                    // For complex expressions without alias, we can't extract a simple name
2338                    // The group_by_col_names will be incomplete, but that's okay -
2339                    // we'll use the Expr itself for sorting
2340                }
2341            }
2342        }
2343    }
2344
2345    let mut filter: Option<Node> = None;
2346    if let Some(wt) = where_tokens {
2347        for chunk in split_tokens(&wt, &Token::Comma) {
2348            if chunk.is_empty() {
2349                continue;
2350            }
2351            let mut or_expr: Option<Node> = None;
2352            for or_chunk in split_tokens(&chunk, &Token::Pipe) {
2353                if or_chunk.is_empty() {
2354                    continue;
2355                }
2356                let e = parse_node(&or_chunk)?;
2357                or_expr = match or_expr {
2358                    Some(curr) => Some(curr.bin(BinOp::Or, e)),
2359                    None => Some(e),
2360                };
2361            }
2362            if let Some(e) = or_expr {
2363                filter = match filter {
2364                    Some(curr) => Some(curr.bin(BinOp::And, e)),
2365                    None => Some(e),
2366                };
2367            }
2368        }
2369    }
2370
2371    Ok(QueryNodes {
2372        cols,
2373        filter,
2374        group_by: group_by_cols,
2375        group_by_names: group_by_col_names,
2376        distinct,
2377    })
2378}
2379
2380/// One expression as Polars runs it.
2381#[cfg(test)]
2382fn parse_expr(tokens: &[Token]) -> Result<Expr, String> {
2383    parse_node(tokens).map(|n| n.to_expr())
2384}
2385
2386#[cfg(test)]
2387mod tests {
2388
2389    use super::*;
2390
2391    #[test]
2392
2393    fn test_tokenize_simple() {
2394        let query = "select a, b where a > 10";
2395
2396        let tokens = tokenize(query).unwrap();
2397
2398        assert_eq!(
2399            tokens,
2400            vec![
2401                Token::Select,
2402                Token::Identifier("a".to_string()),
2403                Token::Comma,
2404                Token::Identifier("b".to_string()),
2405                Token::Where,
2406                Token::Identifier("a".to_string()),
2407                Token::Op(">".to_string()),
2408                Token::Number(10.0),
2409            ]
2410        );
2411    }
2412
2413    #[test]
2414
2415    fn test_tokenize_operators() {
2416        let query = "a != b, c >= d, e <= f, g <> h";
2417
2418        let tokens = tokenize(query).unwrap();
2419
2420        assert_eq!(
2421            tokens,
2422            vec![
2423                Token::Identifier("a".to_string()),
2424                Token::Op("!=".to_string()),
2425                Token::Identifier("b".to_string()),
2426                Token::Comma,
2427                Token::Identifier("c".to_string()),
2428                Token::Op(">=".to_string()),
2429                Token::Identifier("d".to_string()),
2430                Token::Comma,
2431                Token::Identifier("e".to_string()),
2432                Token::Op("<=".to_string()),
2433                Token::Identifier("f".to_string()),
2434                Token::Comma,
2435                Token::Identifier("g".to_string()),
2436                Token::Op("<>".to_string()),
2437                Token::Identifier("h".to_string()),
2438            ]
2439        );
2440    }
2441
2442    #[test]
2443
2444    fn test_parse_simple_expr() {
2445        let tokens = tokenize("a + 1").unwrap();
2446
2447        let expr = parse_expr(&tokens).unwrap();
2448
2449        assert_eq!(expr, col("a").add(lit(1.0)));
2450    }
2451
2452    #[test]
2453
2454    fn test_parse_complex_expr() {
2455        let tokens = tokenize("(a + 1) * 2").unwrap();
2456
2457        let expr = parse_expr(&tokens).unwrap();
2458
2459        assert_eq!(expr, (col("a").add(lit(1.0))).mul(lit(2.0)));
2460    }
2461
2462    #[test]
2463
2464    fn test_parse_not_function() {
2465        let query = "select a where not[a = b]";
2466
2467        let filter = parse_query(query).unwrap().filter;
2468
2469        assert_eq!(filter, Some(col("a").eq(col("b")).not()));
2470    }
2471
2472    #[test]
2473
2474    fn test_parse_not_equivalent_to_neq() {
2475        let query1 = "select a where a != b";
2476
2477        let query2 = "select a where not[a = b]";
2478
2479        let query3 = "select a where not a = b";
2480
2481        let filter1 = parse_query(query1).unwrap().filter;
2482
2483        let filter2 = parse_query(query2).unwrap().filter;
2484
2485        let filter3 = parse_query(query3).unwrap().filter;
2486
2487        // All should produce equivalent expressions
2488
2489        assert_eq!(filter1, Some(col("a").neq(col("b"))));
2490
2491        assert_eq!(filter2, Some(col("a").eq(col("b")).not()));
2492
2493        assert_eq!(filter3, Some(col("a").eq(col("b")).not()));
2494    }
2495
2496    #[test]
2497
2498    fn test_parse_avg_without_brackets() {
2499        let query = "select avg 5+a by category";
2500
2501        let cols = parse_query(query).unwrap().cols;
2502
2503        assert_eq!(cols.len(), 1);
2504
2505        // Should parse as avg[(5+a)]
2506    }
2507
2508    #[test]
2509
2510    fn test_parse_string_literal() {
2511        let query = "select a, b:\"foo\"";
2512
2513        let cols = parse_query(query).unwrap().cols;
2514
2515        assert_eq!(cols.len(), 2);
2516
2517        // First column is a, second is b with literal "foo"
2518
2519        assert_eq!(cols[0], col("a"));
2520
2521        assert_eq!(cols[1], lit("foo").alias("b"));
2522    }
2523
2524    #[test]
2525
2526    fn test_parse_string_in_where() {
2527        let query = "select a where name=\"george\", age > 7";
2528
2529        let filter = parse_query(query).unwrap().filter;
2530
2531        // Should have name="george" AND age > 7
2532
2533        assert!(filter.is_some());
2534    }
2535
2536    #[test]
2537
2538    fn test_parse_col_syntax() {
2539        let query = "select col[\"first name\"]";
2540
2541        let cols = parse_query(query).unwrap().cols;
2542
2543        assert_eq!(cols.len(), 1);
2544
2545        assert_eq!(cols[0], col("first name"));
2546    }
2547
2548    #[test]
2549
2550    fn test_parse_col_syntax_with_alias() {
2551        let query = "select a, b:col[\"first name\"]";
2552
2553        let cols = parse_query(query).unwrap().cols;
2554
2555        assert_eq!(cols.len(), 2);
2556
2557        assert_eq!(cols[0], col("a"));
2558
2559        assert_eq!(cols[1], col("first name").alias("b"));
2560    }
2561
2562    #[test]
2563
2564    fn test_parse_col_syntax_with_string_literal() {
2565        let query = "select col[\"first name\"]:\"derek\", foo where foo > 7";
2566
2567        let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
2568
2569        assert_eq!(cols.len(), 2);
2570
2571        assert_eq!(cols[0], lit("derek").alias("first name"));
2572
2573        assert_eq!(cols[1], col("foo"));
2574
2575        assert!(filter.is_some());
2576    }
2577
2578    #[test]
2579
2580    fn test_parse_string_escape_sequences() {
2581        let query = "select a where name=\"george\\\"s name\"";
2582
2583        let filter = parse_query(query).unwrap().filter;
2584
2585        // Should parse escaped quote correctly
2586
2587        assert!(filter.is_some());
2588    }
2589
2590    #[test]
2591
2592    fn test_parse_query_simple_where() {
2593        let query = "select a where a > 10";
2594
2595        let filter = parse_query(query).unwrap().filter;
2596
2597        assert_eq!(filter, Some(col("a").gt(lit(10.0))));
2598    }
2599
2600    #[test]
2601    fn test_parse_query_unary_minus_in_where() {
2602        // Minus next to literal with operator on other side: -0.5+discount → (-0.5)+discount
2603        let query = "select sum total-1 by product where 0<-0.5+discount";
2604        let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
2605        assert_eq!(cols.len(), 1);
2606        assert!(filter.is_some());
2607        // Filter: 0 < (-0.5) + discount
2608        let expected = lit(0.0).lt(lit(0).sub(lit(0.5)).add(col("discount")));
2609        assert_eq!(filter, Some(expected));
2610    }
2611
2612    #[test]
2613    fn test_parse_query_negative_literal_where() {
2614        let query = "select where 0<-0.1+discount";
2615        let filter = parse_query(query).unwrap().filter;
2616        let expected = lit(0.0).lt(lit(0).sub(lit(0.1)).add(col("discount")));
2617        assert_eq!(filter, Some(expected));
2618    }
2619
2620    #[test]
2621    fn test_parse_unary_plus_minus_expr() {
2622        let tokens = tokenize("-0.5").unwrap();
2623        let expr = parse_expr(&tokens).unwrap();
2624        assert_eq!(expr, lit(0).sub(lit(0.5)));
2625        let tokens = tokenize("+x").unwrap();
2626        let expr = parse_expr(&tokens).unwrap();
2627        assert_eq!(expr, col("x"));
2628    }
2629
2630    #[test]
2631
2632    fn test_parse_query_alias() {
2633        let query = "select my_col:a + 1";
2634
2635        let cols = parse_query(query).unwrap().cols;
2636
2637        assert_eq!(cols, vec![col("a").add(lit(1.0)).alias("my_col")]);
2638    }
2639
2640    #[test]
2641
2642    fn test_parse_query_and_or() {
2643        let query = "select a where a > 10 | a < 5, b = 2";
2644
2645        let filter = parse_query(query).unwrap().filter;
2646
2647        let expected =
2648            (col("a").gt(lit(10.0)).or(col("a").lt(lit(5.0)))).and(col("b").eq(lit(2.0)));
2649
2650        assert_eq!(filter, Some(expected));
2651    }
2652
2653    #[test]
2654
2655    fn test_parse_query_neq() {
2656        let query = "select a where a != 10";
2657
2658        let filter = parse_query(query).unwrap().filter;
2659
2660        assert_eq!(filter, Some(col("a").neq(lit(10.0))));
2661    }
2662
2663    #[test]
2664
2665    fn test_parse_query_gte() {
2666        let query = "select a where a >= 10";
2667
2668        let filter = parse_query(query).unwrap().filter;
2669
2670        assert_eq!(filter, Some(col("a").gt_eq(lit(10.0))));
2671    }
2672
2673    #[test]
2674
2675    fn test_parse_query_lte() {
2676        let query = "select a where a <= 10";
2677
2678        let filter = parse_query(query).unwrap().filter;
2679
2680        assert_eq!(filter, Some(col("a").lt_eq(lit(10.0))));
2681    }
2682
2683    #[test]
2684
2685    fn test_empty_query() {
2686        let query = "select";
2687
2688        let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
2689
2690        assert!(cols.is_empty());
2691
2692        assert!(filter.is_none());
2693    }
2694
2695    #[test]
2696
2697    fn test_select_all_implicit() {
2698        let query = "select where a > 1";
2699
2700        let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
2701
2702        assert!(cols.is_empty());
2703
2704        assert_eq!(filter, Some(col("a").gt(lit(1.0))));
2705    }
2706
2707    #[test]
2708
2709    fn test_invalid_query_no_select() {
2710        let query = "a > 10";
2711
2712        let result = parse_query(query);
2713
2714        assert!(result.is_err());
2715    }
2716
2717    #[test]
2718
2719    fn test_invalid_query_unmatched_paren() {
2720        let query = "select (a + 1";
2721
2722        let result = parse_query(query);
2723
2724        assert!(result.is_err());
2725    }
2726
2727    #[test]
2728
2729    fn test_invalid_query_bad_token() {
2730        let query = "select a where a ? 10";
2731
2732        let result = parse_query(query);
2733
2734        assert!(result.is_err());
2735    }
2736
2737    #[test]
2738    fn test_parse_right_to_left_operator_precedence() {
2739        // Test that operators are evaluated right-to-left
2740        // c>c%n should be parsed as c > (c % n), not (c > c) % n
2741        let query = "select t, v where c>c%n";
2742
2743        let filter = parse_query(query).unwrap().filter;
2744
2745        // Should parse as c > (c % n)
2746        let expected = col("c").gt(col("c").div(col("n")));
2747        assert_eq!(filter, Some(expected));
2748    }
2749
2750    // --- Date/datetime accessor tests ---
2751
2752    #[test]
2753    fn test_tokenize_dot_accessor() {
2754        let tokens = tokenize("foo.date").unwrap();
2755        assert_eq!(
2756            tokens,
2757            vec![
2758                Token::Identifier("foo".to_string()),
2759                Token::Dot,
2760                Token::Identifier("date".to_string()),
2761            ]
2762        );
2763    }
2764
2765    #[test]
2766    fn test_tokenize_decimal_number() {
2767        let tokens = tokenize(".5").unwrap();
2768        assert_eq!(tokens, vec![Token::Number(0.5)]);
2769    }
2770
2771    #[test]
2772    fn test_parse_simple_date_accessor() {
2773        let tokens = tokenize("timestamp.date").unwrap();
2774        let expr = parse_expr(&tokens).unwrap();
2775        assert_eq!(expr, col("timestamp").dt().date().alias("timestamp_date"));
2776    }
2777
2778    #[test]
2779    fn test_parse_col_with_date_accessor() {
2780        let tokens = tokenize("col[\"Created At\"].year").unwrap();
2781        let expr = parse_expr(&tokens).unwrap();
2782        assert_eq!(expr, col("Created At").dt().year().alias("Created At_year"));
2783    }
2784
2785    #[test]
2786    fn test_parse_chained_accessors() {
2787        let tokens = tokenize("dt_col.date.year").unwrap();
2788        let expr = parse_expr(&tokens).unwrap();
2789        assert_eq!(
2790            expr,
2791            col("dt_col")
2792                .dt()
2793                .date()
2794                .dt()
2795                .year()
2796                .alias("dt_col_date_year")
2797        );
2798    }
2799
2800    #[test]
2801    fn test_parse_query_select_with_date_accessor() {
2802        let query = "select event_date: timestamp.date";
2803        let cols = parse_query(query).unwrap().cols;
2804        assert_eq!(cols.len(), 1);
2805        assert_eq!(
2806            cols[0],
2807            col("timestamp")
2808                .dt()
2809                .date()
2810                .alias("timestamp_date")
2811                .alias("event_date")
2812        );
2813    }
2814
2815    #[test]
2816    fn test_parse_query_select_col_with_accessor() {
2817        let query = "select col[\"Event Time\"].date, col[\"Event Time\"].year";
2818        let cols = parse_query(query).unwrap().cols;
2819        assert_eq!(cols.len(), 2);
2820        assert_eq!(
2821            cols[0],
2822            col("Event Time").dt().date().alias("Event Time_date")
2823        );
2824        assert_eq!(
2825            cols[1],
2826            col("Event Time").dt().year().alias("Event Time_year")
2827        );
2828    }
2829
2830    #[test]
2831    fn test_parse_query_where_with_date_accessor() {
2832        let query = "select where created_at.month = 12";
2833        let filter = parse_query(query).unwrap().filter;
2834        assert_eq!(
2835            filter,
2836            Some(
2837                col("created_at")
2838                    .dt()
2839                    .month()
2840                    .alias("created_at_month")
2841                    .eq(lit(12.0))
2842            )
2843        );
2844    }
2845
2846    #[test]
2847    fn test_parse_query_where_dow() {
2848        let query = "select where event_ts.dow = 1";
2849        let filter = parse_query(query).unwrap().filter;
2850        assert_eq!(
2851            filter,
2852            Some(
2853                col("event_ts")
2854                    .dt()
2855                    .weekday()
2856                    .alias("event_ts_dow")
2857                    .eq(lit(1.0))
2858            )
2859        );
2860    }
2861
2862    #[test]
2863    fn test_parse_all_accessors() {
2864        let accessors = [
2865            "date",
2866            "time",
2867            "year",
2868            "month",
2869            "week",
2870            "day",
2871            "dow",
2872            "month_start",
2873            "month_end",
2874        ];
2875        for accessor in accessors {
2876            let query = format!("select x.{}", accessor);
2877            let result = parse_query(&query);
2878            assert!(
2879                result.is_ok(),
2880                "Accessor '{}' should parse: {:?}",
2881                accessor,
2882                result.err()
2883            );
2884        }
2885    }
2886
2887    #[test]
2888    fn test_parse_unknown_accessor() {
2889        let query = "select x.nosuchaccessor";
2890        let result = parse_query(query);
2891        assert!(result.is_err());
2892        let err = result.unwrap_err();
2893        assert!(err.contains("Unknown accessor"));
2894        assert!(err.contains("nosuchaccessor"));
2895    }
2896
2897    #[test]
2898    fn test_parse_date_literal() {
2899        let tokens = tokenize("2021.01.01").unwrap();
2900        assert_eq!(tokens, vec![Token::DateLiteral("2021-01-01".to_string())]);
2901    }
2902
2903    #[test]
2904    fn test_parse_query_where_date_literal() {
2905        let query = "select where dt_col.date > 2021.01.01";
2906        let filter = parse_query(query).unwrap().filter;
2907        assert!(filter.is_some());
2908        // Verify the filter parses without error (date literal 2021.01.01 -> ISO 2021-01-01)
2909    }
2910
2911    #[test]
2912    fn test_number_not_parsed_as_date() {
2913        let tokens = tokenize("2.5").unwrap();
2914        assert_eq!(tokens, vec![Token::Number(2.5)]);
2915    }
2916
2917    #[test]
2918    fn test_sanitize_duplicate_column_error() {
2919        let polars_msg = "duplicate: projections contained duplicate output name 'timestamp'. It's possible that multiple expressions are returning the same default column name. If this is the case, try renaming the columns with `.alias(\"new_name\")` to avoid duplicate column names.";
2920        let sanitized = sanitize_query_error(polars_msg);
2921        assert!(sanitized.contains("Duplicate column name"));
2922        assert!(sanitized.contains("timestamp"));
2923        assert!(sanitized.contains("my_date: timestamp.date"));
2924        assert!(!sanitized.contains(".alias("));
2925    }
2926
2927    #[test]
2928    fn test_parse_timestamp_literal() {
2929        let tokens = tokenize("2021.01.15T14:30:00.123456").unwrap();
2930        assert!(matches!(tokens[0], Token::TimestampLiteral { .. }));
2931    }
2932
2933    #[test]
2934    fn test_parse_null_and_not_null() {
2935        let f1 = parse_query("select where null col1").unwrap().filter;
2936        assert!(f1.is_some());
2937        let f2 = parse_query("select where not null col1").unwrap().filter;
2938        assert!(f2.is_some());
2939    }
2940
2941    #[test]
2942    fn test_parse_coalesce() {
2943        let cols = parse_query("select a: coln^cola^colb").unwrap().cols;
2944        assert_eq!(cols.len(), 1);
2945        // coalesce(coln, coalesce(cola, colb)) - parsing succeeds
2946    }
2947
2948    #[test]
2949    fn test_parse_first_last_aggregation() {
2950        let cols = parse_query("select first[value], last[value] by group")
2951            .unwrap()
2952            .cols;
2953        assert_eq!(cols.len(), 2);
2954    }
2955
2956    #[test]
2957    fn test_parse_string_accessors() {
2958        let filter = parse_query("select where city_name.ends_with[\"lanta\"]")
2959            .unwrap()
2960            .filter;
2961        assert!(filter.is_some());
2962        let cols = parse_query("select name.len, name.upper").unwrap().cols;
2963        assert_eq!(cols.len(), 2);
2964    }
2965
2966    #[test]
2967    fn test_parse_format_accessor() {
2968        let tokens = tokenize("dt_col.format[\"%Y-%m\"]").unwrap();
2969        let expr = parse_expr(&tokens).unwrap();
2970        // dt_col.format["%Y-%m"] parses to dt.to_string - verify we got an expr
2971        assert!(!format!("{:?}", expr).is_empty());
2972    }
2973
2974    #[test]
2975    fn test_parse_by_with_date_accessor() {
2976        let query = "select order_date, count: count id by order_date.year";
2977        let ParsedQuery {
2978            cols,
2979            group_by: group_by_cols,
2980            ..
2981        } = parse_query(query).unwrap();
2982        assert_eq!(cols.len(), 2);
2983        assert_eq!(group_by_cols.len(), 1);
2984        assert_eq!(
2985            group_by_cols[0],
2986            col("order_date").dt().year().alias("order_date_year")
2987        );
2988    }
2989
2990    #[test]
2991    fn test_unaliased_aggregates_of_same_column_coexist() {
2992        let query = "select avg salary, max salary by department";
2993        let ParsedQuery {
2994            cols,
2995            group_by: group_by_cols,
2996            ..
2997        } = parse_query(query).unwrap();
2998        assert_eq!(cols.len(), 2);
2999        assert_eq!(cols[0], col("salary").mean().alias("avg_salary"));
3000        assert_eq!(cols[1], col("salary").max().alias("max_salary"));
3001        assert_eq!(group_by_cols, vec![col("department")]);
3002    }
3003
3004    #[test]
3005    fn test_unaliased_aggregate_bracketed_and_bare_name_alike() {
3006        let bracketed = parse_query("select avg[salary] by department")
3007            .unwrap()
3008            .cols;
3009        let bare = parse_query("select avg salary by department").unwrap().cols;
3010        assert_eq!(bracketed, bare);
3011        assert_eq!(bracketed[0], col("salary").mean().alias("avg_salary"));
3012    }
3013
3014    #[test]
3015    fn test_unaliased_aggregate_col_syntax_auto_alias() {
3016        let cols = parse_query("select sum[col[\"unit price\"]] by region")
3017            .unwrap()
3018            .cols;
3019        assert_eq!(cols[0], col("unit price").sum().alias("sum_unit price"));
3020    }
3021
3022    #[test]
3023    fn test_bare_count_names_itself() {
3024        let cols = parse_query("select count[x] by g").unwrap().cols;
3025        assert_eq!(cols[0], col("x").count().alias("count_x"));
3026    }
3027
3028    #[test]
3029    fn test_explicit_alias_overrides_aggregate_auto_alias() {
3030        let cols = parse_query("select total:sum[price] by region")
3031            .unwrap()
3032            .cols;
3033        // The outer alias is applied last, so the result column is named "total".
3034        assert_eq!(
3035            cols[0],
3036            col("price").sum().alias("sum_price").alias("total")
3037        );
3038    }
3039
3040    #[test]
3041    fn test_aggregate_of_expression_keeps_default_name() {
3042        // No single source column, so there is nothing to build a {fn}_{column} name from.
3043        let cols = parse_query("select sum[price*qty] by region").unwrap().cols;
3044        assert_eq!(cols[0], (col("price").mul(col("qty"))).sum());
3045    }
3046
3047    #[test]
3048    fn test_docs_grouping_example_collects_with_auto_aliases() {
3049        // The example from docs/user-guide/querying-data.md must run as written.
3050        let query = "select avg salary, max salary, count name by department";
3051        let ParsedQuery {
3052            cols,
3053            group_by: group_by_cols,
3054            ..
3055        } = parse_query(query).unwrap();
3056        let df = df!(
3057            "department" => &["eng", "eng", "ops"],
3058            "salary" => &[100.0f64, 200.0, 300.0],
3059            "name" => &["a", "b", "c"],
3060        )
3061        .unwrap();
3062        let out = df
3063            .lazy()
3064            .group_by(group_by_cols)
3065            .agg(cols)
3066            .collect()
3067            .unwrap();
3068        let names: Vec<String> = out
3069            .get_column_names()
3070            .iter()
3071            .map(|n| n.to_string())
3072            .collect();
3073        assert_eq!(
3074            names,
3075            ["department", "avg_salary", "max_salary", "count_name"]
3076        );
3077    }
3078
3079    #[test]
3080    fn test_slash_divides_like_percent() {
3081        let slash = parse_expr(&tokenize("a/b").unwrap()).unwrap();
3082        let percent = parse_expr(&tokenize("a%b").unwrap()).unwrap();
3083        assert_eq!(slash, percent);
3084        assert_eq!(slash, col("a").div(col("b")));
3085    }
3086
3087    #[test]
3088    fn test_slash_right_to_left() {
3089        // Right-to-left like every other operator: 1/c+a is 1/(c+a).
3090        let expr = parse_expr(&tokenize("1/c+a").unwrap()).unwrap();
3091        assert_eq!(expr, lit(1.0).div(col("c").add(col("a"))));
3092    }
3093
3094    #[test]
3095    fn test_slash_in_where_clause() {
3096        // Same shape as the existing % test: c>c/n is c > (c/n).
3097        let filter = parse_query("select t, v where c>c/n").unwrap().filter;
3098        assert_eq!(filter, Some(col("c").gt(col("c").div(col("n")))));
3099    }
3100
3101    #[test]
3102    fn test_by_after_where_errors_with_clause_order() {
3103        // The parser used to drop `by dept` on the floor and filter as if it
3104        // were never typed.
3105        let err = parse_query("select name, salary where x > 1 by dept").unwrap_err();
3106        assert!(
3107            err.contains("Unexpected 'by' after the where clause"),
3108            "{err}"
3109        );
3110        assert!(
3111            err.contains("select [by group] [where conditions]"),
3112            "{err}"
3113        );
3114    }
3115
3116    #[test]
3117    fn test_by_after_where_without_condition_operator() {
3118        let err = parse_query("select where flag by dept").unwrap_err();
3119        assert!(
3120            err.contains("Unexpected 'by' after the where clause"),
3121            "{err}"
3122        );
3123    }
3124
3125    #[test]
3126    fn test_by_inside_parens_in_where_errors_as_stray_token() {
3127        // Nested in parentheses it is not a clause boundary, so the expression
3128        // parser reports it instead.
3129        let err = parse_query("select a where (x by g)").unwrap_err();
3130        assert!(
3131            err.contains("Unexpected 'by' after the expression"),
3132            "{err}"
3133        );
3134    }
3135
3136    #[test]
3137    fn test_trailing_garbage_after_where_errors() {
3138        let err = parse_query("select a where a > 1 2").unwrap_err();
3139        assert!(err.contains("Unexpected '2' after the expression"), "{err}");
3140
3141        let err = parse_query("select a where null col1 foo").unwrap_err();
3142        assert!(
3143            err.contains("Unexpected 'foo' after the expression"),
3144            "{err}"
3145        );
3146    }
3147
3148    #[test]
3149    fn test_trailing_garbage_in_select_errors() {
3150        let err = parse_query("select a b").unwrap_err();
3151        assert!(err.contains("Unexpected 'b' after the expression"), "{err}");
3152
3153        let err = parse_query("select (a, b)").unwrap_err();
3154        assert!(err.contains("Unexpected ',' after the expression"), "{err}");
3155    }
3156
3157    #[test]
3158    fn test_duplicate_clauses_error() {
3159        let err = parse_query("select a where x > 1 where y > 2").unwrap_err();
3160        assert!(err.contains("Unexpected second 'where'"), "{err}");
3161        assert!(err.contains("','"), "{err}");
3162
3163        let err = parse_query("select a by g by h").unwrap_err();
3164        assert!(err.contains("Unexpected second 'by'"), "{err}");
3165    }
3166
3167    #[test]
3168    fn test_operators_that_repeat_an_operand_are_bounded() {
3169        // Found by the `parse_query` fuzz target: each `wavg` repeats its operands, so a
3170        // chain of them grew the expression threefold per link until memory ran out.
3171        let chain = format!("select {}x", "w wavg ".repeat(30));
3172        let err = parse_query(&chain).unwrap_err();
3173        assert!(err.contains("Expression is too large"), "{err}");
3174
3175        let xbar = format!("select {}x{}", "(1 xbar ".repeat(40), ")".repeat(40));
3176        assert!(parse_query(&xbar).is_err());
3177
3178        let inner = format!(
3179            "select {}x{}",
3180            "(".repeat(20),
3181            " in [1, 2, 3, 4])".repeat(20)
3182        );
3183        assert!(parse_query(&inner).is_err());
3184
3185        assert!(parse_query("select w wavg x wavg y by g").is_ok());
3186    }
3187
3188    #[test]
3189    fn test_deeply_nested_expression_is_rejected_not_crashed() {
3190        // Found by the `parse_query` fuzz target: the parser is recursive descent, so a
3191        // long enough chain of unary operators or parentheses recursed until the stack
3192        // ran out and the process died. These must come back as errors.
3193        let unary = format!("select {}x", "-".repeat(5_000));
3194        assert!(
3195            parse_query(&unary).is_err(),
3196            "deep unary chain should error"
3197        );
3198
3199        let parens = format!("select {}x{}", "(".repeat(5_000), ")".repeat(5_000));
3200        assert!(parse_query(&parens).is_err(), "deep nesting should error");
3201
3202        // The counter has to come back down, or the first deep query would poison every
3203        // later one on the same thread.
3204        assert!(
3205            parse_query("select a + b * c").is_ok(),
3206            "an ordinary query must still parse after a rejected one"
3207        );
3208    }
3209
3210    // --- q additions (#367) ---
3211
3212    /// Run a query over `df` the way `DataTableState::query` does.
3213    fn eval(query: &str, df: &DataFrame) -> DataFrame {
3214        let ParsedQuery {
3215            cols,
3216            filter,
3217            group_by: by,
3218            distinct,
3219            ..
3220        } = parse_query_over(query, Some(df.schema().as_ref())).unwrap();
3221        let mut lf = df.clone().lazy();
3222        if let Some(f) = filter {
3223            lf = lf.filter(f);
3224        }
3225        if !by.is_empty() {
3226            let keys = by.len();
3227            lf = lf.group_by(by).agg(cols);
3228            let schema = lf.collect_schema().unwrap();
3229            let sort: Vec<Expr> = schema
3230                .iter_names()
3231                .take(keys)
3232                .map(|n| col(n.as_str()))
3233                .collect();
3234            lf = lf.sort_by_exprs(sort, SortMultipleOptions::default());
3235        } else if !cols.is_empty() {
3236            lf = lf.select(cols);
3237        }
3238        if distinct {
3239            lf = lf.unique_stable(None, UniqueKeepStrategy::First);
3240        }
3241        lf.collect().unwrap()
3242    }
3243
3244    /// One column of the result as display strings, nulls as "null".
3245    fn values(df: &DataFrame, name: &str) -> Vec<String> {
3246        df.column(name)
3247            .unwrap()
3248            .as_materialized_series()
3249            .iter()
3250            .map(|v| match v {
3251                AnyValue::String(s) => s.to_string(),
3252                AnyValue::StringOwned(s) => s.to_string(),
3253                v => v.to_string(),
3254            })
3255            .collect()
3256    }
3257
3258    fn parse_err(query: &str) -> String {
3259        parse_query(query).unwrap_err()
3260    }
3261
3262    #[test]
3263    fn test_time_part_accessors_parse() {
3264        let expr = parse_expr(&tokenize("ts.hour").unwrap()).unwrap();
3265        assert_eq!(expr, col("ts").dt().hour().alias("ts_hour"));
3266        let expr = parse_expr(&tokenize("ts.doy").unwrap()).unwrap();
3267        assert_eq!(expr, col("ts").dt().ordinal_day().alias("ts_doy"));
3268        for accessor in ["hour", "minute", "second", "quarter", "doy"] {
3269            let q = format!("select x.{}", accessor);
3270            assert!(parse_query(&q).is_ok(), "{q}");
3271        }
3272    }
3273
3274    #[test]
3275    fn test_time_part_accessors_evaluate() {
3276        let df = df!("ts" => &["2024-03-15 13:45:30", "2024-12-31 00:00:05"])
3277            .unwrap()
3278            .lazy()
3279            .select([col("ts").str().to_datetime(
3280                None,
3281                None,
3282                StrptimeOptions::default(),
3283                lit("raise"),
3284            )])
3285            .collect()
3286            .unwrap();
3287        let out = eval(
3288            "select ts.hour, ts.minute, ts.second, ts.quarter, ts.doy",
3289            &df,
3290        );
3291        assert_eq!(values(&out, "ts_hour"), ["13", "0"]);
3292        assert_eq!(values(&out, "ts_minute"), ["45", "0"]);
3293        assert_eq!(values(&out, "ts_second"), ["30", "5"]);
3294        assert_eq!(values(&out, "ts_quarter"), ["1", "4"]);
3295        assert_eq!(values(&out, "ts_doy"), ["75", "366"]);
3296    }
3297
3298    #[test]
3299    fn test_hour_groups_trips() {
3300        // The taxi example: trips by pickup hour.
3301        let df =
3302            df!("pickup" => &["2025-01-01 08:10:00", "2025-01-01 08:50:00", "2025-01-01 17:00:00"])
3303                .unwrap()
3304                .lazy()
3305                .with_column(col("pickup").str().to_datetime(
3306                    None,
3307                    None,
3308                    StrptimeOptions::default(),
3309                    lit("raise"),
3310                ))
3311                .collect()
3312                .unwrap();
3313        let out = eval("select trips: count pickup by pickup.hour", &df);
3314        assert_eq!(values(&out, "pickup_hour"), ["8", "17"]);
3315        assert_eq!(values(&out, "trips"), ["2", "1"]);
3316    }
3317
3318    /// A timestamp literal reads as a clock in the zone of the column it meets, as
3319    /// the table shows that column; Polars refuses a zoned/naive comparison otherwise.
3320    #[test]
3321    fn a_timestamp_literal_takes_the_zone_of_its_column() {
3322        let zoned = |zone: &str| {
3323            df!("t" => &["2013-01-15 14:00:00", "2013-01-15 15:00:00"])
3324                .unwrap()
3325                .lazy()
3326                .with_column(col("t").str().to_datetime(
3327                    Some(TimeUnit::Microseconds),
3328                    TimeZone::opt_try_new(Some(zone)).unwrap(),
3329                    StrptimeOptions::default(),
3330                    lit("raise"),
3331                ))
3332                .collect()
3333                .unwrap()
3334        };
3335        for zone in ["UTC", "America/New_York"] {
3336            let df = zoned(zone);
3337            for (query, rows) in [
3338                ("select where t > 2013.01.15T14:30:00.123456", 1),
3339                ("select where 2013.01.15T14:30:00 < t", 1),
3340                ("select where t = 2013.01.15T15:00:00", 1),
3341                (
3342                    "select where t >= 2013.01.15T14:00:00, t < 2013.01.16T00:00:00",
3343                    2,
3344                ),
3345                ("select where t > 2013.01.15", 2),
3346            ] {
3347                assert_eq!(eval(query, &df).height(), rows, "{zone}: {query}");
3348            }
3349            let out = eval("select later: t ^ 2013.01.15T00:00:00", &df);
3350            assert_eq!(out.height(), 2, "{zone}");
3351        }
3352        // A column with no zone is untouched.
3353        let naive = df!("t" => &["2013-01-15 14:00:00"])
3354            .unwrap()
3355            .lazy()
3356            .with_column(col("t").str().to_datetime(
3357                None,
3358                None,
3359                StrptimeOptions::default(),
3360                lit("raise"),
3361            ))
3362            .collect()
3363            .unwrap();
3364        assert_eq!(
3365            eval("select where t < 2013.01.15T14:30:00", &naive).height(),
3366            1
3367        );
3368
3369        // "Copy as Python" says the same.
3370        let schema = zoned("America/New_York").schema().clone();
3371        let mut nodes = parse_nodes("select where t > 2013.01.15T14:30:00").unwrap();
3372        nodes.resolve_time_zones(&schema);
3373        let python = nodes.python_filter().unwrap();
3374        assert!(
3375            python.contains("time_zone=\"America/New_York\", ambiguous=\"earliest\""),
3376            "{python}"
3377        );
3378    }
3379
3380    /// One row of each temporal type, plus text, for the quoted-text errors.
3381    fn temporal_frame() -> DataFrame {
3382        df!("d" => &["2024-01-01"], "s" => &["2024.01.01"])
3383            .unwrap()
3384            .lazy()
3385            .with_columns([
3386                lit("2024-01-01T05:00:00")
3387                    .str()
3388                    .to_datetime(None, None, StrptimeOptions::default(), lit("raise"))
3389                    .alias("ts"),
3390                lit("2024-01-01T05:00:00")
3391                    .str()
3392                    .to_datetime(None, None, StrptimeOptions::default(), lit("raise"))
3393                    .dt()
3394                    .time()
3395                    .alias("t"),
3396                lit(5i64)
3397                    .cast(DataType::Duration(TimeUnit::Milliseconds))
3398                    .alias("dur"),
3399                col("d").str().to_date(StrptimeOptions::default()),
3400            ])
3401            .collect()
3402            .unwrap()
3403    }
3404
3405    /// The error parsing `query` over `df`.
3406    fn parse_error_over(query: &str, df: &DataFrame) -> String {
3407        parse_query_over(query, Some(df.schema().as_ref()))
3408            .err()
3409            .unwrap_or_else(|| panic!("{query} should fail"))
3410    }
3411
3412    #[test]
3413    fn test_quoted_text_against_temporal_column_is_a_q_error() {
3414        let df = temporal_frame();
3415        let cases = [
3416            (
3417                "select where d = \"2024.01.01\"",
3418                "d is a date; \"2024.01.01\" is a string. A date is 2024.01.01",
3419            ),
3420            (
3421                "select where d < \"Jan 1\"",
3422                "d is a date; \"Jan 1\" is a string. A date is 2024.01.01",
3423            ),
3424            (
3425                "select where d = \"2024-01-01\"",
3426                "d is a date; \"2024-01-01\" is a string. A date is 2024.01.01",
3427            ),
3428            (
3429                "select where ts = \"2024.01.01T05:00:00\"",
3430                "ts is a timestamp; \"2024.01.01T05:00:00\" is a string. A timestamp is 2024.01.01T05:00:00",
3431            ),
3432            (
3433                "select where ts < \"2023.06.30T23:59:59.5\"",
3434                "ts is a timestamp; \"2023.06.30T23:59:59.5\" is a string. A timestamp is 2023.06.30T23:59:59.5",
3435            ),
3436            (
3437                "select where ts < \"2024.01.01\"",
3438                "ts is a timestamp; \"2024.01.01\" is a string. A timestamp is 2024.01.01T05:00:00",
3439            ),
3440            (
3441                "select where t = \"05:00:00\"",
3442                "t is a time; \"05:00:00\" is a string. A time has no literal; compare t.hour, t.minute or t.second with a number",
3443            ),
3444            (
3445                "select where t < \"05:00:00\"",
3446                "t is a time; \"05:00:00\" is a string. A time has no literal; compare t.hour, t.minute or t.second with a number",
3447            ),
3448            (
3449                "select where dur = \"5s\"",
3450                "dur is a duration; \"5s\" is a string. A duration has no literal",
3451            ),
3452            (
3453                "select where \"5s\" >= dur",
3454                "dur is a duration; \"5s\" is a string. A duration has no literal",
3455            ),
3456            (
3457                "select where d in [\"2024.01.01\", \"2024.01.02\"]",
3458                "d is a date; \"2024.01.01\" is a string. A date is 2024.01.01",
3459            ),
3460            (
3461                "select x: d != \"x\"",
3462                "d is a date; \"x\" is a string. A date is 2024.01.01",
3463            ),
3464        ];
3465        for (query, want) in cases {
3466            assert_eq!(parse_error_over(query, &df), want, "{query}");
3467        }
3468
3469        // A name that is not a bare word is named as it is typed.
3470        let mut renamed = df.clone();
3471        renamed.rename("d", "start date".into()).unwrap();
3472        assert_eq!(
3473            parse_error_over("select where col[\"start date\"] = \"x\"", &renamed),
3474            "col[\"start date\"] is a date; \"x\" is a string. A date is 2024.01.01"
3475        );
3476
3477        // The remedies run.
3478        assert_eq!(eval("select where d = 2024.01.01", &df).height(), 1);
3479        assert_eq!(eval("select where d in [2024.01.01]", &df).height(), 1);
3480        assert_eq!(
3481            eval("select where ts = 2024.01.01T05:00:00", &df).height(),
3482            1
3483        );
3484        assert_eq!(eval("select where t.hour = 5", &df).height(), 1);
3485    }
3486
3487    #[test]
3488    fn test_quoted_text_against_text_column_still_compares() {
3489        let df = temporal_frame();
3490        assert_eq!(eval("select where s = \"2024.01.01\"", &df).height(), 1);
3491        assert_eq!(eval("select where s < \"2025\"", &df).height(), 1);
3492        assert_eq!(eval("select where s in [\"2024.01.01\"]", &df).height(), 1);
3493        // `like` and string functions are text by name; left to themselves.
3494        assert!(
3495            parse_query_over("select where d like \"2024*\"", Some(df.schema().as_ref())).is_ok()
3496        );
3497        // Without a schema nothing is known, so nothing is refused here.
3498        assert!(parse_query_over("select where d = \"2024.01.01\"", None).is_ok());
3499    }
3500
3501    #[test]
3502    fn test_to_date_and_to_datetime_parse_strings() {
3503        let df = df!(
3504            "DATE" => &["20240101", "20241231", "junk"],
3505            "Date" => &["Sat Sep 12 2020", "Tue Jan 12 2021(P)", "Sun Sep 13 2020"],
3506            "stamp" => &["2024-01-02 03:04", "2024-05-06 07:08", "nope"],
3507        )
3508        .unwrap();
3509        let out = eval(
3510            "select day: DATE.to_date[\"%Y%m%d\"], d: Date.replace[\"(P)\", \"\"].to_date[\"%a %b %d %Y\"], t: stamp.to_datetime[\"%Y-%m-%d %H:%M\"]",
3511            &df,
3512        );
3513        // A value that does not match the format is null, not an error.
3514        assert_eq!(values(&out, "day"), ["2024-01-01", "2024-12-31", "null"]);
3515        assert_eq!(
3516            values(&out, "d"),
3517            ["2020-09-12", "2021-01-12", "2020-09-13"]
3518        );
3519        assert_eq!(
3520            values(&out, "t"),
3521            ["2024-01-02 03:04:00", "2024-05-06 07:08:00", "null"]
3522        );
3523    }
3524
3525    #[test]
3526    fn test_to_date_parses_an_integer_column() {
3527        // NOAA's DATE is 20240101; read from CSV it is an integer.
3528        let df = df!("DATE" => &[20240101i64, 20240229]).unwrap();
3529        let out = eval("select d: DATE.to_date[\"%Y%m%d\"]", &df);
3530        assert_eq!(values(&out, "d"), ["2024-01-01", "2024-02-29"]);
3531    }
3532
3533    #[test]
3534    fn test_casts() {
3535        let df = df!(
3536            "s" => &["3", "4.5", "x"],
3537            "f" => &[1.9f64, -1.9, 3.0],
3538        )
3539        .unwrap();
3540        let out = eval("select a: s.int, b: s.float, c: f.int, d: f.str", &df);
3541        assert_eq!(values(&out, "a"), ["3", "null", "null"]);
3542        assert_eq!(values(&out, "b"), ["3.0", "4.5", "null"]);
3543        assert_eq!(values(&out, "c"), ["1", "-1", "3"]);
3544        assert_eq!(values(&out, "d"), ["1.9", "-1.9", "3.0"]);
3545        assert_eq!(out.column("a").unwrap().dtype(), &DataType::Int64);
3546        assert_eq!(out.column("b").unwrap().dtype(), &DataType::Float64);
3547        assert_eq!(out.column("d").unwrap().dtype(), &DataType::String);
3548    }
3549
3550    #[test]
3551    fn test_string_pieces() {
3552        let df = df!("FT" => &["0–3", "12–1", "  2–2  "]).unwrap();
3553        let out = eval(
3554            "select home: FT.part[\"–\", 0].int, away: FT.part[\"–\", -1].int, none: FT.part[\"–\", 5], head: FT.slice[0, 2], tail: FT.slice[-2], s: FT.strip, r: FT.replace[\"–\", \"-\"]",
3555            &df,
3556        );
3557        assert_eq!(values(&out, "home"), ["0", "12", "null"]);
3558        assert_eq!(values(&out, "away"), ["3", "1", "null"]);
3559        assert_eq!(values(&out, "none"), ["null", "null", "null"]);
3560        assert_eq!(values(&out, "head"), ["0–", "12", "  "]);
3561        assert_eq!(values(&out, "tail"), ["–3", "–1", "  "]);
3562        assert_eq!(values(&out, "s"), ["0–3", "12–1", "2–2"]);
3563        assert_eq!(values(&out, "r"), ["0-3", "12-1", "  2-2  "]);
3564    }
3565
3566    #[test]
3567    fn test_string_pieces_auto_alias() {
3568        let cols = parse_query("select FT.part[\"-\", 0], FT.strip")
3569            .unwrap()
3570            .cols;
3571        let names: Vec<String> = cols
3572            .iter()
3573            .map(|e| e.clone().meta().output_name().unwrap().to_string())
3574            .collect();
3575        assert_eq!(names, ["FT_part_-_0", "FT_strip"]);
3576    }
3577
3578    #[test]
3579    fn test_in_parses_to_equalities() {
3580        let filter = parse_query("select where name in [\"a\", \"b\"]")
3581            .unwrap()
3582            .filter;
3583        assert_eq!(
3584            filter,
3585            Some(col("name").eq(lit("a")).or(col("name").eq(lit("b"))))
3586        );
3587    }
3588
3589    #[test]
3590    fn test_in_filters() {
3591        let df = df!(
3592            "name" => &["Emma", "Jennifer", "Olivia", "Mary"],
3593            "n" => &[1i32, 2, 3, 4],
3594        )
3595        .unwrap();
3596        let out = eval(
3597            "select name where name in [\"Emma\", \"Jennifer\", \"Olivia\"]",
3598            &df,
3599        );
3600        assert_eq!(values(&out, "name"), ["Emma", "Jennifer", "Olivia"]);
3601        let out = eval("select n where n in [2, 4.0, -1]", &df);
3602        assert_eq!(values(&out, "n"), ["2", "4"]);
3603        let out = eval("select name where not name in [\"Mary\"]", &df);
3604        assert_eq!(values(&out, "name"), ["Emma", "Jennifer", "Olivia"]);
3605        // Commas inside the list are not where-clause ANDs.
3606        let out = eval("select name where name in [\"Emma\", \"Mary\"], n > 1", &df);
3607        assert_eq!(values(&out, "name"), ["Mary"]);
3608    }
3609
3610    #[test]
3611    fn test_in_long_list_nests_shallowly() {
3612        let items: Vec<String> = (0..2000).map(|i| i.to_string()).collect();
3613        let q = format!("select where x in [{}]", items.join(", "));
3614        let df = df!("x" => &[5i64, 1999, 2000]).unwrap();
3615        assert_eq!(values(&eval(&q, &df), "x"), ["5", "1999"]);
3616    }
3617
3618    #[test]
3619    fn test_in_a_list_past_the_node_cap_on_a_column() {
3620        // A pasted list of ids: the column is not what multiplies.
3621        let items: Vec<String> = (0..12_000).map(|i| i.to_string()).collect();
3622        let q = format!("select where x in [{}]", items.join(", "));
3623        let df = df!("x" => &[5i64, 11_999, 12_000]).unwrap();
3624        assert_eq!(values(&eval(&q, &df), "x"), ["5", "11999"]);
3625    }
3626
3627    #[test]
3628    fn test_in_errors() {
3629        let err = parse_err("select where x in 1");
3630        assert!(err.contains("in takes a list"), "{err}");
3631        let err = parse_err("select where x in []");
3632        assert!(err.contains("in needs a list of values"), "{err}");
3633        let err = parse_err("select where x in [1,, 2]");
3634        assert!(err.contains("in needs a list of values"), "{err}");
3635        let err = parse_err("select where x in [1] = y");
3636        assert!(err.contains("in takes a list"), "{err}");
3637        let err = parse_err("select where x in [1] + [2]");
3638        assert!(err.contains("in takes a list"), "{err}");
3639    }
3640
3641    #[test]
3642    fn test_like_matches_whole_value() {
3643        let df = df!("item" => &["Crispy Chicken", "Chicken", "Fish", "a.b", "axb"]).unwrap();
3644        let out = eval("select item where item like \"*Chicken*\"", &df);
3645        assert_eq!(values(&out, "item"), ["Crispy Chicken", "Chicken"]);
3646        // Anchored: a prefix pattern does not match mid-string.
3647        let out = eval("select item where item like \"Chick*\"", &df);
3648        assert_eq!(values(&out, "item"), ["Chicken"]);
3649        let out = eval("select item where item like \"F?sh\"", &df);
3650        assert_eq!(values(&out, "item"), ["Fish"]);
3651        // Regex characters are literal.
3652        let out = eval("select item where item like \"a.b\"", &df);
3653        assert_eq!(values(&out, "item"), ["a.b"]);
3654    }
3655
3656    #[test]
3657    fn test_like_errors() {
3658        let err = parse_err("select where item like Chicken");
3659        assert!(err.contains("like takes a quoted pattern"), "{err}");
3660    }
3661
3662    #[test]
3663    fn test_like_regex() {
3664        assert_eq!(like_regex("*a?.b*"), "(?s)^.*a.\\.b.*$");
3665    }
3666
3667    /// One column of each integer width and signedness, `a_*` the larger and `b_*`
3668    /// the smaller, so a difference never wraps an unsigned type.
3669    fn integer_widths() -> (DataFrame, Vec<&'static str>) {
3670        let names = ["i8", "i16", "i32", "i64", "u8", "u16", "u32", "u64"];
3671        let types = [
3672            DataType::Int8,
3673            DataType::Int16,
3674            DataType::Int32,
3675            DataType::Int64,
3676            DataType::UInt8,
3677            DataType::UInt16,
3678            DataType::UInt32,
3679            DataType::UInt64,
3680        ];
3681        let mut columns = Vec::new();
3682        for (name, dtype) in names.iter().zip(types) {
3683            for (side, vals) in [("a", [10i64, 20, 30]), ("b", [1, 2, 3])] {
3684                let c = Column::new(format!("{side}_{name}").into(), vals);
3685                columns.push(c.cast(&dtype).unwrap());
3686            }
3687        }
3688        let df = DataFrame::new_infer_height(columns).unwrap();
3689        (df, names.to_vec())
3690    }
3691
3692    fn as_f64(df: &DataFrame, name: &str) -> Vec<f64> {
3693        df.column(name)
3694            .unwrap()
3695            .cast(&DataType::Float64)
3696            .unwrap()
3697            .f64()
3698            .unwrap()
3699            .into_no_null_iter()
3700            .collect()
3701    }
3702
3703    #[test]
3704    fn mixed_integer_widths_do_arithmetic() {
3705        let (df, names) = integer_widths();
3706        // a = 10, 20, 30 and b = 1, 2, 3 in every width.
3707        let ops: [(&str, [f64; 3]); 5] = [
3708            ("+", [11.0, 22.0, 33.0]),
3709            ("-", [9.0, 18.0, 27.0]),
3710            ("*", [10.0, 40.0, 90.0]),
3711            ("/", [10.0, 10.0, 10.0]),
3712            ("%", [10.0, 10.0, 10.0]),
3713        ];
3714        let mut parts = Vec::new();
3715        let mut expected = Vec::new();
3716        for (i, (op, want)) in ops.iter().enumerate() {
3717            for l in &names {
3718                for r in &names {
3719                    let alias = format!("r{i}_{l}_{r}");
3720                    parts.push(format!("{alias}: a_{l} {op} b_{r}"));
3721                    expected.push((alias, *want));
3722                }
3723            }
3724        }
3725        for l in &names {
3726            for r in &names {
3727                let alias = format!("m_{l}_{r}");
3728                parts.push(format!("{alias}: a_{l} mod b_{r}"));
3729                expected.push((alias, [0.0, 0.0, 0.0]));
3730            }
3731        }
3732        let out = eval(&format!("select {}", parts.join(", ")), &df);
3733        for (alias, want) in expected {
3734            assert_eq!(as_f64(&out, &alias), want, "{alias}");
3735        }
3736    }
3737
3738    #[test]
3739    fn mixed_integer_widths_with_literals() {
3740        let (df, names) = integer_widths();
3741        let mut parts = Vec::new();
3742        let mut expected = Vec::new();
3743        for name in &names {
3744            for (tag, expr, want) in [
3745                ("p", format!("b_{name} + 7"), [8.0, 9.0, 10.0]),
3746                ("s", format!("a_{name} - 7"), [3.0, 13.0, 23.0]),
3747                ("l", format!("7 - b_{name}"), [6.0, 5.0, 4.0]),
3748                ("m", format!("b_{name} * 2.5"), [2.5, 5.0, 7.5]),
3749                ("d", format!("a_{name} / 2"), [5.0, 10.0, 15.0]),
3750                ("q", format!("60 / b_{name}"), [60.0, 30.0, 20.0]),
3751                ("r", format!("a_{name} mod 7"), [3.0, 6.0, 2.0]),
3752            ] {
3753                let alias = format!("{tag}_{name}");
3754                parts.push(format!("{alias}: {expr}"));
3755                expected.push((alias, want));
3756            }
3757        }
3758        let out = eval(&format!("select {}", parts.join(", ")), &df);
3759        for (alias, want) in expected {
3760            assert_eq!(as_f64(&out, &alias), want, "{alias}");
3761        }
3762    }
3763
3764    #[test]
3765    fn mixed_integer_widths_filter_and_group() {
3766        let (df, _) = integer_widths();
3767        let out = eval("select a_u8 where 10 = a_i64 / b_u8, 8 < b_u16 + 7", &df);
3768        assert_eq!(values(&out, "a_u8"), ["20", "30"]);
3769        let out = eval(
3770            "select t: sum a_u8 * b_i16, n: max a_i64 - b_u32 by k: b_u16 mod 2",
3771            &df,
3772        );
3773        assert_eq!(as_f64(&out, "k"), [0.0, 1.0]);
3774        assert_eq!(as_f64(&out, "t"), [40.0, 100.0]);
3775        assert_eq!(as_f64(&out, "n"), [18.0, 27.0]);
3776    }
3777
3778    #[cfg(feature = "sql")]
3779    #[test]
3780    fn mixed_integer_widths_in_sql() {
3781        let (df, _) = integer_widths();
3782        let mut ctx = polars_sql::SQLContext::new();
3783        ctx.register("df", df.lazy());
3784        let out = ctx
3785            .execute(
3786                "SELECT a_i64 / b_u8 AS d, a_u16 % b_u8 AS m, b_u8 + 7 AS p, \
3787                 a_i8 * b_u64 AS x FROM df WHERE a_u8 - b_i16 > 9",
3788            )
3789            .unwrap()
3790            .collect()
3791            .unwrap();
3792        assert_eq!(as_f64(&out, "d"), [10.0, 10.0]);
3793        assert_eq!(as_f64(&out, "m"), [0.0, 0.0]);
3794        assert_eq!(as_f64(&out, "p"), [9.0, 10.0]);
3795        assert_eq!(as_f64(&out, "x"), [40.0, 90.0]);
3796    }
3797
3798    #[test]
3799    fn test_xbar_buckets() {
3800        let df = df!(
3801            "fare" => &[-1.0f64, 0.0, 4.99, 5.0, 12.5],
3802            "n" => &[-1i64, 0, 4, 5, 12],
3803        )
3804        .unwrap();
3805        let out = eval("select f: 5 xbar fare, i: 5 xbar n, h: 0.5 xbar fare", &df);
3806        assert_eq!(values(&out, "f"), ["-5.0", "0.0", "0.0", "5.0", "10.0"]);
3807        // A whole-number bucket keeps an integer column integral.
3808        assert_eq!(values(&out, "i"), ["-5", "0", "0", "5", "10"]);
3809        assert_eq!(out.column("i").unwrap().dtype(), &DataType::Int64);
3810        assert_eq!(values(&out, "h"), ["-1.0", "0.0", "4.5", "5.0", "12.5"]);
3811    }
3812
3813    #[test]
3814    fn test_xbar_groups() {
3815        let df = df!("fare" => &[1.0f64, 3.0, 7.0, 12.0, 14.0]).unwrap();
3816        let out = eval("select trips: count fare by b: 5 xbar fare", &df);
3817        assert_eq!(values(&out, "b"), ["0.0", "5.0", "10.0"]);
3818        assert_eq!(values(&out, "trips"), ["2", "1", "2"]);
3819    }
3820
3821    #[test]
3822    fn test_xbar_errors() {
3823        let err = parse_err("select 0 xbar fare");
3824        assert!(err.contains("positive bucket size"), "{err}");
3825        let err = parse_err("select -5 xbar fare");
3826        assert!(err.contains("positive bucket size"), "{err}");
3827    }
3828
3829    #[test]
3830    fn test_mod() {
3831        let df = df!("n" => &[-7i64, 7, 9], "f" => &[7.5f64, -0.5, 2.0]).unwrap();
3832        let out = eval("select a: n mod 3, b: f mod 2, c: -7 mod 3", &df);
3833        // Floored, as in q: the result takes the sign of the divisor.
3834        assert_eq!(values(&out, "a"), ["2", "1", "0"]);
3835        assert_eq!(values(&out, "b"), ["1.5", "1.5", "0.0"]);
3836        assert_eq!(values(&out, "c"), ["2", "2", "2"]);
3837        // A negative whole divisor is an integer too.
3838        let out = eval("select a: n mod -3", &df);
3839        assert_eq!(values(&out, "a"), ["-1", "-2", "0"]);
3840        assert_eq!(out.column("a").unwrap().dtype(), &DataType::Int64);
3841    }
3842
3843    #[test]
3844    fn test_word_operators_right_to_left() {
3845        let parse = |s: &str| parse_expr(&tokenize(s).unwrap()).unwrap();
3846        // a = b mod 2 is a = (b mod 2).
3847        assert_eq!(parse("a = b mod 2"), col("a").eq(col("b").rem(lit(2i64))));
3848        // 5 xbar x + 1 buckets x + 1; 2 * 5 xbar x doubles the bucket.
3849        assert_eq!(
3850            parse("5 xbar x + 1"),
3851            col("x").add(lit(1.0)).floor_div(lit(5i64)).mul(lit(5i64))
3852        );
3853        assert_eq!(
3854            parse("2 * 5 xbar x"),
3855            lit(2.0).mul(col("x").floor_div(lit(5i64)).mul(lit(5i64)))
3856        );
3857        // flag = name in [...] compares flag with the membership test.
3858        assert_eq!(
3859            parse("flag = name in [\"a\"]"),
3860            col("flag").eq(col("name").eq(lit("a")))
3861        );
3862        assert_eq!(parse("x mod 2 in [1]"), col("x").rem(lit(2.0).eq(lit(1.0))));
3863        // A symbol operator to the left of like takes the whole like as its right side.
3864        assert_eq!(
3865            parse("ok = name like \"a*\""),
3866            col("ok").eq(col("name")
3867                .cast(DataType::String)
3868                .str()
3869                .contains(lit("(?s)^a.*$"), true))
3870        );
3871    }
3872
3873    #[test]
3874    fn test_word_operators_right_to_left_evaluate() {
3875        let df = df!("x" => &[3i64, 4, 9]).unwrap();
3876        // 1 + x mod 4 is 1 + (x mod 4), not (1 + x) mod 4.
3877        let out = eval("select a: 1 + x mod 4", &df);
3878        assert_eq!(values(&out, "a"), ["4.0", "1.0", "2.0"]);
3879        let out = eval("select a: (1 + x) mod 4", &df);
3880        assert_eq!(values(&out, "a"), ["0.0", "1.0", "2.0"]);
3881        // x mod 2 in [1] would be x mod (2 in [1]); parentheses test the remainder.
3882        let out = eval("select x where (x mod 2) in [1]", &df);
3883        assert_eq!(values(&out, "x"), ["3", "9"]);
3884    }
3885
3886    #[test]
3887    fn test_word_operators_are_still_column_names() {
3888        let cols = parse_query("select in, mod, like + xbar").unwrap().cols;
3889        assert_eq!(cols[0], col("in"));
3890        assert_eq!(cols[1], col("mod"));
3891        assert_eq!(cols[2], col("like").add(col("xbar")));
3892        let err = parse_err("select x.in");
3893        assert!(err.contains("Unknown accessor: 'in'"), "{err}");
3894    }
3895
3896    #[test]
3897    fn test_new_aggregates() {
3898        let cols = parse_query("select nunique ID, var x, dev x by g")
3899            .unwrap()
3900            .cols;
3901        assert_eq!(cols[0], col("ID").n_unique().alias("nunique_ID"));
3902        assert_eq!(cols[1], col("x").var(1).alias("var_x"));
3903        assert_eq!(cols[2], col("x").std(1).alias("dev_x"));
3904
3905        let df = df!(
3906            "g" => &["a", "a", "a", "b"],
3907            "ID" => &["s1", "s1", "s2", "s3"],
3908            "x" => &[1.0f64, 2.0, 3.0, 5.0],
3909        )
3910        .unwrap();
3911        let out = eval("select nunique ID, var x, dev[x] by g", &df);
3912        assert_eq!(values(&out, "nunique_ID"), ["2", "1"]);
3913        assert_eq!(values(&out, "var_x"), ["1.0", "null"]);
3914        assert_eq!(values(&out, "dev_x"), ["1.0", "null"]);
3915    }
3916
3917    #[test]
3918    fn test_wavg() {
3919        let df = df!(
3920            "g" => &["a", "a", "a", "b"],
3921            "w" => &[Some(1i64), Some(3), Some(5), Some(2)],
3922            "x" => &[Some(10.0f64), Some(20.0), None, Some(4.0)],
3923        )
3924        .unwrap();
3925        let out = eval("select w wavg x by g", &df);
3926        // The null value's weight stays out of the total: (10 + 60) / 4.
3927        assert_eq!(values(&out, "wavg_x"), ["17.5", "4.0"]);
3928        // A group with no complete pair has no average.
3929        let out = eval("select w wavg x where null x", &df);
3930        assert_eq!(values(&out, "wavg_x"), ["null"]);
3931        for query in ["select wavg[x]", "select wavg x"] {
3932            let err = parse_err(query);
3933            assert!(
3934                err.contains("wavg goes between weights and values"),
3935                "{err}"
3936            );
3937        }
3938    }
3939
3940    #[test]
3941    fn test_round_and_math_functions() {
3942        let df = df!("x" => &[2.25f64, -2.5, 4.0]).unwrap();
3943        let out = eval(
3944            "select r: x.round, r1: x.round[1], s: sqrt x, l: log[x], e: exp 0 * x",
3945            &df,
3946        );
3947        assert_eq!(values(&out, "r"), ["2.0", "-3.0", "4.0"]);
3948        assert_eq!(values(&out, "r1"), ["2.3", "-2.5", "4.0"]);
3949        assert_eq!(values(&out, "s"), ["1.5", "NaN", "2.0"]);
3950        assert!(values(&out, "l")[2].starts_with("1.386"));
3951        assert_eq!(values(&out, "e"), ["1.0", "1.0", "1.0"]);
3952        // An aggregate rounds after it is computed.
3953        let df = df!("g" => &["a", "a"], "d" => &[1.0f64, 2.34]).unwrap();
3954        let out = eval("select m: (avg d).round[1] by g", &df);
3955        assert_eq!(values(&out, "m"), ["1.7"]);
3956    }
3957
3958    #[test]
3959    fn test_select_distinct() {
3960        let ParsedQuery { cols, distinct, .. } =
3961            parse_query("select distinct carrier, origin").unwrap();
3962        assert!(distinct);
3963        assert_eq!(cols, vec![col("carrier"), col("origin")]);
3964        let distinct = parse_query("select carrier").unwrap().distinct;
3965        assert!(!distinct);
3966
3967        let df = df!(
3968            "carrier" => &["UA", "UA", "AA", "UA"],
3969            "origin" => &["EWR", "EWR", "JFK", "LGA"],
3970            "n" => &[1i32, 2, 3, 4],
3971        )
3972        .unwrap();
3973        let out = eval("select distinct carrier, origin", &df);
3974        assert_eq!(values(&out, "carrier"), ["UA", "AA", "UA"]);
3975        assert_eq!(values(&out, "origin"), ["EWR", "JFK", "LGA"]);
3976        let out = eval("select distinct carrier where n > 1", &df);
3977        assert_eq!(values(&out, "carrier"), ["UA", "AA"]);
3978        // A column named distinct is col["distinct"].
3979        let ParsedQuery { cols, distinct, .. } = parse_query("select col[\"distinct\"]").unwrap();
3980        assert!(!distinct);
3981        assert_eq!(cols, vec![col("distinct")]);
3982        // So is a bare `distinct` that is plainly a column or an alias.
3983        let ParsedQuery { cols, distinct, .. } = parse_query("select distinct, n").unwrap();
3984        assert!(!distinct);
3985        assert_eq!(cols, vec![col("distinct"), col("n")]);
3986        let ParsedQuery { cols, distinct, .. } = parse_query("select distinct: n").unwrap();
3987        assert!(!distinct);
3988        assert_eq!(cols, vec![col("n").alias("distinct")]);
3989    }
3990
3991    #[test]
3992    fn test_from_df_is_optional() {
3993        // `from df` sits where q puts it: after the select list and by, before where.
3994        for (with, without) in [
3995            (
3996                "select mean dep_delay by hour from df where origin = \"JFK\"",
3997                "select mean dep_delay by hour where origin = \"JFK\"",
3998            ),
3999            ("select from df where x > 1", "select where x > 1"),
4000            ("select from df", "select"),
4001            ("select a, b from df", "select a, b"),
4002            ("select distinct a from df", "select distinct a"),
4003            ("select n: count a by g from df", "select n: count a by g"),
4004        ] {
4005            assert_eq!(
4006                format!("{:?}", parse_query(with).unwrap()),
4007                format!("{:?}", parse_query(without).unwrap()),
4008                "{with}"
4009            );
4010        }
4011    }
4012
4013    #[test]
4014    fn test_from_names_only_df() {
4015        for query in [
4016            "select from trades",
4017            "select a by g from trades where a > 1",
4018            "select from data.csv",
4019        ] {
4020            assert_eq!(
4021                parse_query(query).unwrap_err(),
4022                "q reads the table on screen, named df: … from df …",
4023                "{query}"
4024            );
4025        }
4026        let err = parse_query("select a where a > 1 from df").unwrap_err();
4027        assert!(err.contains("after the where clause"), "{err}");
4028        let err = parse_query("select a from df by g").unwrap_err();
4029        assert!(err.contains("'by' after 'from df'"), "{err}");
4030    }
4031
4032    #[test]
4033    fn test_from_column_names_and_values() {
4034        // A column named `from` still reads as one wherever a table name cannot follow.
4035        let cols = |q: &str| parse_query(q).unwrap().cols;
4036        assert_eq!(cols("select from"), vec![col("from")]);
4037        assert_eq!(cols("select from, to"), vec![col("from"), col("to")]);
4038        assert_eq!(cols("select from from df"), vec![col("from")]);
4039        assert_eq!(cols("select from + 1"), vec![col("from") + lit(1.0)]);
4040        assert_eq!(cols("select from.year"), cols("select col[\"from\"].year"));
4041        assert_eq!(cols("select max from"), cols("select max col[\"from\"]"));
4042        let ParsedQuery { group_by, .. } = parse_query("select n: count a by from").unwrap();
4043        assert_eq!(group_by, vec![col("from")]);
4044        assert_eq!(
4045            parse_query("select where from = \"df\"").unwrap().filter,
4046            Some(col("from").eq(lit("df")))
4047        );
4048        assert_eq!(
4049            parse_query("select where from in [1, 2]").unwrap().filter,
4050            parse_query("select where col[\"from\"] in [1, 2]")
4051                .unwrap()
4052                .filter
4053        );
4054        // Names that contain the word, and values that are it.
4055        assert_eq!(
4056            cols("select from_city, datefrom from df"),
4057            vec![col("from_city"), col("datefrom")]
4058        );
4059        assert_eq!(
4060            parse_query("select from df where city = \"from df\"")
4061                .unwrap()
4062                .filter,
4063            Some(col("city").eq(lit("from df")))
4064        );
4065        // A column named df is still a column.
4066        assert_eq!(cols("select df from df"), vec![col("df")]);
4067        // Run against data: from df changes nothing.
4068        let df = df!("from" => &[1i64, 2, 3], "df" => &["x", "y", "z"]).unwrap();
4069        let out = eval("select from, df from df where from > 1", &df);
4070        assert_eq!(values(&out, "df"), ["y", "z"]);
4071    }
4072
4073    #[test]
4074    fn test_accessor_argument_count_errors() {
4075        for (query, expected) in [
4076            (
4077                "select x.part[\",\"]",
4078                "part takes 2 arguments, e.g. .part[\"-\", 0]; got 1",
4079            ),
4080            ("select x.slice", "slice takes 1 to 2 arguments"),
4081            ("select x.replace[\"a\"]", "replace takes 2 arguments"),
4082            ("select x.round[1, 2]", "round takes 0 to 1 arguments"),
4083            (
4084                "select x.to_date[\"%Y\", \"%m\"]",
4085                "to_date takes 0 to 1 arguments",
4086            ),
4087            ("select x.hour[1]", "hour takes no arguments"),
4088            ("select x.strip[\" \"]", "strip takes no arguments"),
4089            ("select x.int[1]", "int takes no arguments"),
4090            ("select x.format", "format takes 1 argument"),
4091        ] {
4092            let err = parse_err(query);
4093            assert!(err.contains(expected), "{query}: {err}");
4094        }
4095    }
4096
4097    #[test]
4098    fn test_accessor_argument_type_errors() {
4099        for (query, expected) in [
4100            (
4101                "select x.part[0, \",\"]",
4102                "part: argument 1 must be quoted text",
4103            ),
4104            (
4105                "select x.part[\",\", \"a\"]",
4106                "part: argument 2 must be a whole number",
4107            ),
4108            (
4109                "select x.part[\",\", 1.5]",
4110                "part: argument 2 must be a whole number",
4111            ),
4112            ("select x.round[-1]", "round: decimals cannot be negative"),
4113            (
4114                "select x.slice[0, -1]",
4115                "slice: the length cannot be negative",
4116            ),
4117            ("select x.slice[a + 1]", "slice takes literal arguments"),
4118            ("select x.part[\",\"", "Unmatched bracket after .part"),
4119        ] {
4120            let err = parse_err(query);
4121            assert!(err.contains(expected), "{query}: {err}");
4122        }
4123    }
4124
4125    #[test]
4126    fn test_unknown_accessor_lists_new_names() {
4127        let err = parse_err("select x.nosuch");
4128        for name in [
4129            "hour",
4130            "minute",
4131            "second",
4132            "quarter",
4133            "doy",
4134            "to_date",
4135            "to_datetime",
4136            "part",
4137            "slice",
4138            "replace",
4139            "strip",
4140            "round",
4141            "int",
4142            "float",
4143            "str",
4144        ] {
4145            assert!(err.contains(name), "{name} missing from: {err}");
4146        }
4147    }
4148
4149    #[test]
4150    fn test_nested_functions_parse_in_linear_time() {
4151        // Each argument used to be parsed as an aggregate and then again as a scalar
4152        // function, doubling the work at every level of nesting.
4153        let bare = format!("select {}x", "abs ".repeat(40));
4154        assert!(parse_query(&bare).is_ok());
4155        let bracketed = format!("select {}x{}", "sqrt[".repeat(30), "]".repeat(30));
4156        assert!(parse_query(&bracketed).is_ok());
4157    }
4158
4159    /// One expression's Python code.
4160    fn py(expr: &str) -> String {
4161        parse_node(&tokenize(expr).unwrap()).unwrap().python()
4162    }
4163
4164    #[test]
4165    fn expressions_read_as_python_polars() {
4166        assert_eq!(py("a"), "pl.col(\"a\")");
4167        assert_eq!(py("col[\"first name\"]"), "pl.col(\"first name\")");
4168        // Python's `&` binds tighter than its comparisons, so operands that are
4169        // operations are parenthesized.
4170        assert_eq!(py("a > 1"), "pl.col(\"a\") > 1.0");
4171        assert_eq!(
4172            py("a + b * c"),
4173            "pl.col(\"a\") + (pl.col(\"b\") * pl.col(\"c\"))"
4174        );
4175        assert_eq!(py("-x"), "pl.lit(0) - pl.col(\"x\")");
4176        assert_eq!(py("x mod 3"), "pl.col(\"x\") % 3");
4177        assert_eq!(py("5 xbar fare"), "(pl.col(\"fare\") // 5) * 5");
4178        assert_eq!(py("a ^ 0"), "pl.coalesce(pl.col(\"a\"), pl.lit(0.0))");
4179        assert_eq!(
4180            py("name in [\"Emma\", \"Olivia\"]"),
4181            "(pl.col(\"name\") == \"Emma\") | (pl.col(\"name\") == \"Olivia\")"
4182        );
4183        assert_eq!(
4184            py("item like \"*Chicken*\""),
4185            "pl.col(\"item\").cast(pl.String).str.contains(\"(?s)^.*Chicken.*$\")"
4186        );
4187        assert_eq!(
4188            py("d = 2024.01.31"),
4189            "pl.col(\"d\") == pl.date(2024, 1, 31)"
4190        );
4191        assert_eq!(
4192            py("t > 2024.01.31T10:00:00.5"),
4193            "pl.col(\"t\") > pl.lit(\"2024-01-31T10:00:00.500\").str.to_datetime(\"%Y-%m-%dT%H:%M:%S%.3f\", time_unit=\"ms\")"
4194        );
4195    }
4196
4197    #[test]
4198    fn division_reads_as_polars_runs_it_on_the_types() {
4199        let schema = Schema::from_iter([
4200            Field::new("i".into(), DataType::Int64),
4201            Field::new("j".into(), DataType::Int32),
4202            Field::new("u".into(), DataType::UInt8),
4203            Field::new("f".into(), DataType::Float64),
4204            Field::new("s".into(), DataType::String),
4205        ]);
4206        let py = |expr: &str| {
4207            let mut node = parse_node(&tokenize(expr).unwrap()).unwrap();
4208            node.resolve_division(&schema);
4209            node.python()
4210        };
4211        // Two whole numbers floor-divide, as Polars' `/` on two expressions does.
4212        assert_eq!(py("i / j"), "pl.col(\"i\") // pl.col(\"j\")");
4213        assert_eq!(py("j % i"), "pl.col(\"j\") // pl.col(\"i\")");
4214        assert_eq!(py("i / u"), "pl.col(\"i\") // pl.col(\"u\")");
4215        // A pair Polars has no type for, which fails the query in datui too, and an
4216        // unknown column stay Python's `/`.
4217        assert_eq!(py("i / s"), "pl.col(\"i\") / pl.col(\"s\")");
4218        assert_eq!(py("(i mod 3) / j"), "(pl.col(\"i\") % 3) // pl.col(\"j\")");
4219        // A float on either side, or a number as typed, divides.
4220        assert_eq!(py("i / f"), "pl.col(\"i\") / pl.col(\"f\")");
4221        assert_eq!(py("i / 2"), "pl.col(\"i\") / 2.0");
4222        assert_eq!(
4223            py("sum[i] / count[j]"),
4224            "pl.col(\"i\").sum().alias(\"sum_i\") // pl.col(\"j\").count().alias(\"count_j\")"
4225        );
4226        assert_eq!(py("x / i"), "pl.col(\"x\") / pl.col(\"i\")");
4227    }
4228
4229    #[test]
4230    fn functions_and_accessors_read_as_python_polars() {
4231        assert_eq!(
4232            py("avg salary"),
4233            "pl.col(\"salary\").mean().alias(\"avg_salary\")"
4234        );
4235        assert_eq!(py("not null[x]"), "pl.col(\"x\").is_null().not_()");
4236        assert_eq!(py("log x"), "pl.col(\"x\").log()");
4237        assert_eq!(py("ts.year"), "pl.col(\"ts\").dt.year().alias(\"ts_year\")");
4238        assert_eq!(
4239            py("d.format[\"%Y-%m\"]"),
4240            "pl.col(\"d\").dt.to_string(\"%Y-%m\").alias(\"d_format_%Y-%m\")"
4241        );
4242        assert_eq!(
4243            py("code.part[\"-\", 0]"),
4244            "pl.col(\"code\").cast(pl.String).str.split(\"-\").list.get(0, null_on_oob=True).alias(\"code_part_-_0\")"
4245        );
4246        assert_eq!(
4247            py("s.slice[1]"),
4248            "pl.col(\"s\").cast(pl.String).str.slice(1).alias(\"s_slice_1\")"
4249        );
4250        assert_eq!(
4251            py("s.to_date[\"%Y%m%d\"]"),
4252            "pl.col(\"s\").cast(pl.String).str.to_date(\"%Y%m%d\", strict=False).alias(\"s_to_date_%Y%m%d\")"
4253        );
4254        assert_eq!(
4255            py("x.round[2]"),
4256            "pl.col(\"x\").round(2, mode=\"half_away_from_zero\").alias(\"x_round_2\")"
4257        );
4258        assert_eq!(
4259            py("x.int"),
4260            "pl.col(\"x\").cast(pl.Int64, strict=False).alias(\"x_int\")"
4261        );
4262        assert_eq!(
4263            py("w wavg v"),
4264            "((pl.col(\"w\") * pl.col(\"v\")).sum() / pl.when(pl.col(\"w\").filter((pl.col(\"w\") * pl.col(\"v\")).is_not_null()).sum() != 0).then(pl.col(\"w\").filter((pl.col(\"w\") * pl.col(\"v\")).is_not_null()).sum()).otherwise(pl.lit(None))).alias(\"wavg_v\")"
4265        );
4266    }
4267
4268    #[test]
4269    fn a_whole_query_reads_as_python_steps() {
4270        let steps = |q: &str, keys: &[&str]| {
4271            let keys: Vec<String> = keys.iter().map(|k| k.to_string()).collect();
4272            parse_nodes(q).unwrap().python_steps(&keys)
4273        };
4274        assert_eq!(
4275            steps(
4276                "select name, pay: salary * 1.1 where dept = \"Sales\", age > 30 | senior",
4277                &[]
4278            ),
4279            vec![
4280                ".filter((pl.col(\"dept\") == \"Sales\") & ((pl.col(\"age\") > 30.0) | pl.col(\"senior\")))",
4281                ".select(\"name\", (pl.col(\"salary\") * 1.1).alias(\"pay\"))",
4282            ]
4283        );
4284        assert_eq!(
4285            steps("select by dept", &["dept"]),
4286            vec![
4287                ".group_by(\"dept\")",
4288                ".agg(pl.all().exclude(\"dept\"))",
4289                ".sort(\"dept\", nulls_last=True, maintain_order=True)",
4290            ]
4291        );
4292        assert_eq!(
4293            steps("select distinct dept", &[]),
4294            vec![
4295                ".select(\"dept\")",
4296                ".unique(keep=\"first\", maintain_order=True)",
4297            ]
4298        );
4299        assert!(steps("", &[]).is_empty());
4300    }
4301}