#[cfg(feature = "vad")]
use std::path::{Path, PathBuf};
#[cfg(feature = "vad")]
use std::sync::OnceLock;
use crate::error::DecibriError;
#[cfg(feature = "vad")]
use crate::onnx::{
OnnxInputs, OnnxSession, OnnxSessionBuilder, OnnxTensorData, OnnxTensorOwned, OnnxTensorView,
};
#[derive(Debug, Clone)]
pub struct VadConfig {
pub model_path: PathBuf,
pub sample_rate: u32,
pub threshold: f32,
pub ort_library_path: Option<PathBuf>,
}
impl Default for VadConfig {
fn default() -> Self {
Self {
model_path: PathBuf::from("silero_vad.onnx"),
sample_rate: 16000,
threshold: 0.5,
ort_library_path: None,
}
}
}
impl VadConfig {
pub fn validate(&self) -> Result<usize, DecibriError> {
let window_size = match self.sample_rate {
8000 => 256,
16000 => 512,
_ => return Err(DecibriError::VadSampleRateUnsupported(self.sample_rate)),
};
if !(0.0..=1.0).contains(&self.threshold) {
return Err(DecibriError::VadThresholdOutOfRange(self.threshold));
}
Ok(window_size)
}
}
#[derive(Debug, Clone)]
pub struct VadResult {
pub probability: f32,
pub is_speech: bool,
}
#[cfg(feature = "vad")]
pub struct SileroVad {
session: Box<dyn OnnxSession>,
state: Vec<f32>,
sample_rate: u32,
threshold: f32,
accumulator: Vec<f32>,
window_size: usize,
}
const STATE_SIZE: usize = 256;
#[cfg(feature = "vad")]
static ORT_INIT: OnceLock<Option<PathBuf>> = OnceLock::new();
#[cfg(feature = "vad")]
static ORT_INIT_PID: OnceLock<u32> = OnceLock::new();
#[cfg(feature = "vad")]
fn wrap_init_error(path: Option<&Path>, err: ort::Error) -> DecibriError {
match path {
Some(p) => DecibriError::OrtLoadFailed {
path: p.to_path_buf(),
source: err,
},
None => DecibriError::OrtInitFailed { source: err },
}
}
#[cfg(all(feature = "vad", feature = "ort-load-dynamic"))]
fn do_ort_init(path: Option<&Path>) -> Result<bool, ort::Error> {
match path {
Some(p) => ort::init_from(p).map(|b| b.with_name("decibri").commit()),
None => Ok(ort::init().with_name("decibri").commit()),
}
}
#[cfg(all(feature = "vad", not(feature = "ort-load-dynamic")))]
fn do_ort_init(_path: Option<&Path>) -> Result<bool, ort::Error> {
Ok(ort::init().with_name("decibri").commit())
}
#[cfg(feature = "vad")]
pub(crate) fn init_ort_once(path: Option<&Path>) -> Result<(), DecibriError> {
if ORT_INIT.get().is_some() {
return Ok(());
}
#[cfg(feature = "ort-load-dynamic")]
if let Some(p) = path {
if !p.is_file() {
return Err(DecibriError::OrtPathInvalid {
path: p.to_path_buf(),
reason: "path does not exist or is not a regular file",
});
}
}
match do_ort_init(path) {
Ok(_committed) => {
let _ = ORT_INIT.set(path.map(|p| p.to_path_buf()));
let _ = ORT_INIT_PID.set(std::process::id());
Ok(())
}
Err(e) => Err(wrap_init_error(path, e)),
}
}
#[cfg(feature = "vad")]
pub(crate) fn check_pid_for_ort() -> Result<(), DecibriError> {
if let Some(&init_pid) = ORT_INIT_PID.get() {
let current_pid = std::process::id();
if current_pid != init_pid {
return Err(DecibriError::ForkAfterOrtInit {
init_pid,
current_pid,
});
}
}
Ok(())
}
#[cfg(feature = "vad")]
impl SileroVad {
pub fn new(config: VadConfig) -> Result<Self, DecibriError> {
let window_size = config.validate()?;
init_ort_once(config.ort_library_path.as_deref())?;
let session = OnnxSessionBuilder::from_file(config.model_path.clone())
.with_intra_threads(1)
.build()?;
Ok(Self {
session,
state: vec![0.0f32; STATE_SIZE],
sample_rate: config.sample_rate,
threshold: config.threshold,
accumulator: Vec::new(),
window_size,
})
}
pub fn process(&mut self, samples: &[f32]) -> Result<VadResult, DecibriError> {
check_pid_for_ort()?;
self.accumulator.extend_from_slice(samples);
let mut max_probability: f32 = 0.0;
let mut windows_processed = 0;
while self.accumulator.len() >= self.window_size {
let window: Vec<f32> = self.accumulator.drain(..self.window_size).collect();
let probability = self.infer_window(&window)?;
max_probability = max_probability.max(probability);
windows_processed += 1;
}
if windows_processed == 0 {
return Ok(VadResult {
probability: 0.0,
is_speech: false,
});
}
Ok(VadResult {
probability: max_probability,
is_speech: max_probability >= self.threshold,
})
}
pub fn reset(&mut self) {
self.state.fill(0.0);
self.accumulator.clear();
}
fn infer_window(&mut self, window: &[f32]) -> Result<f32, DecibriError> {
let input_shape = [1i64, self.window_size as i64];
let state_shape = [2i64, 1i64, 128i64];
let sr_shape = [1i64];
let sr_data = [self.sample_rate as i64];
let outputs = self.session.run(OnnxInputs {
items: &[
(
"input",
OnnxTensorView {
shape: &input_shape,
data: OnnxTensorData::F32(window),
},
),
(
"state",
OnnxTensorView {
shape: &state_shape,
data: OnnxTensorData::F32(&self.state),
},
),
(
"sr",
OnnxTensorView {
shape: &sr_shape,
data: OnnxTensorData::I64(&sr_data),
},
),
],
})?;
let prob = outputs
.get("output")
.ok_or_else(|| DecibriError::OnnxBackendFailed {
backend: "ort",
source: "Silero output `output` missing from session run".into(),
})?;
let probability = match &prob.data {
OnnxTensorOwned::F32(v) => v[0],
OnnxTensorOwned::I64(_) => {
return Err(DecibriError::OnnxBackendFailed {
backend: "ort",
source: "Silero output `output` must be f32".into(),
});
}
};
let state_n = outputs
.get("stateN")
.ok_or_else(|| DecibriError::OnnxBackendFailed {
backend: "ort",
source: "Silero output `stateN` missing from session run".into(),
})?;
match &state_n.data {
OnnxTensorOwned::F32(v) => self.state.copy_from_slice(v),
OnnxTensorOwned::I64(_) => {
return Err(DecibriError::OnnxBackendFailed {
backend: "ort",
source: "Silero output `stateN` must be f32".into(),
});
}
}
Ok(probability)
}
}
#[cfg(all(test, feature = "vad"))]
mod tests {
use super::*;
use std::path::Path;
fn model_path() -> PathBuf {
let manifest_dir = env!("CARGO_MANIFEST_DIR");
Path::new(manifest_dir)
.join("..")
.join("..")
.join("models")
.join("silero_vad.onnx")
}
fn default_config() -> VadConfig {
VadConfig {
model_path: model_path(),
sample_rate: 16000,
threshold: 0.5,
ort_library_path: None,
}
}
#[test]
fn test_wrap_init_error_with_path() {
use std::env;
let bogus_path = env::temp_dir().join("does-not-exist-onnxruntime-xyz-test");
let err = wrap_init_error(
Some(bogus_path.as_path()),
ort::Error::new("simulated ort loader failure"),
);
let msg = err.to_string();
assert!(
msg.contains(&bogus_path.display().to_string()),
"error message should contain the attempted path, got: {msg}"
);
assert!(
msg.contains("If ORT_DYLIB_PATH is set"),
"error message should contain actionable guidance phrase, got: {msg}"
);
assert!(
msg.contains("simulated ort loader failure"),
"error message should include the underlying ort error, got: {msg}"
);
}
#[test]
fn test_wrap_init_error_without_path() {
let err = wrap_init_error(None, ort::Error::new("simulated ort init failure"));
let msg = err.to_string();
assert!(
msg.contains("ort_library_path"),
"None-path error should mention the VadConfig field, got: {msg}"
);
assert!(
msg.contains("ORT_DYLIB_PATH"),
"None-path error should mention the env var, got: {msg}"
);
assert!(
msg.contains("ort-download-binaries"),
"None-path error should mention the opt-out feature, got: {msg}"
);
assert!(
msg.contains("simulated ort init failure"),
"error message should include the underlying ort error, got: {msg}"
);
}
#[test]
fn test_vad_config_validation() {
let bad_rate = VadConfig {
sample_rate: 44100,
..default_config()
};
assert!(SileroVad::new(bad_rate).is_err());
}
#[test]
fn test_ort_error_source_chain_preserved() {
use std::error::Error;
let inner = ort::Error::new("simulated underlying ort error");
let err = wrap_init_error(None, inner);
assert!(
err.source().is_some(),
"OrtInitFailed should carry an ort::Error source"
);
let inner_with_path = ort::Error::new("another simulated error");
let path_err = wrap_init_error(Some(Path::new("/tmp/bogus")), inner_with_path);
assert!(
path_err.source().is_some(),
"OrtLoadFailed should carry an ort::Error source"
);
let path_invalid = DecibriError::OrtPathInvalid {
path: PathBuf::from("/tmp/nope"),
reason: "test",
};
assert!(
path_invalid.source().is_none(),
"OrtPathInvalid intentionally has no source (constructing ort::Error \
would trigger the hang the pre-check prevents)"
);
}
#[test]
fn test_is_ort_path_error() {
let load_failed = DecibriError::OrtLoadFailed {
path: PathBuf::from("/tmp/x"),
source: ort::Error::new("test"),
};
assert!(load_failed.is_ort_path_error());
let path_invalid = DecibriError::OrtPathInvalid {
path: PathBuf::from("/tmp/x"),
reason: "test",
};
assert!(path_invalid.is_ort_path_error());
let init_failed = DecibriError::OrtInitFailed {
source: ort::Error::new("test"),
};
assert!(!init_failed.is_ort_path_error());
let sample_rate = DecibriError::SampleRateOutOfRange;
assert!(!sample_rate.is_ort_path_error());
let vad_rate = DecibriError::VadSampleRateUnsupported(44_100);
assert!(!vad_rate.is_ort_path_error());
}
#[test]
fn test_vad_loads_model() {
let vad = SileroVad::new(default_config());
assert!(vad.is_ok(), "Model should load: {:?}", vad.err());
}
#[test]
fn test_vad_silence() {
let mut vad = SileroVad::new(default_config()).unwrap();
let silence = vec![0.0f32; 512];
let result = vad.process(&silence).unwrap();
assert!(
result.probability < 0.5,
"Silence probability should be low, got {}",
result.probability
);
}
#[test]
fn test_vad_state_persistence() {
let mut vad = SileroVad::new(default_config()).unwrap();
let initial_state = vad.state.clone();
let samples = vec![0.0f32; 512];
vad.process(&samples).unwrap();
assert_ne!(
vad.state, initial_state,
"State should change after inference"
);
}
#[test]
fn test_vad_accumulator_windows() {
let mut vad = SileroVad::new(default_config()).unwrap();
let samples = vec![0.0f32; 1600];
let result = vad.process(&samples).unwrap();
assert!(result.probability >= 0.0);
assert_eq!(vad.accumulator.len(), 64);
}
#[test]
fn test_vad_accumulator_carry() {
let mut vad = SileroVad::new(default_config()).unwrap();
let chunk1 = vec![0.0f32; 1600];
vad.process(&chunk1).unwrap();
assert_eq!(vad.accumulator.len(), 64);
let chunk2 = vec![0.0f32; 1600];
vad.process(&chunk2).unwrap();
assert_eq!(vad.accumulator.len(), 128);
}
#[test]
fn test_vad_small_chunk() {
let mut vad = SileroVad::new(default_config()).unwrap();
let samples = vec![0.0f32; 100];
let result = vad.process(&samples).unwrap();
assert_eq!(result.probability, 0.0);
assert_eq!(vad.accumulator.len(), 100);
}
#[test]
fn test_vad_reset() {
let mut vad = SileroVad::new(default_config()).unwrap();
let samples = vec![0.0f32; 512];
vad.process(&samples).unwrap();
vad.accumulator.extend_from_slice(&[0.0; 100]); vad.reset();
assert!(vad.state.iter().all(|&v| v == 0.0), "State should be zeros");
assert!(vad.accumulator.is_empty(), "Accumulator should be empty");
}
}