lean_ctx/core/
pagerank.rs1use 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
66pub 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}