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        let fwd_w = graph
101            .forward_edge_weight(changed_frag, other_frag)
102            .unwrap_or(0.0);
103        let rev_w = graph
104            .forward_edge_weight(other_frag, changed_frag)
105            .unwrap_or(0.0);
106
107        if rev_w > fwd_w {
108            reverse_deps
109                .entry(changed_frag.path.clone())
110                .or_default()
111                .insert(other_frag.path.clone());
112        } else {
113            direct_edge_paths.insert(other_frag.path.clone());
114        }
115    });
116
117    (reverse_deps, direct_edge_paths)
118}
119
120fn find_hub_noise_paths(graph: &Graph, changed_paths: &FxHashSet<Arc<str>>) -> FxHashSet<Arc<str>> {
121    let (reverse_deps, direct_edge_paths) = classify_semantic_edges(graph, changed_paths);
122
123    let changed_dirs: FxHashSet<String> = changed_paths
124        .iter()
125        .filter_map(|p| {
126            Path::new(p.as_ref())
127                .parent()
128                .map(|d| d.to_string_lossy().into_owned())
129        })
130        .collect();
131
132    let mut noise_counts: FxHashMap<Arc<str>, usize> = FxHashMap::default();
133    for (hub_path, deps) in &reverse_deps {
134        if changed_paths.contains(hub_path) {
135            continue;
136        }
137        if deps.len() >= FILTERING.hub_reverse_threshold {
138            for dep in deps {
139                *noise_counts.entry(dep.clone()).or_insert(0) += 1;
140            }
141        }
142    }
143
144    noise_counts
145        .into_iter()
146        .filter(|(p, _count)| {
147            !direct_edge_paths.contains(p)
148                && !changed_dirs.contains(
149                    &Path::new(p.as_ref())
150                        .parent()
151                        .map(|d| d.to_string_lossy().into_owned())
152                        .unwrap_or_default(),
153                )
154        })
155        .map(|(p, _)| p)
156        .collect()
157}
158
159fn find_config_generic_code_files(
160    graph: &Graph,
161    changed_paths: &FxHashSet<Arc<str>>,
162) -> FxHashSet<Arc<str>> {
163    let mut has_real_edge: FxHashSet<Arc<str>> = FxHashSet::default();
164    let mut has_generic_config: FxHashSet<Arc<str>> = FxHashSet::default();
165    let mut generic_edge_count: FxHashMap<Arc<str>, usize> = FxHashMap::default();
166    let config_stems: FxHashSet<String> = changed_paths
167        .iter()
168        .filter_map(|p| {
169            Path::new(p.as_ref())
170                .file_stem()
171                .map(|s| s.to_string_lossy().to_lowercase())
172        })
173        .collect();
174
175    graph.for_each_categorized_edge(|src, dst, category| {
176        let src_changed = changed_paths.contains(&src.path);
177        let dst_changed = changed_paths.contains(&dst.path);
178        if !(src_changed ^ dst_changed) {
179            return;
180        }
181        let other_path = if src_changed { &dst.path } else { &src.path };
182        match category {
183            EdgeCategory::ConfigGeneric => {
184                has_generic_config.insert(other_path.clone());
185                *generic_edge_count.entry(other_path.clone()).or_insert(0) += 1;
186            }
187            EdgeCategory::Semantic | EdgeCategory::Config => {
188                has_real_edge.insert(other_path.clone());
189            }
190            _ => {}
191        }
192    });
193
194    let generic_only: FxHashSet<Arc<str>> = has_generic_config
195        .difference(&has_real_edge)
196        .cloned()
197        .collect();
198
199    generic_only
200        .into_iter()
201        .filter(|p| {
202            let path = Path::new(p.as_ref());
203            let ext = path
204                .extension()
205                .map(|e| format!(".{}", e.to_string_lossy().to_lowercase()))
206                .unwrap_or_default();
207            let stem = path
208                .file_stem()
209                .map(|s| s.to_string_lossy().to_lowercase())
210                .unwrap_or_default();
211            CODE_EXTENSIONS.contains(ext.as_str())
212                && generic_edge_count.get(p).copied().unwrap_or(0) <= 1
213                && !config_stems.contains(&stem)
214        })
215        .collect()
216}
217
218pub fn filter_unrelated_fragments(
219    fragments: &[Fragment],
220    core_ids: &FxHashSet<FragmentId>,
221    graph: &Graph,
222) -> Vec<Fragment> {
223    let changed_paths: FxHashSet<Arc<str>> = core_ids.iter().map(|fid| fid.path.clone()).collect();
224
225    let mut paths_to_remove = find_hub_noise_paths(graph, &changed_paths);
226    let config_generic = find_config_generic_code_files(graph, &changed_paths);
227    for p in config_generic {
228        paths_to_remove.insert(p);
229    }
230    for p in &changed_paths {
231        paths_to_remove.remove(p);
232    }
233
234    fragments
235        .iter()
236        .filter(|f| !paths_to_remove.contains(&f.id.path))
237        .cloned()
238        .collect()
239}
240
241pub fn filter_low_relevance(
242    fragments: Vec<Fragment>,
243    core_ids: &FxHashSet<FragmentId>,
244    rel: &FxHashMap<FragmentId, f64>,
245) -> Vec<Fragment> {
246    fragments
247        .into_iter()
248        .filter(|f| {
249            core_ids.contains(&f.id)
250                || rel.get(&f.id).copied().unwrap_or(0.0)
251                    >= effective_relevance_threshold(f.token_count)
252        })
253        .collect()
254}
255
256pub fn filter_positive_relevance(
257    fragments: Vec<Fragment>,
258    core_ids: &FxHashSet<FragmentId>,
259    rel: &FxHashMap<FragmentId, f64>,
260) -> Vec<Fragment> {
261    fragments
262        .into_iter()
263        .filter(|f| core_ids.contains(&f.id) || rel.get(&f.id).copied().unwrap_or(0.0) > 0.0)
264        .collect()
265}
266
267pub fn cap_context_fragments(
268    fragments: Vec<Fragment>,
269    core_ids: &FxHashSet<FragmentId>,
270    rel: &FxHashMap<FragmentId, f64>,
271) -> Vec<Fragment> {
272    let changed_paths: FxHashSet<Arc<str>> = core_ids.iter().map(|fid| fid.path.clone()).collect();
273
274    let mut ctx_by_path: FxHashMap<Arc<str>, Vec<Fragment>> = FxHashMap::default();
275    let mut result: Vec<Fragment> = Vec::new();
276
277    for f in fragments {
278        if changed_paths.contains(&f.id.path) {
279            result.push(f);
280        } else {
281            ctx_by_path.entry(f.id.path.clone()).or_default().push(f);
282        }
283    }
284
285    for (_path, mut file_frags) in ctx_by_path {
286        if file_frags.len() <= FILTERING.max_context_fragments_per_file {
287            result.extend(file_frags);
288        } else {
289            file_frags.sort_by(|a, b| {
290                let sa = rel.get(&a.id).copied().unwrap_or(0.0);
291                let sb = rel.get(&b.id).copied().unwrap_or(0.0);
292                sb.total_cmp(&sa).then_with(|| a.id.cmp(&b.id))
293            });
294            result.extend(
295                file_frags
296                    .into_iter()
297                    .take(FILTERING.max_context_fragments_per_file),
298            );
299        }
300    }
301
302    result.sort_by(|a, b| a.id.cmp(&b.id));
303    result
304}