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