1use std::path::Path;
2use std::sync::Arc;
3use std::time::Instant;
4
5use rustc_hash::{FxHashMap, FxHashSet};
6
7use crate::config::bm25::BM25;
8use crate::config::limits::{LIMITS, PPR};
9use crate::config::scoring::EGO;
10use crate::config::tokenization::TOKENIZATION;
11use crate::edges;
12use crate::filtering;
13use crate::graph::{self, Graph};
14use crate::ppr::personalized_pagerank;
15use crate::types::{DiffHunk, Fragment, FragmentId, extract_identifier_list};
16
17pub struct ScoringResult {
18 pub rel_scores: FxHashMap<FragmentId, f64>,
19 pub filtered_fragments: Vec<Fragment>,
20 pub graph: Graph,
21 pub graph_build_ms: f64,
26 pub ppr_truncated: bool,
29 pub ppr_forward_pushes: usize,
30 pub ppr_backward_pushes: usize,
31}
32
33pub trait ScoringStrategy: Send + Sync {
34 fn score_and_filter(
35 &self,
36 all_fragments: &[Fragment],
37 core_ids: &FxHashSet<FragmentId>,
38 hunks: &[DiffHunk],
39 repo_root: Option<&Path>,
40 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
41 discovered_paths: Option<&FxHashSet<Arc<str>>>,
42 ) -> ScoringResult;
43}
44
45pub struct PPRScoring {
46 pub alpha: f64,
47 pub low_relevance_filter: bool,
48}
49
50impl PPRScoring {
51 pub fn new(alpha: f64, low_relevance_filter: bool) -> Self {
52 Self {
53 alpha,
54 low_relevance_filter,
55 }
56 }
57}
58
59impl ScoringStrategy for PPRScoring {
60 fn score_and_filter(
61 &self,
62 all_fragments: &[Fragment],
63 core_ids: &FxHashSet<FragmentId>,
64 hunks: &[DiffHunk],
65 repo_root: Option<&Path>,
66 seed_weights: Option<&FxHashMap<FragmentId, f64>>,
67 _discovered_paths: Option<&FxHashSet<Arc<str>>>,
68 ) -> ScoringResult {
69 let skip_expensive = all_fragments.len() > LIMITS.skip_expensive_threshold;
70 let t_graph = Instant::now();
71 let capped = edges::collect_capped_edges(all_fragments, repo_root, skip_expensive);
72 let mut g = graph::build_graph_capped(all_fragments, capped);
73 let graph_build_ms = t_graph.elapsed().as_secs_f64() * 1000.0;
74 let ppr = personalized_pagerank(
75 &mut g,
76 core_ids,
77 self.alpha,
78 PPR.convergence_tolerance,
79 PPR.forward_blend,
80 seed_weights,
81 );
82 let mut rel_scores = ppr.scores;
83 if ppr.truncated {
84 tracing::warn!(
85 "PPR push-cap hit on {} nodes (fwd_pushes={}, bwd_pushes={}); rel_scores biased",
86 g.node_count(),
87 ppr.forward_pushes,
88 ppr.backward_pushes,
89 );
90 }
91 filtering::apply_hunk_proximity_bonus(&mut rel_scores, core_ids, all_fragments, hunks);
92
93 let filtered = filtering::filter_unrelated_fragments(all_fragments, core_ids, &g);
94 let filtered = if self.low_relevance_filter {
95 filtering::filter_low_relevance(filtered, core_ids, &rel_scores)
96 } else {
97 filtering::filter_positive_relevance(filtered, core_ids, &rel_scores)
98 };
99 let filtered = filtering::cap_context_fragments(filtered, core_ids, &rel_scores);
100
101 ScoringResult {
102 rel_scores,
103 filtered_fragments: filtered,
104 graph: g,
105 graph_build_ms,
106 ppr_truncated: ppr.truncated,
107 ppr_forward_pushes: ppr.forward_pushes,
108 ppr_backward_pushes: ppr.backward_pushes,
109 }
110 }
111}
112
113pub struct EgoGraphScoring {
114 pub max_depth: usize,
115}
116
117impl EgoGraphScoring {
118 pub fn new(max_depth: usize) -> Self {
119 Self { max_depth }
120 }
121}
122
123impl ScoringStrategy for EgoGraphScoring {
124 fn score_and_filter(
125 &self,
126 all_fragments: &[Fragment],
127 core_ids: &FxHashSet<FragmentId>,
128 _hunks: &[DiffHunk],
129 repo_root: Option<&Path>,
130 _seed_weights: Option<&FxHashMap<FragmentId, f64>>,
131 _discovered_paths: Option<&FxHashSet<Arc<str>>>,
132 ) -> ScoringResult {
133 let skip_expensive = all_fragments.len() > LIMITS.skip_expensive_threshold;
134 let t_graph = Instant::now();
135 let capped = edges::collect_capped_edges(all_fragments, repo_root, skip_expensive);
136 let g = graph::build_graph_capped(all_fragments, capped);
137 let graph_build_ms = t_graph.elapsed().as_secs_f64() * 1000.0;
138 let mut rel_scores = g.ego_graph(core_ids, self.max_depth);
139
140 let diff_idents: FxHashSet<String> = all_fragments
141 .iter()
142 .filter(|f| core_ids.contains(&f.id))
143 .flat_map(|f| f.identifiers.iter().cloned())
144 .collect();
145
146 if !diff_idents.is_empty() {
147 for frag in all_fragments {
148 if core_ids.contains(&frag.id) || !rel_scores.contains_key(&frag.id) {
149 continue;
150 }
151 let overlap = frag.identifiers.intersection(&diff_idents).count();
152 if overlap > 0 {
153 let bonus = EGO.identifier_overlap_epsilon
154 * overlap.min(EGO.identifier_overlap_cap) as f64
155 / EGO.identifier_overlap_cap as f64;
156 *rel_scores.get_mut(&frag.id).unwrap() += bonus;
157 }
158 }
159 }
160
161 let filtered = filtering::filter_unrelated_fragments(all_fragments, core_ids, &g);
162 let filtered = filtering::filter_positive_relevance(filtered, core_ids, &rel_scores);
163 let filtered = filtering::cap_context_fragments(filtered, core_ids, &rel_scores);
164
165 ScoringResult {
166 rel_scores,
167 filtered_fragments: filtered,
168 graph: g,
169 graph_build_ms,
170 ppr_truncated: false,
171 ppr_forward_pushes: 0,
172 ppr_backward_pushes: 0,
173 }
174 }
175}
176
177pub struct BM25Scoring;
178
179impl ScoringStrategy for BM25Scoring {
180 fn score_and_filter(
181 &self,
182 all_fragments: &[Fragment],
183 core_ids: &FxHashSet<FragmentId>,
184 _hunks: &[DiffHunk],
185 _repo_root: Option<&Path>,
186 _seed_weights: Option<&FxHashMap<FragmentId, f64>>,
187 _discovered_paths: Option<&FxHashSet<Arc<str>>>,
188 ) -> ScoringResult {
189 let query_tokens: Vec<String> = all_fragments
190 .iter()
191 .filter(|f| core_ids.contains(&f.id))
192 .flat_map(|f| {
193 extract_identifier_list(&f.content, TOKENIZATION.query_min_identifier_length)
194 })
195 .collect();
196 let query_set: FxHashSet<String> = query_tokens.into_iter().collect();
197
198 let docs: Vec<(FragmentId, Vec<String>)> = all_fragments
199 .iter()
200 .filter(|f| !core_ids.contains(&f.id))
201 .map(|f| {
202 (
203 f.id.clone(),
204 extract_identifier_list(&f.content, TOKENIZATION.query_min_identifier_length),
205 )
206 })
207 .collect();
208
209 let n_docs = docs.len().max(1);
210 let avgdl = docs.iter().map(|(_, d)| d.len()).sum::<usize>() as f64 / n_docs as f64;
211
212 let mut df: FxHashMap<String, usize> = FxHashMap::default();
213 for (_, doc) in &docs {
214 let unique: FxHashSet<&str> = doc.iter().map(|s| s.as_str()).collect();
215 for term in unique {
216 *df.entry(term.to_string()).or_insert(0) += 1;
217 }
218 }
219
220 let idf: FxHashMap<String, f64> = query_set
221 .iter()
222 .map(|t| {
223 let d = df.get(t).copied().unwrap_or(0) as f64;
224 let val =
225 ((n_docs as f64 - d + BM25.idf_smoothing) / (d + BM25.idf_smoothing)).ln_1p();
226 (t.clone(), val)
227 })
228 .collect();
229
230 let mut rel_scores: FxHashMap<FragmentId, f64> = FxHashMap::default();
231 for frag in all_fragments {
232 if core_ids.contains(&frag.id) {
233 rel_scores.insert(frag.id.clone(), 1.0);
234 }
235 }
236 for (fid, doc) in &docs {
237 let dl = doc.len() as f64;
238 let mut tf: FxHashMap<&str, u32> = FxHashMap::default();
239 for t in doc {
240 *tf.entry(t.as_str()).or_insert(0) += 1;
241 }
242 let mut score = 0.0;
243 for t in &query_set {
244 let freq = tf.get(t.as_str()).copied().unwrap_or(0) as f64;
245 if freq == 0.0 {
246 continue;
247 }
248 let idf_val = idf.get(t).copied().unwrap_or(0.0);
249 score += idf_val * (freq * BM25.k1)
250 / (freq + BM25.k1 * (1.0 - BM25.b + BM25.b * dl / avgdl));
251 }
252 if score > 0.0 {
253 rel_scores.insert(fid.clone(), score);
254 }
255 }
256
257 let max_score = rel_scores.values().copied().fold(0.0f64, f64::max);
258 if max_score > 0.0 {
259 for v in rel_scores.values_mut() {
260 *v /= max_score;
261 }
262 }
263
264 let filtered: Vec<Fragment> = all_fragments
265 .iter()
266 .filter(|f| {
267 core_ids.contains(&f.id) || rel_scores.get(&f.id).copied().unwrap_or(0.0) > 0.0
268 })
269 .cloned()
270 .collect();
271 let filtered = filtering::cap_context_fragments(filtered, core_ids, &rel_scores);
272
273 let g = Graph::new();
274 ScoringResult {
275 rel_scores,
276 filtered_fragments: filtered,
277 graph: g,
278 graph_build_ms: 0.0,
279 ppr_truncated: false,
280 ppr_forward_pushes: 0,
281 ppr_backward_pushes: 0,
282 }
283 }
284}