use std::collections::HashMap;
use serde::{Deserialize, Serialize};
use crate::protocol::AuthMethod;
use crate::types::{ModelCapabilities, ModelInfo, ModelLimits, ModelPricing};
#[derive(Debug, Clone, PartialEq)]
pub enum ProviderType {
Claude,
OpenAiCompatible,
Generic,
}
impl Serialize for ProviderType {
fn serialize<S: serde::Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let s = match self {
ProviderType::Claude => "claude",
ProviderType::OpenAiCompatible => "open_ai_compatible",
ProviderType::Generic => "generic",
};
serializer.serialize_str(s)
}
}
impl<'de> Deserialize<'de> for ProviderType {
fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let s = String::deserialize(deserializer)?;
match s.as_str() {
"claude" => Ok(ProviderType::Claude),
"open_ai" | "open_ai_compatible" | "local" => Ok(ProviderType::OpenAiCompatible),
"generic" => Ok(ProviderType::Generic),
_ => Err(serde::de::Error::unknown_variant(
&s,
&["claude", "open_ai", "open_ai_compatible", "local", "generic"],
)),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ModelConfig {
pub name: String,
#[serde(default)]
pub display_name: Option<String>,
#[serde(default)]
pub capabilities: ModelCapabilities,
#[serde(default)]
pub pricing: ModelPricing,
#[serde(default)]
pub limits: ModelLimits,
}
impl From<ModelConfig> for ModelInfo {
fn from(cfg: ModelConfig) -> Self {
ModelInfo {
name: cfg.name,
display_name: cfg.display_name,
provider: None,
capabilities: cfg.capabilities,
pricing: cfg.pricing,
limits: cfg.limits,
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
#[serde(rename_all = "snake_case")]
pub enum ApiProtocol {
#[default]
ChatCompletions,
Responses,
AnthropicMessages,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderDefinition {
pub provider_type: ProviderType,
#[serde(skip_serializing_if = "Option::is_none")]
pub api_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub base_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub headers: Option<HashMap<String, String>>,
#[serde(default)]
pub protocol: ApiProtocol,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub auth_method: Option<AuthMethod>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub anthropic_version: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub anthropic_beta: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_thinking_type: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default_effort: Option<String>,
pub models: Vec<ModelConfig>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteRule {
pub model: String,
pub provider: Option<String>,
#[serde(default)]
pub temperature: Option<f32>,
#[serde(default)]
pub max_tokens: Option<usize>,
#[serde(default)]
pub fallback: Vec<FallbackEntry>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FallbackEntry {
pub model: String,
pub provider: Option<String>,
#[serde(default)]
pub condition: FallbackCondition,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub enum FallbackCondition {
#[serde(rename = "always")]
#[default]
Always,
#[serde(rename = "rate_limit_only")]
RateLimitOnly,
#[serde(rename = "error_status")]
ErrorStatus(Vec<u16>),
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ProviderConfig {
pub default_model: Option<String>,
pub providers: HashMap<String, ProviderDefinition>,
#[serde(default)]
pub routing: HashMap<String, RouteRule>,
}
impl ProviderConfig {
pub fn from_json(json: &str) -> Result<Self, crate::error::ProviderError> {
let config: Self = serde_json::from_str(json)
.map_err(|e| crate::error::ProviderError::Config(e.to_string()))?;
config.validate()
}
pub fn from_yaml(yaml: &str) -> Result<Self, crate::error::ProviderError> {
let yaml_interpolated = Self::interpolate_env(yaml);
let config: Self = serde_yaml::from_str(&yaml_interpolated)
.map_err(|e| crate::error::ProviderError::Config(e.to_string()))?;
config.validate()
}
pub async fn from_file(
path: impl AsRef<std::path::Path>,
) -> Result<Self, crate::error::ProviderError> {
let content = tokio::fs::read_to_string(path.as_ref())
.await
.map_err(|e| crate::error::ProviderError::Config(format!("读取配置文件失败: {e}")))?;
Self::from_json(&content)
}
pub async fn from_yaml_file(
path: impl AsRef<std::path::Path>,
) -> Result<Self, crate::error::ProviderError> {
let content = tokio::fs::read_to_string(path.as_ref())
.await
.map_err(|e| crate::error::ProviderError::Config(format!("读取配置文件失败: {e}")))?;
Self::from_yaml(&content)
}
fn interpolate_env(input: &str) -> String {
let mut result = input.to_string();
let re = regex_lite::Regex::new(r"\$\{([^:}]+)(?::-(.*?))?\}").ok();
if let Some(re) = re {
for caps in re.captures_iter(input) {
let var_name = caps.get(1).map(|m| m.as_str()).unwrap_or("");
let default_val = caps.get(2).map(|m| m.as_str());
let value = std::env::var(var_name)
.ok()
.or_else(|| default_val.map(|s| s.to_string()))
.unwrap_or_default();
result = result.replace(caps.get(0).map(|m| m.as_str()).unwrap_or(""), &value);
}
}
result
}
fn validate(self) -> Result<Self, crate::error::ProviderError> {
for (name, def) in &self.providers {
if def.models.is_empty() {
return Err(crate::error::ProviderError::Config(format!(
"Provider '{name}' 没有定义任何模型"
)));
}
if def.provider_type == ProviderType::Claude
&& def.api_key.as_ref().is_none_or(|k| k.is_empty())
{
return Err(crate::error::ProviderError::Config(format!(
"Provider '{name}' 缺少 api_key"
)));
}
}
Ok(self)
}
pub fn collect_models(&self) -> Vec<ModelInfo> {
let mut models = Vec::new();
for (provider_name, def) in &self.providers {
for mc in &def.models {
let mut info = ModelInfo::from(mc.clone());
info.provider = Some(provider_name.clone());
models.push(info);
}
}
models
}
}
pub trait ConfigWatcher: Send + Sync {
fn watch(&self) -> futures::stream::BoxStream<'static, ProviderConfig>;
}
#[cfg(test)]
mod tests {
use super::*;
fn valid_openai_json() -> &'static str {
r#"{
"default_model": "gpt-4",
"providers": {
"openai": {
"provider_type": "open_ai",
"api_key": "sk-test",
"models": [
{
"name": "gpt-4",
"capabilities": {
"context_window": 128000,
"max_output_tokens": 4096,
"supports_tool_calling": true
}
}
]
}
}
}"#
}
#[test]
fn test_from_json_valid() {
let config = ProviderConfig::from_json(valid_openai_json()).unwrap();
assert_eq!(config.default_model.unwrap(), "gpt-4");
assert!(config.providers.contains_key("openai"));
assert_eq!(config.providers["openai"].provider_type, ProviderType::OpenAiCompatible);
assert_eq!(config.providers["openai"].models.len(), 1);
assert_eq!(config.providers["openai"].models[0].name, "gpt-4");
}
#[test]
fn test_from_json_invalid_syntax() {
let result = ProviderConfig::from_json("not valid json");
assert!(result.is_err());
}
#[test]
fn test_from_json_empty_models() {
let json = r#"{
"providers": {
"test": {
"provider_type": "open_ai",
"api_key": "sk-test",
"models": []
}
}
}"#;
let result = ProviderConfig::from_json(json);
assert!(result.is_err());
let err = result.unwrap_err();
assert!(format!("{}", err).contains("没有定义任何模型"));
}
#[test]
fn test_from_json_missing_api_key() {
let json = r#"{
"providers": {
"test": {
"provider_type": "claude",
"models": [{"name": "claude-3", "capabilities": {"context_window": 200000, "max_output_tokens": 4096}}]
}
}
}"#;
let result = ProviderConfig::from_json(json);
assert!(result.is_err());
let err = format!("{}", result.unwrap_err());
assert!(err.contains("api_key"), "error: {}", err);
}
#[test]
fn test_from_json_no_api_key_needed() {
let json = r#"{
"providers": {
"custom": {
"provider_type": "open_ai_compatible",
"base_url": "http://localhost:11434",
"models": [{"name": "local-model", "capabilities": {"context_window": 4096, "max_output_tokens": 2048}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
assert!(config.providers.contains_key("custom"));
}
#[test]
fn test_from_yaml_valid() {
let yaml = r#"
default_model: gpt-4
providers:
openai:
provider_type: open_ai
api_key: sk-test
models:
- name: gpt-4
capabilities:
context_window: 128000
max_output_tokens: 4096
"#;
let config = ProviderConfig::from_yaml(yaml).unwrap();
assert_eq!(config.default_model.unwrap(), "gpt-4");
assert!(config.providers.contains_key("openai"));
}
#[test]
fn test_from_yaml_invalid() {
let result = ProviderConfig::from_yaml(": bad yaml : :");
assert!(result.is_err());
}
#[test]
fn test_collect_models() {
let config = ProviderConfig::from_json(valid_openai_json()).unwrap();
let models = config.collect_models();
assert_eq!(models.len(), 1);
assert_eq!(models[0].name, "gpt-4");
assert_eq!(models[0].provider.as_deref(), Some("openai"));
assert!(models[0].capabilities.supports_tool_calling);
}
#[test]
fn test_collect_models_multiple_providers() {
let json = r#"{
"providers": {
"p1": {
"provider_type": "open_ai",
"api_key": "k1",
"models": [{"name": "m1", "capabilities": {"context_window": 100, "max_output_tokens": 100}}]
},
"p2": {
"provider_type": "open_ai_compatible",
"models": [{"name": "m2", "capabilities": {"context_window": 100, "max_output_tokens": 100}}, {"name": "m3", "capabilities": {"context_window": 100, "max_output_tokens": 100}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let models = config.collect_models();
assert_eq!(models.len(), 3);
assert!(models.iter().any(|m| m.provider.as_deref() == Some("p1")));
assert!(models.iter().any(|m| m.provider.as_deref() == Some("p2")));
}
#[test]
fn test_interpolate_env_no_vars() {
let input = "hello world";
let result = ProviderConfig::interpolate_env(input);
assert_eq!(result, "hello world");
}
#[test]
fn test_interpolate_env_with_default() {
let input = r#"api_key: ${MY_KEY:-default_key}"#;
let result = ProviderConfig::interpolate_env(input);
assert_eq!(result, "api_key: default_key");
}
#[test]
fn test_model_config_to_model_info() {
let cfg = ModelConfig {
name: "gpt-4".into(),
display_name: Some("GPT-4".into()),
capabilities: ModelCapabilities {
context_window: 8192,
max_output_tokens: 4096,
..Default::default()
},
pricing: ModelPricing {
input_per_million: 30.0,
output_per_million: 60.0,
..Default::default()
},
limits: ModelLimits::default(),
};
let info: ModelInfo = cfg.into();
assert_eq!(info.name, "gpt-4");
assert_eq!(info.display_name.unwrap(), "GPT-4");
assert_eq!(info.capabilities.context_window, 8192);
assert_eq!(info.pricing.input_per_million, 30.0);
assert!(info.provider.is_none());
}
#[test]
fn test_route_rule_default_fallback() {
let rule = RouteRule {
model: "gpt-4".into(),
provider: Some("openai".into()),
temperature: None,
max_tokens: None,
fallback: vec![],
};
assert_eq!(rule.model, "gpt-4");
}
#[test]
fn test_fallback_condition_default() {
let cond: FallbackCondition = Default::default();
assert!(matches!(cond, FallbackCondition::Always));
}
#[test]
fn test_config_routing() {
let json = r#"{
"default_model": "gpt-4",
"providers": {
"openai": {
"provider_type": "open_ai",
"api_key": "sk-test",
"models": [{"name": "gpt-4", "capabilities": {"context_window": 100, "max_output_tokens": 100}}]
}
},
"routing": {
"chat": {
"model": "gpt-4",
"provider": "openai"
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
assert!(config.routing.contains_key("chat"));
assert_eq!(config.routing["chat"].model, "gpt-4");
}
#[test]
fn test_provider_type_serde() {
assert_eq!(serde_json::to_string(&ProviderType::Claude).unwrap(), r#""claude""#);
assert_eq!(
serde_json::to_string(&ProviderType::OpenAiCompatible).unwrap(),
r#""open_ai_compatible""#
);
assert_eq!(serde_json::to_string(&ProviderType::Generic).unwrap(), r#""generic""#);
let deserialized: ProviderType = serde_json::from_str(r#""open_ai""#).unwrap();
assert_eq!(deserialized, ProviderType::OpenAiCompatible);
}
#[test]
fn test_parse_anthropic_messages_protocol() {
let json = r#"{
"providers": {
"anthropic": {
"provider_type": "claude",
"api_key": "sk-ant-xxx",
"protocol": "anthropic_messages",
"models": [{"name": "claude-3-opus", "capabilities": {"context_window": 200000, "max_output_tokens": 4096}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
assert_eq!(config.providers["anthropic"].protocol, ApiProtocol::AnthropicMessages);
}
#[test]
fn test_parse_open_ai_compatible_backward_compat() {
let json = r#"{
"providers": {
"deepseek": {
"provider_type": "open_ai_compatible",
"api_key": "sk-ds",
"models": [{"name": "deepseek-chat", "capabilities": {"context_window": 64000, "max_output_tokens": 8192}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
assert_eq!(config.providers["deepseek"].provider_type, ProviderType::OpenAiCompatible);
}
#[test]
fn test_parse_auth_method_bearer() {
let json = r#"{
"providers": {
"custom": {
"provider_type": "open_ai_compatible",
"base_url": "https://custom.example.com",
"auth_method": {"type": "bearer", "token": "sk-xxx"},
"models": [{"name": "model-1", "capabilities": {"context_window": 4096, "max_output_tokens": 1024}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let auth = config.providers["custom"].auth_method.as_ref().unwrap();
assert_eq!(*auth, AuthMethod::Bearer { token: "sk-xxx".into() });
}
#[test]
fn test_parse_auth_method_api_key() {
let json = r#"{
"providers": {
"custom": {
"provider_type": "open_ai_compatible",
"base_url": "https://custom.example.com",
"auth_method": {"type": "api_key", "header_name": "x-api-key", "key": "sk-ant-xxx"},
"models": [{"name": "model-1", "capabilities": {"context_window": 4096, "max_output_tokens": 1024}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let auth = config.providers["custom"].auth_method.as_ref().unwrap();
assert_eq!(
*auth,
AuthMethod::ApiKey { header_name: "x-api-key".into(), key: "sk-ant-xxx".into() }
);
}
#[test]
fn test_parse_no_auth_method_defaults_none() {
let json = r#"{
"providers": {
"openai": {
"provider_type": "open_ai",
"api_key": "sk-test",
"models": [{"name": "gpt-4", "capabilities": {"context_window": 128000, "max_output_tokens": 4096}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
assert!(config.providers["openai"].auth_method.is_none());
}
#[test]
fn test_parse_generic_provider_type() {
let json = r#"{
"providers": {
"my_custom": {
"provider_type": "generic",
"base_url": "https://my-custom.api.com",
"auth_method": {"type": "bearer", "token": "sk-custom"},
"protocol": "anthropic_messages",
"models": [{"name": "custom-model", "capabilities": {"context_window": 4096, "max_output_tokens": 1024}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
assert_eq!(config.providers["my_custom"].provider_type, ProviderType::Generic);
}
#[test]
fn test_api_protocol_default() {
assert_eq!(ApiProtocol::default(), ApiProtocol::ChatCompletions);
}
#[test]
fn test_api_protocol_serde() {
assert_eq!(
serde_json::to_string(&ApiProtocol::ChatCompletions).unwrap(),
r#""chat_completions""#
);
assert_eq!(serde_json::to_string(&ApiProtocol::Responses).unwrap(), r#""responses""#);
assert_eq!(
serde_json::to_string(&ApiProtocol::AnthropicMessages).unwrap(),
r#""anthropic_messages""#
);
}
#[test]
fn test_provider_type_generic_serde() {
assert_eq!(serde_json::to_string(&ProviderType::Generic).unwrap(), r#""generic""#);
}
#[test]
fn test_auth_method_serde_roundtrip() {
let original = AuthMethod::Bearer { token: "sk-test".into() };
let json = serde_json::to_string(&original).unwrap();
let deserialized: AuthMethod = serde_json::from_str(&json).unwrap();
assert_eq!(original, deserialized);
let original = AuthMethod::ApiKey { header_name: "x-custom".into(), key: "key-123".into() };
let json = serde_json::to_string(&original).unwrap();
let deserialized: AuthMethod = serde_json::from_str(&json).unwrap();
assert_eq!(original, deserialized);
}
#[test]
fn test_provider_definition_new_fields_roundtrip() {
let json = r#"{
"providers": {
"anthropic": {
"provider_type": "claude",
"api_key": "sk-ant-xxx",
"anthropic_version": "2023-06-01",
"anthropic_beta": ["prompt-caching-2025-02-19", "tools-2025-04-01"],
"default_thinking_type": "extended",
"default_effort": "high",
"models": [{"name": "claude-3-opus", "capabilities": {"context_window": 200000, "max_output_tokens": 4096}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let def = &config.providers["anthropic"];
assert_eq!(def.anthropic_version.as_deref(), Some("2023-06-01"));
let expected_beta: &[String] =
&["prompt-caching-2025-02-19".into(), "tools-2025-04-01".into()];
assert_eq!(def.anthropic_beta.as_deref(), Some(expected_beta));
assert_eq!(def.default_thinking_type.as_deref(), Some("extended"));
assert_eq!(def.default_effort.as_deref(), Some("high"));
let serialized = serde_json::to_string_pretty(&config).unwrap();
let deserialized: ProviderConfig = serde_json::from_str(&serialized).unwrap();
let def2 = &deserialized.providers["anthropic"];
assert_eq!(def2.anthropic_version, def.anthropic_version);
assert_eq!(def2.anthropic_beta, def.anthropic_beta);
assert_eq!(def2.default_thinking_type, def.default_thinking_type);
assert_eq!(def2.default_effort, def.default_effort);
}
#[test]
fn test_provider_definition_without_new_fields_backward_compat() {
let json = r#"{
"providers": {
"openai": {
"provider_type": "open_ai",
"api_key": "sk-test",
"models": [{"name": "gpt-4", "capabilities": {"context_window": 128000, "max_output_tokens": 4096}}]
}
}
}"#;
let config = ProviderConfig::from_json(json).unwrap();
let def = &config.providers["openai"];
assert!(def.anthropic_version.is_none());
assert!(def.anthropic_beta.is_none());
assert!(def.default_thinking_type.is_none());
assert!(def.default_effort.is_none());
}
}