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
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 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}