1use std::path::{Path, PathBuf};
2use std::sync::Arc;
3
4use rustc_hash::{FxHashMap, FxHashSet};
5
6use crate::config::selection::rescue;
7use crate::fragmentation::create_whole_file_fragment;
8use crate::git::CatFileBatch;
9use crate::graph::Graph;
10use crate::interval::IntervalIndex;
11use crate::types::{Fragment, FragmentId};
12
13fn find_dangling_semantic_names(
14 selected: &[Fragment],
15 graph: &Graph,
16 frag_by_id: &FxHashMap<FragmentId, &Fragment>,
17 selected_ids: &FxHashSet<FragmentId>,
18) -> FxHashSet<String> {
19 let mut dangling = FxHashSet::default();
20 for frag in selected {
21 graph.for_each_forward_neighbor(&frag.id, |nbr_id, _w| {
22 if selected_ids.contains(nbr_id) {
23 return;
24 }
25 let cat = graph.edge_category(&frag.id, nbr_id);
26 if cat
27 .map(|c| c != crate::graph::EdgeCategory::Semantic)
28 .unwrap_or(true)
29 {
30 return;
31 }
32 if let Some(nbr_frag) = frag_by_id.get(nbr_id) {
33 if let Some(ref name) = nbr_frag.symbol_name {
34 dangling.insert(name.to_lowercase());
35 }
36 }
37 });
38 }
39 dangling
40}
41
42fn pick_best_fragment<'a>(
43 candidates: &[&'a Fragment],
44 selected_ids: &FxHashSet<FragmentId>,
45) -> Option<&'a Fragment> {
46 let available: Vec<&&'a Fragment> = candidates
47 .iter()
48 .filter(|c| !selected_ids.contains(&c.id))
49 .collect();
50 let full = available.iter().find(|f| !f.kind.is_signature()).copied();
51 let sig = available.iter().find(|f| f.kind.is_signature()).copied();
52 full.or(sig).map(|f| *f)
53}
54
55fn change_coverage_rank(f: &Fragment, core_ids: &FxHashSet<FragmentId>) -> u8 {
56 if core_ids.contains(&f.id) {
57 return 0;
58 }
59 let is_core_stub = f.kind.is_signature()
60 && core_ids
61 .iter()
62 .any(|c| c.path == f.id.path && c.start_line == f.id.start_line);
63 if is_core_stub { 1 } else { 2 }
64}
65
66fn pick_smallest_fitting(
67 candidates: &[Fragment],
68 selected_ids: &FxHashSet<FragmentId>,
69 budget_left: u32,
70 core_ids: &FxHashSet<FragmentId>,
71) -> Option<Fragment> {
72 let mut sorted: Vec<&Fragment> = candidates.iter().collect();
73 sorted.sort_by_key(|f| (change_coverage_rank(f, core_ids), f.token_count));
79 for cand in &sorted {
80 if cand.token_count == 0 || selected_ids.contains(&cand.id) {
81 continue;
82 }
83 if cand.token_count <= budget_left {
84 return Some((*cand).clone());
85 }
86 }
87 None
91}
92
93pub fn coherence_post_pass(
94 selected: &mut Vec<Fragment>,
95 all_fragments: &[Fragment],
96 graph: &Graph,
97 budget: u32,
98) {
99 let selected_ids: FxHashSet<FragmentId> = selected.iter().map(|f| f.id.clone()).collect();
100 let mut interval_idx = IntervalIndex::new();
101 for f in selected.iter() {
102 interval_idx.add(f);
103 }
104 let used: u32 = selected.iter().map(|f| f.token_count).sum();
105 let mut remaining = budget.saturating_sub(used);
106
107 let mut name_to_frags: FxHashMap<String, Vec<&Fragment>> = FxHashMap::default();
108 for f in all_fragments {
109 if let Some(ref name) = f.symbol_name {
110 name_to_frags
111 .entry(name.to_lowercase())
112 .or_default()
113 .push(f);
114 }
115 }
116
117 let frag_by_id: FxHashMap<FragmentId, &Fragment> =
118 all_fragments.iter().map(|f| (f.id.clone(), f)).collect();
119 let dangling_names = find_dangling_semantic_names(selected, graph, &frag_by_id, &selected_ids);
120
121 let mut added_ids = selected_ids;
122 for name in &dangling_names {
123 let candidates = match name_to_frags.get(name) {
124 Some(c) => c,
125 None => continue,
126 };
127 let pick = match pick_best_fragment(candidates, &added_ids) {
128 Some(p) => p,
129 None => continue,
130 };
131 if pick.token_count <= remaining
132 && !added_ids.contains(&pick.id)
133 && !interval_idx.overlaps(pick)
134 {
135 selected.push(pick.clone());
136 added_ids.insert(pick.id.clone());
137 interval_idx.add(pick);
138 remaining = remaining.saturating_sub(pick.token_count);
139 }
140 }
141}
142
143fn compute_rescue_threshold(
144 all_fragments: &[Fragment],
145 rel_scores: &FxHashMap<FragmentId, f64>,
146 core_ids: &FxHashSet<FragmentId>,
147) -> f64 {
148 let mut context_scores: Vec<f64> = all_fragments
149 .iter()
150 .filter(|f| !core_ids.contains(&f.id))
151 .map(|f| rel_scores.get(&f.id).copied().unwrap_or(0.0))
152 .filter(|&s| s > 0.0)
153 .collect();
154 if context_scores.is_empty() {
155 return f64::INFINITY;
156 }
157 context_scores.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
158 let idx = (context_scores.len() as f64 * (1.0 - rescue().min_score_percentile)) as usize;
159 context_scores[idx.min(context_scores.len() - 1)]
160}
161
162pub fn rescue_nontrivial_context(
163 selected: &mut Vec<Fragment>,
164 all_fragments: &[Fragment],
165 rel_scores: &FxHashMap<FragmentId, f64>,
166 core_ids: &FxHashSet<FragmentId>,
167 budget: u32,
168) {
169 let used: u32 = selected.iter().map(|f| f.token_count).sum();
170 let remaining = budget.saturating_sub(used);
171 let rescue_budget = remaining.min((budget as f64 * rescue().budget_fraction) as u32);
172 if rescue_budget == 0 {
173 return;
174 }
175
176 let min_score = compute_rescue_threshold(all_fragments, rel_scores, core_ids);
177 if min_score == f64::INFINITY {
178 return;
179 }
180
181 let selected_ids: FxHashSet<FragmentId> = selected.iter().map(|f| f.id.clone()).collect();
182 let selected_paths: FxHashSet<Arc<str>> = selected.iter().map(|f| f.id.path.clone()).collect();
183 let changed_paths: FxHashSet<Arc<str>> = core_ids.iter().map(|fid| fid.path.clone()).collect();
184
185 let mut candidates: Vec<&Fragment> = all_fragments
186 .iter()
187 .filter(|f| {
188 !selected_ids.contains(&f.id)
189 && !core_ids.contains(&f.id)
190 && !changed_paths.contains(&f.id.path)
191 && !selected_paths.contains(&f.id.path)
192 && rel_scores.get(&f.id).copied().unwrap_or(0.0) >= min_score
193 && f.token_count <= rescue_budget
194 })
195 .collect();
196 candidates.sort_by(|a, b| {
197 let sa = rel_scores.get(&a.id).copied().unwrap_or(0.0);
198 let sb = rel_scores.get(&b.id).copied().unwrap_or(0.0);
199 sb.partial_cmp(&sa).unwrap_or(std::cmp::Ordering::Equal)
200 });
201
202 let mut interval_idx = IntervalIndex::new();
203 for f in selected.iter() {
204 interval_idx.add(f);
205 }
206
207 let mut budget_used = 0u32;
208 for cand in candidates {
209 if budget_used + cand.token_count > rescue_budget {
210 continue;
211 }
212 if interval_idx.overlaps(cand) {
213 continue;
214 }
215 selected.push(cand.clone());
216 interval_idx.add(cand);
217 budget_used += cand.token_count;
218 }
219}
220
221pub fn ensure_changed_files_represented(
222 selected: &mut Vec<Fragment>,
223 all_fragments: &[Fragment],
224 changed_files: &[PathBuf],
225 remaining_budget: u32,
226 root_dir: &Path,
227 preferred_revs: &[String],
228 mut batch_reader: Option<&mut CatFileBatch>,
229 core_ids: &FxHashSet<FragmentId>,
230) {
231 let selected_paths: FxHashSet<String> = selected
232 .iter()
233 .map(|f| f.id.path.as_ref().to_string())
234 .collect();
235 let mut missing_paths: Vec<&PathBuf> = changed_files
236 .iter()
237 .filter(|p| !selected_paths.contains(&p.to_string_lossy().as_ref().to_string()))
238 .collect();
239 missing_paths.sort();
240
241 if missing_paths.is_empty() {
242 return;
243 }
244
245 let mut frags_by_path: FxHashMap<String, Vec<Fragment>> = FxHashMap::default();
246 for f in all_fragments {
247 let path_str = f.id.path.as_ref().to_string();
248 if missing_paths
249 .iter()
250 .any(|p| p.to_string_lossy().as_ref() == path_str)
251 {
252 frags_by_path.entry(path_str).or_default().push(f.clone());
253 }
254 }
255
256 let mut budget_left = remaining_budget;
257 let mut selected_ids: FxHashSet<FragmentId> = selected.iter().map(|f| f.id.clone()).collect();
258 let mut interval_idx = IntervalIndex::new();
259 for f in selected.iter() {
260 interval_idx.add(f);
261 }
262
263 for path in missing_paths.iter().copied() {
264 let path_str = path.to_string_lossy().to_string();
265 let candidates = frags_by_path.get(&path_str).cloned().unwrap_or_default();
266 let candidates = if candidates.is_empty() {
267 match create_whole_file_fragment(
268 path,
269 root_dir,
270 preferred_revs,
271 batch_reader.as_deref_mut(),
272 ) {
273 Some(f) => vec![f],
274 None => continue,
275 }
276 } else {
277 candidates
278 };
279
280 if let Some(picked) =
281 pick_smallest_fitting(&candidates, &selected_ids, budget_left, core_ids)
282 {
283 if !interval_idx.overlaps(&picked) {
284 budget_left = budget_left.saturating_sub(picked.token_count);
285 selected_ids.insert(picked.id.clone());
286 interval_idx.add(&picked);
287 selected.push(picked);
288 }
289 }
290 }
291}
292
293#[cfg(test)]
294mod tests {
295 use super::*;
296
297 fn frag(
298 path: &str,
299 start: u32,
300 end: u32,
301 kind: crate::types::FragmentKind,
302 tokens: u32,
303 ) -> Fragment {
304 Fragment {
305 id: FragmentId::new(Arc::from(path), start, end),
306 kind,
307 content: Arc::from(format!("fragment {path}:{start}-{end}")),
308 identifiers: FxHashSet::default(),
309 token_count: tokens,
310 symbol_name: None,
311 }
312 }
313
314 #[test]
320 fn ensure_changed_files_represented_prefers_core_fragment_when_it_fits() {
321 let core = frag("a.ts", 10, 20, crate::types::FragmentKind::Function, 60);
322 let stub = frag(
323 "a.ts",
324 10,
325 10,
326 crate::types::FragmentKind::FunctionSignature,
327 10,
328 );
329 let all_fragments = vec![core.clone(), stub.clone()];
330 let core_ids: FxHashSet<FragmentId> = std::iter::once(core.id.clone()).collect();
331 let changed_files = vec![PathBuf::from("a.ts")];
332 let mut selected: Vec<Fragment> = Vec::new();
333
334 ensure_changed_files_represented(
335 &mut selected,
336 &all_fragments,
337 &changed_files,
338 100,
339 Path::new("."),
340 &[],
341 None,
342 &core_ids,
343 );
344
345 assert_eq!(selected.len(), 1, "expected exactly one fallback fragment");
346 assert_eq!(
347 selected[0].id, core.id,
348 "fallback picked the signature stub instead of the fragment covering the actual diff hunk"
349 );
350 }
351
352 #[test]
356 fn ensure_changed_files_represented_falls_back_to_stub_when_core_does_not_fit() {
357 let core = frag("a.ts", 10, 20, crate::types::FragmentKind::Function, 60);
358 let stub = frag(
359 "a.ts",
360 10,
361 10,
362 crate::types::FragmentKind::FunctionSignature,
363 10,
364 );
365 let all_fragments = vec![core.clone(), stub.clone()];
366 let core_ids: FxHashSet<FragmentId> = std::iter::once(core.id.clone()).collect();
367 let changed_files = vec![PathBuf::from("a.ts")];
368 let mut selected: Vec<Fragment> = Vec::new();
369
370 ensure_changed_files_represented(
371 &mut selected,
372 &all_fragments,
373 &changed_files,
374 15,
375 Path::new("."),
376 &[],
377 None,
378 &core_ids,
379 );
380
381 assert_eq!(selected.len(), 1);
382 assert_eq!(selected[0].id, stub.id);
383 }
384}