use reblessive::Stk;
use crate::expr::match_plan::{
BindingId, EdgeQuantifier, EdgeStep, ExpandDirection, MatchClausePlan, MatchPredicate,
NodeStep, PathMode as IrPathMode, PathPrefixPlan, PathSearch, PatternPlan,
};
use crate::expr::{BinaryOperator, Expr, Idiom, Part};
use crate::gql::ast::{
BinaryOp, EdgeDirection, EdgePattern, ElementPredicate, GqlExpr, Ident, LabelExpr, MatchClause,
PathMode, PathPattern, PathPatternPrefix, PathSearchKind, Quantifier, QuantifierKind, UnaryOp,
};
use crate::gql::lower::binding::{ClauseBindings, PatternBindings, Registry};
use crate::gql::lower::expr::{self, Scope};
use crate::syn::error::{SyntaxError, bail};
use crate::syn::token::Span;
use crate::val::TableName;
pub(super) async fn lower_clause(
stk: &mut Stk,
clause_bindings: &ClauseBindings<'_>,
registry: &Registry,
) -> Result<MatchClausePlan, SyntaxError> {
let clause = clause_bindings.clause;
let pattern_bindings = clause_bindings.patterns.as_slice();
let mut patterns = Vec::with_capacity(clause.patterns.len());
for (pattern, bindings) in clause.patterns.iter().zip(pattern_bindings.iter()) {
patterns.push(build_pattern_plan(pattern, bindings)?);
}
let mut predicates = lower_predicates(stk, clause, registry, pattern_bindings).await?;
for bindings in pattern_bindings {
for &(first, repeat) in &bindings.node_equalities {
predicates.push(node_id_equality(registry, first, repeat));
}
}
reject_selective_group_predicates(clause, pattern_bindings, &predicates)?;
Ok(MatchClausePlan {
optional_group: clause_bindings.optional_group,
patterns,
predicates,
})
}
fn reject_selective_group_predicates(
clause: &MatchClause,
pattern_bindings: &[PatternBindings],
predicates: &[MatchPredicate],
) -> Result<(), SyntaxError> {
for (pattern, bindings) in clause.patterns.iter().zip(pattern_bindings.iter()) {
let Some(prefix) = pattern.prefix.as_ref() else {
continue;
};
if matches!(prefix.kind, PathSearchKind::All) {
continue;
}
let group_binding = bindings.steps.first().map(|&(edge, _)| edge);
for predicate in predicates {
let touches_group = group_binding.is_some_and(|g| predicate.deps.contains(&g));
let touches_path = bindings.path_var.is_some_and(|p| predicate.deps.contains(&p));
if touches_group || touches_path {
bail!(
"A path search does not support a predicate over its edge group or path value",
@prefix.span => "constrain the endpoint nodes instead"
);
}
}
}
Ok(())
}
fn node_id_equality(registry: &Registry, first: BindingId, repeat: BindingId) -> MatchPredicate {
let id_idiom = |id: BindingId| {
Expr::Idiom(Idiom(vec![
Part::Field(registry.name(id).to_owned().into()),
Part::Field("id".to_owned().into()),
]))
};
let mut deps = vec![first, repeat];
deps.sort_unstable();
MatchPredicate {
expr: Expr::Binary {
left: Box::new(id_idiom(first)),
op: BinaryOperator::Equal,
right: Box::new(id_idiom(repeat)),
},
deps,
}
}
fn build_pattern_plan(
pattern: &PathPattern,
pattern_bindings: &PatternBindings,
) -> Result<PatternPlan, SyntaxError> {
let start = NodeStep {
binding: pattern_bindings.start,
label: label_table(&pattern.start.label)?,
};
let mut steps = Vec::with_capacity(pattern.steps.len());
for (step, &(edge_binding, node_binding)) in
pattern.steps.iter().zip(pattern_bindings.steps.iter())
{
let edge = &step.edge;
let direction = edge_direction(edge)?;
let quantifier = match &edge.quantifier {
None => None,
Some(quantifier) => Some(lower_quantifier(quantifier)?),
};
let edge_step = EdgeStep {
binding: edge_binding,
label: label_table(&edge.label)?,
direction,
quantifier,
};
let node_step = NodeStep {
binding: node_binding,
label: label_table(&step.node.label)?,
};
steps.push((edge_step, node_step));
}
Ok(PatternPlan {
path_var: pattern_bindings.path_var,
search: pattern.prefix.as_ref().map(lower_path_prefix),
start,
steps,
})
}
fn lower_path_prefix(prefix: &PathPatternPrefix) -> PathPrefixPlan {
let search = match prefix.kind {
PathSearchKind::All => PathSearch::All,
PathSearchKind::Any {
count,
} => PathSearch::Any {
count: count.unwrap_or(1),
},
PathSearchKind::AllShortest => PathSearch::AllShortest,
PathSearchKind::AnyShortest => PathSearch::AnyShortest,
PathSearchKind::ShortestCounted {
count,
} => PathSearch::ShortestCounted {
count,
},
PathSearchKind::ShortestGroups {
count,
} => PathSearch::ShortestGroups {
count: count.unwrap_or(1),
},
};
let mode = match prefix.mode {
None | Some(PathMode::Walk) => IrPathMode::Walk,
Some(PathMode::Trail) => IrPathMode::Trail,
Some(PathMode::Simple) => IrPathMode::Simple,
Some(PathMode::Acyclic) => IrPathMode::Acyclic,
};
PathPrefixPlan {
search,
mode,
}
}
fn edge_direction(edge: &EdgePattern) -> Result<ExpandDirection, SyntaxError> {
match edge.direction {
EdgeDirection::Right => Ok(ExpandDirection::Out),
EdgeDirection::Left => Ok(ExpandDirection::In),
EdgeDirection::Undirected
| EdgeDirection::LeftOrUndirected
| EdgeDirection::UndirectedOrRight
| EdgeDirection::LeftOrRight
| EdgeDirection::Any => bail!(
"Undirected and multi-directional edge patterns are not supported yet",
@edge.span => "use a directed edge: `-[…]->` or `<-[…]-`"
),
}
}
fn label_table(label: &Option<LabelExpr>) -> Result<Option<TableName>, SyntaxError> {
match label {
None => Ok(None),
Some(LabelExpr::Name(ident)) => Ok(Some(TableName::new(ident.name.clone()))),
Some(other) => bail!(
"Label expressions (`!`, `&`, `|`, `%`) are not supported yet",
@other.span() => "use a single label name"
),
}
}
fn lower_quantifier(quantifier: &Quantifier) -> Result<EdgeQuantifier, SyntaxError> {
let (min, max) = match quantifier.kind {
QuantifierKind::Star => (0, None),
QuantifierKind::Plus => (1, None),
QuantifierKind::Question => (0, Some(1)),
QuantifierKind::Fixed(n) => (n, Some(n)),
QuantifierKind::Range(min, max) => (min.unwrap_or(0), max),
};
if let Some(max) = max
&& max < min
{
bail!(
"The quantifier maximum must not be smaller than its minimum",
@quantifier.span
);
}
Ok(EdgeQuantifier {
min,
max,
})
}
enum Conjunct<'a> {
Expr {
expr: &'a GqlExpr,
negated: bool,
},
Prop {
binding: BindingId,
binding_name: &'a str,
key: &'a Ident,
value: &'a GqlExpr,
},
}
async fn lower_predicates(
stk: &mut Stk,
clause: &MatchClause,
registry: &Registry,
pattern_bindings: &[PatternBindings],
) -> Result<Vec<MatchPredicate>, SyntaxError> {
let conjuncts = collect_conjuncts(
clause.where_clause.as_ref(),
&clause.patterns,
pattern_bindings,
registry,
);
let scope = Scope {
registry,
allow_aggregates: false,
};
let mut out = Vec::with_capacity(conjuncts.len());
for conjunct in &conjuncts {
let deps = conjunct_deps(conjunct, registry)?;
validate_quantified_edge_conjunct(conjunct, &deps, &clause.patterns, pattern_bindings)?;
let lowered = match conjunct {
Conjunct::Expr {
expr: predicate,
negated,
} => expr::lower_predicate(stk, predicate, *negated, &scope).await?,
Conjunct::Prop {
binding,
binding_name,
key,
value,
} => expr::lower_prop_equality(stk, *binding, binding_name, key, value, &scope).await?,
};
out.push(MatchPredicate {
expr: lowered.into(),
deps,
});
}
Ok(out)
}
fn collect_conjuncts<'a>(
where_clause: Option<&'a GqlExpr>,
patterns: &'a [PathPattern],
pattern_bindings: &[PatternBindings],
registry: &'a Registry,
) -> Vec<Conjunct<'a>> {
let mut out = Vec::new();
if let Some(where_clause) = where_clause {
split_conjuncts(where_clause, &mut out);
}
for pattern in patterns {
collect_element_where(&pattern.start.predicate, &mut out);
for step in &pattern.steps {
collect_element_where(&step.edge.predicate, &mut out);
collect_element_where(&step.node.predicate, &mut out);
}
}
for (pattern, bindings) in patterns.iter().zip(pattern_bindings.iter()) {
collect_element_props(&pattern.start.predicate, bindings.start, registry, &mut out);
for (step, &(edge_binding, node_binding)) in pattern.steps.iter().zip(bindings.steps.iter())
{
collect_element_props(&step.edge.predicate, edge_binding, registry, &mut out);
collect_element_props(&step.node.predicate, node_binding, registry, &mut out);
}
}
out
}
fn collect_element_where<'a>(predicate: &'a Option<ElementPredicate>, out: &mut Vec<Conjunct<'a>>) {
if let Some(ElementPredicate::Where(expr)) = predicate {
split_conjuncts(expr, out);
}
}
fn collect_element_props<'a>(
predicate: &'a Option<ElementPredicate>,
binding: BindingId,
registry: &'a Registry,
out: &mut Vec<Conjunct<'a>>,
) {
if let Some(ElementPredicate::Props(props)) = predicate {
let binding_name = registry.name(binding);
for (key, value) in props {
out.push(Conjunct::Prop {
binding,
binding_name,
key,
value,
});
}
}
}
fn split_conjuncts<'a>(expr: &'a GqlExpr, out: &mut Vec<Conjunct<'a>>) {
let mut stack: Vec<(&'a GqlExpr, bool)> = vec![(expr, false)];
while let Some((expr, negated)) = stack.pop() {
match expr {
GqlExpr::Unary {
op: UnaryOp::Not,
expr,
..
} => stack.push((expr, !negated)),
GqlExpr::Binary {
op: BinaryOp::And,
left,
right,
..
} if !negated => {
stack.push((right, negated));
stack.push((left, negated));
}
GqlExpr::Binary {
op: BinaryOp::Or,
left,
right,
..
} if negated => {
stack.push((right, negated));
stack.push((left, negated));
}
_ => out.push(Conjunct::Expr {
expr,
negated,
}),
}
}
}
fn conjunct_deps(
conjunct: &Conjunct<'_>,
registry: &Registry,
) -> Result<Vec<BindingId>, SyntaxError> {
let mut deps: Vec<BindingId> = Vec::new();
match conjunct {
Conjunct::Expr {
expr,
..
} => {
walk_variables(expr, &mut |ident, _| {
let id = registry.resolve(ident)?;
if !deps.contains(&id) {
deps.push(id);
}
Ok(())
})?;
}
Conjunct::Prop {
binding,
value,
..
} => {
deps.push(*binding);
walk_variables(value, &mut |ident, _| {
let id = registry.resolve(ident)?;
if !deps.contains(&id) {
deps.push(id);
}
Ok(())
})?;
}
}
deps.sort_unstable();
Ok(deps)
}
fn validate_quantified_edge_conjunct(
conjunct: &Conjunct<'_>,
deps: &[BindingId],
patterns: &[PathPattern],
pattern_bindings: &[PatternBindings],
) -> Result<(), SyntaxError> {
for (pattern, bindings) in patterns.iter().zip(pattern_bindings.iter()) {
for (step, &(edge_binding, _)) in pattern.steps.iter().zip(bindings.steps.iter()) {
if step.edge.quantifier.is_none() {
continue;
}
if deps.contains(&edge_binding) && deps.iter().any(|d| *d != edge_binding) {
bail!(
"A predicate inside a quantified edge may only reference that edge",
@conjunct_span(conjunct) => "remove the references to other variables"
);
}
}
}
Ok(())
}
fn conjunct_span(conjunct: &Conjunct<'_>) -> Span {
match conjunct {
Conjunct::Expr {
expr,
..
} => expr.span(),
Conjunct::Prop {
value,
..
} => value.span(),
}
}
fn walk_variables<'a>(
expr: &'a GqlExpr,
visit: &mut impl FnMut(&'a Ident, bool) -> Result<(), SyntaxError>,
) -> Result<(), SyntaxError> {
let mut stack = vec![expr];
while let Some(e) = stack.pop() {
match e {
GqlExpr::Variable(ident) => visit(ident, true)?,
GqlExpr::Property(base, _, _) => {
let mut base = &**base;
while let GqlExpr::Property(inner, _, _) = base {
base = inner;
}
if let GqlExpr::Variable(ident) = base {
visit(ident, false)?;
} else {
stack.push(base);
}
}
GqlExpr::Unary {
expr,
..
}
| GqlExpr::IsBool {
expr,
..
}
| GqlExpr::IsNull {
expr,
..
} => stack.push(expr),
GqlExpr::Binary {
left,
right,
..
} => {
stack.push(right);
stack.push(left);
}
GqlExpr::FunctionCall {
args,
..
} => stack.extend(args.iter()),
GqlExpr::List(items, _) => stack.extend(items.iter()),
GqlExpr::Map(fields, _) => stack.extend(fields.iter().map(|(_, value)| value)),
GqlExpr::Literal(..)
| GqlExpr::Param {
..
} => {}
}
}
Ok(())
}