use super::model::{ApiProtocol, ModelConfig};
use super::traits::*;
use crate::types::*;
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::mpsc;
pub struct ProviderRegistry {
providers: HashMap<ApiProtocol, Arc<dyn StreamProvider>>,
}
impl ProviderRegistry {
pub fn new() -> Self {
Self {
providers: HashMap::new(),
}
}
pub fn register(&mut self, protocol: ApiProtocol, provider: impl StreamProvider + 'static) {
self.providers.insert(protocol, Arc::new(provider));
}
pub fn resolve(&self, protocol: &ApiProtocol) -> Option<Arc<dyn StreamProvider>> {
self.providers.get(protocol).cloned()
}
pub fn get(&self, protocol: &ApiProtocol) -> Option<&dyn StreamProvider> {
self.providers.get(protocol).map(|p| p.as_ref())
}
pub fn has(&self, protocol: &ApiProtocol) -> bool {
self.providers.contains_key(protocol)
}
pub fn protocols(&self) -> Vec<ApiProtocol> {
self.providers.keys().copied().collect()
}
pub async fn stream(
&self,
model: &ModelConfig,
config: StreamConfig,
tx: mpsc::UnboundedSender<StreamEvent>,
cancel: tokio_util::sync::CancellationToken,
) -> Result<Message, ProviderError> {
let provider = self.providers.get(&model.api).ok_or_else(|| {
ProviderError::Other(format!(
"No provider registered for protocol: {}",
model.api
))
})?;
provider.stream(config, tx, cancel).await
}
}
impl Default for ProviderRegistry {
fn default() -> Self {
use crate::provider::{
AnthropicProvider, AzureOpenAiProvider, BedrockProvider, GoogleProvider,
GoogleVertexProvider, OpenAiCompatProvider, OpenAiResponsesProvider,
};
let mut registry = Self::new();
registry.register(ApiProtocol::AnthropicMessages, AnthropicProvider);
registry.register(ApiProtocol::OpenAiCompletions, OpenAiCompatProvider);
registry.register(ApiProtocol::OpenAiResponses, OpenAiResponsesProvider);
registry.register(ApiProtocol::GoogleGenerativeAi, GoogleProvider);
registry.register(ApiProtocol::GoogleVertex, GoogleVertexProvider);
registry.register(ApiProtocol::BedrockConverseStream, BedrockProvider);
registry.register(ApiProtocol::AzureOpenAiResponses, AzureOpenAiProvider);
registry
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_registry_has_all_providers() {
let registry = ProviderRegistry::default();
assert!(registry.has(&ApiProtocol::AnthropicMessages));
assert!(registry.has(&ApiProtocol::OpenAiCompletions));
assert!(registry.has(&ApiProtocol::OpenAiResponses));
assert!(registry.has(&ApiProtocol::GoogleGenerativeAi));
assert!(registry.has(&ApiProtocol::GoogleVertex));
assert!(registry.has(&ApiProtocol::BedrockConverseStream));
assert!(registry.has(&ApiProtocol::AzureOpenAiResponses));
}
#[test]
fn test_registry_protocols() {
let registry = ProviderRegistry::default();
let protocols = registry.protocols();
assert_eq!(protocols.len(), 7);
}
#[test]
fn test_custom_registry() {
let mut registry = ProviderRegistry::new();
assert!(!registry.has(&ApiProtocol::AnthropicMessages));
registry.register(
ApiProtocol::AnthropicMessages,
crate::provider::AnthropicProvider,
);
assert!(registry.has(&ApiProtocol::AnthropicMessages));
}
}
pub fn resolve_api_key(provider: &str) -> Option<String> {
use std::env::var;
let first = |names: &[&str]| {
names.iter().find_map(|n| {
var(n).ok().inspect(|_| {
tracing::debug!("resolved API key for provider '{}' from ${}", provider, n)
})
})
};
match provider {
"anthropic" => first(&["ANTHROPIC_API_KEY"]),
"openai" => first(&["OPENAI_API_KEY"]),
"google" => first(&["GEMINI_API_KEY", "GOOGLE_API_KEY"]),
"xai" => first(&["XAI_API_KEY"]),
"groq" => first(&["GROQ_API_KEY"]),
"deepseek" => first(&["DEEPSEEK_API_KEY"]),
"mistral" => first(&["MISTRAL_API_KEY"]),
"zai" => first(&["ZAI_API_KEY"]),
"minimax" => first(&["MINIMAX_API_KEY"]),
"meta" => first(&["META_API_KEY", "MODEL_API_KEY"]),
"openrouter" => first(&["OPENROUTER_API_KEY"]),
"cerebras" => first(&["CEREBRAS_API_KEY"]),
"qwen" => first(&["DASHSCOPE_API_KEY"]),
"opencode-zen" | "opencode-go" => first(&["OPENCODE_API_KEY"]),
"azure" => first(&["AZURE_OPENAI_API_KEY"]),
"bedrock" => {
let access = var("AWS_ACCESS_KEY_ID").ok()?;
let secret = var("AWS_SECRET_ACCESS_KEY").ok()?;
Some(match var("AWS_SESSION_TOKEN") {
Ok(token) => format!("{}:{}:{}", access, secret, token),
Err(_) => format!("{}:{}", access, secret),
})
}
"vertex" => None,
"local" | "ollama" => Some(String::new()),
_ => first(&["YOAGENT_API_KEY", "API_KEY"]),
}
}
pub(crate) fn resolve_api_key_or_warn(provider: &str) -> String {
match resolve_api_key(provider) {
Some(key) => key,
None => {
tracing::warn!(
"no API key found for provider '{}': {}; requests will fail \
with an authentication error",
provider,
api_key_env_hint(provider)
);
String::new()
}
}
}
fn api_key_env_hint(provider: &str) -> &'static str {
match provider {
"anthropic" => "set ANTHROPIC_API_KEY or call .with_api_key(...)",
"openai" => "set OPENAI_API_KEY or call .with_api_key(...)",
"google" => "set GEMINI_API_KEY (or GOOGLE_API_KEY) or call .with_api_key(...)",
"xai" => "set XAI_API_KEY or call .with_api_key(...)",
"groq" => "set GROQ_API_KEY or call .with_api_key(...)",
"deepseek" => "set DEEPSEEK_API_KEY or call .with_api_key(...)",
"mistral" => "set MISTRAL_API_KEY or call .with_api_key(...)",
"zai" => "set ZAI_API_KEY or call .with_api_key(...)",
"minimax" => "set MINIMAX_API_KEY or call .with_api_key(...)",
"meta" => "set META_API_KEY (or MODEL_API_KEY) or call .with_api_key(...)",
"openrouter" => "set OPENROUTER_API_KEY or call .with_api_key(...)",
"cerebras" => "set CEREBRAS_API_KEY or call .with_api_key(...)",
"qwen" => "set DASHSCOPE_API_KEY or call .with_api_key(...)",
"opencode-zen" | "opencode-go" => "set OPENCODE_API_KEY or call .with_api_key(...)",
"azure" => "set AZURE_OPENAI_API_KEY or call .with_api_key(...)",
"bedrock" => {
"set AWS_ACCESS_KEY_ID + AWS_SECRET_ACCESS_KEY (+ AWS_SESSION_TOKEN) \
or call .with_api_key(\"access:secret[:token]\")"
}
"vertex" => "pass a short-lived OAuth token via .with_api_key(...)",
_ => "set YOAGENT_API_KEY (or API_KEY) or call .with_api_key(...)",
}
}
#[cfg(test)]
mod resolve_key_tests {
use super::resolve_api_key;
#[test]
fn test_deterministic_branches() {
assert_eq!(resolve_api_key("local").as_deref(), Some(""));
assert_eq!(resolve_api_key("ollama").as_deref(), Some(""));
assert_eq!(resolve_api_key("vertex"), None);
}
#[test]
fn test_meta_env_resolution() {
std::env::remove_var("META_API_KEY");
std::env::set_var("MODEL_API_KEY", "official-var");
assert_eq!(resolve_api_key("meta").as_deref(), Some("official-var"));
std::env::set_var("META_API_KEY", "preferred-var");
assert_eq!(resolve_api_key("meta").as_deref(), Some("preferred-var"));
std::env::remove_var("META_API_KEY");
std::env::remove_var("MODEL_API_KEY");
}
#[test]
fn test_env_resolution() {
std::env::set_var("YOAGENT_API_KEY", "from-generic-fallback");
assert_eq!(
resolve_api_key("some-unknown-gateway").as_deref(),
Some("from-generic-fallback")
);
std::env::remove_var("YOAGENT_API_KEY");
}
}
#[cfg(test)]
mod registry_tests {
use super::*;
fn all_protocols() -> Vec<ApiProtocol> {
fn _assert_exhaustive(p: ApiProtocol) {
match p {
ApiProtocol::AnthropicMessages
| ApiProtocol::OpenAiCompletions
| ApiProtocol::OpenAiResponses
| ApiProtocol::AzureOpenAiResponses
| ApiProtocol::GoogleGenerativeAi
| ApiProtocol::GoogleVertex
| ApiProtocol::BedrockConverseStream => {}
}
}
vec![
ApiProtocol::AnthropicMessages,
ApiProtocol::OpenAiCompletions,
ApiProtocol::OpenAiResponses,
ApiProtocol::AzureOpenAiResponses,
ApiProtocol::GoogleGenerativeAi,
ApiProtocol::GoogleVertex,
ApiProtocol::BedrockConverseStream,
]
}
#[test]
fn default_registry_covers_all_protocols() {
let registry = ProviderRegistry::default();
for api in all_protocols() {
assert!(
registry.resolve(&api).is_some(),
"default registry missing a provider for {api}"
);
}
}
#[test]
fn default_registry_maps_each_protocol_to_its_own_provider() {
let registry = ProviderRegistry::default();
for api in all_protocols() {
let provider = registry.resolve(&api).expect("registered");
assert_eq!(
provider.protocol(),
Some(api),
"protocol {api} resolved to a provider that reports {:?}",
provider.protocol()
);
}
}
#[test]
fn empty_registry_resolves_nothing() {
assert!(ProviderRegistry::new()
.resolve(&ApiProtocol::AnthropicMessages)
.is_none());
}
}