use crate::error::{Result, ShimError};
use crate::provider::Provider;
use crate::providers::anthropic::Anthropic;
use crate::providers::chatgpt::{ChatGpt, ChatGptAuth};
use crate::providers::gemini::Gemini;
use crate::providers::openai::OpenAi;
use crate::providers::openai_compat::OpenAiCompatible;
use crate::providers::openrouter::OpenRouter;
use crate::providers::xai::Xai;
use std::collections::HashMap;
pub fn parse_model(model: &str, aliases: &HashMap<String, String>) -> Result<(String, String)> {
let resolved = aliases.get(model).map(|s| s.as_str()).unwrap_or(model);
if let Some((provider, model_name)) = resolved.split_once('/') {
Ok((provider.to_string(), model_name.to_string()))
} else {
let lower = resolved.to_lowercase();
if lower.starts_with("gpt")
|| lower.starts_with("o1")
|| lower.starts_with("o3")
|| lower.starts_with("o4")
{
Ok(("openai".to_string(), resolved.to_string()))
} else if lower.starts_with("claude") {
Ok(("anthropic".to_string(), resolved.to_string()))
} else if lower.starts_with("gemini") {
Ok(("gemini".to_string(), resolved.to_string()))
} else if lower.starts_with("grok") {
Ok(("xai".to_string(), resolved.to_string()))
} else {
Err(ShimError::UnknownProvider(resolved.to_string()))
}
}
}
pub struct Router {
providers: HashMap<String, std::sync::Arc<dyn Provider>>,
pub aliases: HashMap<String, String>,
}
impl Default for Router {
fn default() -> Self {
Self::new()
}
}
impl Router {
pub fn new() -> Self {
Self {
providers: HashMap::new(),
aliases: HashMap::new(),
}
}
pub fn register(mut self, key: &str, provider: Box<dyn Provider>) -> Self {
self.providers.insert(key.to_string(), provider.into());
self
}
pub fn alias(mut self, from: &str, to: &str) -> Self {
self.aliases.insert(from.to_string(), to.to_string());
self
}
pub fn provider_keys(&self) -> Vec<&str> {
self.providers.keys().map(|s| s.as_str()).collect()
}
pub fn get(&self, key: &str) -> Result<&dyn Provider> {
self.providers
.get(key)
.map(|p| p.as_ref())
.ok_or_else(|| ShimError::UnknownProvider(key.to_string()))
}
pub fn from_env() -> Self {
let mut router = Router::new();
if let Ok(catalog) = crate::catalog::global() {
catalog.refresh_in_background();
}
let chatgpt_auth = ChatGptAuth::from_env();
if chatgpt_auth.auth_path().is_file() {
router = router.register("chatgpt", Box::new(ChatGpt::new(chatgpt_auth)));
}
if let Ok(key) = std::env::var("OPENAI_API_KEY") {
router = router.register("openai", Box::new(OpenAi::new(key)));
}
if let Ok(key) = std::env::var("ANTHROPIC_API_KEY") {
router = router.register("anthropic", Box::new(Anthropic::new(key)));
}
if let Ok(key) = std::env::var("GEMINI_API_KEY") {
router = router.register("gemini", Box::new(Gemini::new(key)));
}
if let Ok(key) = std::env::var("XAI_API_KEY") {
router = router.register("xai", Box::new(Xai::new(key)));
}
if let Ok(key) = std::env::var("OPENROUTER_API_KEY") {
router = router.register("openrouter", Box::new(OpenRouter::new(key)));
}
if let Ok(base) = std::env::var("VLLM_BASE_URL") {
let key = std::env::var("VLLM_API_KEY").ok().filter(|k| !k.is_empty());
router = router.register("vllm", Box::new(OpenAiCompatible::new("vllm", base, key)));
}
if let Ok(base) = std::env::var("SGLANG_BASE_URL") {
let key = std::env::var("SGLANG_API_KEY")
.ok()
.filter(|k| !k.is_empty());
router = router.register(
"sglang",
Box::new(OpenAiCompatible::new("sglang", base, key)),
);
}
router
}
pub fn resolve(&self, model: &str) -> Result<(&dyn Provider, String)> {
let (key, model) = self.resolve_key(model)?;
Ok((self.get(&key)?, model))
}
pub fn resolve_owned(&self, model: &str) -> Result<(std::sync::Arc<dyn Provider>, String)> {
let (key, model) = self.resolve_key(model)?;
let provider = self
.providers
.get(&key)
.cloned()
.ok_or(ShimError::UnknownProvider(key))?;
Ok((provider, model))
}
fn resolve_key(&self, model: &str) -> Result<(String, String)> {
crate::catalog::global().map_err(|error| ShimError::ProviderError {
status: 400,
body: format!("invalid model catalog configuration: {error}"),
})?;
let requested = self.aliases.get(model).map(String::as_str).unwrap_or(model);
let metadata = crate::catalog::resolve(requested).filter(|m| {
requested.split_once('/').is_none_or(|(prefix, _)| {
prefix == m.provider
|| crate::catalog::aliases::PROVIDER_ALIASES
.iter()
.any(|(alias, canonical)| prefix == *alias && m.provider == *canonical)
})
});
let (provider_key, model_name) = match metadata {
Some(m) => (m.provider.clone(), m.name.clone()),
None => parse_model(model, &self.aliases)?,
};
Ok((provider_key, model_name))
}
}