From 2525500f2a8b19fc0f9592dd337f758806c634ac Mon Sep 17 00:00:00 2001 From: Sahil Kalgutkar Date: Fri, 28 Aug 2026 14:16:16 -0400 Subject: [PATCH] Add the SQL front end: lexer, syntax tree and parser MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Hand-written throughout. The lexer handles the cases that separate a toy tokeniser from a usable one: quoted identifiers that may hold keywords or spaces, doubled quotes as escapes inside both strings and identifiers, line comments, exponent notation, and an integer too large for i64 falling back to a float rather than failing the query. The parser is recursive descent for statements and precedence climbing for expressions. The precedence table is the part worth reading: NOT sits between AND and comparison so `NOT a = 1` groups as `NOT (a = 1)`, and BETWEEN parses its bounds above AND's precedence so the AND separating them is not mistaken for a conjunction. The tree stays purely syntactic — a column is a name the user typed, not a position — so that resolving names against a schema happens in exactly one place later. --- crates/qf-sql/src/ast.rs | 764 +++++++++++++++++++++++++ crates/qf-sql/src/lexer.rs | 571 ++++++++++++++++++ crates/qf-sql/src/lib.rs | 7 +- crates/qf-sql/src/parser.rs | 1080 +++++++++++++++++++++++++++++++++++ 4 files changed, 2421 insertions(+), 1 deletion(-) create mode 100644 crates/qf-sql/src/ast.rs create mode 100644 crates/qf-sql/src/lexer.rs create mode 100644 crates/qf-sql/src/parser.rs diff --git a/crates/qf-sql/src/ast.rs b/crates/qf-sql/src/ast.rs new file mode 100644 index 0000000..e0f1f49 --- /dev/null +++ b/crates/qf-sql/src/ast.rs @@ -0,0 +1,764 @@ +//! The syntax tree the parser produces and the binder consumes. +//! +//! Nothing here is resolved: a `Column` is a name the user typed, not a +//! position, and a `Function` is a name, not an aggregate. Keeping the AST +//! purely syntactic means the binder is the single place where "does this +//! column exist" is answered, and error messages about missing columns all +//! come from one place. + +use qf_common::{DataType, Value}; +use std::fmt; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum BinaryOp { + Plus, + Minus, + Multiply, + Divide, + Modulo, + Eq, + NotEq, + Lt, + LtEq, + Gt, + GtEq, + And, + Or, +} + +impl BinaryOp { + pub fn is_comparison(self) -> bool { + matches!( + self, + BinaryOp::Eq + | BinaryOp::NotEq + | BinaryOp::Lt + | BinaryOp::LtEq + | BinaryOp::Gt + | BinaryOp::GtEq + ) + } + + pub fn is_arithmetic(self) -> bool { + matches!( + self, + BinaryOp::Plus + | BinaryOp::Minus + | BinaryOp::Multiply + | BinaryOp::Divide + | BinaryOp::Modulo + ) + } + + pub fn is_logical(self) -> bool { + matches!(self, BinaryOp::And | BinaryOp::Or) + } + + /// Mirrors a comparison so that `5 > x` can be rewritten as `x < 5` — the + /// form the optimiser needs before it can push a predicate into a scan. + pub fn swap_operands(self) -> Option { + Some(match self { + BinaryOp::Eq => BinaryOp::Eq, + BinaryOp::NotEq => BinaryOp::NotEq, + BinaryOp::Lt => BinaryOp::Gt, + BinaryOp::LtEq => BinaryOp::GtEq, + BinaryOp::Gt => BinaryOp::Lt, + BinaryOp::GtEq => BinaryOp::LtEq, + _ => return None, + }) + } +} + +impl fmt::Display for BinaryOp { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let s = match self { + BinaryOp::Plus => "+", + BinaryOp::Minus => "-", + BinaryOp::Multiply => "*", + BinaryOp::Divide => "/", + BinaryOp::Modulo => "%", + BinaryOp::Eq => "=", + BinaryOp::NotEq => "<>", + BinaryOp::Lt => "<", + BinaryOp::LtEq => "<=", + BinaryOp::Gt => ">", + BinaryOp::GtEq => ">=", + BinaryOp::And => "AND", + BinaryOp::Or => "OR", + }; + f.write_str(s) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum UnaryOp { + Neg, + Not, +} + +impl fmt::Display for UnaryOp { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + UnaryOp::Neg => f.write_str("-"), + UnaryOp::Not => f.write_str("NOT "), + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub enum Expr { + /// `col` or `table.col`, as written. + Column { + table: Option, + name: String, + }, + Literal(Value), + Binary { + left: Box, + op: BinaryOp, + right: Box, + }, + Unary { + op: UnaryOp, + expr: Box, + }, + /// `f(args)` — resolved to an aggregate or rejected by the binder. + Function { + name: String, + args: Vec, + distinct: bool, + }, + Cast { + expr: Box, + data_type: DataType, + }, + IsNull { + expr: Box, + negated: bool, + }, + Between { + expr: Box, + low: Box, + high: Box, + negated: bool, + }, + InList { + expr: Box, + list: Vec, + negated: bool, + }, + Like { + expr: Box, + pattern: Box, + negated: bool, + }, + Case { + operand: Option>, + branches: Vec<(Expr, Expr)>, + else_result: Option>, + }, + /// The `*` in `count(*)`. + Wildcard, +} + +impl Expr { + pub fn binary(left: Expr, op: BinaryOp, right: Expr) -> Expr { + Expr::Binary { + left: Box::new(left), + op, + right: Box::new(right), + } + } + + pub fn column(name: impl Into) -> Expr { + Expr::Column { + table: None, + name: name.into(), + } + } + + pub fn qualified(table: impl Into, name: impl Into) -> Expr { + Expr::Column { + table: Some(table.into()), + name: name.into(), + } + } + + /// Splits a chain of `AND`s into its parts. Conjunct-at-a-time is the + /// unit predicate pushdown works in: `a AND b` may push `a` into one side + /// of a join and `b` into the other. + pub fn split_conjuncts(&self) -> Vec { + match self { + Expr::Binary { + left, + op: BinaryOp::And, + right, + } => { + let mut out = left.split_conjuncts(); + out.extend(right.split_conjuncts()); + out + } + other => vec![other.clone()], + } + } + + /// Rebuilds a conjunction from its parts. + pub fn join_conjuncts(mut parts: Vec) -> Option { + let first = if parts.is_empty() { + return None; + } else { + parts.remove(0) + }; + Some( + parts + .into_iter() + .fold(first, |acc, p| Expr::binary(acc, BinaryOp::And, p)), + ) + } + + /// Every column this expression reads, in the order encountered. + pub fn columns(&self) -> Vec<(Option, String)> { + let mut out = Vec::new(); + self.walk(&mut |e| { + if let Expr::Column { table, name } = e { + out.push((table.clone(), name.clone())); + } + }); + out + } + + pub fn contains_aggregate(&self) -> bool { + let mut found = false; + self.walk(&mut |e| { + if let Expr::Function { name, .. } = e { + if is_aggregate_name(name) { + found = true; + } + } + }); + found + } + + /// Pre-order walk over every subexpression, including this one. + pub fn walk(&self, f: &mut impl FnMut(&Expr)) { + f(self); + match self { + Expr::Column { .. } | Expr::Literal(_) | Expr::Wildcard => {} + Expr::Binary { left, right, .. } => { + left.walk(f); + right.walk(f); + } + Expr::Unary { expr, .. } | Expr::Cast { expr, .. } | Expr::IsNull { expr, .. } => { + expr.walk(f) + } + Expr::Function { args, .. } => args.iter().for_each(|a| a.walk(f)), + Expr::Between { + expr, low, high, .. + } => { + expr.walk(f); + low.walk(f); + high.walk(f); + } + Expr::InList { expr, list, .. } => { + expr.walk(f); + list.iter().for_each(|e| e.walk(f)); + } + Expr::Like { expr, pattern, .. } => { + expr.walk(f); + pattern.walk(f); + } + Expr::Case { + operand, + branches, + else_result, + } => { + if let Some(o) = operand { + o.walk(f); + } + for (w, t) in branches { + w.walk(f); + t.walk(f); + } + if let Some(e) = else_result { + e.walk(f); + } + } + } + } +} + +pub fn is_aggregate_name(name: &str) -> bool { + matches!( + name.to_ascii_lowercase().as_str(), + "count" | "sum" | "min" | "max" | "avg" + ) +} + +impl fmt::Display for Expr { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Expr::Column { table: None, name } => f.write_str(name), + Expr::Column { + table: Some(t), + name, + } => write!(f, "{t}.{name}"), + Expr::Literal(Value::Utf8(s)) => write!(f, "'{s}'"), + Expr::Literal(v) => write!(f, "{v}"), + Expr::Binary { left, op, right } => write!(f, "({left} {op} {right})"), + Expr::Unary { op, expr } => write!(f, "{op}{expr}"), + Expr::Function { + name, + args, + distinct, + } => { + let inner: Vec = args.iter().map(Expr::to_string).collect(); + let d = if *distinct { "DISTINCT " } else { "" }; + write!(f, "{name}({d}{})", inner.join(", ")) + } + Expr::Cast { expr, data_type } => write!(f, "CAST({expr} AS {data_type})"), + Expr::IsNull { expr, negated } => { + let not = if *negated { "NOT " } else { "" }; + write!(f, "{expr} IS {not}NULL") + } + Expr::Between { + expr, + low, + high, + negated, + } => { + let not = if *negated { "NOT " } else { "" }; + write!(f, "{expr} {not}BETWEEN {low} AND {high}") + } + Expr::InList { + expr, + list, + negated, + } => { + let not = if *negated { "NOT " } else { "" }; + let items: Vec = list.iter().map(Expr::to_string).collect(); + write!(f, "{expr} {not}IN ({})", items.join(", ")) + } + Expr::Like { + expr, + pattern, + negated, + } => { + let not = if *negated { "NOT " } else { "" }; + write!(f, "{expr} {not}LIKE {pattern}") + } + Expr::Case { + operand, + branches, + else_result, + } => { + write!(f, "CASE")?; + if let Some(o) = operand { + write!(f, " {o}")?; + } + for (w, t) in branches { + write!(f, " WHEN {w} THEN {t}")?; + } + if let Some(e) = else_result { + write!(f, " ELSE {e}")?; + } + f.write_str(" END") + } + Expr::Wildcard => f.write_str("*"), + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub enum SelectItem { + /// `*` + Wildcard, + /// `t.*` + QualifiedWildcard(String), + Expr { + expr: Expr, + alias: Option, + }, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum JoinType { + Inner, + Left, + Right, + Full, + Cross, +} + +impl fmt::Display for JoinType { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let s = match self { + JoinType::Inner => "INNER", + JoinType::Left => "LEFT", + JoinType::Right => "RIGHT", + JoinType::Full => "FULL", + JoinType::Cross => "CROSS", + }; + f.write_str(s) + } +} + +#[derive(Debug, Clone, PartialEq)] +pub enum TableRef { + Table { + name: String, + alias: Option, + }, + Join { + left: Box, + right: Box, + join_type: JoinType, + on: Option, + }, +} + +impl TableRef { + /// The names a column can be qualified with in this FROM clause — the + /// alias when there is one, otherwise the table name. + pub fn visible_names(&self) -> Vec { + match self { + TableRef::Table { name, alias } => vec![alias.clone().unwrap_or_else(|| name.clone())], + TableRef::Join { left, right, .. } => { + let mut v = left.visible_names(); + v.extend(right.visible_names()); + v + } + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct OrderByExpr { + pub expr: Expr, + pub ascending: bool, + /// `NULLS FIRST` / `NULLS LAST`. Defaults to SQL's convention: nulls last + /// when ascending, first when descending. + pub nulls_first: bool, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct Query { + pub distinct: bool, + pub projection: Vec, + pub from: Option, + pub selection: Option, + pub group_by: Vec, + pub having: Option, + pub order_by: Vec, + pub limit: Option, + pub offset: Option, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct ColumnDef { + pub name: String, + pub data_type: DataType, + pub nullable: bool, +} + +#[derive(Debug, Clone, PartialEq)] +pub enum Statement { + /// Boxed because a `Query` is an order of magnitude larger than any other + /// statement, and every `Statement` would otherwise pay for it. + Query(Box), + CreateTable { + name: String, + columns: Vec, + }, + Insert { + table: String, + rows: Vec>, + }, + /// `COPY t FROM 'path.csv'` — bulk load, the only way large data gets in. + Copy { + table: String, + path: String, + }, + DropTable { + name: String, + }, + Explain { + analyze: bool, + query: Box, + }, +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn conjunctions_split_and_rejoin() { + let e = Expr::binary( + Expr::binary(Expr::column("a"), BinaryOp::And, Expr::column("b")), + BinaryOp::And, + Expr::column("c"), + ); + let parts = e.split_conjuncts(); + assert_eq!(parts.len(), 3); + assert_eq!( + Expr::join_conjuncts(parts).unwrap().to_string(), + e.to_string() + ); + } + + #[test] + fn a_non_conjunction_splits_into_one_part() { + let e = Expr::binary(Expr::column("a"), BinaryOp::Or, Expr::column("b")); + assert_eq!(e.split_conjuncts().len(), 1); + assert!(Expr::join_conjuncts(vec![]).is_none()); + } + + #[test] + fn comparisons_mirror_when_their_operands_swap() { + assert_eq!(BinaryOp::Lt.swap_operands(), Some(BinaryOp::Gt)); + assert_eq!(BinaryOp::GtEq.swap_operands(), Some(BinaryOp::LtEq)); + assert_eq!(BinaryOp::Eq.swap_operands(), Some(BinaryOp::Eq)); + assert_eq!(BinaryOp::NotEq.swap_operands(), Some(BinaryOp::NotEq)); + assert_eq!(BinaryOp::Plus.swap_operands(), None); + assert_eq!(BinaryOp::And.swap_operands(), None); + } + + #[test] + fn operators_classify_themselves() { + assert!(BinaryOp::Lt.is_comparison()); + assert!(BinaryOp::Multiply.is_arithmetic()); + assert!(BinaryOp::Or.is_logical()); + assert!(!BinaryOp::Or.is_comparison()); + assert!(!BinaryOp::Eq.is_arithmetic()); + } + + #[test] + fn columns_are_collected_from_anywhere_in_an_expression() { + let e = Expr::Case { + operand: None, + branches: vec![( + Expr::binary( + Expr::column("a"), + BinaryOp::Gt, + Expr::Literal(Value::Int64(1)), + ), + Expr::qualified("t", "b"), + )], + else_result: Some(Box::new(Expr::column("c"))), + }; + let cols = e.columns(); + assert_eq!(cols.len(), 3); + assert_eq!(cols[1], (Some("t".to_string()), "b".to_string())); + } + + #[test] + fn aggregates_are_found_however_deeply_nested() { + let e = Expr::binary( + Expr::Literal(Value::Int64(1)), + BinaryOp::Plus, + Expr::Function { + name: "sum".into(), + args: vec![Expr::column("x")], + distinct: false, + }, + ); + assert!(e.contains_aggregate()); + assert!(!Expr::column("x").contains_aggregate()); + assert!(!Expr::Function { + name: "upper".into(), + args: vec![], + distinct: false + } + .contains_aggregate()); + } + + #[test] + fn the_walk_visits_every_node_shape() { + let exprs = vec![ + Expr::Wildcard, + Expr::Unary { + op: UnaryOp::Not, + expr: Box::new(Expr::column("a")), + }, + Expr::Cast { + expr: Box::new(Expr::column("a")), + data_type: DataType::Int64, + }, + Expr::IsNull { + expr: Box::new(Expr::column("a")), + negated: true, + }, + Expr::Between { + expr: Box::new(Expr::column("a")), + low: Box::new(Expr::column("b")), + high: Box::new(Expr::column("c")), + negated: false, + }, + Expr::InList { + expr: Box::new(Expr::column("a")), + list: vec![Expr::column("b")], + negated: false, + }, + Expr::Like { + expr: Box::new(Expr::column("a")), + pattern: Box::new(Expr::Literal(Value::Utf8("x%".into()))), + negated: false, + }, + Expr::Case { + operand: Some(Box::new(Expr::column("a"))), + branches: vec![(Expr::column("b"), Expr::column("c"))], + else_result: None, + }, + ]; + for e in exprs { + let mut seen = 0; + e.walk(&mut |_| seen += 1); + assert!(seen >= 1, "{e} visited nothing"); + } + } + + #[test] + fn expressions_render_back_to_readable_sql() { + assert_eq!( + Expr::binary( + Expr::column("a"), + BinaryOp::Plus, + Expr::Literal(Value::Int64(1)) + ) + .to_string(), + "(a + 1)" + ); + assert_eq!(Expr::qualified("t", "c").to_string(), "t.c"); + assert_eq!(Expr::Literal(Value::Utf8("hi".into())).to_string(), "'hi'"); + assert_eq!( + Expr::Function { + name: "count".into(), + args: vec![Expr::Wildcard], + distinct: false + } + .to_string(), + "count(*)" + ); + assert_eq!( + Expr::Function { + name: "count".into(), + args: vec![Expr::column("x")], + distinct: true + } + .to_string(), + "count(DISTINCT x)" + ); + assert_eq!( + Expr::Cast { + expr: Box::new(Expr::column("a")), + data_type: DataType::Float64 + } + .to_string(), + "CAST(a AS FLOAT64)" + ); + assert_eq!( + Expr::IsNull { + expr: Box::new(Expr::column("a")), + negated: true + } + .to_string(), + "a IS NOT NULL" + ); + assert_eq!( + Expr::Unary { + op: UnaryOp::Neg, + expr: Box::new(Expr::column("a")) + } + .to_string(), + "-a" + ); + } + + #[test] + fn the_remaining_expression_shapes_render_too() { + assert_eq!( + Expr::Between { + expr: Box::new(Expr::column("a")), + low: Box::new(Expr::Literal(Value::Int64(1))), + high: Box::new(Expr::Literal(Value::Int64(9))), + negated: true + } + .to_string(), + "a NOT BETWEEN 1 AND 9" + ); + assert_eq!( + Expr::InList { + expr: Box::new(Expr::column("a")), + list: vec![ + Expr::Literal(Value::Int64(1)), + Expr::Literal(Value::Int64(2)) + ], + negated: false + } + .to_string(), + "a IN (1, 2)" + ); + assert_eq!( + Expr::Like { + expr: Box::new(Expr::column("a")), + pattern: Box::new(Expr::Literal(Value::Utf8("x%".into()))), + negated: true + } + .to_string(), + "a NOT LIKE 'x%'" + ); + assert_eq!( + Expr::Case { + operand: Some(Box::new(Expr::column("a"))), + branches: vec![( + Expr::Literal(Value::Int64(1)), + Expr::Literal(Value::Int64(2)) + )], + else_result: Some(Box::new(Expr::Literal(Value::Int64(3)))) + } + .to_string(), + "CASE a WHEN 1 THEN 2 ELSE 3 END" + ); + assert_eq!(Expr::Wildcard.to_string(), "*"); + } + + #[test] + fn a_from_clause_reports_the_names_columns_may_be_qualified_with() { + let t = TableRef::Join { + left: Box::new(TableRef::Table { + name: "orders".into(), + alias: Some("o".into()), + }), + right: Box::new(TableRef::Table { + name: "customers".into(), + alias: None, + }), + join_type: JoinType::Inner, + on: None, + }; + assert_eq!(t.visible_names(), vec!["o", "customers"]); + } + + #[test] + fn join_types_render_for_plan_output() { + for j in [ + JoinType::Inner, + JoinType::Left, + JoinType::Right, + JoinType::Full, + JoinType::Cross, + ] { + assert!(!j.to_string().is_empty()); + } + assert_eq!(JoinType::Left.to_string(), "LEFT"); + assert_eq!(UnaryOp::Not.to_string(), "NOT "); + } + + #[test] + fn aggregate_names_are_recognised_case_insensitively() { + for n in ["count", "SUM", "Min", "max", "avg"] { + assert!(is_aggregate_name(n)); + } + assert!(!is_aggregate_name("median")); + } +} diff --git a/crates/qf-sql/src/lexer.rs b/crates/qf-sql/src/lexer.rs new file mode 100644 index 0000000..0e86218 --- /dev/null +++ b/crates/qf-sql/src/lexer.rs @@ -0,0 +1,571 @@ +//! Turns SQL text into tokens. +//! +//! Written by hand rather than generated. A SQL lexer is small enough that the +//! interesting decisions — where a keyword stops being a keyword, how a quoted +//! identifier differs from a string literal, what a bare `.` means next to a +//! digit — are worth making explicitly. + +use qf_common::{Error, Result}; +use std::fmt; + +#[derive(Debug, Clone, PartialEq)] +pub enum Token { + /// A bare word that matched no keyword, or a `"quoted identifier"`. + Ident(String), + /// A keyword, stored upper-cased so comparisons need no `eq_ignore_case`. + Keyword(String), + Int(i64), + Float(f64), + /// A `'single quoted'` string literal. + Str(String), + + Plus, + Minus, + Star, + Slash, + Percent, + Eq, + NotEq, + Lt, + LtEq, + Gt, + GtEq, + + Comma, + Dot, + LParen, + RParen, + Semicolon, + + /// End of input. Emitted once so the parser never indexes past the end. + Eof, +} + +impl fmt::Display for Token { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Token::Ident(s) => write!(f, "{s}"), + Token::Keyword(s) => write!(f, "{s}"), + Token::Int(i) => write!(f, "{i}"), + Token::Float(x) => write!(f, "{x}"), + Token::Str(s) => write!(f, "'{s}'"), + Token::Plus => f.write_str("+"), + Token::Minus => f.write_str("-"), + Token::Star => f.write_str("*"), + Token::Slash => f.write_str("/"), + Token::Percent => f.write_str("%"), + Token::Eq => f.write_str("="), + Token::NotEq => f.write_str("<>"), + Token::Lt => f.write_str("<"), + Token::LtEq => f.write_str("<="), + Token::Gt => f.write_str(">"), + Token::GtEq => f.write_str(">="), + Token::Comma => f.write_str(","), + Token::Dot => f.write_str("."), + Token::LParen => f.write_str("("), + Token::RParen => f.write_str(")"), + Token::Semicolon => f.write_str(";"), + Token::Eof => f.write_str("end of input"), + } + } +} + +/// A token plus where it started, so errors can point at the offending text. +#[derive(Debug, Clone, PartialEq)] +pub struct Spanned { + pub token: Token, + pub offset: usize, +} + +const KEYWORDS: &[&str] = &[ + "SELECT", "FROM", "WHERE", "GROUP", "BY", "HAVING", "ORDER", "LIMIT", "OFFSET", "AS", "AND", + "OR", "NOT", "NULL", "IS", "IN", "BETWEEN", "LIKE", "DISTINCT", "JOIN", "INNER", "LEFT", + "RIGHT", "FULL", "OUTER", "CROSS", "ON", "ASC", "DESC", "CASE", "WHEN", "THEN", "ELSE", "END", + "CAST", "TRUE", "FALSE", "CREATE", "TABLE", "INSERT", "INTO", "VALUES", "DROP", "EXPLAIN", + "ANALYZE", "NULLS", "FIRST", "LAST", "COPY", +]; + +pub fn is_keyword(word: &str) -> bool { + let upper = word.to_ascii_uppercase(); + KEYWORDS.contains(&upper.as_str()) +} + +pub struct Lexer<'a> { + src: &'a [u8], + text: &'a str, + pos: usize, +} + +impl<'a> Lexer<'a> { + pub fn new(text: &'a str) -> Self { + Lexer { + src: text.as_bytes(), + text, + pos: 0, + } + } + + /// Tokenises the whole input, always ending with exactly one `Eof`. + pub fn tokenize(mut self) -> Result> { + let mut out = Vec::new(); + loop { + let t = self.next_token()?; + let done = t.token == Token::Eof; + out.push(t); + if done { + return Ok(out); + } + } + } + + fn peek(&self) -> Option { + self.src.get(self.pos).copied() + } + + fn peek_at(&self, n: usize) -> Option { + self.src.get(self.pos + n).copied() + } + + fn skip_trivia(&mut self) { + loop { + match self.peek() { + Some(c) if c.is_ascii_whitespace() => self.pos += 1, + // `-- line comment` + Some(b'-') if self.peek_at(1) == Some(b'-') => { + while let Some(c) = self.peek() { + self.pos += 1; + if c == b'\n' { + break; + } + } + } + _ => return, + } + } + } + + fn next_token(&mut self) -> Result { + self.skip_trivia(); + let offset = self.pos; + let Some(c) = self.peek() else { + return Ok(Spanned { + token: Token::Eof, + offset, + }); + }; + + let token = match c { + b'0'..=b'9' => self.number()?, + b'\'' => self.string()?, + b'"' => self.quoted_ident()?, + c if c == b'_' || c.is_ascii_alphabetic() => self.word(), + b'.' if self.peek_at(1).is_some_and(|d| d.is_ascii_digit()) => self.number()?, + _ => self.operator()?, + }; + Ok(Spanned { token, offset }) + } + + fn word(&mut self) -> Token { + let start = self.pos; + while let Some(c) = self.peek() { + if c == b'_' || c.is_ascii_alphanumeric() { + self.pos += 1; + } else { + break; + } + } + let word = &self.text[start..self.pos]; + if is_keyword(word) { + Token::Keyword(word.to_ascii_uppercase()) + } else { + Token::Ident(word.to_string()) + } + } + + fn number(&mut self) -> Result { + let start = self.pos; + let mut is_float = false; + while let Some(c) = self.peek() { + match c { + b'0'..=b'9' => self.pos += 1, + b'.' if !is_float => { + is_float = true; + self.pos += 1; + } + b'e' | b'E' => { + // Only an exponent if a digit or sign actually follows, + // so `3e` stays a number followed by an identifier. + let next = self.peek_at(1); + let after = self.peek_at(2); + let exponent = matches!(next, Some(d) if d.is_ascii_digit()) + || (matches!(next, Some(b'+' | b'-')) + && matches!(after, Some(d) if d.is_ascii_digit())); + if !exponent { + break; + } + is_float = true; + self.pos += 2; + } + _ => break, + } + } + let text = &self.text[start..self.pos]; + if is_float { + text.parse::() + .map(Token::Float) + .map_err(|_| Error::parse(format!("`{text}` is not a valid number"))) + } else { + match text.parse::() { + Ok(i) => Ok(Token::Int(i)), + // Too big for i64 rather than malformed: keep it as a float + // instead of refusing the query. + Err(_) => text + .parse::() + .map(Token::Float) + .map_err(|_| Error::parse(format!("`{text}` is not a valid number"))), + } + } + } + + fn string(&mut self) -> Result { + let start = self.pos; + self.pos += 1; // opening quote + let mut out = String::new(); + loop { + match self.peek() { + None => { + return Err(Error::parse(format!( + "unterminated string literal starting at offset {start}" + ))) + } + Some(b'\'') => { + // '' inside a literal is an escaped quote. + if self.peek_at(1) == Some(b'\'') { + out.push('\''); + self.pos += 2; + } else { + self.pos += 1; + return Ok(Token::Str(out)); + } + } + Some(_) => { + let ch = self.text[self.pos..].chars().next().unwrap(); + out.push(ch); + self.pos += ch.len_utf8(); + } + } + } + } + + fn quoted_ident(&mut self) -> Result { + let start = self.pos; + self.pos += 1; + let mut out = String::new(); + loop { + match self.peek() { + None => { + return Err(Error::parse(format!( + "unterminated quoted identifier starting at offset {start}" + ))) + } + Some(b'"') => { + if self.peek_at(1) == Some(b'"') { + out.push('"'); + self.pos += 2; + } else { + self.pos += 1; + return Ok(Token::Ident(out)); + } + } + Some(_) => { + let ch = self.text[self.pos..].chars().next().unwrap(); + out.push(ch); + self.pos += ch.len_utf8(); + } + } + } + } + + fn operator(&mut self) -> Result { + let two: Option<&str> = self.text.get(self.pos..self.pos + 2); + if let Some(op) = two { + let token = match op { + "<=" => Some(Token::LtEq), + ">=" => Some(Token::GtEq), + "<>" | "!=" => Some(Token::NotEq), + "||" => None, // reserved for concatenation; not supported yet + _ => None, + }; + if let Some(t) = token { + self.pos += 2; + return Ok(t); + } + } + let c = self.peek().unwrap(); + let token = match c { + b'+' => Token::Plus, + b'-' => Token::Minus, + b'*' => Token::Star, + b'/' => Token::Slash, + b'%' => Token::Percent, + b'=' => Token::Eq, + b'<' => Token::Lt, + b'>' => Token::Gt, + b',' => Token::Comma, + b'.' => Token::Dot, + b'(' => Token::LParen, + b')' => Token::RParen, + b';' => Token::Semicolon, + other => { + return Err(Error::parse(format!( + "unexpected character `{}` at offset {}", + other as char, self.pos + ))) + } + }; + self.pos += 1; + Ok(token) + } +} + +pub fn tokenize(sql: &str) -> Result> { + Lexer::new(sql).tokenize() +} + +#[cfg(test)] +mod tests { + use super::*; + + fn toks(sql: &str) -> Vec { + tokenize(sql) + .unwrap() + .into_iter() + .map(|s| s.token) + .collect() + } + + #[test] + fn a_simple_select_lexes_into_keywords_identifiers_and_punctuation() { + assert_eq!( + toks("SELECT a, b FROM t;"), + vec![ + Token::Keyword("SELECT".into()), + Token::Ident("a".into()), + Token::Comma, + Token::Ident("b".into()), + Token::Keyword("FROM".into()), + Token::Ident("t".into()), + Token::Semicolon, + Token::Eof, + ] + ); + } + + #[test] + fn keywords_are_normalised_to_upper_case_whatever_was_typed() { + assert_eq!( + toks("select Where FROM"), + vec![ + Token::Keyword("SELECT".into()), + Token::Keyword("WHERE".into()), + Token::Keyword("FROM".into()), + Token::Eof, + ] + ); + } + + #[test] + fn identifiers_keep_the_case_they_were_written_in() { + assert_eq!(toks("MyTable")[0], Token::Ident("MyTable".into())); + } + + #[test] + fn a_quoted_identifier_may_hold_a_keyword_or_a_space() { + assert_eq!( + toks(r#""select", "two words", "a""b""#), + vec![ + Token::Ident("select".into()), + Token::Comma, + Token::Ident("two words".into()), + Token::Comma, + Token::Ident("a\"b".into()), + Token::Eof, + ] + ); + } + + #[test] + fn numbers_split_into_integers_and_floats() { + assert_eq!( + toks("1 2.5 .5 1e3 1.5e-2"), + vec![ + Token::Int(1), + Token::Float(2.5), + Token::Float(0.5), + Token::Float(1000.0), + Token::Float(0.015), + Token::Eof, + ] + ); + } + + #[test] + fn a_trailing_e_is_not_swallowed_into_the_number() { + // `3e` is 3 followed by the identifier `e`, not a malformed float. + assert_eq!( + toks("3e"), + vec![Token::Int(3), Token::Ident("e".into()), Token::Eof] + ); + } + + #[test] + fn an_integer_too_large_for_i64_becomes_a_float_rather_than_an_error() { + assert!(matches!(toks("99999999999999999999")[0], Token::Float(_))); + } + + #[test] + fn string_literals_handle_doubled_quotes_and_unicode() { + assert_eq!( + toks("'it''s', 'naïve'"), + vec![ + Token::Str("it's".into()), + Token::Comma, + Token::Str("naïve".into()), + Token::Eof, + ] + ); + } + + #[test] + fn comparison_operators_lex_as_single_tokens() { + assert_eq!( + toks("< <= > >= = <> !="), + vec![ + Token::Lt, + Token::LtEq, + Token::Gt, + Token::GtEq, + Token::Eq, + Token::NotEq, + Token::NotEq, + Token::Eof, + ] + ); + } + + #[test] + fn arithmetic_operators_lex_as_single_tokens() { + assert_eq!( + toks("+ - * / % . ( )"), + vec![ + Token::Plus, + Token::Minus, + Token::Star, + Token::Slash, + Token::Percent, + Token::Dot, + Token::LParen, + Token::RParen, + Token::Eof, + ] + ); + } + + #[test] + fn line_comments_and_whitespace_are_skipped() { + assert_eq!( + toks("SELECT -- everything\n a\n-- trailing comment"), + vec![ + Token::Keyword("SELECT".into()), + Token::Ident("a".into()), + Token::Eof, + ] + ); + } + + #[test] + fn a_minus_that_is_not_a_comment_stays_an_operator() { + assert_eq!( + toks("a - b"), + vec![ + Token::Ident("a".into()), + Token::Minus, + Token::Ident("b".into()), + Token::Eof, + ] + ); + } + + #[test] + fn unterminated_literals_are_reported_with_where_they_started() { + let err = tokenize("SELECT 'oops").unwrap_err(); + assert!(err.to_string().contains("unterminated string")); + assert!(err.to_string().contains("offset 7")); + assert!(tokenize(r#"SELECT "oops"#) + .unwrap_err() + .to_string() + .contains("unterminated quoted identifier")); + } + + #[test] + fn an_unexpected_character_is_reported_with_its_offset() { + let err = tokenize("SELECT a # b").unwrap_err(); + assert!(err.to_string().contains('#')); + assert!(err.to_string().contains("offset 9")); + } + + #[test] + fn the_double_pipe_operator_is_refused_rather_than_silently_split() { + assert!(tokenize("a || b").is_err()); + } + + #[test] + fn spans_point_at_where_each_token_started() { + let spans = tokenize("SELECT a").unwrap(); + assert_eq!(spans[0].offset, 0); + assert_eq!(spans[1].offset, 8); + } + + #[test] + fn empty_input_lexes_to_just_eof() { + assert_eq!(toks(""), vec![Token::Eof]); + assert_eq!(toks(" -- nothing here\n "), vec![Token::Eof]); + } + + #[test] + fn tokens_render_back_to_something_readable_in_errors() { + assert_eq!(Token::Keyword("SELECT".into()).to_string(), "SELECT"); + assert_eq!(Token::Str("x".into()).to_string(), "'x'"); + assert_eq!(Token::NotEq.to_string(), "<>"); + assert_eq!(Token::Eof.to_string(), "end of input"); + assert_eq!(Token::Float(1.5).to_string(), "1.5"); + assert_eq!(Token::Int(2).to_string(), "2"); + assert_eq!(Token::Ident("c".into()).to_string(), "c"); + for t in [ + Token::Plus, + Token::Minus, + Token::Star, + Token::Slash, + Token::Percent, + Token::Eq, + Token::Lt, + Token::LtEq, + Token::Gt, + Token::GtEq, + Token::Comma, + Token::Dot, + Token::LParen, + Token::RParen, + Token::Semicolon, + ] { + assert!(!t.to_string().is_empty()); + } + } + + #[test] + fn the_keyword_list_is_queryable() { + assert!(is_keyword("select")); + assert!(is_keyword("JOIN")); + assert!(!is_keyword("customers")); + } +} diff --git a/crates/qf-sql/src/lib.rs b/crates/qf-sql/src/lib.rs index ff7bd09..2a5ca06 100644 --- a/crates/qf-sql/src/lib.rs +++ b/crates/qf-sql/src/lib.rs @@ -1 +1,6 @@ -// placeholder +pub mod ast; +pub mod lexer; +pub mod parser; + +pub use ast::*; +pub use parser::{parse, Parser}; diff --git a/crates/qf-sql/src/parser.rs b/crates/qf-sql/src/parser.rs new file mode 100644 index 0000000..df65e17 --- /dev/null +++ b/crates/qf-sql/src/parser.rs @@ -0,0 +1,1080 @@ +//! Recursive descent for statements, precedence climbing for expressions. +//! +//! The whole grammar is small enough to read top to bottom, which is the point: +//! when a query parses into something surprising, the rule that did it is a few +//! screens away rather than inside a generated table. + +use crate::ast::*; +use crate::lexer::{tokenize, Spanned, Token}; +use qf_common::{DataType, Error, Result, Value}; + +/// Binding power, loosest first. `NOT` binds tighter than `AND` and looser +/// than any comparison, so `NOT a = b` is `NOT (a = b)`. +const PREC_OR: u8 = 1; +const PREC_AND: u8 = 2; +const PREC_NOT: u8 = 3; +const PREC_COMPARE: u8 = 4; +const PREC_ADD: u8 = 5; +const PREC_MUL: u8 = 6; + +pub struct Parser { + tokens: Vec, + pos: usize, +} + +impl Parser { + pub fn new(sql: &str) -> Result { + Ok(Parser { + tokens: tokenize(sql)?, + pos: 0, + }) + } + + /// Parses every statement in the input, separated by semicolons. + pub fn parse_statements(sql: &str) -> Result> { + let mut p = Parser::new(sql)?; + let mut out = Vec::new(); + loop { + while p.eat(&Token::Semicolon) {} + if p.peek() == &Token::Eof { + return Ok(out); + } + out.push(p.parse_statement()?); + if !matches!(p.peek(), Token::Semicolon | Token::Eof) { + return Err(p.unexpected("a semicolon or the end of the statement")); + } + } + } + + /// Parses exactly one statement and refuses trailing input. + pub fn parse_one(sql: &str) -> Result { + let mut stmts = Parser::parse_statements(sql)?; + match stmts.len() { + 0 => Err(Error::parse("no statement to run")), + 1 => Ok(stmts.remove(0)), + n => Err(Error::parse(format!( + "expected a single statement, found {n}" + ))), + } + } + + fn peek(&self) -> &Token { + &self.tokens[self.pos].token + } + + fn peek_at(&self, n: usize) -> &Token { + let i = (self.pos + n).min(self.tokens.len() - 1); + &self.tokens[i].token + } + + fn next(&mut self) -> Token { + let t = self.tokens[self.pos].token.clone(); + if self.pos < self.tokens.len() - 1 { + self.pos += 1; + } + t + } + + fn eat(&mut self, want: &Token) -> bool { + if self.peek() == want { + self.next(); + true + } else { + false + } + } + + fn eat_keyword(&mut self, word: &str) -> bool { + if matches!(self.peek(), Token::Keyword(k) if k == word) { + self.next(); + true + } else { + false + } + } + + fn peek_keyword(&self, word: &str) -> bool { + matches!(self.peek(), Token::Keyword(k) if k == word) + } + + fn expect(&mut self, want: &Token) -> Result<()> { + if self.eat(want) { + Ok(()) + } else { + Err(self.unexpected(&format!("`{want}`"))) + } + } + + fn expect_keyword(&mut self, word: &str) -> Result<()> { + if self.eat_keyword(word) { + Ok(()) + } else { + Err(self.unexpected(&format!("`{word}`"))) + } + } + + fn unexpected(&self, wanted: &str) -> Error { + let s = &self.tokens[self.pos]; + Error::parse(format!( + "expected {wanted}, found `{}` at offset {}", + s.token, s.offset + )) + } + + /// An identifier, or a non-reserved keyword being used as a name. + fn identifier(&mut self) -> Result { + match self.peek().clone() { + Token::Ident(name) => { + self.next(); + Ok(name) + } + _ => Err(self.unexpected("an identifier")), + } + } + + // ---- statements ---- + + fn parse_statement(&mut self) -> Result { + match self.peek().clone() { + Token::Keyword(k) => match k.as_str() { + "SELECT" => Ok(Statement::Query(Box::new(self.parse_query()?))), + "CREATE" => self.parse_create_table(), + "INSERT" => self.parse_insert(), + "COPY" => self.parse_copy(), + "DROP" => self.parse_drop(), + "EXPLAIN" => self.parse_explain(), + _ => Err(self.unexpected("a statement")), + }, + _ => Err(self.unexpected("a statement")), + } + } + + fn parse_create_table(&mut self) -> Result { + self.expect_keyword("CREATE")?; + self.expect_keyword("TABLE")?; + let name = self.identifier()?; + self.expect(&Token::LParen)?; + let mut columns = Vec::new(); + loop { + let col = self.identifier()?; + let type_name = match self.peek().clone() { + Token::Ident(t) => { + self.next(); + t + } + _ => return Err(self.unexpected("a column type")), + }; + let data_type = DataType::from_sql_name(&type_name)?; + // `NOT NULL` marks the column non-nullable; everything else is. + let nullable = !(self.eat_keyword("NOT") && { + self.expect_keyword("NULL")?; + true + }); + columns.push(ColumnDef { + name: col, + data_type, + nullable, + }); + if !self.eat(&Token::Comma) { + break; + } + } + self.expect(&Token::RParen)?; + if columns.is_empty() { + return Err(Error::parse("a table needs at least one column")); + } + Ok(Statement::CreateTable { name, columns }) + } + + fn parse_insert(&mut self) -> Result { + self.expect_keyword("INSERT")?; + self.expect_keyword("INTO")?; + let table = self.identifier()?; + self.expect_keyword("VALUES")?; + let mut rows = Vec::new(); + loop { + self.expect(&Token::LParen)?; + let mut row = Vec::new(); + loop { + row.push(self.parse_expr(0)?); + if !self.eat(&Token::Comma) { + break; + } + } + self.expect(&Token::RParen)?; + rows.push(row); + if !self.eat(&Token::Comma) { + break; + } + } + Ok(Statement::Insert { table, rows }) + } + + fn parse_copy(&mut self) -> Result { + self.expect_keyword("COPY")?; + let table = self.identifier()?; + self.expect_keyword("FROM")?; + match self.next() { + Token::Str(path) => Ok(Statement::Copy { table, path }), + _ => Err(Error::parse( + "COPY needs a quoted file path, as in COPY t FROM 'data.csv'", + )), + } + } + + fn parse_drop(&mut self) -> Result { + self.expect_keyword("DROP")?; + self.expect_keyword("TABLE")?; + Ok(Statement::DropTable { + name: self.identifier()?, + }) + } + + fn parse_explain(&mut self) -> Result { + self.expect_keyword("EXPLAIN")?; + let analyze = self.eat_keyword("ANALYZE"); + Ok(Statement::Explain { + analyze, + query: Box::new(self.parse_query()?), + }) + } + + // ---- queries ---- + + fn parse_query(&mut self) -> Result { + self.expect_keyword("SELECT")?; + let distinct = self.eat_keyword("DISTINCT"); + let projection = self.parse_projection()?; + + let from = if self.eat_keyword("FROM") { + Some(self.parse_table_ref()?) + } else { + None + }; + + let selection = if self.eat_keyword("WHERE") { + Some(self.parse_expr(0)?) + } else { + None + }; + + let mut group_by = Vec::new(); + if self.eat_keyword("GROUP") { + self.expect_keyword("BY")?; + loop { + group_by.push(self.parse_expr(0)?); + if !self.eat(&Token::Comma) { + break; + } + } + } + + let having = if self.eat_keyword("HAVING") { + Some(self.parse_expr(0)?) + } else { + None + }; + + let mut order_by = Vec::new(); + if self.eat_keyword("ORDER") { + self.expect_keyword("BY")?; + loop { + let expr = self.parse_expr(0)?; + let ascending = if self.eat_keyword("DESC") { + false + } else { + self.eat_keyword("ASC"); + true + }; + // Postgres' default: nulls sort last ascending, first + // descending, so `ORDER BY x DESC` puts real values on top. + let mut nulls_first = !ascending; + if self.eat_keyword("NULLS") { + if self.eat_keyword("FIRST") { + nulls_first = true; + } else { + self.expect_keyword("LAST")?; + nulls_first = false; + } + } + order_by.push(OrderByExpr { + expr, + ascending, + nulls_first, + }); + if !self.eat(&Token::Comma) { + break; + } + } + } + + let limit = if self.eat_keyword("LIMIT") { + Some(self.parse_count("LIMIT")?) + } else { + None + }; + let offset = if self.eat_keyword("OFFSET") { + Some(self.parse_count("OFFSET")?) + } else { + None + }; + + if having.is_some() && group_by.is_empty() { + // A bare HAVING is legal in some dialects over the implicit single + // group, but it is far more often a WHERE that was mistyped. + let projects_aggregate = projection.iter().any(|item| match item { + SelectItem::Expr { expr, .. } => expr.contains_aggregate(), + _ => false, + }); + if !projects_aggregate { + return Err(Error::parse( + "HAVING without GROUP BY and without an aggregate — did you mean WHERE?", + )); + } + } + + Ok(Query { + distinct, + projection, + from, + selection, + group_by, + having, + order_by, + limit, + offset, + }) + } + + fn parse_count(&mut self, clause: &str) -> Result { + match self.next() { + Token::Int(i) if i >= 0 => Ok(i as usize), + other => Err(Error::parse(format!( + "{clause} needs a non-negative integer, found `{other}`" + ))), + } + } + + fn parse_projection(&mut self) -> Result> { + let mut items = Vec::new(); + loop { + items.push(self.parse_select_item()?); + if !self.eat(&Token::Comma) { + break; + } + } + Ok(items) + } + + fn parse_select_item(&mut self) -> Result { + if self.eat(&Token::Star) { + return Ok(SelectItem::Wildcard); + } + // `t.*` + if let (Token::Ident(name), Token::Dot, Token::Star) = ( + self.peek().clone(), + self.peek_at(1).clone(), + self.peek_at(2), + ) { + self.next(); + self.next(); + self.next(); + return Ok(SelectItem::QualifiedWildcard(name)); + } + let expr = self.parse_expr(0)?; + let alias = if self.eat_keyword("AS") { + Some(self.identifier()?) + } else if let Token::Ident(name) = self.peek().clone() { + // A bare identifier straight after an expression is an alias. + self.next(); + Some(name) + } else { + None + }; + Ok(SelectItem::Expr { expr, alias }) + } + + fn parse_table_ref(&mut self) -> Result { + let mut left = self.parse_table_factor()?; + loop { + let join_type = if self.eat_keyword("CROSS") { + self.expect_keyword("JOIN")?; + JoinType::Cross + } else if self.eat_keyword("INNER") { + self.expect_keyword("JOIN")?; + JoinType::Inner + } else if self.eat_keyword("LEFT") { + self.eat_keyword("OUTER"); + self.expect_keyword("JOIN")?; + JoinType::Left + } else if self.eat_keyword("RIGHT") { + self.eat_keyword("OUTER"); + self.expect_keyword("JOIN")?; + JoinType::Right + } else if self.eat_keyword("FULL") { + self.eat_keyword("OUTER"); + self.expect_keyword("JOIN")?; + JoinType::Full + } else if self.eat_keyword("JOIN") { + JoinType::Inner + } else if self.eat(&Token::Comma) { + // `FROM a, b` is a cross join; the WHERE clause usually turns + // it back into an inner join during optimisation. + JoinType::Cross + } else { + return Ok(left); + }; + + let right = self.parse_table_factor()?; + let on = if self.eat_keyword("ON") { + Some(self.parse_expr(0)?) + } else { + None + }; + if on.is_none() && join_type != JoinType::Cross { + return Err(Error::parse(format!("{join_type} JOIN needs an ON clause"))); + } + left = TableRef::Join { + left: Box::new(left), + right: Box::new(right), + join_type, + on, + }; + } + } + + fn parse_table_factor(&mut self) -> Result { + if self.eat(&Token::LParen) { + let inner = self.parse_table_ref()?; + self.expect(&Token::RParen)?; + return Ok(inner); + } + let name = self.identifier()?; + let alias = if self.eat_keyword("AS") { + Some(self.identifier()?) + } else if let Token::Ident(a) = self.peek().clone() { + self.next(); + Some(a) + } else { + None + }; + Ok(TableRef::Table { name, alias }) + } + + // ---- expressions ---- + + fn parse_expr(&mut self, min_prec: u8) -> Result { + let mut left = self.parse_prefix()?; + loop { + // Postfix forms bind at comparison precedence. + if PREC_COMPARE >= min_prec { + if let Some(e) = self.try_parse_postfix(&left)? { + left = e; + continue; + } + } + let Some((op, prec)) = self.peek_binary_op() else { + return Ok(left); + }; + if prec < min_prec { + return Ok(left); + } + self.next(); + // Left-associative: the right operand must bind more tightly. + let right = self.parse_expr(prec + 1)?; + left = Expr::binary(left, op, right); + } + } + + fn peek_binary_op(&self) -> Option<(BinaryOp, u8)> { + Some(match self.peek() { + Token::Keyword(k) if k == "OR" => (BinaryOp::Or, PREC_OR), + Token::Keyword(k) if k == "AND" => (BinaryOp::And, PREC_AND), + Token::Eq => (BinaryOp::Eq, PREC_COMPARE), + Token::NotEq => (BinaryOp::NotEq, PREC_COMPARE), + Token::Lt => (BinaryOp::Lt, PREC_COMPARE), + Token::LtEq => (BinaryOp::LtEq, PREC_COMPARE), + Token::Gt => (BinaryOp::Gt, PREC_COMPARE), + Token::GtEq => (BinaryOp::GtEq, PREC_COMPARE), + Token::Plus => (BinaryOp::Plus, PREC_ADD), + Token::Minus => (BinaryOp::Minus, PREC_ADD), + Token::Star => (BinaryOp::Multiply, PREC_MUL), + Token::Slash => (BinaryOp::Divide, PREC_MUL), + Token::Percent => (BinaryOp::Modulo, PREC_MUL), + _ => return None, + }) + } + + /// `IS [NOT] NULL`, `[NOT] BETWEEN`, `[NOT] IN`, `[NOT] LIKE`. + fn try_parse_postfix(&mut self, left: &Expr) -> Result> { + if self.eat_keyword("IS") { + let negated = self.eat_keyword("NOT"); + self.expect_keyword("NULL")?; + return Ok(Some(Expr::IsNull { + expr: Box::new(left.clone()), + negated, + })); + } + + // A `NOT` here belongs to the postfix form that follows it. + let negated = if self.peek_keyword("NOT") + && matches!(self.peek_at(1), Token::Keyword(k) if k == "BETWEEN" || k == "IN" || k == "LIKE") + { + self.next(); + true + } else { + false + }; + + if self.eat_keyword("BETWEEN") { + // BETWEEN's bounds must not swallow the AND that separates them. + let low = self.parse_expr(PREC_COMPARE)?; + self.expect_keyword("AND")?; + let high = self.parse_expr(PREC_COMPARE)?; + return Ok(Some(Expr::Between { + expr: Box::new(left.clone()), + low: Box::new(low), + high: Box::new(high), + negated, + })); + } + + if self.eat_keyword("IN") { + self.expect(&Token::LParen)?; + let mut list = Vec::new(); + loop { + list.push(self.parse_expr(0)?); + if !self.eat(&Token::Comma) { + break; + } + } + self.expect(&Token::RParen)?; + return Ok(Some(Expr::InList { + expr: Box::new(left.clone()), + list, + negated, + })); + } + + if self.eat_keyword("LIKE") { + let pattern = self.parse_expr(PREC_COMPARE + 1)?; + return Ok(Some(Expr::Like { + expr: Box::new(left.clone()), + pattern: Box::new(pattern), + negated, + })); + } + + if negated { + return Err(self.unexpected("BETWEEN, IN or LIKE after NOT")); + } + Ok(None) + } + + fn parse_prefix(&mut self) -> Result { + match self.peek().clone() { + Token::Minus => { + self.next(); + Ok(Expr::Unary { + op: UnaryOp::Neg, + expr: Box::new(self.parse_expr(PREC_MUL + 1)?), + }) + } + Token::Plus => { + self.next(); + self.parse_expr(PREC_MUL + 1) + } + Token::Keyword(k) if k == "NOT" => { + self.next(); + Ok(Expr::Unary { + op: UnaryOp::Not, + expr: Box::new(self.parse_expr(PREC_NOT)?), + }) + } + _ => self.parse_primary(), + } + } + + fn parse_primary(&mut self) -> Result { + match self.peek().clone() { + Token::Int(i) => { + self.next(); + Ok(Expr::Literal(Value::Int64(i))) + } + Token::Float(x) => { + self.next(); + Ok(Expr::Literal(Value::Float64(x))) + } + Token::Str(s) => { + self.next(); + Ok(Expr::Literal(Value::Utf8(s))) + } + Token::Star => { + self.next(); + Ok(Expr::Wildcard) + } + Token::LParen => { + self.next(); + let e = self.parse_expr(0)?; + self.expect(&Token::RParen)?; + Ok(e) + } + Token::Keyword(k) => match k.as_str() { + "TRUE" => { + self.next(); + Ok(Expr::Literal(Value::Boolean(true))) + } + "FALSE" => { + self.next(); + Ok(Expr::Literal(Value::Boolean(false))) + } + "NULL" => { + self.next(); + Ok(Expr::Literal(Value::Null)) + } + "CAST" => self.parse_cast(), + "CASE" => self.parse_case(), + _ => Err(self.unexpected("an expression")), + }, + Token::Ident(name) => { + self.next(); + if self.eat(&Token::LParen) { + return self.parse_function_args(name); + } + if self.eat(&Token::Dot) { + let col = self.identifier()?; + return Ok(Expr::qualified(name, col)); + } + Ok(Expr::column(name)) + } + _ => Err(self.unexpected("an expression")), + } + } + + fn parse_function_args(&mut self, name: String) -> Result { + let distinct = self.eat_keyword("DISTINCT"); + let mut args = Vec::new(); + if !self.eat(&Token::RParen) { + loop { + args.push(self.parse_expr(0)?); + if !self.eat(&Token::Comma) { + break; + } + } + self.expect(&Token::RParen)?; + } + if distinct && args.len() != 1 { + return Err(Error::parse(format!( + "{name}(DISTINCT ...) takes exactly one argument" + ))); + } + Ok(Expr::Function { + name, + args, + distinct, + }) + } + + fn parse_cast(&mut self) -> Result { + self.expect_keyword("CAST")?; + self.expect(&Token::LParen)?; + let expr = self.parse_expr(0)?; + self.expect_keyword("AS")?; + let type_name = match self.next() { + Token::Ident(t) => t, + other => { + return Err(Error::parse(format!( + "expected a type name in CAST, found `{other}`" + ))) + } + }; + let data_type = DataType::from_sql_name(&type_name)?; + self.expect(&Token::RParen)?; + Ok(Expr::Cast { + expr: Box::new(expr), + data_type, + }) + } + + fn parse_case(&mut self) -> Result { + self.expect_keyword("CASE")?; + let operand = if self.peek_keyword("WHEN") { + None + } else { + Some(Box::new(self.parse_expr(0)?)) + }; + let mut branches = Vec::new(); + while self.eat_keyword("WHEN") { + let when = self.parse_expr(0)?; + self.expect_keyword("THEN")?; + let then = self.parse_expr(0)?; + branches.push((when, then)); + } + if branches.is_empty() { + return Err(Error::parse("CASE needs at least one WHEN branch")); + } + let else_result = if self.eat_keyword("ELSE") { + Some(Box::new(self.parse_expr(0)?)) + } else { + None + }; + self.expect_keyword("END")?; + Ok(Expr::Case { + operand, + branches, + else_result, + }) + } +} + +/// Parses one statement from `sql`. +pub fn parse(sql: &str) -> Result { + Parser::parse_one(sql) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn query(sql: &str) -> Query { + match parse(sql).unwrap() { + Statement::Query(q) => *q, + other => panic!("expected a query, got {other:?}"), + } + } + + fn expr(sql: &str) -> String { + let q = query(&format!("SELECT {sql}")); + match &q.projection[0] { + SelectItem::Expr { expr, .. } => expr.to_string(), + other => panic!("expected an expression, got {other:?}"), + } + } + + #[test] + fn a_minimal_select_parses() { + let q = query("SELECT a, b FROM t"); + assert_eq!(q.projection.len(), 2); + assert!(matches!( + q.from, + Some(TableRef::Table { ref name, .. }) if name == "t" + )); + assert!(q.selection.is_none()); + assert!(!q.distinct); + } + + #[test] + fn multiplication_binds_tighter_than_addition() { + assert_eq!(expr("1 + 2 * 3"), "(1 + (2 * 3))"); + assert_eq!(expr("(1 + 2) * 3"), "((1 + 2) * 3)"); + } + + #[test] + fn arithmetic_is_left_associative() { + assert_eq!(expr("1 - 2 - 3"), "((1 - 2) - 3)"); + assert_eq!(expr("8 / 4 / 2"), "((8 / 4) / 2)"); + } + + #[test] + fn and_binds_tighter_than_or() { + assert_eq!(expr("a OR b AND c"), "(a OR (b AND c))"); + assert_eq!(expr("(a OR b) AND c"), "((a OR b) AND c)"); + } + + #[test] + fn comparison_binds_tighter_than_and() { + assert_eq!(expr("a = 1 AND b = 2"), "((a = 1) AND (b = 2))"); + } + + #[test] + fn not_binds_looser_than_comparison_and_tighter_than_and() { + assert_eq!(expr("NOT a = 1"), "NOT (a = 1)"); + assert_eq!(expr("NOT a AND b"), "(NOT a AND b)"); + } + + #[test] + fn unary_minus_binds_tighter_than_multiplication() { + assert_eq!(expr("-a * b"), "(-a * b)"); + assert_eq!(expr("- 2 + 3"), "(-2 + 3)"); + assert_eq!(expr("+5"), "5"); + } + + #[test] + fn between_does_not_swallow_the_and_that_separates_its_bounds() { + assert_eq!(expr("a BETWEEN 1 AND 5"), "a BETWEEN 1 AND 5"); + // The trailing AND belongs to the outer conjunction, not to BETWEEN. + assert_eq!( + expr("a BETWEEN 1 AND 5 AND b = 2"), + "(a BETWEEN 1 AND 5 AND (b = 2))" + ); + assert_eq!(expr("a NOT BETWEEN 1 AND 5"), "a NOT BETWEEN 1 AND 5"); + } + + #[test] + fn the_postfix_forms_parse_with_and_without_not() { + assert_eq!(expr("a IS NULL"), "a IS NULL"); + assert_eq!(expr("a IS NOT NULL"), "a IS NOT NULL"); + assert_eq!(expr("a IN (1, 2)"), "a IN (1, 2)"); + assert_eq!(expr("a NOT IN (1)"), "a NOT IN (1)"); + assert_eq!(expr("a LIKE 'x%'"), "a LIKE 'x%'"); + assert_eq!(expr("a NOT LIKE 'x%'"), "a NOT LIKE 'x%'"); + } + + #[test] + fn postfix_forms_chain_with_comparisons() { + assert_eq!(expr("a IS NULL AND b IN (1)"), "(a IS NULL AND b IN (1))"); + } + + #[test] + fn a_dangling_not_before_something_else_is_reported() { + assert!(parse("SELECT a NOT b FROM t").is_err()); + } + + #[test] + fn functions_parse_with_wildcards_distinct_and_nesting() { + assert_eq!(expr("count(*)"), "count(*)"); + assert_eq!(expr("count(DISTINCT a)"), "count(DISTINCT a)"); + assert_eq!(expr("sum(a + b)"), "sum((a + b))"); + assert_eq!(expr("coalesce()"), "coalesce()"); + } + + #[test] + fn distinct_with_several_arguments_is_refused() { + assert!(parse("SELECT count(DISTINCT a, b) FROM t").is_err()); + } + + #[test] + fn cast_and_case_parse() { + assert_eq!(expr("CAST(a AS INT)"), "CAST(a AS INT64)"); + assert_eq!( + expr("CASE WHEN a > 1 THEN 'big' ELSE 'small' END"), + "CASE WHEN (a > 1) THEN 'big' ELSE 'small' END" + ); + assert_eq!( + expr("CASE a WHEN 1 THEN 'one' END"), + "CASE a WHEN 1 THEN 'one' END" + ); + } + + #[test] + fn a_case_without_branches_or_a_bad_cast_type_is_refused() { + assert!(parse("SELECT CASE END FROM t").is_err()); + assert!(parse("SELECT CAST(a AS blob) FROM t").is_err()); + assert!(parse("SELECT CAST(a AS 'INT') FROM t").is_err()); + } + + #[test] + fn qualified_columns_and_wildcards_parse() { + let q = query("SELECT t.*, o.id, * FROM t"); + assert!(matches!(q.projection[0], SelectItem::QualifiedWildcard(ref n) if n == "t")); + assert!(matches!(q.projection[2], SelectItem::Wildcard)); + match &q.projection[1] { + SelectItem::Expr { expr, .. } => assert_eq!(expr.to_string(), "o.id"), + other => panic!("unexpected {other:?}"), + } + } + + #[test] + fn aliases_parse_with_and_without_as() { + let q = query("SELECT a AS x, b y FROM t"); + match (&q.projection[0], &q.projection[1]) { + (SelectItem::Expr { alias: Some(a), .. }, SelectItem::Expr { alias: Some(b), .. }) => { + assert_eq!(a, "x"); + assert_eq!(b, "y"); + } + other => panic!("unexpected {other:?}"), + } + } + + #[test] + fn every_join_flavour_parses() { + for (sql, want) in [ + ("SELECT * FROM a JOIN b ON a.i = b.i", JoinType::Inner), + ("SELECT * FROM a INNER JOIN b ON a.i = b.i", JoinType::Inner), + ("SELECT * FROM a LEFT JOIN b ON a.i = b.i", JoinType::Left), + ( + "SELECT * FROM a LEFT OUTER JOIN b ON a.i = b.i", + JoinType::Left, + ), + ("SELECT * FROM a RIGHT JOIN b ON a.i = b.i", JoinType::Right), + ("SELECT * FROM a FULL JOIN b ON a.i = b.i", JoinType::Full), + ("SELECT * FROM a CROSS JOIN b", JoinType::Cross), + ("SELECT * FROM a, b", JoinType::Cross), + ] { + match query(sql).from.unwrap() { + TableRef::Join { join_type, .. } => assert_eq!(join_type, want, "{sql}"), + other => panic!("{sql} produced {other:?}"), + } + } + } + + #[test] + fn an_inner_join_without_on_is_refused() { + assert!(parse("SELECT * FROM a JOIN b") + .unwrap_err() + .to_string() + .contains("needs an ON clause")); + } + + #[test] + fn joins_are_left_deep_and_table_aliases_parse() { + let q = query("SELECT * FROM a x JOIN b AS y ON x.i = y.i JOIN c ON y.i = c.i"); + match q.from.unwrap() { + TableRef::Join { left, .. } => { + assert!(matches!(*left, TableRef::Join { .. })); + } + other => panic!("unexpected {other:?}"), + } + } + + #[test] + fn parenthesised_from_clauses_parse() { + let q = query("SELECT * FROM (a JOIN b ON a.i = b.i)"); + assert!(matches!(q.from, Some(TableRef::Join { .. }))); + } + + #[test] + fn the_full_clause_list_parses_in_order() { + let q = query( + "SELECT region, count(*) c FROM sales WHERE amount > 10 \ + GROUP BY region HAVING count(*) > 2 ORDER BY c DESC LIMIT 5 OFFSET 10", + ); + assert_eq!(q.group_by.len(), 1); + assert!(q.having.is_some()); + assert_eq!(q.order_by.len(), 1); + assert!(!q.order_by[0].ascending); + assert_eq!(q.limit, Some(5)); + assert_eq!(q.offset, Some(10)); + assert!(q.selection.is_some()); + } + + #[test] + fn order_by_defaults_nulls_last_ascending_and_first_descending() { + assert!(!query("SELECT a FROM t ORDER BY a").order_by[0].nulls_first); + assert!(query("SELECT a FROM t ORDER BY a DESC").order_by[0].nulls_first); + assert!(query("SELECT a FROM t ORDER BY a NULLS FIRST").order_by[0].nulls_first); + assert!(!query("SELECT a FROM t ORDER BY a DESC NULLS LAST").order_by[0].nulls_first); + } + + #[test] + fn a_having_that_should_have_been_a_where_is_refused() { + assert!(parse("SELECT a FROM t HAVING a > 1") + .unwrap_err() + .to_string() + .contains("did you mean WHERE")); + // ...but an aggregate makes a bare HAVING meaningful. + assert!(parse("SELECT count(*) FROM t HAVING count(*) > 1").is_ok()); + } + + #[test] + fn a_negative_limit_is_refused() { + assert!(parse("SELECT a FROM t LIMIT -1").is_err()); + assert!(parse("SELECT a FROM t LIMIT 'x'").is_err()); + } + + #[test] + fn select_without_from_parses() { + let q = query("SELECT 1 + 1"); + assert!(q.from.is_none()); + } + + #[test] + fn create_table_parses_types_and_nullability() { + match parse("CREATE TABLE t (id INT NOT NULL, name TEXT, score DOUBLE)").unwrap() { + Statement::CreateTable { name, columns } => { + assert_eq!(name, "t"); + assert_eq!(columns.len(), 3); + assert!(!columns[0].nullable); + assert!(columns[1].nullable); + assert_eq!(columns[2].data_type, DataType::Float64); + } + other => panic!("unexpected {other:?}"), + } + } + + #[test] + fn insert_copy_and_drop_parse() { + match parse("INSERT INTO t VALUES (1, 'a'), (2, 'b')").unwrap() { + Statement::Insert { table, rows } => { + assert_eq!(table, "t"); + assert_eq!(rows.len(), 2); + assert_eq!(rows[0].len(), 2); + } + other => panic!("unexpected {other:?}"), + } + match parse("COPY sales FROM 'sales.csv'").unwrap() { + Statement::Copy { table, path } => { + assert_eq!(table, "sales"); + assert_eq!(path, "sales.csv"); + } + other => panic!("unexpected {other:?}"), + } + assert!(matches!( + parse("DROP TABLE t").unwrap(), + Statement::DropTable { .. } + )); + } + + #[test] + fn copy_without_a_quoted_path_is_refused() { + assert!(parse("COPY t FROM sales.csv") + .unwrap_err() + .to_string() + .contains("quoted file path")); + } + + #[test] + fn explain_and_explain_analyze_parse() { + match parse("EXPLAIN SELECT a FROM t").unwrap() { + Statement::Explain { analyze, .. } => assert!(!analyze), + other => panic!("unexpected {other:?}"), + } + match parse("EXPLAIN ANALYZE SELECT a FROM t").unwrap() { + Statement::Explain { analyze, .. } => assert!(analyze), + other => panic!("unexpected {other:?}"), + } + } + + #[test] + fn several_statements_parse_and_empty_ones_are_skipped() { + let stmts = Parser::parse_statements("SELECT 1; ; SELECT 2;;").unwrap(); + assert_eq!(stmts.len(), 2); + assert!(Parser::parse_statements(" ").unwrap().is_empty()); + assert!(Parser::parse_one("").is_err()); + assert!(Parser::parse_one("SELECT 1; SELECT 2").is_err()); + } + + #[test] + fn errors_name_what_was_expected_and_where() { + let err = parse("SELECT FROM").unwrap_err(); + assert!(err.to_string().contains("expected an expression")); + assert!(err.to_string().contains("offset")); + + assert!(parse("SELECT a FROM").is_err()); + assert!(parse("SELECT a FROM t WHERE").is_err()); + assert!(parse("SELECT (a FROM t").is_err()); + assert!(parse("UPDATE t SET a = 1").is_err()); + assert!(parse("SELECT a FROM t GROUP a").is_err()); + assert!(parse("SELECT a FROM t ORDER a").is_err()); + assert!(parse("CREATE TABLE t").is_err()); + assert!(parse("INSERT INTO t (1)").is_err()); + } + + #[test] + fn a_statement_followed_by_garbage_is_refused() { + assert!(parse("SELECT a FROM t garbage extra").is_err()); + } + + #[test] + fn quoted_identifiers_survive_into_the_tree() { + let q = query(r#"SELECT "odd name" FROM "my table""#); + match &q.projection[0] { + SelectItem::Expr { expr, .. } => assert_eq!(expr.to_string(), "odd name"), + other => panic!("unexpected {other:?}"), + } + assert!(matches!( + q.from, + Some(TableRef::Table { ref name, .. }) if name == "my table" + )); + } +}