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" + )); + } +}