use super::*;
#[test]
fn new_valid() {
let r = Reservoir::new(1024, 10240).unwrap();
assert_eq!(r.size(), 10240);
}
#[test]
fn new_invalid_size() {
assert!(Reservoir::new(1024, 0).is_err());
assert!(Reservoir::new(1024, 200_000).is_err());
}
#[test]
fn step_ok() {
let mut r = Reservoir::new(1024, 10240).unwrap();
let out = r.step(&[0.0; 1024]).unwrap();
assert_eq!(out.state.len(), 10240);
}
#[test]
fn reset_clears() {
let mut r = Reservoir::new(1024, 10240).unwrap();
let _ = r.step(&[1.0; 1024]);
r.reset();
assert!(r.state().iter().all(|x| *x == 0.0));
}
#[test]
fn spectral_radius_bounds() {
let mut r = Reservoir::new(1024, 10240).unwrap();
assert!(r.set_spectral_radius(0.8).is_err());
r.set_spectral_radius(1.0).unwrap();
}
#[test]
fn metrics_steps() {
let mut r = Reservoir::new(1024, 10240).unwrap();
let _ = r.step(&[0.0; 1024]);
assert_eq!(r.metrics_snapshot().reservoir_steps_total, 1);
}
#[test]
fn step_mathematical_correctness() {
let mut r = Reservoir::new_seeded(2, 4, 42).unwrap();
r.alpha = 1.0; r.beta = 0.0; r.update_stride = 1;
let input = [0.5, -0.5];
let out = r.step(&input).unwrap();
assert_eq!(out.state.len(), 4);
for &val in out.state {
assert!((-1.0..=1.0).contains(&val)); }
let manual_norm: f64 = out
.state
.iter()
.map(|&x| (x as f64) * (x as f64))
.sum::<f64>()
.sqrt();
assert!((out.state_norm - manual_norm).abs() < 1e-6);
}
#[test]
#[allow(clippy::float_cmp)]
fn norm_calculation_consistency() {
for i in -100..100 {
let x = i as f64 * 0.01;
assert_eq!((x * x).to_bits(), x.powi(2).to_bits());
}
let mut r = Reservoir::new_seeded(10, 100, 42).unwrap();
let input = vec![0.1; 10];
for i in 0..100 {
let out = r.step(&input).unwrap();
let manual_norm_sq: f64 = out.state.iter().map(|&x| f64::from(x) * f64::from(x)).sum();
let manual_norm = manual_norm_sq.sqrt();
assert!(
(out.state_norm - manual_norm).abs() < 1e-9,
"Norm mismatch at step {}: internal={}, manual={}",
i,
out.state_norm,
manual_norm
);
if r.update_phase == 0 {
}
}
}
#[test]
fn step_short_input_returns_error() {
let mut r = Reservoir::new(1024, 10240).unwrap();
let res = r.step(&[0.0; 512]);
assert!(res.is_err());
let err = res.unwrap_err().to_string();
assert!(err.contains("Input size mismatch"));
}