use std::collections::HashSet;
use super::expansion::{
self, CollectionFilter, VarLenCaps, VarLenCursor, VarLenExpansion, VarLenPattern,
};
use super::types::{BindingRow, ExecutionState, VarLenResume};
use crate::engine::graph::csr::{CsrIndex, Direction, GraphOverlayDelta};
#[derive(Debug, Clone, PartialEq, Eq)]
pub(super) enum NameOrId {
Id(u32),
Name(String),
}
struct NamedBfsSeed {
results: Vec<(NameOrId, String)>,
visited: HashSet<String>,
frontier: Vec<(String, String)>,
start_depth: usize,
}
pub(super) fn expand_named(
csr: &CsrIndex,
source: u32,
pattern: &VarLenPattern<'_>,
caps: VarLenCaps,
overlay: &GraphOverlayDelta,
) -> VarLenExpansion {
let src_name = csr.node_name_raw(source).to_string();
let mut visited: HashSet<String> = HashSet::new();
visited.insert(src_name.clone());
let mut results: Vec<(NameOrId, String)> = Vec::new();
if pattern.min_hops == 0 {
let path = if pattern.want_path {
src_name.clone()
} else {
String::new()
};
results.push((NameOrId::Id(source), path));
}
let seed_path = if pattern.want_path {
src_name.clone()
} else {
String::new()
};
let seed = NamedBfsSeed {
results,
visited,
frontier: vec![(src_name, seed_path)],
start_depth: 1,
};
run_bfs_named(csr, seed, pattern, caps, overlay)
}
pub(super) fn resume_named(
csr: &CsrIndex,
cursor: &VarLenCursor,
pattern: &VarLenPattern<'_>,
caps: VarLenCaps,
overlay: &GraphOverlayDelta,
) -> VarLenExpansion {
let mut visited: HashSet<String> = HashSet::new();
let mut frontier: Vec<(String, String)> = Vec::with_capacity(cursor.frontier.len());
for (name, path) in &cursor.frontier {
let owned =
csr.node_id_raw(name).is_some() || overlay.staged_endpoint_names().any(|n| n == name);
if !owned {
continue;
}
if !visited.insert(name.clone()) {
continue;
}
frontier.push((name.clone(), path.clone()));
}
let seed = NamedBfsSeed {
results: Vec::new(),
visited,
frontier,
start_depth: cursor.depth,
};
run_bfs_named(csr, seed, pattern, caps, overlay)
}
fn run_bfs_named(
csr: &CsrIndex,
seed: NamedBfsSeed,
pattern: &VarLenPattern<'_>,
caps: VarLenCaps,
overlay: &GraphOverlayDelta,
) -> VarLenExpansion {
let NamedBfsSeed {
mut results,
mut visited,
mut frontier,
start_depth,
} = seed;
let mut cursor: Option<VarLenCursor> = None;
let mut boundary: Vec<(String, String, usize)> = Vec::new();
for depth in start_depth..=pattern.max_hops {
if frontier.is_empty() {
break;
}
let mut next_frontier: Vec<(String, String)> = Vec::new();
for (node_name, path) in &frontier {
let src_id = csr.node_id_raw(node_name);
let neighbors = merge_neighbors_named(
csr,
node_name,
src_id,
pattern.label_filter,
pattern.direction,
pattern.collection_filter,
overlay,
);
if neighbors.is_empty() {
boundary.push((node_name.clone(), path.clone(), depth));
continue;
}
for (_label, dst_name) in neighbors {
if !visited.insert(dst_name.clone()) {
continue;
}
let new_path = if pattern.want_path {
format!("{path}->{dst_name}")
} else {
String::new()
};
if depth >= pattern.min_hops {
let bound = match csr.node_id_raw(&dst_name) {
Some(id) => NameOrId::Id(id),
None => NameOrId::Name(dst_name.clone()),
};
results.push((bound, new_path.clone()));
}
if depth < pattern.max_hops {
next_frontier.push((dst_name, new_path));
}
}
}
let cap_hit = results.len() >= caps.max_results || next_frontier.len() >= caps.max_frontier;
if cap_hit {
if depth < pattern.max_hops && !next_frontier.is_empty() {
cursor = Some(VarLenCursor {
frontier: next_frontier,
depth: depth + 1,
});
}
break;
}
frontier = next_frontier;
}
VarLenExpansion {
results: Vec::new(),
named_results: results,
cursor,
boundary,
}
}
pub(super) fn merge_neighbors_named(
csr: &CsrIndex,
src_name: &str,
src_id: Option<u32>,
label_filter: Option<&str>,
direction: Direction,
collection_filter: CollectionFilter,
overlay: &GraphOverlayDelta,
) -> Vec<(String, String)> {
let mut out: Vec<(String, String)> = Vec::new();
let want_out = matches!(direction, Direction::Out | Direction::Both);
let want_in = matches!(direction, Direction::In | Direction::Both);
if let Some(id) = src_id {
if want_out {
out.extend(durable_named(
csr,
id,
label_filter,
Direction::Out,
collection_filter,
src_name,
overlay,
));
}
if want_in {
out.extend(durable_named(
csr,
id,
label_filter,
Direction::In,
collection_filter,
src_name,
overlay,
));
}
}
if want_out {
for (label, dst) in overlay.out_neighbors(src_name, label_filter) {
if !out
.iter()
.any(|(l, n)| l.as_str() == label && n.as_str() == dst)
{
out.push((label.to_string(), dst.to_string()));
}
}
}
if want_in {
for (label, src) in overlay.in_neighbors(src_name, label_filter) {
if !out
.iter()
.any(|(l, n)| l.as_str() == label && n.as_str() == src)
{
out.push((label.to_string(), src.to_string()));
}
}
}
out
}
fn durable_named(
csr: &CsrIndex,
id: u32,
label_filter: Option<&str>,
dir: Direction,
collection_filter: CollectionFilter,
src_name: &str,
overlay: &GraphOverlayDelta,
) -> Vec<(String, String)> {
let mut v = Vec::new();
for (lid, other_id) in
expansion::collect_neighbors(csr, id, label_filter, dir, collection_filter)
{
let label = csr.label_name(lid).to_string();
let other = csr.node_name_raw(other_id).to_string();
let tombstoned = if matches!(dir, Direction::In) {
overlay.is_tombstoned(&other, &label, src_name)
} else {
overlay.is_tombstoned(src_name, &label, &other)
};
if !tombstoned {
v.push((label, other));
}
}
v
}
pub(super) fn record_boundary_resumes(
state: &mut ExecutionState,
triple_idx: usize,
source_row: &BindingRow,
boundary: &[(String, String, usize)],
) {
let Some(pred) = state.is_remote_node else {
return;
};
for (name, path, depth) in boundary {
if pred(name) {
state.record_truncation(VarLenResume {
triple_idx,
source_row: source_row.clone(),
frontier: vec![(name.clone(), path.clone())],
depth: *depth,
});
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn pattern(label: &str, min: usize, max: usize) -> VarLenPattern<'_> {
VarLenPattern {
label_filter: Some(label),
direction: Direction::Out,
min_hops: min,
max_hops: max,
want_path: false,
collection_filter: CollectionFilter::Unscoped,
}
}
fn name_set(exp: &VarLenExpansion, csr: &CsrIndex) -> std::collections::HashSet<String> {
exp.named_results
.iter()
.map(|(b, _)| match b {
NameOrId::Id(id) => csr.node_name_raw(*id).to_string(),
NameOrId::Name(n) => n.clone(),
})
.collect()
}
#[test]
fn staged_edge_traversed_by_name() {
let mut csr = CsrIndex::new();
csr.add_edge("a", "R", "b").unwrap();
let mut ov = GraphOverlayDelta::new();
ov.stage_edge("b", "R", "c");
let src = csr.node_id_raw("a").unwrap();
let exp = expand_named(&csr, src, &pattern("R", 1, 3), VarLenCaps::default(), &ov);
let names = name_set(&exp, &csr);
assert!(names.contains("b"), "durable hop a->b must be reached");
assert!(
names.contains("c"),
"staged hop b->c must be reached through name-keyed BFS; got {names:?}"
);
assert!(
exp.named_results
.iter()
.any(|(b, _)| *b == NameOrId::Name("c".to_string())),
"staged-only node c must be a Name variant"
);
}
#[test]
fn named_resume_union_equals_uncapped() {
let mut csr = CsrIndex::new();
for i in 0..3 {
csr.add_edge(&format!("n{i}"), "R", &format!("n{}", i + 1))
.unwrap();
}
let mut ov = GraphOverlayDelta::new();
ov.stage_edge("n3", "R", "n4");
ov.stage_edge("n4", "R", "n5");
let src = csr.node_id_raw("n0").unwrap();
let pat = pattern("R", 1, 6);
let uncapped = expand_named(&csr, src, &pat, VarLenCaps::default(), &ov);
assert!(uncapped.cursor.is_none());
let full = name_set(&uncapped, &csr);
let caps = VarLenCaps {
max_results: 2,
max_frontier: usize::MAX,
};
let first = expand_named(&csr, src, &pat, caps, &ov);
let mut union = name_set(&first, &csr);
let mut next = first.cursor;
while let Some(c) = next {
let resumed = resume_named(&csr, &c, &pat, caps, &ov);
union.extend(name_set(&resumed, &csr));
next = resumed.cursor;
}
assert_eq!(
union, full,
"first-round ∪ resumed must equal uncapped-with-overlay set"
);
assert!(
full.contains("n5"),
"staged tail must be reachable uncapped"
);
}
#[test]
fn named_resume_foreign_core_skips_unowned() {
let csr = CsrIndex::new();
let mut ov = GraphOverlayDelta::new();
ov.stage_edge("local_only", "R", "x");
let cursor = VarLenCursor {
frontier: vec![
("foreign_a".to_string(), "s->foreign_a".to_string()),
("foreign_b".to_string(), "s->foreign_b".to_string()),
],
depth: 2,
};
let resumed = resume_named(
&csr,
&cursor,
&pattern("R", 1, 6),
VarLenCaps::default(),
&ov,
);
assert!(
resumed.named_results.is_empty(),
"unowned frontier names must yield no results; got {:?}",
resumed.named_results
);
assert!(resumed.boundary.is_empty());
}
#[test]
fn boundary_node_captured_and_recorded() {
let mut csr = CsrIndex::new();
csr.add_edge("a", "R", "b").unwrap();
let mut ov = GraphOverlayDelta::new();
ov.stage_edge("unrelated_src", "R", "unrelated_dst");
let src = csr.node_id_raw("a").unwrap();
let exp = expand_named(&csr, src, &pattern("R", 1, 4), VarLenCaps::default(), &ov);
assert!(
exp.boundary.iter().any(|(n, _, _)| n == "b"),
"zero-out-degree node b must be captured in boundary; got {:?}",
exp.boundary
);
let pred = |n: &str| n == "b";
let mut state = ExecutionState::new(Some(&pred), VarLenCaps::default());
let mut source_row = BindingRow::new();
source_row.insert("a".to_string(), "a".to_string());
record_boundary_resumes(&mut state, 0, &source_row, &exp.boundary);
assert!(
state.truncated(),
"remote boundary node must record a resume"
);
}
#[test]
fn boundary_not_captured_with_staged_out_edge() {
let mut csr = CsrIndex::new();
csr.add_edge("a", "R", "b").unwrap();
let mut ov = GraphOverlayDelta::new();
ov.stage_edge("b", "R", "c");
let src = csr.node_id_raw("a").unwrap();
let exp = expand_named(&csr, src, &pattern("R", 1, 4), VarLenCaps::default(), &ov);
assert!(
!exp.boundary.iter().any(|(n, _, _)| n == "b"),
"b has a staged out-edge and must not be a boundary node; got {:?}",
exp.boundary
);
assert!(exp.boundary.iter().any(|(n, _, _)| n == "c"));
}
#[test]
fn two_anchor_named_expansion() {
let mut csr = CsrIndex::new();
csr.add_edge("a", "R", "b").unwrap();
csr.add_edge("p", "R", "q").unwrap();
let mut ov = GraphOverlayDelta::new();
ov.stage_edge("b", "R", "c");
ov.stage_edge("q", "R", "r");
let a = csr.node_id_raw("a").unwrap();
let p = csr.node_id_raw("p").unwrap();
let pat = pattern("R", 1, 3);
let from_a = name_set(
&expand_named(&csr, a, &pat, VarLenCaps::default(), &ov),
&csr,
);
let from_p = name_set(
&expand_named(&csr, p, &pat, VarLenCaps::default(), &ov),
&csr,
);
assert_eq!(
from_a,
["b", "c"]
.into_iter()
.map(String::from)
.collect::<HashSet<String>>(),
"anchor a reaches only its own tail"
);
assert_eq!(
from_p,
["q", "r"]
.into_iter()
.map(String::from)
.collect::<HashSet<String>>(),
"anchor p reaches only its own tail"
);
}
}