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 DateLiteral(String),
12 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
32fn 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(); 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
83fn 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 chars.next(); 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(); 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 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(); 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 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 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
312fn 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
335const CLAUSE_ORDER: &str = "clause order is select [by group] [where conditions]";
337
338const FROM_ORDER: &str = "clause order is select [by group] [from df] [where conditions]";
340
341const TABLE: &str = "df";
343
344fn 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
366fn 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
410const WORD_OPS: [&str; 5] = ["in", "like", "xbar", "mod", "wavg"];
414
415fn 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#[derive(Debug, Clone, PartialEq)]
433pub(crate) enum Node {
434 Col(String),
435 Num(f64),
437 Int(i64),
439 Str(String),
440 Bool(bool),
441 Null,
442 Date(String),
444 Timestamp {
446 iso: String,
447 format: String,
448 unit: TimeUnit,
449 zone: Option<String>,
452 },
453 Bin(BinOp, Box<Node>, Box<Node>),
454 Coalesce(Box<Node>, Box<Node>),
455 Filter(Box<Node>, Box<Node>),
457 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 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#[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 ContainsRegex(String),
529 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 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
558const MAX_EXPR_NODES: usize = 10_000;
562
563fn 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 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 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 let zone = TimeZone::opt_try_new(zone.as_deref()).ok().flatten();
637 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 Op::Round(decimals) => expr.round(*decimals, RoundMode::HalfAwayFromZero),
1081 Op::Cast(to) => expr.cast(to.dtype()),
1083 }
1084}
1085
1086fn 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
1100fn 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
1113fn 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
1157pub(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
1166fn 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
1174fn 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
1184fn 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(®ex::escape(c.encode_utf8(&mut [0; 4]))),
1193 }
1194 }
1195 re.push('$');
1196 re
1197}
1198
1199fn 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 [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 if left.size() > 1 {
1224 check_copies(&left, items.len())?;
1225 }
1226 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 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 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 check_copies(&weights, 5)?;
1272 check_copies(&values, 3)?;
1273 let weighted = weights.clone().bin(BinOp::Mul, values);
1274 let total = Node::Filter(
1277 Box::new(weights),
1278 Box::new(weighted.clone().op(Op::IsNotNull)),
1279 )
1280 .op(Op::Sum);
1281 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 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 "%" | "/" => 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
1322fn 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
1339const AGG_FUNCTIONS: [&str; 16] = [
1341 "avg", "mean", "min", "max", "count", "std", "stddev", "dev", "var", "med", "median", "sum",
1342 "first", "last", "nunique", "wavg",
1343];
1344
1345const 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
1355fn is_function_name(name: &str) -> bool {
1357 let name = name.to_lowercase();
1359 name != "wavg" && (is_agg_function(&name) || SCALAR_FUNCTIONS.contains(&name.as_str()))
1360}
1361
1362fn 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
1373fn 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 match simple_column_name(args) {
1405 Some(column) => Ok(node.alias(format!("{}_{}", fn_name, column))),
1406 None => Ok(node),
1407 }
1408}
1409
1410fn 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#[derive(Debug, Clone, PartialEq)]
1439enum AccessorArg {
1440 Str(String),
1441 Num(f64),
1442}
1443
1444impl AccessorArg {
1445 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
1454const ACCESSORS: &[(&str, usize, usize, &str)] = &[
1456 ("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 ("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 ("round", 0, 1, ".round[1]"),
1489 ("int", 0, 0, ".int"),
1490 ("float", 0, 0, ".float"),
1491 ("str", 0, 0, ".str"),
1492];
1493
1494const 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
1508fn 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 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 _ => 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
1612fn 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
1631fn 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 if name == "col" && tokens.len() > 1 && tokens[1] == Token::LBracket {
1687 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 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 else if tokens.len() > 1 && tokens[1] == Token::LBracket {
1717 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 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..])), Token::String(s) => Ok((Node::Str(s.clone()), &tokens[1..])), 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 _ => Err(format!(
1778 "Unexpected '{}' where an expression was expected",
1779 token_text(&tokens[0])
1780 )),
1781 }
1782}
1783
1784const MAX_EXPR_DEPTH: u32 = 64;
1800
1801thread_local! {
1802 static EXPR_DEPTH: std::cell::Cell<u32> = const { std::cell::Cell::new(0) };
1803}
1804
1805struct DepthGuard;
1812
1813impl DepthGuard {
1814 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
1833fn 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 if let Token::Identifier(name) = &tokens[0]
1849 && is_function_name(name)
1850 && tokens.len() > 1
1851 && tokens[1] != Token::LBracket
1852 {
1853 return parse_call(name, &tokens[1..]);
1857 }
1858
1859 let mut op_pos = None;
1861 let mut depth = 0;
1862 let mut bracket_depth = 0;
1863
1864 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 let left_tokens = &tokens[..pos];
1882 let right_tokens = &tokens[pos + 1..];
1883
1884 if let Some(op) = infix_op_at(tokens, pos) {
1885 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 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 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 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#[derive(Debug, Default)]
1943pub struct ParsedQuery {
1944 pub cols: Vec<Expr>,
1946 pub filter: Option<Expr>,
1948 pub group_by: Vec<Expr>,
1950 pub group_by_names: Vec<String>,
1952 pub distinct: bool,
1954}
1955
1956impl ParsedQuery {
1957 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
1972pub 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#[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 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 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 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 pub(crate) fn python_filter(&self) -> Option<String> {
2056 self.filter
2057 .as_ref()
2058 .map(|f| format!(".filter({})", f.python()))
2059 }
2060
2061 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
2096fn 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
2113pub 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
2125pub(crate) fn parse_nodes(query: &str) -> Result<QueryNodes, String> {
2127 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 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 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 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 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 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 let alias_tokens = &chunk[..pos];
2219 let expr_tokens = &chunk[pos + 1..];
2220
2221 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 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 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 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 let alias_tokens = &chunk[..pos];
2281 let expr_tokens = &chunk[pos + 1..];
2282
2283 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 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); } else {
2315 let expr = parse_node(&chunk)?;
2316 group_by_cols.push(expr.clone());
2317 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 match &chunk[2] {
2331 Token::String(name) | Token::Identifier(name) => {
2332 group_by_col_names.push(name.clone());
2333 }
2334 _ => {}
2335 }
2336 } else {
2337 }
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#[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 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 }
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 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 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 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 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 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 let query = "select t, v where c>c%n";
2742
2743 let filter = parse_query(query).unwrap().filter;
2744
2745 let expected = col("c").gt(col("c").div(col("n")));
2747 assert_eq!(filter, Some(expected));
2748 }
2749
2750 #[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 }
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 }
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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 #[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 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 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 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 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 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 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 assert!(
3495 parse_query_over("select where d like \"2024*\"", Some(df.schema().as_ref())).is_ok()
3496 );
3497 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 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 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 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 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 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 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 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 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 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 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 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 assert_eq!(parse("a = b mod 2"), col("a").eq(col("b").rem(lit(2i64))));
3848 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 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 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 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 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 assert_eq!(values(&out, "wavg_x"), ["17.5", "4.0"]);
3928 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 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 let ParsedQuery { cols, distinct, .. } = parse_query("select col[\"distinct\"]").unwrap();
3980 assert!(!distinct);
3981 assert_eq!(cols, vec![col("distinct")]);
3982 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 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 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 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 assert_eq!(cols("select df from df"), vec![col("df")]);
4067 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 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 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 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 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 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 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}