use super::super::ast::*;
use super::super::executor::return_item_column_name;
use super::simplification::{
collect_clause_variables, collect_expression_refs, collect_introduced_variables,
collect_predicate_refs,
};
use super::PassCtx;
use crate::graph::storage::GraphRead;
use std::collections::{HashMap, HashSet};
pub(super) fn pass_hoist_with_where(query: &mut CypherQuery, _ctx: &PassCtx) {
hoist_with_where(query)
}
pub(super) fn hoist_with_where(query: &mut CypherQuery) {
let mut i = 0;
while i + 1 < query.clauses.len() {
if !hoistable(&query.clauses, i) {
i += 1;
continue;
}
let predicate = match &mut query.clauses[i + 1] {
Clause::With(w) => {
w.where_clause
.take()
.expect("hoistable() required a WITH-attached WHERE")
.predicate
}
_ => unreachable!("hoistable() required a WITH at i + 1"),
};
query
.clauses
.insert(i + 1, Clause::Where(WhereClause { predicate }));
i += 2;
}
}
fn hoistable(clauses: &[Clause], i: usize) -> bool {
if !matches!(clauses[i], Clause::Match(_)) {
return false;
}
let w = match &clauses[i + 1] {
Clause::With(w) => w,
_ => return false,
};
let where_clause = match &w.where_clause {
Some(wc) => wc,
None => return false,
};
if w.distinct || w.group_limit_hint.is_some() {
return false;
}
if w.items
.iter()
.any(|item| is_aggregate_expression(&item.expression))
|| predicate_has_aggregate(&where_clause.predicate)
{
return false;
}
let mut bound: HashSet<String> = HashSet::new();
for clause in &clauses[..=i] {
collect_introduced_variables(clause, &mut bound);
}
let mut refs: HashSet<String> = HashSet::new();
collect_predicate_refs(&where_clause.predicate, &mut refs);
if !refs.iter().all(|v| bound.contains(v)) {
return false;
}
if w.items
.iter()
.any(|item| matches!(item.expression, Expression::Star))
{
return true;
}
let mut projected: HashSet<String> = HashSet::new();
collect_introduced_variables(&clauses[i + 1], &mut projected);
refs.iter().all(|v| projected.contains(v))
}
fn predicate_has_aggregate(pred: &Predicate) -> bool {
match pred {
Predicate::Comparison { left, right, .. } => {
is_aggregate_expression(left) || is_aggregate_expression(right)
}
Predicate::And(a, b) | Predicate::Or(a, b) | Predicate::Xor(a, b) => {
predicate_has_aggregate(a) || predicate_has_aggregate(b)
}
Predicate::Not(p) => predicate_has_aggregate(p),
Predicate::IsNull(e) | Predicate::IsNotNull(e) => is_aggregate_expression(e),
Predicate::In { expr, list } => {
is_aggregate_expression(expr) || list.iter().any(is_aggregate_expression)
}
Predicate::InLiteralSet { expr, .. } => is_aggregate_expression(expr),
Predicate::StartsWith { expr, pattern }
| Predicate::EndsWith { expr, pattern }
| Predicate::Contains { expr, pattern } => {
is_aggregate_expression(expr) || is_aggregate_expression(pattern)
}
Predicate::Exists { where_clause, .. } => where_clause
.as_ref()
.is_some_and(|p| predicate_has_aggregate(p)),
Predicate::InExpression { expr, list_expr } => {
is_aggregate_expression(expr) || is_aggregate_expression(list_expr)
}
Predicate::LabelCheck { .. } => false,
}
}
pub(super) fn pass_fold_aliasing_with(query: &mut CypherQuery, ctx: &PassCtx) {
fold_aliasing_with(query, !ctx.graph.graph.is_mapped())
}
pub(super) fn pass_hoist_terminal_return_over_with_top_k(query: &mut CypherQuery, ctx: &PassCtx) {
hoist_terminal_return_over_with_top_k(query, !ctx.graph.graph.is_mapped())
}
fn fold_aliasing_with(query: &mut CypherQuery, permit_count_subquery: bool) {
let Some(i) = foldable_with_index(query) else {
return;
};
let tail = &query.clauses[i + 1..];
if !matches!(tail.first(), Some(Clause::Return(_))) || !tail[1..].iter().all(is_ordering_clause)
{
return;
}
let Some(substitutions) = alias_substitutions(query, i, permit_count_subquery) else {
return;
};
let Some(rewritten) = substitute_tail(tail, &substitutions) else {
return;
};
query.clauses.splice(i.., rewritten);
}
fn hoist_terminal_return_over_with_top_k(query: &mut CypherQuery, permit_count_subquery: bool) {
let Some(i) = foldable_with_index(query) else {
return;
};
let tail = &query.clauses[i + 1..];
let ordering = tail.iter().take_while(|c| is_ordering_clause(c)).count();
if ordering == 0 || tail.len() != ordering + 1 {
return;
}
let Some(Clause::Return(ret)) = tail.get(ordering) else {
return;
};
if ret.distinct
|| ret.having.is_some()
|| ret
.items
.iter()
.any(|item| is_aggregate_expression(&item.expression))
|| ret
.items
.iter()
.any(|item| matches!(item.expression, Expression::WindowFunction { .. }))
{
return;
}
let Some(substitutions) = alias_substitutions(query, i, permit_count_subquery) else {
return;
};
let mut reordered: Vec<Clause> = Vec::with_capacity(tail.len());
reordered.push(tail[ordering].clone());
reordered.extend_from_slice(&tail[..ordering]);
let Some(rewritten) = substitute_tail(&reordered, &substitutions) else {
return;
};
query.clauses.splice(i.., rewritten);
}
fn foldable_with_index(query: &CypherQuery) -> Option<usize> {
query
.clauses
.iter()
.position(|c| matches!(c, Clause::With(w) if with_is_1_to_1(w)))
}
fn with_is_1_to_1(w: &WithClause) -> bool {
!w.distinct
&& w.where_clause.is_none()
&& w.group_limit_hint.is_none()
&& !w.items.iter().any(|item| {
is_aggregate_expression(&item.expression)
|| matches!(item.expression, Expression::WindowFunction { .. })
})
}
fn is_ordering_clause(clause: &Clause) -> bool {
matches!(
clause,
Clause::OrderBy(_) | Clause::Skip(_) | Clause::Limit(_)
)
}
fn alias_substitutions(
query: &CypherQuery,
i: usize,
permit_count_subquery: bool,
) -> Option<HashMap<String, Expression>> {
let Clause::With(w) = &query.clauses[i] else {
return None;
};
let mut bound_before: HashSet<String> = HashSet::new();
for clause in &query.clauses[..i] {
collect_introduced_variables(clause, &mut bound_before);
}
let mut map: HashMap<String, Expression> = HashMap::with_capacity(w.items.len());
let mut has_alias = false;
for item in &w.items {
if !is_substitutable_source(&item.expression, permit_count_subquery) {
return None;
}
let mut refs: HashSet<String> = HashSet::new();
collect_expression_refs(&item.expression, &mut refs);
if !refs.iter().all(|v| bound_before.contains(v)) {
return None;
}
let name = match (&item.alias, &item.expression) {
(Some(alias), _) => alias.clone(),
(None, Expression::Variable(v)) => v.clone(),
(None, _) => return None,
};
if !matches!(&item.expression, Expression::Variable(v) if *v == name)
&& bound_before.contains(&name)
{
return None;
}
if !matches!(&item.expression, Expression::Variable(v) if *v == name) {
has_alias = true;
}
map.insert(name, item.expression.clone());
}
if !has_alias {
return None;
}
let mut downstream: HashSet<String> = HashSet::new();
for clause in &query.clauses[i + 1..] {
collect_clause_variables(clause, &mut downstream);
}
if !downstream
.iter()
.filter(|v| bound_before.contains(*v))
.all(|v| map.contains_key(v))
{
return None;
}
Some(map)
}
fn is_substitutable_source(expr: &Expression, permit_count_subquery: bool) -> bool {
match expr {
Expression::Variable(_)
| Expression::PropertyAccess { .. }
| Expression::Literal(_)
| Expression::Parameter(_) => true,
Expression::Add(l, r)
| Expression::Subtract(l, r)
| Expression::Multiply(l, r)
| Expression::Divide(l, r)
| Expression::Modulo(l, r)
| Expression::Concat(l, r) => {
is_substitutable_source(l, permit_count_subquery)
&& is_substitutable_source(r, permit_count_subquery)
}
Expression::Negate(inner) => is_substitutable_source(inner, permit_count_subquery),
Expression::CountSubquery { where_clause, .. }
if permit_count_subquery && where_clause.is_none() =>
{
true
}
Expression::FunctionCall {
name,
args,
distinct,
} if !*distinct
&& matches!(name.as_str(), "vector_score" | "text_score" | "text_bm25")
&& args
.iter()
.all(|arg| is_substitutable_source(arg, permit_count_subquery)) =>
{
true
}
_ => false,
}
}
fn substitute_tail(tail: &[Clause], map: &HashMap<String, Expression>) -> Option<Vec<Clause>> {
let mut out = Vec::with_capacity(tail.len());
for clause in tail {
out.push(match clause {
Clause::Return(r) => {
if r.having.is_some()
|| r.items
.iter()
.any(|item| matches!(item.expression, Expression::Star))
{
return None;
}
let mut items = Vec::with_capacity(r.items.len());
for item in &r.items {
let column = return_item_column_name(item);
if map.contains_key(&column)
&& !matches!(&item.expression, Expression::Variable(v) if *v == column)
{
return None;
}
items.push(ReturnItem {
expression: substitute_expr(&item.expression, map)?,
alias: Some(column),
});
}
Clause::Return(ReturnClause { items, ..r.clone() })
}
Clause::OrderBy(o) => {
let mut items = Vec::with_capacity(o.items.len());
for item in &o.items {
items.push(OrderItem {
expression: substitute_expr(&item.expression, map)?,
ascending: item.ascending,
nulls: item.nulls,
});
}
Clause::OrderBy(OrderByClause { items })
}
Clause::Skip(s) => Clause::Skip(SkipClause {
count: substitute_expr(&s.count, map)?,
}),
Clause::Limit(l) => Clause::Limit(LimitClause {
count: substitute_expr(&l.count, map)?,
}),
_ => return None,
});
}
Some(out)
}
fn substitute_expr(expr: &Expression, map: &HashMap<String, Expression>) -> Option<Expression> {
let sub = |e: &Expression| substitute_expr(e, map);
Some(match expr {
Expression::Variable(v) => match map.get(v) {
Some(replacement) => replacement.clone(),
None => expr.clone(),
},
Expression::PropertyAccess { variable, property } => match map.get(variable) {
Some(Expression::Variable(v)) => Expression::PropertyAccess {
variable: v.clone(),
property: property.clone(),
},
Some(_) => return None,
None => expr.clone(),
},
Expression::Literal(_) | Expression::Parameter(_) | Expression::Star => expr.clone(),
Expression::CountSubquery { .. } => return None,
Expression::FunctionCall {
name,
args,
distinct,
} => Expression::FunctionCall {
name: name.clone(),
args: args.iter().map(sub).collect::<Option<Vec<_>>>()?,
distinct: *distinct,
},
Expression::Add(l, r) => Expression::Add(Box::new(sub(l)?), Box::new(sub(r)?)),
Expression::Subtract(l, r) => Expression::Subtract(Box::new(sub(l)?), Box::new(sub(r)?)),
Expression::Multiply(l, r) => Expression::Multiply(Box::new(sub(l)?), Box::new(sub(r)?)),
Expression::Divide(l, r) => Expression::Divide(Box::new(sub(l)?), Box::new(sub(r)?)),
Expression::Modulo(l, r) => Expression::Modulo(Box::new(sub(l)?), Box::new(sub(r)?)),
Expression::Concat(l, r) => Expression::Concat(Box::new(sub(l)?), Box::new(sub(r)?)),
Expression::Negate(inner) => Expression::Negate(Box::new(sub(inner)?)),
Expression::ListLiteral(items) => {
Expression::ListLiteral(items.iter().map(sub).collect::<Option<Vec<_>>>()?)
}
Expression::IsNull(inner) => Expression::IsNull(Box::new(sub(inner)?)),
Expression::IsNotNull(inner) => Expression::IsNotNull(Box::new(sub(inner)?)),
Expression::IndexAccess { expr, index } => Expression::IndexAccess {
expr: Box::new(sub(expr)?),
index: Box::new(sub(index)?),
},
Expression::MapLiteral(entries) => Expression::MapLiteral(
entries
.iter()
.map(|(k, v)| Some((k.clone(), sub(v)?)))
.collect::<Option<Vec<_>>>()?,
),
_ => return None,
})
}