use std::collections::HashMap;
use rusqlite::Connection;
use crate::cypher::ast::{Expr, ExprKind, ReturnItem};
use crate::cypher::eval::{eval_expr, eval_predicate, expr_to_column_name};
use crate::cypher::executor::{
index_lookup_ids, is_user_visible_field, literal_to_value, node_to_record,
};
use crate::cypher::ir::*;
use crate::cypher::record::NamedRecord;
use crate::node;
use crate::types::{Direction, Result, Value};
pub trait RecordIter {
fn next_record(&mut self) -> Result<Option<NamedRecord>>;
}
pub fn collect_all(iter: &mut dyn RecordIter) -> Result<Vec<NamedRecord>> {
let mut records = Vec::new();
while let Some(rec) = iter.next_record()? {
records.push(rec);
}
Ok(records)
}
pub struct VecIter {
records: std::vec::IntoIter<NamedRecord>,
}
impl VecIter {
pub fn new(records: Vec<NamedRecord>) -> Self {
Self {
records: records.into_iter(),
}
}
}
impl RecordIter for VecIter {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
Ok(self.records.next())
}
}
#[derive(Default)]
pub struct EmptyRowIter {
done: bool,
}
impl EmptyRowIter {
pub fn new() -> Self {
Self::default()
}
}
impl RecordIter for EmptyRowIter {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
if self.done {
Ok(None)
} else {
self.done = true;
Ok(Some(NamedRecord::new()))
}
}
}
pub struct FilterIter<'a> {
input: Box<dyn RecordIter + 'a>,
predicate: Expr,
conn: &'a Connection,
}
impl<'a> RecordIter for FilterIter<'a> {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
while let Some(rec) = self.input.next_record()? {
if eval_predicate(
&self.predicate,
&rec,
crate::cypher::eval::EvalCx::new(self.conn),
)? {
return Ok(Some(rec));
}
}
Ok(None)
}
}
pub struct SkipIter<'a> {
input: Box<dyn RecordIter + 'a>,
remaining_to_skip: u64,
}
impl<'a> RecordIter for SkipIter<'a> {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
while self.remaining_to_skip > 0 {
if self.input.next_record()?.is_none() {
return Ok(None);
}
self.remaining_to_skip -= 1;
}
self.input.next_record()
}
}
pub struct LimitIter<'a> {
input: Box<dyn RecordIter + 'a>,
remaining: u64,
}
impl<'a> RecordIter for LimitIter<'a> {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
if self.remaining == 0 {
return Ok(None);
}
if let Some(rec) = self.input.next_record()? {
self.remaining -= 1;
Ok(Some(rec))
} else {
Ok(None)
}
}
}
pub struct ProjectIter<'a> {
input: Box<dyn RecordIter + 'a>,
items: Vec<ReturnItem>,
emit_compound: bool,
conn: &'a Connection,
}
impl<'a> RecordIter for ProjectIter<'a> {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
let rec = match self.input.next_record()? {
Some(r) => r,
None => return Ok(None),
};
let mut projected = NamedRecord::new();
for item in &self.items {
match &item.expr.kind {
ExprKind::Star => {
if self.emit_compound {
let bound_vars = crate::cypher::executor::compound_binding_vars(&rec);
for var in &bound_vars {
if let Some(compound) =
crate::cypher::executor::build_compound_binding(&rec, var)
{
projected.set(var.clone(), compound);
}
}
for (key, val) in &rec.fields {
if !is_user_visible_field(key) {
continue;
}
let owner = key.split_once('.').map(|(v, _)| v);
if let Some(owner) = owner {
if bound_vars.iter().any(|v| v == owner) {
continue;
}
}
projected.set(key.clone(), val.clone());
}
} else {
for (key, val) in &rec.fields {
projected.set(key.clone(), val.clone());
}
}
}
ExprKind::Variable(var) => {
let col_name = item.alias.clone().unwrap_or_else(|| var.clone());
if self.emit_compound {
if let Some(compound) =
crate::cypher::executor::build_compound_binding(&rec, var)
{
projected.set(col_name, compound);
} else if let Some(existing) = rec.get(&col_name) {
projected.set(col_name, existing.clone());
} else {
let val = eval_expr(
&item.expr,
&rec,
crate::cypher::eval::EvalCx::new(self.conn),
)?;
projected.set(col_name, val);
}
} else {
let src_prefix = format!("{var}.");
let dst_prefix = format!("{col_name}.");
let mut propagated_any = false;
for (key, val) in &rec.fields {
if let Some(rest) = key.strip_prefix(&src_prefix) {
projected.set(format!("{dst_prefix}{rest}"), val.clone());
propagated_any = true;
}
}
if let Some(existing) = rec.get(var) {
projected.set(col_name.clone(), existing.clone());
propagated_any = true;
}
if !propagated_any {
let val = eval_expr(
&item.expr,
&rec,
crate::cypher::eval::EvalCx::new(self.conn),
)?;
projected.set(col_name, val);
}
}
}
_ => {
let col_name = item
.alias
.clone()
.unwrap_or_else(|| expr_to_column_name(&item.expr));
let val = if let Some(existing) = rec.get(&col_name) {
existing.clone()
} else {
eval_expr(
&item.expr,
&rec,
crate::cypher::eval::EvalCx::new(self.conn),
)?
};
projected.set(col_name, val);
}
}
}
Ok(Some(projected))
}
}
pub struct ExpandIter<'a> {
input: Box<dyn RecordIter + 'a>,
conn: &'a Connection,
src_alias: String,
dst_alias: String,
rel_alias: Option<String>,
edge_types: Vec<String>,
direction: Direction,
min_hops: u32,
max_hops: u32,
var_length: bool,
var_length_prop_filters: HashMap<String, Value>,
max_traversal_work: u64,
buffer: std::vec::IntoIter<NamedRecord>,
}
impl<'a> ExpandIter<'a> {
#[allow(clippy::too_many_arguments)]
pub(crate) fn new(
input: Box<dyn RecordIter + 'a>,
conn: &'a Connection,
src_alias: String,
dst_alias: String,
rel_alias: Option<String>,
edge_types: Vec<String>,
direction: Direction,
min_hops: u32,
max_hops: u32,
var_length: bool,
var_length_prop_filters: HashMap<String, Value>,
max_traversal_work: u64,
) -> Self {
Self {
input,
conn,
src_alias,
dst_alias,
rel_alias,
edge_types,
direction,
min_hops,
max_hops,
var_length,
var_length_prop_filters,
max_traversal_work,
buffer: Vec::new().into_iter(),
}
}
}
impl<'a> RecordIter for ExpandIter<'a> {
fn next_record(&mut self) -> Result<Option<NamedRecord>> {
loop {
if let Some(rec) = self.buffer.next() {
return Ok(Some(rec));
}
let rec = match self.input.next_record()? {
Some(r) => r,
None => return Ok(None),
};
let expanded = crate::cypher::executor::expand_record(
self.conn,
&rec,
&self.src_alias,
&self.dst_alias,
self.rel_alias.as_deref(),
&self.edge_types,
self.direction,
self.min_hops,
self.max_hops,
self.var_length,
&self.var_length_prop_filters,
None,
self.max_traversal_work,
)?;
self.buffer = expanded.into_iter();
}
}
}
pub fn build_iter<'a>(
conn: &'a Connection,
plan: &'a LogicalOp,
max_traversal_work: u64,
) -> Result<Box<dyn RecordIter + 'a>> {
match plan {
LogicalOp::EmptyRow => Ok(Box::new(EmptyRowIter::new())),
LogicalOp::Scan { label, alias } => {
let nodes = node::find_nodes_by_label(conn, label)?;
let records: Vec<NamedRecord> =
nodes.iter().map(|n| node_to_record(n, alias)).collect();
Ok(Box::new(VecIter::new(records)))
}
LogicalOp::IndexLookup {
label,
alias,
index_properties,
lookups,
remaining_filters,
} => {
let node_ids = index_lookup_ids(conn, label, index_properties, lookups)?;
let mut records = Vec::new();
for id in node_ids {
let n = node::get_node(conn, id)?;
let rec = node_to_record(&n, alias);
if let Some(filter) = remaining_filters {
if !eval_predicate(filter, &rec, crate::cypher::eval::EvalCx::new(conn))? {
continue;
}
}
records.push(rec);
}
Ok(Box::new(VecIter::new(records)))
}
LogicalOp::Filter { input, predicate } => {
let input_iter = build_iter(conn, input, max_traversal_work)?;
Ok(Box::new(FilterIter {
input: input_iter,
predicate: predicate.clone(),
conn,
}))
}
LogicalOp::Skip { input, count } => {
let input_iter = build_iter(conn, input, max_traversal_work)?;
Ok(Box::new(SkipIter {
input: input_iter,
remaining_to_skip: *count,
}))
}
LogicalOp::Limit { input, count } => {
let input_iter = build_iter(conn, input, max_traversal_work)?;
Ok(Box::new(LimitIter {
input: input_iter,
remaining: *count,
}))
}
LogicalOp::Project {
input,
items,
emit_compound,
} => {
let input_iter = build_iter(conn, input, max_traversal_work)?;
Ok(Box::new(ProjectIter {
input: input_iter,
items: items.clone(),
emit_compound: *emit_compound,
conn,
}))
}
LogicalOp::Expand {
input,
src_alias,
dst_alias,
rel_alias,
edge_types,
direction,
min_hops,
max_hops,
var_length,
var_length_prop_filters,
result_cap: _,
} => {
let input_iter = build_iter(conn, input, max_traversal_work)?;
let prop_filter_values: HashMap<String, Value> = var_length_prop_filters
.iter()
.filter_map(|(k, expr)| match &expr.kind {
ExprKind::Literal(lit) => Some((k.clone(), literal_to_value(lit))),
_ => None,
})
.collect();
Ok(Box::new(ExpandIter {
input: input_iter,
conn,
src_alias: src_alias.clone(),
dst_alias: dst_alias.clone(),
rel_alias: rel_alias.clone(),
edge_types: edge_types.clone(),
direction: *direction,
min_hops: *min_hops,
max_hops: *max_hops,
var_length: *var_length,
var_length_prop_filters: prop_filter_values,
max_traversal_work,
buffer: Vec::new().into_iter(),
}))
}
LogicalOp::Distinct { input } => {
let mut input_iter = build_iter(conn, input, max_traversal_work)?;
let records = collect_all(&mut *input_iter)?;
let mut seen = Vec::new();
let mut deduped = Vec::new();
for rec in records {
if !seen.iter().any(|s: &NamedRecord| s.fields == rec.fields) {
seen.push(rec.clone());
deduped.push(rec);
}
}
Ok(Box::new(VecIter::new(deduped)))
}
LogicalOp::Sort { input, items } => {
let mut input_iter = build_iter(conn, input, max_traversal_work)?;
let mut records = collect_all(&mut *input_iter)?;
records.sort_by(|a, b| {
for item in items {
let va = eval_expr(&item.expr, a, crate::cypher::eval::EvalCx::new(conn))
.unwrap_or(Value::Null);
let vb = eval_expr(&item.expr, b, crate::cypher::eval::EvalCx::new(conn))
.unwrap_or(Value::Null);
let ord = crate::cypher::executor::compare_values_for_sort(&va, &vb);
let ord = if item.descending { ord.reverse() } else { ord };
if ord != std::cmp::Ordering::Equal {
return ord;
}
}
std::cmp::Ordering::Equal
});
Ok(Box::new(VecIter::new(records)))
}
LogicalOp::FullTextLookup {
label,
alias,
property,
op,
term,
remaining_filters,
} => {
use crate::cypher::executor::exec_fulltext_lookup;
let records = exec_fulltext_lookup(
conn,
label,
alias,
property,
*op,
term,
remaining_filters.as_ref(),
&NamedRecord::new(),
)?;
Ok(Box::new(VecIter::new(records)))
}
_ => {
use crate::cypher::executor::exec_pub;
use crate::cypher::executor::ExecContext;
let records = exec_pub(conn, plan, &ExecContext::default())?;
Ok(Box::new(VecIter::new(records)))
}
}
}
pub(crate) fn compare_values_for_sort_pub(a: &Value, b: &Value) -> std::cmp::Ordering {
crate::cypher::executor::compare_values_for_sort(a, b)
}