use std::collections::BTreeMap;
use std::fmt;
use std::sync::Arc;
use dataflow_rs::datavalue::OwnedDataTensor;
use super::manifest::Manifest;
use crate::config::ModelsConfig;
pub mod tract;
pub use self::tract::TractRuntime;
pub const NAMES: &[&str] = &["tract"];
pub const KNOWN_FORMATS: &[&str] = &["onnx"];
const TRACT_DEVICES: &[&str] = &["cpu", "metal", "cuda"];
const TRACT_FORMATS: &[&str] = &["onnx"];
pub fn devices_of(name: &str) -> Option<&'static [&'static str]> {
match name {
"tract" => Some(TRACT_DEVICES),
_ => None,
}
}
pub fn formats_of(name: &str) -> Option<&'static [&'static str]> {
match name {
"tract" => Some(TRACT_FORMATS),
_ => None,
}
}
pub fn intern(name: &str) -> Option<&'static str> {
NAMES.iter().copied().find(|n| *n == name)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LoadError {
pub stage: &'static str,
pub message: String,
}
impl LoadError {
pub fn new(stage: &'static str, message: impl Into<String>) -> Self {
Self {
stage,
message: message.into(),
}
}
}
impl fmt::Display for LoadError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {}", self.stage, self.message)
}
}
impl std::error::Error for LoadError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RunError {
pub stage: &'static str,
pub message: String,
}
impl RunError {
pub fn new(stage: &'static str, message: impl Into<String>) -> Self {
Self {
stage,
message: message.into(),
}
}
}
impl fmt::Display for RunError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}: {}", self.stage, self.message)
}
}
impl std::error::Error for RunError {}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RuntimeSelection {
NoDefault { format: String },
UnknownRuntime { runtime: String, format: String },
NotEnabled { runtime: String, format: String },
DoesNotServe {
runtime: String,
format: String,
serves: &'static [&'static str],
},
}
impl fmt::Display for RuntimeSelection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::NoDefault { format } => write!(
f,
"no default runtime for format '{format}' (models.default_runtime.{format} is \
not set)"
),
Self::UnknownRuntime { runtime, format } => write!(
f,
"runtime '{runtime}' for format '{format}' is unknown: this build lists {}",
NAMES.join(", ")
),
Self::NotEnabled { runtime, format } => write!(
f,
"runtime '{runtime}' for format '{format}' is not enabled on this node \
(models.runtimes.{runtime}.enabled)"
),
Self::DoesNotServe {
runtime,
format,
serves,
} => write!(
f,
"runtime '{runtime}' does not serve format '{format}'; it serves {}",
serves
.iter()
.map(|s| format!("'{s}'"))
.collect::<Vec<_>>()
.join(", ")
),
}
}
}
impl std::error::Error for RuntimeSelection {}
pub trait ModelRuntime: Send + Sync {
fn name(&self) -> &'static str;
fn devices(&self) -> &'static [&'static str];
fn formats(&self) -> &'static [&'static str];
fn load(
&self,
bytes: &[u8],
manifest: &Manifest,
device: &str,
) -> Result<Arc<dyn LoadedModel>, LoadError>;
}
pub trait LoadedModel: Send + Sync {
fn digest(&self) -> &str;
fn resident_bytes(&self) -> usize;
fn run(&self, inputs: Vec<OwnedDataTensor>) -> Result<Vec<OwnedDataTensor>, RunError>;
}
#[derive(Default)]
pub struct ModelRuntimes {
by_name: BTreeMap<&'static str, Arc<dyn ModelRuntime>>,
}
impl ModelRuntimes {
pub fn empty() -> Self {
Self::default()
}
pub fn builtin(config: &ModelsConfig) -> Self {
let mut registry = Self::empty();
if config
.runtimes
.get(self::tract::NAME)
.is_some_and(|runtime| runtime.enabled)
{
registry
.by_name
.insert(self::tract::NAME, Arc::new(TractRuntime));
}
registry
}
pub fn register(&mut self, runtime: Arc<dyn ModelRuntime>) -> Result<(), String> {
let name = runtime.name();
if !NAMES.contains(&name) {
return Err(format!(
"runtime '{name}' is not one this build lists ({}); add it to \
model::runtimes::NAMES before registering it",
NAMES.join(", ")
));
}
if runtime.formats().is_empty() {
return Err(format!(
"runtime '{name}' serves no format, so nothing could ever select it"
));
}
if self.by_name.contains_key(name) {
return Err(format!("runtime '{name}' is registered twice"));
}
self.by_name.insert(name, runtime);
Ok(())
}
pub fn get(&self, name: &str) -> Option<Arc<dyn ModelRuntime>> {
self.by_name.get(name).cloned()
}
pub fn default_for<'c>(
&self,
config: &'c ModelsConfig,
format: &str,
) -> Result<(Arc<dyn ModelRuntime>, &'c str), RuntimeSelection> {
let Some(name) = config.default_runtime_for(format) else {
return Err(RuntimeSelection::NoDefault {
format: format.to_string(),
});
};
self.for_format(config, name, format)
}
pub fn for_format<'c>(
&self,
config: &'c ModelsConfig,
name: &str,
format: &str,
) -> Result<(Arc<dyn ModelRuntime>, &'c str), RuntimeSelection> {
if !NAMES.contains(&name) {
return Err(RuntimeSelection::UnknownRuntime {
runtime: name.to_string(),
format: format.to_string(),
});
}
let enabled = config
.runtimes
.get(name)
.is_some_and(|runtime| runtime.enabled);
let (Some(runtime), Some(device)) =
(self.get(name).filter(|_| enabled), config.device_of(name))
else {
return Err(RuntimeSelection::NotEnabled {
runtime: name.to_string(),
format: format.to_string(),
});
};
let serves = runtime.formats();
if !serves.contains(&format) {
return Err(RuntimeSelection::DoesNotServe {
runtime: name.to_string(),
format: format.to_string(),
serves,
});
}
Ok((runtime, device))
}
pub fn names(&self) -> Vec<&'static str> {
self.by_name.keys().copied().collect()
}
pub fn is_empty(&self) -> bool {
self.by_name.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
struct Stub(&'static str, &'static [&'static str]);
impl ModelRuntime for Stub {
fn name(&self) -> &'static str {
self.0
}
fn devices(&self) -> &'static [&'static str] {
TRACT_DEVICES
}
fn formats(&self) -> &'static [&'static str] {
self.1
}
fn load(
&self,
_bytes: &[u8],
_manifest: &Manifest,
_device: &str,
) -> Result<Arc<dyn LoadedModel>, LoadError> {
Err(LoadError::new("parse", "stub"))
}
}
#[test]
fn every_known_runtime_lists_cpu_and_a_format_and_nothing_else_registers() {
let mut union: Vec<&str> = Vec::new();
for name in NAMES {
let devices = devices_of(name).expect("every known runtime lists its devices");
assert!(devices.contains(&"cpu"), "{name} must run on cpu");
let formats = formats_of(name).expect("every known runtime lists its formats");
assert!(!formats.is_empty(), "{name} must serve a format");
for format in formats {
if !union.contains(format) {
union.push(format);
}
}
assert_eq!(intern(name), Some(*name));
}
union.sort_unstable();
let mut known = KNOWN_FORMATS.to_vec();
known.sort_unstable();
assert_eq!(
known, union,
"KNOWN_FORMATS is the union of formats_of over NAMES"
);
assert!(devices_of("onnxruntime").is_none());
assert!(formats_of("onnxruntime").is_none());
assert!(intern("onnxruntime").is_none());
let mut registry = ModelRuntimes::empty();
assert!(registry.is_empty());
registry
.register(Arc::new(Stub("tract", TRACT_FORMATS)))
.expect("a listed name registers");
let err = registry
.register(Arc::new(Stub("tract", TRACT_FORMATS)))
.expect_err("twice");
assert!(err.contains("registered twice"), "{err}");
let err = registry
.register(Arc::new(Stub("onnxruntime", TRACT_FORMATS)))
.expect_err("unlisted");
assert!(err.contains("model::runtimes::NAMES"), "{err}");
let err = ModelRuntimes::empty()
.register(Arc::new(Stub("tract", &[])))
.expect_err("no format");
assert!(err.contains("serves no format"), "{err}");
assert_eq!(registry.names(), ["tract"]);
assert!(registry.get("tract").is_some());
assert!(registry.get("onnxruntime").is_none());
}
fn refused(
result: Result<(Arc<dyn ModelRuntime>, &str), RuntimeSelection>,
) -> RuntimeSelection {
match result {
Err(selection) => selection,
Ok((runtime, device)) => {
unreachable!("expected a refusal, got {} on {device}", runtime.name())
}
}
}
#[test]
fn selection_follows_the_format_table_and_names_every_refusal() {
let config = ModelsConfig::default();
let registry = ModelRuntimes::builtin(&config);
let (runtime, device) = registry
.default_for(&config, "onnx")
.expect("the default config maps onnx to tract");
assert_eq!(runtime.name(), "tract");
assert_eq!(device, "cpu");
let (runtime, device) = registry
.for_format(&config, "tract", "onnx")
.expect("explicit tract for onnx");
assert_eq!(runtime.name(), "tract");
assert_eq!(device, "cpu");
let mut metal = ModelsConfig::default();
metal.runtimes.get_mut("tract").expect("entry").device = "metal".to_string();
let (_, device) = registry
.default_for(&metal, "onnx")
.expect("device comes from the config");
assert_eq!(device, "metal");
let err = refused(registry.default_for(&config, "nnef"));
assert_eq!(
err,
RuntimeSelection::NoDefault {
format: "nnef".to_string()
}
);
assert_eq!(
err.to_string(),
"no default runtime for format 'nnef' (models.default_runtime.nnef is not set)"
);
let mut ort = ModelsConfig::default();
ort.default_runtime
.insert("onnx".to_string(), "ort".to_string());
let err = refused(registry.default_for(&ort, "onnx"));
assert_eq!(
err,
RuntimeSelection::UnknownRuntime {
runtime: "ort".to_string(),
format: "onnx".to_string()
}
);
assert_eq!(
err.to_string(),
"runtime 'ort' for format 'onnx' is unknown: this build lists tract"
);
let mut disabled = ModelsConfig::default();
disabled.runtimes.get_mut("tract").expect("entry").enabled = false;
let err = refused(ModelRuntimes::builtin(&disabled).default_for(&disabled, "onnx"));
assert_eq!(
err,
RuntimeSelection::NotEnabled {
runtime: "tract".to_string(),
format: "onnx".to_string()
}
);
assert_eq!(
err.to_string(),
"runtime 'tract' for format 'onnx' is not enabled on this node \
(models.runtimes.tract.enabled)"
);
let err = refused(ModelRuntimes::empty().default_for(&config, "onnx"));
assert!(matches!(err, RuntimeSelection::NotEnabled { .. }), "{err}");
let mut by_hand = ModelRuntimes::empty();
by_hand
.register(Arc::new(Stub("tract", TRACT_FORMATS)))
.expect("registers");
let mut absent = ModelsConfig::default();
absent.runtimes.clear();
let err = refused(by_hand.for_format(&absent, "tract", "onnx"));
assert!(matches!(err, RuntimeSelection::NotEnabled { .. }), "{err}");
let err = refused(registry.for_format(&config, "tract", "nnef"));
assert_eq!(
err,
RuntimeSelection::DoesNotServe {
runtime: "tract".to_string(),
format: "nnef".to_string(),
serves: TRACT_FORMATS
}
);
assert_eq!(
err.to_string(),
"runtime 'tract' does not serve format 'nnef'; it serves 'onnx'"
);
}
#[test]
fn builtin_follows_the_config() {
let config = ModelsConfig::default();
let registry = ModelRuntimes::builtin(&config);
assert_eq!(registry.names(), ["tract"]);
let tract = registry.get("tract").expect("registered");
assert_eq!(tract.name(), "tract");
assert!(tract.devices().contains(&"cpu"));
assert_eq!(tract.formats(), formats_of("tract").expect("listed"));
let mut disabled = ModelsConfig::default();
disabled
.runtimes
.get_mut("tract")
.expect("default entry")
.enabled = false;
assert!(ModelRuntimes::builtin(&disabled).is_empty());
let mut absent = ModelsConfig::default();
absent.runtimes.clear();
assert!(ModelRuntimes::builtin(&absent).is_empty());
}
#[test]
fn errors_name_their_stage() {
assert_eq!(
LoadError::new("device", "no metal here").to_string(),
"device: no metal here"
);
assert_eq!(
RunError::new("input", "rank 3 given, rank 4 expected").to_string(),
"input: rank 3 given, rank 4 expected"
);
}
}