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