Skip to main content

_diffctx/
filtering.rs

1use std::path::Path;
2use std::sync::Arc;
3
4use rayon::prelude::*;
5use rustc_hash::{FxHashMap, FxHashSet};
6
7use crate::config::extensions::CODE_EXTENSIONS;
8use crate::config::filtering::FILTERING;
9use crate::graph::{EdgeCategory, Graph};
10use crate::types::{DiffHunk, Fragment, FragmentId, FragmentKind};
11
12fn fragment_hunk_gap(frag_start: u32, frag_end: u32, hunk_start: u32, hunk_end: u32) -> u32 {
13    if frag_end < hunk_start {
14        hunk_start - frag_end
15    } else if frag_start > hunk_end {
16        frag_start - hunk_end
17    } else {
18        0
19    }
20}
21
22fn proximity_score(frag: &Fragment, file_hunks: &[(u32, u32)]) -> f64 {
23    let min_gap = file_hunks
24        .iter()
25        .map(|&(h_start, h_end)| {
26            fragment_hunk_gap(frag.start_line(), frag.end_line(), h_start, h_end)
27        })
28        .min()
29        .unwrap_or(u32::MAX);
30    let half_decay = if frag.kind == FragmentKind::Definition {
31        FILTERING.definition_proximity_half_decay
32    } else {
33        FILTERING.proximity_half_decay
34    };
35    FILTERING.proximity_floor_max / (1.0 + min_gap as f64 / half_decay)
36}
37
38fn effective_relevance_threshold(token_count: u32) -> f64 {
39    let size_factor = (token_count as f64 / FILTERING.size_penalty_base_tokens)
40        .max(1.0)
41        .powf(FILTERING.size_penalty_exponent);
42    FILTERING.low_relevance_threshold * size_factor
43}
44
45pub fn apply_hunk_proximity_bonus(
46    rel: &mut FxHashMap<FragmentId, f64>,
47    core_ids: &FxHashSet<FragmentId>,
48    fragments: &[Fragment],
49    hunks: &[DiffHunk],
50) {
51    let mut hunks_by_path: FxHashMap<&str, Vec<(u32, u32)>> = FxHashMap::default();
52    for h in hunks {
53        let (h_start, h_end) = h.core_selection_range();
54        hunks_by_path
55            .entry(h.path.as_ref())
56            .or_default()
57            .push((h_start, h_end));
58    }
59
60    let bonuses: Vec<(FragmentId, f64)> = fragments
61        .par_iter()
62        .filter(|frag| !core_ids.contains(&frag.id))
63        .filter_map(|frag| {
64            let file_hunks = hunks_by_path.get(frag.path())?;
65            let bonus = proximity_score(frag, file_hunks);
66            Some((frag.id.clone(), bonus))
67        })
68        .collect();
69
70    for (id, bonus) in bonuses {
71        let current = rel.get(&id).copied().unwrap_or(0.0);
72        if current < bonus {
73            rel.insert(id, bonus);
74        }
75    }
76}
77
78fn classify_semantic_edges(
79    graph: &Graph,
80    changed_paths: &FxHashSet<Arc<str>>,
81) -> (
82    FxHashMap<Arc<str>, FxHashSet<Arc<str>>>,
83    FxHashSet<Arc<str>>,
84) {
85    let mut reverse_deps: FxHashMap<Arc<str>, FxHashSet<Arc<str>>> = FxHashMap::default();
86    let mut direct_edge_paths: FxHashSet<Arc<str>> = FxHashSet::default();
87
88    graph.for_each_categorized_edge(|src, dst, category| {
89        if category != EdgeCategory::Semantic {
90            return;
91        }
92        let src_changed = changed_paths.contains(&src.path);
93        let dst_changed = changed_paths.contains(&dst.path);
94        if !(src_changed ^ dst_changed) {
95            return;
96        }
97
98        let (changed_frag, other_frag) = if src_changed { (src, dst) } else { (dst, src) };
99
100        // `graph.edge_categories` is capped in lockstep with the CSR
101        // (see `graph::assemble_graph`), so a categorized edge here is
102        // guaranteed to exist in the CSR too -- `fwd_w == rev_w == 0.0`
103        // can only mean a genuinely near-zero weight, never a
104        // capped-away phantom silently suppressing hub-noise filtering.
105        let fwd_w = graph
106            .forward_edge_weight(changed_frag, other_frag)
107            .unwrap_or(0.0);
108        let rev_w = graph
109            .forward_edge_weight(other_frag, changed_frag)
110            .unwrap_or(0.0);
111
112        if rev_w > fwd_w {
113            reverse_deps
114                .entry(changed_frag.path.clone())
115                .or_default()
116                .insert(other_frag.path.clone());
117        } else {
118            direct_edge_paths.insert(other_frag.path.clone());
119        }
120    });
121
122    (reverse_deps, direct_edge_paths)
123}
124
125fn find_hub_noise_paths(graph: &Graph, changed_paths: &FxHashSet<Arc<str>>) -> FxHashSet<Arc<str>> {
126    let (reverse_deps, direct_edge_paths) = classify_semantic_edges(graph, changed_paths);
127
128    let changed_dirs: FxHashSet<String> = changed_paths
129        .iter()
130        .filter_map(|p| {
131            Path::new(p.as_ref())
132                .parent()
133                .map(|d| d.to_string_lossy().into_owned())
134        })
135        .collect();
136
137    let mut noise_counts: FxHashMap<Arc<str>, usize> = FxHashMap::default();
138    for (hub_path, deps) in &reverse_deps {
139        if changed_paths.contains(hub_path) {
140            continue;
141        }
142        if deps.len() >= FILTERING.hub_reverse_threshold {
143            for dep in deps {
144                *noise_counts.entry(dep.clone()).or_insert(0) += 1;
145            }
146        }
147    }
148
149    noise_counts
150        .into_iter()
151        .filter(|(p, _count)| {
152            !direct_edge_paths.contains(p)
153                && !changed_dirs.contains(
154                    &Path::new(p.as_ref())
155                        .parent()
156                        .map(|d| d.to_string_lossy().into_owned())
157                        .unwrap_or_default(),
158                )
159        })
160        .map(|(p, _)| p)
161        .collect()
162}
163
164fn find_config_generic_code_files(
165    graph: &Graph,
166    changed_paths: &FxHashSet<Arc<str>>,
167) -> FxHashSet<Arc<str>> {
168    let mut has_real_edge: FxHashSet<Arc<str>> = FxHashSet::default();
169    let mut has_generic_config: FxHashSet<Arc<str>> = FxHashSet::default();
170    let mut generic_edge_count: FxHashMap<Arc<str>, usize> = FxHashMap::default();
171    let config_stems: FxHashSet<String> = changed_paths
172        .iter()
173        .filter_map(|p| {
174            Path::new(p.as_ref())
175                .file_stem()
176                .map(|s| s.to_string_lossy().to_lowercase())
177        })
178        .collect();
179
180    graph.for_each_categorized_edge(|src, dst, category| {
181        let src_changed = changed_paths.contains(&src.path);
182        let dst_changed = changed_paths.contains(&dst.path);
183        if !(src_changed ^ dst_changed) {
184            return;
185        }
186        let other_path = if src_changed { &dst.path } else { &src.path };
187        match category {
188            EdgeCategory::ConfigGeneric => {
189                has_generic_config.insert(other_path.clone());
190                *generic_edge_count.entry(other_path.clone()).or_insert(0) += 1;
191            }
192            EdgeCategory::Semantic | EdgeCategory::Config => {
193                has_real_edge.insert(other_path.clone());
194            }
195            _ => {}
196        }
197    });
198
199    let generic_only: FxHashSet<Arc<str>> = has_generic_config
200        .difference(&has_real_edge)
201        .cloned()
202        .collect();
203
204    generic_only
205        .into_iter()
206        .filter(|p| {
207            let path = Path::new(p.as_ref());
208            let ext = path
209                .extension()
210                .map(|e| format!(".{}", e.to_string_lossy().to_lowercase()))
211                .unwrap_or_default();
212            let stem = path
213                .file_stem()
214                .map(|s| s.to_string_lossy().to_lowercase())
215                .unwrap_or_default();
216            CODE_EXTENSIONS.contains(ext.as_str())
217                && generic_edge_count.get(p).copied().unwrap_or(0) <= 1
218                && !config_stems.contains(&stem)
219        })
220        .collect()
221}
222
223pub fn filter_unrelated_fragments(
224    fragments: &[Fragment],
225    core_ids: &FxHashSet<FragmentId>,
226    graph: &Graph,
227) -> Vec<Fragment> {
228    let changed_paths: FxHashSet<Arc<str>> = core_ids.iter().map(|fid| fid.path.clone()).collect();
229
230    let mut paths_to_remove = find_hub_noise_paths(graph, &changed_paths);
231    let config_generic = find_config_generic_code_files(graph, &changed_paths);
232    for p in config_generic {
233        paths_to_remove.insert(p);
234    }
235    for p in &changed_paths {
236        paths_to_remove.remove(p);
237    }
238
239    fragments
240        .iter()
241        .filter(|f| !paths_to_remove.contains(&f.id.path))
242        .cloned()
243        .collect()
244}
245
246pub fn filter_low_relevance(
247    fragments: Vec<Fragment>,
248    core_ids: &FxHashSet<FragmentId>,
249    rel: &FxHashMap<FragmentId, f64>,
250) -> Vec<Fragment> {
251    fragments
252        .into_iter()
253        .filter(|f| {
254            core_ids.contains(&f.id)
255                || rel.get(&f.id).copied().unwrap_or(0.0)
256                    >= effective_relevance_threshold(f.token_count)
257        })
258        .collect()
259}
260
261pub fn filter_positive_relevance(
262    fragments: Vec<Fragment>,
263    core_ids: &FxHashSet<FragmentId>,
264    rel: &FxHashMap<FragmentId, f64>,
265) -> Vec<Fragment> {
266    fragments
267        .into_iter()
268        .filter(|f| core_ids.contains(&f.id) || rel.get(&f.id).copied().unwrap_or(0.0) > 0.0)
269        .collect()
270}
271
272pub fn cap_context_fragments(
273    fragments: Vec<Fragment>,
274    core_ids: &FxHashSet<FragmentId>,
275    rel: &FxHashMap<FragmentId, f64>,
276) -> Vec<Fragment> {
277    let changed_paths: FxHashSet<Arc<str>> = core_ids.iter().map(|fid| fid.path.clone()).collect();
278
279    let mut ctx_by_path: FxHashMap<Arc<str>, Vec<Fragment>> = FxHashMap::default();
280    let mut result: Vec<Fragment> = Vec::new();
281
282    for f in fragments {
283        if changed_paths.contains(&f.id.path) {
284            result.push(f);
285        } else {
286            ctx_by_path.entry(f.id.path.clone()).or_default().push(f);
287        }
288    }
289
290    for (_path, mut file_frags) in ctx_by_path {
291        if file_frags.len() <= FILTERING.max_context_fragments_per_file {
292            result.extend(file_frags);
293        } else {
294            file_frags.sort_by(|a, b| {
295                let sa = rel.get(&a.id).copied().unwrap_or(0.0);
296                let sb = rel.get(&b.id).copied().unwrap_or(0.0);
297                sb.total_cmp(&sa).then_with(|| a.id.cmp(&b.id))
298            });
299            result.extend(
300                file_frags
301                    .into_iter()
302                    .take(FILTERING.max_context_fragments_per_file),
303            );
304        }
305    }
306
307    result.sort_by(|a, b| a.id.cmp(&b.id));
308    result
309}
310
311#[cfg(test)]
312mod tests {
313    use super::*;
314    use crate::types::FragmentKind;
315    use std::sync::Arc;
316
317    fn frag(path: &str, start: u32, end: u32) -> Fragment {
318        Fragment {
319            id: FragmentId::new(Arc::from(path), start, end),
320            kind: FragmentKind::Function,
321            content: Arc::from(""),
322            identifiers: FxHashSet::default(),
323            token_count: 10,
324            symbol_name: None,
325        }
326    }
327
328    #[test]
329    fn cap_context_fragments_output_is_sorted_and_shuffle_invariant() {
330        let changed_path = "changed.rs";
331        let mut core_ids: FxHashSet<FragmentId> = FxHashSet::default();
332        let mut fragments: Vec<Fragment> = Vec::new();
333        for i in 0..5u32 {
334            let f = frag(changed_path, i * 10, i * 10 + 5);
335            core_ids.insert(f.id.clone());
336            fragments.push(f);
337        }
338
339        let mut rel: FxHashMap<FragmentId, f64> = FxHashMap::default();
340        let over_cap_path = "hub.rs";
341        let n_context = FILTERING.max_context_fragments_per_file + 5;
342        for i in 0..n_context {
343            let f = frag(over_cap_path, (i as u32) * 10, (i as u32) * 10 + 5);
344            // Distinct, strictly descending scores: no ties, so the
345            // top-K selection itself is unambiguous and any remaining
346            // non-determinism can only come from the final id sort.
347            rel.insert(f.id.clone(), (n_context - i) as f64);
348            fragments.push(f);
349        }
350
351        let baseline = cap_context_fragments(fragments.clone(), &core_ids, &rel);
352
353        assert_eq!(
354            baseline.len(),
355            5 + FILTERING.max_context_fragments_per_file,
356            "core fragments bypass the per-file cap; context fragments truncate to it"
357        );
358
359        let baseline_ids: Vec<FragmentId> = baseline.iter().map(|f| f.id.clone()).collect();
360        let mut sorted_ids = baseline_ids.clone();
361        sorted_ids.sort();
362        assert_eq!(
363            baseline_ids, sorted_ids,
364            "cap_context_fragments output must be sorted by fragment id"
365        );
366
367        for shuffled in [
368            {
369                let mut v = fragments.clone();
370                v.reverse();
371                v
372            },
373            {
374                let mut v = fragments.clone();
375                v.rotate_left(7);
376                v
377            },
378            {
379                let mut v = fragments.clone();
380                v.sort_by(|a, b| {
381                    rel.get(&a.id)
382                        .copied()
383                        .unwrap_or(0.0)
384                        .total_cmp(&rel.get(&b.id).copied().unwrap_or(0.0))
385                });
386                v
387            },
388        ] {
389            let result = cap_context_fragments(shuffled, &core_ids, &rel);
390            let ids: Vec<FragmentId> = result.iter().map(|f| f.id.clone()).collect();
391            assert_eq!(
392                ids, baseline_ids,
393                "cap_context_fragments must be invariant under input ordering"
394            );
395        }
396    }
397}