Skip to main content

leviath_runtime/
providers.rs

1//! The [`ProviderRegistry`]: a name → [`Provider`] lookup shared by the ECS
2//! pipeline (as the `Providers` resource) and the CLI/daemon spawn path.
3
4use crate::script_provider::ScriptProviderLayer;
5use leviath_providers::Provider;
6use std::collections::HashMap;
7use std::sync::Arc;
8
9/// Registry of inference providers, keyed by provider name (e.g. `"anthropic"`).
10///
11/// The pipeline resolves each agent's stage `ModelConfig` to a concrete
12/// provider through this registry. Native providers are registered eagerly;
13/// script providers are resolved lazily - and hot-reloaded - via
14/// an optional [`ScriptProviderLayer`].
15#[derive(Clone, Default)]
16pub struct ProviderRegistry {
17    providers: HashMap<String, Arc<dyn Provider>>,
18    /// Lazy, hot-reloading resolver for `.rhai` script providers. Shared across
19    /// registry clones (one compile cache daemon-wide).
20    script_layer: Option<Arc<ScriptProviderLayer>>,
21}
22
23impl ProviderRegistry {
24    /// Create a new empty provider registry.
25    pub fn new() -> Self {
26        Self::default()
27    }
28
29    /// Attach a script-provider layer for lazy/hot-reloading `.rhai` providers.
30    pub fn with_script_layer(mut self, layer: Arc<ScriptProviderLayer>) -> Self {
31        self.script_layer = Some(layer);
32        self
33    }
34
35    /// Register a provider by name.
36    pub fn register(&mut self, name: String, provider: Arc<dyn Provider>) {
37        self.providers.insert(name, provider);
38    }
39
40    /// Get a provider by name, returning an owned handle.
41    ///
42    /// A native provider wins; otherwise the script layer is consulted, which
43    /// lazily compiles (or hot-reloads) the matching `.rhai` script.
44    pub fn get(&self, name: &str) -> Option<Arc<dyn Provider>> {
45        if let Some(p) = self.providers.get(name) {
46            return Some(p.clone());
47        }
48        self.script_layer.as_ref()?.get_or_load(name)
49    }
50
51    /// Check if a provider is available: registered natively, or resolvable
52    /// (loadable) as a script provider right now. Used at stage-model selection;
53    /// network-free because script `initialize` runs offline.
54    pub fn has(&self, name: &str) -> bool {
55        self.providers.contains_key(name)
56            || self
57                .script_layer
58                .as_ref()
59                .is_some_and(|l| l.get_or_load(name).is_some())
60    }
61
62    /// Get all *natively-registered* provider names. Script providers are
63    /// resolved on demand and so are not enumerated here.
64    pub fn provider_names(&self) -> Vec<&str> {
65        self.providers.keys().map(|k| k.as_str()).collect()
66    }
67}
68
69#[cfg(test)]
70mod tests {
71    use super::*;
72    use leviath_providers::{
73        InferenceRequest, InferenceResponse, ModelCapabilities, ProviderError,
74    };
75
76    struct StubProvider;
77
78    #[async_trait::async_trait]
79    impl Provider for StubProvider {
80        async fn infer(
81            &self,
82            _request: InferenceRequest,
83        ) -> Result<InferenceResponse, ProviderError> {
84            Err(ProviderError::ApiError("stub".to_string()))
85        }
86        async fn count_tokens(&self, text: &str, _model: &str) -> usize {
87            text.len()
88        }
89        fn max_context_tokens(&self, _model: &str) -> usize {
90            8192
91        }
92        fn name(&self) -> &str {
93            "stub"
94        }
95        fn capabilities(&self, _model: &str) -> ModelCapabilities {
96            ModelCapabilities::default()
97        }
98    }
99
100    fn mock() -> Arc<dyn Provider> {
101        Arc::new(StubProvider)
102    }
103
104    #[test]
105    fn register_get_has_and_names() {
106        let mut reg = ProviderRegistry::new();
107        assert!(!reg.has("anthropic"));
108        assert!(reg.get("anthropic").is_none());
109        reg.register("anthropic".to_string(), mock());
110        assert!(reg.has("anthropic"));
111        assert!(reg.get("anthropic").is_some());
112        assert_eq!(reg.provider_names(), vec!["anthropic"]);
113    }
114
115    #[test]
116    fn default_is_empty() {
117        let reg = ProviderRegistry::default();
118        assert!(reg.provider_names().is_empty());
119    }
120
121    #[test]
122    fn script_layer_resolves_and_native_wins() {
123        use crate::script_provider::ScriptProviderLayer;
124        use std::collections::HashMap;
125
126        let dir = tempfile::tempdir().unwrap();
127        std::fs::write(
128            dir.path().join("groq.rhai"),
129            "fn initialize(config) { #{} }\nfn inference(state, request) { #{ content: \"ok\" } }",
130        )
131        .unwrap();
132        let layer = ScriptProviderLayer::new(
133            dir.path().to_path_buf(),
134            HashMap::new(),
135            HashMap::new(),
136            None,
137            Vec::new(),
138        );
139        let mut reg = ProviderRegistry::new().with_script_layer(Arc::new(layer));
140        reg.register("anthropic".to_string(), mock());
141
142        // Native provider still wins and is found by name.
143        assert!(reg.has("anthropic"));
144        assert!(reg.get("anthropic").is_some());
145
146        // A script provider is resolved lazily through the layer by both has/get.
147        assert!(reg.has("groq"));
148        let p = reg.get("groq").expect("script provider resolves");
149        assert_eq!(p.name(), "groq");
150
151        // An unknown name resolves to nothing (layer returns None).
152        assert!(!reg.has("nope"));
153        assert!(reg.get("nope").is_none());
154    }
155
156    #[tokio::test]
157    async fn stub_provider_methods_are_exercised() {
158        let p = StubProvider;
159        assert_eq!(p.name(), "stub");
160        assert_eq!(p.count_tokens("abcd", "m").await, 4);
161        assert_eq!(p.max_context_tokens("m"), 8192);
162        let _ = p.capabilities("m");
163        let request = InferenceRequest {
164            system: Vec::new(),
165            messages: Vec::new(),
166            model: "m".to_string(),
167            max_tokens: 10,
168            temperature: 0.0,
169            tools: Vec::new(),
170            extra: serde_json::Value::Null,
171            request_timeout_secs: None,
172        };
173        assert!(p.infer(request).await.is_err());
174    }
175}