use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use crate::config::{ApiProtocol, ConfigWatcher, ProviderConfig, ProviderDefinition, ProviderType};
use crate::error::{ProviderError, RetryStrategy};
use crate::protocol::{
AuthMethod, OpenAiChatAdapter, ProtocolAdapter, openai_responses::OpenAiResponsesAdapter,
};
use crate::providers::GenericProvider;
use crate::router::ProviderRouter;
use crate::traits::LlmProvider;
use crate::types::ModelInfo;
#[cfg(feature = "anthropic")]
use crate::protocol::anthropic::AnthropicMessagesAdapter;
pub struct ProviderBuilder {
config: Option<ProviderConfig>,
config_path: Option<PathBuf>,
config_watcher: Option<Box<dyn ConfigWatcher>>,
retry_strategy: RetryStrategy,
http_client: Option<reqwest::Client>,
key_source: Option<std::sync::Arc<dyn crate::key_source::KeySource>>,
}
impl Default for ProviderBuilder {
fn default() -> Self {
Self::new()
}
}
impl ProviderBuilder {
pub fn new() -> Self {
Self {
config: None,
config_path: None,
config_watcher: None,
retry_strategy: RetryStrategy::default(),
http_client: None,
key_source: None,
}
}
pub fn with_config(mut self, config: ProviderConfig) -> Self {
self.config = Some(config);
self
}
pub fn with_config_file(mut self, path: impl Into<PathBuf>) -> Self {
self.config_path = Some(path.into());
self
}
pub fn with_yaml_config_file(mut self, path: impl Into<PathBuf>) -> Self {
self.config_path = Some(path.into());
self
}
pub fn with_retry(mut self, strategy: RetryStrategy) -> Self {
self.retry_strategy = strategy;
self
}
pub fn with_config_watcher(mut self, watcher: impl ConfigWatcher + 'static) -> Self {
self.config_watcher = Some(Box::new(watcher));
self
}
pub fn with_http_client(mut self, client: reqwest::Client) -> Self {
self.http_client = Some(client);
self
}
pub fn with_key_source(
mut self,
key_source: std::sync::Arc<dyn crate::key_source::KeySource>,
) -> Self {
self.key_source = Some(key_source);
self
}
pub async fn build(self) -> Result<ProviderRouter, ProviderError> {
let config = if let Some(cfg) = self.config {
cfg
} else if let Some(path) = self.config_path {
let ext = path.extension().and_then(|e| e.to_str()).unwrap_or("json");
match ext {
"yaml" | "yml" => ProviderConfig::from_yaml_file(path).await?,
_ => ProviderConfig::from_file(path).await?,
}
} else {
return Err(ProviderError::Config("必须提供 ProviderConfig 或配置文件路径".to_owned()));
};
let http_client = match self.http_client {
Some(client) => client,
None => reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(120))
.pool_idle_timeout(std::time::Duration::from_secs(90))
.build()
.map_err(|e| {
ProviderError::Config(format!("Failed to build HTTP client: {}", e))
})?,
};
let mut providers: HashMap<String, Box<dyn LlmProvider>> = HashMap::new();
let ks = self.key_source.clone();
for (name, def) in &config.providers {
let provider: Box<dyn LlmProvider> = match def.provider_type {
#[cfg(feature = "claude")]
ProviderType::Claude => {
let adapter = build_anthropic_adapter(def);
build_generic_provider(name, def, &http_client, adapter, ks.clone())
}
#[cfg(not(feature = "claude"))]
ProviderType::Claude => {
return Err(ProviderError::Config("claude feature 未启用".to_owned()));
}
#[cfg(feature = "openai_compatible")]
ProviderType::OpenAiCompatible => {
let adapter = resolve_adapter_for_openai(def)?;
build_generic_provider(name, def, &http_client, adapter, ks.clone())
}
#[cfg(not(feature = "openai_compatible"))]
ProviderType::OpenAiCompatible => {
return Err(ProviderError::Config(
"openai_compatible feature 未启用".to_owned(),
));
}
ProviderType::Generic => {
let adapter = resolve_adapter_for_generic(def)?;
build_generic_provider(name, def, &http_client, adapter, ks.clone())
}
};
providers.insert(name.clone(), provider);
}
let models = config.collect_models();
let default_model = config
.default_model
.or_else(|| models.first().map(|m| m.name.clone()))
.unwrap_or_default();
Ok(ProviderRouter::new(providers, models, config.routing, default_model, self.key_source))
}
}
fn build_generic_provider(
name: &str,
def: &ProviderDefinition,
client: &reqwest::Client,
adapter: Arc<dyn ProtocolAdapter>,
key_source: Option<Arc<dyn crate::key_source::KeySource>>,
) -> Box<dyn LlmProvider> {
let auth = resolve_auth(def);
let base_url = resolve_base_url(def);
let models: Vec<ModelInfo> = def.models.iter().map(|m| ModelInfo::from(m.clone())).collect();
let extra_headers: Vec<(String, String)> = def
.headers
.as_ref()
.map(|h| h.iter().map(|(k, v)| (k.clone(), v.clone())).collect())
.unwrap_or_default();
Box::new(GenericProvider::with_headers_and_key_source(
name.to_owned(),
adapter,
auth,
base_url,
models,
client.clone(),
extra_headers,
key_source,
))
}
fn resolve_auth(def: &ProviderDefinition) -> AuthMethod {
if let Some(ref auth) = def.auth_method {
return auth.clone();
}
if let Some(ref key) = def.api_key {
match def.provider_type {
ProviderType::Claude => {
AuthMethod::ApiKey { header_name: "x-api-key".to_owned(), key: key.clone() }
}
_ => AuthMethod::Bearer { token: key.clone() },
}
} else {
AuthMethod::None
}
}
fn resolve_base_url(def: &ProviderDefinition) -> String {
match def.provider_type {
ProviderType::OpenAiCompatible => {
def.base_url.clone().unwrap_or_else(|| "https://api.openai.com/v1".to_owned())
}
ProviderType::Claude => {
def.base_url.clone().unwrap_or_else(|| "https://api.anthropic.com/v1".to_owned())
}
ProviderType::Generic => {
def.base_url.clone().unwrap_or_else(|| "https://api.openai.com/v1".to_owned())
}
}
}
fn resolve_adapter_for_openai(
def: &ProviderDefinition,
) -> Result<Arc<dyn ProtocolAdapter>, ProviderError> {
match def.protocol {
ApiProtocol::ChatCompletions => Ok(Arc::new(OpenAiChatAdapter::new())),
ApiProtocol::Responses => Ok(Arc::new(OpenAiResponsesAdapter::new())),
ApiProtocol::AnthropicMessages => Err(ProviderError::Config(format!(
"AnthropicMessages protocol is not compatible with '{:?}' provider type",
def.provider_type
))),
}
}
fn resolve_adapter_for_generic(
def: &ProviderDefinition,
) -> Result<Arc<dyn ProtocolAdapter>, ProviderError> {
match def.protocol {
ApiProtocol::ChatCompletions => Ok(Arc::new(OpenAiChatAdapter::new())),
ApiProtocol::Responses => Ok(Arc::new(OpenAiResponsesAdapter::new())),
ApiProtocol::AnthropicMessages => {
#[cfg(feature = "anthropic")]
{
Ok(build_anthropic_adapter(def))
}
#[cfg(not(feature = "anthropic"))]
{
Err(ProviderError::Config(
"AnthropicMessages protocol requires the 'anthropic' feature".to_owned(),
))
}
}
}
}
#[cfg(feature = "anthropic")]
fn build_anthropic_adapter(def: &ProviderDefinition) -> Arc<dyn ProtocolAdapter> {
use crate::types::{EffortLevel, ThinkingType};
let mut adapter = AnthropicMessagesAdapter::new();
if let Some(beta) = &def.anthropic_beta {
adapter = adapter.with_beta_headers(beta.clone());
}
if let Some(version) = &def.anthropic_version {
adapter = adapter.with_anthropic_version(version.clone());
}
if let Some(effort_str) = &def.default_effort {
if let Some(effort) = match effort_str.as_str() {
"low" => Some(EffortLevel::Low),
"medium" => Some(EffortLevel::Medium),
"high" => Some(EffortLevel::High),
"xhigh" | "x_high" => Some(EffortLevel::XHigh),
"max" => Some(EffortLevel::Max),
_ => None,
} {
adapter = adapter.with_effort_level(effort);
}
}
if let Some(thinking_str) = &def.default_thinking_type {
if let Some(thinking_type) = match thinking_str.as_str() {
"enabled" | "extended" => Some(ThinkingType::Enabled { budget_tokens: None }),
"disabled" => Some(ThinkingType::Disabled),
"adaptive" => Some(ThinkingType::Adaptive),
_ => None,
} {
adapter = adapter.with_default_thinking_type(thinking_type);
}
}
Arc::new(adapter)
}
#[cfg(test)]
mod tests {
use crate::ProviderBuilder;
use crate::config::ProviderConfig;
use crate::router::RouteContext;
use crate::types::{CompletionRequest, Message, RequestOptions};
#[cfg(feature = "openai_compatible")]
#[tokio::test]
async fn test_builder_openai_compatible() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "test-model",
"providers": {{
"custom": {{
"provider_type": "open_ai_compatible",
"base_url": "{base_url}",
"models": [{{"name": "test-model", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "Hello!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 10}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
assert!(!router.model_registry().is_empty());
let route_ctx = RouteContext { model: Some("test-model".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("test-model", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "Hello!");
}
#[cfg(all(feature = "openai_compatible", feature = "claude"))]
#[tokio::test]
async fn test_builder_multi_provider() {
let json = r#"{
"providers": {
"openai": {
"provider_type": "open_ai",
"api_key": "sk-test",
"models": [{"name": "gpt-4", "capabilities": {"context_window": 100, "max_output_tokens": 100}}]
},
"claude": {
"provider_type": "claude",
"api_key": "sk-ant",
"models": [{"name": "claude-3", "capabilities": {"context_window": 200, "max_output_tokens": 200}}]
},
"custom": {
"provider_type": "open_ai_compatible",
"base_url": "http://localhost:11434/v1",
"models": [{"name": "custom-model", "capabilities": {"context_window": 100, "max_output_tokens": 100}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
assert!(router.model_registry().len() >= 3);
}
#[cfg(feature = "openai_compatible")]
#[tokio::test]
async fn test_builder_without_openai_feature() {
let json = r#"{
"providers": {
"custom": {
"provider_type": "open_ai_compatible",
"models": [{"name": "m1", "capabilities": {"context_window": 100, "max_output_tokens": 100}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
assert!(!router.model_registry().is_empty());
}
#[tokio::test]
async fn test_builder_generic_chat() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "test-model",
"providers": {{
"my-provider": {{
"provider_type": "generic",
"base_url": "{base_url}",
"protocol": "chat_completions",
"auth_method": {{"type": "bearer", "token": "sk-test"}},
"models": [{{"name": "test-model", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "generic ok!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 5, "completion_tokens": 10}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let route_ctx = RouteContext { model: Some("test-model".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("test-model", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "generic ok!");
}
#[tokio::test]
async fn test_builder_generic_responses() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "test-model",
"providers": {{
"resp-provider": {{
"provider_type": "generic",
"base_url": "{base_url}",
"protocol": "responses",
"api_key": "sk-test",
"models": [{{"name": "test-model", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/responses"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"id": "resp_1",
"status": "completed",
"output": [{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Responses work!"}]
}],
"usage": {"input_tokens": 5, "output_tokens": 10}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let route_ctx = RouteContext { model: Some("test-model".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("test-model", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "Responses work!");
}
#[tokio::test]
async fn test_builder_custom_headers() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "test-model",
"providers": {{
"hdr-provider": {{
"provider_type": "open_ai_compatible",
"base_url": "{base_url}",
"api_key": "sk-test",
"headers": {{"X-Custom": "custom-val"}},
"models": [{{"name": "test-model", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test"))
.and(wiremock::matchers::header("X-Custom", "custom-val"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "headers work!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 2}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let route_ctx = RouteContext { model: Some("test-model".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("test-model", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "headers work!");
}
#[tokio::test]
async fn test_builder_local_provider() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "test-model",
"providers": {{
"local-provider": {{
"provider_type": "local",
"base_url": "{base_url}",
"models": [{{"name": "test-model", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "local works!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 2}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let route_ctx = RouteContext { model: Some("test-model".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("test-model", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "local works!");
}
#[tokio::test]
async fn test_builder_openai_backward_compat() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "gpt-4",
"providers": {{
"openai": {{
"provider_type": "open_ai",
"base_url": "{base_url}",
"api_key": "sk-test",
"models": [{{"name": "gpt-4", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "backward compat!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 2}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let route_ctx = RouteContext { model: Some("gpt-4".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("gpt-4", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "backward compat!");
}
#[cfg(feature = "openai_compatible")]
#[tokio::test]
async fn test_builder_openai_compatible_responses() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "test-model",
"providers": {{
"custom": {{
"provider_type": "open_ai_compatible",
"base_url": "{base_url}",
"protocol": "responses",
"api_key": "sk-test",
"models": [{{"name": "test-model", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/responses"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"id": "resp_1",
"status": "completed",
"output": [{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": "Responses via openai_compatible!"}]
}],
"usage": {"input_tokens": 5, "output_tokens": 10}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let route_ctx = RouteContext { model: Some("test-model".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("test-model", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "Responses via openai_compatible!");
}
#[tokio::test]
async fn test_builder_routing() {
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
let mock_a = MockServer::start().await;
let mock_b = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "model-a",
"providers": {{
"provider-a": {{
"provider_type": "generic",
"base_url": "{base_url_a}",
"protocol": "chat_completions",
"api_key": "sk-a",
"models": [{{"name": "model-a", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}},
"provider-b": {{
"provider_type": "generic",
"base_url": "{base_url_b}",
"protocol": "chat_completions",
"api_key": "sk-b",
"models": [{{"name": "model-b", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url_a = mock_a.uri(),
base_url_b = mock_b.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-a"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "from A"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1}
})))
.mount(&mock_a)
.await;
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-b"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "from B"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1}
})))
.mount(&mock_b)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new().with_config(config).build().await.unwrap();
let ctx_a = RouteContext { model: Some("model-a".into()), ..Default::default() };
let resp_a = router
.complete(
&ctx_a,
CompletionRequest::new("model-a", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp_a.content.unwrap_or_default(), "from A");
let ctx_b = RouteContext { model: Some("model-b".into()), ..Default::default() };
let resp_b = router
.complete(
&ctx_b,
CompletionRequest::new("model-b", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp_b.content.unwrap_or_default(), "from B");
}
#[tokio::test]
async fn test_builder_key_source_used_in_request() {
use crate::key_source::KeySource;
use async_trait::async_trait;
use std::sync::Arc;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
struct StaticKeySource(String);
#[async_trait]
impl KeySource for StaticKeySource {
async fn get_api_key(&self) -> Result<String, crate::error::ProviderError> {
Ok(self.0.clone())
}
}
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "m1",
"providers": {{
"p1": {{
"provider_type": "open_ai_compatible",
"base_url": "{base_url}",
"models": [{{"name": "m1", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer sk-from-source",
))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "key_source works!"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 2}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new()
.with_config(config)
.with_key_source(Arc::new(StaticKeySource("sk-from-source".into())))
.build()
.await
.unwrap();
let key = router.resolve_api_key("m1").await;
assert_eq!(key.unwrap().unwrap(), "sk-from-source");
let route_ctx = RouteContext { model: Some("m1".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("m1", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "key_source works!");
}
#[tokio::test]
async fn test_key_source_overrides_static_api_key() {
use crate::key_source::KeySource;
use async_trait::async_trait;
use std::sync::Arc;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
struct StaticKeySource(String);
#[async_trait]
impl KeySource for StaticKeySource {
async fn get_api_key(&self) -> Result<String, crate::error::ProviderError> {
Ok(self.0.clone())
}
}
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "m1",
"providers": {{
"p1": {{
"provider_type": "open_ai_compatible",
"base_url": "{base_url}",
"api_key": "sk-baked-from-env",
"models": [{{"name": "m1", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer sk-from-key-source",
))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "dynamic wins"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new()
.with_config(config)
.with_key_source(Arc::new(StaticKeySource("sk-from-key-source".into())))
.build()
.await
.unwrap();
let route_ctx = RouteContext { model: Some("m1".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("m1", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "dynamic wins");
}
#[tokio::test]
async fn test_key_source_empty_falls_back_to_static() {
use crate::key_source::KeySource;
use async_trait::async_trait;
use std::sync::Arc;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
struct EmptyKeySource;
#[async_trait]
impl KeySource for EmptyKeySource {
async fn get_api_key(&self) -> Result<String, crate::error::ProviderError> {
Ok(String::new())
}
}
let mock_server = MockServer::start().await;
let json = format!(
r#"{{
"default_model": "m1",
"providers": {{
"p1": {{
"provider_type": "open_ai_compatible",
"base_url": "{base_url}",
"api_key": "sk-baked-fallback",
"models": [{{"name": "m1", "capabilities": {{"context_window": 100, "max_output_tokens": 100}}}}]
}}
}}
}}"#,
base_url = mock_server.uri()
);
Mock::given(method("POST"))
.and(path("/chat/completions"))
.and(wiremock::matchers::header(
"Authorization",
"Bearer sk-baked-fallback",
))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"choices": [{"message": {"content": "fallback wins"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1}
})))
.mount(&mock_server)
.await;
let config = ProviderConfig::from_json(&json).unwrap();
let router = ProviderBuilder::new()
.with_config(config)
.with_key_source(Arc::new(EmptyKeySource))
.build()
.await
.unwrap();
let route_ctx = RouteContext { model: Some("m1".into()), ..Default::default() };
let resp = router
.complete(
&route_ctx,
CompletionRequest::new("m1", vec![Message::user("Hi")]),
RequestOptions::default(),
)
.await
.unwrap();
assert_eq!(resp.content.unwrap_or_default(), "fallback wins");
}
}