use super::*;
use crate::audio::ced::{Error, WINDOW_SAMPLES};
fn tone_window() -> Vec<f32> {
(0..WINDOW_SAMPLES)
.map(|i| 0.5 * (std::f32::consts::TAU * 1_000.0 * (i as f32 / 16_000.0)).sin())
.collect()
}
fn extract(samples: &[f32]) -> Vec<f32> {
let mel = MelExtractor::new();
let mut out = vec![f32::NAN; N_MELS * N_FRAMES];
mel.extract_into(samples, &mut out).expect("extract");
out
}
#[test]
fn frame_count_matches_the_believed_hop_geometry() {
assert_eq!(N_FRAMES, 1 + WINDOW_SAMPLES / 160);
assert_eq!(N_FRAMES, 1001);
assert_eq!(N_MELS, 64);
let out = extract(&tone_window());
assert!(
out.iter().all(|v| v.is_finite()),
"every mel element must be written and finite"
);
}
#[test]
fn frame_count_fits_upstream_pos_embed_capacity() {
const UPSTREAM_TARGET_LENGTH: usize = 1012;
const {
assert!(
N_FRAMES <= UPSTREAM_TARGET_LENGTH,
"believed mel width must fit upstream pos-embed capacity (target_length 1012)"
);
}
}
#[test]
fn silence_floors_at_amin_db() {
let out = extract(&vec![0.0f32; WINDOW_SAMPLES]);
for (i, &v) in out.iter().enumerate() {
assert_eq!(v, -100.0, "element {i} must sit exactly at the amin floor");
}
}
#[test]
fn tone_energy_lands_in_the_expected_htk_mel_bin() {
let out = extract(&tone_window());
let row_mean =
|m: usize| out[m * N_FRAMES..(m + 1) * N_FRAMES].iter().sum::<f32>() / N_FRAMES as f32;
let peak = (0..N_MELS)
.max_by(|&a, &b| row_mean(a).total_cmp(&row_mean(b)))
.unwrap();
assert_eq!(peak, 22, "1 kHz must peak in HTK mel bin 22");
assert!(
row_mean(22) > row_mean(60) + 30.0,
"tone bin must dominate a far bin by > 30 dB"
);
}
#[test]
fn layout_is_freq_major_rows() {
let mut samples = vec![0.0f32; WINDOW_SAMPLES];
let tone = tone_window();
samples[WINDOW_SAMPLES / 2..].copy_from_slice(&tone[WINDOW_SAMPLES / 2..]);
let out = extract(&samples);
let row = &out[22 * N_FRAMES..23 * N_FRAMES];
let first_half: f32 = row[..N_FRAMES / 2].iter().sum::<f32>() / (N_FRAMES / 2) as f32;
let second_half: f32 = row[N_FRAMES / 2..].iter().sum::<f32>() / (N_FRAMES - N_FRAMES / 2) as f32;
assert!(
second_half > first_half + 30.0,
"peak-bin row must be quiet-then-loud along the frame axis \
(first {first_half} dB, second {second_half} dB)"
);
}
#[test]
fn top_db_couples_the_floor_to_the_window_peak() {
let mut samples = vec![0.0f32; WINDOW_SAMPLES];
let tone = tone_window();
samples[WINDOW_SAMPLES / 2..].copy_from_slice(&tone[WINDOW_SAMPLES / 2..]);
let out = extract(&samples);
let max = out.iter().copied().fold(f32::MIN, f32::max);
let min = out.iter().copied().fold(f32::MAX, f32::min);
assert!(
max - (-100.0) > 120.0,
"the fixture must span more than 120 dB raw (max {max})"
);
assert!(
(min - (max - 120.0)).abs() < 1e-3,
"floor must clamp to max − 120 (max {max}, min {min})"
);
}
#[test]
fn short_input_is_zero_padded_at_the_waveform() {
let half = WINDOW_SAMPLES / 2;
let tone = tone_window();
let short = &tone[..half];
let mut padded = vec![0.0f32; WINDOW_SAMPLES];
padded[..half].copy_from_slice(short);
assert_eq!(extract(short), extract(&padded));
}
#[test]
fn empty_and_overlong_inputs_are_typed_errors() {
let mel = MelExtractor::new();
let mut out = vec![0.0f32; N_MELS * N_FRAMES];
assert!(matches!(
mel.extract_into(&[], &mut out),
Err(Error::EmptyAudio)
));
let long = vec![0.0f32; WINDOW_SAMPLES + 1];
assert!(matches!(
mel.extract_into(&long, &mut out),
Err(Error::AudioTooLong(e)) if e.len() == WINDOW_SAMPLES + 1 && e.max() == WINDOW_SAMPLES
));
}
#[test]
fn htk_mel_map_anchors_and_round_trips() {
assert!((MelExtractor::hz_to_htk_mel(1_000.0) - 999.99).abs() < 0.05);
for hz in [0.0f64, 125.0, 440.0, 1_000.0, 4_000.0, 8_000.0] {
let back = MelExtractor::htk_mel_to_hz(MelExtractor::hz_to_htk_mel(hz));
assert!((back - hz).abs() < 1e-6, "round trip at {hz} gave {back}");
}
}
#[test]
fn hann_window_is_periodic_length_512() {
let win = MelExtractor::periodic_hann(512);
assert_eq!(win.len(), 512);
assert_eq!(win[0], 0.0);
assert_eq!(win[256], 1.0);
assert!(win[511] > 0.0 && win[511] < 1e-3);
}
#[test]
fn filterbank_rows_are_unnormalized_triangles() {
let fb = MelExtractor::build_htk_filterbank(16_000, 512, 64, 0.0, 8_000.0);
assert_eq!(fb.len(), 64 * 257);
assert!(fb.iter().all(|&w| w >= 0.0));
for m in 0..64 {
let row = &fb[m * 257..(m + 1) * 257];
assert!(row.iter().sum::<f64>() > 0.0, "row {m} must carry mass");
let peak = row.iter().copied().fold(0.0f64, f64::max);
assert!(
peak <= 1.0 + 1e-9,
"norm=None peaks at ≤ 1.0, row {m} = {peak}"
);
}
}