use std::collections::HashMap;
use std::sync::{Mutex, OnceLock};
use crate::error::{AvError, AvResult};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum Category {
Modality,
Backbone,
Neck,
Head,
Loss,
PostProc,
Rule,
}
impl Category {
pub fn as_str(self) -> &'static str {
match self {
Category::Modality => "modality",
Category::Backbone => "backbone",
Category::Neck => "neck",
Category::Head => "head",
Category::Loss => "loss",
Category::PostProc => "postproc",
Category::Rule => "rule",
}
}
}
#[derive(Debug, Default)]
pub struct Registry {
entries: HashMap<Category, Vec<String>>,
}
impl Registry {
pub fn register(&mut self, cat: Category, name: &str) -> AvResult<()> {
let bucket = self.entries.entry(cat).or_default();
if bucket.iter().any(|n| n == name) {
return Err(AvError::config(format!(
"插件重复注册: {}/{name}",
cat.as_str()
)));
}
bucket.push(name.to_string());
Ok(())
}
pub fn lookup(&self, cat: Category, name: &str) -> bool {
self.entries
.get(&cat)
.is_some_and(|b| b.iter().any(|n| n == name))
}
pub fn list(&self, cat: Category) -> Vec<String> {
self.entries.get(&cat).cloned().unwrap_or_default()
}
}
static GLOBAL: OnceLock<Mutex<Registry>> = OnceLock::new();
fn global() -> &'static Mutex<Registry> {
GLOBAL.get_or_init(|| Mutex::new(Registry::default()))
}
pub fn register(cat: Category, name: &str) -> AvResult<()> {
global()
.lock()
.expect("registry 互斥锁中毒")
.register(cat, name)
}
pub fn lookup(cat: Category, name: &str) -> bool {
global()
.lock()
.expect("registry 互斥锁中毒")
.lookup(cat, name)
}
pub fn list(cat: Category) -> Vec<String> {
global().lock().expect("registry 互斥锁中毒").list(cat)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn register_lookup_list() {
register(Category::Backbone, "csp-elan").unwrap();
assert!(lookup(Category::Backbone, "csp-elan"));
assert!(!lookup(Category::Backbone, "vit-hybrid"));
assert_eq!(list(Category::Backbone), vec!["csp-elan".to_string()]);
}
#[test]
fn duplicate_registration_errors() {
register(Category::Loss, "ciou").unwrap();
let err = register(Category::Loss, "ciou").unwrap_err();
assert!(err.to_string().contains("重复注册"));
}
}