use super::*;
fn synth_wav(samples: &[i16], sample_rate: u32) -> Vec<u8> {
let num_samples = samples.len();
let bytes_per_sample = 2;
let data_size = num_samples * bytes_per_sample;
let riff_size = 36 + data_size;
let mut out = Vec::with_capacity(44 + data_size);
out.extend_from_slice(b"RIFF");
out.extend_from_slice(&(riff_size as u32).to_le_bytes());
out.extend_from_slice(b"WAVE");
out.extend_from_slice(b"fmt ");
out.extend_from_slice(&16u32.to_le_bytes());
out.extend_from_slice(&1u16.to_le_bytes());
out.extend_from_slice(&1u16.to_le_bytes());
out.extend_from_slice(&sample_rate.to_le_bytes());
out.extend_from_slice(&(sample_rate * 2).to_le_bytes());
out.extend_from_slice(&2u16.to_le_bytes());
out.extend_from_slice(&16u16.to_le_bytes());
out.extend_from_slice(b"data");
out.extend_from_slice(&(data_size as u32).to_le_bytes());
for s in samples {
out.extend_from_slice(&s.to_le_bytes());
}
out
}
#[test]
fn decode_to_pcm_rejects_too_short_input() {
let err = decode_to_pcm(b"short").unwrap_err().to_string();
assert!(err.contains("too short"));
}
#[test]
fn decode_to_pcm_rejects_non_riff_container() {
let mut bytes = vec![0u8; 100];
bytes[..4].copy_from_slice(b"NOPE");
let err = decode_to_pcm(&bytes).unwrap_err().to_string();
assert!(err.contains("RIFF"));
}
#[test]
fn decode_to_pcm_round_trips_synthetic_wav() {
let samples: Vec<i16> = (0..1000)
.map(|i| ((i as f32 * 0.05).sin() * 16000.0) as i16)
.collect();
let wav = synth_wav(&samples, TARGET_SAMPLE_RATE);
let pcm = decode_to_pcm(&wav).expect("decode synth WAV");
assert_eq!(pcm.sample_rate, TARGET_SAMPLE_RATE);
assert_eq!(pcm.samples.len(), samples.len());
for (i, (got, want)) in pcm.samples.iter().zip(&samples).enumerate() {
let want_f = *want as f32 / i16::MAX as f32;
assert!(
(got - want_f).abs() < 1e-4,
"mismatch at {i}: {got} vs {want_f}"
);
}
}
#[test]
fn decode_to_pcm_resamples_non_target_rate_to_target() {
let samples: Vec<i16> = (0..400)
.map(|i| ((i as f32 * 0.1).sin() * 8000.0) as i16)
.collect();
let wav = synth_wav(&samples, 8_000);
let pcm = decode_to_pcm(&wav).expect("decode 8kHz");
assert_eq!(pcm.sample_rate, TARGET_SAMPLE_RATE);
assert!(pcm.samples.len() >= samples.len() * 19 / 10);
assert!(pcm.samples.len() <= samples.len() * 21 / 10);
}
#[test]
fn pcm_duration_secs_handles_zero_sample_rate() {
let p = PcmAudio {
samples: vec![0.0; 100],
sample_rate: 0,
};
assert_eq!(p.duration_secs(), 0.0);
}
#[test]
fn pcm_peak_returns_max_abs_sample() {
let p = PcmAudio::new(vec![0.0, -0.5, 0.3, -0.7, 0.2], 16000);
assert!((p.peak() - 0.7).abs() < 1e-6);
}
#[test]
fn pcm_rms_returns_root_mean_square() {
let p = PcmAudio::new(vec![1.0; 4], 16000);
assert!((p.rms() - 1.0).abs() < 1e-6);
let p = PcmAudio::new(vec![0.0; 4], 16000);
assert_eq!(p.rms(), 0.0);
let p = PcmAudio::new(vec![], 16000);
assert_eq!(p.rms(), 0.0);
}
#[test]
fn peak_normalise_brings_peak_to_target_dbfs() {
let p = PcmAudio::new(vec![0.1, -0.2, 0.05], 16000);
let normalised = peak_normalise(&p, -1.0);
let expected_peak = 10f32.powf(-1.0 / 20.0); assert!((normalised.peak() - expected_peak).abs() < 1e-4);
}
#[test]
fn peak_normalise_handles_silent_input_without_div_by_zero() {
let silent = PcmAudio::new(vec![0.0; 100], 16000);
let out = peak_normalise(&silent, -1.0);
assert_eq!(out.peak(), 0.0);
assert_eq!(out.samples.len(), 100);
}
#[test]
fn bandpass_filter_attenuates_dc_offset() {
let p = PcmAudio::new(vec![0.5; 16_000], 16_000);
let filtered = bandpass_filter(&p, 80.0, 3500.0);
let tail_rms = {
let tail = &filtered.samples[8000..];
let sum_sq: f32 = tail.iter().map(|s| s * s).sum();
(sum_sq / tail.len() as f32).sqrt()
};
assert!(
tail_rms < 0.05,
"DC offset survived bandpass: tail RMS = {tail_rms}"
);
}
#[test]
fn spectral_subtraction_attenuates_constant_noise_below_signal() {
let mut samples = vec![0.05f32; 16_000];
samples[8000] = 0.5;
let p = PcmAudio::new(samples, 16_000);
let denoised = spectral_subtraction_denoise(&p, 200);
assert!(
denoised.samples[8000].abs() > 0.4,
"spike was attenuated: {}",
denoised.samples[8000]
);
let tail_pre = &denoised.samples[1000..7000];
let pre_rms = {
let sum_sq: f32 = tail_pre.iter().map(|s| s * s).sum();
(sum_sq / tail_pre.len() as f32).sqrt()
};
assert!(
pre_rms < 0.05,
"noise floor not attenuated: pre-spike RMS = {pre_rms}"
);
}
#[test]
fn time_stretch_factor_one_returns_unchanged() {
let p = PcmAudio::new(vec![0.1, 0.2, 0.3], 16000);
let out = time_stretch(&p, 1.0);
assert_eq!(out.samples, p.samples);
}
#[test]
fn time_stretch_speeds_up_when_factor_below_one() {
let p = PcmAudio::new(vec![0.1; 16_000], 16_000);
let out = time_stretch(&p, 0.5);
assert!(
out.samples.len() < p.samples.len(),
"factor=0.5 should shorten: {} vs {}",
out.samples.len(),
p.samples.len()
);
}
#[test]
fn encode_wav_pcm16_produces_valid_header() {
let p = PcmAudio::new(vec![0.0, 0.5, -0.5, 0.25], 16000);
let wav = encode_wav_pcm16(&p);
assert_eq!(&wav[..4], b"RIFF");
assert_eq!(&wav[8..12], b"WAVE");
assert_eq!(&wav[12..16], b"fmt ");
assert_eq!(&wav[36..40], b"data");
assert_eq!(wav.len(), 52);
}
#[test]
fn encode_then_decode_round_trips_pcm() {
let original = PcmAudio::new(
(0..1000).map(|i| (i as f32 * 0.05).sin() * 0.5).collect(),
16_000,
);
let wav = encode_wav_pcm16(&original);
let decoded = decode_to_pcm(&wav).expect("round-trip decode");
assert_eq!(decoded.sample_rate, original.sample_rate);
assert_eq!(decoded.samples.len(), original.samples.len());
for (a, b) in decoded.samples.iter().zip(original.samples.iter()) {
assert!((a - b).abs() < 1e-3, "round-trip diff {a} vs {b}");
}
}
#[test]
fn encode_wav_clamps_overdriven_samples_instead_of_overflowing() {
let p = PcmAudio::new(vec![2.0, -2.0, 1.0, -1.0], 16000);
let wav = encode_wav_pcm16(&p);
let decoded = decode_to_pcm(&wav).unwrap();
assert!((decoded.samples[0] - 1.0).abs() < 1e-3);
assert!((decoded.samples[1] - -1.0).abs() < 1e-3);
}
#[test]
fn preprocess_for_stt_runs_full_pipeline_without_panic() {
let raw_samples: Vec<i16> = (0..16_000 * 2)
.map(|i| {
let signal = (i as f32 * 0.05).sin() * 8000.0;
let noise = ((i * 7919) % 100 - 50) as f32 * 80.0;
(signal + noise) as i16
})
.collect();
let wav = synth_wav(&raw_samples, 16_000);
let processed = preprocess_for_stt(&wav).expect("pipeline must not panic");
assert_eq!(processed.sample_rate, TARGET_SAMPLE_RATE);
assert!(!processed.samples.is_empty());
let expected_peak = 10f32.powf(-1.0 / 20.0);
assert!(
(processed.peak() - expected_peak).abs() < 0.05,
"peak normalisation off: {}",
processed.peak()
);
}
#[test]
fn preprocess_for_stt_rejects_invalid_input_with_actionable_error() {
let err = preprocess_for_stt(b"not a wav").unwrap_err().to_string();
assert!(err.contains("too short") || err.contains("RIFF"));
}