1use crate::types::{ScoredCandidate, Topology};
19use xlog_core::RelId;
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
25pub struct ScoredPair {
26 pub topology: Topology,
27 pub left_rel_idx: RelId,
28 pub right_rel_idx: RelId,
29 pub positives_covered: u32,
30 pub negatives_covered: u32,
31}
32
33pub fn reduce_per_topology(
37 scored_pairs: &[ScoredPair],
38 head_rel_idx: RelId,
39 k_per_topology: u32,
40) -> Vec<ScoredCandidate> {
41 let mut result = Vec::new();
42 let k = k_per_topology as usize;
43
44 for topology in Topology::ALL {
45 let mut this_topo: Vec<ScoredPair> = scored_pairs
47 .iter()
48 .copied()
49 .filter(|p| p.topology == topology)
50 .collect();
51
52 this_topo.sort_by(|a, b| {
54 (
55 std::cmp::Reverse(a.positives_covered),
56 a.negatives_covered,
57 a.left_rel_idx.0,
58 a.right_rel_idx.0,
59 )
60 .cmp(&(
61 std::cmp::Reverse(b.positives_covered),
62 b.negatives_covered,
63 b.left_rel_idx.0,
64 b.right_rel_idx.0,
65 ))
66 });
67
68 let positives: Vec<ScoredPair> = this_topo
70 .iter()
71 .copied()
72 .filter(|p| p.positives_covered > 0)
73 .collect();
74
75 let kept_n = std::cmp::min(k, positives.len());
76 for (rank, pair) in positives.iter().take(kept_n).enumerate() {
77 let (next_pos, next_neg) = positives
79 .get(rank + 1)
80 .map(|nxt| (nxt.positives_covered, nxt.negatives_covered))
81 .unwrap_or((0, 0));
82
83 let tie_count = this_topo
85 .iter()
86 .filter(|s| {
87 s.positives_covered == pair.positives_covered
88 && s.negatives_covered == pair.negatives_covered
89 })
90 .count() as u32;
91
92 result.push(ScoredCandidate {
93 topology,
94 head_rel_idx,
95 left_rel_idx: pair.left_rel_idx,
96 right_rel_idx: pair.right_rel_idx,
97 positives_covered: pair.positives_covered,
98 negatives_covered: pair.negatives_covered,
99 local_rank: rank as u32,
100 next_positives_covered: next_pos,
101 next_negatives_covered: next_neg,
102 tie_class_size: tie_count,
103 });
104 }
105 }
106
107 result
108}
109
110#[derive(Debug, Clone, Copy, PartialEq, Eq)]
116pub struct KeptNaryPattern {
117 pub pattern_idx: usize,
118 pub positives_covered: u32,
119 pub negatives_covered: u32,
120 pub local_rank: u32,
121 pub next_positives_covered: u32,
122 pub next_negatives_covered: u32,
123 pub tie_class_size: u32,
124}
125
126pub fn reduce_nary(coverage: &[(u32, u32)], k: u32) -> Vec<KeptNaryPattern> {
140 let mut order: Vec<usize> = (0..coverage.len()).collect();
141 order.sort_by_key(|&i| (std::cmp::Reverse(coverage[i].0), coverage[i].1, i));
142
143 let positives: Vec<usize> = order
144 .iter()
145 .copied()
146 .filter(|&i| coverage[i].0 > 0)
147 .collect();
148
149 let mut tie_counts: std::collections::HashMap<(u32, u32), u32> =
154 std::collections::HashMap::new();
155 for &pair in coverage {
156 *tie_counts.entry(pair).or_insert(0) += 1;
157 }
158
159 let kept_n = std::cmp::min(k as usize, positives.len());
160 let mut result = Vec::with_capacity(kept_n);
161 for (rank, &idx) in positives.iter().take(kept_n).enumerate() {
162 let (pos, neg) = coverage[idx];
163 let (next_pos, next_neg) = positives
164 .get(rank + 1)
165 .map(|&nxt| coverage[nxt])
166 .unwrap_or((0, 0));
167 let tie_count = tie_counts.get(&(pos, neg)).copied().unwrap_or(0);
168 result.push(KeptNaryPattern {
169 pattern_idx: idx,
170 positives_covered: pos,
171 negatives_covered: neg,
172 local_rank: rank as u32,
173 next_positives_covered: next_pos,
174 next_negatives_covered: next_neg,
175 tie_class_size: tie_count,
176 });
177 }
178 result
179}
180
181#[cfg(test)]
182mod tests {
183 use super::*;
184
185 fn pair(topo: Topology, l: u32, r: u32, pos: u32, neg: u32) -> ScoredPair {
186 ScoredPair {
187 topology: topo,
188 left_rel_idx: RelId(l),
189 right_rel_idx: RelId(r),
190 positives_covered: pos,
191 negatives_covered: neg,
192 }
193 }
194
195 const HEAD: RelId = RelId(100);
196
197 #[test]
198 fn empty_input_yields_empty_output() {
199 let result = reduce_per_topology(&[], HEAD, 2);
200 assert!(result.is_empty());
201 }
202
203 #[test]
204 fn zero_k_yields_empty_output() {
205 let pairs = vec![pair(Topology::Chain, 1, 2, 5, 0)];
206 let result = reduce_per_topology(&pairs, HEAD, 0);
207 assert!(result.is_empty());
208 }
209
210 #[test]
211 fn zero_coverage_pairs_are_excluded_from_kept() {
212 let pairs = vec![
213 pair(Topology::Chain, 1, 2, 0, 0),
214 pair(Topology::Chain, 3, 4, 0, 0),
215 ];
216 let result = reduce_per_topology(&pairs, HEAD, 2);
217 assert!(result.is_empty());
218 }
219
220 #[test]
221 fn single_positive_pair_has_tie_one_and_no_next() {
222 let pairs = vec![pair(Topology::Chain, 1, 2, 5, 0)];
223 let result = reduce_per_topology(&pairs, HEAD, 2);
224 assert_eq!(result.len(), 1);
225 let c = &result[0];
226 assert_eq!(c.topology, Topology::Chain);
227 assert_eq!(c.left_rel_idx, RelId(1));
228 assert_eq!(c.right_rel_idx, RelId(2));
229 assert_eq!(c.positives_covered, 5);
230 assert_eq!(c.negatives_covered, 0);
231 assert_eq!(c.local_rank, 0);
232 assert_eq!(c.next_positives_covered, 0);
233 assert_eq!(c.next_negatives_covered, 0);
234 assert_eq!(c.tie_class_size, 1);
235 }
236
237 #[test]
238 fn max_positives_wins_over_higher_negatives() {
239 let pairs = vec![
241 pair(Topology::Chain, 1, 2, 3, 0),
242 pair(Topology::Chain, 3, 4, 5, 2),
243 ];
244 let result = reduce_per_topology(&pairs, HEAD, 1);
245 assert_eq!(result.len(), 1);
246 assert_eq!(result[0].positives_covered, 5);
247 assert_eq!(result[0].negatives_covered, 2);
248 }
249
250 #[test]
251 fn min_negatives_breaks_positives_tie() {
252 let pairs = vec![
254 pair(Topology::Chain, 1, 2, 5, 3),
255 pair(Topology::Chain, 3, 4, 5, 1),
256 ];
257 let result = reduce_per_topology(&pairs, HEAD, 1);
258 assert_eq!(result.len(), 1);
259 assert_eq!(result[0].left_rel_idx, RelId(3));
260 assert_eq!(result[0].negatives_covered, 1);
261 }
262
263 #[test]
264 fn left_idx_breaks_pos_neg_tie() {
265 let pairs = vec![
267 pair(Topology::Chain, 5, 2, 3, 0),
268 pair(Topology::Chain, 1, 2, 3, 0),
269 ];
270 let result = reduce_per_topology(&pairs, HEAD, 1);
271 assert_eq!(result.len(), 1);
272 assert_eq!(result[0].left_rel_idx, RelId(1));
273 }
274
275 #[test]
276 fn right_idx_breaks_all_other_ties() {
277 let pairs = vec![
279 pair(Topology::Chain, 1, 5, 3, 0),
280 pair(Topology::Chain, 1, 2, 3, 0),
281 ];
282 let result = reduce_per_topology(&pairs, HEAD, 1);
283 assert_eq!(result.len(), 1);
284 assert_eq!(result[0].right_rel_idx, RelId(2));
285 }
286
287 #[test]
288 fn top_k_truncation_preserves_order() {
289 let pairs = vec![
291 pair(Topology::Chain, 1, 2, 3, 0),
292 pair(Topology::Chain, 1, 3, 5, 0),
293 pair(Topology::Chain, 1, 4, 4, 0),
294 ];
295 let result = reduce_per_topology(&pairs, HEAD, 2);
296 assert_eq!(result.len(), 2);
297 assert_eq!(result[0].positives_covered, 5);
298 assert_eq!(result[0].local_rank, 0);
299 assert_eq!(result[1].positives_covered, 4);
300 assert_eq!(result[1].local_rank, 1);
301 }
302
303 #[test]
304 fn next_diagnostics_point_to_rank_plus_one_in_positive_list() {
305 let pairs = vec![
307 pair(Topology::Chain, 1, 1, 5, 0), pair(Topology::Chain, 2, 2, 4, 0), pair(Topology::Chain, 3, 3, 3, 0), ];
311 let result = reduce_per_topology(&pairs, HEAD, 2);
312 assert_eq!(result.len(), 2);
313 assert_eq!(result[0].next_positives_covered, 4);
315 assert_eq!(result[0].next_negatives_covered, 0);
316 assert_eq!(result[1].next_positives_covered, 3);
318 assert_eq!(result[1].next_negatives_covered, 0);
319 }
320
321 #[test]
322 fn next_diagnostics_are_zero_when_no_next() {
323 let pairs = vec![pair(Topology::Chain, 1, 1, 5, 0)];
325 let result = reduce_per_topology(&pairs, HEAD, 2);
326 assert_eq!(result[0].next_positives_covered, 0);
327 assert_eq!(result[0].next_negatives_covered, 0);
328 }
329
330 #[test]
331 fn tie_class_size_counts_same_pos_neg_in_full_sorted_list() {
332 let pairs = vec![
334 pair(Topology::Chain, 1, 1, 5, 0),
335 pair(Topology::Chain, 2, 2, 5, 0),
336 pair(Topology::Chain, 3, 3, 5, 0),
337 pair(Topology::Chain, 4, 4, 3, 0),
338 ];
339 let result = reduce_per_topology(&pairs, HEAD, 2);
340 assert_eq!(result[0].tie_class_size, 3);
342 assert_eq!(result[1].tie_class_size, 3);
343 }
344
345 #[test]
346 fn topologies_output_in_all_order() {
347 let pairs = vec![
349 pair(Topology::Fanin, 1, 1, 1, 0),
350 pair(Topology::Fanout, 1, 1, 1, 0),
351 pair(Topology::Star, 1, 1, 1, 0),
352 pair(Topology::Chain, 1, 1, 1, 0),
353 ];
354 let result = reduce_per_topology(&pairs, HEAD, 1);
355 assert_eq!(result.len(), 4);
356 assert_eq!(result[0].topology, Topology::Chain);
357 assert_eq!(result[1].topology, Topology::Star);
358 assert_eq!(result[2].topology, Topology::Fanout);
359 assert_eq!(result[3].topology, Topology::Fanin);
360 }
361
362 #[test]
363 fn topology_filtering_is_per_topology() {
364 let pairs = vec![
366 pair(Topology::Chain, 1, 1, 3, 0),
367 pair(Topology::Chain, 2, 2, 5, 0),
368 pair(Topology::Star, 3, 3, 2, 0),
369 pair(Topology::Star, 4, 4, 4, 0),
370 ];
371 let result = reduce_per_topology(&pairs, HEAD, 1);
372 assert_eq!(result.len(), 2);
373 assert_eq!(result[0].topology, Topology::Chain);
375 assert_eq!(result[0].positives_covered, 5);
376 assert_eq!(result[1].topology, Topology::Star);
378 assert_eq!(result[1].positives_covered, 4);
379 }
380
381 #[test]
382 fn k_larger_than_positives_returns_all_positives() {
383 let pairs = vec![
384 pair(Topology::Chain, 1, 1, 5, 0),
385 pair(Topology::Chain, 2, 2, 3, 0),
386 ];
387 let result = reduce_per_topology(&pairs, HEAD, 10);
388 assert_eq!(result.len(), 2);
389 }
390
391 #[test]
392 fn head_rel_idx_propagates_to_every_candidate() {
393 let pairs = vec![
394 pair(Topology::Chain, 1, 1, 5, 0),
395 pair(Topology::Star, 2, 2, 3, 0),
396 ];
397 let head = RelId(42);
398 let result = reduce_per_topology(&pairs, head, 1);
399 for c in &result {
400 assert_eq!(c.head_rel_idx, head);
401 }
402 }
403
404 #[test]
407 fn nary_empty_input_yields_empty_output() {
408 assert!(reduce_nary(&[], 2).is_empty());
409 }
410
411 #[test]
412 fn nary_zero_k_yields_empty_output() {
413 assert!(reduce_nary(&[(5, 0)], 0).is_empty());
414 }
415
416 #[test]
417 fn nary_zero_coverage_patterns_are_excluded() {
418 assert!(reduce_nary(&[(0, 0), (0, 3)], 2).is_empty());
419 }
420
421 #[test]
422 fn nary_max_positives_wins_over_higher_negatives() {
423 let kept = reduce_nary(&[(3, 0), (5, 2)], 1);
424 assert_eq!(kept.len(), 1);
425 assert_eq!(kept[0].pattern_idx, 1);
426 assert_eq!(kept[0].positives_covered, 5);
427 assert_eq!(kept[0].negatives_covered, 2);
428 }
429
430 #[test]
431 fn nary_min_negatives_breaks_positives_tie() {
432 let kept = reduce_nary(&[(5, 3), (5, 1)], 1);
433 assert_eq!(kept[0].pattern_idx, 1);
434 assert_eq!(kept[0].negatives_covered, 1);
435 }
436
437 #[test]
438 fn nary_enumeration_index_breaks_full_ties() {
439 let kept = reduce_nary(&[(5, 1), (5, 1)], 1);
440 assert_eq!(kept[0].pattern_idx, 0);
441 }
442
443 #[test]
444 fn nary_top_k_truncation_preserves_order_and_ranks() {
445 let kept = reduce_nary(&[(3, 0), (5, 0), (4, 0)], 2);
446 assert_eq!(kept.len(), 2);
447 assert_eq!(kept[0].pattern_idx, 1);
448 assert_eq!(kept[0].local_rank, 0);
449 assert_eq!(kept[1].pattern_idx, 2);
450 assert_eq!(kept[1].local_rank, 1);
451 }
452
453 #[test]
454 fn nary_next_diagnostics_point_past_kept_prefix() {
455 let kept = reduce_nary(&[(5, 0), (4, 0), (3, 7)], 2);
459 assert_eq!(kept.len(), 2);
460 assert_eq!(
461 (
462 kept[0].next_positives_covered,
463 kept[0].next_negatives_covered
464 ),
465 (4, 0)
466 );
467 assert_eq!(
468 (
469 kept[1].next_positives_covered,
470 kept[1].next_negatives_covered
471 ),
472 (3, 7)
473 );
474 }
475
476 #[test]
477 fn nary_next_diagnostics_are_zero_when_no_next() {
478 let kept = reduce_nary(&[(5, 0)], 2);
479 assert_eq!(kept[0].next_positives_covered, 0);
480 assert_eq!(kept[0].next_negatives_covered, 0);
481 }
482
483 #[test]
484 fn nary_tie_class_counts_full_batch_including_zero_coverage_ties() {
485 let kept = reduce_nary(&[(5, 0), (5, 0), (3, 0), (5, 0)], 2);
487 assert_eq!(kept[0].tie_class_size, 3);
488 assert_eq!(kept[1].tie_class_size, 3);
489 }
490
491 #[test]
492 fn nary_k_larger_than_positives_returns_all_positives() {
493 let kept = reduce_nary(&[(5, 0), (0, 0), (3, 0)], 10);
494 assert_eq!(kept.len(), 2);
495 }
496}