1use std::collections::VecDeque;
2
3use rayon;
4use rustc_hash::{FxHashMap, FxHashSet};
5
6use crate::config::limits::PPR;
7use crate::graph::{CsrGraph, Graph};
8use crate::types::FragmentId;
9
10fn init_seed_residuals(
11 csr: &CsrGraph,
12 seeds: &FxHashSet<FragmentId>,
13 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
14) -> Vec<f64> {
15 let n = csr.n;
16 let mut residual = vec![0.0f64; n];
17
18 let valid_seeds: Vec<&FragmentId> = seeds
19 .iter()
20 .filter(|s| csr.node_to_idx.contains_key(*s))
21 .collect();
22
23 if valid_seeds.is_empty() {
24 return residual;
25 }
26
27 if let Some(sw) = seed_weights {
28 let total: f64 = valid_seeds
29 .iter()
30 .map(|s| sw.get(*s).copied().unwrap_or(PPR.default_seed_epsilon))
31 .sum();
32 if total <= 0.0 {
33 return residual;
34 }
35 for s in &valid_seeds {
36 let idx = csr.node_to_idx[*s] as usize;
37 residual[idx] = sw.get(*s).copied().unwrap_or(PPR.default_seed_epsilon) / total;
38 }
39 } else {
40 let weight = 1.0 / valid_seeds.len() as f64;
41 for s in &valid_seeds {
42 let idx = csr.node_to_idx[*s] as usize;
43 residual[idx] = weight;
44 }
45 }
46
47 residual
48}
49
50struct PprPushResult {
57 estimate: Vec<f64>,
58 pushes: usize,
59 truncated: bool,
60}
61
62fn ppr_push_csr(
63 csr: &CsrGraph,
64 seeds: &FxHashSet<FragmentId>,
65 alpha: f64,
66 tol: f64,
67 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
68) -> PprPushResult {
69 let n = csr.n;
70 if n == 0 {
71 return PprPushResult {
72 estimate: Vec::new(),
73 pushes: 0,
74 truncated: false,
75 };
76 }
77
78 let restart = 1.0 - alpha;
79 let mut residual = init_seed_residuals(csr, seeds, seed_weights);
80 let mut estimate = vec![0.0f64; n];
81 let mut in_queue = vec![false; n];
82
83 let mut queue: VecDeque<u32> = VecDeque::new();
84 for i in 0..n {
85 if residual[i] >= tol {
86 queue.push_back(i as u32);
87 in_queue[i] = true;
88 }
89 }
90
91 let max_pushes = (n * PPR.push_scale_factor).min(PPR.max_pushes_cap);
92 let mut pushes: usize = 0;
93 let mut truncated = false;
94
95 while let Some(u) = queue.pop_front() {
96 if pushes >= max_pushes {
97 truncated = true;
98 break;
99 }
100 let ui = u as usize;
101 in_queue[ui] = false;
102
103 let r_u = residual[ui];
104 if r_u < tol {
105 continue;
106 }
107
108 estimate[ui] += restart * r_u;
109 residual[ui] = 0.0;
110
111 let total_w = csr.out_weight_sum[ui];
112 if total_w <= 0.0 {
113 pushes += 1;
114 continue;
115 }
116
117 let propagate = alpha * r_u;
118 let start = csr.indptr[ui] as usize;
119 let end = csr.indptr[ui + 1] as usize;
120
121 for k in start..end {
122 let v = csr.indices[k] as usize;
123 let w = csr.weights[k];
124 let delta = propagate * (w / total_w);
125 residual[v] += delta;
126 if !in_queue[v] && residual[v] >= tol {
127 queue.push_back(v as u32);
128 in_queue[v] = true;
129 }
130 }
131
132 pushes += 1;
133 }
134
135 PprPushResult {
136 estimate,
137 pushes,
138 truncated,
139 }
140}
141
142pub struct PprResult {
151 pub scores: FxHashMap<FragmentId, f64>,
152 pub truncated: bool,
153 pub forward_pushes: usize,
154 pub backward_pushes: usize,
155}
156
157pub fn personalized_pagerank(
158 graph: &mut Graph,
159 seeds: &FxHashSet<FragmentId>,
160 alpha: f64,
161 tol: f64,
162 forward_blend: f64,
163 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
164) -> PprResult {
165 if graph.node_count() == 0 || seeds.is_empty() {
166 return PprResult {
167 scores: FxHashMap::default(),
168 truncated: false,
169 forward_pushes: 0,
170 backward_pushes: 0,
171 };
172 }
173
174 let (fwd_csr, rev_csr) = graph.to_csr();
175
176 let (forward, backward) = rayon::join(
177 || ppr_push_csr(fwd_csr, seeds, alpha, tol, seed_weights),
178 || ppr_push_csr(rev_csr, seeds, alpha, tol, seed_weights),
179 );
180
181 let n = fwd_csr.n;
182 let mut combined = vec![0.0f64; n];
183 for i in 0..n {
184 combined[i] =
185 forward_blend * forward.estimate[i] + (1.0 - forward_blend) * backward.estimate[i];
186 }
187
188 let total: f64 = combined.iter().sum();
189 if total > 0.0 {
190 for v in &mut combined {
191 *v /= total;
192 }
193 }
194
195 let idx_to_node = &fwd_csr.idx_to_node;
196 let mut scores: FxHashMap<FragmentId, f64> = FxHashMap::default();
197 for i in 0..n {
198 if combined[i] > 0.0 {
199 scores.insert(idx_to_node[i].clone(), combined[i]);
200 }
201 }
202
203 PprResult {
204 scores,
205 truncated: forward.truncated || backward.truncated,
206 forward_pushes: forward.pushes,
207 backward_pushes: backward.pushes,
208 }
209}
210
211#[cfg(test)]
212mod tests {
213 use super::*;
214 use std::sync::Arc;
215
216 fn fid(path: &str, start: u32, end: u32) -> FragmentId {
217 FragmentId::new(Arc::from(path), start, end)
218 }
219
220 #[test]
221 fn ppr_empty_graph() {
222 let mut g = Graph::new();
223 let seeds = FxHashSet::default();
224 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-4, 0.4, None).scores;
225 assert!(result.is_empty());
226 }
227
228 #[test]
229 fn ppr_single_node() {
230 let mut g = Graph::new();
231 let a = fid("a.rs", 1, 10);
232 g.add_node(a.clone());
233
234 let mut seeds = FxHashSet::default();
235 seeds.insert(a.clone());
236 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-4, 0.4, None).scores;
237 assert!((result[&a] - 1.0).abs() < 1e-6);
238 }
239
240 #[test]
241 fn ppr_chain_scores_decrease() {
242 let mut g = Graph::new();
243 let a = fid("a.rs", 1, 10);
244 let b = fid("b.rs", 1, 10);
245 let c = fid("c.rs", 1, 10);
246 g.add_node(a.clone());
247 g.add_node(b.clone());
248 g.add_node(c.clone());
249 g.add_edge(a.clone(), b.clone(), 1.0);
250 g.add_edge(b.clone(), c.clone(), 1.0);
251
252 let mut seeds = FxHashSet::default();
253 seeds.insert(a.clone());
254 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, None).scores;
255
256 assert!(result[&a] > result[&b]);
257 assert!(result[&b] > result[&c]);
258 }
259
260 #[test]
261 fn ppr_normalizes_to_one() {
262 let mut g = Graph::new();
263 let a = fid("a.rs", 1, 10);
264 let b = fid("b.rs", 1, 10);
265 let c = fid("c.rs", 1, 10);
266 g.add_node(a.clone());
267 g.add_node(b.clone());
268 g.add_node(c.clone());
269 g.add_edge(a.clone(), b.clone(), 1.0);
270 g.add_edge(b.clone(), c.clone(), 1.0);
271 g.add_edge(c.clone(), a.clone(), 0.5);
272
273 let mut seeds = FxHashSet::default();
274 seeds.insert(a.clone());
275 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, None).scores;
276
277 let total: f64 = result.values().sum();
278 assert!((total - 1.0).abs() < 1e-6);
279 }
280
281 #[test]
282 fn ppr_with_seed_weights() {
283 let mut g = Graph::new();
284 let a = fid("a.rs", 1, 10);
285 let b = fid("b.rs", 1, 10);
286 g.add_node(a.clone());
287 g.add_node(b.clone());
288 g.add_edge(a.clone(), b.clone(), 1.0);
289
290 let mut seeds = FxHashSet::default();
291 seeds.insert(a.clone());
292 seeds.insert(b.clone());
293
294 let mut sw = FxHashMap::default();
295 sw.insert(a.clone(), 0.9);
296 sw.insert(b.clone(), 0.1);
297
298 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, Some(&sw)).scores;
299 assert!(result[&a] > result[&b]);
300 }
301
302 fn build_star_graph() -> (Graph, FragmentId) {
303 let mut g = Graph::new();
304 let center = fid("center.rs", 1, 10);
305 g.add_node(center.clone());
306 for i in 0..5 {
307 let leaf = fid(&format!("leaf_{i}.rs"), 1, 10);
308 g.add_node(leaf.clone());
309 g.add_edge(center.clone(), leaf.clone(), 1.0);
310 g.add_edge(leaf.clone(), center.clone(), 1.0);
311 }
312 (g, center)
313 }
314
315 #[test]
316 fn ppr_is_deterministic_across_calls() {
317 let (mut g1, center) = build_star_graph();
318 let (mut g2, _) = build_star_graph();
319 let seeds: FxHashSet<FragmentId> = std::iter::once(center).collect();
320
321 let r1 = personalized_pagerank(&mut g1, &seeds, 0.6, 1e-6, 0.4, None).scores;
322 let r2 = personalized_pagerank(&mut g2, &seeds, 0.6, 1e-6, 0.4, None).scores;
323
324 assert_eq!(r1.len(), r2.len());
325 for (id, v1) in &r1 {
326 let v2 = r2.get(id).copied().unwrap_or(f64::NAN);
327 assert!((v1 - v2).abs() < 1e-12, "PPR drift at {id}: {v1} vs {v2}");
328 }
329 }
330
331 #[test]
332 fn ppr_converges_under_tighter_tolerance() {
333 let (mut g_loose, center) = build_star_graph();
334 let (mut g_tight, _) = build_star_graph();
335 let seeds: FxHashSet<FragmentId> = std::iter::once(center).collect();
336
337 let loose = personalized_pagerank(&mut g_loose, &seeds, 0.6, 1e-2, 0.4, None).scores;
338 let tight = personalized_pagerank(&mut g_tight, &seeds, 0.6, 1e-6, 0.4, None).scores;
339
340 let max_diff = loose
341 .iter()
342 .map(|(id, v)| (v - tight.get(id).copied().unwrap_or(0.0)).abs())
343 .fold(0.0f64, f64::max);
344 assert!(
345 max_diff < 1e-2,
346 "PPR did not converge: max diff between tol=1e-2 and tol=1e-6 is {max_diff}"
347 );
348 }
349
350 #[test]
351 fn ppr_symmetric_star_assigns_equal_mass_to_leaves() {
352 let (mut g, center) = build_star_graph();
353 let seeds: FxHashSet<FragmentId> = std::iter::once(center.clone()).collect();
354 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None).scores;
355
356 let leaf_scores: Vec<f64> = (0..5)
357 .map(|i| result[&fid(&format!("leaf_{i}.rs"), 1, 10)])
358 .collect();
359 let max_leaf = leaf_scores.iter().cloned().fold(0.0f64, f64::max);
360 let min_leaf = leaf_scores.iter().cloned().fold(f64::INFINITY, f64::min);
361 assert!(
362 (max_leaf - min_leaf) < 1e-6,
363 "Symmetric star should give equal leaf mass; got spread {} (leaves: {leaf_scores:?})",
364 max_leaf - min_leaf
365 );
366 assert!(result[¢er] > max_leaf, "Center must dominate leaves");
367 }
368
369 #[test]
379 fn claim_7_hub_suppression_reduces_hub_mass_without_removal() {
380 use crate::graph::{EdgeCategory, build_graph};
381 use crate::types::{Fragment, FragmentKind};
382
383 let n_leaves = 20usize;
384 let hub = fid("hub.rs", 1, 10);
385 let leaves: Vec<FragmentId> = (0..n_leaves)
386 .map(|i| fid(&format!("leaf_{i}.rs"), 1, 10))
387 .collect();
388
389 let mut naive = Graph::new();
390 naive.add_node(hub.clone());
391 for leaf in &leaves {
392 naive.add_node(leaf.clone());
393 naive.add_edge(leaf.clone(), hub.clone(), 1.0);
394 }
395 for i in 0..n_leaves - 1 {
396 naive.add_edge(leaves[i].clone(), leaves[i + 1].clone(), 1.0);
397 }
398
399 let fragments: Vec<Fragment> = std::iter::once(hub.clone())
400 .chain(leaves.iter().cloned())
401 .map(|id| Fragment {
402 id,
403 kind: FragmentKind::Function,
404 content: Arc::from(""),
405 identifiers: FxHashSet::default(),
406 token_count: 100,
407 symbol_name: None,
408 })
409 .collect();
410 let mut edges: FxHashMap<(FragmentId, FragmentId), f64> = FxHashMap::default();
411 let mut categories: FxHashMap<(FragmentId, FragmentId), EdgeCategory> =
412 FxHashMap::default();
413 for leaf in &leaves {
414 edges.insert((leaf.clone(), hub.clone()), 1.0);
415 categories.insert((leaf.clone(), hub.clone()), EdgeCategory::Generic);
416 }
417 for i in 0..n_leaves - 1 {
418 edges.insert((leaves[i].clone(), leaves[i + 1].clone()), 1.0);
419 categories.insert(
420 (leaves[i].clone(), leaves[i + 1].clone()),
421 EdgeCategory::Generic,
422 );
423 }
424 let mut suppressed = build_graph(&fragments, edges, categories);
425
426 let seeds: FxHashSet<FragmentId> = leaves.iter().take(3).cloned().collect();
427 let alpha = 0.6;
428 let tol = 1e-8;
429 let blend = 1.0;
430
431 let r_naive = personalized_pagerank(&mut naive, &seeds, alpha, tol, blend, None).scores;
432 let r_suppressed =
433 personalized_pagerank(&mut suppressed, &seeds, alpha, tol, blend, None).scores;
434
435 let hub_naive = r_naive.get(&hub).copied().unwrap_or(0.0);
436 let hub_suppressed = r_suppressed.get(&hub).copied().unwrap_or(0.0);
437
438 assert!(
439 hub_suppressed < hub_naive,
440 "Hub suppression did not reduce hub mass: naive={hub_naive}, suppressed={hub_suppressed}"
441 );
442 let reduction_ratio = hub_naive / hub_suppressed.max(1e-12);
443 assert!(
444 reduction_ratio >= 1.5,
445 "Hub suppression effect too small: only {reduction_ratio:.2}× reduction (want ≥1.5×)"
446 );
447
448 assert!(
449 hub_suppressed > 0.0,
450 "Hub mass should be reduced, not removed; got {hub_suppressed}"
451 );
452
453 let mut leaves_present_after = 0;
454 for leaf in &leaves {
455 if r_suppressed.get(leaf).copied().unwrap_or(0.0) > 0.0 {
456 leaves_present_after += 1;
457 }
458 }
459 assert!(
460 leaves_present_after >= leaves.len() / 2,
461 "Suppression should preserve most leaves in the result, only {leaves_present_after}/{} survived",
462 leaves.len()
463 );
464 }
465
466 #[test]
470 fn claim_10_ppr_and_ego_rankings_correlate_on_synthetic_graph() {
471 let mut g = Graph::new();
472 let nodes: Vec<FragmentId> = (0..30).map(|i| fid(&format!("n_{i}.rs"), 1, 10)).collect();
473 for n in &nodes {
474 g.add_node(n.clone());
475 }
476 let mut rng = 0xC0FFEE_u64;
477 let xorshift = |state: &mut u64| -> u64 {
478 *state ^= *state << 13;
479 *state ^= *state >> 7;
480 *state ^= *state << 17;
481 *state
482 };
483 for i in 0..nodes.len() {
484 for _ in 0..3 {
485 let j = (xorshift(&mut rng) as usize) % nodes.len();
486 if i != j {
487 g.add_edge(nodes[i].clone(), nodes[j].clone(), 1.0);
488 }
489 }
490 }
491 for i in 0..nodes.len() {
492 let next = (i + 1) % nodes.len();
493 g.add_edge(nodes[i].clone(), nodes[next].clone(), 1.0);
494 }
495
496 let seeds: FxHashSet<FragmentId> = nodes.iter().take(2).cloned().collect();
497 let ppr_scores = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None).scores;
498 let ego_scores = g.ego_graph(&seeds, 3);
499
500 let common: Vec<FragmentId> = nodes
501 .iter()
502 .filter(|n| ppr_scores.contains_key(n) && ego_scores.contains_key(n))
503 .cloned()
504 .collect();
505 assert!(
506 common.len() >= 10,
507 "Need ≥10 common ranked nodes, got {}",
508 common.len()
509 );
510
511 let ppr_v: Vec<f64> = common.iter().map(|n| ppr_scores[n]).collect();
512 let ego_v: Vec<f64> = common.iter().map(|n| ego_scores[n]).collect();
513
514 let rho = spearman_correlation(&ppr_v, &ego_v);
515 assert!(
516 rho > 0.3,
517 "PPR/EGO Spearman correlation too low: ρ={rho:.3} (paper hypothesizes high correlation, want > 0.3)"
518 );
519 }
520
521 fn rank(values: &[f64]) -> Vec<f64> {
522 let n = values.len();
523 let mut indexed: Vec<(usize, f64)> = values.iter().copied().enumerate().collect();
524 indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
525 let mut ranks = vec![0.0_f64; n];
526 let mut i = 0;
527 while i < n {
528 let mut j = i;
529 while j + 1 < n && indexed[j + 1].1 == indexed[i].1 {
530 j += 1;
531 }
532 let avg_rank = ((i + j) as f64) / 2.0 + 1.0;
533 for k in i..=j {
534 ranks[indexed[k].0] = avg_rank;
535 }
536 i = j + 1;
537 }
538 ranks
539 }
540
541 fn spearman_correlation(x: &[f64], y: &[f64]) -> f64 {
542 assert_eq!(x.len(), y.len());
543 let rx = rank(x);
544 let ry = rank(y);
545 let n = x.len() as f64;
546 let mean_x: f64 = rx.iter().sum::<f64>() / n;
547 let mean_y: f64 = ry.iter().sum::<f64>() / n;
548 let mut cov = 0.0;
549 let mut var_x = 0.0;
550 let mut var_y = 0.0;
551 for i in 0..rx.len() {
552 let dx = rx[i] - mean_x;
553 let dy = ry[i] - mean_y;
554 cov += dx * dy;
555 var_x += dx * dx;
556 var_y += dy * dy;
557 }
558 cov / (var_x.sqrt() * var_y.sqrt()).max(1e-12)
559 }
560
561 #[test]
571 fn ppr_matches_closed_form_on_symmetric_star() {
572 for &(alpha, n_leaves) in &[(0.6, 5usize), (0.5, 4usize), (0.85, 7usize)] {
573 let mut g = Graph::new();
574 let center = fid("center.rs", 1, 10);
575 g.add_node(center.clone());
576 for i in 0..n_leaves {
577 let leaf = fid(&format!("leaf_{i}.rs"), 1, 10);
578 g.add_node(leaf.clone());
579 g.add_edge(center.clone(), leaf.clone(), 1.0);
580 g.add_edge(leaf.clone(), center.clone(), 1.0);
581 }
582 let seeds: FxHashSet<FragmentId> = std::iter::once(center.clone()).collect();
583 let result = personalized_pagerank(&mut g, &seeds, alpha, 1e-10, 0.5, None).scores;
584
585 let expected_center = 1.0 / (1.0 + alpha);
586 let expected_leaf = alpha / (n_leaves as f64 * (1.0 + alpha));
587
588 let actual_center = result[¢er];
589 let center_err = (actual_center - expected_center).abs();
590 assert!(
591 center_err < 1e-3,
592 "α={alpha}, N={n_leaves}: center mass drift |actual - closed-form| = {center_err}; \
593 expected={expected_center}, got={actual_center}"
594 );
595
596 for i in 0..n_leaves {
597 let leaf_id = fid(&format!("leaf_{i}.rs"), 1, 10);
598 let actual_leaf = result[&leaf_id];
599 let leaf_err = (actual_leaf - expected_leaf).abs();
600 assert!(
601 leaf_err < 1e-3,
602 "α={alpha}, N={n_leaves}, leaf_{i}: drift |actual - closed-form| = {leaf_err}; \
603 expected={expected_leaf}, got={actual_leaf}"
604 );
605 }
606 }
607 }
608}