use std::collections::HashMap;
use rudb_common::{Error, Result};
use crate::ast::{
Ast, BinaryOp, CaseArm, Distinct, Expr, ExprRef, JoinKind, LiteralKind, Nulls, Order,
OrderItem, Quantifier, Query, QueryBody, QueryRef, Select, SelectRef, SetOp, Slice, Source,
SourceRef, Statement, StrRef, Target, UnaryOp,
};
use crate::generated::rules::PROGRAM;
use crate::matcher::{NONE, Tree, parse_tokens};
use crate::token::{Kind, Token};
use crate::tokenize::tokenize;
pub fn parse_ast(query: &str) -> Result<Ast> {
let tokens = tokenize(query)?;
let tree = parse_tokens(query, &tokens, PROGRAM, true)?;
transform(query, &tokens, &tree)
}
pub fn transform(query: &str, tokens: &[Token], tree: &Tree) -> Result<Ast> {
let mut transform =
Transform { query, tokens, tree, ast: Ast::default(), interned: HashMap::new() };
transform.program(tree.root())?;
Ok(transform.ast)
}
struct Transform<'a> {
query: &'a str,
tokens: &'a [Token],
tree: &'a Tree,
ast: Ast,
interned: HashMap<String, StrRef>,
}
impl<'a> Transform<'a> {
fn text(&self, node: u32) -> &'a str {
self.tree.text(node, self.query, self.tokens)
}
fn name(&self, node: u32) -> &'static str {
self.tree.name(node)
}
fn kids(&self, node: u32) -> impl Iterator<Item = u32> + use<'a> {
let tree = self.tree;
tree.children(node)
}
fn count(&self, node: u32) -> usize {
self.kids(node).count()
}
fn nth(&self, node: u32, n: usize) -> u32 {
self.kids(node).nth(n).unwrap_or(NONE)
}
fn first(&self, node: u32) -> u32 {
self.nth(node, 0)
}
fn find(&self, node: u32, name: &str) -> u32 {
self.kids(node).find(|&kid| self.name(kid) == name).unwrap_or(NONE)
}
fn leaves(&self, node: u32, out: &mut Vec<u32>) {
let mut any = false;
for kid in self.kids(node) {
any = true;
self.leaves(kid, &mut *out);
}
if !any {
out.push(node);
}
}
fn intern(&mut self, text: &str) -> StrRef {
if let Some(&index) = self.interned.get(text) {
return index;
}
let index = u32::try_from(self.ast.strings.len())
.map_err(|_| Error::internal("more than four billion strings in one query"))
.unwrap_or(NONE);
self.ast.strings.push(text.to_string());
self.interned.insert(text.to_string(), index);
index
}
fn push(&mut self, expr: Expr) -> ExprRef {
let index = self.ast.exprs.len() as u32;
self.ast.exprs.push(expr);
index
}
fn push_source(&mut self, source: Source) -> SourceRef {
let index = self.ast.sources.len() as u32;
self.ast.sources.push(source);
index
}
fn push_query(&mut self, query: Query) -> QueryRef {
let index = self.ast.queries.len() as u32;
self.ast.queries.push(query);
index
}
fn push_select(&mut self, select: Select) -> SelectRef {
let index = self.ast.selects.len() as u32;
self.ast.selects.push(select);
index
}
fn expr_slice(&mut self, items: Vec<ExprRef>) -> Slice {
let start = self.ast.expr_lists.len() as u32;
self.ast.expr_lists.extend(items);
Slice { start, len: self.ast.expr_lists.len() as u32 - start }
}
fn part_slice(&mut self, items: Vec<StrRef>) -> Slice {
let start = self.ast.parts.len() as u32;
self.ast.parts.extend(items);
Slice { start, len: self.ast.parts.len() as u32 - start }
}
fn unsupported<T>(&self, node: u32) -> Result<T> {
let text = self.text(node);
let text = if text.chars().count() > 60 {
let cut = text.char_indices().nth(60).map_or(text.len(), |(at, _)| at);
format!("{}...", &text[..cut])
} else {
text.to_string()
};
Err(Error::not_implemented(format!(
"{text} is not supported yet, the grammar rule is {}",
self.name(node)
)))
}
fn identifier(&mut self, node: u32) -> StrRef {
let mut leaves = Vec::new();
self.leaves(node, &mut leaves);
let text = leaves.last().map_or("", |&leaf| self.text(leaf));
let text = unquote(text.strip_suffix('.').unwrap_or(text));
self.intern(&text)
}
fn name_parts(&mut self, node: u32) -> Slice {
let mut leaves = Vec::new();
self.leaves(node, &mut leaves);
let mut parts = Vec::with_capacity(leaves.len());
for leaf in leaves {
let text = self.text(leaf);
if text.is_empty() || text == "*" {
continue;
}
let text = unquote(text.strip_suffix('.').unwrap_or(text));
let interned = self.intern(&text);
parts.push(interned);
}
self.part_slice(parts)
}
fn program(&mut self, node: u32) -> Result<()> {
for top in self.kids(node) {
let Some(statement) = self.kids(top).find(|&kid| self.name(kid) == "Statement") else {
continue;
};
let statement = self.statement(statement)?;
self.ast.statements.push(statement);
}
Ok(())
}
fn statement(&mut self, node: u32) -> Result<Statement> {
let inner = self.first(node);
match self.name(inner) {
"SelectStatement" => {
let query = self.query(self.first(inner))?;
Ok(Statement::Query(query))
}
_ => self.unsupported(inner),
}
}
fn query(&mut self, node: u32) -> Result<QueryRef> {
if self.find(node, "WithClause") != NONE {
return self.unsupported(self.find(node, "WithClause"));
}
let chain = self.find(node, "SelectSetOpChain");
if chain == NONE {
return self.unsupported(node);
}
let query = self.set_op_chain(chain)?;
let modifiers = self.find(node, "ResultModifiers");
if modifiers != NONE {
self.result_modifiers(query, modifiers)?;
}
Ok(query)
}
fn set_op_chain(&mut self, node: u32) -> Result<QueryRef> {
let mut kids = self.kids(node);
let head = kids.next().unwrap_or(NONE);
let mut left = self.intersect_chain(head)?;
for tail in kids {
let clause = self.first(tail);
let (op, quantifier, by_name) = self.setop_clause(clause)?;
let right = self.intersect_chain(self.nth(tail, 1))?;
left = self.push_query(Query::bare(QueryBody::SetOp {
op,
quantifier,
by_name,
left,
right,
}));
}
Ok(left)
}
fn intersect_chain(&mut self, node: u32) -> Result<QueryRef> {
let mut kids = self.kids(node);
let head = kids.next().unwrap_or(NONE);
let mut left = self.select_atom(head)?;
for tail in kids {
let clause = self.first(tail);
let quantifier = self.quantifier(self.find(clause, "DistinctOrAll"));
let right = self.select_atom(self.nth(tail, 1))?;
left = self.push_query(Query::bare(QueryBody::SetOp {
op: SetOp::Intersect,
quantifier,
by_name: false,
left,
right,
}));
}
Ok(left)
}
fn setop_clause(&mut self, node: u32) -> Result<(SetOp, Quantifier, bool)> {
let kind = self.find(node, "SetopType");
let op = match self.name(self.first(kind)) {
"SetopUnion" => SetOp::Union,
"SetopExcept" => SetOp::Except,
_ => return self.unsupported(kind),
};
let quantifier = self.quantifier(self.find(node, "DistinctOrAll"));
Ok((op, quantifier, self.find(node, "ByName") != NONE))
}
fn quantifier(&self, node: u32) -> Quantifier {
if node == NONE {
return Quantifier::Unstated;
}
match self.name(self.first(node)) {
"DistinctKeyword" => Quantifier::Distinct,
"AllKeyword" => Quantifier::All,
_ => Quantifier::Unstated,
}
}
fn select_atom(&mut self, node: u32) -> Result<QueryRef> {
let inner = self.first(node);
match self.name(inner) {
"SelectParens" => self.query(self.first(inner)),
"SelectStatementType" => {
let kind = self.first(inner);
match self.name(kind) {
"OptionalParensSimpleSelect" => {
let select = self.simple_select(self.unwrap_parens(kind))?;
Ok(self.push_query(Query::bare(QueryBody::Select(select))))
}
_ => self.unsupported(kind),
}
}
_ => self.unsupported(inner),
}
}
fn unwrap_parens(&self, node: u32) -> u32 {
let mut node = self.first(node);
while self.name(node) == "SimpleSelectParens" {
node = self.first(node);
}
node
}
fn result_modifiers(&mut self, query: QueryRef, node: u32) -> Result<()> {
let order = self.find(node, "OrderByClause");
if order != NONE {
let (items, all) = self.order_by(order)?;
let start = self.ast.order_items.len() as u32;
self.ast.order_items.extend(items);
self.ast.queries[query as usize].order_by =
Slice { start, len: self.ast.order_items.len() as u32 - start };
self.ast.queries[query as usize].order_by_all = all;
}
let limit = self.find(node, "LimitOffset");
if limit != NONE {
self.limit_offset(query, self.first(limit))?;
}
Ok(())
}
fn limit_offset(&mut self, query: QueryRef, node: u32) -> Result<()> {
match self.name(node) {
"LimitOffsetClause" | "OffsetLimitClause" => {
let limit = self.find(node, "LimitClause");
if limit != NONE {
self.limit(query, limit)?;
}
let offset = self.find(node, "OffsetClause");
if offset != NONE {
self.offset(query, offset)?;
}
Ok(())
}
_ => self.unsupported(node),
}
}
fn limit(&mut self, query: QueryRef, node: u32) -> Result<()> {
let value = self.first(node);
let inner = self.first(value);
match self.name(inner) {
"LimitAll" => Ok(()),
"LimitExpression" => {
let expr = self.expr(self.first(inner))?;
self.ast.queries[query as usize].limit = expr;
self.ast.queries[query as usize].limit_percent = self.text(inner).ends_with('%');
Ok(())
}
"LimitLiteralPercent" => {
let expr = self.expr(self.first(inner))?;
self.ast.queries[query as usize].limit = expr;
self.ast.queries[query as usize].limit_percent = true;
Ok(())
}
_ => self.unsupported(inner),
}
}
fn offset(&mut self, query: QueryRef, node: u32) -> Result<()> {
let value = self.first(node);
let expr = self.expr(self.first(value))?;
self.ast.queries[query as usize].offset = expr;
Ok(())
}
fn simple_select(&mut self, node: u32) -> Result<SelectRef> {
for name in ["WindowClause", "QualifyClause", "SampleClause"] {
let clause = self.find(node, name);
if clause != NONE {
return self.unsupported(clause);
}
}
let mut select = Select::empty();
self.select_from(&mut select, self.first(node))?;
let filter = self.find(node, "WhereClause");
if filter != NONE {
select.filter = self.expr(self.first(filter))?;
}
let group = self.find(node, "GroupByClause");
if group != NONE {
self.group_by(&mut select, self.first(group))?;
}
let having = self.find(node, "HavingClause");
if having != NONE {
select.having = self.expr(self.first(having))?;
}
Ok(self.push_select(select))
}
fn select_from(&mut self, select: &mut Select, node: u32) -> Result<()> {
let clause = self.first(node);
let targets = self.find(clause, "SelectClause");
let from = self.find(clause, "FromClause");
if from != NONE {
select.from = self.sources(from)?;
}
if targets == NONE {
let star = self.push(Expr::Star { qualifier: Slice::default() });
let start = self.ast.targets.len() as u32;
self.ast.targets.push(Target { expr: star, alias: NONE });
select.targets = Slice { start, len: 1 };
return Ok(());
}
self.select_clause(select, targets)
}
fn select_clause(&mut self, select: &mut Select, node: u32) -> Result<()> {
let distinct = self.find(node, "DistinctClause");
if distinct != NONE {
let inner = self.first(distinct);
select.distinct = match self.name(inner) {
"DistinctAll" => Distinct::No,
"DistinctOn" => {
let on = self.find(inner, "DistinctOnTargets");
if on == NONE {
Distinct::Yes
} else {
let mut items = Vec::new();
for kid in self.kids(on) {
items.push(self.expr(kid)?);
}
Distinct::On(self.expr_slice(items))
}
}
_ => return self.unsupported(inner),
};
}
let list = self.find(node, "TargetList");
if list == NONE {
return Ok(());
}
let mut targets = Vec::new();
for kid in self.kids(list) {
targets.push(self.target(kid)?);
}
let start = self.ast.targets.len() as u32;
self.ast.targets.extend(targets);
select.targets = Slice { start, len: self.ast.targets.len() as u32 - start };
Ok(())
}
fn target(&mut self, node: u32) -> Result<Target> {
let inner = self.first(node);
match self.name(inner) {
"ColIdExpression" => {
let alias = self.identifier(self.first(inner));
let expr = self.expr(self.nth(inner, 1))?;
Ok(Target { expr, alias })
}
"ExpressionAsCollabel" => {
let expr = self.expr(self.first(inner))?;
let alias = self.identifier(self.nth(inner, 1));
Ok(Target { expr, alias })
}
"ExpressionOptIdentifier" => {
let expr = self.expr(self.first(inner))?;
let alias =
if self.count(inner) > 1 { self.identifier(self.nth(inner, 1)) } else { NONE };
Ok(Target { expr, alias })
}
_ => self.unsupported(inner),
}
}
fn group_by(&mut self, select: &mut Select, node: u32) -> Result<()> {
let inner = self.first(node);
match self.name(inner) {
"GroupByAll" => {
select.group_by_all = true;
Ok(())
}
"GroupByList" => {
let mut items = Vec::new();
for kid in self.kids(inner) {
let expression = self.first(kid);
if self.name(expression) != "GroupByBaseExpression" {
return self.unsupported(expression);
}
items.push(self.expr(self.first(expression))?);
}
select.group_by = self.expr_slice(items);
Ok(())
}
_ => self.unsupported(inner),
}
}
fn order_by(&mut self, node: u32) -> Result<(Vec<OrderItem>, bool)> {
let inner = self.first(self.first(node));
match self.name(inner) {
"OrderByAll" => {
let (order, nulls) = self.sort_options(inner);
Ok((vec![OrderItem { expr: NONE, order, nulls }], true))
}
"OrderByExpressionList" => {
let mut items = Vec::new();
for kid in self.kids(inner) {
let expr = self.expr(self.first(kid))?;
let (order, nulls) = self.sort_options(kid);
items.push(OrderItem { expr, order, nulls });
}
Ok((items, false))
}
_ => self.unsupported(inner),
}
}
fn sort_options(&self, node: u32) -> (Order, Nulls) {
let direction = self.find(node, "DescOrAsc");
let order = if direction == NONE {
Order::Unstated
} else if self.name(self.first(direction)) == "DescendingOrder" {
Order::Descending
} else {
Order::Ascending
};
let placement = self.find(node, "NullsFirstOrLast");
let nulls = if placement == NONE {
Nulls::Unstated
} else if self.name(self.first(placement)) == "NullsFirst" {
Nulls::First
} else {
Nulls::Last
};
(order, nulls)
}
fn sources(&mut self, node: u32) -> Result<Slice> {
let mut items = Vec::new();
for kid in self.kids(node) {
items.push(self.table_ref(kid)?);
}
let start = self.ast.source_lists.len() as u32;
self.ast.source_lists.extend(items);
Ok(Slice { start, len: self.ast.source_lists.len() as u32 - start })
}
fn table_ref(&mut self, node: u32) -> Result<SourceRef> {
let mut kids = self.kids(node);
let head = kids.next().unwrap_or(NONE);
let mut left = self.inner_table_ref(head)?;
for tail in kids {
let clause = self.first(tail);
if self.name(clause) != "JoinClause" {
return self.unsupported(clause);
}
left = self.join(left, self.first(clause))?;
}
Ok(left)
}
fn inner_table_ref(&mut self, node: u32) -> Result<SourceRef> {
let inner = if self.name(node) == "InnerTableRef" { self.first(node) } else { node };
match self.name(inner) {
"BaseTableRef" => {
if self.find(inner, "TableAliasColon") != NONE {
return self.unsupported(inner);
}
for name in ["AtClause", "SampleClause"] {
let clause = self.find(inner, name);
if clause != NONE {
return self.unsupported(clause);
}
}
let name = self.name_parts(self.find(inner, "BaseTableName"));
let (alias, columns) = self.table_alias(self.find(inner, "TableAlias"));
Ok(self.push_source(Source::Table { name, alias, columns }))
}
"TableSubquery" => {
if self.find(inner, "TableAliasColon") != NONE
|| self.find(inner, "Lateral") != NONE
{
return self.unsupported(inner);
}
let reference = self.find(inner, "SubqueryReference");
let query = self.query(self.first(reference))?;
let (alias, columns) = self.table_alias(self.find(inner, "TableAlias"));
Ok(self.push_source(Source::Subquery { query, alias, columns }))
}
"ParensTableRef" => {
if self.find(inner, "TableAliasColon") != NONE
|| self.find(inner, "SampleClause") != NONE
|| self.find(inner, "TableAlias") != NONE
{
return self.unsupported(inner);
}
self.table_ref(self.find(inner, "TableRef"))
}
_ => self.unsupported(inner),
}
}
fn table_alias(&mut self, node: u32) -> (StrRef, Slice) {
if node == NONE {
return (NONE, Slice::default());
}
let inner = self.first(node);
let alias = self.identifier(self.first(inner));
let list = self.find(inner, "ColumnAliases");
if list == NONE {
return (alias, Slice::default());
}
let mut columns = Vec::new();
for kid in self.kids(list) {
let name = self.identifier(kid);
columns.push(name);
}
(alias, self.part_slice(columns))
}
fn join(&mut self, left: SourceRef, node: u32) -> Result<SourceRef> {
match self.name(node) {
"RegularJoinClause" => {
if self.find(node, "Asof") != NONE {
return self.unsupported(node);
}
let kind = self.join_type(self.find(node, "JoinType"));
let right = self.table_ref(self.find(node, "TableRef"))?;
let (on, using) = self.join_qualifier(self.find(node, "JoinQualifier"))?;
Ok(self.push_source(Source::Join { left, right, kind, natural: false, on, using }))
}
"JoinWithoutOnClause" => {
let prefix = self.first(self.find(node, "JoinPrefix"));
let (kind, natural) = match self.name(prefix) {
"CrossJoinPrefix" => (JoinKind::Cross, false),
"PositionalJoinPrefix" => (JoinKind::Positional, false),
"NaturalJoinPrefix" => (self.join_type(self.find(prefix, "JoinType")), true),
_ => return self.unsupported(prefix),
};
let right = self.inner_table_ref(self.find(node, "InnerTableRef"))?;
Ok(self.push_source(Source::Join {
left,
right,
kind,
natural,
on: NONE,
using: Slice::default(),
}))
}
_ => self.unsupported(node),
}
}
fn join_type(&self, node: u32) -> JoinKind {
if node == NONE {
return JoinKind::Inner;
}
match self.name(self.first(node)) {
"FullJoin" => JoinKind::Full,
"LeftJoin" => JoinKind::Left,
"RightJoin" => JoinKind::Right,
"SemiJoin" => JoinKind::Semi,
"AntiJoin" => JoinKind::Anti,
_ => JoinKind::Inner,
}
}
fn join_qualifier(&mut self, node: u32) -> Result<(ExprRef, Slice)> {
let inner = self.first(node);
match self.name(inner) {
"OnClause" => Ok((self.expr(self.first(inner))?, Slice::default())),
"UsingClause" => {
let mut columns = Vec::new();
for kid in self.kids(inner) {
let name = self.identifier(kid);
columns.push(name);
}
Ok((NONE, self.part_slice(columns)))
}
_ => self.unsupported(inner),
}
}
fn expr(&mut self, node: u32) -> Result<ExprRef> {
let mut node = node;
loop {
let count = self.count(node);
let name = self.name(node);
match name {
"LogicalOrExpression" if count > 1 => return self.logical(node, BinaryOp::Or),
"LogicalAndExpression" if count > 1 => return self.logical(node, BinaryOp::And),
"LogicalNotExpression" if count > 1 => return self.logical_not(node),
"IsExpression" if count > 1 => return self.is_expression(node),
"BetweenInLikeExpression" if count > 1 => return self.between_in_like(node),
"PrefixExpression" if count > 1 => return self.prefix(node),
"BaseExpression" if count > 1 => return self.indirection(node),
"LambdaArrowExpression"
| "IsDistinctFromExpression"
| "ComparisonExpression"
| "OtherOperatorExpression"
| "BitwiseExpression"
| "AdditiveExpression"
| "MultiplicativeExpression"
| "ExponentiationExpression"
| "CollateExpression"
| "AtTimeZoneExpression"
if count > 1 =>
{
return self.tail_chain(node);
}
"ColumnReference" => {
let name = self.name_parts(node);
return Ok(self.push(Expr::Column { name }));
}
"StarExpression" => return self.star(node),
"NumberLiteral" => {
let text = self.text(node).to_string();
let text = self.intern(&text);
return Ok(self.push(Expr::Literal { kind: LiteralKind::Number, text }));
}
"StringLiteral" => {
let text = self.string_value(node);
let text = self.intern(&text);
return Ok(self.push(Expr::Literal { kind: LiteralKind::String, text }));
}
"NullLiteral" | "TrueLiteral" | "FalseLiteral" => {
let kind = match name {
"NullLiteral" => LiteralKind::Null,
"TrueLiteral" => LiteralKind::True,
_ => LiteralKind::False,
};
return Ok(self.push(Expr::Literal { kind, text: NONE }));
}
"FunctionExpression" => return self.function(node),
"CastExpression" => return self.cast(node),
"CaseExpression" => return self.case(node),
"ParenthesisExpression" => return self.row(node),
"SubqueryExpression" => return self.subquery(node),
_ if count == 1 => node = self.first(node),
_ => return self.unsupported(node),
}
}
}
fn tail_chain(&mut self, node: u32) -> Result<ExprRef> {
let mut kids = self.kids(node);
let head = kids.next().unwrap_or(NONE);
let mut left = self.expr(head)?;
for tail in kids {
let operator = self.first(tail);
let op = self.binary_op(operator)?;
let operand = self.kids(tail).last().unwrap_or(NONE);
if self.count(tail) > 2 {
return self.unsupported(tail);
}
let right = self.expr(operand)?;
left = self.push(Expr::Binary { op, left, right });
}
Ok(left)
}
fn binary_op(&mut self, node: u32) -> Result<BinaryOp> {
let mut leaf = node;
while self.count(leaf) == 1 {
leaf = self.first(leaf);
}
let text = self.text(node);
let upper = text.to_ascii_uppercase();
let op = match upper.as_str() {
"OR" => BinaryOp::Or,
"AND" => BinaryOp::And,
"=" | "==" => BinaryOp::Eq,
"!=" | "<>" => BinaryOp::NotEq,
"<" => BinaryOp::Lt,
">" => BinaryOp::Gt,
"<=" => BinaryOp::LtEq,
">=" => BinaryOp::GtEq,
"+" => BinaryOp::Add,
"-" => BinaryOp::Subtract,
"*" => BinaryOp::Multiply,
"/" => BinaryOp::Divide,
"//" => BinaryOp::IntegerDivide,
"%" => BinaryOp::Modulo,
"^" | "**" => BinaryOp::Power,
"&" => BinaryOp::BitAnd,
"|" => BinaryOp::BitOr,
"<<" => BinaryOp::ShiftLeft,
">>" => BinaryOp::ShiftRight,
"||" => BinaryOp::Concat,
"COLLATE" => BinaryOp::Collate,
"->" => BinaryOp::Arrow,
"->>" => BinaryOp::LongArrow,
"@>" => BinaryOp::Contains,
"<@" => BinaryOp::ContainedBy,
"&&" => BinaryOp::Overlaps,
"^@" => BinaryOp::StartsWith,
"<<=" => BinaryOp::InetContainedByOrEq,
">>=" => BinaryOp::InetContainsOrEq,
_ if self.name(leaf) == "AtTimeZoneOperator" => BinaryOp::AtTimeZone,
_ if self.name(leaf) == "IsDistinctFromOp" => {
if upper.split_whitespace().any(|word| word == "NOT") {
BinaryOp::IsNotDistinctFrom
} else {
BinaryOp::IsDistinctFrom
}
}
_ if self.name(leaf) == "OperatorLiteral" => {
let interned = self.intern(text);
BinaryOp::Named(interned)
}
_ => return self.unsupported(node),
};
Ok(op)
}
fn logical(&mut self, node: u32, op: BinaryOp) -> Result<ExprRef> {
let mut kids = self.kids(node);
let head = kids.next().unwrap_or(NONE);
let mut left = self.expr(head)?;
for tail in kids {
let right = self.expr(self.first(tail))?;
left = self.push(Expr::Binary { op, left, right });
}
Ok(left)
}
fn logical_not(&mut self, node: u32) -> Result<ExprRef> {
let negations = self.count(self.first(node));
let mut expr = self.expr(self.nth(node, 1))?;
for _ in 0..negations {
expr = self.push(Expr::Unary { op: UnaryOp::Not, operand: expr });
}
Ok(expr)
}
fn is_expression(&mut self, node: u32) -> Result<ExprRef> {
let mut kids = self.kids(node);
let head = kids.next().unwrap_or(NONE);
let mut expr = self.expr(head)?;
for test in kids {
let inner = self.first(test);
let negated = self.text(inner).to_ascii_uppercase().contains("NOT");
let op = match self.name(inner) {
"NotNull" => UnaryOp::IsNotNull,
"IsNull" => UnaryOp::IsNull,
"IsLiteral" => match self.name(self.first(self.first(inner))) {
"NullLiteral" if negated => UnaryOp::IsNotNull,
"NullLiteral" => UnaryOp::IsNull,
"TrueLiteral" if negated => UnaryOp::IsNotTrue,
"TrueLiteral" => UnaryOp::IsTrue,
"FalseLiteral" if negated => UnaryOp::IsNotFalse,
"FalseLiteral" => UnaryOp::IsFalse,
"UnknownLiteral" if negated => UnaryOp::IsNotUnknown,
"UnknownLiteral" => UnaryOp::IsUnknown,
_ => return self.unsupported(inner),
},
_ => return self.unsupported(inner),
};
expr = self.push(Expr::Unary { op, operand: expr });
}
Ok(expr)
}
fn between_in_like(&mut self, node: u32) -> Result<ExprRef> {
let operand = self.expr(self.first(node))?;
let op = self.nth(node, 1);
let negated = self.text(op).to_ascii_uppercase().starts_with("NOT");
let inner = self.first(self.first(op));
match self.name(inner) {
"BetweenClause" => {
let low = self.expr(self.first(inner))?;
let high = self.expr(self.nth(inner, 1))?;
Ok(self.push(Expr::Between { operand, low, high, negated }))
}
"InClause" => {
let expression = self.first(self.first(inner));
match self.name(expression) {
"InExpressionList" => {
let mut items = Vec::new();
for kid in self.kids(expression) {
items.push(self.expr(kid)?);
}
let list = self.expr_slice(items);
Ok(self.push(Expr::In { operand, list, negated }))
}
_ => self.unsupported(expression),
}
}
"LikeClause" => {
if self.find(inner, "EscapeClause") != NONE {
return self.unsupported(inner);
}
let variation = self.name(self.first(self.first(inner)));
let op = match (variation, negated) {
("LikeToken", false) | ("NotLikeOp", true) => BinaryOp::Like,
("LikeToken", true) | ("NotLikeOp", false) => BinaryOp::NotLike,
("ILikeToken", false) | ("NotILikeOp", true) => BinaryOp::ILike,
("ILikeToken", true) | ("NotILikeOp", false) => BinaryOp::NotILike,
("GlobToken", _) => BinaryOp::Glob,
("RegexMatchToken", _) => BinaryOp::Regex,
("SimilarToToken", false) | ("NotSimilarToOp", true) => BinaryOp::SimilarTo,
("SimilarToToken", true) | ("NotSimilarToOp", false) => BinaryOp::NotSimilarTo,
("RegexInsensitiveMatchToken", false)
| ("NotRegexInsensitiveMatchOp", true) => BinaryOp::RegexInsensitive,
("RegexInsensitiveMatchToken", true)
| ("NotRegexInsensitiveMatchOp", false) => BinaryOp::NotRegexInsensitive,
_ => return self.unsupported(inner),
};
let right = self.expr(self.nth(inner, 1))?;
let expr = self.push(Expr::Binary { op, left: operand, right });
if negated && matches!(op, BinaryOp::Glob | BinaryOp::Regex) {
return Ok(self.push(Expr::Unary { op: UnaryOp::Not, operand: expr }));
}
Ok(expr)
}
_ => self.unsupported(inner),
}
}
fn prefix(&mut self, node: u32) -> Result<ExprRef> {
let kids: Vec<u32> = self.kids(node).collect();
let mut expr = self.expr(kids[kids.len() - 1])?;
for &operator in kids[..kids.len() - 1].iter().rev() {
let op = match self.name(self.first(operator)) {
"MinusPrefixOperator" => UnaryOp::Negate,
"PlusPrefixOperator" => UnaryOp::Plus,
"TildePrefixOperator" => UnaryOp::BitNot,
_ => return self.unsupported(operator),
};
expr = self.push(Expr::Unary { op, operand: expr });
}
Ok(expr)
}
fn indirection(&mut self, node: u32) -> Result<ExprRef> {
let mut expr = self.expr(self.first(node))?;
for step in self.kids(self.nth(node, 1)) {
let inner = self.first(step);
expr = match self.name(inner) {
"CastOperator" => {
let text = self.text(self.first(inner)).to_string();
let ty = self.intern(&text);
self.push(Expr::Cast { operand: expr, ty, try_cast: false })
}
"DotOperator" => {
let dot = self.first(inner);
match self.name(dot) {
"DotColumnOperator" => {
let field = self.identifier(self.first(dot));
let text = self.ast.string(field).to_string();
let literal = self.intern(&text);
let key = self
.push(Expr::Literal { kind: LiteralKind::String, text: literal });
let name = self.function_name("struct_extract");
let args = self.expr_slice(vec![expr, key]);
self.push(Expr::Function { name, args, distinct: false })
}
"DotMethodOperator" => {
let method = self.first(dot);
let text = self.text(self.first(method)).to_string();
let text = unquote(&text);
let name = self.function_name(&text);
let mut args = vec![expr];
let list = self.find(method, "MethodExpressionArguments");
if list != NONE {
let inner = self.first(list);
let arguments = self.find(inner, "MethodFunctionArguments");
if arguments != NONE {
for kid in self.kids(arguments) {
args.push(self.argument(kid)?);
}
}
}
let args = self.expr_slice(args);
self.push(Expr::Function { name, args, distinct: false })
}
_ => return self.unsupported(dot),
}
}
"SliceExpression" => {
let bound = self.first(inner);
let has_end = self.find(bound, "EndSliceBound") != NONE;
let has_step = self.find(bound, "StepSliceBound") != NONE;
if has_end || has_step {
return self.unsupported(inner);
}
let index = self.expr(self.first(bound))?;
let name = self.function_name("array_extract");
let args = self.expr_slice(vec![expr, index]);
self.push(Expr::Function { name, args, distinct: false })
}
"PostfixOperator" => {
self.push(Expr::Unary { op: UnaryOp::Factorial, operand: expr })
}
_ => return self.unsupported(inner),
};
}
Ok(expr)
}
fn function_name(&mut self, name: &str) -> Slice {
let interned = self.intern(name);
self.part_slice(vec![interned])
}
fn star(&mut self, node: u32) -> Result<ExprRef> {
for name in ["ExcludeList", "ReplaceList", "RenameList"] {
let list = self.find(node, name);
if list != NONE {
return self.unsupported(list);
}
}
let qualifier = self.find(node, "StarQualifierList");
let qualifier =
if qualifier == NONE { Slice::default() } else { self.name_parts(qualifier) };
Ok(self.push(Expr::Star { qualifier }))
}
fn function(&mut self, node: u32) -> Result<ExprRef> {
for name in ["WithinGroupClause", "FilterClause", "ExportClause", "OverClause"] {
let clause = self.find(node, name);
if clause != NONE {
return self.unsupported(clause);
}
}
let name = self.name_parts(self.first(node));
let list = self.first(self.nth(node, 1));
for name in ["OrderByClause", "IgnoreOrRespectNulls"] {
let clause = self.find(list, name);
if clause != NONE {
return self.unsupported(clause);
}
}
let distinct = self.quantifier(self.find(list, "DistinctOrAll")) == Quantifier::Distinct;
let mut args = Vec::new();
let arguments = self.find(list, "FunctionArgumentList");
if arguments != NONE {
for kid in self.kids(arguments) {
args.push(self.argument(kid)?);
}
}
let args = self.expr_slice(args);
Ok(self.push(Expr::Function { name, args, distinct }))
}
fn argument(&mut self, node: u32) -> Result<ExprRef> {
let inner = self.first(node);
match self.name(inner) {
"PositionalFunctionArgument" => self.expr(self.first(inner)),
_ => self.unsupported(inner),
}
}
fn cast(&mut self, node: u32) -> Result<ExprRef> {
let try_cast = self.name(self.first(self.first(node))) == "TryCastKeyword";
let arguments = self.nth(node, 1);
let operand = self.expr(self.first(arguments))?;
let text = self.text(self.nth(arguments, 1)).to_string();
let ty = self.intern(&text);
Ok(self.push(Expr::Cast { operand, ty, try_cast }))
}
fn case(&mut self, node: u32) -> Result<ExprRef> {
let mut operand = NONE;
let mut arms = Vec::new();
let mut otherwise = NONE;
for kid in self.kids(node) {
match self.name(kid) {
"CaseWhenThen" => {
let when = self.expr(self.first(kid))?;
let then = self.expr(self.nth(kid, 1))?;
arms.push(CaseArm { when, then });
}
"CaseElse" => otherwise = self.expr(self.first(kid))?,
_ => operand = self.expr(kid)?,
}
}
let start = self.ast.case_arms.len() as u32;
self.ast.case_arms.extend(arms);
let arms = Slice { start, len: self.ast.case_arms.len() as u32 - start };
Ok(self.push(Expr::Case { operand, arms, otherwise }))
}
fn row(&mut self, node: u32) -> Result<ExprRef> {
let mut items = Vec::new();
for kid in self.kids(node) {
items.push(self.expr(kid)?);
}
if items.len() == 1 {
return Ok(items[0]);
}
let items = self.expr_slice(items);
Ok(self.push(Expr::Row { items }))
}
fn subquery(&mut self, node: u32) -> Result<ExprRef> {
if self.find(node, "SubqueryNot") != NONE || self.find(node, "SubqueryExists") != NONE {
return self.unsupported(node);
}
let reference = self.find(node, "SubqueryReference");
let query = self.query(self.first(reference))?;
Ok(self.push(Expr::Subquery { query }))
}
fn string_value(&self, node: u32) -> String {
let span = self.tree.node(node);
let mut value = String::new();
for token in &self.tokens[span.start as usize..span.end as usize] {
if token.kind != Kind::String {
continue;
}
let text = token.text(self.query);
match text.strip_prefix('\'').and_then(|rest| rest.strip_suffix('\'')) {
Some(body) => value.push_str(&body.replace("''", "'")),
None => value.push_str(text),
}
}
value
}
}
fn unquote(text: &str) -> String {
match text.strip_prefix('"').and_then(|rest| rest.strip_suffix('"')) {
Some(body) => body.replace("\"\"", "\""),
None => text.to_string(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::corpus::CORPUS;
use crate::matcher::parse;
fn show(ast: &Ast, expr: ExprRef) -> String {
if expr == NONE {
return "-".to_string();
}
let list = |slice: Slice| {
ast.expr_list(slice).iter().map(|&item| show(ast, item)).collect::<Vec<_>>().join(", ")
};
match ast.expr(expr) {
Expr::Star { qualifier } if qualifier.is_empty() => "*".to_string(),
Expr::Star { qualifier } => format!("{}.*", ast.name_text(qualifier)),
Expr::Column { name } => ast.name_text(name),
Expr::Literal { kind, text } => match kind {
LiteralKind::Number => ast.string(text).to_string(),
LiteralKind::String => format!("'{}'", ast.string(text)),
other => format!("{other:?}").to_uppercase(),
},
Expr::Unary { op, operand } => format!("({op:?} {})", show(ast, operand)),
Expr::Binary { op, left, right } => {
let op = match op {
BinaryOp::Named(name) => ast.string(name).to_string(),
other => format!("{other:?}"),
};
format!("({} {op} {})", show(ast, left), show(ast, right))
}
Expr::Function { name, args, distinct } => {
let distinct = if distinct { "DISTINCT " } else { "" };
format!("{}({distinct}{})", ast.name_text(name), list(args))
}
Expr::Cast { operand, ty, try_cast } => {
let word = if try_cast { "TRY_CAST" } else { "CAST" };
format!("{word}({} AS {})", show(ast, operand), ast.string(ty))
}
Expr::Case { operand, arms, otherwise } => {
let arms = ast
.arm_list(arms)
.iter()
.map(|arm| format!("WHEN {} THEN {}", show(ast, arm.when), show(ast, arm.then)))
.collect::<Vec<_>>()
.join(" ");
format!("CASE {} {arms} ELSE {} END", show(ast, operand), show(ast, otherwise))
}
Expr::Between { operand, low, high, negated } => {
let not = if negated { "NOT " } else { "" };
format!(
"({not}{} BETWEEN {} AND {})",
show(ast, operand),
show(ast, low),
show(ast, high)
)
}
Expr::In { operand, list: items, negated } => {
let not = if negated { "NOT " } else { "" };
format!("({not}{} IN [{}])", show(ast, operand), list(items))
}
Expr::Row { items } => format!("ROW({})", list(items)),
Expr::Subquery { query } => format!("({})", show_query(ast, query)),
}
}
fn show_source(ast: &Ast, source: SourceRef) -> String {
let alias = |alias: StrRef| match alias {
NONE => String::new(),
other => format!(" AS {}", ast.string(other)),
};
match ast.source(source) {
Source::Table { name, alias: name_alias, .. } => {
format!("{}{}", ast.name_text(name), alias(name_alias))
}
Source::Subquery { query, alias: query_alias, .. } => {
format!("({}){}", show_query(ast, query), alias(query_alias))
}
Source::Join { left, right, kind, natural, on, using } => {
let natural = if natural { "NATURAL " } else { "" };
let on = if on == NONE { String::new() } else { format!(" ON {}", show(ast, on)) };
let using = if using.is_empty() {
String::new()
} else {
format!(" USING ({})", ast.name_text(using))
};
format!(
"({} {natural}{kind:?} JOIN {}{on}{using})",
show_source(ast, left),
show_source(ast, right)
)
}
}
}
fn show_query(ast: &Ast, index: QueryRef) -> String {
let query = ast.query(index);
let list = |slice: Slice| {
ast.expr_list(slice).iter().map(|&item| show(ast, item)).collect::<Vec<_>>().join(", ")
};
let mut out = match query.body {
QueryBody::SetOp { op, quantifier, by_name, left, right } => {
let by_name = if by_name { " BY NAME" } else { "" };
format!(
"({} {op:?} {quantifier:?}{by_name} {})",
show_query(ast, left),
show_query(ast, right)
)
}
QueryBody::Select(index) => {
let select = ast.select(index);
let distinct = match select.distinct {
Distinct::No => String::new(),
Distinct::Yes => " DISTINCT".to_string(),
Distinct::On(on) => format!(" DISTINCT ON ({})", list(on)),
};
let targets = ast
.target_list(select.targets)
.iter()
.map(|target| match target.alias {
NONE => show(ast, target.expr),
alias => format!("{} AS {}", show(ast, target.expr), ast.string(alias)),
})
.collect::<Vec<_>>()
.join(", ");
let mut out = format!("SELECT{distinct} {targets}");
if !select.from.is_empty() {
let from = ast
.source_list(select.from)
.iter()
.map(|&source| show_source(ast, source))
.collect::<Vec<_>>()
.join(", ");
out += &format!(" FROM {from}");
}
if select.filter != NONE {
out += &format!(" WHERE {}", show(ast, select.filter));
}
if select.group_by_all {
out += " GROUP BY ALL";
} else if !select.group_by.is_empty() {
out += &format!(" GROUP BY {}", list(select.group_by));
}
if select.having != NONE {
out += &format!(" HAVING {}", show(ast, select.having));
}
out
}
};
if query.order_by_all {
out += " ORDER BY ALL";
} else if !query.order_by.is_empty() {
let items = ast
.order_list(query.order_by)
.iter()
.map(|item| format!("{} {:?} {:?}", show(ast, item.expr), item.order, item.nulls))
.collect::<Vec<_>>()
.join(", ");
out += &format!(" ORDER BY {items}");
}
if query.limit != NONE {
let percent = if query.limit_percent { "%" } else { "" };
out += &format!(" LIMIT {}{percent}", show(ast, query.limit));
}
if query.offset != NONE {
out += &format!(" OFFSET {}", show(ast, query.offset));
}
out
}
fn round(query: &str) -> String {
let ast = parse_ast(query).unwrap_or_else(|error| panic!("{query}: {error}"));
assert_eq!(ast.statements.len(), 1, "{query} is one statement");
let Statement::Query(index) = ast.statements[0];
show_query(&ast, index)
}
#[test]
fn the_query_m0_has_to_run_transforms() {
assert_eq!(round("SELECT * FROM t WHERE x > 5"), "SELECT * FROM t WHERE (x Gt 5)");
}
#[test]
fn every_statement_in_the_corpus_gets_a_defined_answer() {
let mut done = 0;
for query in CORPUS {
match parse_ast(query) {
Ok(ast) => {
assert_eq!(ast.statements.len(), 1, "{query}");
done += 1;
}
Err(error) => {
let message = error.to_string();
assert!(
message.starts_with("Not implemented Error"),
"{query} failed with {message}, which is not a not-implemented error"
);
}
}
}
assert!(done >= 19, "only {done} of the corpus transforms, which is fewer than it was");
}
#[test]
fn the_ast_is_far_smaller_than_the_parse_tree() {
let query = CORPUS[4];
let tree = parse(query).unwrap();
let ast = parse_ast(query).unwrap();
assert!(
ast.node_count() * 20 < tree.arena_len(),
"{} ast nodes against {} parse nodes",
ast.node_count(),
tree.arena_len()
);
}
#[test]
fn precedence_comes_out_of_the_chain_and_into_the_tree() {
assert_eq!(round("SELECT 1 + 2 * 3"), "SELECT (1 Add (2 Multiply 3))");
assert_eq!(round("SELECT (1 + 2) * 3"), "SELECT ((1 Add 2) Multiply 3)");
assert_eq!(round("SELECT 1 + 2 + 3"), "SELECT ((1 Add 2) Add 3)");
assert_eq!(round("SELECT 1 - 2 - 3"), "SELECT ((1 Subtract 2) Subtract 3)");
assert_eq!(
round("SELECT a OR b AND c"),
"SELECT (a Or (b And c))",
"and binds tighter than or"
);
}
#[test]
fn a_double_negation_is_two_nodes_and_not_none() {
assert_eq!(round("SELECT NOT NOT a"), "SELECT (Not (Not a))");
}
#[test]
fn a_parenthesised_single_expression_is_not_a_row() {
assert_eq!(round("SELECT (a)"), "SELECT a");
assert_eq!(round("SELECT (a, b)"), "SELECT ROW(a, b)");
}
#[test]
fn the_three_ways_to_write_an_alias_all_arrive() {
assert_eq!(round("SELECT a AS b"), "SELECT a AS b");
assert_eq!(round("SELECT a b"), "SELECT a AS b");
assert_eq!(round("SELECT b: a"), "SELECT a AS b");
assert_eq!(round("SELECT a"), "SELECT a", "and no alias when none was written");
}
#[test]
fn a_from_with_no_select_selects_everything() {
assert_eq!(round("FROM t"), "SELECT * FROM t");
assert_eq!(round("FROM t SELECT a"), "SELECT a FROM t");
}
#[test]
fn joins_nest_to_the_left() {
assert_eq!(
round("SELECT * FROM a JOIN b ON a.i = b.i LEFT JOIN c USING (k)"),
"SELECT * FROM ((a Inner JOIN b ON (a.i Eq b.i)) Left JOIN c USING (k))"
);
assert_eq!(
round("SELECT * FROM a NATURAL JOIN b"),
"SELECT * FROM (a NATURAL Inner JOIN b)"
);
assert_eq!(round("SELECT * FROM a CROSS JOIN b"), "SELECT * FROM (a Cross JOIN b)");
assert_eq!(
round("SELECT * FROM a POSITIONAL JOIN b"),
"SELECT * FROM (a Positional JOIN b)"
);
assert_eq!(round("SELECT * FROM a, b"), "SELECT * FROM a, b", "a comma is not a join node");
}
#[test]
fn a_qualified_name_keeps_its_parts_however_it_was_spelled() {
assert_eq!(round("SELECT a"), "SELECT a");
assert_eq!(round("SELECT t.a"), "SELECT t.a");
assert_eq!(round("SELECT s.t.a"), "SELECT s.t.a");
assert_eq!(round("SELECT c.s.t.a"), "SELECT c.s.t.a");
assert_eq!(round("SELECT * FROM s.t"), "SELECT * FROM s.t");
}
#[test]
fn a_star_can_be_qualified() {
assert_eq!(round("SELECT *"), "SELECT *");
assert_eq!(round("SELECT t.*"), "SELECT t.*");
assert_eq!(round("SELECT s.t.*"), "SELECT s.t.*");
}
#[test]
fn a_quoted_identifier_keeps_its_case_and_loses_its_quotes() {
let ast = parse_ast("SELECT \"Mixed Case\", \"a\"\"b\"").unwrap();
assert_eq!(ast.strings[0], "Mixed Case");
assert_eq!(ast.strings[1], "a\"b");
}
#[test]
fn a_string_literal_is_decoded_and_adjacent_ones_are_joined() {
assert_eq!(round("SELECT 'it''s'"), "SELECT 'it's'");
assert_eq!(round("SELECT 'a'\n'b'"), "SELECT 'ab'", "the standard's adjacency rule");
}
#[test]
fn the_null_and_boolean_tests_are_postfix_unary_operators() {
assert_eq!(round("SELECT x IS NULL"), "SELECT (IsNull x)");
assert_eq!(round("SELECT x IS NOT NULL"), "SELECT (IsNotNull x)");
assert_eq!(round("SELECT x ISNULL"), "SELECT (IsNull x)");
assert_eq!(round("SELECT x NOTNULL"), "SELECT (IsNotNull x)");
assert_eq!(round("SELECT x IS TRUE"), "SELECT (IsTrue x)");
assert_eq!(round("SELECT x IS NOT FALSE"), "SELECT (IsNotFalse x)");
assert_eq!(round("SELECT x IS DISTINCT FROM y"), "SELECT (x IsDistinctFrom y)");
assert_eq!(round("SELECT x IS NOT DISTINCT FROM y"), "SELECT (x IsNotDistinctFrom y)");
}
#[test]
fn the_like_family_folds_its_negation_into_the_operator() {
assert_eq!(round("SELECT x LIKE 'a'"), "SELECT (x Like 'a')");
assert_eq!(round("SELECT x NOT LIKE 'a'"), "SELECT (x NotLike 'a')");
assert_eq!(round("SELECT x ILIKE 'a'"), "SELECT (x ILike 'a')");
assert_eq!(round("SELECT x ~~ 'a'"), "SELECT (x Like 'a')", "the operator spelling");
assert_eq!(round("SELECT x !~~ 'a'"), "SELECT (x NotLike 'a')");
assert_eq!(round("SELECT x SIMILAR TO 'a'"), "SELECT (x SimilarTo 'a')");
assert_eq!(round("SELECT x NOT GLOB 'a'"), "SELECT (Not (x Glob 'a'))");
}
#[test]
fn between_and_in_carry_their_negation_as_a_flag() {
assert_eq!(round("SELECT x BETWEEN 1 AND 2"), "SELECT (x BETWEEN 1 AND 2)");
assert_eq!(round("SELECT x NOT BETWEEN 1 AND 2"), "SELECT (NOT x BETWEEN 1 AND 2)");
assert_eq!(round("SELECT x IN (1, 2)"), "SELECT (x IN [1, 2])");
assert_eq!(round("SELECT x NOT IN (1, 2)"), "SELECT (NOT x IN [1, 2])");
}
#[test]
fn both_spellings_of_a_cast_are_the_same_node() {
assert_eq!(round("SELECT CAST(x AS BIGINT)"), "SELECT CAST(x AS BIGINT)");
assert_eq!(round("SELECT x::BIGINT"), "SELECT CAST(x AS BIGINT)");
assert_eq!(round("SELECT TRY_CAST(x AS BIGINT)"), "SELECT TRY_CAST(x AS BIGINT)");
assert_eq!(
round("SELECT x::DECIMAL(18, 3)"),
"SELECT CAST(x AS DECIMAL(18, 3))",
"the type is kept as text because parsing it is the type system's job"
);
}
#[test]
fn a_case_keeps_its_arms_in_order() {
assert_eq!(
round("SELECT CASE WHEN a THEN 1 WHEN b THEN 2 ELSE 3 END"),
"SELECT CASE - WHEN a THEN 1 WHEN b THEN 2 ELSE 3 END"
);
assert_eq!(
round("SELECT CASE x WHEN 1 THEN 'a' END"),
"SELECT CASE x WHEN 1 THEN 'a' ELSE - END",
"a simple case keeps the operand and a missing else is not an implicit null yet"
);
}
#[test]
fn a_field_access_and_a_method_call_are_ordinary_function_calls() {
assert_eq!(round("SELECT (f(x)).y"), "SELECT struct_extract(f(x), 'y')");
assert_eq!(round("SELECT a[1]"), "SELECT array_extract(a, 1)");
}
#[test]
fn an_aggregate_keeps_its_distinct() {
assert_eq!(round("SELECT count(*)"), "SELECT count(*)");
assert_eq!(round("SELECT count(DISTINCT x)"), "SELECT count(DISTINCT x)");
assert_eq!(round("SELECT count(ALL x)"), "SELECT count(x)");
assert_eq!(round("SELECT main.count(x)"), "SELECT main.count(x)");
}
#[test]
fn the_modifiers_hang_off_the_query_and_not_off_the_select() {
assert_eq!(
round("SELECT 1 UNION ALL SELECT 2 ORDER BY 1"),
"(SELECT 1 Union All SELECT 2) ORDER BY 1 Unstated Unstated"
);
assert_eq!(
round("SELECT a FROM t UNION SELECT b FROM u EXCEPT SELECT c FROM v"),
"((SELECT a FROM t Union Unstated SELECT b FROM u) Except Unstated SELECT c FROM v)",
"set operators are left associative"
);
assert_eq!(
round("SELECT 1 UNION SELECT 2 INTERSECT SELECT 3"),
"(SELECT 1 Union Unstated (SELECT 2 Intersect Unstated SELECT 3))",
"and intersect binds tighter than the other two"
);
}
#[test]
fn the_sort_and_limit_clauses_keep_what_was_written() {
assert_eq!(
round("SELECT a FROM t ORDER BY a"),
"SELECT a FROM t ORDER BY a Unstated Unstated"
);
assert_eq!(
round("SELECT a FROM t ORDER BY a DESC NULLS LAST"),
"SELECT a FROM t ORDER BY a Descending Last"
);
assert_eq!(round("SELECT a FROM t ORDER BY ALL"), "SELECT a FROM t ORDER BY ALL");
assert_eq!(round("SELECT a FROM t GROUP BY ALL"), "SELECT a FROM t GROUP BY ALL");
assert_eq!(round("SELECT a FROM t LIMIT 10 OFFSET 5"), "SELECT a FROM t LIMIT 10 OFFSET 5");
assert_eq!(round("SELECT a FROM t OFFSET 5 LIMIT 10"), "SELECT a FROM t LIMIT 10 OFFSET 5");
assert_eq!(round("SELECT a FROM t LIMIT 10%"), "SELECT a FROM t LIMIT 10%");
assert_eq!(round("SELECT a FROM t LIMIT ALL"), "SELECT a FROM t", "which is no limit");
}
#[test]
fn a_subquery_appears_in_both_places_it_can() {
assert_eq!(
round("SELECT * FROM (SELECT x FROM t) AS s"),
"SELECT * FROM (SELECT x FROM t) AS s"
);
assert_eq!(round("SELECT (SELECT 1)"), "SELECT (SELECT 1)");
}
#[test]
fn distinct_on_keeps_its_expressions() {
assert_eq!(round("SELECT DISTINCT a"), "SELECT DISTINCT a");
assert_eq!(round("SELECT ALL a"), "SELECT a", "which is the default written out");
assert_eq!(round("SELECT DISTINCT ON (a, b) a"), "SELECT DISTINCT ON (a, b) a");
}
#[test]
fn an_operator_the_dialect_does_not_name_is_kept_by_name() {
assert_eq!(round("SELECT a <=> b"), "SELECT (a <=> b)");
assert!(parse_ast("SELECT a foo b").is_err(), "a bare word is not an operator");
}
#[test]
fn a_script_is_a_list_of_statements() {
let ast = parse_ast("SELECT 1; SELECT 2;").unwrap();
assert_eq!(ast.statements.len(), 2);
let Statement::Query(second) = ast.statements[1];
assert_eq!(show_query(&ast, second), "SELECT 2");
}
#[test]
fn an_unsupported_construct_names_itself_and_what_was_written() {
let error = parse_ast("CREATE TABLE t (a INTEGER)").unwrap_err().to_string();
assert!(error.starts_with("Not implemented Error"), "{error}");
assert!(error.contains("CREATE TABLE t (a INTEGER)"), "{error}");
assert!(error.contains("CreateStatement"), "{error}");
}
#[test]
fn a_long_construct_is_cut_short_in_the_message() {
let query =
format!("CREATE TABLE t AS SELECT {} FROM u", "averylongcolumnname, ".repeat(8));
let error = parse_ast(&query).unwrap_err().to_string();
assert!(error.contains("..."), "{error}");
assert!(error.len() < 200, "{error}");
}
#[test]
fn the_transformer_never_panics_on_anything_the_matcher_accepts() {
for query in [
"SELECT",
"FROM t SELECT",
"SELECT * FROM t WHERE",
"SELECT ()",
"SELECT a FROM t GROUP BY ()",
] {
let answer = parse_ast(query);
if let Err(error) = answer {
let message = error.to_string();
assert!(
message.starts_with("Not implemented Error")
|| message.starts_with("Parser Error"),
"{query} failed with {message}"
);
}
}
}
#[test]
fn interning_means_a_name_written_twice_is_stored_once() {
let ast = parse_ast("SELECT a, a, a FROM t WHERE a = a").unwrap();
assert_eq!(ast.strings.iter().filter(|text| *text == "a").count(), 1);
}
}