Skip to main content

_diffctx/
scoring.rs

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    /// Wall time spent constructing the typed dependency graph (edge
22    /// builders + dedup + hub suppression + per-source cap). Reported
23    /// separately so `scoring_ms` stays pure rank computation. Zero for
24    /// BM25 (no graph built).
25    pub graph_build_ms: f64,
26    /// PPR push-iteration was cut by `max_pushes_cap` before convergence.
27    /// Always false for non-PPR strategies (EGO/BM25).
28    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}