Skip to main content

openkind_engine/
registry.rs

1//! Engine registry for model alias routing and metadata discovery.
2
3use std::collections::HashMap;
4use std::sync::{Arc, RwLock};
5
6use openkind_core::ModelInfo;
7
8use crate::engine::DecisionEngine;
9
10/// A registry mapping model alias → engine. Lets the server dispatch by
11/// the `model` field in the request without the engine itself knowing.
12#[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    /// Construct an empty `EngineRegistry`.
27    pub fn new() -> Self {
28        Self::default()
29    }
30
31    /// Register a decision engine under the specified model alias (e.g. `"jev-latest"`).
32    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        // Backend cleanup can be slow or consult the registry itself.
40        drop(replaced);
41    }
42
43    /// Look up a decision engine by its registered model alias.
44    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    /// Sorted list of registered alias names.
53    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    /// `GET /v1/models` payload — one entry per registered alias,
66    /// pulling metadata from the engine itself.
67    pub fn list_models(&self) -> Vec<ModelInfo> {
68        // Snapshot the handles before asking backends for metadata. No registry
69        // lock is held during backend code or inference.
70        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    /// Publish a loaded engine without replacing a live alias. Cloned registries
84    /// share updates so HTTP and gRPC observe the same model lifecycle.
85    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    /// Stop new dispatches. Existing requests retain their engine handle until
102    /// completion, so unloading does not interrupt an accepted request.
103    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}