Skip to main content

tokenmiser_providers/
registry.rs

1//! Resolves a model name to a concrete `Provider`, in order: configured
2//! alias, `provider:model` prefix, model-family heuristic, then the opt-in
3//! `routing.default_provider`.
4
5use 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)>, // model -> (provider_name, real_model)
18    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    /// Pick a provider + real model name for an incoming `model` request.
61    pub fn resolve(&self, model: &str) -> Result<(Arc<dyn Provider>, String), ProviderError> {
62        // Explicit alias.
63        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        // `provider:model` prefix.
70        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        // Model-family heuristic.
77        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        // Opt-in only: falling through to a guessed default would route
104        // typos and unsupported models to an arbitrary configured provider,
105        // which then fails with an unrelated missing-API-key error.
106        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        // Gemini goes through its OpenAI-compatibility layer.
127        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}