use crate::model_package::{ModelPackage, ModelTask};
use latexsnipper_foundation::Result;
use std::path::Path;
pub trait ModelPlugin: Send + Sync {
fn name(&self) -> &str;
fn version(&self) -> &str;
fn supported_adapters(&self) -> Vec<&str>;
fn create_package(
&self,
adapter: &str,
manifest: &dyn ModelManifestView,
model_dir: &Path,
) -> Result<Box<dyn ModelPackage>>;
fn init(&mut self) -> Result<()> {
Ok(())
}
fn cleanup(&mut self) -> Result<()> {
Ok(())
}
}
pub trait ModelManifestView {
fn id(&self) -> &str;
fn task(&self) -> ModelTask;
fn adapter(&self) -> &str;
fn version(&self) -> &str;
fn input_name(&self) -> &str;
fn input_shape(&self) -> &[i64];
fn get_file(&self, name: &str) -> Option<&str>;
}
pub struct ModelPluginRegistry {
plugins: Vec<Box<dyn ModelPlugin>>,
adapter_to_plugin: std::collections::HashMap<String, usize>,
}
impl ModelPluginRegistry {
pub fn new() -> Self {
Self {
plugins: Vec::new(),
adapter_to_plugin: std::collections::HashMap::new(),
}
}
pub fn register(&mut self, mut plugin: Box<dyn ModelPlugin>) -> Result<()> {
plugin.init()?;
let idx = self.plugins.len();
let adapters = plugin.supported_adapters();
for adapter in adapters {
self.adapter_to_plugin.insert(adapter.to_string(), idx);
}
self.plugins.push(plugin);
Ok(())
}
pub fn find_plugin(&self, adapter: &str) -> Option<&dyn ModelPlugin> {
let idx = self.adapter_to_plugin.get(adapter)?;
self.plugins.get(*idx).map(|p| p.as_ref())
}
pub fn create_package(
&self,
adapter: &str,
manifest: &dyn ModelManifestView,
model_dir: &Path,
) -> Result<Option<Box<dyn ModelPackage>>> {
let plugin = match self.find_plugin(adapter) {
Some(p) => p,
None => return Ok(None),
};
let package = plugin.create_package(adapter, manifest, model_dir)?;
Ok(Some(package))
}
pub fn registered_adapters(&self) -> Vec<&str> {
self.adapter_to_plugin.keys().map(|s| s.as_str()).collect()
}
pub fn plugins(&self) -> Vec<&dyn ModelPlugin> {
self.plugins.iter().map(|p| p.as_ref()).collect()
}
}
impl Default for ModelPluginRegistry {
fn default() -> Self {
Self::new()
}
}
pub struct ManifestAdapter<'a> {
manifest: &'a crate::model_registry::ModelManifest,
}
impl<'a> ManifestAdapter<'a> {
pub fn new(manifest: &'a crate::model_registry::ModelManifest) -> Self {
Self { manifest }
}
}
impl<'a> ModelManifestView for ManifestAdapter<'a> {
fn id(&self) -> &str {
&self.manifest.id
}
fn task(&self) -> ModelTask {
self.manifest.task
}
fn adapter(&self) -> &str {
&self.manifest.adapter
}
fn version(&self) -> &str {
&self.manifest.version
}
fn input_name(&self) -> &str {
&self.manifest.input.name
}
fn input_shape(&self) -> &[i64] {
&self.manifest.input.shape
}
fn get_file(&self, name: &str) -> Option<&str> {
match name {
"primary" => self.manifest.files.primary.as_deref(),
"encoder" => self.manifest.files.encoder.as_deref(),
"decoder" => self.manifest.files.decoder.as_deref(),
"tokenizer" => self.manifest.files.tokenizer.as_deref(),
"config" => self.manifest.files.config.as_deref(),
_ => None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
struct TestPlugin {
name: String,
}
impl ModelPlugin for TestPlugin {
fn name(&self) -> &str {
&self.name
}
fn version(&self) -> &str {
"1.0.0"
}
fn supported_adapters(&self) -> Vec<&str> {
vec!["test-adapter-v1"]
}
fn create_package(
&self,
_adapter: &str,
_manifest: &dyn ModelManifestView,
_model_dir: &Path,
) -> Result<Box<dyn ModelPackage>> {
Err(latexsnipper_foundation::SnipperError::Other(
"Not implemented".into(),
))
}
}
#[test]
fn test_plugin_registry() {
let mut registry = ModelPluginRegistry::new();
let plugin = Box::new(TestPlugin {
name: "test".to_string(),
});
registry.register(plugin).unwrap();
assert!(registry.find_plugin("test-adapter-v1").is_some());
assert!(registry.find_plugin("unknown").is_none());
assert_eq!(registry.registered_adapters(), vec!["test-adapter-v1"]);
}
}