Skip to main content

lean_ctx/core/
pagerank.rs

1//! PageRank computation on the Property Graph.
2//!
3//! Provides a reusable `compute` function that can be called by
4//! `ctx_architecture`, `ctx_overview`, and `ctx_fill` for importance-weighted
5//! context selection.
6
7use std::collections::{HashMap, HashSet};
8
9use rusqlite::Connection;
10
11pub struct PageRankInput {
12    pub files: HashSet<String>,
13    pub forward: HashMap<String, Vec<String>>,
14}
15
16impl PageRankInput {
17    pub fn from_connection(conn: &Connection) -> Self {
18        let mut files: HashSet<String> = HashSet::new();
19        let mut forward: HashMap<String, Vec<String>> = HashMap::new();
20
21        if let Ok(mut stmt) =
22            conn.prepare("SELECT DISTINCT file_path FROM nodes WHERE kind = 'file'")
23            && let Ok(rows) = stmt.query_map([], |row| row.get::<_, String>(0))
24        {
25            for f in rows.flatten() {
26                files.insert(f);
27            }
28        }
29
30        let edge_sql = "
31            SELECT DISTINCT n1.file_path, n2.file_path
32            FROM edges e
33            JOIN nodes n1 ON e.source_id = n1.id
34            JOIN nodes n2 ON e.target_id = n2.id
35            WHERE n1.kind = 'file' AND n2.kind = 'file'
36              AND n1.file_path != n2.file_path
37        ";
38        if let Ok(mut stmt) = conn.prepare(edge_sql)
39            && let Ok(rows) = stmt.query_map([], |row| {
40                Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?))
41            })
42        {
43            for row in rows.flatten() {
44                let (src, tgt) = row;
45                forward.entry(src).or_default().push(tgt);
46            }
47        }
48
49        for deps in forward.values_mut() {
50            deps.sort();
51            deps.dedup();
52        }
53
54        Self { files, forward }
55    }
56}
57
58pub fn compute(input: &PageRankInput, damping: f64, iterations: usize) -> HashMap<String, f64> {
59    compute_personalized(input, damping, iterations, &[])
60}
61
62/// Personalized PageRank: if `seed_files` is non-empty, teleportation bias goes
63/// to those files instead of uniform distribution. Handles dangling nodes
64/// (nodes with no outgoing edges) by redistributing their rank.
65pub fn compute_personalized(
66    input: &PageRankInput,
67    damping: f64,
68    iterations: usize,
69    seed_files: &[String],
70) -> HashMap<String, f64> {
71    let n = input.files.len();
72    if n == 0 {
73        return HashMap::new();
74    }
75
76    let personalization: HashMap<String, f64> = if seed_files.is_empty() {
77        let uniform = 1.0 / n as f64;
78        input.files.iter().map(|f| (f.clone(), uniform)).collect()
79    } else {
80        let valid_seeds: Vec<&String> = seed_files
81            .iter()
82            .filter(|f| input.files.contains(*f))
83            .collect();
84        if valid_seeds.is_empty() {
85            let uniform = 1.0 / n as f64;
86            input.files.iter().map(|f| (f.clone(), uniform)).collect()
87        } else {
88            let weight = 1.0 / valid_seeds.len() as f64;
89            let mut p = HashMap::new();
90            for f in &valid_seeds {
91                p.insert((*f).clone(), weight);
92            }
93            p
94        }
95    };
96
97    let dangling: HashSet<&String> = input
98        .files
99        .iter()
100        .filter(|f| !input.forward.contains_key(*f) || input.forward[*f].is_empty())
101        .collect();
102
103    let init = 1.0 / n as f64;
104    let mut rank: HashMap<String, f64> = input.files.iter().map(|f| (f.clone(), init)).collect();
105
106    let eps = 1e-8;
107    for _ in 0..iterations {
108        let dangling_sum: f64 = dangling
109            .iter()
110            .map(|f| rank.get(*f).copied().unwrap_or(0.0))
111            .sum();
112
113        let mut new_rank: HashMap<String, f64> = HashMap::with_capacity(n);
114
115        for f in &input.files {
116            let teleport = personalization.get(f).copied().unwrap_or(0.0);
117            let dangling_contrib = personalization.get(f).copied().unwrap_or(0.0) * dangling_sum;
118            new_rank.insert(
119                f.clone(),
120                (1.0 - damping) * teleport + damping * dangling_contrib,
121            );
122        }
123
124        for (node, neighbors) in &input.forward {
125            if neighbors.is_empty() {
126                continue;
127            }
128            let share = rank.get(node).copied().unwrap_or(0.0) / neighbors.len() as f64;
129            for neighbor in neighbors {
130                if let Some(nr) = new_rank.get_mut(neighbor) {
131                    *nr += damping * share;
132                }
133            }
134        }
135
136        let diff: f64 = input
137            .files
138            .iter()
139            .map(|f| {
140                (rank.get(f).copied().unwrap_or(0.0) - new_rank.get(f).copied().unwrap_or(0.0))
141                    .abs()
142            })
143            .sum();
144        rank = new_rank;
145
146        if diff < eps {
147            break;
148        }
149    }
150
151    rank
152}
153
154pub fn top_files(conn: &Connection, limit: usize) -> Vec<(String, f64)> {
155    top_files_personalized(conn, limit, &[])
156}
157
158pub fn top_files_personalized(
159    conn: &Connection,
160    limit: usize,
161    seed_files: &[String],
162) -> Vec<(String, f64)> {
163    let input = PageRankInput::from_connection(conn);
164    let ranks = compute_personalized(&input, 0.85, 50, seed_files);
165    let mut sorted: Vec<(String, f64)> = ranks.into_iter().collect();
166    sorted.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
167    sorted.truncate(limit);
168    sorted
169}
170
171#[cfg(test)]
172mod tests {
173    use super::*;
174    use crate::core::property_graph::{CodeGraph, Edge, EdgeKind, Node};
175
176    #[test]
177    fn pagerank_basic() {
178        let g = CodeGraph::open_in_memory().unwrap();
179        let a = g.upsert_node(&Node::file("a.rs")).unwrap();
180        let b = g.upsert_node(&Node::file("b.rs")).unwrap();
181        let c = g.upsert_node(&Node::file("c.rs")).unwrap();
182
183        g.upsert_edge(&Edge::new(a, b, EdgeKind::Imports)).unwrap();
184        g.upsert_edge(&Edge::new(a, c, EdgeKind::Imports)).unwrap();
185        g.upsert_edge(&Edge::new(b, c, EdgeKind::Imports)).unwrap();
186
187        let input = PageRankInput::from_connection(g.connection());
188        let ranks = compute(&input, 0.85, 30);
189
190        assert_eq!(ranks.len(), 3);
191        let rank_c = ranks.get("c.rs").copied().unwrap_or(0.0);
192        let rank_a = ranks.get("a.rs").copied().unwrap_or(0.0);
193        assert!(
194            rank_c > rank_a,
195            "c.rs should rank higher (more incoming): c={rank_c} a={rank_a}"
196        );
197    }
198
199    #[test]
200    fn top_files_limit() {
201        let g = CodeGraph::open_in_memory().unwrap();
202        for i in 0..10 {
203            g.upsert_node(&Node::file(&format!("f{i}.rs"))).unwrap();
204        }
205        let top = top_files(g.connection(), 3);
206        assert!(top.len() <= 3);
207    }
208
209    #[test]
210    fn empty_graph() {
211        let g = CodeGraph::open_in_memory().unwrap();
212        let top = top_files(g.connection(), 10);
213        assert!(top.is_empty());
214    }
215
216    #[test]
217    fn personalized_pagerank_boosts_seed() {
218        let g = CodeGraph::open_in_memory().unwrap();
219        let a = g.upsert_node(&Node::file("a.rs")).unwrap();
220        let b = g.upsert_node(&Node::file("b.rs")).unwrap();
221        let c = g.upsert_node(&Node::file("c.rs")).unwrap();
222
223        g.upsert_edge(&Edge::new(a, b, EdgeKind::Imports)).unwrap();
224        g.upsert_edge(&Edge::new(b, c, EdgeKind::Imports)).unwrap();
225
226        let input = PageRankInput::from_connection(g.connection());
227
228        let uniform = compute_personalized(&input, 0.85, 50, &[]);
229        let seeded = compute_personalized(&input, 0.85, 50, &["a.rs".to_string()]);
230
231        let a_uniform = uniform.get("a.rs").copied().unwrap_or(0.0);
232        let a_seeded = seeded.get("a.rs").copied().unwrap_or(0.0);
233
234        assert!(
235            a_seeded > a_uniform,
236            "seeded a.rs ({a_seeded}) should rank higher than uniform ({a_uniform})"
237        );
238    }
239
240    #[test]
241    fn early_convergence() {
242        let g = CodeGraph::open_in_memory().unwrap();
243        let a = g.upsert_node(&Node::file("a.rs")).unwrap();
244        let b = g.upsert_node(&Node::file("b.rs")).unwrap();
245        g.upsert_edge(&Edge::new(a, b, EdgeKind::Imports)).unwrap();
246        g.upsert_edge(&Edge::new(b, a, EdgeKind::Imports)).unwrap();
247
248        let input = PageRankInput::from_connection(g.connection());
249        let ranks = compute(&input, 0.85, 1000);
250        assert_eq!(ranks.len(), 2);
251    }
252}