1use std::collections::HashMap;
25
26const FIRING_THRESHOLD: f64 = 1e-4;
29
30pub fn spread(
41 seeds: &HashMap<String, f64>,
42 adjacency: &HashMap<String, Vec<(String, f64)>>,
43 decay: f64,
44 iterations: usize,
45) -> HashMap<String, f64> {
46 let decay = decay.clamp(0.0, 0.999);
47 let mut activation: HashMap<String, f64> = seeds.clone();
48 let mut frontier: HashMap<String, f64> = seeds.clone();
49
50 for _ in 0..iterations {
51 let mut next: HashMap<String, f64> = HashMap::new();
52 for (node, &energy) in &frontier {
53 if energy < FIRING_THRESHOLD {
54 continue;
55 }
56 let Some(edges) = adjacency.get(node) else {
57 continue;
58 };
59 let total: f64 = edges.iter().map(|(_, w)| w.max(0.0)).sum();
60 if total <= 0.0 {
61 continue;
62 }
63 for (nbr, w) in edges {
64 let w = w.max(0.0);
65 if w <= 0.0 {
66 continue;
67 }
68 let delta = energy * decay * (w / total);
69 if delta >= FIRING_THRESHOLD {
70 *next.entry(nbr.clone()).or_insert(0.0) += delta;
71 }
72 }
73 }
74 if next.is_empty() {
75 break;
76 }
77 for (node, e) in &next {
78 *activation.entry(node.clone()).or_insert(0.0) += e;
79 }
80 frontier = next;
81 }
82
83 activation
84}
85
86pub fn related_ranked(
90 seeds: &HashMap<String, f64>,
91 adjacency: &HashMap<String, Vec<(String, f64)>>,
92 decay: f64,
93 iterations: usize,
94 top_k: usize,
95) -> Vec<(String, f64)> {
96 let activation = spread(seeds, adjacency, decay, iterations);
97 let mut ranked: Vec<(String, f64)> = activation
98 .into_iter()
99 .filter(|(node, _)| !seeds.contains_key(node))
100 .collect();
101 ranked.sort_by(|a, b| b.1.total_cmp(&a.1));
102 ranked.truncate(top_k);
103 ranked
104}
105
106#[cfg(test)]
107mod tests {
108 use super::*;
109
110 fn edge(adj: &mut HashMap<String, Vec<(String, f64)>>, from: &str, to: &str, w: f64) {
111 adj.entry(from.to_string())
112 .or_default()
113 .push((to.to_string(), w));
114 adj.entry(to.to_string())
115 .or_default()
116 .push((from.to_string(), w));
117 }
118
119 #[test]
120 fn activation_reaches_connected_nodes_only() {
121 let mut adj = HashMap::new();
122 edge(&mut adj, "a", "b", 1.0);
123 edge(&mut adj, "b", "c", 1.0);
124 adj.entry("island".to_string()).or_default();
126
127 let seeds = HashMap::from([("a".to_string(), 1.0)]);
128 let act = spread(&seeds, &adj, 0.7, 5);
129
130 assert!(act.contains_key("b"));
131 assert!(act.contains_key("c"));
132 assert!(!act.contains_key("island"));
133 }
134
135 #[test]
136 fn closer_nodes_get_more_activation() {
137 let mut adj = HashMap::new();
138 edge(&mut adj, "a", "b", 1.0); edge(&mut adj, "b", "c", 1.0); let seeds = HashMap::from([("a".to_string(), 1.0)]);
142 let act = spread(&seeds, &adj, 0.7, 5);
143
144 assert!(act["b"] > act["c"]);
146 }
147
148 #[test]
149 fn stronger_edges_transmit_more() {
150 let mut adj = HashMap::new();
151 edge(&mut adj, "seed", "strong", 9.0);
152 edge(&mut adj, "seed", "weak", 1.0);
153
154 let seeds = HashMap::from([("seed".to_string(), 1.0)]);
155 let act = spread(&seeds, &adj, 0.7, 3);
156 assert!(act["strong"] > act["weak"]);
157 }
158
159 #[test]
160 fn terminates_and_stays_bounded_on_cycles() {
161 let mut adj = HashMap::new();
163 edge(&mut adj, "a", "b", 1.0);
164 edge(&mut adj, "b", "c", 1.0);
165 edge(&mut adj, "c", "a", 1.0);
166
167 let seeds = HashMap::from([("a".to_string(), 1.0)]);
168 let act = spread(&seeds, &adj, 0.9, 1000);
169 let total: f64 = act.values().sum();
171 assert!(total.is_finite());
172 assert!(total < 100.0, "energy must not blow up on cycles: {total}");
173 }
174
175 #[test]
176 fn related_ranked_excludes_seeds() {
177 let mut adj = HashMap::new();
178 edge(&mut adj, "a", "b", 1.0);
179 edge(&mut adj, "a", "c", 1.0);
180
181 let seeds = HashMap::from([("a".to_string(), 1.0)]);
182 let ranked = related_ranked(&seeds, &adj, 0.7, 3, 10);
183 assert!(ranked.iter().all(|(n, _)| n != "a"));
184 assert_eq!(ranked.len(), 2);
185 }
186
187 #[test]
188 fn empty_seeds_yield_empty_result() {
189 let adj = HashMap::new();
190 let seeds = HashMap::new();
191 assert!(spread(&seeds, &adj, 0.7, 5).is_empty());
192 }
193}