#![allow(clippy::result_large_err)]
use std::collections::HashMap;
use super::ast::*;
use crate::datatypes::values::Value;
use crate::error::KgError;
use crate::graph::core::pattern_matching::{ParamLabel, Pattern, PatternElement};
pub fn resolve(query: &mut CypherQuery, params: &HashMap<String, Value>) -> Result<(), KgError> {
resolve_clauses(&mut query.clauses, params)
}
fn bind<'a>(params: &'a HashMap<String, Value>, param: &str) -> Result<&'a str, KgError> {
match params.get(param) {
Some(Value::String(name)) => Ok(name),
Some(other) => Err(execution_error(format!(
"Parameter ${param} is used as a label or relationship type, so it must be a \
string, but a {} was supplied.",
other.type_name()
))),
None => Err(execution_error(format!(
"Missing parameter: ${param} (used as a label or relationship type)"
))),
}
}
fn execution_error(message: String) -> KgError {
KgError::CypherExecution {
message,
position: None,
}
}
fn apply(
markers: &mut Vec<ParamLabel>,
params: &HashMap<String, Value>,
mut slot: impl FnMut(usize, &str),
) -> Result<(), KgError> {
for marker in markers.iter() {
let name = bind(params, &marker.param)?;
slot(marker.slot, name);
}
markers.clear();
Ok(())
}
fn apply_one(
label: &mut String,
marker: &mut Option<String>,
params: &HashMap<String, Value>,
) -> Result<(), KgError> {
if let Some(param) = marker.take() {
*label = bind(params, ¶m)?.to_string();
}
Ok(())
}
fn resolve_clauses(clauses: &mut [Clause], params: &HashMap<String, Value>) -> Result<(), KgError> {
for clause in clauses.iter_mut() {
match clause {
Clause::Match(m) | Clause::OptionalMatch(m) => {
for pattern in &mut m.patterns {
resolve_pattern(pattern, params)?;
}
if let Some(wc) = &mut m.where_clause {
resolve_predicate(&mut wc.predicate, params)?;
}
}
Clause::Where(w) => resolve_predicate(&mut w.predicate, params)?,
Clause::With(w) => {
for item in &mut w.items {
resolve_return_item(item, params)?;
}
if let Some(wc) = &mut w.where_clause {
resolve_predicate(&mut wc.predicate, params)?;
}
}
Clause::Return(r) => {
for item in &mut r.items {
resolve_return_item(item, params)?;
}
}
Clause::Create(c) => {
for pattern in &mut c.patterns {
resolve_create_elements(&mut pattern.elements, params)?;
}
}
Clause::Merge(m) => {
resolve_create_elements(&mut m.pattern.elements, params)?;
for items in [m.on_create.as_mut(), m.on_match.as_mut()]
.into_iter()
.flatten()
{
resolve_set_items(items, params)?;
}
}
Clause::Set(s) => resolve_set_items(&mut s.items, params)?,
Clause::Remove(r) => {
for item in &mut r.items {
if let RemoveItem::Label {
label, label_param, ..
} = item
{
apply_one(label, label_param, params)?;
}
}
}
Clause::Foreach { body, .. } => resolve_clauses(body, params)?,
Clause::CallSubquery { body, .. } => resolve_clauses(&mut body.clauses, params)?,
Clause::Union(u) => resolve_clauses(&mut u.query.clauses, params)?,
_ => {}
}
}
Ok(())
}
fn resolve_pattern(pattern: &mut Pattern, params: &HashMap<String, Value>) -> Result<(), KgError> {
for element in &mut pattern.elements {
match element {
PatternElement::Node(node) => {
if node.label_params.is_empty() {
continue;
}
let mut markers = std::mem::take(&mut node.label_params);
apply(&mut markers, params, |slot, name| match slot {
0 => node.node_type = Some(name.to_string()),
n => {
if let Some(extra) = node.extra_labels.get_mut(n - 1) {
*extra = name.to_string();
}
}
})?;
}
PatternElement::Edge(edge) => {
if edge.type_params.is_empty() {
continue;
}
let mut markers = std::mem::take(&mut edge.type_params);
apply(&mut markers, params, |slot, name| {
if let Some(types) = &mut edge.connection_types {
if let Some(ty) = types.get_mut(slot) {
*ty = name.to_string();
}
}
if slot == 0 {
edge.connection_type = Some(name.to_string());
}
})?;
}
}
}
Ok(())
}
fn resolve_create_elements(
elements: &mut [CreateElement],
params: &HashMap<String, Value>,
) -> Result<(), KgError> {
for element in elements.iter_mut() {
match element {
CreateElement::Node(node) => {
if node.label_params.is_empty() {
continue;
}
let mut markers = std::mem::take(&mut node.label_params);
apply(&mut markers, params, |slot, name| match slot {
0 => node.label = Some(name.to_string()),
n => {
if let Some(extra) = node.extra_labels.get_mut(n - 1) {
*extra = name.to_string();
}
}
})?;
}
CreateElement::Edge(edge) => {
if let Some(param) = edge.type_param.take() {
edge.connection_type = bind(params, ¶m)?.to_string();
}
}
}
}
Ok(())
}
fn resolve_set_items(
items: &mut [SetItem],
params: &HashMap<String, Value>,
) -> Result<(), KgError> {
for item in items.iter_mut() {
match item {
SetItem::Label {
label, label_param, ..
} => apply_one(label, label_param, params)?,
SetItem::Property { expression, .. } | SetItem::Map { expression, .. } => {
resolve_expression(expression, params)?
}
}
}
Ok(())
}
fn resolve_return_item(
item: &mut ReturnItem,
params: &HashMap<String, Value>,
) -> Result<(), KgError> {
resolve_expression(&mut item.expression, params)
}
fn resolve_predicate(pred: &mut Predicate, params: &HashMap<String, Value>) -> Result<(), KgError> {
match pred {
Predicate::LabelCheck {
label, label_param, ..
} => apply_one(label, label_param, params)?,
Predicate::Exists {
patterns,
where_clause,
..
} => {
for pattern in patterns.iter_mut() {
resolve_pattern(pattern, params)?;
}
if let Some(inner) = where_clause {
resolve_predicate(inner, params)?;
}
}
Predicate::And(l, r) | Predicate::Or(l, r) | Predicate::Xor(l, r) => {
resolve_predicate(l, params)?;
resolve_predicate(r, params)?;
}
Predicate::Not(inner) => resolve_predicate(inner, params)?,
Predicate::Comparison { left, right, .. } => {
resolve_expression(left, params)?;
resolve_expression(right, params)?;
}
Predicate::IsNull(expr)
| Predicate::IsNotNull(expr)
| Predicate::InLiteralSet { expr, .. } => resolve_expression(expr, params)?,
Predicate::In { expr, list } => {
resolve_expression(expr, params)?;
for item in list.iter_mut() {
resolve_expression(item, params)?;
}
}
Predicate::StartsWith { expr, pattern }
| Predicate::EndsWith { expr, pattern }
| Predicate::Contains { expr, pattern } => {
resolve_expression(expr, params)?;
resolve_expression(pattern, params)?;
}
Predicate::InExpression { expr, list_expr } => {
resolve_expression(expr, params)?;
resolve_expression(list_expr, params)?;
}
}
Ok(())
}
fn resolve_expression(
expr: &mut Expression,
params: &HashMap<String, Value>,
) -> Result<(), KgError> {
match expr {
Expression::PredicateExpr(pred) => resolve_predicate(pred, params)?,
Expression::CountSubquery {
patterns,
where_clause,
..
} => {
for pattern in patterns.iter_mut() {
resolve_pattern(pattern, params)?;
}
if let Some(inner) = where_clause {
resolve_predicate(inner, params)?;
}
}
Expression::Case {
operand,
when_clauses,
else_expr,
} => {
for inner in operand.iter_mut().chain(else_expr.iter_mut()) {
resolve_expression(inner, params)?;
}
for (when, then) in when_clauses.iter_mut() {
match when {
CaseCondition::Predicate(pred) => resolve_predicate(pred, params)?,
CaseCondition::Expression(expr) => resolve_expression(expr, params)?,
}
resolve_expression(then, params)?;
}
}
Expression::ListComprehension {
list_expr,
filter,
map_expr,
..
} => {
resolve_expression(list_expr, params)?;
if let Some(filter) = filter {
resolve_predicate(filter, params)?;
}
if let Some(map_expr) = map_expr {
resolve_expression(map_expr, params)?;
}
}
Expression::QuantifiedList {
list_expr, filter, ..
} => {
resolve_expression(list_expr, params)?;
resolve_predicate(filter, params)?;
}
other => resolve_operand_expressions(other, params)?,
}
Ok(())
}
fn resolve_operand_expressions(
expr: &mut Expression,
params: &HashMap<String, Value>,
) -> Result<(), KgError> {
match expr {
Expression::Add(l, r)
| Expression::Subtract(l, r)
| Expression::Multiply(l, r)
| Expression::Divide(l, r)
| Expression::Modulo(l, r)
| Expression::Concat(l, r) => {
resolve_expression(l, params)?;
resolve_expression(r, params)?;
}
Expression::Negate(inner)
| Expression::IsNull(inner)
| Expression::IsNotNull(inner)
| Expression::ExprPropertyAccess { expr: inner, .. } => resolve_expression(inner, params)?,
Expression::FunctionCall { args, .. } | Expression::ListLiteral(args) => {
for arg in args.iter_mut() {
resolve_expression(arg, params)?;
}
}
Expression::IndexAccess { expr, index } => {
resolve_expression(expr, params)?;
resolve_expression(index, params)?;
}
Expression::ListSlice { expr, start, end } => {
resolve_expression(expr, params)?;
for inner in start.iter_mut().chain(end.iter_mut()) {
resolve_expression(inner, params)?;
}
}
Expression::MapLiteral(entries) => {
for (_, value) in entries.iter_mut() {
resolve_expression(value, params)?;
}
}
Expression::MapProjection { items, .. } => {
for item in items.iter_mut() {
if let MapProjectionItem::Alias { expr, .. } = item {
resolve_expression(expr, params)?;
}
}
}
Expression::Reduce {
init,
list_expr,
body,
..
} => {
resolve_expression(init, params)?;
resolve_expression(list_expr, params)?;
resolve_expression(body, params)?;
}
Expression::WindowFunction {
partition_by,
order_by,
..
} => {
for inner in partition_by.iter_mut() {
resolve_expression(inner, params)?;
}
for item in order_by.iter_mut() {
resolve_expression(&mut item.expression, params)?;
}
}
_ => {}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::core::pattern_matching::PatternElement;
fn parse(query: &str) -> CypherQuery {
super::super::parser::parse_cypher(query).expect("parse")
}
fn params(pairs: &[(&str, Value)]) -> HashMap<String, Value> {
pairs
.iter()
.map(|(k, v)| (k.to_string(), v.clone()))
.collect()
}
fn first_node_type(query: &CypherQuery) -> Option<String> {
for clause in &query.clauses {
if let Clause::Match(m) = clause {
if let Some(PatternElement::Node(n)) = m.patterns[0].elements.first() {
return n.node_type.clone();
}
}
}
None
}
#[test]
fn an_unresolved_slot_parks_the_source_spelling() {
let query = parse("MATCH (n:$label) RETURN n");
assert_eq!(first_node_type(&query).as_deref(), Some("$label"));
}
#[test]
fn resolution_writes_the_bound_name_and_clears_the_marker() {
let mut query = parse("MATCH (n:$label) RETURN n");
resolve(
&mut query,
¶ms(&[("label", Value::String("Person".into()))]),
)
.unwrap();
assert_eq!(first_node_type(&query).as_deref(), Some("Person"));
let Clause::Match(m) = &query.clauses[0] else {
panic!("expected MATCH");
};
let Some(PatternElement::Node(n)) = m.patterns[0].elements.first() else {
panic!("expected node");
};
assert!(n.label_params.is_empty(), "marker must be cleared");
}
#[test]
fn a_backticked_dollar_label_is_never_treated_as_a_reference() {
let mut query = parse("MATCH (n:`$label`) RETURN n");
resolve(
&mut query,
¶ms(&[("label", Value::String("Person".into()))]),
)
.unwrap();
assert_eq!(first_node_type(&query).as_deref(), Some("$label"));
}
#[test]
fn a_missing_parameter_is_an_error() {
let mut query = parse("MATCH (n:$label) RETURN n");
let err = resolve(&mut query, &HashMap::new()).unwrap_err();
assert!(err.to_string().contains("$label"), "{err}");
}
#[test]
fn a_non_string_parameter_is_an_error() {
let mut query = parse("MATCH (n:$label) RETURN n");
let err = resolve(&mut query, ¶ms(&[("label", Value::Int64(7))])).unwrap_err();
assert!(err.to_string().contains("string"), "{err}");
}
#[test]
fn resolution_reaches_every_name_position() {
let bound = params(&[
("label", Value::String("Person".into())),
("type", Value::String("KNOWS".into())),
]);
for query in [
"MATCH (n:$label) RETURN n",
"MATCH (n:Person:$label) RETURN n",
"MATCH (a)-[:$type]->(b) RETURN a",
"MATCH (a)-[:KNOWS|$type]->(b) RETURN a",
"MATCH (a) WHERE EXISTS { MATCH (a)-[:$type]->(:$label) } RETURN a",
"MATCH (a) WHERE a:$label RETURN a",
"MATCH (a) RETURN COUNT { (a)-[:$type]->(:$label) } AS n",
"CREATE (n:$label {id: 1})",
"MATCH (a), (b) CREATE (a)-[:$type]->(b)",
"MERGE (n:$label {id: 1})",
"MATCH (n) SET n:$label",
"MATCH (n) REMOVE n:$label",
"MATCH (n) FOREACH (x IN [1] | SET n:$label)",
"CALL { MATCH (n:$label) RETURN n } RETURN n",
"MATCH (n:$label) RETURN n UNION MATCH (m:$label) RETURN m",
] {
let mut parsed = parse(query);
resolve(&mut parsed, &bound).unwrap_or_else(|e| panic!("{query}: {e}"));
let rendered = format!("{:?}", parsed);
assert!(
!rendered.contains("$label") && !rendered.contains("$type"),
"{query} left an unresolved name position: {rendered}"
);
}
}
}