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 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
250pub 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 (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 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}