Skip to main content

_diffctx/
ppr.rs

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
50/// Mirrors `init_seed_residuals`'s own validity check without building the
51/// residual vector, so `personalized_pagerank` can tell "no seed mass ever
52/// entered the push" apart from "the push ran and every score genuinely
53/// converged to zero" -- both cases otherwise return the same empty
54/// `scores` map with `forward_pushes == 0`, which is indistinguishable to
55/// a caller that just sees an empty selection and assumes it means the
56/// diff genuinely has no related context (a seed absent from the graph and
57/// an explicit all-zero `seed_weights` map, reachable on a deletion-only
58/// hunk, both hit this path).
59fn 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
81/// Result of one PPR push pass. `truncated` flags whether the
82/// `max_pushes` budget cut the iteration short — when true, the
83/// returned estimate is biased toward seeds and the absolute scores
84/// are not comparable to a converged run on the same graph. Surfaced
85/// to Python via `LatencyBreakdown.ppr_truncated` so calibration /
86/// final-eval rows can be filtered or flagged for the paper.
87struct 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
173/// Public PPR result. `truncated` is logical-OR of forward + backward
174/// push truncation flags. When true, downstream renormalization
175/// (sum-to-1 across nodes) hides the fact that the iteration was cut
176/// short by `max_pushes_cap`; the absolute relevance scores are
177/// biased toward seeds and not directly comparable across instances.
178/// We surface this to Python so calibration / final-eval rows can be
179/// filtered in post-analysis (PolyBench / Multi-SWE-bench instances
180/// with >20k fragments are the primary suspects).
181pub struct PprResult {
182    pub scores: FxHashMap<FragmentId, f64>,
183    pub truncated: bool,
184    pub forward_pushes: usize,
185    pub backward_pushes: usize,
186    /// False when no seed mass ever entered the push (no seed matched a
187    /// graph node, or `seed_weights` summed to <= 0 across matched seeds).
188    /// Lets the caller tell "nothing to seed from" apart from "seeded and
189    /// converged to a genuinely empty `scores` map" -- both produce the
190    /// same empty map and `forward_pushes == 0` otherwise.
191    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[&center] > max_leaf, "Center must dominate leaves");
407    }
408
409    /// Claim 7 (paper §4.4): hub suppression $w'_{uv} = w_{uv} / \ln(1 + \text{in\_deg}(v))$
410    /// reduces PPR mass concentration on hub nodes without removing them.
411    ///
412    /// Two graphs are compared:
413    ///   - Naive: a leaf-to-hub graph built directly via `Graph::add_edge`, no suppression.
414    ///   - Suppressed: same topology built via `build_graph` with non-exempt category,
415    ///     which triggers `apply_hub_suppression` for in-degree above the median.
416    ///
417    /// Expected: hub mass is materially reduced; non-hub mass is largely preserved.
418    #[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    /// Claim 10 (paper §4.4 hypothesis): PPR with $\alpha \in [0.5, 0.65]$ ranks nodes
507    /// similarly to ego-graph BFS scoring (hop-decay 1/(1+d)). Spearman rank correlation
508    /// of the two score vectors should be high on a connected graph.
509    #[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    /// Claim 6B (paper §4.4): personalized PageRank converges to the closed-form
602    /// stationary distribution π = (1-α)(I - αM)^(-1) p, normalized to sum to 1.
603    ///
604    /// For a symmetric star with center c and N leaves, all bidirectional weight 1:
605    ///     π_center = 1 / (1 + α)
606    ///     π_leaf   = α / (N · (1 + α))
607    /// Derivation:
608    ///     π_c = (1-α) + α · Σ_leaf π_leaf,   π_leaf = α · (1/N) · π_c
609    ///     ⇒ π_c · (1 + α) = 1 after normalization (Σ π = 1).
610    #[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[&center];
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    /// `max_pushes = n * PPR.push_scale_factor` (3*100=300 for this
664    /// fixture). A high alpha (slow restart) plus a tight tolerance
665    /// forces far more pushes than that budget, so the push loop must
666    /// stop early and report `truncated`. Deleting `truncated = true;`
667    /// in `ppr_push_csr` would make this pass silently -- the pinned
668    /// `forward_pushes` count is what actually catches that deletion,
669    /// since the flag alone can't distinguish "stopped early" from "ran
670    /// exactly to convergence at push 300".
671    #[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        // a -> b -> c, c has no outgoing edges (dangling). None of the
707        // existing fixtures (full cycle, symmetric star) has a node with
708        // zero out-weight; `ppr_push_csr`'s `total_w <= 0.0` branch is
709        // only exercised here.
710        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        // Baseline: a real, present, positively-weighted seed produces a
738        // non-empty, seeded result.
739        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        // A seed id that matches no node in the graph.
749        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        // A seed that IS in the graph, but `seed_weights` zeroes it out --
763        // reachable when every changed line in a hunk maps to a deleted
764        // fragment. Same empty `scores` / zero pushes as the absent-seed
765        // case above, but `seeded` must still tell them apart.
766        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); // self-loop
818        cats.insert((a.clone(), a.clone()), EdgeCategory::Semantic);
819        edges.insert((b.clone(), b.clone()), f64::INFINITY); // self-loop AND non-finite
820        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}