#[cfg(feature = "vad")]
use std::path::{Path, PathBuf};
#[cfg(feature = "vad")]
use std::sync::OnceLock;
use crate::error::DecibriError;
#[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: ort::session::Session,
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")]
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")]
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()));
Ok(())
}
Err(e) => Err(wrap_init_error(path, e)),
}
}
#[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 = ort::session::Session::builder()
.map_err(DecibriError::OrtSessionBuildFailed)?
.with_intra_threads(1)
.map_err(|e| DecibriError::OrtThreadsConfigFailed(e.into()))?
.commit_from_file(&config.model_path)
.map_err(|e| DecibriError::VadModelLoadFailed {
path: config.model_path.clone(),
source: e,
})?;
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> {
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_tensor =
ort::value::Tensor::from_array(([1i64, self.window_size as i64], window.to_vec()))
.map_err(|e| DecibriError::OrtTensorCreateFailed {
kind: "input",
source: e,
})?;
let state_tensor =
ort::value::Tensor::from_array(([2i64, 1i64, 128i64], self.state.clone())).map_err(
|e| DecibriError::OrtTensorCreateFailed {
kind: "state",
source: e,
},
)?;
let sr_tensor = ort::value::Tensor::from_array(([1i64], vec![self.sample_rate as i64]))
.map_err(|e| DecibriError::OrtTensorCreateFailed {
kind: "sr",
source: e,
})?;
let input_values = ort::inputs![
"input" => input_tensor,
"state" => state_tensor,
"sr" => sr_tensor,
];
let outputs = self
.session
.run(input_values)
.map_err(DecibriError::OrtInferenceFailed)?;
let prob_tensor = outputs["output"].try_extract_tensor::<f32>().map_err(|e| {
DecibriError::OrtTensorExtractFailed {
kind: "output",
source: e,
}
})?;
let probability = prob_tensor.1[0];
let state_tensor = outputs["stateN"].try_extract_tensor::<f32>().map_err(|e| {
DecibriError::OrtTensorExtractFailed {
kind: "state",
source: e,
}
})?;
self.state.copy_from_slice(state_tensor.1);
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");
}
}