use std::path::Path;
#[cfg(not(any(feature = "onnx", feature = "backend-tract")))]
compile_error!("feature `infer` requires `onnx` (ort) and/or `backend-tract`");
mod factory;
#[cfg(feature = "onnx")]
mod ort_session;
#[cfg(all(test, feature = "backend-tract", feature = "onnx"))]
mod parity;
mod runtime;
#[cfg(feature = "backend-tract")]
mod tract_session;
pub use factory::{InferenceBackend, RuntimeSession};
#[cfg(feature = "onnx")]
pub use ort_session::OrtSession;
pub use runtime::{InferenceError, InferenceRuntime, InferenceTensor, NamedTensor, TensorData};
#[cfg(feature = "backend-tract")]
pub use tract_session::TractSession;
pub const ONNX_MIN_HEADER_BYTES: usize = 64;
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum ExecutionProvider {
Cpu,
CoreMl,
Nnapi,
Cuda,
XnnPack,
}
impl ExecutionProvider {
pub fn auto() -> Self {
#[cfg(all(target_os = "macos", target_arch = "aarch64"))]
return Self::CoreMl;
#[cfg(all(target_os = "linux", target_arch = "aarch64"))]
return Self::XnnPack;
#[cfg(not(any(
all(target_os = "macos", target_arch = "aarch64"),
all(target_os = "linux", target_arch = "aarch64"),
)))]
return Self::Cpu;
}
pub fn is_available(self) -> bool {
match self {
Self::Cpu => true,
Self::CoreMl => {
cfg!(all(
feature = "coreml",
target_os = "macos",
target_arch = "aarch64"
))
}
Self::XnnPack => cfg!(feature = "xnnpack"),
Self::Nnapi | Self::Cuda => false,
}
}
}
pub fn build_session_with_ep(
model_path: &Path,
ep: ExecutionProvider,
intra_threads: Option<usize>,
) -> Result<RuntimeSession, OnnxError> {
RuntimeSession::from_path(model_path, ep, intra_threads)
}
pub fn resolve_session_pool_size(configured: usize) -> usize {
std::env::var("POLYVOICE_SESSION_POOL_SIZE")
.ok()
.and_then(|s| s.parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(configured.max(1))
.max(1)
}
pub fn resolve_intra_threads(pool_size: usize) -> usize {
if let Some(n) = std::env::var("POLYVOICE_INTRA_THREADS")
.ok()
.and_then(|s| s.parse::<usize>().ok())
.filter(|&n| n > 0)
{
return n;
}
let pool = pool_size.max(1);
std::thread::available_parallelism()
.map(|n| (n.get() / pool).max(1))
.unwrap_or(1)
}
pub fn read_model_metadata_props(
path: &Path,
) -> Result<std::collections::HashMap<String, String>, OnnxError> {
validate_onnx_header(path)?;
#[cfg(feature = "onnx")]
{
let session = OrtSession::from_path(path, ExecutionProvider::Cpu, Some(1))?;
session.custom_metadata_props()
}
#[cfg(not(feature = "onnx"))]
{
Ok(std::collections::HashMap::new())
}
}
#[derive(Clone, thiserror::Error, Debug)]
pub enum OnnxError {
#[error(transparent)]
Validation(#[from] OnnxValidationError),
#[error("failed to build inference session for {path}: {detail}")]
SessionBuild {
path: std::path::PathBuf,
detail: String,
},
#[error("failed to read ONNX metadata_props: {detail}")]
Metadata { detail: String },
}
#[derive(Clone, thiserror::Error, Debug)]
#[error("ONNX header validation failed for {path}: {detail}")]
pub struct OnnxValidationError {
pub path: std::path::PathBuf,
pub detail: String,
}
pub fn validate_onnx_header(path: &Path) -> Result<(), OnnxValidationError> {
let metadata = std::fs::metadata(path).map_err(|e| OnnxValidationError {
path: path.to_path_buf(),
detail: format!("cannot read metadata: {e}"),
})?;
if metadata.len() < ONNX_MIN_HEADER_BYTES as u64 {
return Err(OnnxValidationError {
path: path.to_path_buf(),
detail: format!(
"file too small ({} bytes, need at least {ONNX_MIN_HEADER_BYTES})",
metadata.len()
),
});
}
let mut file = std::fs::File::open(path).map_err(|e| OnnxValidationError {
path: path.to_path_buf(),
detail: format!("cannot open file: {e}"),
})?;
let mut header = [0u8; ONNX_MIN_HEADER_BYTES];
let n = std::io::Read::read(&mut file, &mut header).map_err(|e| OnnxValidationError {
path: path.to_path_buf(),
detail: format!("cannot read header: {e}"),
})?;
if n < ONNX_MIN_HEADER_BYTES {
return Err(OnnxValidationError {
path: path.to_path_buf(),
detail: format!("short read ({n} bytes, need at least {ONNX_MIN_HEADER_BYTES})"),
});
}
let has_onnx_magic = header[..16].windows(4).any(|w| w == b"ONNX");
let has_protobuf_header = header[0] == 0x08;
if !has_onnx_magic && !has_protobuf_header {
return Err(OnnxValidationError {
path: path.to_path_buf(),
detail: "ONNX magic bytes not found and file does not start with a valid ONNX protobuf header".to_string(),
});
}
Ok(())
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod tests {
use super::*;
use std::io::Write;
#[test]
#[cfg_attr(miri, ignore)]
fn valid_onnx_file_passes_validation() {
let path = std::path::Path::new("models/silero_vad.onnx");
if !path.exists() {
return;
}
assert!(validate_onnx_header(path).is_ok());
}
#[test]
#[cfg_attr(miri, ignore)]
fn random_64_bytes_fails_validation() {
let mut tmp = tempfile::NamedTempFile::new().unwrap();
tmp.write_all(&[0xAB; 64]).unwrap();
let result = validate_onnx_header(tmp.path());
assert!(result.is_err());
let err = result.unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("ONNX magic") || msg.contains("protobuf header"),
"unexpected error message: {msg}"
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn empty_file_fails_validation() {
let tmp = tempfile::NamedTempFile::new().unwrap();
let result = validate_onnx_header(tmp.path());
assert!(result.is_err());
let err = result.unwrap_err();
assert!(
err.to_string().contains("too small"),
"unexpected error: {err}"
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn file_with_onnx_magic_passes() {
let mut tmp = tempfile::NamedTempFile::new().unwrap();
let mut data = vec![0u8; 64];
data[4..8].copy_from_slice(b"ONNX");
tmp.write_all(&data).unwrap();
assert!(validate_onnx_header(tmp.path()).is_ok());
}
#[test]
#[cfg_attr(miri, ignore)]
fn file_with_protobuf_header_passes() {
let mut tmp = tempfile::NamedTempFile::new().unwrap();
let mut data = vec![0u8; 64];
data[0] = 0x08; data[1] = 0x08; tmp.write_all(&data).unwrap();
assert!(validate_onnx_header(tmp.path()).is_ok());
}
#[test]
#[cfg_attr(miri, ignore)]
fn build_session_with_ep_rejects_garbage_before_ort() {
let mut tmp = tempfile::NamedTempFile::new().unwrap();
tmp.write_all(&[0xAB; 64]).unwrap();
let err = build_session_with_ep(tmp.path(), ExecutionProvider::Cpu, None)
.expect_err("garbage must fail header validation");
assert!(err.to_string().contains("ONNX header validation failed"));
}
#[test]
#[cfg(feature = "onnx")]
#[cfg_attr(miri, ignore)]
fn build_session_with_ep_cpu_and_unwired_ep_build_ok() {
let path = std::path::Path::new("models/silero_vad.onnx");
if !path.exists() {
return;
}
InferenceBackend::force(Some(InferenceBackend::Ort));
let built = build_session_with_ep(path, ExecutionProvider::Cpu, None);
assert!(
built.is_ok(),
"ort session build failed: {:?}",
built.err().map(|e| e.to_string())
);
assert!(build_session_with_ep(path, ExecutionProvider::Cpu, Some(1)).is_ok());
assert!(build_session_with_ep(path, ExecutionProvider::Cuda, None).is_ok());
assert!(build_session_with_ep(path, ExecutionProvider::Nnapi, None).is_ok());
assert!(build_session_with_ep(path, ExecutionProvider::auto(), None).is_ok());
InferenceBackend::force(None);
}
#[test]
#[cfg(feature = "onnx")]
#[cfg_attr(miri, ignore)]
fn build_session_with_ep_optional_providers_build_ok() {
let path = std::path::Path::new("models/silero_vad.onnx");
if !path.exists() {
return;
}
InferenceBackend::force(Some(InferenceBackend::Ort));
assert!(build_session_with_ep(path, ExecutionProvider::CoreMl, None).is_ok());
assert!(build_session_with_ep(path, ExecutionProvider::XnnPack, None).is_ok());
InferenceBackend::force(None);
}
#[test]
fn resolve_session_pool_size_is_at_least_one() {
assert!(resolve_session_pool_size(0) >= 1);
assert!(resolve_session_pool_size(4) >= 1);
}
#[test]
fn resolve_intra_threads_is_at_least_one() {
assert!(resolve_intra_threads(1) >= 1);
assert!(resolve_intra_threads(4) >= 1);
}
#[test]
fn execution_provider_auto_matches_platform() {
let auto = ExecutionProvider::auto();
#[cfg(all(target_os = "macos", target_arch = "aarch64"))]
assert_eq!(auto, ExecutionProvider::CoreMl);
#[cfg(all(target_os = "linux", target_arch = "aarch64"))]
assert_eq!(auto, ExecutionProvider::XnnPack);
#[cfg(not(any(
all(target_os = "macos", target_arch = "aarch64"),
all(target_os = "linux", target_arch = "aarch64"),
)))]
assert_eq!(auto, ExecutionProvider::Cpu);
let copied = auto;
assert_eq!(copied, auto);
assert!(!format!("{auto:?}").is_empty());
}
#[test]
fn execution_provider_is_available_matches_wiring() {
assert!(ExecutionProvider::Cpu.is_available());
assert!(!ExecutionProvider::Nnapi.is_available());
assert!(!ExecutionProvider::Cuda.is_available());
assert_eq!(
ExecutionProvider::CoreMl.is_available(),
cfg!(all(
feature = "coreml",
target_os = "macos",
target_arch = "aarch64"
))
);
assert_eq!(
ExecutionProvider::XnnPack.is_available(),
cfg!(feature = "xnnpack")
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn read_model_metadata_props_real_model() {
let path = std::path::Path::new("models/silero_vad.onnx");
if !path.exists() {
return;
}
let props = read_model_metadata_props(path).unwrap();
assert!(props.keys().all(|k| !k.is_empty()));
}
#[test]
#[cfg_attr(miri, ignore)]
fn read_model_metadata_props_rejects_missing_file() {
let err =
read_model_metadata_props(std::path::Path::new("models/definitely_not_a_model.onnx"))
.expect_err("missing file must fail validation");
assert!(
matches!(err, OnnxError::Validation(_)),
"unexpected error: {err}"
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn read_model_metadata_props_rejects_garbage() {
let mut tmp = tempfile::NamedTempFile::new().unwrap();
tmp.write_all(&[0xAB; 64]).unwrap();
let err =
read_model_metadata_props(tmp.path()).expect_err("garbage must fail header validation");
assert!(
matches!(err, OnnxError::Validation(_)),
"unexpected error: {err}"
);
}
#[test]
#[cfg_attr(miri, ignore)]
fn validate_onnx_header_missing_file() {
let err = validate_onnx_header(std::path::Path::new("models/no_such_file.onnx"))
.expect_err("missing file must fail");
assert!(err.detail.contains("cannot read metadata"));
assert!(err.to_string().contains("no_such_file.onnx"));
}
#[test]
fn onnx_error_display_variants() {
let build = OnnxError::SessionBuild {
path: std::path::PathBuf::from("m.onnx"),
detail: "parse failed".to_string(),
};
assert_eq!(
build.to_string(),
"failed to build inference session for m.onnx: parse failed"
);
let meta = OnnxError::Metadata {
detail: "no meta".to_string(),
};
assert_eq!(
meta.to_string(),
"failed to read ONNX metadata_props: no meta"
);
let validation = OnnxError::Validation(OnnxValidationError {
path: std::path::PathBuf::from("bad.onnx"),
detail: "too small".to_string(),
});
assert_eq!(
validation.to_string(),
"ONNX header validation failed for bad.onnx: too small"
);
let cloned = validation.clone();
assert_eq!(cloned.to_string(), validation.to_string());
}
}