use super::params::VadParams;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VadState {
Quiet,
Starting,
Speaking,
Stopping,
}
impl Default for VadState {
fn default() -> Self {
Self::Quiet
}
}
pub fn calculate_audio_volume(audio: &[u8]) -> f32 {
let samples: Vec<i16> = audio
.chunks_exact(2)
.map(|b| i16::from_le_bytes([b[0], b[1]]))
.collect();
if samples.is_empty() {
return 0.0;
}
let sum_sq: f64 = samples.iter().map(|&s| (s as f64).powi(2)).sum();
let rms = (sum_sq / samples.len() as f64).sqrt() as f32 / 32768.0;
if rms < 1e-9 {
return 0.0;
}
let db = (20.0 * rms.log10()).clamp(-60.0, 0.0);
(db + 60.0) / 60.0
}
#[inline]
pub fn exp_smoothing(current: f32, prev: f32, factor: f32) -> f32 {
current * factor + prev * (1.0 - factor)
}
const SMOOTHING_FACTOR: f32 = 0.2;
pub struct StateMachine {
params: VadParams,
frames_required: usize,
bytes_required: usize,
start_frames: usize,
stop_frames: usize,
starting_count: usize,
stopping_count: usize,
prev_volume: f32,
buffer: Vec<u8>,
pub state: VadState,
}
impl StateMachine {
pub fn new(sample_rate: u32, params: VadParams) -> Self {
let frames_required: usize = if sample_rate == 16000 { 512 } else { 256 };
let bytes_required = frames_required * 2;
let frames_per_sec = frames_required as f32 / sample_rate as f32;
let start_frames = (params.start_secs / frames_per_sec).round() as usize;
let stop_frames = (params.stop_secs / frames_per_sec).round() as usize;
Self {
params,
frames_required,
bytes_required,
start_frames,
stop_frames,
starting_count: 0,
stopping_count: 0,
prev_volume: 0.0,
buffer: Vec::with_capacity(bytes_required * 2),
state: VadState::Quiet,
}
}
pub fn next_window(&mut self, chunk: &[u8]) -> Option<Vec<u8>> {
self.buffer.extend_from_slice(chunk);
if self.buffer.len() >= self.bytes_required {
let window: Vec<u8> = self.buffer.drain(..self.bytes_required).collect();
Some(window)
} else {
None
}
}
pub fn advance(&mut self, confidence: f32, audio_window: &[u8]) -> VadState {
let volume = exp_smoothing(
calculate_audio_volume(audio_window),
self.prev_volume,
SMOOTHING_FACTOR,
);
self.prev_volume = volume;
let speaking = confidence >= self.params.confidence
&& volume >= self.params.min_volume;
if speaking {
match self.state {
VadState::Quiet => {
self.state = VadState::Starting;
self.starting_count = 1;
}
VadState::Starting => {
self.starting_count += 1;
}
VadState::Stopping => {
self.state = VadState::Speaking;
self.stopping_count = 0;
}
VadState::Speaking => {}
}
} else {
match self.state {
VadState::Starting => {
self.state = VadState::Quiet;
self.starting_count = 0;
}
VadState::Speaking => {
self.state = VadState::Stopping;
self.stopping_count = 1;
}
VadState::Stopping => {
self.stopping_count += 1;
}
VadState::Quiet => {}
}
}
if self.state == VadState::Starting
&& self.starting_count >= self.start_frames
{
self.state = VadState::Speaking;
self.starting_count = 0;
}
if self.state == VadState::Stopping
&& self.stopping_count >= self.stop_frames
{
self.state = VadState::Quiet;
self.stopping_count = 0;
}
self.state
}
}
#[cfg(test)]
mod tests {
use super::*;
fn windows_when_drained(frame_samples: usize, n_frames: usize) -> usize {
let mut m = StateMachine::new(16_000, VadParams::default());
let chunk = vec![0u8; frame_samples * 2];
let mut windows = 0;
for _ in 0..n_frames {
let mut pending: &[u8] = &chunk;
while {
let w = m.next_window(pending);
pending = &[];
w
}
.is_some()
{
windows += 1;
}
}
windows
}
fn windows_when_called_once(frame_samples: usize, n_frames: usize) -> usize {
let mut m = StateMachine::new(16_000, VadParams::default());
let chunk = vec![0u8; frame_samples * 2];
(0..n_frames).filter(|_| m.next_window(&chunk).is_some()).count()
}
#[test]
fn drain_loop_keeps_up_with_wide_frames() {
for frame in [160usize, 320, 512, 960, 4800] {
let n = 200;
let analysed = windows_when_drained(frame, n) * 512;
let fed = frame * n;
assert!(
fed - analysed < 512,
"frame={frame}: fed {fed} samples, analysed {analysed} — \
a drain loop must leave less than one window unconsumed"
);
}
}
#[test]
fn single_call_per_frame_starves_on_wide_frames() {
let analysed = windows_when_called_once(960, 200) * 512;
let fed = 960 * 200;
assert!(
(analysed as f32) / (fed as f32) < 0.55,
"expected the old pattern to fall behind, got {analysed}/{fed}"
);
}
#[test]
fn drain_loop_is_a_no_op_for_narrow_frames() {
for frame in [160usize, 320, 512] {
assert_eq!(
windows_when_drained(frame, 200),
windows_when_called_once(frame, 200),
"frame={frame}: drain loop must be behaviour-identical here"
);
}
}
}