use super::{FeatureSelectionError, top_k_mask, validate_matrix};
#[kani::proof]
fn fs_validate_matrix_rejects_non_finite() {
let a: f64 = kani::any();
let b: f64 = kani::any();
let c: f64 = kani::any();
let d: f64 = kani::any();
let x = vec![vec![a, b], vec![c, d]];
let all_finite = a.is_finite() && b.is_finite() && c.is_finite() && d.is_finite();
match validate_matrix(&x) {
Ok((rows, cols)) => {
assert!(all_finite, "Ok returned for a non-finite matrix");
assert!(rows == 2 && cols == 2, "shape was not the expected 2x2");
}
Err(FeatureSelectionError::NonFinite) => {
assert!(!all_finite, "NonFinite returned for an all-finite matrix");
}
Err(_) => assert!(false, "validate_matrix returned an unreachable variant"),
}
}
#[kani::proof]
#[kani::unwind(6)]
fn fs_top_k_mask_selects_min_k_n() {
let s0: f64 = kani::any();
let s1: f64 = kani::any();
let s2: f64 = kani::any();
for s in [s0, s1, s2] {
kani::assume(s.is_finite());
}
let k: usize = kani::any();
kani::assume(k <= 4);
let scores = [s0, s1, s2];
let mask = top_k_mask(&scores, k);
let selected = mask.iter().filter(|&&flag| flag).count();
assert!(selected == k.min(3), "selected count was not min(k, n)");
}