#![allow(dead_code)]
use crate::ast::{
is_aggregate_name, ArithOp, CallClause, CallYield, CompareOp, Expr, Literal, MergeClause,
NodePattern, Pattern, PropAccess, QuantifierKind, QueryClause, QueryPart, RelDirection,
RelPattern, RemoveItem, ReturnExpr, ReturnItem, ReturnTail, SetItem, SortDir, Statement, Tail,
UnwindClause, UnwindSource, WithClause, WithExpr,
};
use crate::error::QueryError;
use crate::generated::cypherparser::{
AddSubExpressionContext, AndExpressionContext, AndExpressionContextAttrs, AtomContext,
AtomContextAttrs, AtomicExpressionContext, AtomicExpressionContextAll,
AtomicExpressionContextAttrs, BoolLitContext, BoolLitContextAttrs, CaseExpressionContext,
CharLitContext, CharLitContextAttrs, ComparisonExpressionContext,
ComparisonExpressionContextAttrs, ComparisonSignsContextAll, ComparisonSignsContextAttrs,
CountAllContext, CreateIndexStContext, CreateIndexStContextAttrs, CreateStContext,
CreateStContextAttrs, DeleteStContext, DeleteStContextAttrs, ExplainStContext,
ExplainStContextAttrs, ExpressionChainContextAttrs, ExpressionContext, ExpressionContextAttrs,
FilterExpressionContext, FilterExpressionContextAttrs, FilterWithContext,
FilterWithContextAttrs, FunctionInvocationContext, FunctionInvocationContextAttrs,
InExpressionContextAttrs, InvocationNameContextAll, InvocationNameContextAttrs,
LhsContextAttrs, LimitStContextAttrs, ListComprehensionContext, ListComprehensionContextAttrs,
ListExpressionContextAll, ListExpressionContextAttrs, ListLitContext, ListLitContextAttrs,
LiteralContext, LiteralContextAttrs, MapLitContext, MapLitContextAttrs, MapPairContextAttrs,
MatchStContext, MatchStContextAttrs, MergeActionContextAll, MergeActionContextAttrs,
MergeStContext, MergeStContextAttrs, MultDivExpressionContext, MultiPartQContext,
MultiPartQContextAttrs, NameContextAll, NameContextAttrs, NodeLabelsContextAttrs,
NodePatternContext, NodePatternContextAttrs, NotExpressionContext, NotExpressionContextAttrs,
NullExpressionContextAttrs, NumLitContext, NumLitContextAll, NumLitContextAttrs,
OrderItemContextAttrs, OrderStContext, OrderStContextAttrs, ParameterContext,
ParameterContextAttrs, ParenExpressionChainContextAll, ParenExpressionChainContextAttrs,
ParenthesizedExpressionContext, ParenthesizedExpressionContextAttrs,
PatternComprehensionContext, PatternComprehensionContextAttrs, PatternContextAttrs,
PatternElemChainContextAttrs, PatternElemContext, PatternElemContextAttrs,
PatternPartContextAttrs, PatternWhereContextAttrs, PowerExpressionContext,
PowerExpressionContextAttrs, ProjectionBodyContext, ProjectionBodyContextAttrs,
ProjectionItemContextAttrs, ProjectionItemsContextAttrs, PropertiesContextAll,
PropertiesContextAttrs, PropertyExpressionContext, PropertyExpressionContextAttrs,
PropertyOrLabelExpressionContext, PropertyOrLabelExpressionContextAttrs, QueryCallStContextAll,
QueryCallStContextAttrs, ReadingStatementContextAll, ReadingStatementContextAttrs,
RegularQueryContext, RegularQueryContextAttrs, RelationDetailContext,
RelationDetailContextAttrs, RelationshipPatternContext, RelationshipPatternContextAttrs,
RelationshipTypesContextAttrs, RelationshipsChainPatternContext,
RelationshipsChainPatternContextAttrs, RemoveItemContextAll, RemoveItemContextAttrs,
RemoveStContext, RemoveStContextAttrs, ReturnStContext, ReturnStContextAttrs,
SetItemContextAll, SetItemContextAttrs, SetStContext, SetStContextAttrs,
ShortestPathWrapperContextAttrs, SinglePartQContext, SinglePartQContextAttrs,
SkipStContextAttrs, StandaloneCallContext, StandaloneCallContextAttrs,
StringExpPrefixContextAll, StringExpPrefixContextAttrs, StringExpressionContextAll,
StringExpressionContextAttrs, StringListNullExpressionContext,
StringListNullExpressionContextAttrs, StringLitContext, StringLitContextAttrs,
SubqueryExistContext, SubqueryExistContextAttrs, SymbolContextAll, SymbolContextAttrs,
UnaryAddSubExpressionContext, UnaryAddSubExpressionContextAttrs, UnionStContextAttrs,
UnwindStContext, UnwindStContextAttrs, UpdatingStatementContextAll,
UpdatingStatementContextAttrs, WhereContextAttrs, WithStContext, WithStContextAttrs,
XorExpressionContext, XorExpressionContextAttrs, YieldItemContextAttrs, YieldItemsContextAll,
YieldItemsContextAttrs,
};
use crate::generated::cypherparservisitor::CypherParserVisitorCompat;
use crate::parse_helpers::{
group_into_linear_patterns, parse_int_literal, parse_rel_range, unescape_string,
validate_named_path_pattern, validate_shortest_path_pattern,
};
use antlr4rust::parser_rule_context::ParserRuleContext;
use antlr4rust::token::Token;
use antlr4rust::tree::{ParseTree, ParseTreeVisitorCompat, Tree};
use std::rc::Rc;
#[derive(Debug, Default)]
pub(crate) enum AstNode {
#[default]
None,
Literal(Literal),
NodePattern(NodePattern),
RelPattern(RelPattern),
Pattern(Pattern),
QueryParts(Vec<QueryPart>),
ReturnExpr(ReturnExpr),
ReturnClause(ParsedReturnClause),
WithClause(WithClause),
UnwindClause(UnwindClause),
SetItems(Vec<SetItem>),
DeleteItems(ParsedDelete),
RemoveItems(Vec<RemoveItem>),
CreatePatterns(Vec<Pattern>),
MergeClause(MergeClause),
Statement(Statement),
Err(QueryError),
}
#[derive(Debug)]
pub(crate) struct ParsedDelete {
pub items: Vec<ReturnExpr>,
pub detach: bool,
}
#[derive(Debug)]
pub(crate) struct ParsedReturnClause {
pub tail: Tail,
pub order_by: Option<Vec<(ReturnExpr, SortDir)>>,
pub skip: Option<ReturnExpr>,
pub limit: Option<ReturnExpr>,
}
macro_rules! ast_node_into {
($name:ident, $variant:ident, $ty:ty) => {
fn $name(self) -> Result<$ty, QueryError> {
match self {
AstNode::$variant(v) => Ok(v),
AstNode::Err(e) => Err(e),
other => unreachable!("expected AstNode::{}, got {other:?}", stringify!($variant)),
}
}
};
}
impl AstNode {
ast_node_into!(into_literal, Literal, Literal);
ast_node_into!(into_node_pattern, NodePattern, NodePattern);
ast_node_into!(into_rel_pattern, RelPattern, RelPattern);
ast_node_into!(into_pattern, Pattern, Pattern);
ast_node_into!(into_query_parts, QueryParts, Vec<QueryPart>);
ast_node_into!(into_return_expr, ReturnExpr, ReturnExpr);
fn into_return_expr_lenient(self) -> Result<ReturnExpr, QueryError> {
match self {
AstNode::Literal(l) => Ok(ReturnExpr::Lit(l)),
AstNode::ReturnExpr(e) => Ok(e),
AstNode::Err(e) => Err(e),
other => {
unreachable!("expected AstNode::Literal or AstNode::ReturnExpr, got {other:?}")
}
}
}
ast_node_into!(into_return_clause, ReturnClause, ParsedReturnClause);
ast_node_into!(into_with_clause, WithClause, WithClause);
ast_node_into!(into_unwind_clause, UnwindClause, UnwindClause);
ast_node_into!(into_set_items, SetItems, Vec<SetItem>);
ast_node_into!(into_delete_items, DeleteItems, ParsedDelete);
ast_node_into!(into_remove_items, RemoveItems, Vec<RemoveItem>);
ast_node_into!(into_create_patterns, CreatePatterns, Vec<Pattern>);
ast_node_into!(into_merge_clause, MergeClause, MergeClause);
ast_node_into!(into_statement, Statement, Statement);
}
pub(crate) struct AstBuilder {
result: AstNode,
}
impl AstBuilder {
pub(crate) fn new() -> Self {
AstBuilder {
result: AstNode::default(),
}
}
}
impl<'input> ParseTreeVisitorCompat<'input> for AstBuilder {
type Node = crate::generated::cypherparser::CypherParserContextType;
type Return = AstNode;
fn temp_result(&mut self) -> &mut Self::Return {
&mut self.result
}
}
impl<'input> CypherParserVisitorCompat<'input> for AstBuilder {
fn visit_literal(&mut self, ctx: &LiteralContext<'input>) -> Self::Return {
if ctx.NULL_W().is_some() {
return AstNode::Literal(Literal::Null);
}
self.visit_children(ctx)
}
fn visit_boolLit(&mut self, ctx: &BoolLitContext<'input>) -> Self::Return {
AstNode::Literal(Literal::Bool(ctx.TRUE().is_some()))
}
fn visit_numLit(&mut self, ctx: &NumLitContext<'input>) -> Self::Return {
let text = ctx
.DIGIT()
.expect("numLit context always has a DIGIT token")
.get_text();
match parse_num_lit_text(&text) {
Ok(lit) => AstNode::Literal(lit),
Err(e) => AstNode::Err(e),
}
}
fn visit_stringLit(&mut self, ctx: &StringLitContext<'input>) -> Self::Return {
let text = ctx
.STRING_LITERAL()
.expect("stringLit context always has a STRING_LITERAL token")
.get_text();
match unescape_string(&text[1..text.len() - 1]) {
Ok(s) => AstNode::Literal(Literal::String(s)),
Err(e) => AstNode::Err(e),
}
}
fn visit_charLit(&mut self, ctx: &CharLitContext<'input>) -> Self::Return {
let text = ctx
.CHAR_LITERAL()
.expect("charLit context always has a CHAR_LITERAL token")
.get_text();
match unescape_string(&text[1..text.len() - 1]) {
Ok(s) => AstNode::Literal(Literal::String(s)),
Err(e) => AstNode::Err(e),
}
}
fn visit_nodePattern(&mut self, ctx: &NodePatternContext<'input>) -> Self::Return {
match self.build_node_pattern(ctx) {
Ok(n) => AstNode::NodePattern(n),
Err(e) => AstNode::Err(e),
}
}
fn visit_relationDetail(&mut self, ctx: &RelationDetailContext<'input>) -> Self::Return {
match self.build_rel_detail(ctx) {
Ok(r) => AstNode::RelPattern(r),
Err(e) => AstNode::Err(e),
}
}
fn visit_relationshipPattern(
&mut self,
ctx: &RelationshipPatternContext<'input>,
) -> Self::Return {
match self.build_relationship_pattern(ctx) {
Ok(r) => AstNode::RelPattern(r),
Err(e) => AstNode::Err(e),
}
}
fn visit_patternElem(&mut self, ctx: &PatternElemContext<'input>) -> Self::Return {
match self.build_pattern_elem(ctx) {
Ok(p) => AstNode::Pattern(p),
Err(e) => AstNode::Err(e),
}
}
fn visit_matchSt(&mut self, ctx: &MatchStContext<'input>) -> Self::Return {
match self.build_match_st(ctx) {
Ok(parts) => AstNode::QueryParts(parts),
Err(e) => AstNode::Err(e),
}
}
fn visit_expression(&mut self, ctx: &ExpressionContext<'input>) -> Self::Return {
let mut operands = ctx.xorExpression_all().into_iter();
let mut lhs = match self
.visit(
&*operands
.next()
.expect("expression has at least one xorExpression"),
)
.into_return_expr()
{
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
for rhs_ctx in operands {
let rhs = match self.visit(&*rhs_ctx).into_return_expr() {
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
lhs = ReturnExpr::Or(Box::new(lhs), Box::new(rhs));
}
AstNode::ReturnExpr(lhs)
}
fn visit_xorExpression(&mut self, ctx: &XorExpressionContext<'input>) -> Self::Return {
let mut operands = ctx.andExpression_all().into_iter();
let mut lhs = match self
.visit(
&*operands
.next()
.expect("xorExpression has at least one andExpression"),
)
.into_return_expr()
{
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
for rhs_ctx in operands {
let rhs = match self.visit(&*rhs_ctx).into_return_expr() {
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
lhs = ReturnExpr::Xor(Box::new(lhs), Box::new(rhs));
}
AstNode::ReturnExpr(lhs)
}
fn visit_andExpression(&mut self, ctx: &AndExpressionContext<'input>) -> Self::Return {
let mut operands = ctx.notExpression_all().into_iter();
let mut lhs = match self
.visit(
&*operands
.next()
.expect("andExpression has at least one notExpression"),
)
.into_return_expr()
{
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
for rhs_ctx in operands {
let rhs = match self.visit(&*rhs_ctx).into_return_expr() {
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
lhs = ReturnExpr::And(Box::new(lhs), Box::new(rhs));
}
AstNode::ReturnExpr(lhs)
}
fn visit_notExpression(&mut self, ctx: &NotExpressionContext<'input>) -> Self::Return {
let inner = ctx
.comparisonExpression()
.expect("notExpression always has a comparisonExpression");
match self.visit(&*inner).into_return_expr() {
Ok(mut expr) => {
for _ in ctx.NOT_all() {
expr = ReturnExpr::Not(Box::new(expr));
}
AstNode::ReturnExpr(expr)
}
Err(e) => AstNode::Err(e),
}
}
fn visit_comparisonExpression(
&mut self,
ctx: &ComparisonExpressionContext<'input>,
) -> Self::Return {
match self.build_comparison_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_stringListNullExpression(
&mut self,
ctx: &StringListNullExpressionContext<'input>,
) -> Self::Return {
match self.build_string_list_null_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_addSubExpression(&mut self, ctx: &AddSubExpressionContext<'input>) -> Self::Return {
match self.build_add_sub_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_multDivExpression(&mut self, ctx: &MultDivExpressionContext<'input>) -> Self::Return {
match self.build_mult_div_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_powerExpression(&mut self, ctx: &PowerExpressionContext<'input>) -> Self::Return {
let mut operands = ctx.unaryAddSubExpression_all().into_iter();
let mut lhs = match self
.visit(
&*operands
.next()
.expect("powerExpression has at least one unaryAddSubExpression"),
)
.into_return_expr()
{
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
for rhs_ctx in operands {
let rhs = match self.visit(&*rhs_ctx).into_return_expr() {
Ok(e) => e,
Err(e) => return AstNode::Err(e),
};
lhs = ReturnExpr::Arith(Box::new(lhs), ArithOp::Pow, Box::new(rhs));
}
AstNode::ReturnExpr(lhs)
}
fn visit_unaryAddSubExpression(
&mut self,
ctx: &UnaryAddSubExpressionContext<'input>,
) -> Self::Return {
match self.build_unary_add_sub_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_atomicExpression(&mut self, ctx: &AtomicExpressionContext<'input>) -> Self::Return {
match self.build_atomic_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_propertyOrLabelExpression(
&mut self,
ctx: &PropertyOrLabelExpressionContext<'input>,
) -> Self::Return {
match self.build_property_or_label_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_propertyExpression(
&mut self,
ctx: &PropertyExpressionContext<'input>,
) -> Self::Return {
match self.build_property_expression(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_atom(&mut self, ctx: &AtomContext<'input>) -> Self::Return {
match self.build_atom(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_parenthesizedExpression(
&mut self,
ctx: &ParenthesizedExpressionContext<'input>,
) -> Self::Return {
let inner = ctx
.expression()
.expect("parenthesizedExpression always has an expression");
self.visit(&*inner)
}
fn visit_functionInvocation(
&mut self,
ctx: &FunctionInvocationContext<'input>,
) -> Self::Return {
match self.build_function_invocation(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_parameter(&mut self, ctx: &ParameterContext<'input>) -> Self::Return {
match self.build_parameter(ctx) {
Ok(expr) => AstNode::ReturnExpr(expr),
Err(e) => AstNode::Err(e),
}
}
fn visit_countAll(&mut self, _ctx: &CountAllContext<'input>) -> Self::Return {
AstNode::ReturnExpr(ReturnExpr::CountStar)
}
fn visit_returnSt(&mut self, ctx: &ReturnStContext<'input>) -> Self::Return {
let body_ctx = ctx
.projectionBody()
.expect("returnSt always has a projectionBody");
match self.build_projection_body(&body_ctx) {
Ok(c) => AstNode::ReturnClause(c),
Err(e) => AstNode::Err(e),
}
}
fn visit_withSt(&mut self, ctx: &WithStContext<'input>) -> Self::Return {
match self.build_with_clause(ctx) {
Ok(c) => AstNode::WithClause(c),
Err(e) => AstNode::Err(e),
}
}
fn visit_unwindSt(&mut self, ctx: &UnwindStContext<'input>) -> Self::Return {
match self.build_unwind_st(ctx) {
Ok(c) => AstNode::UnwindClause(c),
Err(e) => AstNode::Err(e),
}
}
fn visit_setSt(&mut self, ctx: &SetStContext<'input>) -> Self::Return {
match self.build_set_st(ctx) {
Ok(items) => AstNode::SetItems(items),
Err(e) => AstNode::Err(e),
}
}
fn visit_deleteSt(&mut self, ctx: &DeleteStContext<'input>) -> Self::Return {
match self.build_delete_st(ctx) {
Ok(d) => AstNode::DeleteItems(d),
Err(e) => AstNode::Err(e),
}
}
fn visit_removeSt(&mut self, ctx: &RemoveStContext<'input>) -> Self::Return {
match self.build_remove_st(ctx) {
Ok(items) => AstNode::RemoveItems(items),
Err(e) => AstNode::Err(e),
}
}
fn visit_createSt(&mut self, ctx: &CreateStContext<'input>) -> Self::Return {
match self.build_create_st(ctx) {
Ok(patterns) => AstNode::CreatePatterns(patterns),
Err(e) => AstNode::Err(e),
}
}
fn visit_mergeSt(&mut self, ctx: &MergeStContext<'input>) -> Self::Return {
match self.build_merge_st(ctx) {
Ok(c) => AstNode::MergeClause(c),
Err(e) => AstNode::Err(e),
}
}
fn visit_singlePartQ(&mut self, ctx: &SinglePartQContext<'input>) -> Self::Return {
match self.build_single_part_q(ctx) {
Ok(s) => AstNode::Statement(s),
Err(e) => AstNode::Err(e),
}
}
fn visit_multiPartQ(&mut self, ctx: &MultiPartQContext<'input>) -> Self::Return {
match self.build_multi_part_q(ctx) {
Ok(s) => AstNode::Statement(s),
Err(e) => AstNode::Err(e),
}
}
fn visit_regularQuery(&mut self, ctx: &RegularQueryContext<'input>) -> Self::Return {
match self.build_regular_query(ctx) {
Ok(s) => AstNode::Statement(s),
Err(e) => AstNode::Err(e),
}
}
fn visit_standaloneCall(&mut self, ctx: &StandaloneCallContext<'input>) -> Self::Return {
match self.build_standalone_call(ctx) {
Ok(s) => AstNode::Statement(s),
Err(e) => AstNode::Err(e),
}
}
fn visit_explainSt(&mut self, ctx: &ExplainStContext<'input>) -> Self::Return {
match self.build_explain_st(ctx) {
Ok(s) => AstNode::Statement(s),
Err(e) => AstNode::Err(e),
}
}
fn visit_createIndexSt(&mut self, ctx: &CreateIndexStContext<'input>) -> Self::Return {
match self.build_create_index_st(ctx) {
Ok(s) => AstNode::Statement(s),
Err(e) => AstNode::Err(e),
}
}
fn visit_listLit(&mut self, ctx: &ListLitContext<'input>) -> Self::Return {
let mut items = Vec::new();
if let Some(chain_ctx) = ctx.expressionChain() {
for expr_ctx in chain_ctx.expression_all() {
match self.visit(&*expr_ctx).into_return_expr() {
Ok(e) => items.push(e),
Err(e) => return AstNode::Err(e),
}
}
}
AstNode::ReturnExpr(ReturnExpr::ListLit(items))
}
fn visit_mapLit(&mut self, ctx: &MapLitContext<'input>) -> Self::Return {
let mut items = Vec::new();
for pair_ctx in ctx.mapPair_all() {
let name_ctx = pair_ctx.name().expect("mapPair always has a name");
let expr_ctx = pair_ctx
.expression()
.expect("mapPair always has an expression");
let value = match self.visit(&*expr_ctx).into_return_expr() {
Ok(v) => v,
Err(e) => return AstNode::Err(e),
};
items.push((name_text(&name_ctx), value));
}
AstNode::ReturnExpr(ReturnExpr::MapLit(items))
}
}
fn symbol_text(ctx: &SymbolContextAll) -> String {
match ctx.ESC_LITERAL() {
Some(t) => {
let text = t.get_text();
text[1..text.len() - 1].to_string()
}
None => ctx.get_text(),
}
}
fn name_text(ctx: &NameContextAll) -> String {
match ctx.symbol() {
Some(s) => symbol_text(&s),
None => ctx.get_text(),
}
}
fn parse_num_lit_text(text: &str) -> Result<Literal, QueryError> {
let unsigned = text.strip_prefix('-').unwrap_or(text);
let is_hex_or_octal = unsigned
.as_bytes()
.get(1)
.is_some_and(|b| matches!(b, b'x' | b'X' | b'o' | b'O'))
&& unsigned.starts_with('0');
let is_float = !is_hex_or_octal
&& (text.contains('.')
|| text.ends_with(['f', 'F', 'd', 'D'])
|| text
.rfind(['e', 'E'])
.is_some_and(|i| text[..i].chars().all(|c| c.is_ascii_digit() || c == '-')));
if is_float {
let f: f64 = text
.parse()
.map_err(|_| QueryError::Syntax(format!("invalid float literal '{text}'")))?;
if f.is_infinite() {
Err(QueryError::Syntax(format!(
"float literal '{text}' is too large to represent"
)))
} else {
Ok(Literal::Float(f))
}
} else {
parse_int_literal(text).map(Literal::Int)
}
}
fn compare_sign(ctx: &ComparisonSignsContextAll) -> CompareOp {
if ctx.LE().is_some() {
CompareOp::Le
} else if ctx.GE().is_some() {
CompareOp::Ge
} else if ctx.GT().is_some() {
CompareOp::Gt
} else if ctx.LT().is_some() {
CompareOp::Lt
} else if ctx.NOT_EQUAL().is_some() {
CompareOp::Ne
} else {
CompareOp::Eq
}
}
fn string_exp_op(ctx: &StringExpPrefixContextAll) -> CompareOp {
if ctx.STARTS().is_some() {
CompareOp::StartsWith
} else if ctx.ENDS().is_some() {
CompareOp::EndsWith
} else {
CompareOp::Contains
}
}
fn invocation_name_text(ctx: &InvocationNameContextAll) -> String {
ctx.symbol_all()
.iter()
.map(|s| symbol_text(s))
.collect::<Vec<_>>()
.join(".")
}
fn bare_num_lit<'i>(
ctx: &AtomicExpressionContextAll<'i>,
) -> Option<std::rc::Rc<NumLitContextAll<'i>>> {
if !ctx.listExpression_all().is_empty() {
return None;
}
let prop_or_label = ctx.propertyOrLabelExpression()?;
if prop_or_label.nodeLabels().is_some() {
return None;
}
let prop_expr = prop_or_label.propertyExpression()?;
if !prop_expr.name_all().is_empty() {
return None;
}
prop_expr.atom()?.literal()?.numLit()
}
fn list_expr_bound_is_before_range(ctx: &ListExpressionContextAll) -> bool {
let mut seen_range = false;
for child in ctx.get_children() {
match child.get_text().as_str() {
"[" | "]" => continue,
".." => seen_range = true,
_ => return !seen_range,
}
}
unreachable!("listExpression slice form always has exactly one expression child")
}
impl AstBuilder {
fn build_properties(
&mut self,
ctx: Option<Rc<PropertiesContextAll>>,
) -> Result<Vec<(String, ReturnExpr)>, QueryError> {
let Some(ctx) = ctx else {
return Ok(Vec::new());
};
let Some(map_ctx) = ctx.mapLit() else {
return Err(QueryError::Syntax(
"a parameter can't be used as a pattern's whole properties map".into(),
));
};
let expr = self.visit(&*map_ctx).into_return_expr()?;
let ReturnExpr::MapLit(items) = expr else {
unreachable!("mapLit always builds a ReturnExpr::MapLit");
};
Ok(items)
}
fn build_node_pattern(&mut self, ctx: &NodePatternContext) -> Result<NodePattern, QueryError> {
let var = ctx.symbol().map(|s| symbol_text(&s));
let labels = ctx
.nodeLabels()
.map(|nl| nl.name_all().iter().map(|n| name_text(n)).collect())
.unwrap_or_default();
let has_explicit_props = ctx.properties().is_some();
let props = self.build_properties(ctx.properties())?;
Ok(NodePattern {
var,
labels,
props,
has_explicit_props,
})
}
fn build_rel_detail(&mut self, ctx: &RelationDetailContext) -> Result<RelPattern, QueryError> {
let var = ctx.symbol().map(|s| symbol_text(&s));
let rel_types = ctx
.relationshipTypes()
.map(|rt| rt.name_all().iter().map(|n| name_text(n)).collect())
.unwrap_or_default();
let props = self.build_properties(ctx.properties())?;
let hop_range = ctx
.rangeLit()
.map(|r| parse_rel_range(&r.get_text()))
.transpose()?;
Ok(RelPattern {
var,
rel_types,
props,
direction: RelDirection::Either,
hop_range,
capture_path_segment: false,
rel_list_var: None,
})
}
fn build_relationship_pattern(
&mut self,
ctx: &RelationshipPatternContext,
) -> Result<RelPattern, QueryError> {
let mut rel = match ctx.relationDetail() {
Some(rd) => self.visit(&*rd).into_rel_pattern()?,
None => RelPattern {
var: None,
rel_types: Vec::new(),
props: Vec::new(),
direction: RelDirection::Either,
hop_range: None,
capture_path_segment: false,
rel_list_var: None,
},
};
rel.direction = match (ctx.LT().is_some(), ctx.GT().is_some()) {
(true, false) => RelDirection::Left,
(false, true) => RelDirection::Right,
(true, true) | (false, false) => RelDirection::Either,
};
Ok(rel)
}
fn build_pattern_elem(&mut self, ctx: &PatternElemContext) -> Result<Pattern, QueryError> {
if ctx.LPAREN().is_some() || !ctx.qppElemChain_all().is_empty() {
return Err(QueryError::Syntax(
"quantified path patterns aren't supported yet".into(),
));
}
let node_ctx = ctx
.nodePattern()
.expect("patternElem always starts with a nodePattern in the non-QPP alternative");
let start = self.visit(&*node_ctx).into_node_pattern()?;
let mut hops = Vec::new();
for chain in ctx.patternElemChain_all() {
let rel_ctx = chain
.relationshipPattern()
.expect("patternElemChain always has a relationshipPattern");
let node_ctx = chain
.nodePattern()
.expect("patternElemChain always has a nodePattern");
let rel = self.visit(&*rel_ctx).into_rel_pattern()?;
let node = self.visit(&*node_ctx).into_node_pattern()?;
hops.push((rel, node));
}
Ok(Pattern { start, hops })
}
fn build_relationships_chain_pattern(
&mut self,
ctx: &RelationshipsChainPatternContext,
) -> Result<Pattern, QueryError> {
let node_ctx = ctx
.nodePattern()
.expect("relationshipsChainPattern always has a nodePattern");
let start = self.visit(&*node_ctx).into_node_pattern()?;
let mut hops = Vec::new();
for chain in ctx.patternElemChain_all() {
let rel_ctx = chain
.relationshipPattern()
.expect("patternElemChain always has a relationshipPattern");
let node_ctx = chain
.nodePattern()
.expect("patternElemChain always has a nodePattern");
let rel = self.visit(&*rel_ctx).into_rel_pattern()?;
let node = self.visit(&*node_ctx).into_node_pattern()?;
hops.push((rel, node));
}
Ok(Pattern { start, hops })
}
fn build_match_st(&mut self, ctx: &MatchStContext) -> Result<Vec<QueryPart>, QueryError> {
let optional = ctx.OPTIONAL().is_some();
let pw = ctx
.patternWhere()
.expect("matchSt always has a patternWhere");
let where_clause = match pw.where_() {
Some(where_ctx) => {
let expr_ctx = where_ctx
.expression()
.expect("where always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
Some(return_expr_to_expr(expr)?)
}
None => None,
};
let pattern_ctx = pw.pattern().expect("patternWhere always has a pattern");
let mut path_var = None;
let mut shortest_path = false;
let mut patterns = Vec::new();
for (i, part) in pattern_ctx.patternPart_all().into_iter().enumerate() {
let pattern = match part.shortestPathWrapper() {
Some(sp_ctx) => {
if i != 0 {
return Err(QueryError::Syntax(
"shortestPath() must be the first (and only) comma-separated pattern"
.into(),
));
}
shortest_path = true;
let elem_ctx = sp_ctx
.patternElem()
.expect("shortestPathWrapper always has a patternElem");
self.visit(&*elem_ctx).into_pattern()?
}
None => {
let elem_ctx = part.patternElem().expect(
"patternPart always has a patternElem when shortestPathWrapper is absent",
);
self.visit(&*elem_ctx).into_pattern()?
}
};
patterns.push(pattern);
if part.ASSIGN().is_some() {
if path_var.is_some() {
return Err(QueryError::Syntax(
"at most one comma-separated pattern part can have a named-path variable"
.into(),
));
}
let symbol_ctx = part
.symbol()
.expect("patternPart with ASSIGN always has a symbol");
path_var = Some(symbol_text(&symbol_ctx));
}
}
let groups = group_into_linear_patterns(patterns)?;
if groups.len() > 1 && (shortest_path || path_var.is_some()) {
return Err(QueryError::Syntax(
"a named path/shortestPath() can't span a comma-separated cross join".into(),
));
}
if shortest_path {
validate_shortest_path_pattern(&groups[0])?;
} else if path_var.is_some() {
validate_named_path_pattern(&groups[0])?;
}
let last = groups.len() - 1;
Ok(groups
.into_iter()
.enumerate()
.map(|(i, pattern)| QueryPart {
optional,
path_var: if i == 0 { path_var.clone() } else { None },
shortest_path: i == 0 && shortest_path,
pattern,
where_clause: if i == last {
where_clause.clone()
} else {
None
},
with: None,
})
.collect())
}
fn build_comparison_expression(
&mut self,
ctx: &ComparisonExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let mut operands = Vec::new();
for operand_ctx in ctx.stringListNullExpression_all() {
operands.push(self.visit(&*operand_ctx).into_return_expr()?);
}
let mut ops = Vec::new();
for sign_ctx in ctx.comparisonSigns_all() {
ops.push(compare_sign(&sign_ctx));
}
if ops.is_empty() {
return Ok(operands
.into_iter()
.next()
.expect("comparisonExpression has at least one stringListNullExpression"));
}
let mut pairs = operands.windows(2).zip(&ops).map(|(pair, op)| {
ReturnExpr::Compare(Box::new(pair[0].clone()), *op, Box::new(pair[1].clone()))
});
let mut acc = pairs
.next()
.expect("a comparison chain has at least one pair");
for next in pairs {
acc = ReturnExpr::And(Box::new(acc), Box::new(next));
}
Ok(acc)
}
fn build_add_sub_expression(
&mut self,
ctx: &AddSubExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let mut children = ctx.get_children();
let mut lhs = self
.visit(
&*children
.next()
.expect("addSubExpression has at least one multDivExpression"),
)
.into_return_expr()?;
while let Some(op_node) = children.next() {
let op = match op_node.get_text().as_str() {
"+" => ArithOp::Add,
"-" => ArithOp::Sub,
other => unreachable!("unexpected addSubExpression operator {other:?}"),
};
let rhs_node = children
.next()
.expect("addSubExpression operator has a following multDivExpression");
let rhs = self.visit(&*rhs_node).into_return_expr()?;
lhs = ReturnExpr::Arith(Box::new(lhs), op, Box::new(rhs));
}
Ok(lhs)
}
fn build_mult_div_expression(
&mut self,
ctx: &MultDivExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let mut children = ctx.get_children();
let mut lhs = self
.visit(
&*children
.next()
.expect("multDivExpression has at least one powerExpression"),
)
.into_return_expr()?;
while let Some(op_node) = children.next() {
let op = match op_node.get_text().as_str() {
"*" => ArithOp::Mul,
"/" => ArithOp::Div,
"%" => ArithOp::Mod,
other => unreachable!("unexpected multDivExpression operator {other:?}"),
};
let rhs_node = children
.next()
.expect("multDivExpression operator has a following powerExpression");
let rhs = self.visit(&*rhs_node).into_return_expr()?;
lhs = ReturnExpr::Arith(Box::new(lhs), op, Box::new(rhs));
}
Ok(lhs)
}
fn build_unary_add_sub_expression(
&mut self,
ctx: &UnaryAddSubExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let atomic_ctx = ctx
.atomicExpression()
.expect("unaryAddSubExpression always has an atomicExpression");
if ctx.SUB().is_some() {
if let Some(numlit_ctx) = bare_num_lit(&atomic_ctx) {
let text = numlit_ctx
.DIGIT()
.expect("numLit context always has a DIGIT token")
.get_text();
return parse_num_lit_text(&format!("-{text}")).map(ReturnExpr::Lit);
}
let operand = self.visit(&*atomic_ctx).into_return_expr()?;
return Ok(ReturnExpr::Neg(Box::new(operand)));
}
self.visit(&*atomic_ctx).into_return_expr()
}
fn build_atomic_expression(
&mut self,
ctx: &AtomicExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let base_ctx = ctx
.propertyOrLabelExpression()
.expect("atomicExpression always has a propertyOrLabelExpression");
let mut base = self.visit(&*base_ctx).into_return_expr()?;
for l in ctx.listExpression_all() {
base = self.build_list_expression(&l, base)?;
}
Ok(base)
}
fn build_string_list_null_expression(
&mut self,
ctx: &StringListNullExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let base_ctx = ctx
.addSubExpression()
.expect("stringListNullExpression always has an addSubExpression");
let base = self.visit(&*base_ctx).into_return_expr()?;
if let Some(s) = ctx.stringExpression() {
return self.build_string_expression(&s, base);
}
if let Some(i) = ctx.inExpression() {
let rhs_ctx = i
.addSubExpression()
.expect("inExpression always has an addSubExpression");
let rhs = self.visit(&*rhs_ctx).into_return_expr()?;
return Ok(ReturnExpr::In(Box::new(base), Box::new(rhs)));
}
let Some(n) = ctx.nullExpression() else {
return Ok(base);
};
Ok(if n.NOT().is_some() {
ReturnExpr::Not(Box::new(ReturnExpr::IsNull(Box::new(base))))
} else {
ReturnExpr::IsNull(Box::new(base))
})
}
fn build_string_expression(
&mut self,
ctx: &StringExpressionContextAll,
base: ReturnExpr,
) -> Result<ReturnExpr, QueryError> {
let prefix_ctx = ctx
.stringExpPrefix()
.expect("stringExpression always has a stringExpPrefix");
let op = string_exp_op(&prefix_ctx);
let rhs_ctx = ctx
.addSubExpression()
.expect("stringExpression always has an addSubExpression");
let rhs = self.visit(&*rhs_ctx).into_return_expr()?;
Ok(ReturnExpr::Compare(Box::new(base), op, Box::new(rhs)))
}
fn build_list_expression(
&mut self,
ctx: &ListExpressionContextAll,
base: ReturnExpr,
) -> Result<ReturnExpr, QueryError> {
let exprs = ctx.expression_all();
if ctx.RANGE().is_some() {
let (start, end) = match exprs.len() {
0 => (None, None),
1 => {
let before_range = list_expr_bound_is_before_range(ctx);
let e = self.visit(&*exprs[0].clone()).into_return_expr()?;
if before_range {
(Some(Box::new(e)), None)
} else {
(None, Some(Box::new(e)))
}
}
2 => {
let start = self.visit(&*exprs[0].clone()).into_return_expr()?;
let end = self.visit(&*exprs[1].clone()).into_return_expr()?;
(Some(Box::new(start)), Some(Box::new(end)))
}
n => unreachable!("listExpression slice form has {n} expressions, expected 0-2"),
};
return Ok(ReturnExpr::Slice(Box::new(base), start, end));
}
let index_ctx = exprs
.into_iter()
.next()
.expect("non-slice listExpression always has exactly one expression");
let index = self.visit(&*index_ctx).into_return_expr()?;
Ok(ReturnExpr::Index(Box::new(base), Box::new(index)))
}
fn build_property_or_label_expression(
&mut self,
ctx: &PropertyOrLabelExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let prop_ctx = ctx
.propertyExpression()
.expect("propertyOrLabelExpression always has a propertyExpression");
let base = self.visit(&*prop_ctx).into_return_expr()?;
let Some(labels_ctx) = ctx.nodeLabels() else {
return Ok(base);
};
let ReturnExpr::Var(var) = base else {
return Err(QueryError::Syntax(
"a label check (`x:Label`) only applies to a bare variable".into(),
));
};
let labels = labels_ctx.name_all().iter().map(|n| name_text(n)).collect();
Ok(ReturnExpr::HasLabel(var, labels))
}
fn build_property_expression(
&mut self,
ctx: &PropertyExpressionContext,
) -> Result<ReturnExpr, QueryError> {
let atom_ctx = ctx.atom().expect("propertyExpression always has an atom");
let base = self.visit(&*atom_ctx).into_return_expr()?;
let mut names = ctx.name_all().into_iter();
let Some(first) = names.next() else {
return Ok(base);
};
let mut expr = match base {
ReturnExpr::Var(var) => ReturnExpr::Prop(PropAccess {
var,
prop: name_text(&first),
}),
other => ReturnExpr::PropOf(Box::new(other), name_text(&first)),
};
for name in names {
expr = ReturnExpr::PropOf(Box::new(expr), name_text(&name));
}
Ok(expr)
}
fn build_atom(&mut self, ctx: &AtomContext) -> Result<ReturnExpr, QueryError> {
if let Some(lit_ctx) = ctx.literal() {
return self.visit(&*lit_ctx).into_return_expr_lenient();
}
if let Some(param_ctx) = ctx.parameter() {
return self.build_parameter(¶m_ctx);
}
if let Some(paren_ctx) = ctx.parenthesizedExpression() {
return self.visit(&*paren_ctx).into_return_expr();
}
if let Some(func_ctx) = ctx.functionInvocation() {
return self.build_function_invocation(&func_ctx);
}
if let Some(count_ctx) = ctx.countAll() {
let _ = self.visit(&*count_ctx);
return Ok(ReturnExpr::CountStar);
}
if let Some(sym_ctx) = ctx.symbol() {
return Ok(ReturnExpr::Var(symbol_text(&sym_ctx)));
}
if let Some(filter_ctx) = ctx.filterWith() {
return self.build_filter_with(&filter_ctx);
}
if let Some(lc_ctx) = ctx.listComprehension() {
return self.build_list_comprehension(&lc_ctx);
}
if let Some(case_ctx) = ctx.caseExpression() {
return self.build_case_expression(&case_ctx);
}
if let Some(pc_ctx) = ctx.patternComprehension() {
return self.build_pattern_comprehension(&pc_ctx);
}
if let Some(rcp_ctx) = ctx.relationshipsChainPattern() {
return Ok(ReturnExpr::PatternPredicate(
self.build_relationships_chain_pattern(&rcp_ctx)?,
));
}
if let Some(se_ctx) = ctx.subqueryExist() {
return self.build_subquery_exist(&se_ctx);
}
Err(QueryError::Syntax(
"this expression form (path-as-expression) isn't supported by the ANTLR parser yet"
.into(),
))
}
fn build_pattern_comprehension(
&mut self,
ctx: &PatternComprehensionContext,
) -> Result<ReturnExpr, QueryError> {
let path_var = ctx
.lhs()
.and_then(|lhs| lhs.symbol())
.map(|s| symbol_text(&s));
let rcp_ctx = ctx
.relationshipsChainPattern()
.expect("patternComprehension always has a relationshipsChainPattern");
let pattern = self.build_relationships_chain_pattern(&rcp_ctx)?;
let where_clause = match ctx.where_() {
Some(where_ctx) => {
let expr_ctx = where_ctx
.expression()
.expect("where always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
Some(Box::new(return_expr_to_expr(expr)?))
}
None => None,
};
let proj_ctx = ctx
.expression()
.expect("patternComprehension always has a projection expression");
let projection = self.visit(&*proj_ctx).into_return_expr()?;
Ok(ReturnExpr::PatternComprehension {
path_var,
pattern: Box::new(pattern),
where_clause,
projection: Box::new(projection),
})
}
fn build_subquery_exist(
&mut self,
ctx: &SubqueryExistContext,
) -> Result<ReturnExpr, QueryError> {
if let Some(rq_ctx) = ctx.regularQuery() {
let stmt = self.build_regular_query(&rq_ctx)?;
return Ok(ReturnExpr::ExistsSubquery(Box::new(stmt)));
}
let pw_ctx = ctx
.patternWhere()
.expect("subqueryExist always has a regularQuery or patternWhere");
let pattern_ctx = pw_ctx.pattern().expect("patternWhere always has a pattern");
let mut parts = pattern_ctx.patternPart_all().into_iter();
let part = parts
.next()
.expect("pattern always has at least one patternPart");
if parts.next().is_some() {
return Err(QueryError::Syntax(
"exists {} with more than one comma-separated pattern isn't supported yet".into(),
));
}
if part.ASSIGN().is_some() || part.shortestPathWrapper().is_some() {
return Err(QueryError::Syntax(
"exists {} doesn't support a named path or shortestPath()".into(),
));
}
let elem_ctx = part
.patternElem()
.expect("a patternPart without ASSIGN/shortestPathWrapper always has a patternElem");
let pattern = self.visit(&*elem_ctx).into_pattern()?;
let where_clause = match pw_ctx.where_() {
Some(where_ctx) => {
let expr_ctx = where_ctx
.expression()
.expect("where always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
Some(Box::new(return_expr_to_expr(expr)?))
}
None => None,
};
Ok(ReturnExpr::ExistsPattern {
pattern: Box::new(pattern),
where_clause,
})
}
fn build_case_expression(
&mut self,
ctx: &CaseExpressionContext,
) -> Result<ReturnExpr, QueryError> {
#[derive(PartialEq)]
enum Pos {
BeforeFirstWhen,
AfterWhen,
AfterThen,
AfterElse,
}
let mut pos = Pos::BeforeFirstWhen;
let mut test = None;
let mut whens: Vec<(ReturnExpr, ReturnExpr)> = Vec::new();
let mut pending_when: Option<ReturnExpr> = None;
let mut else_ = None;
for child in ctx.get_children() {
match child.get_text().to_ascii_uppercase().as_str() {
"CASE" | "END" => continue,
"WHEN" => pos = Pos::AfterWhen,
"THEN" => pos = Pos::AfterThen,
"ELSE" => pos = Pos::AfterElse,
_ => {
let expr = self.visit(&*child).into_return_expr()?;
match pos {
Pos::BeforeFirstWhen => test = Some(Box::new(expr)),
Pos::AfterWhen => pending_when = Some(expr),
Pos::AfterThen => {
let w = pending_when
.take()
.expect("a THEN expression always follows a WHEN expression");
whens.push((w, expr));
}
Pos::AfterElse => else_ = Some(Box::new(expr)),
}
}
}
}
Ok(ReturnExpr::Case { test, whens, else_ })
}
fn build_filter_expression(
&mut self,
ctx: &FilterExpressionContext,
) -> Result<(String, ReturnExpr, Option<Box<ReturnExpr>>), QueryError> {
let var_ctx = ctx.symbol().expect("filterExpression always has a symbol");
let var = symbol_text(&var_ctx);
let source_ctx = ctx
.expression()
.expect("filterExpression always has an expression");
let source = self.visit(&*source_ctx).into_return_expr()?;
let where_clause = match ctx.where_() {
Some(where_ctx) => {
let expr_ctx = where_ctx
.expression()
.expect("where always has an expression");
Some(Box::new(self.visit(&*expr_ctx).into_return_expr()?))
}
None => None,
};
Ok((var, source, where_clause))
}
fn build_filter_with(&mut self, ctx: &FilterWithContext) -> Result<ReturnExpr, QueryError> {
let kind = if ctx.ALL().is_some() {
QuantifierKind::All
} else if ctx.ANY().is_some() {
QuantifierKind::Any
} else if ctx.NONE().is_some() {
QuantifierKind::None
} else {
ctx.SINGLE()
.expect("filterWith always has one of ALL/ANY/NONE/SINGLE");
QuantifierKind::Single
};
let fe_ctx = ctx
.filterExpression()
.expect("filterWith always has a filterExpression");
let (var, source, where_clause) = self.build_filter_expression(&fe_ctx)?;
Ok(ReturnExpr::Quantifier {
kind,
var,
source: Box::new(source),
where_clause,
})
}
fn build_list_comprehension(
&mut self,
ctx: &ListComprehensionContext,
) -> Result<ReturnExpr, QueryError> {
let fe_ctx = ctx
.filterExpression()
.expect("listComprehension always has a filterExpression");
let (var, source, where_clause) = self.build_filter_expression(&fe_ctx)?;
let project = match ctx.expression() {
Some(expr_ctx) => Some(Box::new(self.visit(&*expr_ctx).into_return_expr()?)),
None => None,
};
Ok(ReturnExpr::ListComp {
var,
source: Box::new(source),
where_clause,
project,
})
}
fn build_function_invocation(
&mut self,
ctx: &FunctionInvocationContext,
) -> Result<ReturnExpr, QueryError> {
let name_ctx = ctx
.invocationName()
.expect("functionInvocation always has an invocationName");
let name = invocation_name_text(&name_ctx);
let distinct = ctx.DISTINCT().is_some();
let mut args = Vec::new();
if let Some(chain_ctx) = ctx.expressionChain() {
for arg_ctx in chain_ctx.expression_all() {
args.push(self.visit(&*arg_ctx).into_return_expr()?);
}
}
if distinct && !is_aggregate_name(&name) {
return Err(QueryError::Syntax(format!(
"'{name}(DISTINCT ...)' isn't valid — DISTINCT is only meaningful inside an aggregate function"
)));
}
Ok(ReturnExpr::Call {
name,
args,
distinct,
})
}
fn build_standalone_call(
&mut self,
ctx: &StandaloneCallContext,
) -> Result<Statement, QueryError> {
let name_ctx = ctx
.invocationName()
.expect("standaloneCall always has an invocationName");
let name = invocation_name_text(&name_ctx);
let args = match ctx.parenExpressionChain() {
Some(paren_ctx) => Some(self.build_call_args(&paren_ctx)?),
None => None,
};
let yield_items = if ctx.MULT().is_some() {
Some(CallYield::Star)
} else if let Some(yi_ctx) = ctx.yieldItems() {
Some(self.build_yield_items(&yi_ctx)?)
} else {
None
};
Ok(Statement::StandaloneCall(Box::new(CallClause {
name,
args,
with: None,
yield_items,
})))
}
fn build_query_call_st(
&mut self,
ctx: &QueryCallStContextAll,
) -> Result<CallClause, QueryError> {
let name_ctx = ctx
.invocationName()
.expect("queryCallSt always has an invocationName");
let name = invocation_name_text(&name_ctx);
let paren_ctx = ctx
.parenExpressionChain()
.expect("queryCallSt always has a parenExpressionChain");
let args = Some(self.build_call_args(&paren_ctx)?);
let yield_items = match ctx.yieldItems() {
Some(yi_ctx) => Some(self.build_yield_items(&yi_ctx)?),
None => None,
};
Ok(CallClause {
name,
args,
with: None,
yield_items,
})
}
fn build_call_args(
&mut self,
ctx: &ParenExpressionChainContextAll,
) -> Result<Vec<ReturnExpr>, QueryError> {
let mut args = Vec::new();
if let Some(chain_ctx) = ctx.expressionChain() {
for arg_ctx in chain_ctx.expression_all() {
args.push(self.visit(&*arg_ctx).into_return_expr()?);
}
}
Ok(args)
}
fn build_yield_items(&mut self, ctx: &YieldItemsContextAll) -> Result<CallYield, QueryError> {
let mut items = Vec::new();
for item_ctx in ctx.yieldItem_all() {
let symbols = item_ctx.symbol_all();
let (name, alias) = match symbols.len() {
1 => (symbol_text(&symbols[0]), None),
2 => (symbol_text(&symbols[0]), Some(symbol_text(&symbols[1]))),
other => unreachable!("yieldItem always has 1 or 2 symbols, got {other}"),
};
items.push((name, alias));
}
let where_clause = match ctx.where_() {
Some(where_ctx) => {
let expr_ctx = where_ctx
.expression()
.expect("where always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
Some(Box::new(return_expr_to_expr(expr)?))
}
None => None,
};
Ok(CallYield::Items(items, where_clause))
}
fn build_parameter(&mut self, ctx: &ParameterContext) -> Result<ReturnExpr, QueryError> {
let name = if let Some(sym_ctx) = ctx.symbol() {
symbol_text(&sym_ctx)
} else if let Some(num_ctx) = ctx.numLit() {
num_ctx
.DIGIT()
.expect("numLit context always has a DIGIT token")
.get_text()
} else {
unreachable!("parameter always has a symbol or numLit")
};
Ok(ReturnExpr::Lit(Literal::Param(name)))
}
fn build_projection_body(
&mut self,
ctx: &ProjectionBodyContext,
) -> Result<ParsedReturnClause, QueryError> {
let distinct = ctx.DISTINCT().is_some();
let items_ctx = ctx
.projectionItems()
.expect("projectionBody always has projectionItems");
let tail = if items_ctx.MULT().is_some() {
if !items_ctx.projectionItem_all().is_empty() {
return Err(QueryError::Syntax(
"RETURN * can't be combined with additional items".into(),
));
}
Tail::ReturnStar(distinct)
} else {
let mut items = Vec::new();
for item_ctx in items_ctx.projectionItem_all() {
let expr_ctx = item_ctx
.expression()
.expect("projectionItem always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
let alias = item_ctx.symbol().map(|s| symbol_text(&s));
items.push(ReturnItem { expr, alias });
}
Tail::Return(items, distinct)
};
let (order_by, skip, limit) = self.build_order_skip_limit(ctx)?;
Ok(ParsedReturnClause {
tail,
order_by,
skip,
limit,
})
}
#[allow(clippy::type_complexity)]
fn build_order_skip_limit(
&mut self,
ctx: &ProjectionBodyContext,
) -> Result<
(
Option<Vec<(ReturnExpr, SortDir)>>,
Option<ReturnExpr>,
Option<ReturnExpr>,
),
QueryError,
> {
let order_by = match ctx.orderSt() {
Some(order_ctx) => Some(self.build_order_by(&order_ctx)?),
None => None,
};
let skip = match ctx.skipSt() {
Some(skip_ctx) => {
let expr_ctx = skip_ctx
.expression()
.expect("skipSt always has an expression");
Some(self.visit(&*expr_ctx).into_return_expr()?)
}
None => None,
};
let limit = match ctx.limitSt() {
Some(limit_ctx) => {
let expr_ctx = limit_ctx
.expression()
.expect("limitSt always has an expression");
Some(self.visit(&*expr_ctx).into_return_expr()?)
}
None => None,
};
Ok((order_by, skip, limit))
}
fn build_order_by(
&mut self,
ctx: &OrderStContext,
) -> Result<Vec<(ReturnExpr, SortDir)>, QueryError> {
let mut items = Vec::new();
for item_ctx in ctx.orderItem_all() {
let expr_ctx = item_ctx
.expression()
.expect("orderItem always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
let dir = if item_ctx.DESC().is_some() || item_ctx.DESCENDING().is_some() {
SortDir::Desc
} else {
SortDir::Asc
};
items.push((expr, dir));
}
Ok(items)
}
fn build_with_clause(&mut self, ctx: &WithStContext) -> Result<WithClause, QueryError> {
let body_ctx = ctx
.projectionBody()
.expect("withSt always has a projectionBody");
let distinct = body_ctx.DISTINCT().is_some();
let items_ctx = body_ctx
.projectionItems()
.expect("projectionBody always has projectionItems");
let star = items_ctx.MULT().is_some();
let mut items = Vec::new();
for item_ctx in items_ctx.projectionItem_all() {
let expr_ctx = item_ctx
.expression()
.expect("projectionItem always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
let alias = item_ctx.symbol().map(|s| symbol_text(&s));
items.push(ReturnItem { expr, alias });
}
let (order_by, skip, limit) = self.build_order_skip_limit(&body_ctx)?;
let where_clause = match ctx.where_() {
Some(where_ctx) => {
let expr_ctx = where_ctx
.expression()
.expect("where always has an expression");
let expr = self.visit(&*expr_ctx).into_return_expr()?;
Some(return_expr_to_with_expr(expr))
}
None => None,
};
Ok(WithClause {
items,
star,
distinct,
where_clause,
order_by,
skip,
limit,
})
}
fn build_unwind_st(&mut self, ctx: &UnwindStContext) -> Result<UnwindClause, QueryError> {
let expr_ctx = ctx.expression().expect("unwindSt always has an expression");
let source = UnwindSource(self.visit(&*expr_ctx).into_return_expr()?);
let var_ctx = ctx.symbol().expect("unwindSt always has a symbol");
Ok(UnwindClause {
source,
var: symbol_text(&var_ctx),
where_clause: None,
with: None,
})
}
fn build_set_st(&mut self, ctx: &SetStContext) -> Result<Vec<SetItem>, QueryError> {
ctx.setItem_all()
.into_iter()
.map(|item_ctx| self.build_set_item(&item_ctx))
.collect()
}
fn build_set_item(&mut self, ctx: &SetItemContextAll) -> Result<SetItem, QueryError> {
if let Some(prop_ctx) = ctx.propertyExpression() {
let expr_ctx = ctx
.expression()
.expect("setItem's propertyExpression form always has an expression");
return match self.build_property_expression(&prop_ctx)? {
ReturnExpr::Prop(prop) => {
let value = self.visit(&*expr_ctx).into_return_expr()?;
Ok(SetItem::Prop(prop, value))
}
ReturnExpr::Var(var) => {
let value = self.visit(&*expr_ctx).into_return_expr()?;
Ok(SetItem::MapAssign {
var,
value,
merge: false,
})
}
_ => Err(QueryError::Syntax(
"expected a property access (x.prop) or variable on the left of SET's `=`"
.into(),
)),
};
}
let sym_ctx = ctx
.symbol()
.expect("setItem always has a propertyExpression or symbol");
let var = symbol_text(&sym_ctx);
if let Some(labels_ctx) = ctx.nodeLabels() {
let labels = labels_ctx.name_all().iter().map(|n| name_text(n)).collect();
return Ok(SetItem::Labels(var, labels));
}
let expr_ctx = ctx
.expression()
.expect("setItem's symbol-assign form always has an expression");
let value = self.visit(&*expr_ctx).into_return_expr()?;
Ok(SetItem::MapAssign {
var,
value,
merge: ctx.ADD_ASSIGN().is_some(),
})
}
fn build_delete_st(&mut self, ctx: &DeleteStContext) -> Result<ParsedDelete, QueryError> {
let chain_ctx = ctx
.expressionChain()
.expect("deleteSt always has an expressionChain");
let mut items = Vec::new();
for expr_ctx in chain_ctx.expression_all() {
items.push(self.visit(&*expr_ctx).into_return_expr()?);
}
Ok(ParsedDelete {
items,
detach: ctx.DETACH().is_some(),
})
}
fn build_remove_st(&mut self, ctx: &RemoveStContext) -> Result<Vec<RemoveItem>, QueryError> {
ctx.removeItem_all()
.into_iter()
.map(|item_ctx| self.build_remove_item(&item_ctx))
.collect()
}
fn build_remove_item(&mut self, ctx: &RemoveItemContextAll) -> Result<RemoveItem, QueryError> {
if let Some(prop_ctx) = ctx.propertyExpression() {
return Ok(RemoveItem::Prop(self.build_prop_access(&prop_ctx)?));
}
let sym_ctx = ctx
.symbol()
.expect("removeItem always has a symbol+nodeLabels or a propertyExpression");
let labels_ctx = ctx
.nodeLabels()
.expect("removeItem's symbol form always has nodeLabels");
let labels = labels_ctx.name_all().iter().map(|n| name_text(n)).collect();
Ok(RemoveItem::Labels(symbol_text(&sym_ctx), labels))
}
fn build_prop_access(
&mut self,
ctx: &PropertyExpressionContext,
) -> Result<PropAccess, QueryError> {
match self.build_property_expression(ctx)? {
ReturnExpr::Prop(p) => Ok(p),
_ => Err(QueryError::Syntax(
"expected a property access (x.prop)".into(),
)),
}
}
fn build_create_st(&mut self, ctx: &CreateStContext) -> Result<Vec<Pattern>, QueryError> {
let pattern_ctx = ctx.pattern().expect("createSt always has a pattern");
pattern_ctx
.patternPart_all()
.into_iter()
.map(|part_ctx| {
if part_ctx.ASSIGN().is_some() {
return Err(QueryError::Syntax(
"named-path capture (`p = ...`) isn't supported on CREATE".into(),
));
}
if part_ctx.shortestPathWrapper().is_some() {
return Err(QueryError::Syntax(
"shortestPath() isn't valid in CREATE".into(),
));
}
let elem_ctx = part_ctx.patternElem().expect(
"patternPart always has a patternElem when shortestPathWrapper is absent",
);
self.visit(&*elem_ctx).into_pattern()
})
.collect()
}
fn build_merge_st(&mut self, ctx: &MergeStContext) -> Result<MergeClause, QueryError> {
let part_ctx = ctx.patternPart().expect("mergeSt always has a patternPart");
let path_var = if part_ctx.ASSIGN().is_some() {
let symbol_ctx = part_ctx
.symbol()
.expect("patternPart with ASSIGN always has a symbol");
Some(symbol_text(&symbol_ctx))
} else {
None
};
if part_ctx.shortestPathWrapper().is_some() {
return Err(QueryError::Syntax(
"shortestPath() isn't valid in MERGE".into(),
));
}
let elem_ctx = part_ctx
.patternElem()
.expect("patternPart always has a patternElem when shortestPathWrapper is absent");
let pattern = self.visit(&*elem_ctx).into_pattern()?;
if pattern.hops.len() > 1 {
return Err(QueryError::Syntax(
"MERGE with more than one relationship hop isn't supported yet — split it into a MATCH \
for the already-known part and a MERGE for one new hop"
.into(),
));
}
let mut on_create = Vec::new();
let mut on_match = Vec::new();
for action_ctx in ctx.mergeAction_all() {
let set_items = self.build_merge_action(&action_ctx)?;
if action_ctx.MATCH().is_some() {
if !on_match.is_empty() {
return Err(QueryError::Syntax(
"MERGE can have at most one ON MATCH SET clause".into(),
));
}
on_match = set_items;
} else {
if !on_create.is_empty() {
return Err(QueryError::Syntax(
"MERGE can have at most one ON CREATE SET clause".into(),
));
}
on_create = set_items;
}
}
Ok(MergeClause {
pattern,
path_var,
on_create,
on_match,
with: None,
})
}
fn build_merge_action(
&mut self,
ctx: &MergeActionContextAll,
) -> Result<Vec<SetItem>, QueryError> {
let set_ctx = ctx.setSt().expect("mergeAction always has a setSt");
self.build_set_st(&set_ctx)
}
fn append_reading_statement(
&mut self,
ctx: &ReadingStatementContextAll,
clauses: &mut Vec<QueryClause>,
) -> Result<(), QueryError> {
if let Some(match_ctx) = ctx.matchSt() {
let parts = self.visit(&*match_ctx).into_query_parts()?;
clauses.extend(parts.into_iter().map(QueryClause::Match));
return Ok(());
}
if let Some(unwind_ctx) = ctx.unwindSt() {
let clause = self.visit(&*unwind_ctx).into_unwind_clause()?;
clauses.push(QueryClause::Unwind(clause));
return Ok(());
}
let call_ctx = ctx
.queryCallSt()
.expect("readingStatement is matchSt | unwindSt | queryCallSt");
let call = self.build_query_call_st(&call_ctx)?;
clauses.push(QueryClause::Call(call));
Ok(())
}
fn build_updating_statement_as_clause(
&mut self,
ctx: &UpdatingStatementContextAll,
) -> Result<QueryClause, QueryError> {
if let Some(create_ctx) = ctx.createSt() {
return Ok(QueryClause::Create(
self.visit(&*create_ctx).into_create_patterns()?,
));
}
if let Some(merge_ctx) = ctx.mergeSt() {
return Ok(QueryClause::Merge(
self.visit(&*merge_ctx).into_merge_clause()?,
));
}
if let Some(delete_ctx) = ctx.deleteSt() {
let d = self.visit(&*delete_ctx).into_delete_items()?;
return Ok(QueryClause::Delete {
items: d.items,
detach: d.detach,
});
}
if let Some(set_ctx) = ctx.setSt() {
return Ok(QueryClause::Set(self.visit(&*set_ctx).into_set_items()?));
}
let remove_ctx = ctx
.removeSt()
.expect("updatingStatement always has one of its 5 alternatives");
Ok(QueryClause::Remove(
self.visit(&*remove_ctx).into_remove_items()?,
))
}
#[allow(clippy::type_complexity)]
fn build_mutating_tail(
&mut self,
ctx: &UpdatingStatementContextAll,
return_ctx: Option<&ReturnStContext>,
) -> Result<
(
Tail,
Option<Vec<(ReturnExpr, SortDir)>>,
Option<ReturnExpr>,
Option<ReturnExpr>,
),
QueryError,
> {
let mut order_by = None;
let mut skip = None;
let mut limit = None;
let ret = match return_ctx {
Some(return_ctx) => {
let c = self.visit(return_ctx).into_return_clause()?;
order_by = c.order_by;
skip = c.skip;
limit = c.limit;
let Tail::Return(items, distinct) = c.tail else {
return Err(QueryError::Syntax(
"RETURN * isn't supported as a mutating clause's own trailing RETURN"
.into(),
));
};
Some(ReturnTail { items, distinct })
}
None => None,
};
let tail = if let Some(create_ctx) = ctx.createSt() {
Tail::Create(self.visit(&*create_ctx).into_create_patterns()?, ret)
} else if let Some(delete_ctx) = ctx.deleteSt() {
let d = self.visit(&*delete_ctx).into_delete_items()?;
if d.detach {
Tail::DetachDelete(d.items, ret)
} else {
Tail::Delete(d.items, ret)
}
} else if let Some(set_ctx) = ctx.setSt() {
Tail::Set(self.visit(&*set_ctx).into_set_items()?, ret)
} else {
let remove_ctx = ctx
.removeSt()
.expect("build_mutating_tail's caller already excluded mergeSt");
Tail::Remove(self.visit(&*remove_ctx).into_remove_items()?, ret)
};
Ok((tail, order_by, skip, limit))
}
fn build_single_part_q(&mut self, ctx: &SinglePartQContext) -> Result<Statement, QueryError> {
let mut clauses = Vec::new();
for rs_ctx in ctx.readingStatement_all() {
self.append_reading_statement(&rs_ctx, &mut clauses)?;
}
let updating = ctx.updatingStatement_all();
let return_ctx = ctx.returnSt();
if clauses.is_empty() && return_ctx.is_none() && updating.len() == 1 {
if let Some(create_ctx) = updating[0].createSt() {
let patterns = self.visit(&*create_ctx).into_create_patterns()?;
return Ok(Statement::Create(patterns));
}
}
let mut tail = None;
let mut order_by = None;
let mut skip = None;
let mut limit = None;
let mut consumed_return = false;
if let Some((last, earlier)) = updating.split_last() {
for us_ctx in earlier {
clauses.push(self.build_updating_statement_as_clause(us_ctx)?);
}
if last.mergeSt().is_some() {
clauses.push(self.build_updating_statement_as_clause(last)?);
} else {
let (t, ob, sk, lim) = self.build_mutating_tail(last, return_ctx.as_deref())?;
tail = Some(t);
order_by = ob;
skip = sk;
limit = lim;
consumed_return = return_ctx.is_some();
}
}
if !consumed_return {
if let Some(return_ctx) = return_ctx {
let c = self.visit(&*return_ctx).into_return_clause()?;
tail = Some(c.tail);
order_by = c.order_by;
skip = c.skip;
limit = c.limit;
}
}
if tail.is_none() && !clauses.iter().any(|c| matches!(c, QueryClause::Merge(_))) {
return Err(QueryError::Syntax(
"a query needs a RETURN/DELETE/SET tail, unless it has a MERGE clause with nothing after it".into(),
));
}
Ok(Statement::Match {
clauses,
tail,
order_by,
skip: skip.map(Box::new),
limit: limit.map(Box::new),
})
}
fn build_multi_part_q(&mut self, ctx: &MultiPartQContext) -> Result<Statement, QueryError> {
enum Item<'i> {
Reading(Rc<ReadingStatementContextAll<'i>>),
Updating(Rc<UpdatingStatementContextAll<'i>>),
With(Rc<WithStContext<'i>>),
}
let mut items: Vec<(isize, Item)> = Vec::new();
for rs in ctx.readingStatement_all() {
let idx = rs.start().get_token_index();
items.push((idx, Item::Reading(rs)));
}
for us in ctx.updatingStatement_all() {
let idx = us.start().get_token_index();
items.push((idx, Item::Updating(us)));
}
for w in ctx.withSt_all() {
let idx = w.start().get_token_index();
items.push((idx, Item::With(w)));
}
items.sort_by_key(|(idx, _)| *idx);
let mut clauses: Vec<QueryClause> = Vec::new();
let mut attach_target: Option<usize> = None;
for (_, item) in items {
match item {
Item::Reading(rs) => {
self.append_reading_statement(&rs, &mut clauses)?;
attach_target = Some(clauses.len() - 1);
}
Item::Updating(us) => {
let clause = self.build_updating_statement_as_clause(&us)?;
let can_attach = matches!(clause, QueryClause::Merge(_));
clauses.push(clause);
attach_target = can_attach.then_some(clauses.len() - 1);
}
Item::With(w) => {
let with = self.visit(&*w).into_with_clause()?;
match attach_target.take() {
Some(i) => match &mut clauses[i] {
QueryClause::Match(part) => part.with = Some(with),
QueryClause::Unwind(u) => u.with = Some(with),
QueryClause::Merge(m) => m.with = Some(with),
QueryClause::Call(call) => call.with = Some(with),
_ => unreachable!(
"attach_target is only ever set right after pushing a Match/Unwind/Merge/Call clause"
),
},
None => clauses.push(QueryClause::With(with)),
}
}
}
}
let sp_ctx = ctx
.singlePartQ()
.expect("multiPartQ always ends in a singlePartQ");
let (tail_clauses, tail, order_by, skip, limit) =
match self.build_single_part_q(&sp_ctx)? {
Statement::Match {
clauses,
tail,
order_by,
skip,
limit,
} => (clauses, tail, order_by, skip, limit),
Statement::Create(patterns) => {
(Vec::new(), Some(Tail::Create(patterns, None)), None, None, None)
}
other => unreachable!(
"build_single_part_q only ever returns Statement::Match or Statement::Create, got {other:?}"
),
};
clauses.extend(tail_clauses);
Ok(Statement::Match {
clauses,
tail,
order_by,
skip,
limit,
})
}
fn build_explain_st(&mut self, ctx: &ExplainStContext) -> Result<Statement, QueryError> {
let inner = match ctx.createIndexSt() {
Some(ci_ctx) => self.build_create_index_st(&ci_ctx)?,
None => {
let rq_ctx = ctx
.regularQuery()
.expect("explainSt always has a createIndexSt or regularQuery");
self.visit(&*rq_ctx).into_statement()?
}
};
Ok(Statement::Explain(Box::new(inner)))
}
fn build_create_index_st(
&mut self,
ctx: &CreateIndexStContext,
) -> Result<Statement, QueryError> {
let names = ctx.name_all();
let label = name_text(
names
.first()
.expect("createIndexSt always has a label name"),
);
let prop = name_text(
names
.get(1)
.expect("createIndexSt always has a property name"),
);
Ok(Statement::CreateIndex {
label,
prop,
unique: ctx.UNIQUE().is_some(),
})
}
fn build_regular_query(&mut self, ctx: &RegularQueryContext) -> Result<Statement, QueryError> {
let sq_ctx = ctx
.singleQuery()
.expect("regularQuery always has a singleQuery");
let first = self.visit(&*sq_ctx).into_statement()?;
let unions = ctx.unionSt_all();
if unions.is_empty() {
return Ok(first);
}
let mut parts = vec![first];
let mut all: Option<bool> = None;
for u_ctx in unions {
let this_all = u_ctx.ALL().is_some();
match all {
None => all = Some(this_all),
Some(prev) if prev != this_all => {
return Err(QueryError::Syntax(
"can't mix UNION and UNION ALL in the same statement".into(),
));
}
Some(_) => {}
}
let part_sq = u_ctx
.singleQuery()
.expect("unionSt always has a singleQuery");
parts.push(self.visit(&*part_sq).into_statement()?);
}
Ok(Statement::Union {
parts,
all: all.unwrap_or(false),
})
}
}
pub fn parse_antlr(input: &str) -> Result<Statement, QueryError> {
use crate::generated::cypherlexer::CypherLexer;
use crate::generated::cypherparser::{CypherParser, ScriptContextAttrs};
use antlr4rust::common_token_stream::CommonTokenStream;
use antlr4rust::error_listener::ErrorListener;
use antlr4rust::recognizer::Recognizer;
use antlr4rust::token_factory::TokenFactory;
use antlr4rust::InputStream;
use antlr4rust::Parser as _;
use std::cell::RefCell;
struct CollectErrors(Rc<RefCell<Vec<String>>>);
impl<'a, T: Recognizer<'a>> ErrorListener<'a, T> for CollectErrors {
fn syntax_error(
&self,
_recognizer: &T,
_offending_symbol: Option<&<T::TF as TokenFactory<'a>>::Inner>,
line: isize,
column: isize,
msg: &str,
_e: Option<&antlr4rust::errors::ANTLRError>,
) {
self.0
.borrow_mut()
.push(format!("line {line}:{column} {msg}"));
}
}
let errors = Rc::new(RefCell::new(Vec::new()));
let stream = InputStream::new(input);
let mut lexer = CypherLexer::new(stream);
lexer.remove_error_listeners();
lexer.add_error_listener(Box::new(CollectErrors(errors.clone())));
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
parser.remove_error_listeners();
parser.add_error_listener(Box::new(CollectErrors(errors.clone())));
let ctx = parser
.script()
.map_err(|e| QueryError::Syntax(e.to_string()))?;
if let Some(msg) = errors.borrow().first() {
return Err(QueryError::Syntax(format!("syntax error: {msg}")));
}
let query_ctx = ctx.query().expect("script always has a query");
AstBuilder::new().visit(&*query_ctx).into_statement()
}
pub fn parse_antlr_many(input: &str) -> Result<Vec<Statement>, QueryError> {
let trimmed = input.trim_end();
let trimmed = trimmed.strip_suffix(';').unwrap_or(trimmed);
split_statements(trimmed)
.into_iter()
.map(parse_antlr)
.collect()
}
fn split_statements(input: &str) -> Vec<&str> {
let bytes = input.as_bytes();
let mut starts = vec![0usize];
let mut semicolons = Vec::new();
let mut quote: Option<u8> = None;
let mut i = 0;
while i < bytes.len() {
let b = bytes[i];
match quote {
Some(q) => {
if b == b'\\' && q != b'`' {
i += 1; } else if b == q {
quote = None;
}
}
None => match b {
b'\'' | b'"' | b'`' => quote = Some(b),
b';' => {
semicolons.push(i);
starts.push(i + 1);
}
_ => {}
},
}
i += 1;
}
starts
.iter()
.enumerate()
.map(|(idx, &start)| {
let end = semicolons.get(idx).copied().unwrap_or(bytes.len());
&input[start..end]
})
.collect()
}
fn return_expr_to_with_expr(expr: ReturnExpr) -> WithExpr {
match expr {
ReturnExpr::And(l, r) => WithExpr::And(
Box::new(return_expr_to_with_expr(*l)),
Box::new(return_expr_to_with_expr(*r)),
),
ReturnExpr::Or(l, r) => WithExpr::Or(
Box::new(return_expr_to_with_expr(*l)),
Box::new(return_expr_to_with_expr(*r)),
),
ReturnExpr::Not(inner) => WithExpr::Not(Box::new(return_expr_to_with_expr(*inner))),
ReturnExpr::Compare(l, op, r) => WithExpr::Compare(*l, op, *r),
ReturnExpr::IsNull(inner) => WithExpr::IsNull(*inner),
other => WithExpr::Bare(other),
}
}
fn return_expr_to_expr(expr: ReturnExpr) -> Result<Expr, QueryError> {
Ok(match expr {
ReturnExpr::And(l, r) => Expr::And(
Box::new(return_expr_to_expr(*l)?),
Box::new(return_expr_to_expr(*r)?),
),
ReturnExpr::Or(l, r) => Expr::Or(
Box::new(return_expr_to_expr(*l)?),
Box::new(return_expr_to_expr(*r)?),
),
ReturnExpr::Not(inner) => Expr::Not(Box::new(return_expr_to_expr(*inner)?)),
ReturnExpr::Compare(l, op, r) => match (*l, *r) {
(ReturnExpr::Prop(pa), ReturnExpr::Lit(lit)) => Expr::Compare(pa, op, lit),
(ReturnExpr::Prop(pa1), ReturnExpr::Prop(pa2)) => Expr::PropCompare(pa1, op, pa2),
(ReturnExpr::Var(a), ReturnExpr::Var(b)) => match op {
CompareOp::Eq => Expr::VarEq(a, b),
CompareOp::Ne => Expr::Not(Box::new(Expr::VarEq(a, b))),
_ => {
return Err(QueryError::Syntax(format!(
"{a} {op:?} {b}: only = and <> are meaningful for comparing two \
nodes/relationships by identity (no ordering exists between them)"
)))
}
},
(l, r) => Expr::GeneralCompare(l, op, r),
},
ReturnExpr::IsNull(inner) => match *inner {
ReturnExpr::Prop(pa) => Expr::IsNull(pa),
other => Expr::GeneralIsNull(other),
},
ReturnExpr::HasLabel(var, labels) => {
let mut labels = labels.into_iter();
let first = labels
.next()
.expect("HasLabel always carries at least one label");
labels.fold(Expr::HasLabel(var.clone(), first), |acc, label| {
Expr::And(Box::new(acc), Box::new(Expr::HasLabel(var.clone(), label)))
})
}
ReturnExpr::PatternPredicate(pattern) => Expr::Pattern(pattern),
ReturnExpr::ExistsPattern {
pattern,
where_clause,
} => Expr::Exists {
pattern,
where_clause,
},
ReturnExpr::ExistsSubquery(stmt) => Expr::ExistsSubquery(stmt),
other => Expr::GeneralBare(other),
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::generated::cypherlexer::CypherLexer;
use crate::generated::cypherparser::CypherParser;
use antlr4rust::common_token_stream::CommonTokenStream;
use antlr4rust::InputStream;
fn parse_literal_expr(input: &str) -> Result<Literal, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.literal()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `literal`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_literal()
}
fn parse_pattern(input: &str) -> Result<Pattern, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.patternElem()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `patternElem`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_pattern()
}
fn parse_match(input: &str) -> Result<Vec<QueryPart>, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.matchSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `matchSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_query_parts()
}
fn parse_expr(input: &str) -> Result<ReturnExpr, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.expression()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `expression`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_return_expr()
}
fn parse_return(input: &str) -> Result<ParsedReturnClause, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.returnSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `returnSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_return_clause()
}
fn parse_with(input: &str) -> Result<WithClause, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.withSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `withSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_with_clause()
}
fn parse_unwind(input: &str) -> Result<UnwindClause, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.unwindSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `unwindSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_unwind_clause()
}
fn parse_set(input: &str) -> Result<Vec<SetItem>, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.setSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `setSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_set_items()
}
fn parse_delete(input: &str) -> Result<ParsedDelete, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.deleteSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `deleteSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_delete_items()
}
fn parse_remove(input: &str) -> Result<Vec<RemoveItem>, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.removeSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `removeSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_remove_items()
}
fn parse_create(input: &str) -> Result<Vec<Pattern>, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.createSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `createSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_create_patterns()
}
fn parse_merge(input: &str) -> Result<MergeClause, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.mergeSt()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `mergeSt`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_merge_clause()
}
fn parse_statement(input: &str) -> Result<Statement, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.singlePartQ()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `singlePartQ`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_statement()
}
fn parse_multi_part_statement(input: &str) -> Result<Statement, QueryError> {
let stream = InputStream::new(input);
let lexer = CypherLexer::new(stream);
let tokens = CommonTokenStream::new(lexer);
let mut parser = CypherParser::new(tokens);
let ctx = parser
.multiPartQ()
.unwrap_or_else(|e| panic!("failed to parse {input:?} as `multiPartQ`: {e:?}"));
AstBuilder::new().visit(&*ctx).into_statement()
}
#[test]
fn bool_literals() {
assert_eq!(parse_literal_expr("true").unwrap(), Literal::Bool(true));
assert_eq!(parse_literal_expr("FALSE").unwrap(), Literal::Bool(false));
}
#[test]
fn null_literal() {
assert_eq!(parse_literal_expr("null").unwrap(), Literal::Null);
}
#[test]
fn decimal_int() {
assert_eq!(parse_literal_expr("42").unwrap(), Literal::Int(42));
assert_eq!(parse_literal_expr("007").unwrap(), Literal::Int(7));
}
#[test]
fn hex_and_octal_int() {
assert_eq!(parse_literal_expr("0x1A").unwrap(), Literal::Int(26));
assert_eq!(parse_literal_expr("0o17").unwrap(), Literal::Int(15));
}
#[test]
fn float_literals() {
assert_eq!(parse_literal_expr("2.5").unwrap(), Literal::Float(2.5));
assert_eq!(parse_literal_expr("1e10").unwrap(), Literal::Float(1e10));
assert_eq!(parse_literal_expr(".5").unwrap(), Literal::Float(0.5));
}
#[test]
fn float_overflow_errors() {
assert!(parse_literal_expr("1e999").is_err());
}
#[test]
fn string_and_char_literals() {
assert_eq!(
parse_literal_expr("\"hello\"").unwrap(),
Literal::String("hello".to_string())
);
assert_eq!(
parse_literal_expr("'a string with spaces and a hyphen-in-it'").unwrap(),
Literal::String("a string with spaces and a hyphen-in-it".to_string())
);
}
#[test]
fn string_escapes() {
assert_eq!(
parse_literal_expr(r#"'line1\nline2'"#).unwrap(),
Literal::String("line1\nline2".to_string())
);
assert_eq!(
parse_literal_expr(r#"'é'"#).unwrap(),
Literal::String("é".to_string())
);
}
#[test]
fn single_node() {
let p = parse_pattern("(a:Person)").unwrap();
assert_eq!(p.start.var.as_deref(), Some("a"));
assert_eq!(p.start.labels, vec!["Person".to_string()]);
assert!(p.hops.is_empty());
}
#[test]
fn anonymous_node() {
let p = parse_pattern("()").unwrap();
assert_eq!(p.start.var, None);
assert!(p.start.labels.is_empty());
}
#[test]
fn multiple_labels() {
let p = parse_pattern("(a:Person:Employee)").unwrap();
assert_eq!(
p.start.labels,
vec!["Person".to_string(), "Employee".to_string()]
);
}
#[test]
fn escaped_identifier() {
let p = parse_pattern("(`weird name`)").unwrap();
assert_eq!(p.start.var.as_deref(), Some("weird name"));
}
#[test]
fn directions() {
assert_eq!(
parse_pattern("(a)-->(b)").unwrap().hops[0].0.direction,
RelDirection::Right
);
assert_eq!(
parse_pattern("(a)<--(b)").unwrap().hops[0].0.direction,
RelDirection::Left
);
assert_eq!(
parse_pattern("(a)--(b)").unwrap().hops[0].0.direction,
RelDirection::Either
);
assert_eq!(
parse_pattern("(a)<-->(b)").unwrap().hops[0].0.direction,
RelDirection::Either
);
}
#[test]
fn rel_type_and_var() {
let p = parse_pattern("(a)-[r:KNOWS]->(b)").unwrap();
let (rel, node) = &p.hops[0];
assert_eq!(rel.var.as_deref(), Some("r"));
assert_eq!(rel.rel_types, vec!["KNOWS".to_string()]);
assert_eq!(node.var.as_deref(), Some("b"));
assert_eq!(rel.hop_range, None);
}
#[test]
fn multiple_rel_types() {
let p = parse_pattern("(a)-[:KNOWS|LIKES]->(b)").unwrap();
assert_eq!(
p.hops[0].0.rel_types,
vec!["KNOWS".to_string(), "LIKES".to_string()]
);
}
#[test]
fn var_length_bounds() {
assert_eq!(
parse_pattern("(a)-[*0]->(b)").unwrap().hops[0].0.hop_range,
Some((0, Some(0)))
);
assert_eq!(
parse_pattern("(a)-[*2]->(b)").unwrap().hops[0].0.hop_range,
Some((2, Some(2)))
);
assert_eq!(
parse_pattern("(a)-[*1..3]->(b)").unwrap().hops[0]
.0
.hop_range,
Some((1, Some(3)))
);
assert_eq!(
parse_pattern("(a)-[*]->(b)").unwrap().hops[0].0.hop_range,
Some((1, None))
);
}
#[test]
fn multi_hop_chain() {
let p = parse_pattern("(a)-[:KNOWS]->(b)<-[:LIKES]-(c)").unwrap();
assert_eq!(p.hops.len(), 2);
assert_eq!(p.hops[0].0.direction, RelDirection::Right);
assert_eq!(p.hops[1].0.direction, RelDirection::Left);
}
#[test]
fn node_pattern_properties() {
let pattern = parse_pattern("(a {name: 'x', age: 1 + 1})").unwrap();
assert_eq!(
pattern.start.props,
vec![
(
"name".to_string(),
ReturnExpr::Lit(Literal::String("x".to_string()))
),
(
"age".to_string(),
ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(1))),
ArithOp::Add,
Box::new(ReturnExpr::Lit(Literal::Int(1))),
)
),
]
);
}
#[test]
fn rel_pattern_properties() {
let pattern = parse_pattern("(a)-[:T {weight: 5}]->(b)").unwrap();
assert_eq!(
pattern.hops[0].0.props,
vec![("weight".to_string(), ReturnExpr::Lit(Literal::Int(5)))]
);
}
#[test]
fn pattern_properties_parameter_not_supported() {
assert!(parse_pattern("(a $props)").is_err());
}
#[test]
fn simple_match() {
let parts = parse_match("MATCH (a:Person)-[:KNOWS]->(b)").unwrap();
assert_eq!(parts.len(), 1);
assert!(!parts[0].optional);
assert_eq!(parts[0].path_var, None);
assert_eq!(parts[0].pattern.start.var.as_deref(), Some("a"));
assert_eq!(parts[0].pattern.hops.len(), 1);
}
#[test]
fn optional_match() {
let parts = parse_match("OPTIONAL MATCH (a)").unwrap();
assert!(parts[0].optional);
}
#[test]
fn named_path() {
let parts = parse_match("MATCH p = (a)-->(b)").unwrap();
assert_eq!(parts.len(), 1);
assert_eq!(parts[0].path_var.as_deref(), Some("p"));
}
#[test]
fn comma_pattern_shared_node_merges_into_one_linear_chain() {
let parts = parse_match("MATCH (a)-->(b), (b)-->(c)").unwrap();
assert_eq!(parts.len(), 1);
assert_eq!(parts[0].pattern.hops.len(), 2);
}
#[test]
fn comma_pattern_disjoint_becomes_multiple_query_parts() {
let parts = parse_match("MATCH (a), (b)").unwrap();
assert_eq!(parts.len(), 2);
}
#[test]
fn named_path_over_disjoint_cross_join_errors() {
assert!(parse_match("MATCH p = (a), (b)").is_err());
}
#[test]
fn shortest_path() {
let parts = parse_match("MATCH shortestPath((a)-[*1..3]->(b))").unwrap();
assert_eq!(parts.len(), 1);
assert!(parts[0].shortest_path);
assert_eq!(parts[0].pattern.hops.len(), 1);
}
#[test]
fn shortest_path_with_named_path_capture() {
let parts = parse_match("MATCH p = shortestPath((a)-[*1..3]->(b))").unwrap();
assert_eq!(parts[0].path_var.as_deref(), Some("p"));
assert!(parts[0].shortest_path);
}
#[test]
fn shortest_path_requires_variable_length_hop() {
assert!(parse_match("MATCH shortestPath((a)-->(b))").is_err());
}
#[test]
fn shortest_path_not_first_in_cross_join_errors() {
assert!(parse_match("MATCH (c), shortestPath((a)-[*1..3]->(b))").is_err());
}
#[test]
fn shortest_path_over_disjoint_cross_join_errors() {
assert!(parse_match("MATCH shortestPath((a)-[*1..3]->(b)), (c)").is_err());
}
#[test]
fn shortest_path_not_valid_in_create() {
assert!(parse_statement("CREATE shortestPath((a)-[*1..3]->(b))").is_err());
}
#[test]
fn shortest_path_not_valid_in_merge() {
assert!(parse_merge("MERGE shortestPath((a)-[*1..3]->(b))").is_err());
}
#[test]
fn named_path_over_a_single_variable_length_hop_is_supported() {
let parts = parse_match("MATCH p = (a)-[*1..3]->(b)").unwrap();
assert_eq!(parts[0].path_var.as_deref(), Some("p"));
}
#[test]
fn named_path_over_variable_length_mixed_with_another_hop_is_supported() {
let parts = parse_match("MATCH p = (a)-[*1..3]->(b)-->(c)").unwrap();
assert_eq!(parts[0].path_var.as_deref(), Some("p"));
}
#[test]
fn match_where() {
let parts = parse_match("MATCH (a) WHERE a.x = 1").unwrap();
assert_eq!(parts.len(), 1);
assert!(matches!(
parts[0].where_clause,
Some(Expr::Compare(
PropAccess { .. },
CompareOp::Eq,
Literal::Int(1),
))
));
}
#[test]
fn match_where_var_eq() {
let parts = parse_match("MATCH (a), (b) WHERE a = b").unwrap();
assert!(matches!(parts[1].where_clause, Some(Expr::VarEq(_, _))));
}
#[test]
fn match_where_label_predicate() {
let parts = parse_match("MATCH (a) WHERE a:A:B").unwrap();
assert!(matches!(parts[0].where_clause, Some(Expr::And(_, _))));
}
#[test]
fn match_where_pattern_predicate() {
let parts = parse_match("MATCH (n) WHERE (n)-[]->() RETURN n")
.unwrap_or_else(|e| panic!("expected pattern predicate to parse, got {e:?}"));
let Some(Expr::Pattern(pattern)) = &parts[0].where_clause else {
panic!("expected Expr::Pattern");
};
assert_eq!(pattern.hops.len(), 1);
}
#[test]
fn match_where_pattern_predicate_combined_with_and() {
let parts = parse_match("MATCH (n) WHERE (n)-->() AND n.x = 1").unwrap();
let Some(Expr::And(l, r)) = &parts[0].where_clause else {
panic!("expected Expr::And");
};
assert!(matches!(**l, Expr::Pattern(_)));
assert!(matches!(**r, Expr::Compare(..)));
}
#[test]
fn pattern_predicate_outside_where_still_parses() {
let expr = parse_expr("(n)-->()").unwrap();
assert!(matches!(expr, ReturnExpr::PatternPredicate(_)));
}
#[test]
fn match_where_on_last_group_of_cross_join() {
let parts = parse_match("MATCH (a), (b) WHERE b.x = 1").unwrap();
assert_eq!(parts.len(), 2);
assert!(parts[0].where_clause.is_none());
assert!(parts[1].where_clause.is_some());
}
#[test]
fn arithmetic_precedence() {
assert_eq!(
parse_expr("1 + 2 * 3").unwrap(),
ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(1))),
ArithOp::Add,
Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(2))),
ArithOp::Mul,
Box::new(ReturnExpr::Lit(Literal::Int(3))),
)),
)
);
}
#[test]
fn arithmetic_left_associative() {
assert_eq!(
parse_expr("10 - 2 - 3").unwrap(),
ReturnExpr::Arith(
Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(10))),
ArithOp::Sub,
Box::new(ReturnExpr::Lit(Literal::Int(2))),
)),
ArithOp::Sub,
Box::new(ReturnExpr::Lit(Literal::Int(3))),
)
);
}
#[test]
fn power_left_associative() {
assert_eq!(
parse_expr("4 ^ 3 ^ 2").unwrap(),
ReturnExpr::Arith(
Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(4))),
ArithOp::Pow,
Box::new(ReturnExpr::Lit(Literal::Int(3))),
)),
ArithOp::Pow,
Box::new(ReturnExpr::Lit(Literal::Int(2))),
)
);
}
#[test]
fn binary_minus_no_whitespace() {
assert_eq!(
parse_expr("5-1").unwrap(),
ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(5))),
ArithOp::Sub,
Box::new(ReturnExpr::Lit(Literal::Int(1))),
)
);
}
#[test]
fn unary_minus_on_variable() {
assert_eq!(
parse_expr("-x").unwrap(),
ReturnExpr::Neg(Box::new(ReturnExpr::Var("x".to_string())))
);
}
#[test]
fn unary_minus_folds_into_literal() {
assert_eq!(parse_expr("-5").unwrap(), ReturnExpr::Lit(Literal::Int(-5)));
assert_eq!(
parse_expr("-5.5").unwrap(),
ReturnExpr::Lit(Literal::Float(-5.5))
);
}
#[test]
fn unary_minus_int_min_two_complement_edge_case() {
assert_eq!(
parse_expr("-9223372036854775808").unwrap(),
ReturnExpr::Lit(Literal::Int(i64::MIN))
);
}
#[test]
fn comparison_chain_folds_into_nested_and() {
assert_eq!(
parse_expr("1 < x < 3").unwrap(),
ReturnExpr::And(
Box::new(ReturnExpr::Compare(
Box::new(ReturnExpr::Lit(Literal::Int(1))),
CompareOp::Lt,
Box::new(ReturnExpr::Var("x".to_string())),
)),
Box::new(ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Lt,
Box::new(ReturnExpr::Lit(Literal::Int(3))),
)),
)
);
}
#[test]
fn boolean_operators() {
assert_eq!(
parse_expr("true AND false").unwrap(),
ReturnExpr::And(
Box::new(ReturnExpr::Lit(Literal::Bool(true))),
Box::new(ReturnExpr::Lit(Literal::Bool(false))),
)
);
assert_eq!(
parse_expr("true OR false").unwrap(),
ReturnExpr::Or(
Box::new(ReturnExpr::Lit(Literal::Bool(true))),
Box::new(ReturnExpr::Lit(Literal::Bool(false))),
)
);
assert_eq!(
parse_expr("true XOR false").unwrap(),
ReturnExpr::Xor(
Box::new(ReturnExpr::Lit(Literal::Bool(true))),
Box::new(ReturnExpr::Lit(Literal::Bool(false))),
)
);
}
#[test]
fn double_negation() {
assert_eq!(
parse_expr("NOT NOT true").unwrap(),
ReturnExpr::Not(Box::new(ReturnExpr::Not(Box::new(ReturnExpr::Lit(
Literal::Bool(true)
)))))
);
}
#[test]
fn is_null() {
assert_eq!(
parse_expr("x IS NULL").unwrap(),
ReturnExpr::IsNull(Box::new(ReturnExpr::Var("x".to_string())))
);
assert_eq!(
parse_expr("x IS NOT NULL").unwrap(),
ReturnExpr::Not(Box::new(ReturnExpr::IsNull(Box::new(ReturnExpr::Var(
"x".to_string()
)))))
);
}
#[test]
fn in_operator() {
assert_eq!(
parse_expr("x IN y").unwrap(),
ReturnExpr::In(
Box::new(ReturnExpr::Var("x".to_string())),
Box::new(ReturnExpr::Var("y".to_string())),
)
);
}
#[test]
fn is_null_binds_looser_than_arithmetic() {
assert_eq!(
parse_expr("x + 0 IS NULL").unwrap(),
ReturnExpr::IsNull(Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Var("x".to_string())),
ArithOp::Add,
Box::new(ReturnExpr::Lit(Literal::Int(0))),
)))
);
}
#[test]
fn in_binds_looser_than_arithmetic_and_operand_can_be_sliced() {
assert_eq!(
parse_expr("3 IN [1, 2, 3][0..2]").unwrap(),
ReturnExpr::In(
Box::new(ReturnExpr::Lit(Literal::Int(3))),
Box::new(ReturnExpr::Slice(
Box::new(ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
ReturnExpr::Lit(Literal::Int(3)),
])),
Some(Box::new(ReturnExpr::Lit(Literal::Int(0)))),
Some(Box::new(ReturnExpr::Lit(Literal::Int(2)))),
))
)
);
}
#[test]
fn starts_with_operand_can_be_an_arithmetic_expression() {
assert_eq!(
parse_expr("x STARTS WITH y + z").unwrap(),
ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::StartsWith,
Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Var("y".to_string())),
ArithOp::Add,
Box::new(ReturnExpr::Var("z".to_string())),
)),
)
);
}
#[test]
fn chained_index_postfix_still_works() {
assert_eq!(
parse_expr("[[1, 2], [3, 4]][0][1]").unwrap(),
ReturnExpr::Index(
Box::new(ReturnExpr::Index(
Box::new(ReturnExpr::ListLit(vec![
ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
]),
ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(3)),
ReturnExpr::Lit(Literal::Int(4)),
]),
])),
Box::new(ReturnExpr::Lit(Literal::Int(0))),
)),
Box::new(ReturnExpr::Lit(Literal::Int(1))),
)
);
}
#[test]
fn case_searched_form() {
assert_eq!(
parse_expr("CASE WHEN x > 1 THEN 'big' WHEN x > 0 THEN 'small' ELSE 'none' END")
.unwrap(),
ReturnExpr::Case {
test: None,
whens: vec![
(
ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Gt,
Box::new(ReturnExpr::Lit(Literal::Int(1))),
),
ReturnExpr::Lit(Literal::String("big".to_string())),
),
(
ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Gt,
Box::new(ReturnExpr::Lit(Literal::Int(0))),
),
ReturnExpr::Lit(Literal::String("small".to_string())),
),
],
else_: Some(Box::new(ReturnExpr::Lit(Literal::String(
"none".to_string()
)))),
}
);
}
#[test]
fn case_simple_form_with_test_no_else() {
assert_eq!(
parse_expr("CASE x WHEN 1 THEN 'one' WHEN 2 THEN 'two' END").unwrap(),
ReturnExpr::Case {
test: Some(Box::new(ReturnExpr::Var("x".to_string()))),
whens: vec![
(
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::String("one".to_string())),
),
(
ReturnExpr::Lit(Literal::Int(2)),
ReturnExpr::Lit(Literal::String("two".to_string())),
),
],
else_: None,
}
);
}
#[test]
fn quantifier_none() {
assert_eq!(
parse_expr("none(x IN [1,2] WHERE x > 1)").unwrap(),
ReturnExpr::Quantifier {
kind: QuantifierKind::None,
var: "x".to_string(),
source: Box::new(ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
])),
where_clause: Some(Box::new(ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Gt,
Box::new(ReturnExpr::Lit(Literal::Int(1))),
))),
}
);
}
#[test]
fn quantifier_all_any_single_no_where() {
assert!(matches!(
parse_expr("all(x IN [1]) ").unwrap(),
ReturnExpr::Quantifier {
kind: QuantifierKind::All,
where_clause: None,
..
}
));
assert!(matches!(
parse_expr("any(x IN [1])").unwrap(),
ReturnExpr::Quantifier {
kind: QuantifierKind::Any,
..
}
));
assert!(matches!(
parse_expr("single(x IN [1])").unwrap(),
ReturnExpr::Quantifier {
kind: QuantifierKind::Single,
..
}
));
}
#[test]
fn list_comprehension_with_projection() {
assert_eq!(
parse_expr("[x IN [1,2] WHERE x > 1 | x * 2]").unwrap(),
ReturnExpr::ListComp {
var: "x".to_string(),
source: Box::new(ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
])),
where_clause: Some(Box::new(ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Gt,
Box::new(ReturnExpr::Lit(Literal::Int(1))),
))),
project: Some(Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Var("x".to_string())),
ArithOp::Mul,
Box::new(ReturnExpr::Lit(Literal::Int(2))),
))),
}
);
}
#[test]
fn list_comprehension_with_where_no_project() {
assert_eq!(
parse_expr("[x IN [1,2] WHERE x > 1]").unwrap(),
ReturnExpr::ListComp {
var: "x".to_string(),
source: Box::new(ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
])),
where_clause: Some(Box::new(ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Gt,
Box::new(ReturnExpr::Lit(Literal::Int(1))),
))),
project: None,
}
);
}
#[test]
fn list_comprehension_bare_identity_no_where_no_project() {
assert_eq!(
parse_expr("[x IN [1, 2, 3]]").unwrap(),
ReturnExpr::ListComp {
var: "x".to_string(),
source: Box::new(ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
ReturnExpr::Lit(Literal::Int(3)),
])),
where_clause: None,
project: None,
}
);
}
#[test]
fn string_predicates() {
assert_eq!(
parse_expr("x STARTS WITH y").unwrap(),
ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::StartsWith,
Box::new(ReturnExpr::Var("y".to_string())),
)
);
assert_eq!(
parse_expr("x ENDS WITH y").unwrap(),
ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::EndsWith,
Box::new(ReturnExpr::Var("y".to_string())),
)
);
assert_eq!(
parse_expr("x CONTAINS y").unwrap(),
ReturnExpr::Compare(
Box::new(ReturnExpr::Var("x".to_string())),
CompareOp::Contains,
Box::new(ReturnExpr::Var("y".to_string())),
)
);
}
#[test]
fn index_and_slice() {
assert_eq!(
parse_expr("list[0]").unwrap(),
ReturnExpr::Index(
Box::new(ReturnExpr::Var("list".to_string())),
Box::new(ReturnExpr::Lit(Literal::Int(0))),
)
);
assert_eq!(
parse_expr("list[1..3]").unwrap(),
ReturnExpr::Slice(
Box::new(ReturnExpr::Var("list".to_string())),
Some(Box::new(ReturnExpr::Lit(Literal::Int(1)))),
Some(Box::new(ReturnExpr::Lit(Literal::Int(3)))),
)
);
assert_eq!(
parse_expr("list[..3]").unwrap(),
ReturnExpr::Slice(
Box::new(ReturnExpr::Var("list".to_string())),
None,
Some(Box::new(ReturnExpr::Lit(Literal::Int(3)))),
)
);
assert_eq!(
parse_expr("list[1..]").unwrap(),
ReturnExpr::Slice(
Box::new(ReturnExpr::Var("list".to_string())),
Some(Box::new(ReturnExpr::Lit(Literal::Int(1)))),
None,
)
);
}
#[test]
fn property_access() {
assert_eq!(
parse_expr("n.name").unwrap(),
ReturnExpr::Prop(PropAccess {
var: "n".to_string(),
prop: "name".to_string(),
})
);
}
#[test]
fn property_access_with_backtick_escaped_name() {
assert_eq!(
parse_expr("n.`weird name`").unwrap(),
ReturnExpr::Prop(PropAccess {
var: "n".to_string(),
prop: "weird name".to_string(),
})
);
}
#[test]
fn property_access_on_computed_expr_becomes_prop_of() {
let expr = parse_expr("duration.between(a, b).days").unwrap();
let ReturnExpr::PropOf(base, prop) = expr else {
panic!("expected PropOf, got {expr:?}");
};
assert_eq!(prop, "days");
assert!(matches!(*base, ReturnExpr::Call { .. }));
}
#[test]
fn chained_property_access_folds_left_to_right() {
let expr = parse_expr("a.b.c").unwrap();
let ReturnExpr::PropOf(base, prop) = expr else {
panic!("expected PropOf, got {expr:?}");
};
assert_eq!(prop, "c");
assert_eq!(
*base,
ReturnExpr::Prop(PropAccess {
var: "a".to_string(),
prop: "b".to_string(),
})
);
}
#[test]
fn has_label() {
assert_eq!(
parse_expr("n:Person").unwrap(),
ReturnExpr::HasLabel("n".to_string(), vec!["Person".to_string()])
);
}
#[test]
fn function_call() {
assert_eq!(
parse_expr("size(list)").unwrap(),
ReturnExpr::Call {
name: "size".to_string(),
args: vec![ReturnExpr::Var("list".to_string())],
distinct: false,
}
);
}
#[test]
fn namespaced_function_call() {
assert_eq!(
parse_expr("duration.between(a, b)").unwrap(),
ReturnExpr::Call {
name: "duration.between".to_string(),
args: vec![
ReturnExpr::Var("a".to_string()),
ReturnExpr::Var("b".to_string())
],
distinct: false,
}
);
}
#[test]
fn count_star() {
assert_eq!(parse_expr("count(*)").unwrap(), ReturnExpr::CountStar);
}
#[test]
fn aggregate_distinct() {
assert_eq!(
parse_expr("count(DISTINCT x)").unwrap(),
ReturnExpr::Call {
name: "count".to_string(),
args: vec![ReturnExpr::Var("x".to_string())],
distinct: true,
}
);
}
#[test]
fn distinct_on_non_aggregate_errors() {
assert!(parse_expr("size(DISTINCT x)").is_err());
}
#[test]
fn distinct_on_namespaced_call_errors() {
assert!(parse_expr("duration.between(DISTINCT a, b)").is_err());
}
#[test]
fn parameter_by_name() {
assert_eq!(
parse_expr("$name").unwrap(),
ReturnExpr::Lit(Literal::Param("name".to_string()))
);
}
#[test]
fn parameter_by_position() {
assert_eq!(
parse_expr("$0").unwrap(),
ReturnExpr::Lit(Literal::Param("0".to_string()))
);
}
#[test]
fn parenthesized_expression() {
assert_eq!(
parse_expr("(1 + 2) * 3").unwrap(),
ReturnExpr::Arith(
Box::new(ReturnExpr::Arith(
Box::new(ReturnExpr::Lit(Literal::Int(1))),
ArithOp::Add,
Box::new(ReturnExpr::Lit(Literal::Int(2))),
)),
ArithOp::Mul,
Box::new(ReturnExpr::Lit(Literal::Int(3))),
)
);
}
#[test]
fn return_simple_items() {
let c = parse_return("RETURN a, b.name AS name").unwrap();
let Tail::Return(items, distinct) = c.tail else {
panic!("expected Tail::Return");
};
assert!(!distinct);
assert_eq!(items.len(), 2);
assert_eq!(items[0].expr, ReturnExpr::Var("a".to_string()));
assert_eq!(items[0].alias, None);
assert_eq!(
items[1].expr,
ReturnExpr::Prop(PropAccess {
var: "b".to_string(),
prop: "name".to_string(),
})
);
assert_eq!(items[1].alias.as_deref(), Some("name"));
}
#[test]
fn return_distinct() {
let c = parse_return("RETURN DISTINCT a").unwrap();
let Tail::Return(_, distinct) = c.tail else {
panic!("expected Tail::Return");
};
assert!(distinct);
}
#[test]
fn return_star() {
let c = parse_return("RETURN *").unwrap();
assert!(matches!(c.tail, Tail::ReturnStar(false)));
}
#[test]
fn return_order_by_skip_limit() {
let c = parse_return("RETURN a ORDER BY a DESC SKIP 5 LIMIT 10").unwrap();
let order_by = c.order_by.unwrap();
assert_eq!(order_by.len(), 1);
assert_eq!(order_by[0].0, ReturnExpr::Var("a".to_string()));
assert_eq!(order_by[0].1, SortDir::Desc);
assert_eq!(c.skip, Some(ReturnExpr::Lit(Literal::Int(5))));
assert_eq!(c.limit, Some(ReturnExpr::Lit(Literal::Int(10))));
}
#[test]
fn order_by_default_ascending() {
let c = parse_return("RETURN a ORDER BY a").unwrap();
assert_eq!(c.order_by.unwrap()[0].1, SortDir::Asc);
}
#[test]
fn limit_accepts_arbitrary_expression() {
let c = parse_return("RETURN a LIMIT 1 + 1").unwrap();
assert!(c.limit.is_some());
}
#[test]
fn return_star_with_extra_items_errors() {
assert!(parse_return("RETURN *, x AS y").is_err());
}
#[test]
fn with_items() {
let c = parse_with("WITH a, b.name AS name").unwrap();
assert!(!c.star);
assert!(!c.distinct);
assert_eq!(c.items.len(), 2);
assert_eq!(c.items[0].expr, ReturnExpr::Var("a".to_string()));
assert_eq!(c.items[1].alias.as_deref(), Some("name"));
}
#[test]
fn with_star() {
let c = parse_with("WITH *").unwrap();
assert!(c.star);
assert!(c.items.is_empty());
}
#[test]
fn with_star_and_items() {
let c = parse_with("WITH *, x AS y").unwrap();
assert!(c.star);
assert_eq!(c.items.len(), 1);
assert_eq!(c.items[0].alias.as_deref(), Some("y"));
}
#[test]
fn with_distinct_order_skip_limit() {
let c = parse_with("WITH DISTINCT a ORDER BY a SKIP 1 LIMIT 2").unwrap();
assert!(c.distinct);
assert!(c.order_by.is_some());
assert_eq!(c.skip, Some(ReturnExpr::Lit(Literal::Int(1))));
assert_eq!(c.limit, Some(ReturnExpr::Lit(Literal::Int(2))));
}
#[test]
fn with_where_compare() {
let c = parse_with("WITH a WHERE a.x = 1").unwrap();
let WithExpr::Compare(lhs, op, rhs) = c.where_clause.unwrap() else {
panic!("expected WithExpr::Compare");
};
assert_eq!(
lhs,
ReturnExpr::Prop(PropAccess {
var: "a".to_string(),
prop: "x".to_string()
})
);
assert_eq!(op, CompareOp::Eq);
assert_eq!(rhs, ReturnExpr::Lit(Literal::Int(1)));
}
#[test]
fn with_where_and_or_not() {
let c = parse_with("WITH a WHERE NOT (a.x = 1 AND a.y = 2)").unwrap();
assert!(matches!(c.where_clause.unwrap(), WithExpr::Not(_)));
let c = parse_with("WITH a WHERE a.x = 1 OR a.y = 2").unwrap();
assert!(matches!(c.where_clause.unwrap(), WithExpr::Or(_, _)));
}
#[test]
fn with_where_is_null() {
let c = parse_with("WITH a WHERE a IS NULL").unwrap();
assert!(matches!(c.where_clause.unwrap(), WithExpr::IsNull(_)));
}
#[test]
fn with_where_bare_expression() {
let c = parse_with("WITH n WHERE n:Person").unwrap();
assert!(matches!(c.where_clause.unwrap(), WithExpr::Bare(_)));
}
#[test]
fn with_where_xor_becomes_bare() {
let c = parse_with("WITH a WHERE a.x XOR a.y").unwrap();
assert!(matches!(c.where_clause.unwrap(), WithExpr::Bare(_)));
}
#[test]
fn unwind_basic() {
let c = parse_unwind("UNWIND [1, 2, 3] AS x").unwrap();
assert_eq!(c.var, "x");
assert_eq!(
c.source.0,
ReturnExpr::ListLit(vec![
ReturnExpr::Lit(Literal::Int(1)),
ReturnExpr::Lit(Literal::Int(2)),
ReturnExpr::Lit(Literal::Int(3)),
])
);
assert!(c.where_clause.is_none());
assert!(c.with.is_none());
}
#[test]
fn set_prop() {
let items = parse_set("SET n.name = 'x'").unwrap();
assert_eq!(items.len(), 1);
let SetItem::Prop(prop, value) = &items[0] else {
panic!("expected SetItem::Prop");
};
assert_eq!(prop.var, "n");
assert_eq!(prop.prop, "name");
assert_eq!(*value, ReturnExpr::Lit(Literal::String("x".to_string())));
}
#[test]
fn set_labels() {
let items = parse_set("SET n:A:B").unwrap();
let SetItem::Labels(var, labels) = &items[0] else {
panic!("expected SetItem::Labels");
};
assert_eq!(var, "n");
assert_eq!(labels, &vec!["A".to_string(), "B".to_string()]);
}
#[test]
fn set_map_assign() {
let items = parse_set("SET n = {a: 1}").unwrap();
let SetItem::MapAssign { var, merge, .. } = &items[0] else {
panic!("expected SetItem::MapAssign");
};
assert_eq!(var, "n");
assert!(!merge);
let items = parse_set("SET n += {a: 1}").unwrap();
let SetItem::MapAssign { merge, .. } = &items[0] else {
panic!("expected SetItem::MapAssign");
};
assert!(merge);
}
#[test]
fn set_multiple_items() {
assert_eq!(parse_set("SET n.a = 1, n.b = 2").unwrap().len(), 2);
}
#[test]
fn delete_items() {
let d = parse_delete("DELETE n, r").unwrap();
assert!(!d.detach);
assert_eq!(d.items.len(), 2);
}
#[test]
fn detach_delete() {
let d = parse_delete("DETACH DELETE n").unwrap();
assert!(d.detach);
}
#[test]
fn remove_prop() {
let items = parse_remove("REMOVE n.name").unwrap();
let RemoveItem::Prop(prop) = &items[0] else {
panic!("expected RemoveItem::Prop");
};
assert_eq!(prop.var, "n");
assert_eq!(prop.prop, "name");
}
#[test]
fn remove_labels() {
let items = parse_remove("REMOVE n:A:B").unwrap();
let RemoveItem::Labels(var, labels) = &items[0] else {
panic!("expected RemoveItem::Labels");
};
assert_eq!(var, "n");
assert_eq!(labels, &vec!["A".to_string(), "B".to_string()]);
}
#[test]
fn create_single_pattern() {
let patterns = parse_create("CREATE (a:Person)").unwrap();
assert_eq!(patterns.len(), 1);
assert_eq!(patterns[0].start.var.as_deref(), Some("a"));
}
#[test]
fn create_comma_patterns_stay_separate() {
let patterns = parse_create("CREATE (a), (a)-->(b)").unwrap();
assert_eq!(patterns.len(), 2);
}
#[test]
fn create_named_path_errors() {
assert!(parse_create("CREATE p = (a)-->(b)").is_err());
}
#[test]
fn merge_single_hop() {
let m = parse_merge("MERGE (a)-[:KNOWS]->(b)").unwrap();
assert_eq!(m.pattern.hops.len(), 1);
assert!(m.on_create.is_empty());
assert!(m.on_match.is_empty());
}
#[test]
fn merge_multi_hop_errors() {
assert!(parse_merge("MERGE (a)-->(b)-->(c)").is_err());
}
#[test]
fn merge_named_path_capture() {
let m = parse_merge("MERGE p = (a)-->(b)").unwrap();
assert_eq!(m.path_var.as_deref(), Some("p"));
}
#[test]
fn merge_on_create_on_match() {
let m = parse_merge("MERGE (a) ON CREATE SET a.created = true ON MATCH SET a.seen = true")
.unwrap();
assert_eq!(m.on_create.len(), 1);
assert_eq!(m.on_match.len(), 1);
}
#[test]
fn merge_duplicate_on_create_errors() {
assert!(parse_merge("MERGE (a) ON CREATE SET a.x = 1 ON CREATE SET a.y = 2").is_err());
}
#[test]
fn merge_duplicate_on_match_errors() {
assert!(parse_merge("MERGE (a) ON MATCH SET a.x = 1 ON MATCH SET a.y = 2").is_err());
}
#[test]
fn statement_match_return() {
let s = parse_statement("MATCH (a) RETURN a").unwrap();
let Statement::Match {
clauses,
tail,
order_by,
skip,
limit,
} = s
else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 1);
assert!(matches!(clauses[0], QueryClause::Match(_)));
assert!(matches!(tail, Some(Tail::Return(_, false))));
assert!(order_by.is_none());
assert!(skip.is_none());
assert!(limit.is_none());
}
#[test]
fn statement_return_star() {
let s = parse_statement("MATCH (a) RETURN *").unwrap();
let Statement::Match { tail, .. } = s else {
panic!("expected Statement::Match");
};
assert!(matches!(tail, Some(Tail::ReturnStar(false))));
}
#[test]
fn statement_order_by_skip_limit_on_bare_return() {
let s = parse_statement("MATCH (a) RETURN a ORDER BY a SKIP 1 LIMIT 2").unwrap();
let Statement::Match {
order_by,
skip,
limit,
..
} = s
else {
panic!("expected Statement::Match");
};
assert!(order_by.is_some());
assert_eq!(skip, Some(Box::new(ReturnExpr::Lit(Literal::Int(1)))));
assert_eq!(limit, Some(Box::new(ReturnExpr::Lit(Literal::Int(2)))));
}
#[test]
fn statement_multiple_reading_clauses() {
let s = parse_statement("MATCH (a) UNWIND [1,2] AS x RETURN a, x").unwrap();
let Statement::Match { clauses, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 2);
assert!(matches!(clauses[0], QueryClause::Match(_)));
assert!(matches!(clauses[1], QueryClause::Unwind(_)));
}
#[test]
fn statement_set_becomes_tail_with_return_tail() {
let s = parse_statement("MATCH (n) SET n.x = 1 RETURN n").unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 1);
let Some(Tail::Set(items, Some(ret))) = tail else {
panic!("expected Tail::Set with a ReturnTail");
};
assert_eq!(items.len(), 1);
assert_eq!(ret.items.len(), 1);
}
#[test]
fn statement_set_without_trailing_return() {
let s = parse_statement("MATCH (n) SET n.x = 1").unwrap();
let Statement::Match { tail, .. } = s else {
panic!("expected Statement::Match");
};
assert!(matches!(tail, Some(Tail::Set(_, None))));
}
#[test]
fn statement_detach_delete_tail() {
let s = parse_statement("MATCH (n) DETACH DELETE n").unwrap();
let Statement::Match { tail, .. } = s else {
panic!("expected Statement::Match");
};
assert!(matches!(tail, Some(Tail::DetachDelete(_, None))));
}
#[test]
fn statement_two_updating_clauses_last_becomes_tail() {
let s = parse_statement("MATCH (n) SET n.x = 1 DELETE n RETURN count(n)").unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 2);
assert!(matches!(clauses[1], QueryClause::Set(_)));
assert!(matches!(tail, Some(Tail::Delete(_, Some(_)))));
}
#[test]
fn statement_bare_merge_no_tail() {
let s = parse_statement("MERGE (a)").unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert!(matches!(clauses[0], QueryClause::Merge(_)));
assert!(tail.is_none());
}
#[test]
fn statement_merge_with_trailing_return() {
let s = parse_statement("MERGE (a) RETURN a ORDER BY a").unwrap();
let Statement::Match {
clauses,
tail,
order_by,
..
} = s
else {
panic!("expected Statement::Match");
};
assert!(matches!(clauses[0], QueryClause::Merge(_)));
assert!(matches!(tail, Some(Tail::Return(_, false))));
assert!(order_by.is_some());
}
#[test]
fn statement_bare_match_without_tail_errors() {
assert!(parse_statement("MATCH (n)").is_err());
}
#[test]
fn statement_mutating_tail_order_by_skip_limit_apply_at_statement_level() {
let s =
parse_statement("MATCH (n) SET n.x = 1 RETURN n ORDER BY n.x SKIP 1 LIMIT 2").unwrap();
let Statement::Match {
tail,
order_by,
skip,
limit,
..
} = s
else {
panic!("expected Statement::Match");
};
assert!(matches!(tail, Some(Tail::Set(_, Some(_)))));
assert!(order_by.is_some());
assert_eq!(skip, Some(Box::new(ReturnExpr::Lit(Literal::Int(1)))));
assert_eq!(limit, Some(Box::new(ReturnExpr::Lit(Literal::Int(2)))));
}
#[test]
fn statement_mutating_tail_return_star_errors() {
assert!(parse_statement("MATCH (n) SET n.x = 1 RETURN *").is_err());
}
#[test]
fn statement_create_tail() {
let s = parse_statement("CREATE (a) RETURN a").unwrap();
let Statement::Match { tail, .. } = s else {
panic!("expected Statement::Match");
};
assert!(matches!(tail, Some(Tail::Create(_, Some(_)))));
}
#[test]
fn statement_bare_create_is_not_wrapped_in_match() {
let s = parse_antlr("CREATE (a);").unwrap();
assert!(matches!(s, Statement::Create(_)));
}
#[test]
fn statement_remove_tail() {
let s = parse_statement("MATCH (n) REMOVE n.x").unwrap();
let Statement::Match { tail, .. } = s else {
panic!("expected Statement::Match");
};
assert!(matches!(tail, Some(Tail::Remove(_, None))));
}
#[test]
fn multi_part_with_attaches_to_preceding_match() {
let s = parse_multi_part_statement("MATCH (a:A) WITH a MATCH (b:B) RETURN a, b").unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 2);
let QueryClause::Match(first) = &clauses[0] else {
panic!("expected first clause to be Match");
};
assert!(first.with.is_some());
assert!(matches!(clauses[1], QueryClause::Match(_)));
assert!(matches!(tail, Some(Tail::Return(_, false))));
}
#[test]
fn multi_part_chained_with_second_one_standalone() {
let s = parse_multi_part_statement("MATCH (a:A) WITH a.num AS x WITH x % 3 AS x RETURN x")
.unwrap();
let Statement::Match { clauses, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 2);
let QueryClause::Match(first) = &clauses[0] else {
panic!("expected first clause to be Match");
};
assert!(first.with.is_some());
assert!(matches!(clauses[1], QueryClause::With(_)));
}
#[test]
fn multi_part_set_then_with_stays_separate_entries() {
let s = parse_multi_part_statement(
"MATCH (n:N) WITH n, n.num AS num DELETE n WITH num WHERE num % 2 = 0 RETURN num",
)
.unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 3);
assert!(matches!(clauses[0], QueryClause::Match(_)));
assert!(matches!(clauses[1], QueryClause::Delete { .. }));
assert!(matches!(clauses[2], QueryClause::With(_)));
assert!(matches!(tail, Some(Tail::Return(_, false))));
}
#[test]
fn multi_part_create_with_star_create_create_tail() {
let s =
parse_multi_part_statement("CREATE (a) WITH a WITH * CREATE (b) CREATE (a)<-[:T]-(b)")
.unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 4);
assert!(matches!(clauses[0], QueryClause::Create(_)));
assert!(matches!(clauses[1], QueryClause::With(_)));
assert!(matches!(clauses[2], QueryClause::With(_)));
assert!(matches!(clauses[3], QueryClause::Create(_)));
assert!(matches!(tail, Some(Tail::Create(_, None))));
}
#[test]
fn multi_part_merge_with_attaches() {
let s = parse_multi_part_statement("MERGE (a:A) WITH a MATCH (b:B) RETURN a, b").unwrap();
let Statement::Match { clauses, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 2);
let QueryClause::Merge(m) = &clauses[0] else {
panic!("expected first clause to be Merge");
};
assert!(m.with.is_some());
}
#[test]
fn multi_part_trailing_bare_create_becomes_tail_not_top_level_statement() {
let s = parse_multi_part_statement("MATCH (a) WITH a CREATE (b)").unwrap();
let Statement::Match { clauses, tail, .. } = s else {
panic!("expected Statement::Match");
};
assert_eq!(clauses.len(), 1);
assert!(matches!(clauses[0], QueryClause::Match(_)));
assert!(matches!(tail, Some(Tail::Create(_, None))));
}
#[test]
fn parse_antlr_no_union_passes_through() {
let s = parse_antlr("MATCH (a) RETURN a;").unwrap();
assert!(matches!(s, Statement::Match { .. }));
}
#[test]
fn parse_antlr_union() {
let s = parse_antlr("MATCH (a) RETURN a UNION MATCH (b) RETURN b;").unwrap();
let Statement::Union { parts, all } = s else {
panic!("expected Statement::Union");
};
assert_eq!(parts.len(), 2);
assert!(!all);
}
#[test]
fn parse_antlr_union_all() {
let s = parse_antlr("MATCH (a) RETURN a UNION ALL MATCH (b) RETURN b;").unwrap();
let Statement::Union { parts, all } = s else {
panic!("expected Statement::Union");
};
assert_eq!(parts.len(), 2);
assert!(all);
}
#[test]
fn parse_antlr_union_three_parts() {
let s =
parse_antlr("MATCH (a) RETURN a UNION MATCH (b) RETURN b UNION MATCH (c) RETURN c;")
.unwrap();
let Statement::Union { parts, .. } = s else {
panic!("expected Statement::Union");
};
assert_eq!(parts.len(), 3);
}
#[test]
fn parse_antlr_mixed_union_and_union_all_errors() {
let err = parse_antlr(
"MATCH (a) RETURN a UNION MATCH (b) RETURN b UNION ALL MATCH (c) RETURN c;",
)
.unwrap_err();
assert!(matches!(err, QueryError::Syntax(_)));
}
#[test]
fn parse_antlr_standalone_call() {
let stmt = parse_antlr("CALL db.labels() YIELD label").unwrap();
let Statement::StandaloneCall(call) = stmt else {
panic!("expected a Statement::StandaloneCall, got {stmt:?}");
};
assert_eq!(call.name, "db.labels");
assert_eq!(call.args, Some(vec![]));
assert!(matches!(
call.yield_items,
Some(CallYield::Items(items, None)) if items == vec![("label".to_string(), None)]
));
}
#[test]
fn parse_antlr_syntax_error() {
assert!(parse_antlr("MATCH (a RETURN a;").is_err());
}
#[test]
fn parse_antlr_many_basic() {
let stmts = parse_antlr_many("CREATE (a); CREATE (b); MATCH (n) RETURN n").unwrap();
assert_eq!(stmts.len(), 3);
assert!(matches!(stmts[0], Statement::Create(_)));
assert!(matches!(stmts[2], Statement::Match { .. }));
}
#[test]
fn parse_antlr_many_single_statement() {
let stmts = parse_antlr_many("RETURN 1").unwrap();
assert_eq!(stmts.len(), 1);
}
#[test]
fn parse_antlr_many_strips_single_trailing_semicolon() {
let stmts = parse_antlr_many("CREATE (a);").unwrap();
assert_eq!(stmts.len(), 1);
}
#[test]
fn parse_antlr_many_semicolon_inside_string_literal_not_a_separator() {
let stmts = parse_antlr_many("RETURN ';'").unwrap();
assert_eq!(stmts.len(), 1);
}
#[test]
fn split_statements_respects_all_three_quote_forms() {
assert_eq!(
split_statements("RETURN ';'; RETURN 1"),
vec!["RETURN ';'", " RETURN 1"]
);
assert_eq!(
split_statements(r#"RETURN ";"; RETURN 1"#),
vec![r#"RETURN ";""#, " RETURN 1"]
);
assert_eq!(
split_statements("MATCH (`a;b`) RETURN 1; RETURN 2"),
vec!["MATCH (`a;b`) RETURN 1", " RETURN 2"]
);
}
#[test]
fn split_statements_handles_escaped_quotes_inside_a_literal() {
assert_eq!(
split_statements(r"RETURN 'it\'s; a test'; RETURN 1"),
vec![r"RETURN 'it\'s; a test'", " RETURN 1"]
);
}
#[test]
fn split_statements_backtick_literal_has_no_escapes() {
assert_eq!(
split_statements(r"MATCH (`a\`) RETURN 1; RETURN 2"),
vec![r"MATCH (`a\`) RETURN 1", " RETURN 2"]
);
}
#[test]
fn parse_antlr_create_index() {
let s = parse_antlr("CREATE INDEX ON :Person(name);").unwrap();
let Statement::CreateIndex {
label,
prop,
unique,
} = s
else {
panic!("expected Statement::CreateIndex");
};
assert_eq!(label, "Person");
assert_eq!(prop, "name");
assert!(!unique);
}
#[test]
fn parse_antlr_create_index_unique() {
let s = parse_antlr("CREATE INDEX ON :Person(name) UNIQUE;").unwrap();
let Statement::CreateIndex { unique, .. } = s else {
panic!("expected Statement::CreateIndex");
};
assert!(unique);
}
#[test]
fn parse_antlr_explain_match() {
let s = parse_antlr("EXPLAIN MATCH (a) RETURN a;").unwrap();
let Statement::Explain(inner) = s else {
panic!("expected Statement::Explain");
};
assert!(matches!(*inner, Statement::Match { .. }));
}
#[test]
fn parse_antlr_explain_create_index() {
let s = parse_antlr("EXPLAIN CREATE INDEX ON :Person(name);").unwrap();
let Statement::Explain(inner) = s else {
panic!("expected Statement::Explain");
};
assert!(matches!(*inner, Statement::CreateIndex { .. }));
}
#[test]
fn parse_antlr_index_still_usable_as_property_name() {
let s = parse_antlr("MATCH (a) RETURN a.index;").unwrap();
assert!(matches!(s, Statement::Match { .. }));
}
}