1use std::collections::{BTreeSet, HashMap, VecDeque};
15
16use crate::file::Rete;
17use crate::terms::NodeId;
18
19pub fn build_adjacency(rete: &Rete, pred: &str) -> HashMap<NodeId, Vec<NodeId>> {
24 let mut adj: HashMap<NodeId, Vec<NodeId>> = HashMap::new();
25 for (s, o) in rete.predicate_pairs(pred) {
26 adj.entry(s).or_default().push(o);
27 }
28 adj
29}
30
31pub fn reach_one(adj: &HashMap<NodeId, Vec<NodeId>>, seed: NodeId) -> BTreeSet<NodeId> {
35 let mut visited = BTreeSet::new();
36 let mut queue = VecDeque::new();
37 if let Some(succ) = adj.get(&seed) {
38 for &n in succ {
39 if visited.insert(n) {
40 queue.push_back(n);
41 }
42 }
43 }
44 while let Some(n) = queue.pop_front() {
45 if let Some(succ) = adj.get(&n) {
46 for &m in succ {
47 if visited.insert(m) {
48 queue.push_back(m);
49 }
50 }
51 }
52 }
53 visited
54}
55
56pub fn batch_reach_serial(
60 adj: &HashMap<NodeId, Vec<NodeId>>,
61 seeds: &[NodeId],
62) -> Vec<BTreeSet<NodeId>> {
63 seeds.iter().map(|&s| reach_one(adj, s)).collect()
64}
65
66#[cfg(test)]
67mod tests {
68 use super::*;
69 use crate::dictionary::DictionaryBuilder;
70 use crate::file::{build_pyramid_meta, write_dataset, Rete, DEFAULT_TILE_BUDGET};
71 use crate::index::GraphIndexBuilder;
72
73 fn fixture() -> Rete {
75 let edges = [
76 ("A", "B"),
77 ("B", "C"),
78 ("A", "C"),
79 ("C", "A"),
80 ("D", "E"),
81 ("E", "F"),
82 ("D", "F"),
83 ("F", "D"),
84 ("C", "D"),
85 ];
86 let mut db = DictionaryBuilder::new();
87 for (s, o) in edges {
88 db.observe(s, "knows", o);
89 }
90 let dict = db.build();
91 let mut triples: Vec<(u32, u32, u32)> = edges
92 .iter()
93 .map(|(s, o)| dict.encode(s, "knows", o).unwrap())
94 .collect();
95 triples.sort_unstable();
96 triples.dedup();
97
98 let mut def = GraphIndexBuilder::new();
99 for &t in &triples {
100 def.push(t);
101 }
102 let (meta, levels) = build_pyramid_meta(&dict, &triples, DEFAULT_TILE_BUDGET);
103 let bytes = write_dataset(&dict, &def.build(), &[], false, &meta, levels);
104 Rete::open(&bytes).unwrap()
105 }
106
107 #[test]
108 fn build_adjacency_and_reach() {
109 let rete = fixture();
110 let adj = build_adjacency(&rete, "knows");
111 let dict = rete.dictionary();
112 let a = dict.node_of_term("A").unwrap();
113 let reached = reach_one(&adj, a);
114 assert_eq!(reached.len(), 6);
117 }
118
119 #[test]
120 fn batch_serial_matches_single() {
121 let rete = fixture();
122 let adj = build_adjacency(&rete, "knows");
123 let seeds: Vec<u32> = adj.keys().copied().collect();
124 let batch = batch_reach_serial(&adj, &seeds);
125 for (i, &s) in seeds.iter().enumerate() {
126 assert_eq!(batch[i], reach_one(&adj, s));
127 }
128 assert!(batch.iter().any(|r| r.len() >= 4));
129 }
130}