openkind_engine/
registry.rs1use std::collections::HashMap;
4use std::sync::{Arc, RwLock};
5
6use openkind_core::ModelInfo;
7
8use crate::engine::DecisionEngine;
9
10#[derive(Default, Clone)]
13pub struct EngineRegistry {
14 engines: Arc<RwLock<HashMap<String, Arc<dyn DecisionEngine>>>>,
15}
16
17impl std::fmt::Debug for EngineRegistry {
18 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
19 f.debug_struct("EngineRegistry")
20 .field("models", &self.models())
21 .finish()
22 }
23}
24
25impl EngineRegistry {
26 pub fn new() -> Self {
28 Self::default()
29 }
30
31 pub fn register(&mut self, alias: impl Into<String>, engine: Arc<dyn DecisionEngine>) {
33 let alias = alias.into();
34 let replaced = self
35 .engines
36 .write()
37 .expect("registry lock poisoned")
38 .insert(alias, engine);
39 drop(replaced);
41 }
42
43 pub fn get(&self, model: &str) -> Option<Arc<dyn DecisionEngine>> {
45 self.engines
46 .read()
47 .expect("registry lock poisoned")
48 .get(model)
49 .cloned()
50 }
51
52 pub fn models(&self) -> Vec<String> {
54 let mut v: Vec<String> = self
55 .engines
56 .read()
57 .expect("registry lock poisoned")
58 .keys()
59 .cloned()
60 .collect();
61 v.sort();
62 v
63 }
64
65 pub fn list_models(&self) -> Vec<ModelInfo> {
68 let engines = self.engines.read().expect("registry lock poisoned").clone();
71 let mut models: Vec<_> = engines
72 .into_iter()
73 .map(|(name, engine)| {
74 let mut metadata = engine.model_metadata();
75 metadata.name = name;
76 metadata
77 })
78 .collect();
79 models.sort_by(|a, b| a.name.cmp(&b.name));
80 models
81 }
82
83 pub fn register_if_absent(&self, alias: String, engine: Arc<dyn DecisionEngine>) -> bool {
86 use std::collections::hash_map::Entry;
87 match self
88 .engines
89 .write()
90 .expect("registry lock poisoned")
91 .entry(alias)
92 {
93 Entry::Vacant(entry) => {
94 entry.insert(engine);
95 true
96 }
97 Entry::Occupied(_) => false,
98 }
99 }
100
101 pub fn unregister(&self, alias: &str) -> Option<Arc<dyn DecisionEngine>> {
104 self.engines
105 .write()
106 .expect("registry lock poisoned")
107 .remove(alias)
108 }
109}
110
111#[cfg(test)]
112mod tests {
113 use std::sync::atomic::{AtomicBool, Ordering};
114
115 use async_trait::async_trait;
116 use openkind_core::{SystemRequest, SystemResponse};
117
118 use super::*;
119 use crate::{EngineResult, MockEngine};
120
121 struct CleanupBackend {
122 registry: EngineRegistry,
123 cleanup_had_access: Arc<AtomicBool>,
124 }
125
126 #[async_trait]
127 impl DecisionEngine for CleanupBackend {
128 fn backend_id(&self) -> &str {
129 "cleanup"
130 }
131
132 async fn evaluate(&self, _req: SystemRequest) -> EngineResult<SystemResponse> {
133 unreachable!("cleanup-only test")
134 }
135 }
136
137 impl Drop for CleanupBackend {
138 fn drop(&mut self) {
139 self.cleanup_had_access
140 .store(self.registry.engines.try_write().is_ok(), Ordering::SeqCst);
141 }
142 }
143
144 #[test]
145 fn replacement_releases_registry_lock_before_backend_cleanup() {
146 let mut registry = EngineRegistry::new();
147 let cleanup_had_access = Arc::new(AtomicBool::new(false));
148 registry.register(
149 "model",
150 Arc::new(CleanupBackend {
151 registry: registry.clone(),
152 cleanup_had_access: cleanup_had_access.clone(),
153 }),
154 );
155
156 registry.register("model", Arc::new(MockEngine::new()));
157
158 assert!(cleanup_had_access.load(Ordering::SeqCst));
159 assert_eq!(registry.get("model").unwrap().backend_id(), "mock");
160 }
161
162 #[test]
163 fn rejected_registration_releases_lock_before_backend_cleanup() {
164 let mut registry = EngineRegistry::new();
165 registry.register("model", Arc::new(MockEngine::new()));
166 let cleanup_had_access = Arc::new(AtomicBool::new(false));
167
168 assert!(!registry.register_if_absent(
169 "model".into(),
170 Arc::new(CleanupBackend {
171 registry: registry.clone(),
172 cleanup_had_access: cleanup_had_access.clone(),
173 }),
174 ));
175
176 assert!(cleanup_had_access.load(Ordering::SeqCst));
177 assert_eq!(registry.get("model").unwrap().backend_id(), "mock");
178 }
179}