use std::{collections::BTreeMap, sync::Arc};
use thiserror::Error;
use crate::{Inference, InferenceError};
pub type Factory = Arc<dyn Fn() -> Result<Arc<dyn Inference>, InferenceError> + Send + Sync>;
#[derive(Default)]
pub struct Registry {
factories: BTreeMap<String, Factory>,
}
impl Registry {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn register(&mut self, kind: &str, factory: Factory) -> Result<(), RegistryError> {
let kind = Self::normalize_kind(kind).map_err(|_| RegistryError::EmptyKind)?;
if self.factories.contains_key(&kind) {
return Err(RegistryError::DuplicateKind(kind));
}
self.factories.insert(kind, factory);
Ok(())
}
pub fn build(&self, kind: &str) -> Result<Arc<dyn Inference>, InferenceError> {
let kind = Self::normalize_kind(kind)?;
let factory = self
.factories
.get(&kind)
.ok_or_else(|| InferenceError::Decode(format!("unknown inference adapter {kind:?}")))?;
factory()
}
fn normalize_kind(kind: &str) -> Result<String, InferenceError> {
let kind = kind.trim();
if kind.is_empty() {
return Err(InferenceError::InvalidInput(
"inference adapter kind is required".to_owned(),
));
}
Ok(kind.to_owned())
}
#[must_use]
pub fn kinds(&self) -> Vec<String> {
self.factories.keys().cloned().collect()
}
}
#[derive(Debug, Error, PartialEq, Eq)]
#[non_exhaustive]
pub enum RegistryError {
#[error("inference adapter kind is required")]
EmptyKind,
#[error("inference adapter {0:?} already registered")]
DuplicateKind(String),
}
#[must_use]
pub fn default_registry() -> Registry {
Registry::new()
}
#[cfg(test)]
mod tests {
use rskit_errors::ErrorCode;
use super::{Registry, RegistryError};
use crate::InferenceError;
#[test]
fn register_rejects_empty_kind() {
let mut registry = Registry::new();
let factory = std::sync::Arc::new(|| unreachable!("factory should not run"));
let err = registry.register(" \t ", factory).unwrap_err();
assert_eq!(err, RegistryError::EmptyKind);
}
#[test]
fn build_rejects_empty_kind_as_invalid_input() {
let Err(err) = Registry::new().build(" \t ") else {
panic!("empty inference adapter kind should be rejected");
};
assert!(matches!(err, InferenceError::InvalidInput(_)));
let app_error = rskit_errors::AppError::from(err);
assert_eq!(app_error.code(), ErrorCode::InvalidInput);
}
}