use super::*;
fn cfg(threshold: f32, min_silence_ms: u32, min_speech_ms: u32, speech_pad_ms: u32) -> VadConfig {
VadConfig {
threshold,
min_silence_ms,
min_speech_ms,
speech_pad_ms,
}
}
#[test]
fn test_ms_to_samples_16khz() {
assert_eq!(VadConfig::ms_to_samples(1000), 16000);
assert_eq!(VadConfig::ms_to_samples(500), 8000);
assert_eq!(VadConfig::ms_to_samples(0), 0);
}
#[test]
fn test_regions_empty_probs_is_empty() {
let c = VadConfig::default();
assert!(regions_from_probs(&[], 512, 0, &c).is_empty());
assert!(regions_from_probs(&[0.9, 0.9], 512, 0, &c).is_empty());
}
#[test]
fn test_regions_all_silence_is_empty() {
let c = cfg(0.5, 0, 0, 0);
let probs = vec![0.1f32; 10];
assert!(regions_from_probs(&probs, 512, 10 * 512, &c).is_empty());
}
#[test]
fn test_regions_single_block_no_pad_no_mins() {
let c = cfg(0.5, 0, 0, 0);
let probs = [0.1, 0.9, 0.9, 0.1];
let r = regions_from_probs(&probs, 100, 400, &c);
assert_eq!(r, vec![(100, 300)]);
}
#[test]
fn test_regions_trailing_speech_clamps_to_total() {
let c = cfg(0.5, 0, 0, 0);
let probs = [0.1, 0.9, 0.9];
let r = regions_from_probs(&probs, 100, 250, &c);
assert_eq!(r, vec![(100, 250)]);
}
#[test]
fn test_regions_min_silence_merges_short_gap() {
let c = cfg(0.5, 100, 0, 0); let probs = [0.9, 0.1, 0.9];
let r = regions_from_probs(&probs, 100, 300, &c);
assert_eq!(r, vec![(0, 300)]);
}
#[test]
fn test_regions_long_gap_keeps_two_regions() {
let c = cfg(0.5, 0, 0, 0);
let probs = [0.9, 0.1, 0.1, 0.9];
let r = regions_from_probs(&probs, 100, 400, &c);
assert_eq!(r, vec![(0, 100), (300, 400)]);
}
#[test]
fn test_regions_min_speech_drops_short_blip() {
let c = cfg(0.5, 0, 100, 0);
let probs = [0.1, 0.9, 0.1];
assert!(regions_from_probs(&probs, 100, 300, &c).is_empty());
}
#[test]
fn test_regions_padding_extends_and_clamps() {
let c = cfg(0.5, 0, 0, 10); let probs = [0.1, 0.9, 0.1];
let r = regions_from_probs(&probs, 100, 1000, &c);
assert_eq!(r, vec![(0, 360)]);
}
#[test]
fn test_regions_padding_merges_overlapping_neighbours() {
let c = cfg(0.5, 0, 0, 50); let probs = [0.9, 0.1, 0.1, 0.9, 0.1];
let r = regions_from_probs(&probs, 100, 2000, &c);
assert_eq!(r, vec![(0, 1200)]);
}
#[test]
fn test_hangover_fires_once_after_min_silence() {
let c = cfg(0.5, 100, 0, 0); let mut h = Hangover::new(&c);
assert!(!h.update(0.9, 512));
assert!(!h.update(0.1, 512)); assert!(!h.update(0.1, 512)); assert!(!h.update(0.1, 512)); assert!(h.update(0.1, 512)); assert!(!h.update(0.1, 512));
}
#[test]
fn test_hangover_no_fire_before_any_speech() {
let c = cfg(0.5, 0, 0, 0);
let mut h = Hangover::new(&c);
for _ in 0..10 {
assert!(!h.update(0.1, 512));
}
}
#[test]
fn test_hangover_rearms_for_next_utterance() {
let c = cfg(0.5, 50, 0, 0); let mut h = Hangover::new(&c);
h.update(0.9, 512); assert!(!h.update(0.1, 512)); assert!(h.update(0.1, 512)); assert!(!h.update(0.9, 512));
assert!(!h.update(0.1, 512)); assert!(h.update(0.1, 512)); }
#[cfg(feature = "file-decode")]
mod segmenter {
use super::*;
fn stream(
probs: &[f32],
total: usize,
cfg: &VadConfig,
) -> (Vec<(usize, usize)>, Vec<f32>, usize) {
let raw: Vec<f32> = (0..total).map(|i| i as f32).collect();
let mut seg = VadSegmenter::new(cfg);
let mut out = Vec::new();
let mut it = probs.iter().copied();
let mut peak = 0usize;
let mut i = 0usize;
let mut chunk = 1usize;
while i < total {
let end = (i + chunk).min(total);
seg.push_with(&raw[i..end], &mut out, |_, _| Ok(it.next().unwrap_or(0.0)))
.expect("push");
peak = peak.max(seg.retained());
i = end;
chunk = chunk % 977 + 1;
}
seg.finish_with(total, &mut out, |_, _| Ok(it.next().unwrap_or(0.0)))
.expect("finish");
(seg.regions().to_vec(), out, peak)
}
fn batch(probs: &[f32], total: usize, cfg: &VadConfig) -> (Vec<(usize, usize)>, Vec<f32>) {
let regions = regions_from_probs(probs, VAD_FRAME_SAMPLES, total, cfg);
let out = regions
.iter()
.flat_map(|&(s, e)| (s..e).map(|i| i as f32))
.collect();
(regions, out)
}
fn assert_stream_matches_batch(probs: &[f32], total: usize, cfg: &VadConfig) {
let (got_regions, got_samples, _) = stream(probs, total, cfg);
let (want_regions, want_samples) = batch(probs, total, cfg);
assert_eq!(
got_regions, want_regions,
"regions diverged (total={total})"
);
assert_eq!(
got_samples, want_samples,
"compressed buffer diverged (total={total})"
);
}
fn probs_for(total: usize, f: impl Fn(usize) -> f32) -> Vec<f32> {
(0..total.div_ceil(VAD_FRAME_SAMPLES)).map(f).collect()
}
#[test]
fn test_segmenter_matches_batch_on_shaped_sequences() {
let c = VadConfig::default();
let fs = VAD_FRAME_SAMPLES;
for period in [1usize, 2, 3, 5, 8, 16, 20, 31, 64] {
let total = 200 * fs + 137; let probs = probs_for(total, |i| if (i / period) % 2 == 0 { 0.9 } else { 0.1 });
assert_stream_matches_batch(&probs, total, &c);
}
for level in [0.1f32, 0.9] {
let total = 97 * fs;
let probs = probs_for(total, |_| level);
assert_stream_matches_batch(&probs, total, &c);
}
}
#[test]
fn test_segmenter_matches_batch_on_degenerate_configs() {
let fs = VAD_FRAME_SAMPLES;
let total = 120 * fs + 11;
let probs = probs_for(total, |i| if (i / 7) % 3 == 0 { 0.9 } else { 0.1 });
for c in [
cfg(0.5, 0, 0, 0),
cfg(0.5, 0, 0, 200),
cfg(0.5, 10, 0, 500),
cfg(0.5, 1000, 2000, 100),
cfg(0.5, 40, 40, 40),
] {
assert_stream_matches_batch(&probs, total, &c);
}
}
#[test]
fn test_segmenter_matches_batch_on_short_and_empty_inputs() {
let c = VadConfig::default();
for total in [
0usize,
1,
2,
VAD_FRAME_SAMPLES - 1,
VAD_FRAME_SAMPLES,
VAD_FRAME_SAMPLES + 1,
] {
for level in [0.1f32, 0.9] {
assert_stream_matches_batch(&probs_for(total, |_| level), total, &c);
}
}
}
#[cfg(not(miri))]
proptest::proptest! {
#![proptest_config(proptest::prelude::ProptestConfig::with_cases(256))]
#[test]
fn prop_segmenter_matches_batch(
probs in proptest::collection::vec(0.0f32..=1.0, 1..60),
tail in 1usize..=VAD_FRAME_SAMPLES,
threshold in 0.1f32..0.9,
min_silence_ms in 0u32..800,
min_speech_ms in 0u32..500,
speech_pad_ms in 0u32..400,
) {
let total = (probs.len() - 1) * VAD_FRAME_SAMPLES + tail;
let c = cfg(threshold, min_silence_ms, min_speech_ms, speech_pad_ms);
let (got_regions, got_samples, _) = stream(&probs, total, &c);
let (want_regions, want_samples) = batch(&probs, total, &c);
proptest::prop_assert_eq!(got_regions, want_regions);
proptest::prop_assert_eq!(got_samples, want_samples);
}
}
#[test]
fn test_segmenter_retains_bounded_pcm_on_unbroken_speech() {
let c = VadConfig::default();
let total = 16000 * 3600;
let probs = probs_for(total, |_| 0.9);
let (regions, out, peak) = stream(&probs, total, &c);
assert_eq!(regions, vec![(0, total)]);
assert_eq!(out.len(), total);
let bound = VadConfig::ms_to_samples(c.min_speech_ms + c.min_silence_ms + c.speech_pad_ms)
+ VAD_FRAME_SAMPLES
+ 977; assert!(
peak <= bound,
"retained {peak} samples, expected at most {bound}"
);
}
#[test]
fn test_segmenter_retains_bounded_pcm_on_sparse_speech() {
let c = VadConfig::default();
let total = 16000 * 3600 * 3;
let probs = probs_for(total, |i| if (i / 40) % 5 == 0 { 0.9 } else { 0.1 });
let (regions, out, peak) = stream(&probs, total, &c);
assert!(!regions.is_empty());
assert_eq!(out.len(), regions.iter().map(|(s, e)| e - s).sum::<usize>());
let bound = VadConfig::ms_to_samples(c.min_speech_ms + c.min_silence_ms + c.speech_pad_ms)
+ VAD_FRAME_SAMPLES
+ 977;
assert!(
peak <= bound,
"retained {peak} samples, expected at most {bound}"
);
}
}
#[test]
fn test_remap_no_regions_is_identity() {
assert_eq!(remap_compressed_seconds(1.5, &[], 16000.0), 1.5);
}
#[test]
fn test_remap_single_region_offsets_by_start() {
let regions = [(16000usize, 32000usize)];
assert_eq!(remap_compressed_seconds(0.0, ®ions, 16000.0), 1.0);
assert_eq!(remap_compressed_seconds(0.5, ®ions, 16000.0), 1.5);
}
#[test]
fn test_remap_second_region_skips_silence_gap() {
let regions = [(0usize, 16000usize), (48000usize, 64000usize)];
assert_eq!(remap_compressed_seconds(0.5, ®ions, 16000.0), 0.5);
assert_eq!(remap_compressed_seconds(1.5, ®ions, 16000.0), 3.5);
}
#[test]
fn test_remap_past_end_clamps_to_last_region_end() {
let regions = [(0usize, 16000usize), (48000usize, 64000usize)];
assert_eq!(remap_compressed_seconds(10.0, ®ions, 16000.0), 4.0);
}
#[test]
#[ignore = "requires the Silero VAD model at ~/.gigastt/models/vad/silero_vad.onnx"]
fn test_silero_silence_low_prob_and_runs() {
let home = std::env::var("HOME").expect("HOME");
let path = std::path::PathBuf::from(home).join(".gigastt/models/vad/silero_vad.onnx");
if !path.exists() {
eprintln!("skipping {}: Silero VAD model not present", path.display());
return;
}
let vad = SileroVad::load(&path).expect("load silero");
let silence = vec![0.0f32; 16000];
let probs = vad.frame_probs(&silence).expect("frame_probs");
assert!(!probs.is_empty(), "expected at least one frame");
for p in &probs {
assert!((0.0..=1.0).contains(p), "prob {p} out of range");
}
let max_silence = probs.iter().cloned().fold(0.0f32, f32::max);
assert!(
max_silence < 0.5,
"silence should be below threshold, got {max_silence}"
);
let tone: Vec<f32> = (0..16000)
.map(|i| 0.5 * (2.0 * std::f32::consts::PI * 200.0 * i as f32 / 16000.0).sin())
.collect();
let probs2 = vad.frame_probs(&tone).expect("frame_probs tone");
for p in &probs2 {
assert!((0.0..=1.0).contains(p), "tone prob {p} out of range");
}
assert!(
vad.speech_regions(&silence, &VadConfig::default())
.expect("regions")
.is_empty()
);
}
fn silero_model_path() -> std::path::PathBuf {
let home = std::env::var("HOME").expect("HOME");
std::path::PathBuf::from(home).join(".gigastt/models/vad/silero_vad.onnx")
}
#[test]
#[ignore = "requires the Silero VAD model at ~/.gigastt/models/vad/silero_vad.onnx"]
fn test_endpointer_buffers_subframe_chunks_across_pushes() {
let path = silero_model_path();
if !path.exists() {
eprintln!("skipping {}: Silero VAD model not present", path.display());
return;
}
let vad = SileroVad::load(&path).expect("load silero");
let c = VadConfig::default();
let mut ep = VadEndpointer::new(&c);
let part = vec![0.0f32; 200];
assert!(!ep.push(&vad, &part).expect("push part 1"));
assert!(!ep.push(&vad, &part).expect("push part 2"));
let rest = vec![0.0f32; 200];
assert!(!ep.push(&vad, &rest).expect("push part 3")); }
#[test]
#[ignore = "requires the Silero VAD model at ~/.gigastt/models/vad/silero_vad.onnx"]
fn test_endpointer_no_endpoint_on_leading_silence() {
let path = silero_model_path();
if !path.exists() {
eprintln!("skipping {}: Silero VAD model not present", path.display());
return;
}
let vad = SileroVad::load(&path).expect("load silero");
let c = VadConfig::default();
let mut ep = VadEndpointer::new(&c);
let silence = vec![0.0f32; 16000];
assert!(
!ep.push(&vad, &silence).expect("push silence"),
"leading silence must not endpoint"
);
assert!(!ep.push(&vad, &[]).expect("push empty"));
}