use crate::vad::{VadError, VoiceActivityDetector};
pub const ADAPTER_TYPE: &str = "earshot";
pub const FRAME_SIZE: usize = 256;
pub const SAMPLE_RATE: u32 = 16_000;
pub struct EarshotVad {
detector: Box<earshot::Detector>,
}
impl EarshotVad {
pub fn new() -> Self {
Self {
detector: earshot::Detector::default_boxed(),
}
}
pub fn frame_size(&self) -> usize {
FRAME_SIZE
}
fn score_frame(&mut self, frame: &[f32]) -> Result<f32, VadError> {
debug_assert_eq!(frame.len(), FRAME_SIZE);
let score = self.detector.predict_f32(frame);
if !(0.0..=1.0).contains(&score) {
return Err(VadError::Model(format!(
"earshot returned out-of-range score {score}"
)));
}
Ok(score)
}
}
impl Default for EarshotVad {
fn default() -> Self {
Self::new()
}
}
impl VoiceActivityDetector for EarshotVad {
fn reset(&mut self) {
self.detector.reset();
}
fn process(&mut self, samples: &[f32]) -> Result<Vec<f32>, VadError> {
if !samples.len().is_multiple_of(FRAME_SIZE) {
return Err(VadError::InvalidChunkSize {
expected: FRAME_SIZE,
got: samples.len(),
});
}
let mut probs = Vec::with_capacity(samples.len() / FRAME_SIZE);
for frame in samples.chunks_exact(FRAME_SIZE) {
probs.push(self.score_frame(frame)?);
}
Ok(probs)
}
fn sample_rate(&self) -> u32 {
SAMPLE_RATE
}
}
#[cfg(feature = "download")]
pub fn register_with(
registry: &mut crate::models::AdapterRegistry,
) -> Result<(), crate::models::AdapterError> {
use crate::models::{AdapterFactory, AdapterStage, BuiltinAdapter};
use std::sync::Arc;
let factory: AdapterFactory = Arc::new(|| {
Box::new(BuiltinAdapter {
stage: AdapterStage::Vad,
id: ADAPTER_TYPE.to_owned(),
})
});
registry.register(AdapterStage::Vad, ADAPTER_TYPE, factory)?;
Ok(())
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
use crate::vad::VoiceActivityDetector;
fn speechish(n: usize) -> Vec<f32> {
(0..n)
.map(|i| {
let t = i as f32 / SAMPLE_RATE as f32;
let f0 = 180.0;
let carrier = (2.0 * std::f32::consts::PI * f0 * t).sin()
+ 0.5 * (2.0 * std::f32::consts::PI * (2.0 * f0) * t).sin()
+ 0.25 * (2.0 * std::f32::consts::PI * (3.0 * f0) * t).sin();
let envelope = 0.5 + 0.5 * (2.0 * std::f32::consts::PI * 4.0 * t).sin();
(carrier * envelope * 0.35).clamp(-1.0, 1.0)
})
.collect()
}
#[test]
fn construct_and_sample_rate() {
let vad = EarshotVad::new();
assert_eq!(vad.sample_rate(), 16_000);
assert_eq!(vad.frame_size(), 256);
assert_eq!(ADAPTER_TYPE, "earshot");
}
#[test]
fn silence_scores_low() {
let mut vad = EarshotVad::new();
let silence = vec![0.0f32; FRAME_SIZE * 31];
let probs = vad.process(&silence).expect("silence process");
assert_eq!(probs.len(), 31);
assert!(probs.iter().all(|&p| (0.0..=1.0).contains(&p)));
let mean = probs.iter().sum::<f32>() / probs.len() as f32;
assert!(
mean < 0.5,
"expected silence mean score < 0.5, got {mean} ({probs:?})"
);
}
#[test]
fn speech_scores_higher_than_silence() {
let mut vad = EarshotVad::new();
let silence = vec![0.0f32; FRAME_SIZE * 40];
let speech = speechish(FRAME_SIZE * 40);
let silence_probs = vad.process(&silence).expect("silence");
vad.reset();
let speech_probs = vad.process(&speech).expect("speech");
assert_eq!(silence_probs.len(), speech_probs.len());
let silence_mean = silence_probs.iter().sum::<f32>() / silence_probs.len() as f32;
let speech_mean = speech_probs.iter().sum::<f32>() / speech_probs.len() as f32;
assert!(
speech_mean > silence_mean,
"speech mean ({speech_mean}) should exceed silence mean ({silence_mean})"
);
assert!(speech_probs.iter().all(|&p| (0.0..=1.0).contains(&p)));
}
#[test]
fn reset_restores_initial_state() {
let mut vad = EarshotVad::new();
let odd = vec![0.1f32; 100];
let err = vad
.process(&odd)
.expect_err("partial chunk must be rejected");
assert!(matches!(
err,
VadError::InvalidChunkSize {
expected: FRAME_SIZE,
got: 100
}
));
let silence = vec![0.0f32; FRAME_SIZE * 8];
let first = vad.process(&silence).expect("silence");
vad.reset();
let second = vad.process(&silence).expect("post-reset silence");
assert_eq!(first, second, "reset must restore initial detector state");
}
#[test]
fn partial_chunks_are_rejected() {
let mut vad = EarshotVad::new();
for &n in &[1usize, 17, 100, 255, 257, 511, 1000] {
let chunk = speechish(n);
let err = vad.process(&chunk).expect_err("partial chunk");
assert!(matches!(
err,
VadError::InvalidChunkSize {
expected: FRAME_SIZE,
got
} if got == n
));
}
let full = speechish(FRAME_SIZE * 3);
let probs = vad.process(&full).expect("aligned chunk");
assert_eq!(probs.len(), 3);
}
#[test]
fn aligned_input_yields_one_prob_per_frame() {
let mut vad = EarshotVad::new();
for frames in [1usize, 2, 5] {
let audio = speechish(FRAME_SIZE * frames);
let probs = vad.process(&audio).expect("aligned");
assert_eq!(probs.len(), frames);
assert!(probs.iter().all(|&p| (0.0..=1.0).contains(&p)));
}
}
#[test]
fn empty_input_returns_empty() {
let mut vad = EarshotVad::new();
let probs = vad.process(&[]).expect("empty");
assert!(probs.is_empty());
}
#[test]
fn multi_chunk_equals_single_pass() {
let audio = speechish(FRAME_SIZE * 10);
let mut a = EarshotVad::new();
let mut b = EarshotVad::new();
let once = a.process(&audio).expect("once");
let mid = FRAME_SIZE * 4;
let mut streamed = b.process(&audio[..mid]).expect("part1");
streamed.extend(b.process(&audio[mid..]).expect("part2"));
assert_eq!(once.len(), streamed.len());
for (i, (x, y)) in once.iter().zip(streamed.iter()).enumerate() {
assert!(
(x - y).abs() < 1e-5,
"frame {i}: single-pass {x} vs streamed {y}"
);
}
}
#[cfg(feature = "download")]
#[test]
fn register_with_empty_registry() {
let mut reg = crate::models::AdapterRegistry::new();
register_with(&mut reg).unwrap();
assert!(reg.contains(crate::models::AdapterStage::Vad, ADAPTER_TYPE));
let err = register_with(&mut reg).expect_err("duplicate");
assert!(matches!(
err,
crate::models::AdapterError::AlreadyRegistered { .. }
));
}
}