use polars::prelude::StrptimeOptions;
use polars::prelude::*;
use std::ops::{Add, Div, Mul, Rem, Sub};
#[derive(Debug, Clone, PartialEq)]
enum Token {
Identifier(String),
Number(f64),
String(String),
DateLiteral(String),
TimestampLiteral {
iso: String,
format_str: String,
time_unit: TimeUnit,
},
Op(String),
LParen,
RParen,
LBracket,
RBracket,
Comma,
Colon,
Pipe,
Dot,
Select,
Where,
By,
}
fn parse_timestamp_literal(
date_part: &str,
chars: &mut std::iter::Peekable<std::str::Chars<'_>>,
) -> Option<(String, String, TimeUnit)> {
if chars.peek() != Some(&'T') {
return None;
}
chars.next(); let mut time_part = String::new();
while let Some(&c) = chars.peek() {
if c.is_ascii_digit() || c == ':' || c == '.' {
time_part.push(c);
chars.next();
} else {
break;
}
}
let parts: Vec<&str> = time_part.split(':').collect();
if parts.len() != 3 {
return None;
}
let (h, m, s) = (parts[0], parts[1], parts[2]);
if h.len() != 2 || m.len() != 2 || s.len() < 2 {
return None;
}
let (sec_part, frac) = match s.split_once('.') {
Some((a, f)) => (a, f),
None => (s, ""),
};
let (time_unit, format_str) = match frac.len() {
0 => (TimeUnit::Microseconds, "%Y-%m-%dT%H:%M:%S".to_string()),
1..=3 => (TimeUnit::Milliseconds, "%Y-%m-%dT%H:%M:%S%.3f".to_string()),
4..=6 => (TimeUnit::Microseconds, "%Y-%m-%dT%H:%M:%S%.6f".to_string()),
7..=9 => (TimeUnit::Nanoseconds, "%Y-%m-%dT%H:%M:%S%.9f".to_string()),
_ => (TimeUnit::Nanoseconds, "%Y-%m-%dT%H:%M:%S%.9f".to_string()),
};
let iso_date = parse_date_literal(date_part)?;
let frac_padded = match time_unit {
TimeUnit::Milliseconds => format!("{:0<3}", frac),
TimeUnit::Microseconds => format!("{:0<6}", frac),
TimeUnit::Nanoseconds => format!("{:0<9}", frac),
};
let iso = if frac.is_empty() {
format!("{}T{}:{}:{}", iso_date, h, m, sec_part)
} else {
format!("{}T{}:{}:{}.{}", iso_date, h, m, sec_part, frac_padded)
};
Some((iso, format_str, time_unit))
}
fn parse_date_literal(s: &str) -> Option<String> {
let parts: Vec<&str> = s.split('.').collect();
if parts.len() != 3 {
return None;
}
let year: u32 = parts[0].parse().ok()?;
let month: u32 = parts[1].parse().ok()?;
let day: u32 = parts[2].parse().ok()?;
if parts[0].len() != 4 || !(1000..=9999).contains(&year) {
return None;
}
if !(1..=12).contains(&month) || !(1..=31).contains(&day) {
return None;
}
Some(format!("{:04}-{:02}-{:02}", year, month, day))
}
fn tokenize(input: &str) -> Result<Vec<Token>, String> {
let mut tokens = Vec::new();
let mut chars = input.chars().peekable();
while let Some(&c) = chars.peek() {
match c {
' ' | '\t' | '\n' | '\r' => {
chars.next();
}
',' => {
tokens.push(Token::Comma);
chars.next();
}
':' => {
tokens.push(Token::Colon);
chars.next();
}
'|' => {
tokens.push(Token::Pipe);
chars.next();
}
'(' => {
tokens.push(Token::LParen);
chars.next();
}
')' => {
tokens.push(Token::RParen);
chars.next();
}
'[' => {
tokens.push(Token::LBracket);
chars.next();
}
']' => {
tokens.push(Token::RBracket);
chars.next();
}
'"' => {
chars.next(); let mut string_val = String::new();
let mut found_closing_quote = false;
while let Some(&c) = chars.peek() {
if c == '\\' {
chars.next(); if let Some(&next_c) = chars.peek() {
match next_c {
'n' => {
string_val.push('\n');
chars.next();
}
't' => {
string_val.push('\t');
chars.next();
}
'r' => {
string_val.push('\r');
chars.next();
}
'\\' => {
string_val.push('\\');
chars.next();
}
'"' => {
string_val.push('"');
chars.next();
}
_ => {
string_val.push('\\');
string_val.push(next_c);
chars.next();
}
}
} else {
return Err("Unterminated escape sequence in string".to_string());
}
} else if c == '"' {
chars.next(); found_closing_quote = true;
break;
} else {
string_val.push(c);
chars.next();
}
}
if !found_closing_quote {
return Err("Unterminated string literal".to_string());
}
tokens.push(Token::String(string_val));
}
'^' => {
tokens.push(Token::Op("^".to_string()));
chars.next();
}
'+' | '-' | '*' | '%' | '/' | '=' | '<' | '>' | '!' => {
let mut op = c.to_string();
chars.next();
if let Some(&next_c) = chars.peek()
&& ((c == '<' && (next_c == '=' || next_c == '>'))
|| (c == '>' && next_c == '=')
|| (c == '!' && next_c == '='))
{
op.push(next_c);
chars.next();
}
tokens.push(Token::Op(op));
}
'.' => {
chars.next();
if chars.peek().is_some_and(|nc| nc.is_ascii_digit()) {
let mut num_str = String::from('.');
while let Some(&nc) = chars.peek() {
if nc.is_ascii_digit() {
num_str.push(nc);
chars.next();
} else {
break;
}
}
if let Ok(n) = num_str.parse::<f64>() {
tokens.push(Token::Number(n));
} else {
return Err(format!("Invalid number: {}", num_str));
}
} else {
tokens.push(Token::Dot);
}
}
'0'..='9' => {
let mut num_str = String::new();
while let Some(&nc) = chars.peek() {
if nc.is_ascii_digit() || nc == '.' {
num_str.push(nc);
chars.next();
} else {
break;
}
}
let is_timestamp =
parse_date_literal(&num_str).is_some() && chars.peek() == Some(&'T');
if is_timestamp
&& let Some((iso, format_str, time_unit)) =
parse_timestamp_literal(&num_str, &mut chars)
{
tokens.push(Token::TimestampLiteral {
iso,
format_str,
time_unit,
});
continue;
}
if let Some(iso) = parse_date_literal(&num_str) {
tokens.push(Token::DateLiteral(iso));
} else if let Ok(n) = num_str.parse::<f64>() {
tokens.push(Token::Number(n));
} else {
return Err(format!("Invalid number: {}", num_str));
}
}
_ if c.is_alphabetic() || c == '_' => {
let mut ident = String::new();
while let Some(&nc) = chars.peek() {
if nc.is_alphanumeric() || nc == '_' {
ident.push(nc);
chars.next();
} else {
break;
}
}
match ident.as_str() {
"select" => tokens.push(Token::Select),
"where" => tokens.push(Token::Where),
"by" => tokens.push(Token::By),
_ => tokens.push(Token::Identifier(ident)),
}
}
_ => return Err(format!("Unexpected character: {}", c)),
}
}
Ok(tokens)
}
fn split_tokens(tokens: &[Token], delimiter: &Token) -> Vec<Vec<Token>> {
let mut result = Vec::new();
let mut current = Vec::new();
let mut depth = 0;
let mut bracket_depth = 0;
for token in tokens {
match token {
Token::LParen => depth += 1,
Token::RParen => depth -= 1,
Token::LBracket => bracket_depth += 1,
Token::RBracket => bracket_depth -= 1,
_ => {}
}
if depth == 0 && bracket_depth == 0 && token == delimiter {
result.push(current);
current = Vec::new();
} else {
current.push(token.clone());
}
}
result.push(current);
result
}
fn token_text(token: &Token) -> String {
match token {
Token::Identifier(s) => s.clone(),
Token::Number(n) => n.to_string(),
Token::String(s) => format!("\"{}\"", s),
Token::DateLiteral(iso) => iso.clone(),
Token::TimestampLiteral { iso, .. } => iso.clone(),
Token::Op(op) => op.clone(),
Token::LParen => "(".to_string(),
Token::RParen => ")".to_string(),
Token::LBracket => "[".to_string(),
Token::RBracket => "]".to_string(),
Token::Comma => ",".to_string(),
Token::Colon => ":".to_string(),
Token::Pipe => "|".to_string(),
Token::Dot => ".".to_string(),
Token::Select => "select".to_string(),
Token::Where => "where".to_string(),
Token::By => "by".to_string(),
}
}
const CLAUSE_ORDER: &str = "clause order is select [by group] [where conditions]";
const FROM_ORDER: &str = "clause order is select [by group] [from df] [where conditions]";
const TABLE: &str = "df";
fn from_table_at(tokens: &[Token], i: usize) -> Option<(String, usize)> {
if tokens.get(i) != Some(&Token::Identifier("from".to_string())) {
return None;
}
let mut name = match tokens.get(i + 1) {
Some(Token::Identifier(n)) if !WORD_OPS.contains(&n.as_str()) => n.clone(),
_ => return None,
};
let mut end = i + 2;
while let (Some(Token::Dot), Some(Token::Identifier(part))) =
(tokens.get(end), tokens.get(end + 1))
{
name.push('.');
name.push_str(part);
end += 2;
}
Some((name, end))
}
fn strip_from(body: &[Token]) -> Result<Vec<Token>, String> {
let mut depth = 0i32;
let mut found: Option<(usize, usize)> = None;
for (i, token) in body.iter().enumerate() {
match token {
Token::LParen | Token::LBracket => depth += 1,
Token::RParen | Token::RBracket => depth -= 1,
_ => {}
}
if depth != 0 {
continue;
}
let Some((name, end)) = from_table_at(body, i) else {
continue;
};
let after_where = body[..i].contains(&Token::Where);
match body.get(end) {
None | Some(Token::Where) | Some(Token::By) => {}
_ => continue,
}
if name != TABLE {
return Err("q reads the table on screen, named df: … from df …".to_string());
}
if after_where {
return Err(format!(
"Unexpected 'from df' after the where clause: {FROM_ORDER}"
));
}
if body.get(end) == Some(&Token::By) {
return Err(format!("Unexpected 'by' after 'from df': {FROM_ORDER}"));
}
found = Some((i, end));
}
let mut body = body.to_vec();
if let Some((start, end)) = found {
body.drain(start..end);
}
Ok(body)
}
const WORD_OPS: [&str; 5] = ["in", "like", "xbar", "mod", "wavg"];
fn infix_op_at(tokens: &[Token], i: usize) -> Option<&str> {
match tokens.get(i)? {
Token::Op(op) => Some(op.as_str()),
Token::Identifier(word)
if i > 0 && tokens[i - 1] != Token::Dot && WORD_OPS.contains(&word.as_str()) =>
{
Some(word.as_str())
}
_ => None,
}
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) enum Node {
Col(String),
Num(f64),
Int(i64),
Str(String),
Bool(bool),
Null,
Date(String),
Timestamp {
iso: String,
format: String,
unit: TimeUnit,
zone: Option<String>,
},
Bin(BinOp, Box<Node>, Box<Node>),
Coalesce(Box<Node>, Box<Node>),
Filter(Box<Node>, Box<Node>),
When(Box<Node>, Box<Node>, Box<Node>),
Op(Box<Node>, Op),
Alias(Box<Node>, String),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum BinOp {
Add,
Sub,
Mul,
Div,
TrueDiv,
FloorDiv,
Rem,
Eq,
Neq,
Lt,
Gt,
LtEq,
GtEq,
And,
Or,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) enum Op {
Mean,
Min,
Max,
Count,
Std,
Var,
Median,
Sum,
First,
Last,
NUnique,
Not,
IsNull,
IsNotNull,
LenChars,
Upper,
Lower,
Abs,
Floor,
Ceil,
Sqrt,
Ln,
Exp,
Date,
Time,
Year,
Quarter,
Month,
Week,
Day,
OrdinalDay,
Weekday,
Hour,
Minute,
Second,
MonthStart,
MonthEnd,
DtFormat(String),
StartsWith(String),
EndsWith(String),
ContainsLiteral(String),
ContainsRegex(String),
Part(String, i64),
Slice(i64, Option<u64>),
ReplaceAll(String, String),
Strip,
ToDate(Option<String>),
ToDatetime(Option<String>),
Round(u32),
Cast(CastTo),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum CastTo {
Int64,
Float64,
String,
}
impl CastTo {
fn dtype(self) -> DataType {
match self {
CastTo::Int64 => DataType::Int64,
CastTo::Float64 => DataType::Float64,
CastTo::String => DataType::String,
}
}
}
const MAX_EXPR_NODES: usize = 10_000;
fn check_copies(node: &Node, copies: usize) -> Result<(), String> {
if node.size().saturating_mul(copies) > MAX_EXPR_NODES {
return Err(
"Expression is too large: nested wavg, xbar or in repeat what they are \
given. Simplify it or split it into steps."
.to_string(),
);
}
Ok(())
}
impl Node {
fn size(&self) -> usize {
1 + match self {
Node::Col(_)
| Node::Num(_)
| Node::Int(_)
| Node::Str(_)
| Node::Bool(_)
| Node::Null
| Node::Date(_)
| Node::Timestamp { .. } => 0,
Node::Bin(_, a, b) | Node::Coalesce(a, b) | Node::Filter(a, b) => a.size() + b.size(),
Node::When(a, b, c) => a.size() + b.size() + c.size(),
Node::Op(a, _) | Node::Alias(a, _) => a.size(),
}
}
fn op(self, op: Op) -> Node {
Node::Op(Box::new(self), op)
}
fn bin(self, op: BinOp, right: Node) -> Node {
Node::Bin(op, Box::new(self), Box::new(right))
}
fn alias(self, name: impl Into<String>) -> Node {
Node::Alias(Box::new(self), name.into())
}
fn cast_text(self) -> Node {
self.op(Op::Cast(CastTo::String))
}
pub(crate) fn to_expr(&self) -> Expr {
match self {
Node::Col(name) => col(name),
Node::Num(n) => lit(*n),
Node::Int(n) => lit(*n),
Node::Str(s) => lit(s.as_str()),
Node::Bool(b) => lit(*b),
Node::Null => lit(NULL),
Node::Date(iso) => {
let opts = StrptimeOptions {
format: Some("%Y-%m-%d".into()),
..Default::default()
};
lit(iso.as_str()).str().to_date(opts)
}
Node::Timestamp {
iso,
format,
unit,
zone,
} => {
let opts = StrptimeOptions {
format: Some(format.as_str().into()),
..Default::default()
};
let zone = TimeZone::opt_try_new(zone.as_deref()).ok().flatten();
lit(iso.as_str())
.str()
.to_datetime(Some(*unit), zone, opts, lit("earliest"))
}
Node::Bin(op, left, right) => {
let (left, right) = (left.to_expr(), right.to_expr());
match op {
BinOp::Add => left.add(right),
BinOp::Sub => left.sub(right),
BinOp::Mul => left.mul(right),
BinOp::Div => left.div(right),
BinOp::TrueDiv => left.true_div(right),
BinOp::FloorDiv => left.floor_div(right),
BinOp::Rem => left.rem(right),
BinOp::Eq => left.eq(right),
BinOp::Neq => left.neq(right),
BinOp::Lt => left.lt(right),
BinOp::Gt => left.gt(right),
BinOp::LtEq => left.lt_eq(right),
BinOp::GtEq => left.gt_eq(right),
BinOp::And => left.and(right),
BinOp::Or => left.or(right),
}
}
Node::Coalesce(left, right) => coalesce(&[left.to_expr(), right.to_expr()]),
Node::Filter(values, predicate) => values.to_expr().filter(predicate.to_expr()),
Node::When(condition, then, otherwise) => when(condition.to_expr())
.then(then.to_expr())
.otherwise(otherwise.to_expr()),
Node::Op(inner, op) => apply_op_expr(inner.to_expr(), op),
Node::Alias(inner, name) => inner.to_expr().alias(name.as_str()),
}
}
pub(crate) fn without_aliases(&self) -> Node {
let strip = |n: &Node| Box::new(n.without_aliases());
match self {
Node::Alias(inner, _) => inner.without_aliases(),
Node::Bin(op, l, r) => Node::Bin(*op, strip(l), strip(r)),
Node::Coalesce(l, r) => Node::Coalesce(strip(l), strip(r)),
Node::Filter(v, p) => Node::Filter(strip(v), strip(p)),
Node::When(c, t, o) => Node::When(strip(c), strip(t), strip(o)),
Node::Op(inner, op) => Node::Op(strip(inner), op.clone()),
leaf => leaf.clone(),
}
}
pub(crate) fn resolve_division(&mut self, schema: &Schema) {
match self {
Node::Bin(_, left, right) | Node::Coalesce(left, right) | Node::Filter(left, right) => {
left.resolve_division(schema);
right.resolve_division(schema);
}
Node::When(c, t, o) => {
c.resolve_division(schema);
t.resolve_division(schema);
o.resolve_division(schema);
}
Node::Op(inner, _) | Node::Alias(inner, _) => inner.resolve_division(schema),
_ => {}
}
if let Node::Bin(BinOp::Div, ..) = self {
let quotient = DataFrame::empty_with_schema(schema)
.lazy()
.select([self.to_expr()])
.collect_schema()
.ok()
.and_then(|s| s.get_at_index(0).map(|(_, dtype)| dtype.is_integer()));
if let (Some(whole), Node::Bin(op, ..)) = (quotient, self) {
*op = if whole {
BinOp::FloorDiv
} else {
BinOp::TrueDiv
};
}
}
}
pub(crate) fn resolve_time_zones(&mut self, schema: &Schema) {
match self {
Node::Bin(_, left, right) | Node::Coalesce(left, right) => {
left.resolve_time_zones(schema);
right.resolve_time_zones(schema);
Self::share_zone(left, right, schema);
}
Node::Filter(values, predicate) => {
values.resolve_time_zones(schema);
predicate.resolve_time_zones(schema);
}
Node::When(c, t, o) => {
c.resolve_time_zones(schema);
t.resolve_time_zones(schema);
o.resolve_time_zones(schema);
Self::share_zone(t, o, schema);
}
Node::Op(inner, _) | Node::Alias(inner, _) => inner.resolve_time_zones(schema),
_ => {}
}
}
fn check_quoted_temporal(&self, schema: &Schema) -> Result<(), String> {
match self {
Node::Bin(op, left, right) => {
left.check_quoted_temporal(schema)?;
right.check_quoted_temporal(schema)?;
let compares = matches!(
op,
BinOp::Eq | BinOp::Neq | BinOp::Lt | BinOp::Gt | BinOp::LtEq | BinOp::GtEq
);
if compares
&& let Some(err) = quoted_temporal(left, right, schema)
.or_else(|| quoted_temporal(right, left, schema))
{
return Err(err);
}
}
Node::Coalesce(left, right) | Node::Filter(left, right) => {
left.check_quoted_temporal(schema)?;
right.check_quoted_temporal(schema)?;
}
Node::When(c, t, o) => {
c.check_quoted_temporal(schema)?;
t.check_quoted_temporal(schema)?;
o.check_quoted_temporal(schema)?;
}
Node::Op(inner, _) | Node::Alias(inner, _) => inner.check_quoted_temporal(schema)?,
_ => {}
}
Ok(())
}
fn share_zone(a: &mut Node, b: &mut Node, schema: &Schema) {
if !Self::take_zone(a, b, schema) {
Self::take_zone(b, a, schema);
}
}
fn take_zone(literal: &mut Node, other: &Node, schema: &Schema) -> bool {
if let Node::Timestamp {
zone: zone @ None, ..
} = literal
&& let Some(DataType::Datetime(_, Some(tz))) = other.dtype(schema)
{
*zone = Some(tz.to_string());
return true;
}
false
}
fn dtype(&self, schema: &Schema) -> Option<DataType> {
DataFrame::empty_with_schema(schema)
.lazy()
.select([self.to_expr()])
.collect_schema()
.ok()
.and_then(|s| s.get_at_index(0).map(|(_, dtype)| dtype.clone()))
}
fn python_literal(&self) -> Option<String> {
Some(match self {
Node::Num(n) => crate::python_script::py_float(*n),
Node::Int(n) => n.to_string(),
Node::Str(s) => crate::python_script::py_str(s),
Node::Bool(b) => crate::python_script::py_bool(*b).to_string(),
Node::Null => "None".to_string(),
_ => return None,
})
}
pub(crate) fn python(&self) -> String {
use crate::python_script::py_str;
if let Some(literal) = self.python_literal() {
return format!("pl.lit({literal})");
}
match self {
Node::Col(name) => format!("pl.col({})", py_str(name)),
Node::Date(iso) => {
let parts: Vec<String> = iso
.split('-')
.map(|p| p.trim_start_matches('0').to_string())
.map(|p| if p.is_empty() { "0".to_string() } else { p })
.collect();
format!("pl.date({})", parts.join(", "))
}
Node::Timestamp {
iso,
format,
unit,
zone,
} => format!(
"pl.lit({}).str.to_datetime({}, time_unit={}{})",
py_str(iso),
py_str(format),
py_str(time_unit_name(*unit)),
zone.as_ref().map_or(String::new(), |zone| format!(
", time_zone={}, ambiguous=\"earliest\"",
py_str(zone)
))
),
Node::Bin(op, left, right) => {
let right = match right.python_literal() {
Some(literal) => literal,
None => right.python_operand(),
};
format!("{} {} {}", left.python_operand(), op.python(), right)
}
Node::Coalesce(left, right) => {
format!("pl.coalesce({}, {})", left.python(), right.python())
}
Node::Filter(values, predicate) => {
format!("{}.filter({})", values.python_operand(), predicate.python())
}
Node::When(condition, then, otherwise) => format!(
"pl.when({}).then({}).otherwise({})",
condition.python(),
then.python(),
otherwise.python()
),
Node::Op(inner, op) => format!("{}{}", inner.python_operand(), op.python()),
Node::Alias(inner, name) => {
let mut inner = inner.as_ref();
while let Node::Alias(deeper, _) = inner {
inner = deeper;
}
format!("{}.alias({})", inner.python_operand(), py_str(name))
}
_ => unreachable!("literals return above"),
}
}
fn python_operand(&self) -> String {
match self {
Node::Bin(..) => format!("({})", self.python()),
_ => self.python(),
}
}
}
fn time_unit_name(unit: TimeUnit) -> &'static str {
match unit {
TimeUnit::Milliseconds => "ms",
TimeUnit::Microseconds => "us",
TimeUnit::Nanoseconds => "ns",
}
}
impl BinOp {
fn python(self) -> &'static str {
match self {
BinOp::Add => "+",
BinOp::Sub => "-",
BinOp::Mul => "*",
BinOp::Div | BinOp::TrueDiv => "/",
BinOp::FloorDiv => "//",
BinOp::Rem => "%",
BinOp::Eq => "==",
BinOp::Neq => "!=",
BinOp::Lt => "<",
BinOp::Gt => ">",
BinOp::LtEq => "<=",
BinOp::GtEq => ">=",
BinOp::And => "&",
BinOp::Or => "|",
}
}
}
impl Op {
fn python(&self) -> String {
use crate::python_script::py_str;
let fixed = match self {
Op::Mean => ".mean()",
Op::Min => ".min()",
Op::Max => ".max()",
Op::Count => ".count()",
Op::Std => ".std()",
Op::Var => ".var()",
Op::Median => ".median()",
Op::Sum => ".sum()",
Op::First => ".first()",
Op::Last => ".last()",
Op::NUnique => ".n_unique()",
Op::Not => ".not_()",
Op::IsNull => ".is_null()",
Op::IsNotNull => ".is_not_null()",
Op::LenChars => ".str.len_chars()",
Op::Upper => ".str.to_uppercase()",
Op::Lower => ".str.to_lowercase()",
Op::Abs => ".abs()",
Op::Floor => ".floor()",
Op::Ceil => ".ceil()",
Op::Sqrt => ".sqrt()",
Op::Ln => ".log()",
Op::Exp => ".exp()",
Op::Date => ".dt.date()",
Op::Time => ".dt.time()",
Op::Year => ".dt.year()",
Op::Quarter => ".dt.quarter()",
Op::Month => ".dt.month()",
Op::Week => ".dt.week()",
Op::Day => ".dt.day()",
Op::OrdinalDay => ".dt.ordinal_day()",
Op::Weekday => ".dt.weekday()",
Op::Hour => ".dt.hour()",
Op::Minute => ".dt.minute()",
Op::Second => ".dt.second()",
Op::MonthStart => ".dt.month_start()",
Op::MonthEnd => ".dt.month_end()",
Op::Strip => ".str.strip_chars()",
Op::Cast(CastTo::Int64) => ".cast(pl.Int64, strict=False)",
Op::Cast(CastTo::Float64) => ".cast(pl.Float64, strict=False)",
Op::Cast(CastTo::String) => ".cast(pl.String)",
_ => "",
};
if !fixed.is_empty() {
return fixed.to_string();
}
let format_arg = |format: &Option<String>| match format {
Some(f) => format!("{}, strict=False", py_str(f)),
None => "strict=False".to_string(),
};
match self {
Op::DtFormat(f) => format!(".dt.to_string({})", py_str(f)),
Op::StartsWith(s) => format!(".str.starts_with({})", py_str(s)),
Op::EndsWith(s) => format!(".str.ends_with({})", py_str(s)),
Op::ContainsLiteral(s) => format!(".str.contains({}, literal=True)", py_str(s)),
Op::ContainsRegex(r) => format!(".str.contains({})", py_str(r)),
Op::Part(sep, i) => format!(
".str.split({}).list.get({i}, null_on_oob=True)",
py_str(sep)
),
Op::Slice(start, Some(len)) => format!(".str.slice({start}, {len})"),
Op::Slice(start, None) => format!(".str.slice({start})"),
Op::ReplaceAll(from, to) => format!(
".str.replace_all({}, {}, literal=True)",
py_str(from),
py_str(to)
),
Op::ToDate(format) => format!(".str.to_date({})", format_arg(format)),
Op::ToDatetime(format) => format!(".str.to_datetime({})", format_arg(format)),
Op::Round(d) => format!(".round({d}, mode=\"half_away_from_zero\")"),
_ => unreachable!("fixed calls return above"),
}
}
}
fn apply_op_expr(expr: Expr, op: &Op) -> Expr {
let strptime = |format: &Option<String>| StrptimeOptions {
format: format.as_deref().map(Into::into),
strict: false,
..Default::default()
};
match op {
Op::Mean => expr.mean(),
Op::Min => expr.min(),
Op::Max => expr.max(),
Op::Count => expr.count(),
Op::Std => expr.std(1),
Op::Var => expr.var(1),
Op::Median => expr.median(),
Op::Sum => expr.sum(),
Op::First => expr.first(),
Op::Last => expr.last(),
Op::NUnique => expr.n_unique(),
Op::Not => expr.not(),
Op::IsNull => expr.is_null(),
Op::IsNotNull => expr.is_not_null(),
Op::LenChars => expr.str().len_chars(),
Op::Upper => expr.str().to_uppercase(),
Op::Lower => expr.str().to_lowercase(),
Op::Abs => expr.abs(),
Op::Floor => expr.floor(),
Op::Ceil => expr.ceil(),
Op::Sqrt => expr.sqrt(),
Op::Ln => expr.log(lit(std::f64::consts::E)),
Op::Exp => expr.exp(),
Op::Date => expr.dt().date(),
Op::Time => expr.dt().time(),
Op::Year => expr.dt().year(),
Op::Quarter => expr.dt().quarter(),
Op::Month => expr.dt().month(),
Op::Week => expr.dt().week(),
Op::Day => expr.dt().day(),
Op::OrdinalDay => expr.dt().ordinal_day(),
Op::Weekday => expr.dt().weekday(),
Op::Hour => expr.dt().hour(),
Op::Minute => expr.dt().minute(),
Op::Second => expr.dt().second(),
Op::MonthStart => expr.dt().month_start(),
Op::MonthEnd => expr.dt().month_end(),
Op::DtFormat(f) => expr.dt().to_string(f),
Op::StartsWith(s) => expr.str().starts_with(lit(s.as_str())),
Op::EndsWith(s) => expr.str().ends_with(lit(s.as_str())),
Op::ContainsLiteral(s) => expr.str().contains_literal(lit(s.as_str())),
Op::ContainsRegex(r) => expr.str().contains(lit(r.as_str()), true),
Op::Part(sep, i) => expr
.str()
.split(lit(sep.as_str()))
.list()
.get(lit(*i), true),
Op::Slice(start, length) => {
let length = length.map_or_else(|| lit(NULL), lit);
expr.str().slice(lit(*start), length)
}
Op::ReplaceAll(from, to) => {
expr.str()
.replace_all(lit(from.as_str()), lit(to.as_str()), true)
}
Op::Strip => expr.str().strip_chars(lit(NULL)),
Op::ToDate(format) => expr.str().to_date(strptime(format)),
Op::ToDatetime(format) => {
expr.str()
.to_datetime(None, None, strptime(format), lit("raise"))
}
Op::Round(decimals) => expr.round(*decimals, RoundMode::HalfAwayFromZero),
Op::Cast(to) => expr.cast(to.dtype()),
}
}
fn int_or_node(tokens: &[Token]) -> Result<Node, String> {
let whole = |n: f64| n.fract() == 0.0 && n.abs() < i64::MAX as f64;
match tokens {
[Token::Number(n)] if whole(*n) => Ok(Node::Int(*n as i64)),
[Token::Op(minus), Token::Number(n)] if minus == "-" && whole(*n) => {
Ok(Node::Int(-(*n as i64)))
}
_ => parse_node(tokens),
}
}
fn brackets_balanced(tokens: &[Token]) -> bool {
let mut depth = 0usize;
tokens.iter().all(|t| match t {
Token::LBracket => {
depth += 1;
true
}
Token::RBracket => depth.checked_sub(1).map(|d| depth = d).is_some(),
_ => true,
})
}
fn quoted_temporal(column: &Node, text: &Node, schema: &Schema) -> Option<String> {
let (Node::Col(name), Node::Str(s)) = (column, text) else {
return None;
};
let shown = q_name(name);
let unquoted = tokenize(s).ok();
let literal = |is_kind: fn(&Token) -> bool, example: &str| match unquoted.as_deref() {
Some([token]) if is_kind(token) => s.trim().to_string(),
_ => example.to_string(),
};
let (kind, remedy) = match schema.get(name)? {
DataType::Date => (
"date",
format!(
"A date is {}",
literal(|t| matches!(t, Token::DateLiteral(_)), "2024.01.01")
),
),
DataType::Datetime(..) => (
"timestamp",
format!(
"A timestamp is {}",
literal(
|t| matches!(t, Token::TimestampLiteral { .. }),
"2024.01.01T05:00:00"
)
),
),
DataType::Time => (
"time",
format!(
"A time has no literal; compare {shown}.hour, {shown}.minute or {shown}.second with a number"
),
),
DataType::Duration(_) => ("duration", "A duration has no literal".to_string()),
_ => return None,
};
Some(format!(
"{shown} is a {kind}; \"{s}\" is a string. {remedy}"
))
}
pub(crate) fn q_name(name: &str) -> String {
if is_plain_name(name) {
name.to_string()
} else {
format!("col[\"{name}\"]")
}
}
fn is_plain_name(name: &str) -> bool {
let mut chars = name.chars();
chars.next().is_some_and(|c| c.is_alphabetic() || c == '_')
&& chars.all(|c| c.is_alphanumeric() || c == '_')
&& !matches!(name, "select" | "where" | "by")
}
fn any_of(mut conditions: Vec<Node>) -> Node {
if conditions.len() <= 1 {
return conditions.pop().unwrap_or(Node::Bool(false));
}
let right = conditions.split_off(conditions.len() / 2);
any_of(conditions).bin(BinOp::Or, any_of(right))
}
fn like_regex(pattern: &str) -> String {
let mut re = String::from("(?s)^");
for c in pattern.chars() {
match c {
'*' => re.push_str(".*"),
'?' => re.push('.'),
_ => re.push_str(®ex::escape(c.encode_utf8(&mut [0; 4]))),
}
}
re.push('$');
re
}
fn apply_infix(left_tokens: &[Token], op: &str, right_tokens: &[Token]) -> Result<Node, String> {
match op {
"in" => {
let list = match right_tokens {
[Token::LBracket, inner @ .., Token::RBracket] if brackets_balanced(inner) => inner,
_ => {
return Err(
"in takes a list on its right, e.g. name in [\"Emma\", \"Olivia\"]"
.to_string(),
);
}
};
let items = split_tokens(list, &Token::Comma);
if items.iter().any(|item| item.is_empty()) {
return Err(
"in needs a list of values, e.g. name in [\"Emma\", \"Olivia\"]".to_string(),
);
}
let left = parse_node(left_tokens)?;
if left.size() > 1 {
check_copies(&left, items.len())?;
}
let conditions = items
.iter()
.map(|item| Ok(left.clone().bin(BinOp::Eq, parse_node(item)?)))
.collect::<Result<Vec<_>, String>>()?;
Ok(any_of(conditions))
}
"like" => {
let [Token::String(pattern)] = right_tokens else {
return Err(
"like takes a quoted pattern on its right, e.g. item like \"*Chicken*\""
.to_string(),
);
};
let left = parse_node(left_tokens)?;
Ok(left.cast_text().op(Op::ContainsRegex(like_regex(pattern))))
}
"xbar" => {
if let [Token::Number(n)] = left_tokens
&& *n <= 0.0
{
return Err(
"xbar needs a positive bucket size, e.g. 5 xbar fare_amount".to_string()
);
}
let right = parse_node(right_tokens)?;
let size = int_or_node(left_tokens)?;
check_copies(&size, 2)?;
Ok(right
.bin(BinOp::FloorDiv, size.clone())
.bin(BinOp::Mul, size))
}
"mod" => {
let right = int_or_node(right_tokens)?;
let left = int_or_node(left_tokens)?;
Ok(left.bin(BinOp::Rem, right))
}
"wavg" => {
let values = parse_node(right_tokens)?;
let weights = parse_node(left_tokens)?;
check_copies(&weights, 5)?;
check_copies(&values, 3)?;
let weighted = weights.clone().bin(BinOp::Mul, values);
let total = Node::Filter(
Box::new(weights),
Box::new(weighted.clone().op(Op::IsNotNull)),
)
.op(Op::Sum);
let total = Node::When(
Box::new(total.clone().bin(BinOp::Neq, Node::Int(0))),
Box::new(total),
Box::new(Node::Null),
);
let node = weighted.op(Op::Sum).bin(BinOp::TrueDiv, total);
Ok(match simple_column_name(right_tokens) {
Some(column) => node.alias(format!("wavg_{}", column)),
None => node,
})
}
_ => {
let right = parse_node(right_tokens)?;
let left = parse_node(left_tokens)?;
apply_op(left, op, right)
}
}
}
fn apply_op(left: Node, op: &str, right: Node) -> Result<Node, String> {
let op = match op {
"+" => BinOp::Add,
"-" => BinOp::Sub,
"*" => BinOp::Mul,
"%" | "/" => BinOp::Div,
"^" => return Ok(Node::Coalesce(Box::new(left), Box::new(right))),
"=" => BinOp::Eq,
"<" => BinOp::Lt,
">" => BinOp::Gt,
"<=" => BinOp::LtEq,
">=" => BinOp::GtEq,
"<>" | "!=" => BinOp::Neq,
_ => return Err(format!("Unknown operator: {}", op)),
};
Ok(left.bin(op, right))
}
fn simple_column_name(tokens: &[Token]) -> Option<String> {
match tokens {
[Token::Identifier(name)] => Some(name.clone()),
[
Token::Identifier(c),
Token::LBracket,
Token::String(name) | Token::Identifier(name),
Token::RBracket,
] if c == "col" => Some(name.clone()),
_ => None,
}
}
const WAVG_USAGE: &str = "wavg goes between weights and values, e.g. passengers wavg fare";
const AGG_FUNCTIONS: [&str; 16] = [
"avg", "mean", "min", "max", "count", "std", "stddev", "dev", "var", "med", "median", "sum",
"first", "last", "nunique", "wavg",
];
const SCALAR_FUNCTIONS: [&str; 13] = [
"len", "length", "not", "null", "upper", "lower", "abs", "floor", "ceil", "ceiling", "sqrt",
"log", "exp",
];
fn is_agg_function(name: &str) -> bool {
AGG_FUNCTIONS.contains(&name.to_lowercase().as_str())
}
fn is_function_name(name: &str) -> bool {
let name = name.to_lowercase();
name != "wavg" && (is_agg_function(&name) || SCALAR_FUNCTIONS.contains(&name.as_str()))
}
fn parse_call(name: &str, args: &[Token]) -> Result<Node, String> {
if is_agg_function(name) {
parse_agg_function(name, args)
} else {
parse_function(name, args)
}
}
fn parse_agg_function(name: &str, args: &[Token]) -> Result<Node, String> {
if args.is_empty() {
return Err(format!(
"Aggregation function {} requires an argument",
name
));
}
let fn_name = name.to_lowercase();
if fn_name == "wavg" {
return Err(WAVG_USAGE.to_string());
}
let node = parse_node(args)?;
let op = match fn_name.as_str() {
"avg" | "mean" => Op::Mean,
"min" => Op::Min,
"max" => Op::Max,
"count" => Op::Count,
"std" | "stddev" | "dev" => Op::Std,
"var" => Op::Var,
"med" | "median" => Op::Median,
"sum" => Op::Sum,
"first" => Op::First,
"last" => Op::Last,
"nunique" => Op::NUnique,
_ => return Err(format!("Unknown aggregation function: {}", name)),
};
let node = node.op(op);
match simple_column_name(args) {
Some(column) => Ok(node.alias(format!("{}_{}", fn_name, column))),
None => Ok(node),
}
}
fn parse_function(name: &str, args: &[Token]) -> Result<Node, String> {
if args.is_empty() {
return Err(format!("Function {} requires an argument", name));
}
let name_lower = name.to_lowercase();
if !SCALAR_FUNCTIONS.contains(&name_lower.as_str()) {
return Err(format!("Unknown function: {}", name));
}
let node = parse_node(args)?;
let op = match name_lower.as_str() {
"not" => Op::Not,
"null" => Op::IsNull,
"len" | "length" => Op::LenChars,
"upper" => Op::Upper,
"lower" => Op::Lower,
"abs" => Op::Abs,
"floor" => Op::Floor,
"ceil" | "ceiling" => Op::Ceil,
"sqrt" => Op::Sqrt,
"log" => Op::Ln,
"exp" => Op::Exp,
_ => return Err(format!("Unknown function: {}", name)),
};
Ok(node.op(op))
}
#[derive(Debug, Clone, PartialEq)]
enum AccessorArg {
Str(String),
Num(f64),
}
impl AccessorArg {
fn alias_text(&self) -> String {
match self {
AccessorArg::Str(s) => s.clone(),
AccessorArg::Num(n) => n.to_string(),
}
}
}
const ACCESSORS: &[(&str, usize, usize, &str)] = &[
("date", 0, 0, ".date"),
("time", 0, 0, ".time"),
("year", 0, 0, ".year"),
("quarter", 0, 0, ".quarter"),
("month", 0, 0, ".month"),
("week", 0, 0, ".week"),
("day", 0, 0, ".day"),
("doy", 0, 0, ".doy"),
("dow", 0, 0, ".dow"),
("weekday", 0, 0, ".weekday"),
("hour", 0, 0, ".hour"),
("minute", 0, 0, ".minute"),
("second", 0, 0, ".second"),
("month_start", 0, 0, ".month_start"),
("month_end", 0, 0, ".month_end"),
("format", 1, 1, ".format[\"%Y-%m\"]"),
("len", 0, 0, ".len"),
("length", 0, 0, ".length"),
("upper", 0, 0, ".upper"),
("lower", 0, 0, ".lower"),
("starts_with", 1, 1, ".starts_with[\"x\"]"),
("ends_with", 1, 1, ".ends_with[\"x\"]"),
("contains", 1, 1, ".contains[\"x\"]"),
("part", 2, 2, ".part[\"-\", 0]"),
("slice", 1, 2, ".slice[0, 4]"),
("replace", 2, 2, ".replace[\"(P)\", \"\"]"),
("strip", 0, 0, ".strip"),
("to_date", 0, 1, ".to_date[\"%Y%m%d\"]"),
("to_datetime", 0, 1, ".to_datetime[\"%Y-%m-%d %H:%M\"]"),
("round", 0, 1, ".round[1]"),
("int", 0, 0, ".int"),
("float", 0, 0, ".float"),
("str", 0, 0, ".str"),
];
const ACCESSOR_HELP: &str = "Valid date/time: date, time, year, quarter, month, week, day, doy, dow, hour, minute, second, month_start, month_end, format. \
Valid string: len, upper, lower, starts_with, ends_with, contains, part, slice, replace, strip, to_date, to_datetime. \
Valid number: round, int, float, str";
fn arg_count_text(min: usize, max: usize) -> String {
match (min, max) {
(0, 0) => "no arguments".to_string(),
(1, 1) => "1 argument".to_string(),
(a, b) if a == b => format!("{} arguments", a),
(a, b) => format!("{} to {} arguments", a, b),
}
}
fn apply_accessor(node: Node, accessor: &str, args: &[AccessorArg]) -> Result<Node, String> {
let name = accessor.to_lowercase();
let Some(&(_, min, max, usage)) = ACCESSORS.iter().find(|(n, ..)| *n == name) else {
return Err(format!(
"Unknown accessor: '{}'. {}",
accessor, ACCESSOR_HELP
));
};
if args.len() < min || args.len() > max {
return Err(format!(
"{} takes {}, e.g. {}; got {}",
name,
arg_count_text(min, max),
usage,
args.len()
));
}
let text = |i: usize| match args.get(i) {
Some(AccessorArg::Str(s)) => Ok(s.clone()),
_ => Err(format!(
"{}: argument {} must be quoted text, e.g. {}",
name,
i + 1,
usage
)),
};
let int = |i: usize| match args.get(i) {
Some(AccessorArg::Num(n)) if n.fract() == 0.0 && n.abs() <= u32::MAX as f64 => {
Ok(*n as i64)
}
_ => Err(format!(
"{}: argument {} must be a whole number, e.g. {}",
name,
i + 1,
usage
)),
};
let as_str = || node.clone().cast_text();
Ok(match name.as_str() {
"date" => node.op(Op::Date),
"time" => node.op(Op::Time),
"year" => node.op(Op::Year),
"quarter" => node.op(Op::Quarter),
"month" => node.op(Op::Month),
"week" => node.op(Op::Week),
"day" => node.op(Op::Day),
"doy" => node.op(Op::OrdinalDay),
"dow" | "weekday" => node.op(Op::Weekday),
"hour" => node.op(Op::Hour),
"minute" => node.op(Op::Minute),
"second" => node.op(Op::Second),
"month_start" => node.op(Op::MonthStart),
"month_end" => node.op(Op::MonthEnd),
"format" => node.op(Op::DtFormat(text(0)?)),
"len" | "length" => node.op(Op::LenChars),
"upper" => node.op(Op::Upper),
"lower" => node.op(Op::Lower),
"starts_with" => node.op(Op::StartsWith(text(0)?)),
"ends_with" => node.op(Op::EndsWith(text(0)?)),
"contains" => node.op(Op::ContainsLiteral(text(0)?)),
"part" => as_str().op(Op::Part(text(0)?, int(1)?)),
"slice" => {
let start = int(0)?;
let length = match args.len() {
2 => {
let n = int(1)?;
if n < 0 {
return Err(format!(
"slice: the length cannot be negative, e.g. {}",
usage
));
}
Some(n as u64)
}
_ => None,
};
as_str().op(Op::Slice(start, length))
}
"replace" => as_str().op(Op::ReplaceAll(text(0)?, text(1)?)),
"strip" => as_str().op(Op::Strip),
"to_date" => as_str().op(Op::ToDate(args.first().map(|_| text(0)).transpose()?)),
"to_datetime" => as_str().op(Op::ToDatetime(args.first().map(|_| text(0)).transpose()?)),
"round" => {
let decimals = if args.is_empty() { 0 } else { int(0)? };
let decimals = u32::try_from(decimals)
.map_err(|_| format!("round: decimals cannot be negative, e.g. {}", usage))?;
node.op(Op::Round(decimals))
}
"int" => node.op(Op::Cast(CastTo::Int64)),
"float" => node.op(Op::Cast(CastTo::Float64)),
"str" => node.op(Op::Cast(CastTo::String)),
_ => {
return Err(format!(
"Unknown accessor: '{}'. {}",
accessor, ACCESSOR_HELP
));
}
})
}
fn parse_accessor_args(accessor: &str, tokens: &[Token]) -> Result<Vec<AccessorArg>, String> {
if tokens.is_empty() {
return Ok(Vec::new());
}
split_tokens(tokens, &Token::Comma)
.iter()
.map(|arg| match arg.as_slice() {
[Token::String(s)] | [Token::Identifier(s)] => Ok(AccessorArg::Str(s.clone())),
[Token::Number(n)] => Ok(AccessorArg::Num(*n)),
[Token::Op(minus), Token::Number(n)] if minus == "-" => Ok(AccessorArg::Num(-n)),
_ => Err(format!(
"{} takes literal arguments, quoted text or numbers, e.g. .part[\"-\", 0]",
accessor
)),
})
.collect()
}
fn parse_accessors<'a>(
mut expr: Node,
mut tokens: &'a [Token],
base_name: Option<&str>,
) -> Result<(Node, &'a [Token]), String> {
let mut alias_suffix = String::new();
while let [Token::Dot, Token::Identifier(accessor), rest @ ..] = tokens {
let (args, consumed) = if rest.first() == Some(&Token::LBracket) {
let mut depth = 0;
let close = rest
.iter()
.position(|t| {
match t {
Token::LBracket => depth += 1,
Token::RBracket => depth -= 1,
_ => {}
}
depth == 0
})
.ok_or_else(|| format!("Unmatched bracket after .{}", accessor))?;
(parse_accessor_args(accessor, &rest[1..close])?, close + 3)
} else {
(Vec::new(), 2)
};
expr = apply_accessor(expr, accessor, &args)?;
if !alias_suffix.is_empty() {
alias_suffix.push('_');
}
alias_suffix.push_str(accessor);
for arg in &args {
alias_suffix.push('_');
alias_suffix.push_str(&arg.alias_text());
}
tokens = &tokens[consumed..];
}
if !alias_suffix.is_empty() {
let alias = match base_name {
Some(name) => format!("{}_{}", name, alias_suffix),
None => alias_suffix,
};
expr = expr.alias(alias);
}
Ok((expr, tokens))
}
fn parse_term(tokens: &[Token]) -> Result<(Node, &[Token]), String> {
if tokens.is_empty() {
return Err("Unexpected end of expression".to_string());
}
match &tokens[0] {
Token::Identifier(name) => {
if name == "col" && tokens.len() > 1 && tokens[1] == Token::LBracket {
let mut depth = 1;
let mut i = 2;
while i < tokens.len() && depth > 0 {
match tokens[i] {
Token::LBracket => depth += 1,
Token::RBracket => depth -= 1,
_ => {}
}
i += 1;
}
if depth > 0 {
return Err("Unmatched bracket in col[]".to_string());
}
let col_name_tokens = &tokens[2..i - 1];
if col_name_tokens.len() != 1 {
return Err("col[] must contain a single string or identifier".to_string());
}
let col_name = match &col_name_tokens[0] {
Token::String(s) => s.clone(),
Token::Identifier(id) => id.clone(),
_ => return Err("col[] must contain a string or identifier".to_string()),
};
let expr = Node::Col(col_name.clone());
let (expr, remaining) = parse_accessors(expr, &tokens[i..], Some(&col_name))?;
Ok((expr, remaining))
}
else if tokens.len() > 1 && tokens[1] == Token::LBracket {
let mut depth = 1;
let mut i = 2;
while i < tokens.len() && depth > 0 {
match tokens[i] {
Token::LBracket => depth += 1,
Token::RBracket => depth -= 1,
_ => {}
}
i += 1;
}
if depth > 0 {
return Err("Unmatched bracket in function call".to_string());
}
let expr = parse_call(name, &tokens[2..i - 1])?;
parse_accessors(expr, &tokens[i..], None)
} else {
let expr = Node::Col(name.clone());
let (expr, remaining) = parse_accessors(expr, &tokens[1..], Some(name))?;
Ok((expr, remaining))
}
}
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..])),
Token::TimestampLiteral {
iso,
format_str,
time_unit,
} => Ok((
Node::Timestamp {
iso: iso.clone(),
format: format_str.clone(),
unit: *time_unit,
zone: None,
},
&tokens[1..],
)),
Token::LParen => {
let mut depth = 1;
let mut i = 1;
while i < tokens.len() && depth > 0 {
match tokens[i] {
Token::LParen => depth += 1,
Token::RParen => depth -= 1,
_ => {}
}
i += 1;
}
if depth > 0 {
return Err("Unmatched parenthesis".to_string());
}
let inner = parse_node(&tokens[1..i - 1])?;
let (expr, remaining) = parse_accessors(inner, &tokens[i..], None)?;
Ok((expr, remaining))
}
_ => Err(format!(
"Unexpected '{}' where an expression was expected",
token_text(&tokens[0])
)),
}
}
const MAX_EXPR_DEPTH: u32 = 64;
thread_local! {
static EXPR_DEPTH: std::cell::Cell<u32> = const { std::cell::Cell::new(0) };
}
struct DepthGuard;
impl DepthGuard {
fn enter() -> Option<Self> {
EXPR_DEPTH.with(|depth| {
let next = depth.get() + 1;
if next > MAX_EXPR_DEPTH {
return None;
}
depth.set(next);
Some(DepthGuard)
})
}
}
impl Drop for DepthGuard {
fn drop(&mut self) {
EXPR_DEPTH.with(|depth| depth.set(depth.get().saturating_sub(1)));
}
}
fn parse_node(tokens: &[Token]) -> Result<Node, String> {
let Some(_depth_guard) = DepthGuard::enter() else {
return Err(
"Expression is nested too deeply. Simplify it or split it into steps.".to_string(),
);
};
if tokens.is_empty() {
return Err("Empty expression".to_string());
}
if let Token::Identifier(name) = &tokens[0]
&& is_function_name(name)
&& tokens.len() > 1
&& tokens[1] != Token::LBracket
{
return parse_call(name, &tokens[1..]);
}
let mut op_pos = None;
let mut depth = 0;
let mut bracket_depth = 0;
for (i, token) in tokens.iter().enumerate() {
match token {
Token::LParen => depth += 1,
Token::RParen => depth -= 1,
Token::LBracket => bracket_depth += 1,
Token::RBracket => bracket_depth -= 1,
_ if depth == 0 && bracket_depth == 0 && infix_op_at(tokens, i).is_some() => {
op_pos = Some(i);
break;
}
_ => {}
}
}
if let Some(pos) = op_pos {
let left_tokens = &tokens[..pos];
let right_tokens = &tokens[pos + 1..];
if let Some(op) = infix_op_at(tokens, pos) {
if left_tokens.is_empty()
&& op == "-"
&& !right_tokens.is_empty()
&& matches!(right_tokens[0], Token::Number(_))
&& let Token::Number(n) = right_tokens[0]
{
if right_tokens.len() >= 3
&& let Some(bin_op) = infix_op_at(right_tokens, 1)
{
if WORD_OPS.contains(&bin_op) {
return apply_infix(&[Token::Number(-n)], bin_op, &right_tokens[2..]);
}
let right = parse_node(&right_tokens[2..])?;
return apply_op(Node::Int(0).bin(BinOp::Sub, Node::Num(n)), bin_op, right);
}
if right_tokens.len() == 1 {
return Ok(Node::Int(0).bin(BinOp::Sub, Node::Num(n)));
}
}
if left_tokens.is_empty() && (op == "+" || op == "-") {
let inner = parse_node(right_tokens)?;
return if op == "-" {
Ok(Node::Int(0).bin(BinOp::Sub, inner))
} else {
Ok(inner)
};
}
if left_tokens.is_empty() {
return Err("Missing left operand".to_string());
}
apply_infix(left_tokens, op, right_tokens)
} else {
Err("Expected operator".to_string())
}
} else {
let (expr, remaining) = parse_term(tokens)?;
if let Some(extra) = remaining.first() {
if matches!(&tokens[0], Token::Identifier(w) if w == "wavg") {
return Err(WAVG_USAGE.to_string());
}
return Err(format!(
"Unexpected '{}' after the expression",
token_text(extra)
));
}
Ok(expr)
}
}
#[derive(Debug, Default)]
pub struct ParsedQuery {
pub cols: Vec<Expr>,
pub filter: Option<Expr>,
pub group_by: Vec<Expr>,
pub group_by_names: Vec<String>,
pub distinct: bool,
}
impl ParsedQuery {
pub fn past_calendar_safe(self, schema: Option<&Schema>) -> Self {
let guard = |e: Expr| crate::past_calendar::guard_expr(e, schema);
Self {
cols: self.cols.into_iter().map(guard).collect(),
filter: self.filter.map(guard),
group_by: self.group_by.into_iter().map(guard).collect(),
..self
}
}
}
pub fn sanitize_query_error(msg: &str) -> String {
let msg_lower = msg.to_lowercase();
if msg_lower.contains("duplicate")
&& (msg_lower.contains("output name") || msg_lower.contains("projection"))
{
let name = msg
.split('\'')
.nth(1)
.map(|s| s.to_string())
.unwrap_or_else(|| "column".to_string());
return format!(
"Duplicate column name '{}' in result. Use aliases to rename columns, e.g. `select my_date: timestamp.date`",
name
);
}
if msg_lower.contains(".alias(") || msg_lower.contains("try renaming") {
return "Duplicate column names in result. Use aliases to rename columns, e.g. `select my_date: timestamp.date`"
.to_string();
}
msg.to_string()
}
#[derive(Debug, Default)]
pub(crate) struct QueryNodes {
pub cols: Vec<Node>,
pub filter: Option<Node>,
pub group_by: Vec<Node>,
pub group_by_names: Vec<String>,
pub distinct: bool,
}
impl QueryNodes {
fn into_parsed(self) -> ParsedQuery {
let lower = |nodes: Vec<Node>| nodes.iter().map(Node::to_expr).collect();
ParsedQuery {
cols: lower(self.cols),
filter: self.filter.as_ref().map(Node::to_expr),
group_by: lower(self.group_by),
group_by_names: self.group_by_names,
distinct: self.distinct,
}
}
pub(crate) fn resolve_division(&mut self, schema: &Schema) {
let nodes = self
.cols
.iter_mut()
.chain(self.filter.iter_mut())
.chain(self.group_by.iter_mut());
for node in nodes {
node.resolve_division(schema);
}
}
pub(crate) fn resolve_time_zones(&mut self, schema: &Schema) {
let nodes = self
.cols
.iter_mut()
.chain(self.filter.iter_mut())
.chain(self.group_by.iter_mut());
for node in nodes {
node.resolve_time_zones(schema);
}
}
fn check_quoted_temporal(&self, schema: &Schema) -> Result<(), String> {
self.cols
.iter()
.chain(self.filter.iter())
.chain(self.group_by.iter())
.try_for_each(|node| node.check_quoted_temporal(schema))
}
pub(crate) fn python_filter(&self) -> Option<String> {
self.filter
.as_ref()
.map(|f| format!(".filter({})", f.python()))
}
pub(crate) fn python_steps(&self, key_names: &[String]) -> Vec<String> {
let mut steps: Vec<String> = self.python_filter().into_iter().collect();
if !self.group_by.is_empty() {
let keys = python_list(&self.group_by);
let aggs = if !self.cols.is_empty() {
python_list(&self.cols)
} else if self.group_by_names.is_empty() {
"pl.all()".to_string()
} else {
let names: Vec<String> = self
.group_by_names
.iter()
.map(|n| crate::python_script::py_str(n))
.collect();
format!("pl.all().exclude({})", names.join(", "))
};
steps.push(format!(".group_by({keys})"));
steps.push(format!(".agg({aggs})"));
steps.push(crate::python_script::sort_call(
key_names,
&vec![false; key_names.len()],
));
} else if !self.cols.is_empty() {
steps.push(format!(".select({})", python_list(&self.cols)));
}
if self.distinct {
steps.push(".unique(keep=\"first\", maintain_order=True)".to_string());
}
steps
}
}
fn python_list(nodes: &[Node]) -> String {
nodes
.iter()
.map(|n| match n {
Node::Col(name) => crate::python_script::py_str(name),
n => n.python(),
})
.collect::<Vec<_>>()
.join(", ")
}
pub fn parse_query(query: &str) -> Result<ParsedQuery, String> {
parse_nodes(query).map(QueryNodes::into_parsed)
}
pub fn parse_query_over(query: &str, schema: Option<&Schema>) -> Result<ParsedQuery, String> {
let mut nodes = parse_nodes(query)?;
if let Some(schema) = schema {
nodes.resolve_time_zones(schema);
nodes.check_quoted_temporal(schema)?;
}
Ok(nodes.into_parsed())
}
pub(crate) fn parse_nodes(query: &str) -> Result<QueryNodes, String> {
let trimmed = query.trim();
if trimmed.is_empty() {
return Ok(QueryNodes::default());
}
let tokens = tokenize(query)?;
if tokens.is_empty() || tokens[0] != Token::Select {
return Err("Query must start with 'select'".to_string());
}
let distinct = tokens.get(1) == Some(&Token::Identifier("distinct".to_string()))
&& !matches!(
tokens.get(2),
Some(Token::Colon | Token::Comma | Token::Dot | Token::Op(_))
);
let body = strip_from(&tokens[if distinct { 2 } else { 1 }..])?;
let body = &body[..];
let mut parts = split_tokens(body, &Token::Where);
let select_by_tokens = parts.remove(0);
let where_tokens = if !parts.is_empty() {
Some(parts.remove(0))
} else {
None
};
if !parts.is_empty() {
return Err(
"Unexpected second 'where': combine conditions with ',' (and) or '|' (or)".to_string(),
);
}
if let Some(ref wt) = where_tokens {
let mut depth = 0;
let mut bracket_depth = 0;
for token in wt {
match token {
Token::LParen => depth += 1,
Token::RParen => depth -= 1,
Token::LBracket => bracket_depth += 1,
Token::RBracket => bracket_depth -= 1,
Token::By if depth == 0 && bracket_depth == 0 => {
return Err(format!(
"Unexpected 'by' after the where clause: {}",
CLAUSE_ORDER
));
}
_ => {}
}
}
}
let mut select_by_parts = split_tokens(&select_by_tokens, &Token::By);
let cols_tokens = select_by_parts.remove(0);
let by_tokens = if !select_by_parts.is_empty() {
Some(select_by_parts.remove(0))
} else {
None
};
if !select_by_parts.is_empty() {
return Err(format!("Unexpected second 'by': {}", CLAUSE_ORDER));
}
let mut cols = Vec::new();
if !cols_tokens.is_empty() {
for chunk in split_tokens(&cols_tokens, &Token::Comma) {
if chunk.is_empty() {
continue;
}
let mut colon_pos = None;
let mut depth = 0;
for (i, token) in chunk.iter().enumerate() {
match token {
Token::LBracket => depth += 1,
Token::RBracket => depth -= 1,
Token::Colon if depth == 0 => {
colon_pos = Some(i);
break;
}
_ => {}
}
}
if let Some(pos) = colon_pos {
let alias_tokens = &chunk[..pos];
let expr_tokens = &chunk[pos + 1..];
let alias_name = if alias_tokens.len() == 1 {
if let Token::Identifier(name) = &alias_tokens[0] {
name.clone()
} else {
return Err("Expected identifier or col[] for alias".to_string());
}
} else if alias_tokens.len() == 4
&& alias_tokens[0] == Token::Identifier("col".to_string())
&& alias_tokens[1] == Token::LBracket
&& alias_tokens[3] == Token::RBracket
{
match &alias_tokens[2] {
Token::String(name) | Token::Identifier(name) => name.clone(),
_ => {
return Err(
"Expected string or identifier in col[] for alias".to_string()
);
}
}
} else {
return Err("Alias must be an identifier or col[]".to_string());
};
let expr = parse_node(expr_tokens)?;
cols.push(expr.alias(alias_name));
} else {
cols.push(parse_node(&chunk)?);
}
}
}
let mut group_by_cols = Vec::new();
let mut group_by_col_names = Vec::new();
if let Some(bt) = by_tokens {
for chunk in split_tokens(&bt, &Token::Comma) {
if chunk.is_empty() {
continue;
}
let mut colon_pos = None;
let mut depth = 0;
for (i, token) in chunk.iter().enumerate() {
match token {
Token::LBracket => depth += 1,
Token::RBracket => depth -= 1,
Token::Colon if depth == 0 => {
colon_pos = Some(i);
break;
}
_ => {}
}
}
if let Some(pos) = colon_pos {
let alias_tokens = &chunk[..pos];
let expr_tokens = &chunk[pos + 1..];
let alias_name = if alias_tokens.len() == 1 {
if let Token::Identifier(name) = &alias_tokens[0] {
name.clone()
} else {
return Err(
"Expected identifier or col[] for alias in by clause".to_string()
);
}
} else if alias_tokens.len() == 4
&& alias_tokens[0] == Token::Identifier("col".to_string())
&& alias_tokens[1] == Token::LBracket
&& alias_tokens[3] == Token::RBracket
{
match &alias_tokens[2] {
Token::String(name) | Token::Identifier(name) => name.clone(),
_ => {
return Err(
"Expected string or identifier in col[] for alias in by clause"
.to_string(),
);
}
}
} else {
return Err("Alias must be an identifier or col[] in by clause".to_string());
};
let expr = parse_node(expr_tokens)?;
group_by_cols.push(expr.alias(alias_name.clone()));
group_by_col_names.push(alias_name); } else {
let expr = parse_node(&chunk)?;
group_by_cols.push(expr.clone());
if chunk.len() == 1 {
if let Token::Identifier(name) = &chunk[0] {
group_by_col_names.push(name.clone());
}
} else if chunk.len() == 4
&& chunk[0] == Token::Identifier("col".to_string())
&& chunk[1] == Token::LBracket
&& chunk[3] == Token::RBracket
{
match &chunk[2] {
Token::String(name) | Token::Identifier(name) => {
group_by_col_names.push(name.clone());
}
_ => {}
}
} else {
}
}
}
}
let mut filter: Option<Node> = None;
if let Some(wt) = where_tokens {
for chunk in split_tokens(&wt, &Token::Comma) {
if chunk.is_empty() {
continue;
}
let mut or_expr: Option<Node> = None;
for or_chunk in split_tokens(&chunk, &Token::Pipe) {
if or_chunk.is_empty() {
continue;
}
let e = parse_node(&or_chunk)?;
or_expr = match or_expr {
Some(curr) => Some(curr.bin(BinOp::Or, e)),
None => Some(e),
};
}
if let Some(e) = or_expr {
filter = match filter {
Some(curr) => Some(curr.bin(BinOp::And, e)),
None => Some(e),
};
}
}
}
Ok(QueryNodes {
cols,
filter,
group_by: group_by_cols,
group_by_names: group_by_col_names,
distinct,
})
}
#[cfg(test)]
fn parse_expr(tokens: &[Token]) -> Result<Expr, String> {
parse_node(tokens).map(|n| n.to_expr())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tokenize_simple() {
let query = "select a, b where a > 10";
let tokens = tokenize(query).unwrap();
assert_eq!(
tokens,
vec![
Token::Select,
Token::Identifier("a".to_string()),
Token::Comma,
Token::Identifier("b".to_string()),
Token::Where,
Token::Identifier("a".to_string()),
Token::Op(">".to_string()),
Token::Number(10.0),
]
);
}
#[test]
fn test_tokenize_operators() {
let query = "a != b, c >= d, e <= f, g <> h";
let tokens = tokenize(query).unwrap();
assert_eq!(
tokens,
vec![
Token::Identifier("a".to_string()),
Token::Op("!=".to_string()),
Token::Identifier("b".to_string()),
Token::Comma,
Token::Identifier("c".to_string()),
Token::Op(">=".to_string()),
Token::Identifier("d".to_string()),
Token::Comma,
Token::Identifier("e".to_string()),
Token::Op("<=".to_string()),
Token::Identifier("f".to_string()),
Token::Comma,
Token::Identifier("g".to_string()),
Token::Op("<>".to_string()),
Token::Identifier("h".to_string()),
]
);
}
#[test]
fn test_parse_simple_expr() {
let tokens = tokenize("a + 1").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(expr, col("a").add(lit(1.0)));
}
#[test]
fn test_parse_complex_expr() {
let tokens = tokenize("(a + 1) * 2").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(expr, (col("a").add(lit(1.0))).mul(lit(2.0)));
}
#[test]
fn test_parse_not_function() {
let query = "select a where not[a = b]";
let filter = parse_query(query).unwrap().filter;
assert_eq!(filter, Some(col("a").eq(col("b")).not()));
}
#[test]
fn test_parse_not_equivalent_to_neq() {
let query1 = "select a where a != b";
let query2 = "select a where not[a = b]";
let query3 = "select a where not a = b";
let filter1 = parse_query(query1).unwrap().filter;
let filter2 = parse_query(query2).unwrap().filter;
let filter3 = parse_query(query3).unwrap().filter;
assert_eq!(filter1, Some(col("a").neq(col("b"))));
assert_eq!(filter2, Some(col("a").eq(col("b")).not()));
assert_eq!(filter3, Some(col("a").eq(col("b")).not()));
}
#[test]
fn test_parse_avg_without_brackets() {
let query = "select avg 5+a by category";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols.len(), 1);
}
#[test]
fn test_parse_string_literal() {
let query = "select a, b:\"foo\"";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols.len(), 2);
assert_eq!(cols[0], col("a"));
assert_eq!(cols[1], lit("foo").alias("b"));
}
#[test]
fn test_parse_string_in_where() {
let query = "select a where name=\"george\", age > 7";
let filter = parse_query(query).unwrap().filter;
assert!(filter.is_some());
}
#[test]
fn test_parse_col_syntax() {
let query = "select col[\"first name\"]";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols.len(), 1);
assert_eq!(cols[0], col("first name"));
}
#[test]
fn test_parse_col_syntax_with_alias() {
let query = "select a, b:col[\"first name\"]";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols.len(), 2);
assert_eq!(cols[0], col("a"));
assert_eq!(cols[1], col("first name").alias("b"));
}
#[test]
fn test_parse_col_syntax_with_string_literal() {
let query = "select col[\"first name\"]:\"derek\", foo where foo > 7";
let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
assert_eq!(cols.len(), 2);
assert_eq!(cols[0], lit("derek").alias("first name"));
assert_eq!(cols[1], col("foo"));
assert!(filter.is_some());
}
#[test]
fn test_parse_string_escape_sequences() {
let query = "select a where name=\"george\\\"s name\"";
let filter = parse_query(query).unwrap().filter;
assert!(filter.is_some());
}
#[test]
fn test_parse_query_simple_where() {
let query = "select a where a > 10";
let filter = parse_query(query).unwrap().filter;
assert_eq!(filter, Some(col("a").gt(lit(10.0))));
}
#[test]
fn test_parse_query_unary_minus_in_where() {
let query = "select sum total-1 by product where 0<-0.5+discount";
let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
assert_eq!(cols.len(), 1);
assert!(filter.is_some());
let expected = lit(0.0).lt(lit(0).sub(lit(0.5)).add(col("discount")));
assert_eq!(filter, Some(expected));
}
#[test]
fn test_parse_query_negative_literal_where() {
let query = "select where 0<-0.1+discount";
let filter = parse_query(query).unwrap().filter;
let expected = lit(0.0).lt(lit(0).sub(lit(0.1)).add(col("discount")));
assert_eq!(filter, Some(expected));
}
#[test]
fn test_parse_unary_plus_minus_expr() {
let tokens = tokenize("-0.5").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(expr, lit(0).sub(lit(0.5)));
let tokens = tokenize("+x").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(expr, col("x"));
}
#[test]
fn test_parse_query_alias() {
let query = "select my_col:a + 1";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols, vec![col("a").add(lit(1.0)).alias("my_col")]);
}
#[test]
fn test_parse_query_and_or() {
let query = "select a where a > 10 | a < 5, b = 2";
let filter = parse_query(query).unwrap().filter;
let expected =
(col("a").gt(lit(10.0)).or(col("a").lt(lit(5.0)))).and(col("b").eq(lit(2.0)));
assert_eq!(filter, Some(expected));
}
#[test]
fn test_parse_query_neq() {
let query = "select a where a != 10";
let filter = parse_query(query).unwrap().filter;
assert_eq!(filter, Some(col("a").neq(lit(10.0))));
}
#[test]
fn test_parse_query_gte() {
let query = "select a where a >= 10";
let filter = parse_query(query).unwrap().filter;
assert_eq!(filter, Some(col("a").gt_eq(lit(10.0))));
}
#[test]
fn test_parse_query_lte() {
let query = "select a where a <= 10";
let filter = parse_query(query).unwrap().filter;
assert_eq!(filter, Some(col("a").lt_eq(lit(10.0))));
}
#[test]
fn test_empty_query() {
let query = "select";
let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
assert!(cols.is_empty());
assert!(filter.is_none());
}
#[test]
fn test_select_all_implicit() {
let query = "select where a > 1";
let ParsedQuery { cols, filter, .. } = parse_query(query).unwrap();
assert!(cols.is_empty());
assert_eq!(filter, Some(col("a").gt(lit(1.0))));
}
#[test]
fn test_invalid_query_no_select() {
let query = "a > 10";
let result = parse_query(query);
assert!(result.is_err());
}
#[test]
fn test_invalid_query_unmatched_paren() {
let query = "select (a + 1";
let result = parse_query(query);
assert!(result.is_err());
}
#[test]
fn test_invalid_query_bad_token() {
let query = "select a where a ? 10";
let result = parse_query(query);
assert!(result.is_err());
}
#[test]
fn test_parse_right_to_left_operator_precedence() {
let query = "select t, v where c>c%n";
let filter = parse_query(query).unwrap().filter;
let expected = col("c").gt(col("c").div(col("n")));
assert_eq!(filter, Some(expected));
}
#[test]
fn test_tokenize_dot_accessor() {
let tokens = tokenize("foo.date").unwrap();
assert_eq!(
tokens,
vec![
Token::Identifier("foo".to_string()),
Token::Dot,
Token::Identifier("date".to_string()),
]
);
}
#[test]
fn test_tokenize_decimal_number() {
let tokens = tokenize(".5").unwrap();
assert_eq!(tokens, vec![Token::Number(0.5)]);
}
#[test]
fn test_parse_simple_date_accessor() {
let tokens = tokenize("timestamp.date").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(expr, col("timestamp").dt().date().alias("timestamp_date"));
}
#[test]
fn test_parse_col_with_date_accessor() {
let tokens = tokenize("col[\"Created At\"].year").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(expr, col("Created At").dt().year().alias("Created At_year"));
}
#[test]
fn test_parse_chained_accessors() {
let tokens = tokenize("dt_col.date.year").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert_eq!(
expr,
col("dt_col")
.dt()
.date()
.dt()
.year()
.alias("dt_col_date_year")
);
}
#[test]
fn test_parse_query_select_with_date_accessor() {
let query = "select event_date: timestamp.date";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols.len(), 1);
assert_eq!(
cols[0],
col("timestamp")
.dt()
.date()
.alias("timestamp_date")
.alias("event_date")
);
}
#[test]
fn test_parse_query_select_col_with_accessor() {
let query = "select col[\"Event Time\"].date, col[\"Event Time\"].year";
let cols = parse_query(query).unwrap().cols;
assert_eq!(cols.len(), 2);
assert_eq!(
cols[0],
col("Event Time").dt().date().alias("Event Time_date")
);
assert_eq!(
cols[1],
col("Event Time").dt().year().alias("Event Time_year")
);
}
#[test]
fn test_parse_query_where_with_date_accessor() {
let query = "select where created_at.month = 12";
let filter = parse_query(query).unwrap().filter;
assert_eq!(
filter,
Some(
col("created_at")
.dt()
.month()
.alias("created_at_month")
.eq(lit(12.0))
)
);
}
#[test]
fn test_parse_query_where_dow() {
let query = "select where event_ts.dow = 1";
let filter = parse_query(query).unwrap().filter;
assert_eq!(
filter,
Some(
col("event_ts")
.dt()
.weekday()
.alias("event_ts_dow")
.eq(lit(1.0))
)
);
}
#[test]
fn test_parse_all_accessors() {
let accessors = [
"date",
"time",
"year",
"month",
"week",
"day",
"dow",
"month_start",
"month_end",
];
for accessor in accessors {
let query = format!("select x.{}", accessor);
let result = parse_query(&query);
assert!(
result.is_ok(),
"Accessor '{}' should parse: {:?}",
accessor,
result.err()
);
}
}
#[test]
fn test_parse_unknown_accessor() {
let query = "select x.nosuchaccessor";
let result = parse_query(query);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(err.contains("Unknown accessor"));
assert!(err.contains("nosuchaccessor"));
}
#[test]
fn test_parse_date_literal() {
let tokens = tokenize("2021.01.01").unwrap();
assert_eq!(tokens, vec![Token::DateLiteral("2021-01-01".to_string())]);
}
#[test]
fn test_parse_query_where_date_literal() {
let query = "select where dt_col.date > 2021.01.01";
let filter = parse_query(query).unwrap().filter;
assert!(filter.is_some());
}
#[test]
fn test_number_not_parsed_as_date() {
let tokens = tokenize("2.5").unwrap();
assert_eq!(tokens, vec![Token::Number(2.5)]);
}
#[test]
fn test_sanitize_duplicate_column_error() {
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.";
let sanitized = sanitize_query_error(polars_msg);
assert!(sanitized.contains("Duplicate column name"));
assert!(sanitized.contains("timestamp"));
assert!(sanitized.contains("my_date: timestamp.date"));
assert!(!sanitized.contains(".alias("));
}
#[test]
fn test_parse_timestamp_literal() {
let tokens = tokenize("2021.01.15T14:30:00.123456").unwrap();
assert!(matches!(tokens[0], Token::TimestampLiteral { .. }));
}
#[test]
fn test_parse_null_and_not_null() {
let f1 = parse_query("select where null col1").unwrap().filter;
assert!(f1.is_some());
let f2 = parse_query("select where not null col1").unwrap().filter;
assert!(f2.is_some());
}
#[test]
fn test_parse_coalesce() {
let cols = parse_query("select a: coln^cola^colb").unwrap().cols;
assert_eq!(cols.len(), 1);
}
#[test]
fn test_parse_first_last_aggregation() {
let cols = parse_query("select first[value], last[value] by group")
.unwrap()
.cols;
assert_eq!(cols.len(), 2);
}
#[test]
fn test_parse_string_accessors() {
let filter = parse_query("select where city_name.ends_with[\"lanta\"]")
.unwrap()
.filter;
assert!(filter.is_some());
let cols = parse_query("select name.len, name.upper").unwrap().cols;
assert_eq!(cols.len(), 2);
}
#[test]
fn test_parse_format_accessor() {
let tokens = tokenize("dt_col.format[\"%Y-%m\"]").unwrap();
let expr = parse_expr(&tokens).unwrap();
assert!(!format!("{:?}", expr).is_empty());
}
#[test]
fn test_parse_by_with_date_accessor() {
let query = "select order_date, count: count id by order_date.year";
let ParsedQuery {
cols,
group_by: group_by_cols,
..
} = parse_query(query).unwrap();
assert_eq!(cols.len(), 2);
assert_eq!(group_by_cols.len(), 1);
assert_eq!(
group_by_cols[0],
col("order_date").dt().year().alias("order_date_year")
);
}
#[test]
fn test_unaliased_aggregates_of_same_column_coexist() {
let query = "select avg salary, max salary by department";
let ParsedQuery {
cols,
group_by: group_by_cols,
..
} = parse_query(query).unwrap();
assert_eq!(cols.len(), 2);
assert_eq!(cols[0], col("salary").mean().alias("avg_salary"));
assert_eq!(cols[1], col("salary").max().alias("max_salary"));
assert_eq!(group_by_cols, vec![col("department")]);
}
#[test]
fn test_unaliased_aggregate_bracketed_and_bare_name_alike() {
let bracketed = parse_query("select avg[salary] by department")
.unwrap()
.cols;
let bare = parse_query("select avg salary by department").unwrap().cols;
assert_eq!(bracketed, bare);
assert_eq!(bracketed[0], col("salary").mean().alias("avg_salary"));
}
#[test]
fn test_unaliased_aggregate_col_syntax_auto_alias() {
let cols = parse_query("select sum[col[\"unit price\"]] by region")
.unwrap()
.cols;
assert_eq!(cols[0], col("unit price").sum().alias("sum_unit price"));
}
#[test]
fn test_bare_count_names_itself() {
let cols = parse_query("select count[x] by g").unwrap().cols;
assert_eq!(cols[0], col("x").count().alias("count_x"));
}
#[test]
fn test_explicit_alias_overrides_aggregate_auto_alias() {
let cols = parse_query("select total:sum[price] by region")
.unwrap()
.cols;
assert_eq!(
cols[0],
col("price").sum().alias("sum_price").alias("total")
);
}
#[test]
fn test_aggregate_of_expression_keeps_default_name() {
let cols = parse_query("select sum[price*qty] by region").unwrap().cols;
assert_eq!(cols[0], (col("price").mul(col("qty"))).sum());
}
#[test]
fn test_docs_grouping_example_collects_with_auto_aliases() {
let query = "select avg salary, max salary, count name by department";
let ParsedQuery {
cols,
group_by: group_by_cols,
..
} = parse_query(query).unwrap();
let df = df!(
"department" => &["eng", "eng", "ops"],
"salary" => &[100.0f64, 200.0, 300.0],
"name" => &["a", "b", "c"],
)
.unwrap();
let out = df
.lazy()
.group_by(group_by_cols)
.agg(cols)
.collect()
.unwrap();
let names: Vec<String> = out
.get_column_names()
.iter()
.map(|n| n.to_string())
.collect();
assert_eq!(
names,
["department", "avg_salary", "max_salary", "count_name"]
);
}
#[test]
fn test_slash_divides_like_percent() {
let slash = parse_expr(&tokenize("a/b").unwrap()).unwrap();
let percent = parse_expr(&tokenize("a%b").unwrap()).unwrap();
assert_eq!(slash, percent);
assert_eq!(slash, col("a").div(col("b")));
}
#[test]
fn test_slash_right_to_left() {
let expr = parse_expr(&tokenize("1/c+a").unwrap()).unwrap();
assert_eq!(expr, lit(1.0).div(col("c").add(col("a"))));
}
#[test]
fn test_slash_in_where_clause() {
let filter = parse_query("select t, v where c>c/n").unwrap().filter;
assert_eq!(filter, Some(col("c").gt(col("c").div(col("n")))));
}
#[test]
fn test_by_after_where_errors_with_clause_order() {
let err = parse_query("select name, salary where x > 1 by dept").unwrap_err();
assert!(
err.contains("Unexpected 'by' after the where clause"),
"{err}"
);
assert!(
err.contains("select [by group] [where conditions]"),
"{err}"
);
}
#[test]
fn test_by_after_where_without_condition_operator() {
let err = parse_query("select where flag by dept").unwrap_err();
assert!(
err.contains("Unexpected 'by' after the where clause"),
"{err}"
);
}
#[test]
fn test_by_inside_parens_in_where_errors_as_stray_token() {
let err = parse_query("select a where (x by g)").unwrap_err();
assert!(
err.contains("Unexpected 'by' after the expression"),
"{err}"
);
}
#[test]
fn test_trailing_garbage_after_where_errors() {
let err = parse_query("select a where a > 1 2").unwrap_err();
assert!(err.contains("Unexpected '2' after the expression"), "{err}");
let err = parse_query("select a where null col1 foo").unwrap_err();
assert!(
err.contains("Unexpected 'foo' after the expression"),
"{err}"
);
}
#[test]
fn test_trailing_garbage_in_select_errors() {
let err = parse_query("select a b").unwrap_err();
assert!(err.contains("Unexpected 'b' after the expression"), "{err}");
let err = parse_query("select (a, b)").unwrap_err();
assert!(err.contains("Unexpected ',' after the expression"), "{err}");
}
#[test]
fn test_duplicate_clauses_error() {
let err = parse_query("select a where x > 1 where y > 2").unwrap_err();
assert!(err.contains("Unexpected second 'where'"), "{err}");
assert!(err.contains("','"), "{err}");
let err = parse_query("select a by g by h").unwrap_err();
assert!(err.contains("Unexpected second 'by'"), "{err}");
}
#[test]
fn test_operators_that_repeat_an_operand_are_bounded() {
let chain = format!("select {}x", "w wavg ".repeat(30));
let err = parse_query(&chain).unwrap_err();
assert!(err.contains("Expression is too large"), "{err}");
let xbar = format!("select {}x{}", "(1 xbar ".repeat(40), ")".repeat(40));
assert!(parse_query(&xbar).is_err());
let inner = format!(
"select {}x{}",
"(".repeat(20),
" in [1, 2, 3, 4])".repeat(20)
);
assert!(parse_query(&inner).is_err());
assert!(parse_query("select w wavg x wavg y by g").is_ok());
}
#[test]
fn test_deeply_nested_expression_is_rejected_not_crashed() {
let unary = format!("select {}x", "-".repeat(5_000));
assert!(
parse_query(&unary).is_err(),
"deep unary chain should error"
);
let parens = format!("select {}x{}", "(".repeat(5_000), ")".repeat(5_000));
assert!(parse_query(&parens).is_err(), "deep nesting should error");
assert!(
parse_query("select a + b * c").is_ok(),
"an ordinary query must still parse after a rejected one"
);
}
fn eval(query: &str, df: &DataFrame) -> DataFrame {
let ParsedQuery {
cols,
filter,
group_by: by,
distinct,
..
} = parse_query_over(query, Some(df.schema().as_ref())).unwrap();
let mut lf = df.clone().lazy();
if let Some(f) = filter {
lf = lf.filter(f);
}
if !by.is_empty() {
let keys = by.len();
lf = lf.group_by(by).agg(cols);
let schema = lf.collect_schema().unwrap();
let sort: Vec<Expr> = schema
.iter_names()
.take(keys)
.map(|n| col(n.as_str()))
.collect();
lf = lf.sort_by_exprs(sort, SortMultipleOptions::default());
} else if !cols.is_empty() {
lf = lf.select(cols);
}
if distinct {
lf = lf.unique_stable(None, UniqueKeepStrategy::First);
}
lf.collect().unwrap()
}
fn values(df: &DataFrame, name: &str) -> Vec<String> {
df.column(name)
.unwrap()
.as_materialized_series()
.iter()
.map(|v| match v {
AnyValue::String(s) => s.to_string(),
AnyValue::StringOwned(s) => s.to_string(),
v => v.to_string(),
})
.collect()
}
fn parse_err(query: &str) -> String {
parse_query(query).unwrap_err()
}
#[test]
fn test_time_part_accessors_parse() {
let expr = parse_expr(&tokenize("ts.hour").unwrap()).unwrap();
assert_eq!(expr, col("ts").dt().hour().alias("ts_hour"));
let expr = parse_expr(&tokenize("ts.doy").unwrap()).unwrap();
assert_eq!(expr, col("ts").dt().ordinal_day().alias("ts_doy"));
for accessor in ["hour", "minute", "second", "quarter", "doy"] {
let q = format!("select x.{}", accessor);
assert!(parse_query(&q).is_ok(), "{q}");
}
}
#[test]
fn test_time_part_accessors_evaluate() {
let df = df!("ts" => &["2024-03-15 13:45:30", "2024-12-31 00:00:05"])
.unwrap()
.lazy()
.select([col("ts").str().to_datetime(
None,
None,
StrptimeOptions::default(),
lit("raise"),
)])
.collect()
.unwrap();
let out = eval(
"select ts.hour, ts.minute, ts.second, ts.quarter, ts.doy",
&df,
);
assert_eq!(values(&out, "ts_hour"), ["13", "0"]);
assert_eq!(values(&out, "ts_minute"), ["45", "0"]);
assert_eq!(values(&out, "ts_second"), ["30", "5"]);
assert_eq!(values(&out, "ts_quarter"), ["1", "4"]);
assert_eq!(values(&out, "ts_doy"), ["75", "366"]);
}
#[test]
fn test_hour_groups_trips() {
let df =
df!("pickup" => &["2025-01-01 08:10:00", "2025-01-01 08:50:00", "2025-01-01 17:00:00"])
.unwrap()
.lazy()
.with_column(col("pickup").str().to_datetime(
None,
None,
StrptimeOptions::default(),
lit("raise"),
))
.collect()
.unwrap();
let out = eval("select trips: count pickup by pickup.hour", &df);
assert_eq!(values(&out, "pickup_hour"), ["8", "17"]);
assert_eq!(values(&out, "trips"), ["2", "1"]);
}
#[test]
fn a_timestamp_literal_takes_the_zone_of_its_column() {
let zoned = |zone: &str| {
df!("t" => &["2013-01-15 14:00:00", "2013-01-15 15:00:00"])
.unwrap()
.lazy()
.with_column(col("t").str().to_datetime(
Some(TimeUnit::Microseconds),
TimeZone::opt_try_new(Some(zone)).unwrap(),
StrptimeOptions::default(),
lit("raise"),
))
.collect()
.unwrap()
};
for zone in ["UTC", "America/New_York"] {
let df = zoned(zone);
for (query, rows) in [
("select where t > 2013.01.15T14:30:00.123456", 1),
("select where 2013.01.15T14:30:00 < t", 1),
("select where t = 2013.01.15T15:00:00", 1),
(
"select where t >= 2013.01.15T14:00:00, t < 2013.01.16T00:00:00",
2,
),
("select where t > 2013.01.15", 2),
] {
assert_eq!(eval(query, &df).height(), rows, "{zone}: {query}");
}
let out = eval("select later: t ^ 2013.01.15T00:00:00", &df);
assert_eq!(out.height(), 2, "{zone}");
}
let naive = df!("t" => &["2013-01-15 14:00:00"])
.unwrap()
.lazy()
.with_column(col("t").str().to_datetime(
None,
None,
StrptimeOptions::default(),
lit("raise"),
))
.collect()
.unwrap();
assert_eq!(
eval("select where t < 2013.01.15T14:30:00", &naive).height(),
1
);
let schema = zoned("America/New_York").schema().clone();
let mut nodes = parse_nodes("select where t > 2013.01.15T14:30:00").unwrap();
nodes.resolve_time_zones(&schema);
let python = nodes.python_filter().unwrap();
assert!(
python.contains("time_zone=\"America/New_York\", ambiguous=\"earliest\""),
"{python}"
);
}
fn temporal_frame() -> DataFrame {
df!("d" => &["2024-01-01"], "s" => &["2024.01.01"])
.unwrap()
.lazy()
.with_columns([
lit("2024-01-01T05:00:00")
.str()
.to_datetime(None, None, StrptimeOptions::default(), lit("raise"))
.alias("ts"),
lit("2024-01-01T05:00:00")
.str()
.to_datetime(None, None, StrptimeOptions::default(), lit("raise"))
.dt()
.time()
.alias("t"),
lit(5i64)
.cast(DataType::Duration(TimeUnit::Milliseconds))
.alias("dur"),
col("d").str().to_date(StrptimeOptions::default()),
])
.collect()
.unwrap()
}
fn parse_error_over(query: &str, df: &DataFrame) -> String {
parse_query_over(query, Some(df.schema().as_ref()))
.err()
.unwrap_or_else(|| panic!("{query} should fail"))
}
#[test]
fn test_quoted_text_against_temporal_column_is_a_q_error() {
let df = temporal_frame();
let cases = [
(
"select where d = \"2024.01.01\"",
"d is a date; \"2024.01.01\" is a string. A date is 2024.01.01",
),
(
"select where d < \"Jan 1\"",
"d is a date; \"Jan 1\" is a string. A date is 2024.01.01",
),
(
"select where d = \"2024-01-01\"",
"d is a date; \"2024-01-01\" is a string. A date is 2024.01.01",
),
(
"select where ts = \"2024.01.01T05:00:00\"",
"ts is a timestamp; \"2024.01.01T05:00:00\" is a string. A timestamp is 2024.01.01T05:00:00",
),
(
"select where ts < \"2023.06.30T23:59:59.5\"",
"ts is a timestamp; \"2023.06.30T23:59:59.5\" is a string. A timestamp is 2023.06.30T23:59:59.5",
),
(
"select where ts < \"2024.01.01\"",
"ts is a timestamp; \"2024.01.01\" is a string. A timestamp is 2024.01.01T05:00:00",
),
(
"select where t = \"05:00:00\"",
"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",
),
(
"select where t < \"05:00:00\"",
"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",
),
(
"select where dur = \"5s\"",
"dur is a duration; \"5s\" is a string. A duration has no literal",
),
(
"select where \"5s\" >= dur",
"dur is a duration; \"5s\" is a string. A duration has no literal",
),
(
"select where d in [\"2024.01.01\", \"2024.01.02\"]",
"d is a date; \"2024.01.01\" is a string. A date is 2024.01.01",
),
(
"select x: d != \"x\"",
"d is a date; \"x\" is a string. A date is 2024.01.01",
),
];
for (query, want) in cases {
assert_eq!(parse_error_over(query, &df), want, "{query}");
}
let mut renamed = df.clone();
renamed.rename("d", "start date".into()).unwrap();
assert_eq!(
parse_error_over("select where col[\"start date\"] = \"x\"", &renamed),
"col[\"start date\"] is a date; \"x\" is a string. A date is 2024.01.01"
);
assert_eq!(eval("select where d = 2024.01.01", &df).height(), 1);
assert_eq!(eval("select where d in [2024.01.01]", &df).height(), 1);
assert_eq!(
eval("select where ts = 2024.01.01T05:00:00", &df).height(),
1
);
assert_eq!(eval("select where t.hour = 5", &df).height(), 1);
}
#[test]
fn test_quoted_text_against_text_column_still_compares() {
let df = temporal_frame();
assert_eq!(eval("select where s = \"2024.01.01\"", &df).height(), 1);
assert_eq!(eval("select where s < \"2025\"", &df).height(), 1);
assert_eq!(eval("select where s in [\"2024.01.01\"]", &df).height(), 1);
assert!(
parse_query_over("select where d like \"2024*\"", Some(df.schema().as_ref())).is_ok()
);
assert!(parse_query_over("select where d = \"2024.01.01\"", None).is_ok());
}
#[test]
fn test_to_date_and_to_datetime_parse_strings() {
let df = df!(
"DATE" => &["20240101", "20241231", "junk"],
"Date" => &["Sat Sep 12 2020", "Tue Jan 12 2021(P)", "Sun Sep 13 2020"],
"stamp" => &["2024-01-02 03:04", "2024-05-06 07:08", "nope"],
)
.unwrap();
let out = eval(
"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\"]",
&df,
);
assert_eq!(values(&out, "day"), ["2024-01-01", "2024-12-31", "null"]);
assert_eq!(
values(&out, "d"),
["2020-09-12", "2021-01-12", "2020-09-13"]
);
assert_eq!(
values(&out, "t"),
["2024-01-02 03:04:00", "2024-05-06 07:08:00", "null"]
);
}
#[test]
fn test_to_date_parses_an_integer_column() {
let df = df!("DATE" => &[20240101i64, 20240229]).unwrap();
let out = eval("select d: DATE.to_date[\"%Y%m%d\"]", &df);
assert_eq!(values(&out, "d"), ["2024-01-01", "2024-02-29"]);
}
#[test]
fn test_casts() {
let df = df!(
"s" => &["3", "4.5", "x"],
"f" => &[1.9f64, -1.9, 3.0],
)
.unwrap();
let out = eval("select a: s.int, b: s.float, c: f.int, d: f.str", &df);
assert_eq!(values(&out, "a"), ["3", "null", "null"]);
assert_eq!(values(&out, "b"), ["3.0", "4.5", "null"]);
assert_eq!(values(&out, "c"), ["1", "-1", "3"]);
assert_eq!(values(&out, "d"), ["1.9", "-1.9", "3.0"]);
assert_eq!(out.column("a").unwrap().dtype(), &DataType::Int64);
assert_eq!(out.column("b").unwrap().dtype(), &DataType::Float64);
assert_eq!(out.column("d").unwrap().dtype(), &DataType::String);
}
#[test]
fn test_string_pieces() {
let df = df!("FT" => &["0–3", "12–1", " 2–2 "]).unwrap();
let out = eval(
"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[\"–\", \"-\"]",
&df,
);
assert_eq!(values(&out, "home"), ["0", "12", "null"]);
assert_eq!(values(&out, "away"), ["3", "1", "null"]);
assert_eq!(values(&out, "none"), ["null", "null", "null"]);
assert_eq!(values(&out, "head"), ["0–", "12", " "]);
assert_eq!(values(&out, "tail"), ["–3", "–1", " "]);
assert_eq!(values(&out, "s"), ["0–3", "12–1", "2–2"]);
assert_eq!(values(&out, "r"), ["0-3", "12-1", " 2-2 "]);
}
#[test]
fn test_string_pieces_auto_alias() {
let cols = parse_query("select FT.part[\"-\", 0], FT.strip")
.unwrap()
.cols;
let names: Vec<String> = cols
.iter()
.map(|e| e.clone().meta().output_name().unwrap().to_string())
.collect();
assert_eq!(names, ["FT_part_-_0", "FT_strip"]);
}
#[test]
fn test_in_parses_to_equalities() {
let filter = parse_query("select where name in [\"a\", \"b\"]")
.unwrap()
.filter;
assert_eq!(
filter,
Some(col("name").eq(lit("a")).or(col("name").eq(lit("b"))))
);
}
#[test]
fn test_in_filters() {
let df = df!(
"name" => &["Emma", "Jennifer", "Olivia", "Mary"],
"n" => &[1i32, 2, 3, 4],
)
.unwrap();
let out = eval(
"select name where name in [\"Emma\", \"Jennifer\", \"Olivia\"]",
&df,
);
assert_eq!(values(&out, "name"), ["Emma", "Jennifer", "Olivia"]);
let out = eval("select n where n in [2, 4.0, -1]", &df);
assert_eq!(values(&out, "n"), ["2", "4"]);
let out = eval("select name where not name in [\"Mary\"]", &df);
assert_eq!(values(&out, "name"), ["Emma", "Jennifer", "Olivia"]);
let out = eval("select name where name in [\"Emma\", \"Mary\"], n > 1", &df);
assert_eq!(values(&out, "name"), ["Mary"]);
}
#[test]
fn test_in_long_list_nests_shallowly() {
let items: Vec<String> = (0..2000).map(|i| i.to_string()).collect();
let q = format!("select where x in [{}]", items.join(", "));
let df = df!("x" => &[5i64, 1999, 2000]).unwrap();
assert_eq!(values(&eval(&q, &df), "x"), ["5", "1999"]);
}
#[test]
fn test_in_a_list_past_the_node_cap_on_a_column() {
let items: Vec<String> = (0..12_000).map(|i| i.to_string()).collect();
let q = format!("select where x in [{}]", items.join(", "));
let df = df!("x" => &[5i64, 11_999, 12_000]).unwrap();
assert_eq!(values(&eval(&q, &df), "x"), ["5", "11999"]);
}
#[test]
fn test_in_errors() {
let err = parse_err("select where x in 1");
assert!(err.contains("in takes a list"), "{err}");
let err = parse_err("select where x in []");
assert!(err.contains("in needs a list of values"), "{err}");
let err = parse_err("select where x in [1,, 2]");
assert!(err.contains("in needs a list of values"), "{err}");
let err = parse_err("select where x in [1] = y");
assert!(err.contains("in takes a list"), "{err}");
let err = parse_err("select where x in [1] + [2]");
assert!(err.contains("in takes a list"), "{err}");
}
#[test]
fn test_like_matches_whole_value() {
let df = df!("item" => &["Crispy Chicken", "Chicken", "Fish", "a.b", "axb"]).unwrap();
let out = eval("select item where item like \"*Chicken*\"", &df);
assert_eq!(values(&out, "item"), ["Crispy Chicken", "Chicken"]);
let out = eval("select item where item like \"Chick*\"", &df);
assert_eq!(values(&out, "item"), ["Chicken"]);
let out = eval("select item where item like \"F?sh\"", &df);
assert_eq!(values(&out, "item"), ["Fish"]);
let out = eval("select item where item like \"a.b\"", &df);
assert_eq!(values(&out, "item"), ["a.b"]);
}
#[test]
fn test_like_errors() {
let err = parse_err("select where item like Chicken");
assert!(err.contains("like takes a quoted pattern"), "{err}");
}
#[test]
fn test_like_regex() {
assert_eq!(like_regex("*a?.b*"), "(?s)^.*a.\\.b.*$");
}
fn integer_widths() -> (DataFrame, Vec<&'static str>) {
let names = ["i8", "i16", "i32", "i64", "u8", "u16", "u32", "u64"];
let types = [
DataType::Int8,
DataType::Int16,
DataType::Int32,
DataType::Int64,
DataType::UInt8,
DataType::UInt16,
DataType::UInt32,
DataType::UInt64,
];
let mut columns = Vec::new();
for (name, dtype) in names.iter().zip(types) {
for (side, vals) in [("a", [10i64, 20, 30]), ("b", [1, 2, 3])] {
let c = Column::new(format!("{side}_{name}").into(), vals);
columns.push(c.cast(&dtype).unwrap());
}
}
let df = DataFrame::new_infer_height(columns).unwrap();
(df, names.to_vec())
}
fn as_f64(df: &DataFrame, name: &str) -> Vec<f64> {
df.column(name)
.unwrap()
.cast(&DataType::Float64)
.unwrap()
.f64()
.unwrap()
.into_no_null_iter()
.collect()
}
#[test]
fn mixed_integer_widths_do_arithmetic() {
let (df, names) = integer_widths();
let ops: [(&str, [f64; 3]); 5] = [
("+", [11.0, 22.0, 33.0]),
("-", [9.0, 18.0, 27.0]),
("*", [10.0, 40.0, 90.0]),
("/", [10.0, 10.0, 10.0]),
("%", [10.0, 10.0, 10.0]),
];
let mut parts = Vec::new();
let mut expected = Vec::new();
for (i, (op, want)) in ops.iter().enumerate() {
for l in &names {
for r in &names {
let alias = format!("r{i}_{l}_{r}");
parts.push(format!("{alias}: a_{l} {op} b_{r}"));
expected.push((alias, *want));
}
}
}
for l in &names {
for r in &names {
let alias = format!("m_{l}_{r}");
parts.push(format!("{alias}: a_{l} mod b_{r}"));
expected.push((alias, [0.0, 0.0, 0.0]));
}
}
let out = eval(&format!("select {}", parts.join(", ")), &df);
for (alias, want) in expected {
assert_eq!(as_f64(&out, &alias), want, "{alias}");
}
}
#[test]
fn mixed_integer_widths_with_literals() {
let (df, names) = integer_widths();
let mut parts = Vec::new();
let mut expected = Vec::new();
for name in &names {
for (tag, expr, want) in [
("p", format!("b_{name} + 7"), [8.0, 9.0, 10.0]),
("s", format!("a_{name} - 7"), [3.0, 13.0, 23.0]),
("l", format!("7 - b_{name}"), [6.0, 5.0, 4.0]),
("m", format!("b_{name} * 2.5"), [2.5, 5.0, 7.5]),
("d", format!("a_{name} / 2"), [5.0, 10.0, 15.0]),
("q", format!("60 / b_{name}"), [60.0, 30.0, 20.0]),
("r", format!("a_{name} mod 7"), [3.0, 6.0, 2.0]),
] {
let alias = format!("{tag}_{name}");
parts.push(format!("{alias}: {expr}"));
expected.push((alias, want));
}
}
let out = eval(&format!("select {}", parts.join(", ")), &df);
for (alias, want) in expected {
assert_eq!(as_f64(&out, &alias), want, "{alias}");
}
}
#[test]
fn mixed_integer_widths_filter_and_group() {
let (df, _) = integer_widths();
let out = eval("select a_u8 where 10 = a_i64 / b_u8, 8 < b_u16 + 7", &df);
assert_eq!(values(&out, "a_u8"), ["20", "30"]);
let out = eval(
"select t: sum a_u8 * b_i16, n: max a_i64 - b_u32 by k: b_u16 mod 2",
&df,
);
assert_eq!(as_f64(&out, "k"), [0.0, 1.0]);
assert_eq!(as_f64(&out, "t"), [40.0, 100.0]);
assert_eq!(as_f64(&out, "n"), [18.0, 27.0]);
}
#[cfg(feature = "sql")]
#[test]
fn mixed_integer_widths_in_sql() {
let (df, _) = integer_widths();
let mut ctx = polars_sql::SQLContext::new();
ctx.register("df", df.lazy());
let out = ctx
.execute(
"SELECT a_i64 / b_u8 AS d, a_u16 % b_u8 AS m, b_u8 + 7 AS p, \
a_i8 * b_u64 AS x FROM df WHERE a_u8 - b_i16 > 9",
)
.unwrap()
.collect()
.unwrap();
assert_eq!(as_f64(&out, "d"), [10.0, 10.0]);
assert_eq!(as_f64(&out, "m"), [0.0, 0.0]);
assert_eq!(as_f64(&out, "p"), [9.0, 10.0]);
assert_eq!(as_f64(&out, "x"), [40.0, 90.0]);
}
#[test]
fn test_xbar_buckets() {
let df = df!(
"fare" => &[-1.0f64, 0.0, 4.99, 5.0, 12.5],
"n" => &[-1i64, 0, 4, 5, 12],
)
.unwrap();
let out = eval("select f: 5 xbar fare, i: 5 xbar n, h: 0.5 xbar fare", &df);
assert_eq!(values(&out, "f"), ["-5.0", "0.0", "0.0", "5.0", "10.0"]);
assert_eq!(values(&out, "i"), ["-5", "0", "0", "5", "10"]);
assert_eq!(out.column("i").unwrap().dtype(), &DataType::Int64);
assert_eq!(values(&out, "h"), ["-1.0", "0.0", "4.5", "5.0", "12.5"]);
}
#[test]
fn test_xbar_groups() {
let df = df!("fare" => &[1.0f64, 3.0, 7.0, 12.0, 14.0]).unwrap();
let out = eval("select trips: count fare by b: 5 xbar fare", &df);
assert_eq!(values(&out, "b"), ["0.0", "5.0", "10.0"]);
assert_eq!(values(&out, "trips"), ["2", "1", "2"]);
}
#[test]
fn test_xbar_errors() {
let err = parse_err("select 0 xbar fare");
assert!(err.contains("positive bucket size"), "{err}");
let err = parse_err("select -5 xbar fare");
assert!(err.contains("positive bucket size"), "{err}");
}
#[test]
fn test_mod() {
let df = df!("n" => &[-7i64, 7, 9], "f" => &[7.5f64, -0.5, 2.0]).unwrap();
let out = eval("select a: n mod 3, b: f mod 2, c: -7 mod 3", &df);
assert_eq!(values(&out, "a"), ["2", "1", "0"]);
assert_eq!(values(&out, "b"), ["1.5", "1.5", "0.0"]);
assert_eq!(values(&out, "c"), ["2", "2", "2"]);
let out = eval("select a: n mod -3", &df);
assert_eq!(values(&out, "a"), ["-1", "-2", "0"]);
assert_eq!(out.column("a").unwrap().dtype(), &DataType::Int64);
}
#[test]
fn test_word_operators_right_to_left() {
let parse = |s: &str| parse_expr(&tokenize(s).unwrap()).unwrap();
assert_eq!(parse("a = b mod 2"), col("a").eq(col("b").rem(lit(2i64))));
assert_eq!(
parse("5 xbar x + 1"),
col("x").add(lit(1.0)).floor_div(lit(5i64)).mul(lit(5i64))
);
assert_eq!(
parse("2 * 5 xbar x"),
lit(2.0).mul(col("x").floor_div(lit(5i64)).mul(lit(5i64)))
);
assert_eq!(
parse("flag = name in [\"a\"]"),
col("flag").eq(col("name").eq(lit("a")))
);
assert_eq!(parse("x mod 2 in [1]"), col("x").rem(lit(2.0).eq(lit(1.0))));
assert_eq!(
parse("ok = name like \"a*\""),
col("ok").eq(col("name")
.cast(DataType::String)
.str()
.contains(lit("(?s)^a.*$"), true))
);
}
#[test]
fn test_word_operators_right_to_left_evaluate() {
let df = df!("x" => &[3i64, 4, 9]).unwrap();
let out = eval("select a: 1 + x mod 4", &df);
assert_eq!(values(&out, "a"), ["4.0", "1.0", "2.0"]);
let out = eval("select a: (1 + x) mod 4", &df);
assert_eq!(values(&out, "a"), ["0.0", "1.0", "2.0"]);
let out = eval("select x where (x mod 2) in [1]", &df);
assert_eq!(values(&out, "x"), ["3", "9"]);
}
#[test]
fn test_word_operators_are_still_column_names() {
let cols = parse_query("select in, mod, like + xbar").unwrap().cols;
assert_eq!(cols[0], col("in"));
assert_eq!(cols[1], col("mod"));
assert_eq!(cols[2], col("like").add(col("xbar")));
let err = parse_err("select x.in");
assert!(err.contains("Unknown accessor: 'in'"), "{err}");
}
#[test]
fn test_new_aggregates() {
let cols = parse_query("select nunique ID, var x, dev x by g")
.unwrap()
.cols;
assert_eq!(cols[0], col("ID").n_unique().alias("nunique_ID"));
assert_eq!(cols[1], col("x").var(1).alias("var_x"));
assert_eq!(cols[2], col("x").std(1).alias("dev_x"));
let df = df!(
"g" => &["a", "a", "a", "b"],
"ID" => &["s1", "s1", "s2", "s3"],
"x" => &[1.0f64, 2.0, 3.0, 5.0],
)
.unwrap();
let out = eval("select nunique ID, var x, dev[x] by g", &df);
assert_eq!(values(&out, "nunique_ID"), ["2", "1"]);
assert_eq!(values(&out, "var_x"), ["1.0", "null"]);
assert_eq!(values(&out, "dev_x"), ["1.0", "null"]);
}
#[test]
fn test_wavg() {
let df = df!(
"g" => &["a", "a", "a", "b"],
"w" => &[Some(1i64), Some(3), Some(5), Some(2)],
"x" => &[Some(10.0f64), Some(20.0), None, Some(4.0)],
)
.unwrap();
let out = eval("select w wavg x by g", &df);
assert_eq!(values(&out, "wavg_x"), ["17.5", "4.0"]);
let out = eval("select w wavg x where null x", &df);
assert_eq!(values(&out, "wavg_x"), ["null"]);
for query in ["select wavg[x]", "select wavg x"] {
let err = parse_err(query);
assert!(
err.contains("wavg goes between weights and values"),
"{err}"
);
}
}
#[test]
fn test_round_and_math_functions() {
let df = df!("x" => &[2.25f64, -2.5, 4.0]).unwrap();
let out = eval(
"select r: x.round, r1: x.round[1], s: sqrt x, l: log[x], e: exp 0 * x",
&df,
);
assert_eq!(values(&out, "r"), ["2.0", "-3.0", "4.0"]);
assert_eq!(values(&out, "r1"), ["2.3", "-2.5", "4.0"]);
assert_eq!(values(&out, "s"), ["1.5", "NaN", "2.0"]);
assert!(values(&out, "l")[2].starts_with("1.386"));
assert_eq!(values(&out, "e"), ["1.0", "1.0", "1.0"]);
let df = df!("g" => &["a", "a"], "d" => &[1.0f64, 2.34]).unwrap();
let out = eval("select m: (avg d).round[1] by g", &df);
assert_eq!(values(&out, "m"), ["1.7"]);
}
#[test]
fn test_select_distinct() {
let ParsedQuery { cols, distinct, .. } =
parse_query("select distinct carrier, origin").unwrap();
assert!(distinct);
assert_eq!(cols, vec![col("carrier"), col("origin")]);
let distinct = parse_query("select carrier").unwrap().distinct;
assert!(!distinct);
let df = df!(
"carrier" => &["UA", "UA", "AA", "UA"],
"origin" => &["EWR", "EWR", "JFK", "LGA"],
"n" => &[1i32, 2, 3, 4],
)
.unwrap();
let out = eval("select distinct carrier, origin", &df);
assert_eq!(values(&out, "carrier"), ["UA", "AA", "UA"]);
assert_eq!(values(&out, "origin"), ["EWR", "JFK", "LGA"]);
let out = eval("select distinct carrier where n > 1", &df);
assert_eq!(values(&out, "carrier"), ["UA", "AA"]);
let ParsedQuery { cols, distinct, .. } = parse_query("select col[\"distinct\"]").unwrap();
assert!(!distinct);
assert_eq!(cols, vec![col("distinct")]);
let ParsedQuery { cols, distinct, .. } = parse_query("select distinct, n").unwrap();
assert!(!distinct);
assert_eq!(cols, vec![col("distinct"), col("n")]);
let ParsedQuery { cols, distinct, .. } = parse_query("select distinct: n").unwrap();
assert!(!distinct);
assert_eq!(cols, vec![col("n").alias("distinct")]);
}
#[test]
fn test_from_df_is_optional() {
for (with, without) in [
(
"select mean dep_delay by hour from df where origin = \"JFK\"",
"select mean dep_delay by hour where origin = \"JFK\"",
),
("select from df where x > 1", "select where x > 1"),
("select from df", "select"),
("select a, b from df", "select a, b"),
("select distinct a from df", "select distinct a"),
("select n: count a by g from df", "select n: count a by g"),
] {
assert_eq!(
format!("{:?}", parse_query(with).unwrap()),
format!("{:?}", parse_query(without).unwrap()),
"{with}"
);
}
}
#[test]
fn test_from_names_only_df() {
for query in [
"select from trades",
"select a by g from trades where a > 1",
"select from data.csv",
] {
assert_eq!(
parse_query(query).unwrap_err(),
"q reads the table on screen, named df: … from df …",
"{query}"
);
}
let err = parse_query("select a where a > 1 from df").unwrap_err();
assert!(err.contains("after the where clause"), "{err}");
let err = parse_query("select a from df by g").unwrap_err();
assert!(err.contains("'by' after 'from df'"), "{err}");
}
#[test]
fn test_from_column_names_and_values() {
let cols = |q: &str| parse_query(q).unwrap().cols;
assert_eq!(cols("select from"), vec![col("from")]);
assert_eq!(cols("select from, to"), vec![col("from"), col("to")]);
assert_eq!(cols("select from from df"), vec![col("from")]);
assert_eq!(cols("select from + 1"), vec![col("from") + lit(1.0)]);
assert_eq!(cols("select from.year"), cols("select col[\"from\"].year"));
assert_eq!(cols("select max from"), cols("select max col[\"from\"]"));
let ParsedQuery { group_by, .. } = parse_query("select n: count a by from").unwrap();
assert_eq!(group_by, vec![col("from")]);
assert_eq!(
parse_query("select where from = \"df\"").unwrap().filter,
Some(col("from").eq(lit("df")))
);
assert_eq!(
parse_query("select where from in [1, 2]").unwrap().filter,
parse_query("select where col[\"from\"] in [1, 2]")
.unwrap()
.filter
);
assert_eq!(
cols("select from_city, datefrom from df"),
vec![col("from_city"), col("datefrom")]
);
assert_eq!(
parse_query("select from df where city = \"from df\"")
.unwrap()
.filter,
Some(col("city").eq(lit("from df")))
);
assert_eq!(cols("select df from df"), vec![col("df")]);
let df = df!("from" => &[1i64, 2, 3], "df" => &["x", "y", "z"]).unwrap();
let out = eval("select from, df from df where from > 1", &df);
assert_eq!(values(&out, "df"), ["y", "z"]);
}
#[test]
fn test_accessor_argument_count_errors() {
for (query, expected) in [
(
"select x.part[\",\"]",
"part takes 2 arguments, e.g. .part[\"-\", 0]; got 1",
),
("select x.slice", "slice takes 1 to 2 arguments"),
("select x.replace[\"a\"]", "replace takes 2 arguments"),
("select x.round[1, 2]", "round takes 0 to 1 arguments"),
(
"select x.to_date[\"%Y\", \"%m\"]",
"to_date takes 0 to 1 arguments",
),
("select x.hour[1]", "hour takes no arguments"),
("select x.strip[\" \"]", "strip takes no arguments"),
("select x.int[1]", "int takes no arguments"),
("select x.format", "format takes 1 argument"),
] {
let err = parse_err(query);
assert!(err.contains(expected), "{query}: {err}");
}
}
#[test]
fn test_accessor_argument_type_errors() {
for (query, expected) in [
(
"select x.part[0, \",\"]",
"part: argument 1 must be quoted text",
),
(
"select x.part[\",\", \"a\"]",
"part: argument 2 must be a whole number",
),
(
"select x.part[\",\", 1.5]",
"part: argument 2 must be a whole number",
),
("select x.round[-1]", "round: decimals cannot be negative"),
(
"select x.slice[0, -1]",
"slice: the length cannot be negative",
),
("select x.slice[a + 1]", "slice takes literal arguments"),
("select x.part[\",\"", "Unmatched bracket after .part"),
] {
let err = parse_err(query);
assert!(err.contains(expected), "{query}: {err}");
}
}
#[test]
fn test_unknown_accessor_lists_new_names() {
let err = parse_err("select x.nosuch");
for name in [
"hour",
"minute",
"second",
"quarter",
"doy",
"to_date",
"to_datetime",
"part",
"slice",
"replace",
"strip",
"round",
"int",
"float",
"str",
] {
assert!(err.contains(name), "{name} missing from: {err}");
}
}
#[test]
fn test_nested_functions_parse_in_linear_time() {
let bare = format!("select {}x", "abs ".repeat(40));
assert!(parse_query(&bare).is_ok());
let bracketed = format!("select {}x{}", "sqrt[".repeat(30), "]".repeat(30));
assert!(parse_query(&bracketed).is_ok());
}
fn py(expr: &str) -> String {
parse_node(&tokenize(expr).unwrap()).unwrap().python()
}
#[test]
fn expressions_read_as_python_polars() {
assert_eq!(py("a"), "pl.col(\"a\")");
assert_eq!(py("col[\"first name\"]"), "pl.col(\"first name\")");
assert_eq!(py("a > 1"), "pl.col(\"a\") > 1.0");
assert_eq!(
py("a + b * c"),
"pl.col(\"a\") + (pl.col(\"b\") * pl.col(\"c\"))"
);
assert_eq!(py("-x"), "pl.lit(0) - pl.col(\"x\")");
assert_eq!(py("x mod 3"), "pl.col(\"x\") % 3");
assert_eq!(py("5 xbar fare"), "(pl.col(\"fare\") // 5) * 5");
assert_eq!(py("a ^ 0"), "pl.coalesce(pl.col(\"a\"), pl.lit(0.0))");
assert_eq!(
py("name in [\"Emma\", \"Olivia\"]"),
"(pl.col(\"name\") == \"Emma\") | (pl.col(\"name\") == \"Olivia\")"
);
assert_eq!(
py("item like \"*Chicken*\""),
"pl.col(\"item\").cast(pl.String).str.contains(\"(?s)^.*Chicken.*$\")"
);
assert_eq!(
py("d = 2024.01.31"),
"pl.col(\"d\") == pl.date(2024, 1, 31)"
);
assert_eq!(
py("t > 2024.01.31T10:00:00.5"),
"pl.col(\"t\") > pl.lit(\"2024-01-31T10:00:00.500\").str.to_datetime(\"%Y-%m-%dT%H:%M:%S%.3f\", time_unit=\"ms\")"
);
}
#[test]
fn division_reads_as_polars_runs_it_on_the_types() {
let schema = Schema::from_iter([
Field::new("i".into(), DataType::Int64),
Field::new("j".into(), DataType::Int32),
Field::new("u".into(), DataType::UInt8),
Field::new("f".into(), DataType::Float64),
Field::new("s".into(), DataType::String),
]);
let py = |expr: &str| {
let mut node = parse_node(&tokenize(expr).unwrap()).unwrap();
node.resolve_division(&schema);
node.python()
};
assert_eq!(py("i / j"), "pl.col(\"i\") // pl.col(\"j\")");
assert_eq!(py("j % i"), "pl.col(\"j\") // pl.col(\"i\")");
assert_eq!(py("i / u"), "pl.col(\"i\") // pl.col(\"u\")");
assert_eq!(py("i / s"), "pl.col(\"i\") / pl.col(\"s\")");
assert_eq!(py("(i mod 3) / j"), "(pl.col(\"i\") % 3) // pl.col(\"j\")");
assert_eq!(py("i / f"), "pl.col(\"i\") / pl.col(\"f\")");
assert_eq!(py("i / 2"), "pl.col(\"i\") / 2.0");
assert_eq!(
py("sum[i] / count[j]"),
"pl.col(\"i\").sum().alias(\"sum_i\") // pl.col(\"j\").count().alias(\"count_j\")"
);
assert_eq!(py("x / i"), "pl.col(\"x\") / pl.col(\"i\")");
}
#[test]
fn functions_and_accessors_read_as_python_polars() {
assert_eq!(
py("avg salary"),
"pl.col(\"salary\").mean().alias(\"avg_salary\")"
);
assert_eq!(py("not null[x]"), "pl.col(\"x\").is_null().not_()");
assert_eq!(py("log x"), "pl.col(\"x\").log()");
assert_eq!(py("ts.year"), "pl.col(\"ts\").dt.year().alias(\"ts_year\")");
assert_eq!(
py("d.format[\"%Y-%m\"]"),
"pl.col(\"d\").dt.to_string(\"%Y-%m\").alias(\"d_format_%Y-%m\")"
);
assert_eq!(
py("code.part[\"-\", 0]"),
"pl.col(\"code\").cast(pl.String).str.split(\"-\").list.get(0, null_on_oob=True).alias(\"code_part_-_0\")"
);
assert_eq!(
py("s.slice[1]"),
"pl.col(\"s\").cast(pl.String).str.slice(1).alias(\"s_slice_1\")"
);
assert_eq!(
py("s.to_date[\"%Y%m%d\"]"),
"pl.col(\"s\").cast(pl.String).str.to_date(\"%Y%m%d\", strict=False).alias(\"s_to_date_%Y%m%d\")"
);
assert_eq!(
py("x.round[2]"),
"pl.col(\"x\").round(2, mode=\"half_away_from_zero\").alias(\"x_round_2\")"
);
assert_eq!(
py("x.int"),
"pl.col(\"x\").cast(pl.Int64, strict=False).alias(\"x_int\")"
);
assert_eq!(
py("w wavg v"),
"((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\")"
);
}
#[test]
fn a_whole_query_reads_as_python_steps() {
let steps = |q: &str, keys: &[&str]| {
let keys: Vec<String> = keys.iter().map(|k| k.to_string()).collect();
parse_nodes(q).unwrap().python_steps(&keys)
};
assert_eq!(
steps(
"select name, pay: salary * 1.1 where dept = \"Sales\", age > 30 | senior",
&[]
),
vec![
".filter((pl.col(\"dept\") == \"Sales\") & ((pl.col(\"age\") > 30.0) | pl.col(\"senior\")))",
".select(\"name\", (pl.col(\"salary\") * 1.1).alias(\"pay\"))",
]
);
assert_eq!(
steps("select by dept", &["dept"]),
vec![
".group_by(\"dept\")",
".agg(pl.all().exclude(\"dept\"))",
".sort(\"dept\", nulls_last=True, maintain_order=True)",
]
);
assert_eq!(
steps("select distinct dept", &[]),
vec![
".select(\"dept\")",
".unique(keep=\"first\", maintain_order=True)",
]
);
assert!(steps("", &[]).is_empty());
}
}