use std::path::{Path, PathBuf};
use std::sync::OnceLock;
use crate::error::DecibriError;
#[allow(dead_code)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
#[non_exhaustive]
pub(crate) enum ExecutionProvider {
#[default]
Cpu,
CoreMl,
Cuda,
DirectMl,
Rocm,
}
static ORT_INIT: OnceLock<Option<PathBuf>> = OnceLock::new();
static ORT_INIT_PID: OnceLock<u32> = OnceLock::new();
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(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(not(feature = "ort-load-dynamic"))]
fn do_ort_init(_path: Option<&Path>) -> Result<bool, ort::Error> {
Ok(ort::init().with_name("decibri").commit())
}
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)),
}
}
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(())
}
pub(crate) struct OnnxTensorView<'a> {
pub shape: &'a [i64],
pub data: OnnxTensorData<'a>,
}
pub(crate) enum OnnxTensorData<'a> {
F32(&'a [f32]),
I64(&'a [i64]),
}
pub(crate) struct OnnxInputs<'a> {
pub items: &'a [(&'a str, OnnxTensorView<'a>)],
}
pub(crate) struct OnnxOutputs {
pub tensors: Vec<(String, OnnxOutputTensor)>,
}
pub(crate) struct OnnxOutputTensor {
#[allow(dead_code)]
pub shape: Vec<i64>,
pub data: OnnxTensorOwned,
}
pub(crate) enum OnnxTensorOwned {
F32(Vec<f32>),
#[allow(dead_code)]
I64(Vec<i64>),
}
impl OnnxOutputs {
pub(crate) fn get(&self, name: &str) -> Option<&OnnxOutputTensor> {
self.tensors.iter().find(|(n, _)| n == name).map(|(_, t)| t)
}
}
pub(crate) trait OnnxSession: Send + Sync {
fn run(&mut self, inputs: OnnxInputs<'_>) -> Result<OnnxOutputs, DecibriError>;
}
pub(crate) struct OnnxSessionBuilder {
model_path: PathBuf,
intra_threads: usize,
execution_provider: ExecutionProvider,
}
impl OnnxSessionBuilder {
pub(crate) fn from_file(path: impl Into<PathBuf>) -> Self {
Self {
model_path: path.into(),
intra_threads: 1,
execution_provider: ExecutionProvider::Cpu,
}
}
pub(crate) fn with_intra_threads(mut self, n: usize) -> Self {
self.intra_threads = n;
self
}
#[allow(dead_code)]
pub(crate) fn with_execution_provider(mut self, ep: ExecutionProvider) -> Self {
self.execution_provider = ep;
self
}
pub(crate) fn build(self) -> Result<Box<dyn OnnxSession>, DecibriError> {
let session = ort_impl::OrtSession::open(
&self.model_path,
self.intra_threads,
self.execution_provider,
)?;
Ok(Box::new(session))
}
}
mod ort_impl {
use std::borrow::Cow;
use std::path::Path;
use ort::session::{Session, SessionInputValue};
use ort::value::{Tensor, TensorElementType};
use super::{
DecibriError, ExecutionProvider, OnnxInputs, OnnxOutputTensor, OnnxOutputs, OnnxSession,
OnnxTensorData, OnnxTensorOwned,
};
type Epd = ort::execution_providers::ExecutionProviderDispatch;
macro_rules! accelerator_dispatch {
($name:ident, $feature:literal, $ty:ident, $display:literal) => {
#[cfg(feature = $feature)]
fn $name() -> Result<Epd, DecibriError> {
Ok(ort::execution_providers::$ty::default().build())
}
#[cfg(not(feature = $feature))]
fn $name() -> Result<Epd, DecibriError> {
Err(DecibriError::OnnxBackendFailed {
backend: $display,
source: concat!(
"execution provider not available in this build; rebuild decibri \
with the `",
$feature,
"` feature enabled"
)
.into(),
})
}
};
}
accelerator_dispatch!(coreml_dispatch, "coreml", CoreML, "CoreML");
accelerator_dispatch!(cuda_dispatch, "cuda", CUDA, "CUDA");
accelerator_dispatch!(directml_dispatch, "directml", DirectML, "DirectML");
accelerator_dispatch!(rocm_dispatch, "rocm", ROCm, "ROCm");
fn provider_dispatch(ep: ExecutionProvider) -> Result<Option<Vec<Epd>>, DecibriError> {
let accelerator = match ep {
ExecutionProvider::Cpu => return Ok(None),
ExecutionProvider::CoreMl => coreml_dispatch()?,
ExecutionProvider::Cuda => cuda_dispatch()?,
ExecutionProvider::DirectMl => directml_dispatch()?,
ExecutionProvider::Rocm => rocm_dispatch()?,
};
let cpu = ort::execution_providers::CPU::default()
.with_arena_allocator(true)
.build();
Ok(Some(vec![accelerator, cpu]))
}
pub(super) struct OrtSession {
inner: Session,
}
impl OrtSession {
pub(super) fn open(
path: &Path,
intra: usize,
execution_provider: ExecutionProvider,
) -> Result<Self, DecibriError> {
let providers = provider_dispatch(execution_provider)?;
let builder = Session::builder()
.map_err(DecibriError::OrtSessionBuildFailed)?
.with_intra_threads(intra)
.map_err(|e| DecibriError::OrtThreadsConfigFailed(e.into()))?;
let mut builder = match providers {
Some(providers) => builder
.with_execution_providers(providers)
.map_err(|e| DecibriError::OrtSessionBuildFailed(e.into()))?,
None => builder,
};
let session =
builder
.commit_from_file(path)
.map_err(|e| DecibriError::VadModelLoadFailed {
path: path.to_path_buf(),
source: e,
})?;
Ok(Self { inner: session })
}
}
fn known_kind(name: &str) -> &'static str {
match name {
"input" => "input",
"state" => "state",
"sr" => "sr",
"output" => "output",
"stateN" => "state",
_ => "trait_tensor",
}
}
impl OnnxSession for OrtSession {
fn run(&mut self, inputs: OnnxInputs<'_>) -> Result<OnnxOutputs, DecibriError> {
let mut input_pairs: Vec<(Cow<'_, str>, SessionInputValue<'_>)> =
Vec::with_capacity(inputs.items.len());
for (name, view) in inputs.items.iter() {
let kind = known_kind(name);
let shape: Vec<i64> = view.shape.to_vec();
let value: SessionInputValue<'_> = match &view.data {
OnnxTensorData::F32(slice) => {
let tensor = Tensor::from_array((shape, slice.to_vec()))
.map_err(|e| DecibriError::OrtTensorCreateFailed { kind, source: e })?;
tensor.into()
}
OnnxTensorData::I64(slice) => {
let tensor = Tensor::from_array((shape, slice.to_vec()))
.map_err(|e| DecibriError::OrtTensorCreateFailed { kind, source: e })?;
tensor.into()
}
};
input_pairs.push((Cow::Owned((*name).to_string()), value));
}
let outputs = self
.inner
.run(input_pairs)
.map_err(DecibriError::OrtInferenceFailed)?;
let names: Vec<String> = outputs.keys().map(|k| k.to_string()).collect();
let mut tensors: Vec<(String, OnnxOutputTensor)> = Vec::with_capacity(names.len());
for name in names {
let kind = known_kind(&name);
let value = &outputs[name.as_str()];
let dtype = value.dtype().tensor_type();
let owned = match dtype {
Some(TensorElementType::Float32) => {
let (shape, data) = value.try_extract_tensor::<f32>().map_err(|e| {
DecibriError::OrtTensorExtractFailed { kind, source: e }
})?;
OnnxOutputTensor {
shape: shape.to_vec(),
data: OnnxTensorOwned::F32(data.to_vec()),
}
}
Some(TensorElementType::Int64) => {
let (shape, data) = value.try_extract_tensor::<i64>().map_err(|e| {
DecibriError::OrtTensorExtractFailed { kind, source: e }
})?;
OnnxOutputTensor {
shape: shape.to_vec(),
data: OnnxTensorOwned::I64(data.to_vec()),
}
}
_ => {
return Err(DecibriError::OnnxBackendFailed {
backend: "ort",
source: format!(
"unsupported output tensor element type for {name}: {dtype:?}"
)
.into(),
});
}
};
tensors.push((name, owned));
}
Ok(OnnxOutputs { tensors })
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[allow(dead_code)]
fn assert_send_sync<T: Send + Sync>() {}
#[test]
fn box_dyn_session_is_send_sync() {
assert_send_sync::<Box<dyn OnnxSession>>();
}
#[test]
fn execution_provider_default_is_cpu() {
assert_eq!(ExecutionProvider::default(), ExecutionProvider::Cpu);
let all = [
ExecutionProvider::Cpu,
ExecutionProvider::CoreMl,
ExecutionProvider::Cuda,
ExecutionProvider::DirectMl,
ExecutionProvider::Rocm,
];
for (i, a) in all.iter().enumerate() {
for b in &all[i + 1..] {
assert_ne!(a, b, "execution provider variants must be distinct");
}
}
}
#[cfg(not(feature = "cuda"))]
#[test]
fn builder_rejects_unavailable_execution_provider() {
let result = OnnxSessionBuilder::from_file("does-not-need-to-exist.onnx")
.with_execution_provider(ExecutionProvider::Cuda)
.build();
match result {
Err(DecibriError::OnnxBackendFailed { backend, source }) => {
assert_eq!(backend, "CUDA");
let msg = source.to_string();
assert!(
msg.contains("cuda"),
"error should name the `cuda` feature to enable, got: {msg}"
);
}
Err(other) => panic!("expected OnnxBackendFailed, got error: {other:?}"),
Ok(_) => panic!("expected OnnxBackendFailed, got a built session"),
}
}
#[test]
fn wrap_init_error_with_path_is_load_failed() {
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 wrap_init_error_without_path_is_init_failed() {
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 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)"
);
}
struct MockSession {
expected_input_count: usize,
canned_outputs: Vec<(String, OnnxOutputTensor)>,
}
impl OnnxSession for MockSession {
fn run(&mut self, inputs: OnnxInputs<'_>) -> Result<OnnxOutputs, DecibriError> {
assert_eq!(
inputs.items.len(),
self.expected_input_count,
"MockSession got unexpected number of inputs"
);
let mut tensors = Vec::with_capacity(self.canned_outputs.len());
for (name, t) in self.canned_outputs.iter() {
let cloned = OnnxOutputTensor {
shape: t.shape.clone(),
data: match &t.data {
OnnxTensorOwned::F32(v) => OnnxTensorOwned::F32(v.clone()),
OnnxTensorOwned::I64(v) => OnnxTensorOwned::I64(v.clone()),
},
};
tensors.push((name.clone(), cloned));
}
Ok(OnnxOutputs { tensors })
}
}
#[test]
fn mock_session_round_trip() {
let canned = vec![(
"output".to_string(),
OnnxOutputTensor {
shape: vec![1, 1],
data: OnnxTensorOwned::F32(vec![0.42]),
},
)];
let mut session: Box<dyn OnnxSession> = Box::new(MockSession {
expected_input_count: 3,
canned_outputs: canned,
});
let audio = vec![0.0f32; 512];
let state = vec![0.0f32; 256];
let sr = vec![16000i64];
let inputs = OnnxInputs {
items: &[
(
"input",
OnnxTensorView {
shape: &[1, 512],
data: OnnxTensorData::F32(&audio),
},
),
(
"state",
OnnxTensorView {
shape: &[2, 1, 128],
data: OnnxTensorData::F32(&state),
},
),
(
"sr",
OnnxTensorView {
shape: &[1],
data: OnnxTensorData::I64(&sr),
},
),
],
};
let outputs = session.run(inputs).expect("MockSession should succeed");
assert_eq!(outputs.tensors.len(), 1);
let probe = outputs.get("output").expect("output present");
assert_eq!(probe.shape, vec![1, 1]);
match &probe.data {
OnnxTensorOwned::F32(v) => assert_eq!(v.as_slice(), &[0.42f32]),
OnnxTensorOwned::I64(_) => panic!("expected F32"),
}
}
#[test]
fn output_lookup_by_name_returns_none_for_missing() {
let outputs = OnnxOutputs {
tensors: vec![(
"output".to_string(),
OnnxOutputTensor {
shape: vec![1],
data: OnnxTensorOwned::F32(vec![1.0]),
},
)],
};
assert!(outputs.get("output").is_some());
assert!(outputs.get("not_a_name").is_none());
}
#[test]
fn input_slice_preserves_order_and_names() {
let a = vec![1.0f32, 2.0, 3.0];
let b = vec![10i64, 20];
let inputs = OnnxInputs {
items: &[
(
"first",
OnnxTensorView {
shape: &[3],
data: OnnxTensorData::F32(&a),
},
),
(
"second",
OnnxTensorView {
shape: &[2],
data: OnnxTensorData::I64(&b),
},
),
],
};
assert_eq!(inputs.items.len(), 2);
assert_eq!(inputs.items[0].0, "first");
assert_eq!(inputs.items[1].0, "second");
match inputs.items[0].1.data {
OnnxTensorData::F32(s) => assert_eq!(s, a.as_slice()),
OnnxTensorData::I64(_) => panic!("expected F32"),
}
match inputs.items[1].1.data {
OnnxTensorData::I64(s) => assert_eq!(s, b.as_slice()),
OnnxTensorData::F32(_) => panic!("expected I64"),
}
}
#[test]
fn ort_backed_session_runs_silero_inference() {
use std::path::Path;
let manifest_dir = env!("CARGO_MANIFEST_DIR");
let model = Path::new(manifest_dir)
.join("..")
.join("..")
.join("models")
.join("silero_vad.onnx");
if !model.is_file() {
eprintln!(
"skipping ort_backed_session_runs_silero_inference: model not found at {}",
model.display()
);
return;
}
super::init_ort_once(None).expect("ORT init should succeed");
let mut session = OnnxSessionBuilder::from_file(&model)
.with_intra_threads(1)
.build()
.expect("ORT session build should succeed for bundled model");
let audio = vec![0.0f32; 512];
let state = vec![0.0f32; 256];
let sr = vec![16000i64];
let outputs = session
.run(OnnxInputs {
items: &[
(
"input",
OnnxTensorView {
shape: &[1, 512],
data: OnnxTensorData::F32(&audio),
},
),
(
"state",
OnnxTensorView {
shape: &[2, 1, 128],
data: OnnxTensorData::F32(&state),
},
),
(
"sr",
OnnxTensorView {
shape: &[1],
data: OnnxTensorData::I64(&sr),
},
),
],
})
.expect("ORT session run should succeed");
let probe = outputs
.get("output")
.expect("Silero emits an `output` tensor");
match &probe.data {
OnnxTensorOwned::F32(v) => assert_eq!(v.len(), 1, "Silero `output` is one f32"),
OnnxTensorOwned::I64(_) => panic!("Silero `output` is f32, not i64"),
}
let state_n = outputs
.get("stateN")
.expect("Silero emits a `stateN` tensor");
match &state_n.data {
OnnxTensorOwned::F32(v) => assert_eq!(v.len(), 256, "Silero `stateN` is 256 f32s"),
OnnxTensorOwned::I64(_) => panic!("Silero `stateN` is f32, not i64"),
}
}
}