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
50fn has_valid_seed_mass(
60 csr: &CsrGraph,
61 seeds: &FxHashSet<FragmentId>,
62 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
63) -> bool {
64 let has_valid_seed = seeds.iter().any(|s| csr.node_to_idx.contains_key(s));
65 if !has_valid_seed {
66 return false;
67 }
68 match seed_weights {
69 Some(sw) => {
70 let total: f64 = seeds
71 .iter()
72 .filter(|s| csr.node_to_idx.contains_key(*s))
73 .map(|s| sw.get(s).copied().unwrap_or(PPR.default_seed_epsilon))
74 .sum();
75 total > 0.0
76 }
77 None => true,
78 }
79}
80
81struct PprPushResult {
88 estimate: Vec<f64>,
89 pushes: usize,
90 truncated: bool,
91}
92
93fn ppr_push_csr(
94 csr: &CsrGraph,
95 seeds: &FxHashSet<FragmentId>,
96 alpha: f64,
97 tol: f64,
98 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
99) -> PprPushResult {
100 let n = csr.n;
101 if n == 0 {
102 return PprPushResult {
103 estimate: Vec::new(),
104 pushes: 0,
105 truncated: false,
106 };
107 }
108
109 let restart = 1.0 - alpha;
110 let mut residual = init_seed_residuals(csr, seeds, seed_weights);
111 let mut estimate = vec![0.0f64; n];
112 let mut in_queue = vec![false; n];
113
114 let mut queue: VecDeque<u32> = VecDeque::new();
115 for i in 0..n {
116 if residual[i] >= tol {
117 queue.push_back(i as u32);
118 in_queue[i] = true;
119 }
120 }
121
122 let max_pushes = (n * PPR.push_scale_factor).min(PPR.max_pushes_cap);
123 let mut pushes: usize = 0;
124 let mut truncated = false;
125
126 while let Some(u) = queue.pop_front() {
127 if pushes >= max_pushes {
128 truncated = true;
129 break;
130 }
131 let ui = u as usize;
132 in_queue[ui] = false;
133
134 let r_u = residual[ui];
135 if r_u < tol {
136 continue;
137 }
138
139 estimate[ui] += restart * r_u;
140 residual[ui] = 0.0;
141
142 let total_w = csr.out_weight_sum[ui];
143 if total_w <= 0.0 {
144 pushes += 1;
145 continue;
146 }
147
148 let propagate = alpha * r_u;
149 let start = csr.indptr[ui] as usize;
150 let end = csr.indptr[ui + 1] as usize;
151
152 for k in start..end {
153 let v = csr.indices[k] as usize;
154 let w = csr.weights[k];
155 let delta = propagate * (w / total_w);
156 residual[v] += delta;
157 if !in_queue[v] && residual[v] >= tol {
158 queue.push_back(v as u32);
159 in_queue[v] = true;
160 }
161 }
162
163 pushes += 1;
164 }
165
166 PprPushResult {
167 estimate,
168 pushes,
169 truncated,
170 }
171}
172
173pub struct PprResult {
182 pub scores: FxHashMap<FragmentId, f64>,
183 pub truncated: bool,
184 pub forward_pushes: usize,
185 pub backward_pushes: usize,
186 pub seeded: bool,
192}
193
194pub fn personalized_pagerank(
195 graph: &mut Graph,
196 seeds: &FxHashSet<FragmentId>,
197 alpha: f64,
198 tol: f64,
199 forward_blend: f64,
200 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
201) -> PprResult {
202 if graph.node_count() == 0 || seeds.is_empty() {
203 return PprResult {
204 scores: FxHashMap::default(),
205 truncated: false,
206 forward_pushes: 0,
207 backward_pushes: 0,
208 seeded: false,
209 };
210 }
211
212 let (fwd_csr, rev_csr) = graph.to_csr();
213 let seeded = has_valid_seed_mass(fwd_csr, seeds, seed_weights);
214
215 let (forward, backward) = rayon::join(
216 || ppr_push_csr(fwd_csr, seeds, alpha, tol, seed_weights),
217 || ppr_push_csr(rev_csr, seeds, alpha, tol, seed_weights),
218 );
219
220 let n = fwd_csr.n;
221 let mut combined = vec![0.0f64; n];
222 for i in 0..n {
223 combined[i] =
224 forward_blend * forward.estimate[i] + (1.0 - forward_blend) * backward.estimate[i];
225 }
226
227 let total: f64 = combined.iter().sum();
228 if total > 0.0 {
229 for v in &mut combined {
230 *v /= total;
231 }
232 }
233
234 let idx_to_node = &fwd_csr.idx_to_node;
235 let mut scores: FxHashMap<FragmentId, f64> = FxHashMap::default();
236 for i in 0..n {
237 if combined[i] > 0.0 {
238 scores.insert(idx_to_node[i].clone(), combined[i]);
239 }
240 }
241
242 PprResult {
243 scores,
244 truncated: forward.truncated || backward.truncated,
245 forward_pushes: forward.pushes,
246 backward_pushes: backward.pushes,
247 seeded,
248 }
249}
250
251#[cfg(test)]
252mod tests {
253 use super::*;
254 use std::sync::Arc;
255
256 fn fid(path: &str, start: u32, end: u32) -> FragmentId {
257 FragmentId::new(Arc::from(path), start, end)
258 }
259
260 #[test]
261 fn ppr_empty_graph() {
262 let mut g = Graph::new();
263 let seeds = FxHashSet::default();
264 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-4, 0.4, None).scores;
265 assert!(result.is_empty());
266 }
267
268 #[test]
269 fn ppr_single_node() {
270 let mut g = Graph::new();
271 let a = fid("a.rs", 1, 10);
272 g.add_node(a.clone());
273
274 let mut seeds = FxHashSet::default();
275 seeds.insert(a.clone());
276 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-4, 0.4, None).scores;
277 assert!((result[&a] - 1.0).abs() < 1e-6);
278 }
279
280 #[test]
281 fn ppr_chain_scores_decrease() {
282 let mut g = Graph::new();
283 let a = fid("a.rs", 1, 10);
284 let b = fid("b.rs", 1, 10);
285 let c = fid("c.rs", 1, 10);
286 g.add_node(a.clone());
287 g.add_node(b.clone());
288 g.add_node(c.clone());
289 g.add_edge(a.clone(), b.clone(), 1.0);
290 g.add_edge(b.clone(), c.clone(), 1.0);
291
292 let mut seeds = FxHashSet::default();
293 seeds.insert(a.clone());
294 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, None).scores;
295
296 assert!(result[&a] > result[&b]);
297 assert!(result[&b] > result[&c]);
298 }
299
300 #[test]
301 fn ppr_normalizes_to_one() {
302 let mut g = Graph::new();
303 let a = fid("a.rs", 1, 10);
304 let b = fid("b.rs", 1, 10);
305 let c = fid("c.rs", 1, 10);
306 g.add_node(a.clone());
307 g.add_node(b.clone());
308 g.add_node(c.clone());
309 g.add_edge(a.clone(), b.clone(), 1.0);
310 g.add_edge(b.clone(), c.clone(), 1.0);
311 g.add_edge(c.clone(), a.clone(), 0.5);
312
313 let mut seeds = FxHashSet::default();
314 seeds.insert(a.clone());
315 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, None).scores;
316
317 let total: f64 = result.values().sum();
318 assert!((total - 1.0).abs() < 1e-6);
319 }
320
321 #[test]
322 fn ppr_with_seed_weights() {
323 let mut g = Graph::new();
324 let a = fid("a.rs", 1, 10);
325 let b = fid("b.rs", 1, 10);
326 g.add_node(a.clone());
327 g.add_node(b.clone());
328 g.add_edge(a.clone(), b.clone(), 1.0);
329
330 let mut seeds = FxHashSet::default();
331 seeds.insert(a.clone());
332 seeds.insert(b.clone());
333
334 let mut sw = FxHashMap::default();
335 sw.insert(a.clone(), 0.9);
336 sw.insert(b.clone(), 0.1);
337
338 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.4, Some(&sw)).scores;
339 assert!(result[&a] > result[&b]);
340 }
341
342 fn build_star_graph() -> (Graph, FragmentId) {
343 let mut g = Graph::new();
344 let center = fid("center.rs", 1, 10);
345 g.add_node(center.clone());
346 for i in 0..5 {
347 let leaf = fid(&format!("leaf_{i}.rs"), 1, 10);
348 g.add_node(leaf.clone());
349 g.add_edge(center.clone(), leaf.clone(), 1.0);
350 g.add_edge(leaf.clone(), center.clone(), 1.0);
351 }
352 (g, center)
353 }
354
355 #[test]
356 fn ppr_is_deterministic_across_calls() {
357 let (mut g1, center) = build_star_graph();
358 let (mut g2, _) = build_star_graph();
359 let seeds: FxHashSet<FragmentId> = std::iter::once(center).collect();
360
361 let r1 = personalized_pagerank(&mut g1, &seeds, 0.6, 1e-6, 0.4, None).scores;
362 let r2 = personalized_pagerank(&mut g2, &seeds, 0.6, 1e-6, 0.4, None).scores;
363
364 assert_eq!(r1.len(), r2.len());
365 for (id, v1) in &r1 {
366 let v2 = r2.get(id).copied().unwrap_or(f64::NAN);
367 assert!((v1 - v2).abs() < 1e-12, "PPR drift at {id}: {v1} vs {v2}");
368 }
369 }
370
371 #[test]
372 fn ppr_converges_under_tighter_tolerance() {
373 let (mut g_loose, center) = build_star_graph();
374 let (mut g_tight, _) = build_star_graph();
375 let seeds: FxHashSet<FragmentId> = std::iter::once(center).collect();
376
377 let loose = personalized_pagerank(&mut g_loose, &seeds, 0.6, 1e-2, 0.4, None).scores;
378 let tight = personalized_pagerank(&mut g_tight, &seeds, 0.6, 1e-6, 0.4, None).scores;
379
380 let max_diff = loose
381 .iter()
382 .map(|(id, v)| (v - tight.get(id).copied().unwrap_or(0.0)).abs())
383 .fold(0.0f64, f64::max);
384 assert!(
385 max_diff < 1e-2,
386 "PPR did not converge: max diff between tol=1e-2 and tol=1e-6 is {max_diff}"
387 );
388 }
389
390 #[test]
391 fn ppr_symmetric_star_assigns_equal_mass_to_leaves() {
392 let (mut g, center) = build_star_graph();
393 let seeds: FxHashSet<FragmentId> = std::iter::once(center.clone()).collect();
394 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None).scores;
395
396 let leaf_scores: Vec<f64> = (0..5)
397 .map(|i| result[&fid(&format!("leaf_{i}.rs"), 1, 10)])
398 .collect();
399 let max_leaf = leaf_scores.iter().cloned().fold(0.0f64, f64::max);
400 let min_leaf = leaf_scores.iter().cloned().fold(f64::INFINITY, f64::min);
401 assert!(
402 (max_leaf - min_leaf) < 1e-6,
403 "Symmetric star should give equal leaf mass; got spread {} (leaves: {leaf_scores:?})",
404 max_leaf - min_leaf
405 );
406 assert!(result[¢er] > max_leaf, "Center must dominate leaves");
407 }
408
409 #[test]
419 fn claim_7_hub_suppression_reduces_hub_mass_without_removal() {
420 use crate::graph::{EdgeCategory, build_graph};
421 use crate::types::{Fragment, FragmentKind};
422
423 let n_leaves = 20usize;
424 let hub = fid("hub.rs", 1, 10);
425 let leaves: Vec<FragmentId> = (0..n_leaves)
426 .map(|i| fid(&format!("leaf_{i}.rs"), 1, 10))
427 .collect();
428
429 let mut naive = Graph::new();
430 naive.add_node(hub.clone());
431 for leaf in &leaves {
432 naive.add_node(leaf.clone());
433 naive.add_edge(leaf.clone(), hub.clone(), 1.0);
434 }
435 for i in 0..n_leaves - 1 {
436 naive.add_edge(leaves[i].clone(), leaves[i + 1].clone(), 1.0);
437 }
438
439 let fragments: Vec<Fragment> = std::iter::once(hub.clone())
440 .chain(leaves.iter().cloned())
441 .map(|id| Fragment {
442 id,
443 kind: FragmentKind::Function,
444 content: Arc::from(""),
445 identifiers: FxHashSet::default(),
446 token_count: 100,
447 symbol_name: None,
448 })
449 .collect();
450 let mut edges: FxHashMap<(FragmentId, FragmentId), f64> = FxHashMap::default();
451 let mut categories: FxHashMap<(FragmentId, FragmentId), EdgeCategory> =
452 FxHashMap::default();
453 for leaf in &leaves {
454 edges.insert((leaf.clone(), hub.clone()), 1.0);
455 categories.insert((leaf.clone(), hub.clone()), EdgeCategory::Generic);
456 }
457 for i in 0..n_leaves - 1 {
458 edges.insert((leaves[i].clone(), leaves[i + 1].clone()), 1.0);
459 categories.insert(
460 (leaves[i].clone(), leaves[i + 1].clone()),
461 EdgeCategory::Generic,
462 );
463 }
464 let mut suppressed = build_graph(&fragments, edges, categories);
465
466 let seeds: FxHashSet<FragmentId> = leaves.iter().take(3).cloned().collect();
467 let alpha = 0.6;
468 let tol = 1e-8;
469 let blend = 1.0;
470
471 let r_naive = personalized_pagerank(&mut naive, &seeds, alpha, tol, blend, None).scores;
472 let r_suppressed =
473 personalized_pagerank(&mut suppressed, &seeds, alpha, tol, blend, None).scores;
474
475 let hub_naive = r_naive.get(&hub).copied().unwrap_or(0.0);
476 let hub_suppressed = r_suppressed.get(&hub).copied().unwrap_or(0.0);
477
478 assert!(
479 hub_suppressed < hub_naive,
480 "Hub suppression did not reduce hub mass: naive={hub_naive}, suppressed={hub_suppressed}"
481 );
482 let reduction_ratio = hub_naive / hub_suppressed.max(1e-12);
483 assert!(
484 reduction_ratio >= 1.5,
485 "Hub suppression effect too small: only {reduction_ratio:.2}× reduction (want ≥1.5×)"
486 );
487
488 assert!(
489 hub_suppressed > 0.0,
490 "Hub mass should be reduced, not removed; got {hub_suppressed}"
491 );
492
493 let mut leaves_present_after = 0;
494 for leaf in &leaves {
495 if r_suppressed.get(leaf).copied().unwrap_or(0.0) > 0.0 {
496 leaves_present_after += 1;
497 }
498 }
499 assert!(
500 leaves_present_after >= leaves.len() / 2,
501 "Suppression should preserve most leaves in the result, only {leaves_present_after}/{} survived",
502 leaves.len()
503 );
504 }
505
506 #[test]
510 fn claim_10_ppr_and_ego_rankings_correlate_on_synthetic_graph() {
511 let mut g = Graph::new();
512 let nodes: Vec<FragmentId> = (0..30).map(|i| fid(&format!("n_{i}.rs"), 1, 10)).collect();
513 for n in &nodes {
514 g.add_node(n.clone());
515 }
516 let mut rng = 0xC0FFEE_u64;
517 let xorshift = |state: &mut u64| -> u64 {
518 *state ^= *state << 13;
519 *state ^= *state >> 7;
520 *state ^= *state << 17;
521 *state
522 };
523 for i in 0..nodes.len() {
524 for _ in 0..3 {
525 let j = (xorshift(&mut rng) as usize) % nodes.len();
526 if i != j {
527 g.add_edge(nodes[i].clone(), nodes[j].clone(), 1.0);
528 }
529 }
530 }
531 for i in 0..nodes.len() {
532 let next = (i + 1) % nodes.len();
533 g.add_edge(nodes[i].clone(), nodes[next].clone(), 1.0);
534 }
535
536 let seeds: FxHashSet<FragmentId> = nodes.iter().take(2).cloned().collect();
537 let ppr_scores = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None).scores;
538 let ego_scores = g.ego_graph(&seeds, 3);
539
540 let common: Vec<FragmentId> = nodes
541 .iter()
542 .filter(|n| ppr_scores.contains_key(n) && ego_scores.contains_key(n))
543 .cloned()
544 .collect();
545 assert!(
546 common.len() >= 10,
547 "Need ≥10 common ranked nodes, got {}",
548 common.len()
549 );
550
551 let ppr_v: Vec<f64> = common.iter().map(|n| ppr_scores[n]).collect();
552 let ego_v: Vec<f64> = common.iter().map(|n| ego_scores[n]).collect();
553
554 let rho = spearman_correlation(&ppr_v, &ego_v);
555 assert!(
556 rho > 0.3,
557 "PPR/EGO Spearman correlation too low: ρ={rho:.3} (paper hypothesizes high correlation, want > 0.3)"
558 );
559 }
560
561 fn rank(values: &[f64]) -> Vec<f64> {
562 let n = values.len();
563 let mut indexed: Vec<(usize, f64)> = values.iter().copied().enumerate().collect();
564 indexed.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
565 let mut ranks = vec![0.0_f64; n];
566 let mut i = 0;
567 while i < n {
568 let mut j = i;
569 while j + 1 < n && indexed[j + 1].1 == indexed[i].1 {
570 j += 1;
571 }
572 let avg_rank = ((i + j) as f64) / 2.0 + 1.0;
573 for k in i..=j {
574 ranks[indexed[k].0] = avg_rank;
575 }
576 i = j + 1;
577 }
578 ranks
579 }
580
581 fn spearman_correlation(x: &[f64], y: &[f64]) -> f64 {
582 assert_eq!(x.len(), y.len());
583 let rx = rank(x);
584 let ry = rank(y);
585 let n = x.len() as f64;
586 let mean_x: f64 = rx.iter().sum::<f64>() / n;
587 let mean_y: f64 = ry.iter().sum::<f64>() / n;
588 let mut cov = 0.0;
589 let mut var_x = 0.0;
590 let mut var_y = 0.0;
591 for i in 0..rx.len() {
592 let dx = rx[i] - mean_x;
593 let dy = ry[i] - mean_y;
594 cov += dx * dy;
595 var_x += dx * dx;
596 var_y += dy * dy;
597 }
598 cov / (var_x.sqrt() * var_y.sqrt()).max(1e-12)
599 }
600
601 #[test]
611 fn ppr_matches_closed_form_on_symmetric_star() {
612 for &(alpha, n_leaves) in &[(0.6, 5usize), (0.5, 4usize), (0.85, 7usize)] {
613 let mut g = Graph::new();
614 let center = fid("center.rs", 1, 10);
615 g.add_node(center.clone());
616 for i in 0..n_leaves {
617 let leaf = fid(&format!("leaf_{i}.rs"), 1, 10);
618 g.add_node(leaf.clone());
619 g.add_edge(center.clone(), leaf.clone(), 1.0);
620 g.add_edge(leaf.clone(), center.clone(), 1.0);
621 }
622 let seeds: FxHashSet<FragmentId> = std::iter::once(center.clone()).collect();
623 let result = personalized_pagerank(&mut g, &seeds, alpha, 1e-10, 0.5, None).scores;
624
625 let expected_center = 1.0 / (1.0 + alpha);
626 let expected_leaf = alpha / (n_leaves as f64 * (1.0 + alpha));
627
628 let actual_center = result[¢er];
629 let center_err = (actual_center - expected_center).abs();
630 assert!(
631 center_err < 1e-3,
632 "α={alpha}, N={n_leaves}: center mass drift |actual - closed-form| = {center_err}; \
633 expected={expected_center}, got={actual_center}"
634 );
635
636 for i in 0..n_leaves {
637 let leaf_id = fid(&format!("leaf_{i}.rs"), 1, 10);
638 let actual_leaf = result[&leaf_id];
639 let leaf_err = (actual_leaf - expected_leaf).abs();
640 assert!(
641 leaf_err < 1e-3,
642 "α={alpha}, N={n_leaves}, leaf_{i}: drift |actual - closed-form| = {leaf_err}; \
643 expected={expected_leaf}, got={actual_leaf}"
644 );
645 }
646 }
647 }
648
649 fn build_three_cycle() -> Graph {
650 let mut g = Graph::new();
651 let a = fid("a.rs", 1, 10);
652 let b = fid("b.rs", 1, 10);
653 let c = fid("c.rs", 1, 10);
654 g.add_node(a.clone());
655 g.add_node(b.clone());
656 g.add_node(c.clone());
657 g.add_edge(a, b.clone(), 1.0);
658 g.add_edge(b, c.clone(), 1.0);
659 g.add_edge(c, fid("a.rs", 1, 10), 1.0);
660 g
661 }
662
663 #[test]
672 fn ppr_truncates_under_tight_tolerance_and_high_alpha() {
673 let mut g = build_three_cycle();
674 let a = fid("a.rs", 1, 10);
675 let seeds: FxHashSet<FragmentId> = std::iter::once(a).collect();
676
677 let result = personalized_pagerank(&mut g, &seeds, 0.99, 1e-12, 0.5, None);
678
679 assert!(
680 result.truncated,
681 "alpha=0.99, tol=1e-12 on a 3-cycle must exhaust the push budget"
682 );
683 assert_eq!(
684 result.forward_pushes, 300,
685 "pinned push count for this fixture; a regression in max_pushes or the push loop \
686 would move this number"
687 );
688 }
689
690 #[test]
691 fn ppr_does_not_truncate_under_loose_tolerance() {
692 let mut g = build_three_cycle();
693 let a = fid("a.rs", 1, 10);
694 let seeds: FxHashSet<FragmentId> = std::iter::once(a).collect();
695
696 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-6, 0.5, None);
697
698 assert!(
699 !result.truncated,
700 "alpha=0.6, tol=1e-6 on a 3-cycle must converge well within the push budget"
701 );
702 }
703
704 #[test]
705 fn ppr_chain_with_dangling_terminal_node_sums_to_one() {
706 let mut g = Graph::new();
711 let a = fid("a.rs", 1, 10);
712 let b = fid("b.rs", 1, 10);
713 let c = fid("c.rs", 1, 10);
714 g.add_node(a.clone());
715 g.add_node(b.clone());
716 g.add_node(c.clone());
717 g.add_edge(a.clone(), b.clone(), 1.0);
718 g.add_edge(b.clone(), c.clone(), 1.0);
719
720 let seeds: FxHashSet<FragmentId> = std::iter::once(a).collect();
721 let result = personalized_pagerank(&mut g, &seeds, 0.6, 1e-8, 0.5, None);
722
723 assert!(result.seeded);
724 let total: f64 = result.scores.values().sum();
725 assert!(
726 (total - 1.0).abs() < 1e-6,
727 "mass must still sum to 1.0 despite the dangling terminal node; got {total}"
728 );
729 }
730
731 #[test]
732 fn ppr_seed_absent_from_graph_is_distinguishable_from_zero_seed_weights() {
733 let a = fid("a.rs", 1, 10);
734 let b = fid("b.rs", 1, 10);
735 let ghost = fid("ghost.rs", 1, 10);
736
737 let mut g_present = Graph::new();
740 g_present.add_node(a.clone());
741 g_present.add_node(b.clone());
742 g_present.add_edge(a.clone(), b.clone(), 1.0);
743 let seeds_present: FxHashSet<FragmentId> = std::iter::once(a.clone()).collect();
744 let present = personalized_pagerank(&mut g_present, &seeds_present, 0.6, 1e-6, 0.5, None);
745 assert!(present.seeded);
746 assert!(!present.scores.is_empty());
747
748 let mut g_absent = Graph::new();
750 g_absent.add_node(a.clone());
751 g_absent.add_node(b.clone());
752 g_absent.add_edge(a.clone(), b.clone(), 1.0);
753 let seeds_absent: FxHashSet<FragmentId> = std::iter::once(ghost).collect();
754 let absent = personalized_pagerank(&mut g_absent, &seeds_absent, 0.6, 1e-6, 0.5, None);
755 assert!(
756 !absent.seeded,
757 "a seed id absent from the graph must not be reported as seeded"
758 );
759 assert!(absent.scores.is_empty());
760 assert_eq!(absent.forward_pushes, 0);
761
762 let mut g_zeroed = Graph::new();
767 g_zeroed.add_node(a.clone());
768 g_zeroed.add_node(b.clone());
769 g_zeroed.add_edge(a.clone(), b.clone(), 1.0);
770 let seeds_zeroed: FxHashSet<FragmentId> = std::iter::once(a.clone()).collect();
771 let mut zero_weights: FxHashMap<FragmentId, f64> = FxHashMap::default();
772 zero_weights.insert(a, 0.0);
773 let zeroed = personalized_pagerank(
774 &mut g_zeroed,
775 &seeds_zeroed,
776 0.6,
777 1e-6,
778 0.5,
779 Some(&zero_weights),
780 );
781 assert!(
782 !zeroed.seeded,
783 "an all-zero seed_weights map must not be reported as seeded"
784 );
785 assert!(zeroed.scores.is_empty());
786 assert_eq!(zeroed.forward_pushes, 0);
787
788 assert_eq!(
789 (absent.scores.len(), absent.forward_pushes, absent.seeded),
790 (zeroed.scores.len(), zeroed.forward_pushes, zeroed.seeded),
791 "both degenerate cases produce identical scores/pushes -- `seeded` is currently the \
792 only signal telling them apart from each other and from real convergence-to-zero"
793 );
794 }
795
796 #[test]
797 fn ppr_excludes_self_loop_and_infinite_weight_and_stays_finite() {
798 use crate::graph::{EdgeCategory, build_graph};
799 use crate::types::{Fragment, FragmentKind};
800
801 let a = fid("a.rs", 1, 10);
802 let b = fid("b.rs", 1, 10);
803 let plain = |id: FragmentId| Fragment {
804 id,
805 kind: FragmentKind::Function,
806 content: Arc::from(""),
807 identifiers: FxHashSet::default(),
808 token_count: 10,
809 symbol_name: None,
810 };
811 let frags = vec![plain(a.clone()), plain(b.clone())];
812
813 let mut edges: FxHashMap<(FragmentId, FragmentId), f64> = FxHashMap::default();
814 let mut cats: FxHashMap<(FragmentId, FragmentId), EdgeCategory> = FxHashMap::default();
815 edges.insert((a.clone(), b.clone()), 1.0);
816 cats.insert((a.clone(), b.clone()), EdgeCategory::Semantic);
817 edges.insert((a.clone(), a.clone()), 5.0); cats.insert((a.clone(), a.clone()), EdgeCategory::Semantic);
819 edges.insert((b.clone(), b.clone()), f64::INFINITY); cats.insert((b.clone(), b.clone()), EdgeCategory::Semantic);
821
822 let mut graph = build_graph(&frags, edges, cats);
823 let seeds: FxHashSet<FragmentId> = std::iter::once(a).collect();
824 let result = personalized_pagerank(&mut graph, &seeds, 0.6, 1e-6, 0.5, None);
825
826 assert!(result.seeded);
827 assert!(
828 !result.scores.is_empty(),
829 "a self-loop / infinite-weight edge must not poison out_weight_sum into NaN and \
830 empty out the whole score map"
831 );
832 for (id, score) in &result.scores {
833 assert!(
834 score.is_finite(),
835 "score for {id} must be finite, got {score}"
836 );
837 }
838 }
839}