use std::collections::HashMap;
use marsdb_graph::PropertyValue;
use crate::ast::{
CallClause, CallYield, Expr, Literal, MergeClause, NodePattern, Pattern, QueryClause,
QueryPart, ReturnExpr, ReturnTail, SetItem, Statement, Tail, UnwindClause, WithClause,
WithExpr,
};
use crate::error::QueryError;
pub fn substitute_params(
stmt: &mut Statement,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match stmt {
Statement::Create(patterns) => {
for pattern in patterns {
substitute_pattern(pattern, params)?;
}
}
Statement::CreateIndex { .. } => {}
Statement::Explain(inner) => substitute_params(inner, params)?,
Statement::Match {
clauses,
tail,
order_by,
skip,
limit,
} => {
for clause in clauses {
substitute_query_clause(clause, params)?;
}
if let Some(tail) = tail {
substitute_tail(tail, params)?;
}
if let Some(items) = order_by {
for (expr, _) in items {
substitute_return_expr(expr, params)?;
}
}
if let Some(expr) = skip {
substitute_return_expr(expr, params)?;
}
if let Some(expr) = limit {
substitute_return_expr(expr, params)?;
}
}
Statement::Union { parts, .. } => {
for part in parts {
substitute_params(part, params)?;
}
}
Statement::StandaloneCall(call) => substitute_call_clause(call, params)?,
}
Ok(())
}
fn substitute_call_clause(
call: &mut CallClause,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
if let Some(args) = &mut call.args {
for arg in args {
substitute_return_expr(arg, params)?;
}
}
if let Some(CallYield::Items(_, Some(where_expr))) = &mut call.yield_items {
substitute_expr(where_expr, params)?;
}
Ok(())
}
fn substitute_query_clause(
clause: &mut QueryClause,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match clause {
QueryClause::Match(part) => substitute_query_part(part, params),
QueryClause::Unwind(u) => substitute_unwind_clause(u, params),
QueryClause::Merge(m) => substitute_merge_clause(m, params),
QueryClause::With(with) => substitute_with_clause(with, params),
QueryClause::Set(items) => {
for item in items {
substitute_set_item(item, params)?;
}
Ok(())
}
QueryClause::Delete { items, detach: _ } => {
for expr in items {
substitute_return_expr(expr, params)?;
}
Ok(())
}
QueryClause::Remove(_) => Ok(()),
QueryClause::Create(patterns) => {
for pattern in patterns {
substitute_pattern(pattern, params)?;
}
Ok(())
}
QueryClause::Call(call) => substitute_call_clause(call, params),
}
}
fn substitute_set_item(
item: &mut SetItem,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match item {
SetItem::Prop(_, value) | SetItem::MapAssign { value, .. } => {
substitute_return_expr(value, params)
}
SetItem::Labels(..) => Ok(()),
}
}
fn substitute_merge_clause(
m: &mut MergeClause,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
substitute_pattern(&mut m.pattern, params)?;
for item in m.on_create.iter_mut().chain(m.on_match.iter_mut()) {
substitute_set_item(item, params)?;
}
if let Some(with) = &mut m.with {
substitute_with_clause(with, params)?;
}
Ok(())
}
fn substitute_query_part(
part: &mut QueryPart,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
substitute_pattern(&mut part.pattern, params)?;
if let Some(expr) = &mut part.where_clause {
substitute_expr(expr, params)?;
}
if let Some(with) = &mut part.with {
substitute_with_clause(with, params)?;
}
Ok(())
}
fn substitute_unwind_clause(
u: &mut UnwindClause,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
substitute_return_expr(&mut u.source.0, params)?;
if let Some(expr) = &mut u.where_clause {
substitute_with_expr(expr, params)?;
}
if let Some(with) = &mut u.with {
substitute_with_clause(with, params)?;
}
Ok(())
}
fn substitute_with_clause(
with: &mut WithClause,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
for item in &mut with.items {
substitute_return_expr(&mut item.expr, params)?;
}
if let Some(where_clause) = &mut with.where_clause {
substitute_with_expr(where_clause, params)?;
}
if let Some(items) = &mut with.order_by {
for (expr, _) in items {
substitute_return_expr(expr, params)?;
}
}
if let Some(expr) = &mut with.skip {
substitute_return_expr(expr, params)?;
}
if let Some(expr) = &mut with.limit {
substitute_return_expr(expr, params)?;
}
Ok(())
}
fn substitute_with_expr(
expr: &mut WithExpr,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match expr {
WithExpr::And(l, r) | WithExpr::Or(l, r) => {
substitute_with_expr(l, params)?;
substitute_with_expr(r, params)?;
}
WithExpr::Not(e) => substitute_with_expr(e, params)?,
WithExpr::Compare(lhs, _, rhs) => {
substitute_return_expr(lhs, params)?;
substitute_return_expr(rhs, params)?;
}
WithExpr::IsNull(e) => substitute_return_expr(e, params)?,
WithExpr::Bare(e) => substitute_return_expr(e, params)?,
}
Ok(())
}
fn substitute_pattern(
pattern: &mut Pattern,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
substitute_node(&mut pattern.start, params)?;
for (rel, node) in &mut pattern.hops {
for (_, expr) in &mut rel.props {
substitute_return_expr(expr, params)?;
}
substitute_node(node, params)?;
}
Ok(())
}
fn substitute_node(
node: &mut NodePattern,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
for (_, expr) in &mut node.props {
substitute_return_expr(expr, params)?;
}
Ok(())
}
fn substitute_expr(
expr: &mut Expr,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match expr {
Expr::And(l, r) | Expr::Or(l, r) => {
substitute_expr(l, params)?;
substitute_expr(r, params)?;
}
Expr::Not(e) => substitute_expr(e, params)?,
Expr::Compare(_, _, lit) => substitute_literal(lit, params)?,
Expr::PropCompare(_, _, _) => {}
Expr::IsNull(_) => {}
Expr::HasLabel(_, _) => {}
Expr::VarEq(_, _) => {}
Expr::EdgeNotInSet { .. } => {}
Expr::GeneralCompare(lhs, _, rhs) => {
substitute_return_expr(lhs, params)?;
substitute_return_expr(rhs, params)?;
}
Expr::GeneralIsNull(e) => substitute_return_expr(e, params)?,
Expr::GeneralBare(e) => substitute_return_expr(e, params)?,
Expr::Pattern(pattern) => substitute_pattern(pattern, params)?,
Expr::Exists {
pattern,
where_clause,
} => {
substitute_pattern(pattern, params)?;
if let Some(w) = where_clause {
substitute_expr(w, params)?;
}
}
Expr::ExistsSubquery(stmt) => substitute_params(stmt, params)?,
}
Ok(())
}
fn substitute_tail(
tail: &mut Tail,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match tail {
Tail::Return(items, _) => {
for item in items {
substitute_return_expr(&mut item.expr, params)?;
}
}
Tail::ReturnStar(_) => {}
Tail::Delete(exprs, ret) | Tail::DetachDelete(exprs, ret) => {
for expr in exprs {
substitute_return_expr(expr, params)?;
}
substitute_return_tail(ret, params)?;
}
Tail::Remove(_, ret) => {
substitute_return_tail(ret, params)?;
}
Tail::Set(items, ret) => {
for item in items {
substitute_set_item(item, params)?;
}
substitute_return_tail(ret, params)?;
}
Tail::Create(patterns, ret) => {
for pattern in patterns {
substitute_pattern(pattern, params)?;
}
substitute_return_tail(ret, params)?;
}
}
Ok(())
}
fn substitute_return_tail(
ret: &mut Option<ReturnTail>,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
if let Some(rt) = ret {
for item in &mut rt.items {
substitute_return_expr(&mut item.expr, params)?;
}
}
Ok(())
}
fn substitute_return_expr(
expr: &mut ReturnExpr,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
match expr {
ReturnExpr::Var(_) | ReturnExpr::Prop(_) | ReturnExpr::CountStar => {}
ReturnExpr::PatternPredicate(pattern) => substitute_pattern(pattern, params)?,
ReturnExpr::Lit(Literal::Param(name)) => {
let value = params
.get(name)
.ok_or_else(|| QueryError::MissingParam(name.clone()))?
.clone();
*expr = property_value_to_return_expr(name, &value)?;
}
ReturnExpr::Lit(_) => {}
ReturnExpr::Call { args, .. } => {
for arg in args {
substitute_return_expr(arg, params)?;
}
}
ReturnExpr::Case { test, whens, else_ } => {
if let Some(t) = test {
substitute_return_expr(t, params)?;
}
for (when, then) in whens {
substitute_return_expr(when, params)?;
substitute_return_expr(then, params)?;
}
if let Some(e) = else_ {
substitute_return_expr(e, params)?;
}
}
ReturnExpr::Arith(l, _, r) => {
substitute_return_expr(l, params)?;
substitute_return_expr(r, params)?;
}
ReturnExpr::Neg(e) => substitute_return_expr(e, params)?,
ReturnExpr::ListLit(items) => {
for item in items {
substitute_return_expr(item, params)?;
}
}
ReturnExpr::Index(base, index) => {
substitute_return_expr(base, params)?;
substitute_return_expr(index, params)?;
}
ReturnExpr::PropOf(base, _) => substitute_return_expr(base, params)?,
ReturnExpr::Slice(base, start, end) => {
substitute_return_expr(base, params)?;
if let Some(s) = start {
substitute_return_expr(s, params)?;
}
if let Some(e) = end {
substitute_return_expr(e, params)?;
}
}
ReturnExpr::ListComp {
source,
where_clause,
project,
..
} => {
substitute_return_expr(source, params)?;
if let Some(w) = where_clause {
substitute_return_expr(w, params)?;
}
if let Some(p) = project {
substitute_return_expr(p, params)?;
}
}
ReturnExpr::Quantifier {
source,
where_clause,
..
} => {
substitute_return_expr(source, params)?;
if let Some(w) = where_clause {
substitute_return_expr(w, params)?;
}
}
ReturnExpr::MapLit(entries) => {
for (_, v) in entries {
substitute_return_expr(v, params)?;
}
}
ReturnExpr::And(l, r) | ReturnExpr::Or(l, r) | ReturnExpr::Xor(l, r) => {
substitute_return_expr(l, params)?;
substitute_return_expr(r, params)?;
}
ReturnExpr::Not(e) => substitute_return_expr(e, params)?,
ReturnExpr::Compare(l, _, r) => {
substitute_return_expr(l, params)?;
substitute_return_expr(r, params)?;
}
ReturnExpr::IsNull(e) => substitute_return_expr(e, params)?,
ReturnExpr::In(needle, haystack) => {
substitute_return_expr(needle, params)?;
substitute_return_expr(haystack, params)?;
}
ReturnExpr::HasLabel(..) => {}
ReturnExpr::PatternComprehension {
pattern,
where_clause,
projection,
..
} => {
substitute_pattern(pattern, params)?;
if let Some(w) = where_clause {
substitute_expr(w, params)?;
}
substitute_return_expr(projection, params)?;
}
ReturnExpr::ExistsPattern {
pattern,
where_clause,
} => {
substitute_pattern(pattern, params)?;
if let Some(w) = where_clause {
substitute_expr(w, params)?;
}
}
ReturnExpr::ExistsSubquery(stmt) => substitute_params(stmt, params)?,
}
Ok(())
}
fn substitute_literal(
lit: &mut Literal,
params: &HashMap<String, PropertyValue>,
) -> Result<(), QueryError> {
if let Literal::Param(name) = lit {
let value = params
.get(name)
.ok_or_else(|| QueryError::MissingParam(name.clone()))?;
*lit = property_value_to_literal(name, value)?;
}
Ok(())
}
fn property_value_to_literal(name: &str, pv: &PropertyValue) -> Result<Literal, QueryError> {
Ok(match pv {
PropertyValue::Null => Literal::Null,
PropertyValue::Bool(b) => Literal::Bool(*b),
PropertyValue::Int(i) => Literal::Int(*i),
PropertyValue::Float(f) => Literal::Float(*f),
PropertyValue::String(s) => Literal::String(s.clone()),
PropertyValue::Date(_)
| PropertyValue::Duration { .. }
| PropertyValue::LocalTime(_)
| PropertyValue::Time { .. }
| PropertyValue::LocalDateTime { .. }
| PropertyValue::DateTime { .. } => {
return Err(QueryError::Type(format!(
"${name}: passing a temporal value as a query parameter isn't supported in a \
pattern-level property comparison (only in ordinary expression position)"
)))
}
PropertyValue::List(_) => {
return Err(QueryError::Type(format!(
"${name}: a list-valued query parameter can't be used here (only in ordinary \
expression position, not a pattern-level property comparison)"
)))
}
PropertyValue::Map(_) => {
return Err(QueryError::Type(format!(
"${name}: a map-valued query parameter can't be used here (only in ordinary \
expression position, not a pattern-level property comparison)"
)))
}
})
}
fn property_value_to_return_expr(name: &str, pv: &PropertyValue) -> Result<ReturnExpr, QueryError> {
Ok(match pv {
PropertyValue::List(items) => ReturnExpr::ListLit(
items
.iter()
.map(|item| property_value_to_return_expr(name, item))
.collect::<Result<Vec<_>, _>>()?,
),
PropertyValue::Map(entries) => ReturnExpr::MapLit(
entries
.iter()
.map(|(key, value)| Ok((key.clone(), property_value_to_return_expr(name, value)?)))
.collect::<Result<Vec<_>, QueryError>>()?,
),
PropertyValue::Date(d) => temporal_call("date", crate::temporal::format_date(*d)),
PropertyValue::Duration {
months,
days,
seconds,
nanos,
} => temporal_call(
"duration",
crate::temporal::format_duration(*months, *days, *seconds, *nanos),
),
PropertyValue::LocalTime(nanos_of_day) => temporal_call(
"localtime",
crate::temporal::format_local_time(*nanos_of_day),
),
PropertyValue::Time {
nanos_of_day,
offset_seconds,
} => temporal_call(
"time",
crate::temporal::format_time(*nanos_of_day, *offset_seconds),
),
PropertyValue::LocalDateTime {
epoch_seconds,
nanos,
} => temporal_call(
"localdatetime",
crate::temporal::format_local_date_time(*epoch_seconds, *nanos),
),
PropertyValue::DateTime {
epoch_seconds,
nanos,
zone,
} => temporal_call(
"datetime",
crate::temporal::format_date_time(
*epoch_seconds,
*nanos,
&crate::executor::tz_from_graph(zone),
),
),
other => ReturnExpr::Lit(property_value_to_literal(name, other)?),
})
}
fn temporal_call(name: &str, formatted: String) -> ReturnExpr {
ReturnExpr::Call {
name: name.to_string(),
args: vec![ReturnExpr::Lit(Literal::String(formatted))],
distinct: false,
}
}