use std::sync::atomic::{AtomicUsize, Ordering};
use rusqlite::Connection;
use crate::cypher::ast::*;
use crate::cypher::ir::*;
use crate::index;
use crate::types::Value;
static ANON_COUNTER: AtomicUsize = AtomicUsize::new(0);
pub fn plan(conn: &Connection, stmt: &Statement) -> crate::types::Result<LogicalOp> {
let mut op = plan_inner(conn, stmt, false)?;
apply_post_passes(conn, &mut op);
Ok(op)
}
pub(in crate::cypher::planner) fn apply_post_passes(conn: &Connection, op: &mut LogicalOp) {
push_limit_into_var_length_expand(op);
rewrite_id_filter_to_lookup(op);
rewrite_text_filter_to_fts(conn, op);
}
pub(in crate::cypher::planner) fn rewrite_id_filter_to_lookup(op: &mut LogicalOp) {
walk_children_mut(op, rewrite_id_filter_to_lookup);
if !matches!(op, LogicalOp::Filter { .. }) {
return;
}
let placeholder = LogicalOp::SingleRow;
let LogicalOp::Filter { input, predicate } = std::mem::replace(op, placeholder) else {
unreachable!()
};
if let LogicalOp::Scan { label, alias } = input.as_ref() {
if label.is_empty() {
if let Some(value_expr) = extract_id_eq_alias(&predicate, alias) {
if !expr_references_var(&value_expr, alias) {
*op = LogicalOp::IdLookup {
alias: alias.clone(),
value_expr,
};
return;
}
}
}
}
if let LogicalOp::CorrelatedJoin { right, .. } = input.as_ref() {
if let LogicalOp::Scan { label, alias } = right.as_ref() {
if label.is_empty() {
if let Some(value_expr) = extract_id_eq_alias(&predicate, alias) {
if !expr_references_var(&value_expr, alias) {
let alias = alias.clone();
let LogicalOp::CorrelatedJoin {
input: join_input,
same_match,
..
} = *input
else {
unreachable!("matched above")
};
*op = LogicalOp::CorrelatedJoin {
input: join_input,
right: Box::new(LogicalOp::IdLookup { alias, value_expr }),
same_match,
};
return;
}
}
}
}
}
*op = LogicalOp::Filter { input, predicate };
}
fn extract_id_eq_alias(predicate: &Expr, alias: &str) -> Option<Expr> {
let (left, right) = match &predicate.kind {
ExprKind::BinaryOp {
left,
op: BinOp::Eq,
right,
} => (left.as_ref(), right.as_ref()),
_ => return None,
};
if matches_id_of_alias(left, alias) {
return Some(right.clone());
}
if matches_id_of_alias(right, alias) {
return Some(left.clone());
}
None
}
fn matches_id_of_alias(expr: &Expr, alias: &str) -> bool {
let ExprKind::FunctionCall { name, args, .. } = &expr.kind else {
return false;
};
if !name.eq_ignore_ascii_case("id") || args.len() != 1 {
return false;
}
matches!(&args[0].kind, ExprKind::Variable(v) if v == alias)
}
pub(in crate::cypher::planner) fn rewrite_text_filter_to_fts(
conn: &Connection,
op: &mut LogicalOp,
) {
match op {
LogicalOp::Expand { input, .. }
| LogicalOp::Filter { input, .. }
| LogicalOp::Project { input, .. }
| LogicalOp::Aggregate { input, .. }
| LogicalOp::Sort { input, .. }
| LogicalOp::Distinct { input }
| LogicalOp::Skip { input, .. }
| LogicalOp::Limit { input, .. }
| LogicalOp::MatchCreate { input, .. }
| LogicalOp::Delete { input, .. }
| LogicalOp::SetProperty { input, .. }
| LogicalOp::SetLabel { input, .. }
| LogicalOp::SetProperties { input, .. }
| LogicalOp::Remove { input, .. }
| LogicalOp::MatchMerge { input, .. }
| LogicalOp::MaterializePath { input, .. }
| LogicalOp::Unwind { input, .. }
| LogicalOp::Call { input, .. }
| LogicalOp::ShortestPath { input, .. } => rewrite_text_filter_to_fts(conn, input),
LogicalOp::CrossProduct { left, right, .. } => {
rewrite_text_filter_to_fts(conn, left);
rewrite_text_filter_to_fts(conn, right);
}
LogicalOp::CorrelatedJoin { input, right, .. }
| LogicalOp::LeftOuterJoin { input, right, .. } => {
rewrite_text_filter_to_fts(conn, input);
rewrite_text_filter_to_fts(conn, right);
}
LogicalOp::Union { inputs, .. } => {
for inp in inputs {
rewrite_text_filter_to_fts(conn, inp);
}
}
LogicalOp::CreateSequence { ops } => {
for inner in ops {
rewrite_text_filter_to_fts(conn, inner);
}
}
LogicalOp::SingleRow
| LogicalOp::Scan { .. }
| LogicalOp::IndexLookup { .. }
| LogicalOp::IdLookup { .. }
| LogicalOp::FullTextLookup { .. }
| LogicalOp::CreateNode { .. }
| LogicalOp::CreateEdge { .. }
| LogicalOp::Merge { .. }
| LogicalOp::CreateIndex { .. }
| LogicalOp::DropIndex { .. }
| LogicalOp::EmptyRow => {}
}
if !matches!(op, LogicalOp::Filter { .. }) {
return;
}
let placeholder = LogicalOp::SingleRow;
let LogicalOp::Filter { input, predicate } = std::mem::replace(op, placeholder) else {
unreachable!()
};
if let LogicalOp::Scan { label, alias } = input.as_ref() {
if !label.is_empty() {
if let Some((prop, fts_op, term, residual, needs_ci)) =
extract_fts_predicate(&predicate, alias)
{
if fts_rewrite_is_usable(conn, label, &prop, needs_ci) {
*op = LogicalOp::FullTextLookup {
label: label.clone(),
alias: alias.clone(),
property: prop,
op: fts_op,
term,
remaining_filters: residual,
};
return;
}
}
if let Some(union) = try_rewrite_or_chain_to_union(conn, label, alias, &predicate) {
*op = union;
return;
}
}
}
if let LogicalOp::CorrelatedJoin { right, .. } = input.as_ref() {
if let LogicalOp::Scan { label, alias } = right.as_ref() {
if !label.is_empty() {
if let Some((prop, fts_op, term, residual, needs_ci)) =
extract_fts_predicate(&predicate, alias)
{
if fts_rewrite_is_usable(conn, label, &prop, needs_ci) {
let label = label.clone();
let alias = alias.clone();
let LogicalOp::CorrelatedJoin {
input: join_input,
same_match,
..
} = *input
else {
unreachable!("matched above")
};
let new_right = LogicalOp::FullTextLookup {
label,
alias,
property: prop,
op: fts_op,
term,
remaining_filters: residual,
};
*op = LogicalOp::CorrelatedJoin {
input: join_input,
right: Box::new(new_right),
same_match,
};
return;
}
}
}
}
}
*op = LogicalOp::Filter { input, predicate };
}
enum FtsKind {
CaseSensitive,
CaseInsensitive,
}
fn fts_kind_for(conn: &Connection, label: &str, property: &str) -> Option<FtsKind> {
match crate::fts::fts_tokenizer_kind(conn, label, property) {
Ok(crate::fts::FtsTokenizerKind::TrigramCaseSensitive) => Some(FtsKind::CaseSensitive),
Ok(crate::fts::FtsTokenizerKind::TrigramCaseInsensitive) => Some(FtsKind::CaseInsensitive),
Ok(crate::fts::FtsTokenizerKind::Word) => None,
Err(_) => None,
}
}
fn fts_rewrite_is_usable(conn: &Connection, label: &str, property: &str, needs_ci: bool) -> bool {
match fts_kind_for(conn, label, property) {
Some(FtsKind::CaseInsensitive) => true,
Some(FtsKind::CaseSensitive) => !needs_ci,
None => false,
}
}
struct MatchedFts {
property: String,
op: crate::cypher::ir::FullTextOp,
term: Expr,
needs_ci: bool,
}
fn extract_fts_predicate(
predicate: &Expr,
alias: &str,
) -> Option<(
String,
crate::cypher::ir::FullTextOp,
Expr,
Option<Expr>,
bool,
)> {
let conjuncts = flatten_top_level_and(predicate);
for (idx, c) in conjuncts.iter().enumerate() {
if let Some(matched) = match_fts_binop(c, alias) {
let MatchedFts {
property: prop,
op: fts_op,
term,
needs_ci,
} = matched;
let residual_parts: Vec<&Expr> = conjuncts
.iter()
.enumerate()
.filter(|(i, _)| *i != idx)
.map(|(_, e)| *e)
.collect();
let residual = rebuild_and(&residual_parts);
return Some((prop, fts_op, term, residual, needs_ci));
}
}
None
}
fn try_rewrite_or_chain_to_union(
conn: &Connection,
label: &str,
alias: &str,
predicate: &Expr,
) -> Option<LogicalOp> {
let disjuncts = flatten_top_level_or(predicate);
if disjuncts.len() < 2 {
return None;
}
let mut inputs: Vec<LogicalOp> = Vec::with_capacity(disjuncts.len());
for d in disjuncts {
let matched = match_fts_binop(d, alias)?;
let MatchedFts {
property: prop,
op: fts_op,
term,
needs_ci,
} = matched;
if !fts_rewrite_is_usable(conn, label, &prop, needs_ci) {
return None;
}
inputs.push(LogicalOp::FullTextLookup {
label: label.to_string(),
alias: alias.to_string(),
property: prop,
op: fts_op,
term,
remaining_filters: None,
});
}
Some(LogicalOp::Union { inputs, all: false })
}
fn match_fts_binop(e: &Expr, alias: &str) -> Option<MatchedFts> {
use crate::cypher::ir::FullTextOp;
let ExprKind::BinaryOp { left, op, right } = &e.kind else {
return None;
};
let fts_op = match op {
BinOp::Contains => FullTextOp::Contains,
BinOp::StartsWith => FullTextOp::StartsWith,
BinOp::EndsWith => FullTextOp::EndsWith,
_ => return None,
};
if let (Some(prop), Some(inner_rhs)) = (
unwrap_tolower_of_property(left, alias),
unwrap_tolower(right),
) {
return Some(MatchedFts {
property: prop,
op: fts_op,
term: inner_rhs,
needs_ci: true,
});
}
let ExprKind::Property(var, prop) = &left.kind else {
return None;
};
if var != alias {
return None;
}
Some(MatchedFts {
property: prop.clone(),
op: fts_op,
term: (**right).clone(),
needs_ci: false,
})
}
fn unwrap_tolower_of_property(e: &Expr, alias: &str) -> Option<String> {
let ExprKind::FunctionCall { name, args, .. } = &e.kind else {
return None;
};
if !name.eq_ignore_ascii_case("toLower") {
return None;
}
if args.len() != 1 {
return None;
}
let ExprKind::Property(var, prop) = &args[0].kind else {
return None;
};
if var != alias {
return None;
}
Some(prop.clone())
}
fn unwrap_tolower(e: &Expr) -> Option<Expr> {
let ExprKind::FunctionCall { name, args, .. } = &e.kind else {
return None;
};
if !name.eq_ignore_ascii_case("toLower") {
return None;
}
if args.len() != 1 {
return None;
}
Some(args[0].clone())
}
fn flatten_top_level_and(e: &Expr) -> Vec<&Expr> {
let mut out = Vec::new();
fn walk<'a>(e: &'a Expr, out: &mut Vec<&'a Expr>) {
if let ExprKind::BinaryOp {
left,
op: BinOp::And,
right,
} = &e.kind
{
walk(left, out);
walk(right, out);
} else {
out.push(e);
}
}
walk(e, &mut out);
out
}
fn flatten_top_level_or(e: &Expr) -> Vec<&Expr> {
let mut out = Vec::new();
fn walk<'a>(e: &'a Expr, out: &mut Vec<&'a Expr>) {
if let ExprKind::BinaryOp {
left,
op: BinOp::Or,
right,
} = &e.kind
{
walk(left, out);
walk(right, out);
} else {
out.push(e);
}
}
walk(e, &mut out);
out
}
fn rebuild_and(parts: &[&Expr]) -> Option<Expr> {
let mut iter = parts.iter().copied().cloned();
let first = iter.next()?;
Some(iter.fold(first, |acc, e| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(acc),
op: BinOp::And,
right: Box::new(e),
})
}))
}
fn expr_references_var(expr: &Expr, name: &str) -> bool {
use ExprKind::*;
match &expr.kind {
Variable(v) => v == name,
Property(v, _) => v == name,
BinaryOp { left, right, .. } => {
expr_references_var(left, name) || expr_references_var(right, name)
}
Not(inner) | IsNull(inner) | IsNotNull(inner) => expr_references_var(inner, name),
FunctionCall { args, .. } => args.iter().any(|a| expr_references_var(a, name)),
Case {
operand,
alternatives,
default,
} => {
operand
.as_deref()
.is_some_and(|e| expr_references_var(e, name))
|| alternatives
.iter()
.any(|(c, r)| expr_references_var(c, name) || expr_references_var(r, name))
|| default
.as_deref()
.is_some_and(|e| expr_references_var(e, name))
}
List(items) => items.iter().any(|e| expr_references_var(e, name)),
_ => false,
}
}
pub(in crate::cypher::planner) fn push_limit_into_var_length_expand(op: &mut LogicalOp) {
if let LogicalOp::Limit { input, count } = op {
if let Some(LogicalOp::Expand { result_cap, .. }) = find_pushdown_target(input) {
let new_cap = match *result_cap {
Some(existing) => existing.min(*count),
None => *count,
};
*result_cap = Some(new_cap);
}
}
walk_children_mut(op, push_limit_into_var_length_expand);
}
pub(in crate::cypher::planner) fn find_pushdown_target(
op: &mut LogicalOp,
) -> Option<&mut LogicalOp> {
match op {
LogicalOp::Project { input, .. } => find_pushdown_target(input),
LogicalOp::Expand { var_length, .. } if *var_length => Some(op),
_ => None,
}
}
pub(in crate::cypher::planner) fn walk_children_mut(op: &mut LogicalOp, f: fn(&mut LogicalOp)) {
match op {
LogicalOp::Expand { input, .. }
| LogicalOp::Filter { input, .. }
| LogicalOp::Project { input, .. }
| LogicalOp::Aggregate { input, .. }
| LogicalOp::Sort { input, .. }
| LogicalOp::Distinct { input }
| LogicalOp::Skip { input, .. }
| LogicalOp::Limit { input, .. }
| LogicalOp::MatchCreate { input, .. }
| LogicalOp::Delete { input, .. }
| LogicalOp::SetProperty { input, .. }
| LogicalOp::SetLabel { input, .. }
| LogicalOp::SetProperties { input, .. }
| LogicalOp::Remove { input, .. }
| LogicalOp::MatchMerge { input, .. }
| LogicalOp::MaterializePath { input, .. }
| LogicalOp::Unwind { input, .. }
| LogicalOp::Call { input, .. }
| LogicalOp::ShortestPath { input, .. } => f(input),
LogicalOp::CrossProduct { left, right, .. } => {
f(left);
f(right);
}
LogicalOp::CorrelatedJoin { input, right, .. }
| LogicalOp::LeftOuterJoin { input, right, .. } => {
f(input);
f(right);
}
LogicalOp::Union { inputs, .. } => {
for inp in inputs {
f(inp);
}
}
LogicalOp::CreateSequence { ops } => {
for inner in ops {
f(inner);
}
}
LogicalOp::SingleRow
| LogicalOp::Scan { .. }
| LogicalOp::IndexLookup { .. }
| LogicalOp::IdLookup { .. }
| LogicalOp::FullTextLookup { .. }
| LogicalOp::CreateNode { .. }
| LogicalOp::CreateEdge { .. }
| LogicalOp::Merge { .. }
| LogicalOp::CreateIndex { .. }
| LogicalOp::DropIndex { .. }
| LogicalOp::EmptyRow => {}
}
}
pub fn plan_subquery(conn: &Connection, stmt: &Statement) -> crate::types::Result<LogicalOp> {
let mut op = plan_inner(conn, stmt, true)?;
apply_post_passes(conn, &mut op);
Ok(op)
}
pub fn plan_with_procedures(
conn: &Connection,
stmt: &Statement,
procedures: &crate::cypher::procedure::ProcedureRegistry,
params: Option<&std::collections::HashMap<String, Value>>,
) -> crate::types::Result<LogicalOp> {
match stmt {
Statement::Call {
procedure_name,
args,
implicit_args,
yield_items,
yield_star,
return_clause,
order_by,
skip,
limit,
} => plan_call(
conn,
procedure_name,
args,
*implicit_args,
yield_items.as_deref(),
*yield_star,
return_clause.as_ref(),
order_by,
skip.as_deref(),
limit.as_deref(),
procedures,
params,
),
Statement::Explain(inner) => plan_with_procedures(conn, inner, procedures, params),
_ => plan(conn, stmt),
}
.map(|mut op| {
apply_post_passes(conn, &mut op);
op
})
}
#[cfg(test)]
mod limit_pushdown_tests {
use super::*;
use crate::cypher::parser;
use rusqlite::Connection;
fn plan_query(query: &str) -> LogicalOp {
let conn = Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
let stmt = parser::parse(query).unwrap();
plan(&conn, &stmt).unwrap()
}
fn find_var_length_expand(op: &LogicalOp) -> Option<&LogicalOp> {
match op {
LogicalOp::Expand {
var_length: true, ..
} => Some(op),
LogicalOp::Limit { input, .. }
| LogicalOp::Project { input, .. }
| LogicalOp::Filter { input, .. }
| LogicalOp::Sort { input, .. }
| LogicalOp::Distinct { input } => find_var_length_expand(input),
_ => None,
}
}
#[test]
fn pushdown_applies_to_simple_var_length_with_limit() {
let plan = plan_query("MATCH (a)-[*1..3]->(b) RETURN b LIMIT 10");
let expand = find_var_length_expand(&plan).expect("expected var-length Expand");
let LogicalOp::Expand { result_cap, .. } = expand else {
unreachable!()
};
assert_eq!(*result_cap, Some(10));
}
#[test]
fn pushdown_skipped_when_sort_intervenes() {
let plan = plan_query("MATCH (a)-[*1..3]->(b) RETURN b ORDER BY b LIMIT 10");
let expand = find_var_length_expand(&plan).expect("expected var-length Expand");
let LogicalOp::Expand { result_cap, .. } = expand else {
unreachable!()
};
assert_eq!(
*result_cap, None,
"Sort between Limit and Expand should block pushdown"
);
}
#[test]
fn pushdown_skipped_for_fixed_length_expand() {
let plan = plan_query("MATCH (a)-[r]->(b) RETURN b LIMIT 10");
match &plan {
LogicalOp::Limit { input, count } => {
assert_eq!(*count, 10);
fn check(op: &LogicalOp) {
if let LogicalOp::Expand { result_cap, .. } = op {
assert_eq!(*result_cap, None);
}
match op {
LogicalOp::Project { input, .. }
| LogicalOp::Expand { input, .. }
| LogicalOp::Filter { input, .. } => check(input),
_ => {}
}
}
check(input);
}
_ => panic!("expected Limit at top, got {:?}", plan.op_name()),
}
}
}
#[cfg(test)]
mod plan_tests {
use super::plan;
use crate::cypher::ir::*;
use crate::cypher::parser::parse;
use crate::types::Direction;
use rusqlite::Connection;
fn plan_query(q: &str) -> LogicalOp {
let conn = Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
let stmt = parse(q).unwrap();
plan(&conn, &stmt).unwrap()
}
#[test]
fn plan_simple_scan() {
let op = plan_query("MATCH (n:Person) RETURN n");
match op {
LogicalOp::Project { input, items, .. } => {
assert_eq!(items.len(), 1);
match *input {
LogicalOp::Scan {
ref label,
ref alias,
} => {
assert_eq!(label, "Person");
assert_eq!(alias, "n");
}
_ => panic!("expected Scan, got {input:?}"),
}
}
_ => panic!("expected Project, got {op:?}"),
}
}
#[test]
fn plan_scan_with_expand() {
let op = plan_query("MATCH (a:Person)-[:KNOWS]->(b:Person) RETURN b");
match op {
LogicalOp::Project { input, .. } => match *input {
LogicalOp::Filter { input, .. } => match *input {
LogicalOp::Expand {
ref src_alias,
ref dst_alias,
ref edge_types,
direction,
min_hops,
max_hops,
..
} => {
assert_eq!(src_alias, "a");
assert_eq!(dst_alias, "b");
assert_eq!(edge_types.first().map(|s| s.as_str()), Some("KNOWS"));
assert_eq!(direction, Direction::Outgoing);
assert_eq!(min_hops, 1);
assert_eq!(max_hops, 1);
}
_ => panic!("expected Expand"),
},
_ => panic!("expected Filter for destination label"),
},
_ => panic!("expected Project"),
}
}
#[test]
fn plan_variable_length_expand() {
let op = plan_query("MATCH (a)-[:CALLS*1..5]->(b) RETURN b");
match op {
LogicalOp::Project { input, .. } => match *input {
LogicalOp::Expand {
min_hops, max_hops, ..
} => {
assert_eq!(min_hops, 1);
assert_eq!(max_hops, 5);
}
_ => panic!("expected Expand"),
},
_ => panic!("expected Project"),
}
}
#[test]
fn plan_with_filter() {
let op = plan_query("MATCH (n:Person) WHERE n.age = 30 RETURN n");
match op {
LogicalOp::Project { input, .. } => match *input {
LogicalOp::Filter { .. } => {}
_ => panic!("expected Filter"),
},
_ => panic!("expected Project"),
}
}
#[test]
fn plan_with_aggregate() {
let op = plan_query("MATCH (n:Person) RETURN count(*) AS cnt");
match op {
LogicalOp::Project { input, .. } => match *input {
LogicalOp::Aggregate { ref aggregates, .. } => {
assert_eq!(aggregates.len(), 1);
assert_eq!(aggregates[0].function, AggregateFunction::Count);
}
_ => panic!("expected Aggregate"),
},
_ => panic!("expected Project"),
}
}
#[test]
fn plan_with_order_by_and_limit() {
let op = plan_query("MATCH (n:Person) RETURN n.name ORDER BY n.name LIMIT 5");
match op {
LogicalOp::Limit { input, count } => {
assert_eq!(count, 5);
match *input {
LogicalOp::Project { input, .. } => match *input {
LogicalOp::Sort { .. } => {}
_ => panic!("expected Sort"),
},
_ => panic!("expected Project"),
}
}
_ => panic!("expected Limit"),
}
}
#[test]
fn plan_create_node() {
let op = plan_query("CREATE (n:Person {name: 'Alice'})");
match op {
LogicalOp::CreateNode {
labels,
alias,
properties,
} => {
assert_eq!(labels, vec!["Person".to_string()]);
assert_eq!(alias.as_deref(), Some("n"));
assert_eq!(properties.len(), 1);
}
_ => panic!("expected CreateNode, got {op:?}"),
}
}
#[test]
fn plan_create_edge() {
let op = plan_query("CREATE (a:Person {name: 'Alice'})-[:KNOWS]->(b:Person {name: 'Bob'})");
match op {
LogicalOp::CreateSequence { ref ops } => {
assert_eq!(ops.len(), 3);
assert!(matches!(ops[0], LogicalOp::CreateNode { .. }));
assert!(matches!(ops[1], LogicalOp::CreateNode { .. }));
assert!(matches!(ops[2], LogicalOp::CreateEdge { .. }));
}
_ => panic!("expected CreateSequence, got {op:?}"),
}
}
#[test]
fn plan_delete() {
let op = plan_query("MATCH (n:Person) WHERE n.name = 'Alice' DELETE n");
match op {
LogicalOp::Delete { exprs, .. } => {
assert_eq!(exprs.len(), 1);
assert!(matches!(
&exprs[0].kind,
crate::cypher::ast::ExprKind::Variable(v) if v == "n"
));
}
_ => panic!("expected Delete"),
}
}
#[test]
fn plan_set_property() {
let op = plan_query("MATCH (n:Person) WHERE n.name = 'Alice' SET n.age = 31");
match op {
LogicalOp::SetProperty { assignments, .. } => {
assert_eq!(assignments.len(), 1);
assert_eq!(assignments[0].property, "age");
}
_ => panic!("expected SetProperty"),
}
}
#[test]
fn plan_merge() {
let op = plan_query(
"MERGE (n:Person {name: 'Alice'}) ON CREATE SET n.created = true ON MATCH SET n.seen = true",
);
match op {
LogicalOp::Merge {
on_create,
on_match,
..
} => {
assert_eq!(on_create.len(), 1);
assert_eq!(on_match.len(), 1);
}
_ => panic!("expected Merge"),
}
}
}
#[cfg(test)]
mod fts_or_chain_tests {
use super::*;
#[test]
fn flatten_top_level_or_splits_chain() {
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let lit = |s: &str| Expr::synthetic(ExprKind::Literal(LiteralValue::String(s.to_string())));
let or = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Or,
right: Box::new(r),
})
};
let chain = or(or(lit("a"), lit("b")), lit("c"));
let parts = flatten_top_level_or(&chain);
assert_eq!(parts.len(), 3);
}
#[test]
fn flatten_top_level_or_returns_single_for_non_or() {
use crate::cypher::ast::{Expr, ExprKind, LiteralValue};
let lit = Expr::synthetic(ExprKind::Literal(LiteralValue::String("x".to_string())));
let parts = flatten_top_level_or(&lit);
assert_eq!(parts.len(), 1);
}
#[test]
fn try_rewrite_or_chain_two_indexed_disjuncts_returns_union() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "title").unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "body").unwrap();
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let prop = |alias: &str, p: &str| {
Expr::synthetic(ExprKind::Property(alias.to_string(), p.to_string()))
};
let lit = |s: &str| Expr::synthetic(ExprKind::Literal(LiteralValue::String(s.to_string())));
let contains = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Contains,
right: Box::new(r),
})
};
let or = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Or,
right: Box::new(r),
})
};
let predicate = or(
contains(prop("n", "title"), lit("x")),
contains(prop("n", "body"), lit("x")),
);
let result = try_rewrite_or_chain_to_union(&conn, "Doc", "n", &predicate);
let plan = result.expect("expected Some(Union)");
match plan {
LogicalOp::Union { inputs, all: false } => {
assert_eq!(inputs.len(), 2);
for inp in &inputs {
assert!(
matches!(inp, LogicalOp::FullTextLookup { .. }),
"expected FullTextLookup, got {inp:?}"
);
}
}
other => panic!("expected Union, got {other:?}"),
}
}
#[test]
fn try_rewrite_or_chain_non_or_returns_none() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "title").unwrap();
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let predicate = Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(Expr::synthetic(ExprKind::Property(
"n".to_string(),
"title".to_string(),
))),
op: BinOp::Contains,
right: Box::new(Expr::synthetic(ExprKind::Literal(LiteralValue::String(
"x".to_string(),
)))),
});
assert!(try_rewrite_or_chain_to_union(&conn, "Doc", "n", &predicate).is_none());
}
#[test]
fn try_rewrite_or_chain_missing_index_returns_none() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "title").unwrap();
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let prop = |alias: &str, p: &str| {
Expr::synthetic(ExprKind::Property(alias.to_string(), p.to_string()))
};
let lit = |s: &str| Expr::synthetic(ExprKind::Literal(LiteralValue::String(s.to_string())));
let contains = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Contains,
right: Box::new(r),
})
};
let or = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Or,
right: Box::new(r),
})
};
let predicate = or(
contains(prop("n", "title"), lit("x")),
contains(prop("n", "body"), lit("x")),
);
assert!(try_rewrite_or_chain_to_union(&conn, "Doc", "n", &predicate).is_none());
}
#[test]
fn try_rewrite_or_chain_non_fts_disjunct_returns_none() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "title").unwrap();
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let prop = |alias: &str, p: &str| {
Expr::synthetic(ExprKind::Property(alias.to_string(), p.to_string()))
};
let lit_str =
|s: &str| Expr::synthetic(ExprKind::Literal(LiteralValue::String(s.to_string())));
let lit_int = |n: i64| Expr::synthetic(ExprKind::Literal(LiteralValue::I64(n)));
let contains = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Contains,
right: Box::new(r),
})
};
let eq = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Eq,
right: Box::new(r),
})
};
let or = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Or,
right: Box::new(r),
})
};
let predicate = or(
contains(prop("n", "title"), lit_str("x")),
eq(prop("n", "id"), lit_int(5)),
);
assert!(try_rewrite_or_chain_to_union(&conn, "Doc", "n", &predicate).is_none());
}
#[test]
fn rewrite_text_filter_to_fts_handles_or_chain() {
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "title").unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "body").unwrap();
let prop = |alias: &str, p: &str| {
Expr::synthetic(ExprKind::Property(alias.to_string(), p.to_string()))
};
let lit = |s: &str| Expr::synthetic(ExprKind::Literal(LiteralValue::String(s.to_string())));
let contains = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Contains,
right: Box::new(r),
})
};
let or = |l: Expr, r: Expr| {
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(l),
op: BinOp::Or,
right: Box::new(r),
})
};
let mut plan = LogicalOp::Filter {
input: Box::new(LogicalOp::Scan {
label: "Doc".to_string(),
alias: "n".to_string(),
}),
predicate: or(
contains(prop("n", "title"), lit("x")),
contains(prop("n", "body"), lit("x")),
),
};
rewrite_text_filter_to_fts(&conn, &mut plan);
assert!(
matches!(&plan, LogicalOp::Union { inputs, all: false } if inputs.len() == 2),
"expected Union(2), got {plan:?}",
);
}
#[test]
fn rewrite_text_filter_to_fts_single_contains_still_works() {
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "title").unwrap();
let mut plan = LogicalOp::Filter {
input: Box::new(LogicalOp::Scan {
label: "Doc".to_string(),
alias: "n".to_string(),
}),
predicate: Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(Expr::synthetic(ExprKind::Property(
"n".to_string(),
"title".to_string(),
))),
op: BinOp::Contains,
right: Box::new(Expr::synthetic(ExprKind::Literal(LiteralValue::String(
"x".to_string(),
)))),
}),
};
rewrite_text_filter_to_fts(&conn, &mut plan);
assert!(
matches!(plan, LogicalOp::FullTextLookup { .. }),
"regression: single CONTAINS must still use AND-chain rewrite"
);
}
fn build_tolower_contains(alias: &str, prop: &str, term: &str) -> Expr {
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(Expr::synthetic(ExprKind::FunctionCall {
name: "toLower".to_string(),
args: vec![Expr::synthetic(ExprKind::Property(
alias.to_string(),
prop.to_string(),
))],
distinct: false,
original_text: None,
})),
op: BinOp::Contains,
right: Box::new(Expr::synthetic(ExprKind::FunctionCall {
name: "toLower".to_string(),
args: vec![Expr::synthetic(ExprKind::Literal(LiteralValue::String(
term.to_string(),
)))],
distinct: false,
original_text: None,
})),
})
}
#[test]
fn tolower_contains_rewrites_when_ci_index_present() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index_ci(&conn, "Doc", "body").unwrap();
let mut op = LogicalOp::Filter {
input: Box::new(LogicalOp::Scan {
label: "Doc".to_string(),
alias: "n".to_string(),
}),
predicate: build_tolower_contains("n", "body", "Alice"),
};
rewrite_text_filter_to_fts(&conn, &mut op);
assert!(
matches!(op, LogicalOp::FullTextLookup { .. }),
"expected FullTextLookup, got {op:?}",
);
}
#[test]
fn tolower_contains_falls_back_when_only_cs_index() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index(&conn, "Doc", "body").unwrap();
let mut op = LogicalOp::Filter {
input: Box::new(LogicalOp::Scan {
label: "Doc".to_string(),
alias: "n".to_string(),
}),
predicate: build_tolower_contains("n", "body", "Alice"),
};
rewrite_text_filter_to_fts(&conn, &mut op);
assert!(
matches!(op, LogicalOp::Filter { .. }),
"CS-only index must not satisfy a needs_ci predicate; got {op:?}",
);
}
#[test]
fn tolower_contains_falls_back_when_no_index() {
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
let mut op = LogicalOp::Filter {
input: Box::new(LogicalOp::Scan {
label: "Doc".to_string(),
alias: "n".to_string(),
}),
predicate: build_tolower_contains("n", "body", "Alice"),
};
rewrite_text_filter_to_fts(&conn, &mut op);
assert!(matches!(op, LogicalOp::Filter { .. }));
}
#[test]
fn asymmetric_tolower_does_not_rewrite_even_with_ci_index() {
use crate::cypher::ast::{BinOp, Expr, ExprKind, LiteralValue};
let conn = rusqlite::Connection::open_in_memory().unwrap();
crate::schema::init_schema(&conn).unwrap();
crate::fts::create_fulltext_index_ci(&conn, "Doc", "body").unwrap();
let predicate = Expr::synthetic(ExprKind::BinaryOp {
left: Box::new(Expr::synthetic(ExprKind::FunctionCall {
name: "toLower".to_string(),
args: vec![Expr::synthetic(ExprKind::Property(
"n".to_string(),
"body".to_string(),
))],
distinct: false,
original_text: None,
})),
op: BinOp::Contains,
right: Box::new(Expr::synthetic(ExprKind::Literal(LiteralValue::String(
"Alice".to_string(),
)))),
});
let mut op = LogicalOp::Filter {
input: Box::new(LogicalOp::Scan {
label: "Doc".to_string(),
alias: "n".to_string(),
}),
predicate,
};
rewrite_text_filter_to_fts(&conn, &mut op);
assert!(matches!(op, LogicalOp::Filter { .. }));
}
}
mod helpers;
mod multi;
mod pattern;
mod statement;
mod validation;
pub use pattern::plan_patterns;
use statement::{plan_call, plan_inner};