leviath_runtime/
providers.rs1use crate::script_provider::ScriptProviderLayer;
5use leviath_providers::Provider;
6use std::collections::HashMap;
7use std::sync::Arc;
8
9#[derive(Clone, Default)]
16pub struct ProviderRegistry {
17 providers: HashMap<String, Arc<dyn Provider>>,
18 script_layer: Option<Arc<ScriptProviderLayer>>,
21}
22
23impl ProviderRegistry {
24 pub fn new() -> Self {
26 Self::default()
27 }
28
29 pub fn with_script_layer(mut self, layer: Arc<ScriptProviderLayer>) -> Self {
31 self.script_layer = Some(layer);
32 self
33 }
34
35 pub fn register(&mut self, name: String, provider: Arc<dyn Provider>) {
37 self.providers.insert(name, provider);
38 }
39
40 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 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 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 assert!(reg.has("anthropic"));
144 assert!(reg.get("anthropic").is_some());
145
146 assert!(reg.has("groq"));
148 let p = reg.get("groq").expect("script provider resolves");
149 assert_eq!(p.name(), "groq");
150
151 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}