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