use super::{is_aggregate_expression, Clause, CypherQuery, Expression, PassCtx, ReturnItem};
use crate::graph::core::pattern_matching::Pattern;
use crate::graph::core::pattern_matching::PatternElement;
use crate::graph::languages::cypher::ast::{CaseCondition, MapProjectionItem, Predicate};
use crate::graph::schema::DirGraph;
use std::collections::HashSet;
pub(super) fn pass_mark_fast_var_length_paths(query: &mut CypherQuery, _ctx: &PassCtx) {
mark_fast_var_length_paths(query)
}
pub(super) fn pass_mark_disjoint_fixed_trails(query: &mut CypherQuery, _ctx: &PassCtx) {
mark_disjoint_fixed_trails(query)
}
pub(super) fn pass_mark_skip_target_type_check(query: &mut CypherQuery, ctx: &PassCtx) {
mark_skip_target_type_check(query, ctx.graph)
}
fn mark_fast_var_length_paths(query: &mut CypherQuery) {
let clause_count = query.clauses.len();
let consumer_safe: Vec<bool> = (0..clause_count)
.map(|idx| {
matches!(
query.clauses[idx],
Clause::Match(_) | Clause::OptionalMatch(_)
) && consumer_is_dedup_safe(&query.clauses, idx)
})
.collect();
for (idx, clause) in query.clauses.iter_mut().enumerate() {
for_each_exists_subquery(clause, &mut |patterns| {
mark_private_var_length_edges(patterns)
});
if !consumer_safe[idx] {
continue;
}
let mc = match clause {
Clause::Match(mc) | Clause::OptionalMatch(mc) => mc,
_ => continue,
};
if !mc.path_assignments.is_empty() {
continue;
}
mark_private_var_length_edges(&mut mc.patterns);
}
}
fn mark_private_var_length_edges(patterns: &mut [Pattern]) {
let type_sets: Vec<Option<Vec<String>>> = patterns
.iter()
.flat_map(|pattern| pattern.elements.iter())
.filter_map(|element| match element {
PatternElement::Edge(ep) => Some(segment_types(ep)),
PatternElement::Node(_) => None,
})
.collect();
let mut position = 0usize;
for pattern in patterns.iter_mut() {
for element in &mut pattern.elements {
let PatternElement::Edge(ep) = element else {
continue;
};
let here = position;
position += 1;
if ep.variable.is_some() || !ep.var_length.is_some_and(|(min, _)| min <= 1) {
continue;
}
if relationships_are_private(&type_sets, here) {
ep.needs_path_info = false;
}
}
}
}
fn segment_types(edge: &crate::graph::core::pattern_matching::EdgePattern) -> Option<Vec<String>> {
match &edge.connection_types {
Some(types) if !types.is_empty() => Some(types.clone()),
Some(_) => None,
None => edge.connection_type.as_ref().map(|ty| vec![ty.clone()]),
}
}
fn relationships_are_private(type_sets: &[Option<Vec<String>>], here: usize) -> bool {
if type_sets.len() == 1 {
return true;
}
let Some(mine) = type_sets[here].as_ref() else {
return false;
};
type_sets.iter().enumerate().all(|(other, types)| {
other == here
|| types
.as_ref()
.is_some_and(|theirs| theirs.iter().all(|ty| !mine.contains(ty)))
})
}
fn for_each_exists_subquery(clause: &mut Clause, visit: &mut impl FnMut(&mut [Pattern])) {
match clause {
Clause::Match(mc) | Clause::OptionalMatch(mc) => {
if let Some(wc) = &mut mc.where_clause {
walk_predicate(&mut wc.predicate, visit);
}
}
Clause::Where(wc) => walk_predicate(&mut wc.predicate, visit),
Clause::With(wc) => {
for item in &mut wc.items {
walk_expression(&mut item.expression, visit);
}
if let Some(inner) = &mut wc.where_clause {
walk_predicate(&mut inner.predicate, visit);
}
}
Clause::Return(rc) => {
for item in &mut rc.items {
walk_expression(&mut item.expression, visit);
}
if let Some(having) = &mut rc.having {
walk_predicate(having, visit);
}
}
Clause::OrderBy(ob) => {
for item in &mut ob.items {
walk_expression(&mut item.expression, visit);
}
}
Clause::Unwind(uc) => walk_expression(&mut uc.expression, visit),
_ => {}
}
}
fn walk_predicate(pred: &mut Predicate, visit: &mut impl FnMut(&mut [Pattern])) {
match pred {
Predicate::Exists {
patterns,
where_clause,
..
} => {
visit(patterns);
if let Some(inner) = where_clause {
walk_predicate(inner, visit);
}
}
Predicate::And(a, b) | Predicate::Or(a, b) | Predicate::Xor(a, b) => {
walk_predicate(a, visit);
walk_predicate(b, visit);
}
Predicate::Not(inner) => walk_predicate(inner, visit),
Predicate::Comparison { left, right, .. }
| Predicate::StartsWith {
expr: left,
pattern: right,
}
| Predicate::EndsWith {
expr: left,
pattern: right,
}
| Predicate::Contains {
expr: left,
pattern: right,
}
| Predicate::InExpression {
expr: left,
list_expr: right,
} => {
walk_expression(left, visit);
walk_expression(right, visit);
}
Predicate::IsNull(expr)
| Predicate::IsNotNull(expr)
| Predicate::InLiteralSet { expr, .. } => walk_expression(expr, visit),
Predicate::In { expr, list } => {
walk_expression(expr, visit);
for item in list {
walk_expression(item, visit);
}
}
Predicate::LabelCheck { .. } => {}
}
}
fn walk_expression(expr: &mut Expression, visit: &mut impl FnMut(&mut [Pattern])) {
match expr {
Expression::PredicateExpr(pred) => walk_predicate(pred, visit),
Expression::Add(a, b)
| Expression::Subtract(a, b)
| Expression::Multiply(a, b)
| Expression::Divide(a, b)
| Expression::Modulo(a, b)
| Expression::Concat(a, b)
| Expression::IndexAccess { expr: a, index: b } => {
walk_expression(a, visit);
walk_expression(b, visit);
}
Expression::Negate(inner)
| Expression::IsNull(inner)
| Expression::IsNotNull(inner)
| Expression::ExprPropertyAccess { expr: inner, .. } => walk_expression(inner, visit),
Expression::FunctionCall { args, .. } => {
for arg in args {
walk_expression(arg, visit);
}
}
Expression::ListLiteral(items) => {
for item in items {
walk_expression(item, visit);
}
}
Expression::MapLiteral(entries) => {
for (_, value) in entries {
walk_expression(value, visit);
}
}
Expression::MapProjection { items, .. } => {
for item in items {
if let MapProjectionItem::Alias { expr, .. } = item {
walk_expression(expr, visit);
}
}
}
Expression::Case {
operand,
when_clauses,
else_expr,
} => {
if let Some(operand) = operand {
walk_expression(operand, visit);
}
for (condition, result) in when_clauses {
match condition {
CaseCondition::Predicate(pred) => walk_predicate(pred, visit),
CaseCondition::Expression(expr) => walk_expression(expr, visit),
}
walk_expression(result, visit);
}
if let Some(else_expr) = else_expr {
walk_expression(else_expr, visit);
}
}
Expression::ListComprehension {
list_expr,
filter,
map_expr,
..
} => {
walk_expression(list_expr, visit);
if let Some(filter) = filter {
walk_predicate(filter, visit);
}
if let Some(map_expr) = map_expr {
walk_expression(map_expr, visit);
}
}
Expression::ListSlice { expr, start, end } => {
walk_expression(expr, visit);
if let Some(start) = start {
walk_expression(start, visit);
}
if let Some(end) = end {
walk_expression(end, visit);
}
}
Expression::QuantifiedList {
list_expr, filter, ..
} => {
walk_expression(list_expr, visit);
walk_predicate(filter, visit);
}
Expression::Reduce {
init,
list_expr,
body,
..
} => {
walk_expression(init, visit);
walk_expression(list_expr, visit);
walk_expression(body, visit);
}
Expression::WindowFunction {
partition_by,
order_by,
..
} => {
for expr in partition_by {
walk_expression(expr, visit);
}
for item in order_by {
walk_expression(&mut item.expression, visit);
}
}
Expression::CountSubquery { where_clause, .. } => {
if let Some(inner) = where_clause {
walk_predicate(inner, visit);
}
}
Expression::PropertyAccess { .. }
| Expression::Variable(_)
| Expression::Literal(_)
| Expression::Star
| Expression::Parameter(_) => {}
}
}
fn consumer_is_dedup_safe(clauses: &[Clause], idx: usize) -> bool {
for clause in &clauses[idx + 1..] {
match clause {
Clause::Return(r) => return projection_is_dedup_safe(r.distinct, &r.items),
Clause::With(w) => return projection_is_dedup_safe(w.distinct, &w.items),
Clause::Match(_)
| Clause::OptionalMatch(_)
| Clause::Where(_)
| Clause::Unwind(_)
| Clause::OrderBy(_) => continue,
_ => return false,
}
}
false
}
fn projection_is_dedup_safe(distinct: bool, items: &[ReturnItem]) -> bool {
if items.is_empty() {
return false;
}
if distinct {
return items.iter().all(|item| {
!is_aggregate_expression(&item.expression)
|| is_distinct_safe_aggregate(&item.expression)
});
}
items
.iter()
.all(|item| is_distinct_safe_aggregate(&item.expression))
}
fn mark_disjoint_fixed_trails(query: &mut CypherQuery) {
for clause in &mut query.clauses {
let mc = match clause {
Clause::Match(mc) | Clause::OptionalMatch(mc) => mc,
_ => continue,
};
if !mc.path_assignments.is_empty() || mc.patterns.len() != 1 {
continue;
}
let pattern = &mut mc.patterns[0];
if !fixed_edge_types_are_pairwise_disjoint(pattern) {
continue;
}
for element in &mut pattern.elements {
if let PatternElement::Edge(edge) = element {
edge.needs_path_info = false;
}
}
}
}
pub(super) fn fixed_edge_types_are_pairwise_disjoint(
pattern: &crate::graph::core::pattern_matching::Pattern,
) -> bool {
let mut seen = HashSet::new();
let mut edge_count = 0usize;
for element in &pattern.elements {
let PatternElement::Edge(edge) = element else {
continue;
};
edge_count += 1;
if edge.var_length.is_some() {
return false;
}
if let Some(types) = &edge.connection_types {
if types.is_empty() || types.iter().any(|ty| !seen.insert(ty.as_str())) {
return false;
}
} else if let Some(ty) = &edge.connection_type {
if !seen.insert(ty.as_str()) {
return false;
}
} else {
return false;
}
}
edge_count > 0
}
fn is_distinct_safe_aggregate(expr: &Expression) -> bool {
if let Expression::FunctionCall {
name,
args: _,
distinct,
} = expr
{
let nm = name.to_lowercase();
if matches!(nm.as_str(), "min" | "max") {
return true;
}
if *distinct && matches!(nm.as_str(), "count" | "collect") {
return true;
}
}
false
}
fn mark_skip_target_type_check(query: &mut CypherQuery, graph: &DirGraph) {
use crate::graph::core::pattern_matching::EdgeDirection;
for clause in &mut query.clauses {
let mc = match clause {
Clause::Match(mc) | Clause::OptionalMatch(mc) => mc,
_ => continue,
};
for pattern in &mut mc.patterns {
let elements = &mut pattern.elements;
let len = elements.len();
for i in 0..len {
if i + 2 >= len {
break;
}
let (conn_types, direction, target_node_type) = {
let edge = match &elements[i + 1] {
PatternElement::Edge(ep) => ep,
_ => continue,
};
let target = match &elements[i + 2] {
PatternElement::Node(np) => np,
_ => continue,
};
if !target.extra_labels.is_empty() {
continue;
}
let types: Vec<String> = match &edge.connection_types {
Some(types) if !types.is_empty() => types.clone(),
_ => match &edge.connection_type {
Some(ct) => vec![ct.clone()],
None => continue,
},
};
match &target.node_type {
Some(nt) => (types, edge.direction, nt.clone()),
None => continue,
}
};
let guaranteed = conn_types.iter().all(|conn_type| {
graph.connection_type_metadata.get(conn_type).is_some_and(
|info| match direction {
EdgeDirection::Outgoing => {
info.target_types.len() == 1
&& info.target_types.contains(&target_node_type)
}
EdgeDirection::Incoming => {
info.source_types.len() == 1
&& info.source_types.contains(&target_node_type)
}
EdgeDirection::Both => false, },
)
});
if guaranteed {
if let PatternElement::Edge(ep) = &mut elements[i + 1] {
ep.skip_target_type_check = true;
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::{fixed_edge_types_are_pairwise_disjoint, mark_fast_var_length_paths};
use crate::graph::core::pattern_matching::{parse_pattern, PatternElement};
use crate::graph::languages::cypher::parser::parse_cypher;
fn var_length_marks(query: &str) -> Vec<bool> {
let mut parsed = parse_cypher(query).unwrap_or_else(|e| panic!("{query}: {e}"));
mark_fast_var_length_paths(&mut parsed);
let mut marks = Vec::new();
for clause in &parsed.clauses {
let mc = match clause {
super::Clause::Match(mc) | super::Clause::OptionalMatch(mc) => mc,
_ => continue,
};
for pattern in &mc.patterns {
for element in &pattern.elements {
if let PatternElement::Edge(ep) = element {
if ep.var_length.is_some() {
marks.push(ep.needs_path_info);
}
}
}
}
}
marks
}
#[test]
fn dedup_safe_consumer_marks_a_min_one_segment() {
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*1..3]->(b:N) RETURN DISTINCT b.id"),
vec![false]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*0..3]->(b:N) RETURN count(DISTINCT b) AS n"),
vec![false]
);
}
#[test]
fn min_hops_of_two_or_more_is_never_marked() {
for query in [
"MATCH (a:N)-[:R*2..2]->(b:N) RETURN DISTINCT b.id",
"MATCH (a:N)-[:R*2..3]-(b:N) RETURN count(DISTINCT b) AS n",
"MATCH (a:N)-[:R*3..3]->(b:N) RETURN DISTINCT b.id",
"MATCH (a:N)-[:R*5]->(b:N) RETURN DISTINCT b.id",
] {
assert_eq!(var_length_marks(query), vec![true], "{query}");
}
}
#[test]
fn each_clause_is_proved_against_its_own_consumer() {
assert_eq!(
var_length_marks(
"MATCH (a:N)-[:R*1..2]->(b:N) WITH DISTINCT b MATCH (b)-[:R*1..2]->(c:N) RETURN c.id"
),
vec![false, true]
);
assert_eq!(
var_length_marks(
"MATCH (a:N)-[:R*1..2]->(b:N) WITH b MATCH (b)-[:R*1..2]->(c:N) RETURN DISTINCT c.id"
),
vec![true, false]
);
}
#[test]
fn distinct_over_a_row_counting_aggregate_is_not_dedup_safe() {
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*1..3]->(b:N) RETURN DISTINCT a.id, count(b) AS n"),
vec![true]
);
assert_eq!(
var_length_marks(
"MATCH (a:N)-[:R*1..3]->(b:N) RETURN DISTINCT a.id, count(DISTINCT b) AS n"
),
vec![false]
);
}
#[test]
fn a_write_between_the_match_and_its_projection_is_a_barrier() {
assert_eq!(
var_length_marks(
"MATCH (a:N)-[:R*1..2]->(b:N) CREATE (:Log {t: b.id}) RETURN DISTINCT b.id"
),
vec![true]
);
}
#[test]
fn a_path_assignment_or_edge_variable_keeps_the_exact_trail() {
assert_eq!(
var_length_marks("MATCH p = (a:N)-[:R*1..2]->(b:N) RETURN DISTINCT b.id"),
vec![true]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[r:R*1..2]->(b:N) RETURN DISTINCT b.id"),
vec![true]
);
}
fn exists_var_length_marks(query: &str) -> Vec<bool> {
let mut parsed = parse_cypher(query).unwrap_or_else(|e| panic!("{query}: {e}"));
mark_fast_var_length_paths(&mut parsed);
let mut marks = Vec::new();
for clause in &mut parsed.clauses {
super::for_each_exists_subquery(clause, &mut |patterns| {
for pattern in patterns.iter() {
for element in &pattern.elements {
if let PatternElement::Edge(ep) = element {
if ep.var_length.is_some() {
marks.push(ep.needs_path_info);
}
}
}
}
});
}
marks
}
#[test]
fn an_exists_subquery_is_dedup_safe_wherever_it_is_written() {
for query in [
"MATCH (a:N) WHERE EXISTS { (a)-[:R*1..3]->(:N) } RETURN a.id",
"MATCH (a:N) WHERE NOT EXISTS { (a)-[:R*1..3]->(:N) } RETURN a.id",
"MATCH (a:N) RETURN EXISTS { (a)-[:R*1..3]->(:N) } AS reachable",
"MATCH (a:N) RETURN CASE WHEN EXISTS { (a)-[:R*1..3]->(:N) } THEN 1 ELSE 0 END AS r",
"MATCH (a:N) WITH a WHERE EXISTS { (a)-[:R*0..3]->(:N) } RETURN a.id",
"MATCH (a:N) WHERE EXISTS { (a)-[:R*1..3]->(:N) } AND a.id > 1 RETURN a.id",
"MATCH (a:N) OPTIONAL MATCH (a)-[:S]->(b) WHERE EXISTS { (b)-[:R*1..2]->(:N) } RETURN a.id",
] {
assert_eq!(exists_var_length_marks(query), vec![false], "{query}");
}
}
#[test]
fn an_exists_subquery_keeps_the_gates_that_are_about_correctness() {
for query in [
"MATCH (a:N) WHERE EXISTS { (a)-[:R*2..3]->(:N) } RETURN a.id",
"MATCH (a:N) WHERE EXISTS { (a)-[r:R*1..3]->(:N) } RETURN a.id",
] {
assert_eq!(exists_var_length_marks(query), vec![true], "{query}");
}
let mut parsed =
parse_cypher("MATCH (a:N) RETURN COUNT { (a)-[:R*1..3]->(:N) } AS n").unwrap();
mark_fast_var_length_paths(&mut parsed);
assert!(count_subquery_var_length_marks(&parsed)
.iter()
.all(|marked| *marked));
}
fn count_subquery_var_length_marks(query: &super::CypherQuery) -> Vec<bool> {
use super::Expression;
let mut marks = Vec::new();
for clause in &query.clauses {
let super::Clause::Return(rc) = clause else {
continue;
};
for item in &rc.items {
if let Expression::CountSubquery { patterns, .. } = &item.expression {
for pattern in patterns {
for element in &pattern.elements {
if let PatternElement::Edge(ep) = element {
if ep.var_length.is_some() {
marks.push(ep.needs_path_info);
}
}
}
}
}
}
}
assert!(!marks.is_empty(), "no COUNT subquery segment found");
marks
}
#[test]
fn a_segment_sharing_relationships_with_a_sibling_keeps_its_trail() {
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*1..2]->(x)-[:R]->(c) RETURN DISTINCT c.id"),
vec![true]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*1..2]->(x)-->(c) RETURN DISTINCT c.id"),
vec![true]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[*1..2]->(x)-[:R]->(c) RETURN DISTINCT c.id"),
vec![true]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*1..2]->(x), (y)-[:R]->(c) RETURN DISTINCT c.id"),
vec![true]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[:R*1..2]->(x)-[:S]->(c) RETURN DISTINCT c.id"),
vec![false]
);
assert_eq!(
var_length_marks("MATCH (a:N)-[*1..2]->(c) RETURN DISTINCT c.id"),
vec![false]
);
}
#[test]
fn disjoint_fixed_edge_types_need_no_trail() {
let pattern = parse_pattern("(a)-[:JUDGED_BY]-(b)-[:CITES]->(c)").unwrap();
assert!(fixed_edge_types_are_pairwise_disjoint(&pattern));
let single = parse_pattern("(a)-[:CITES]->(b)").unwrap();
assert!(fixed_edge_types_are_pairwise_disjoint(&single));
}
#[test]
fn overlapping_or_unbounded_edge_types_keep_trail() {
for text in [
"(a)-[:CITES]->(b)-[:CITES]->(c)",
"(a)-[:CITES|REFERS_TO]->(b)-[:REFERS_TO]->(c)",
"(a)-->(b)-[:CITES]->(c)",
"(a)-[:CITES*1..2]->(b)-[:REFERS_TO]->(c)",
] {
let pattern = parse_pattern(text).unwrap();
assert!(!fixed_edge_types_are_pairwise_disjoint(&pattern), "{text}");
}
}
}