tokenmiser_providers/
registry.rs1use std::collections::HashMap;
6use std::sync::Arc;
7
8use tokenmiser_config::{ProviderConfig, ProviderKind, TokenmiserConfig};
9
10use crate::{
11 anthropic::AnthropicProvider, ollama::OllamaProvider, openai::OpenAIProvider, Provider,
12 ProviderError,
13};
14
15pub struct ProviderRegistry {
16 providers: HashMap<String, Arc<dyn Provider>>,
17 aliases: HashMap<String, (String, String)>, default_provider: Option<String>,
19}
20
21impl ProviderRegistry {
22 pub fn from_config(cfg: &TokenmiserConfig) -> Self {
23 let mut providers: HashMap<String, Arc<dyn Provider>> = HashMap::new();
24 for p in &cfg.providers {
25 let provider: Arc<dyn Provider> = build_provider(p.clone());
26 providers.insert(p.name.clone(), provider);
27 }
28
29 let aliases = cfg
30 .routing
31 .aliases
32 .iter()
33 .map(|(model, target)| {
34 (
35 model.clone(),
36 (target.provider.clone(), target.model.clone()),
37 )
38 })
39 .collect();
40
41 Self {
42 providers,
43 aliases,
44 default_provider: cfg.routing.default_provider.clone(),
45 }
46 }
47
48 pub fn get(&self, name: &str) -> Option<Arc<dyn Provider>> {
49 self.providers.get(name).cloned()
50 }
51
52 pub fn register(&mut self, name: String, provider: Arc<dyn Provider>) {
53 self.providers.insert(name, provider);
54 }
55
56 pub fn names(&self) -> Vec<String> {
57 self.providers.keys().cloned().collect()
58 }
59
60 pub fn resolve(&self, model: &str) -> Result<(Arc<dyn Provider>, String), ProviderError> {
62 if let Some((provider_name, real_model)) = self.aliases.get(model) {
64 if let Some(p) = self.providers.get(provider_name) {
65 return Ok((p.clone(), real_model.clone()));
66 }
67 }
68
69 if let Some((prefix, rest)) = model.split_once(':') {
71 if let Some(p) = self.providers.get(prefix) {
72 return Ok((p.clone(), rest.to_string()));
73 }
74 }
75
76 let lower = model.to_lowercase();
78 let guess: Option<&str> = if lower.starts_with("claude") {
79 Some("anthropic")
80 } else if lower.starts_with("gpt") || lower.starts_with("o1") || lower.starts_with("o3") {
81 Some("openai")
82 } else if lower.starts_with("gemini") {
83 Some("gemini")
84 } else if lower.starts_with("deepseek") {
85 Some("deepseek")
86 } else if lower.starts_with("llama")
87 || lower.starts_with("qwen")
88 || lower.starts_with("mistral")
89 || lower.starts_with("phi")
90 || lower.starts_with("gemma")
91 {
92 Some("ollama")
93 } else {
94 None
95 };
96
97 if let Some(name) = guess {
98 if let Some(p) = self.providers.get(name) {
99 return Ok((p.clone(), model.to_string()));
100 }
101 }
102
103 if let Some(name) = &self.default_provider {
107 if let Some(p) = self.providers.get(name) {
108 return Ok((p.clone(), model.to_string()));
109 }
110 }
111
112 let mut known: Vec<String> = self.providers.keys().cloned().collect();
113 known.sort();
114 Err(ProviderError::UnknownModel {
115 model: model.to_string(),
116 known_providers: known.join(", "),
117 })
118 }
119}
120
121fn build_provider(cfg: ProviderConfig) -> Arc<dyn Provider> {
122 match cfg.kind {
123 ProviderKind::OpenAI | ProviderKind::DeepSeek => Arc::new(OpenAIProvider::new(cfg)),
124 ProviderKind::Anthropic => Arc::new(AnthropicProvider::new(cfg)),
125 ProviderKind::Ollama => Arc::new(OllamaProvider::new(cfg)),
126 ProviderKind::Gemini => Arc::new(OpenAIProvider::new(cfg)),
128 }
129}
130
131#[cfg(test)]
132mod tests {
133 use super::*;
134 use tokenmiser_config::{ListenConfig, ModelTarget, RoutingConfig};
135
136 fn cfg() -> TokenmiserConfig {
137 let mut aliases = HashMap::new();
138 aliases.insert(
139 "gpt-5".into(),
140 ModelTarget {
141 provider: "openai".into(),
142 model: "gpt-5".into(),
143 },
144 );
145 TokenmiserConfig {
146 listen: ListenConfig::default(),
147 providers: vec![
148 ProviderConfig::openai(),
149 ProviderConfig::anthropic(),
150 ProviderConfig::ollama_local(),
151 ],
152 routing: RoutingConfig {
153 aliases,
154 default_provider: Some("openai".into()),
155 },
156 cache: Default::default(),
157 budget: Default::default(),
158 security: Default::default(),
159 }
160 }
161
162 #[test]
163 fn resolves_claude_to_anthropic() {
164 let reg = ProviderRegistry::from_config(&cfg());
165 let (p, m) = reg.resolve("claude-sonnet-4-6").unwrap();
166 assert_eq!(p.name(), "anthropic");
167 assert_eq!(m, "claude-sonnet-4-6");
168 }
169
170 #[test]
171 fn resolves_provider_prefix() {
172 let reg = ProviderRegistry::from_config(&cfg());
173 let (p, m) = reg.resolve("ollama:llama3.2").unwrap();
174 assert_eq!(p.name(), "ollama");
175 assert_eq!(m, "llama3.2");
176 }
177
178 #[test]
179 fn resolves_alias() {
180 let reg = ProviderRegistry::from_config(&cfg());
181 let (p, m) = reg.resolve("gpt-5").unwrap();
182 assert_eq!(p.name(), "openai");
183 assert_eq!(m, "gpt-5");
184 }
185
186 #[test]
187 fn falls_back_to_default_provider() {
188 let reg = ProviderRegistry::from_config(&cfg());
189 let (p, _) = reg.resolve("some-obscure-model").unwrap();
190 assert_eq!(p.name(), "openai");
191 }
192
193 #[test]
194 fn unknown_model_with_no_default_returns_actionable_error() {
195 let mut c = cfg();
196 c.routing.default_provider = None;
197 let reg = ProviderRegistry::from_config(&c);
198 match reg.resolve("totally-unknown-model") {
199 Ok(_) => panic!("expected UnknownModel error, got Ok"),
200 Err(ProviderError::UnknownModel {
201 model,
202 known_providers,
203 }) => {
204 assert_eq!(model, "totally-unknown-model");
205 assert!(known_providers.contains("openai"));
206 assert!(known_providers.contains("anthropic"));
207 assert!(known_providers.contains("ollama"));
208 }
209 Err(other) => panic!("expected UnknownModel, got {:?}", other),
210 }
211 }
212}