use crate::config::paths::Paths; use crate::config::Config; use crate::providers::anthropic::AnthropicProvider; use crate::providers::base::{ModelInfo, ProviderType}; use crate::providers::ollama::OllamaProvider; use crate::providers::openai::OpenAiProvider; use anyhow::Result; use include_dir::{include_dir, Dir}; use once_cell::sync::Lazy; use serde::{Deserialize, Serialize}; use std::collections::HashMap; use std::path::Path; use std::sync::Mutex; use utoipa::ToSchema; static FIXED_PROVIDERS: Dir = include_dir!("$CARGO_MANIFEST_DIR/src/providers/declarative"); pub fn custom_providers_dir() -> std::path::PathBuf { Paths::config_dir().join("custom_providers") } #[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] #[serde(rename_all = "lowercase")] pub enum ProviderEngine { OpenAI, Ollama, Anthropic, } #[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] pub struct DeclarativeProviderConfig { pub name: String, pub engine: ProviderEngine, pub display_name: String, pub description: Option, #[serde(default)] pub api_key_env: String, pub base_url: String, pub models: Vec, pub headers: Option>, pub timeout_seconds: Option, pub supports_streaming: Option, #[serde(default = "default_requires_auth")] pub requires_auth: bool, #[serde(default)] pub catalog_provider_id: Option, } fn default_requires_auth() -> bool { true } impl DeclarativeProviderConfig { pub fn id(&self) -> &str { &self.name } pub fn display_name(&self) -> &str { &self.display_name } pub fn models(&self) -> &[ModelInfo] { &self.models } } #[derive(Debug, Clone, Serialize, Deserialize, ToSchema)] pub struct LoadedProvider { pub config: DeclarativeProviderConfig, pub is_editable: bool, } static ID_GENERATION_LOCK: Lazy> = Lazy::new(|| Mutex::new(())); pub fn generate_id(display_name: &str) -> String { let _guard = ID_GENERATION_LOCK.lock().unwrap(); let normalized = display_name.to_lowercase().replace(' ', "_"); let base_id = format!("custom_{}", normalized); let custom_dir = custom_providers_dir(); let mut candidate_id = base_id.clone(); let mut counter = 1; while custom_dir.join(format!("{}.json", candidate_id)).exists() { candidate_id = format!("{}_{}", base_id, counter); counter += 1; } candidate_id } pub fn generate_api_key_name(id: &str) -> String { format!("{}_API_KEY", id.to_uppercase()) } #[derive(Debug, Clone)] pub struct CreateCustomProviderParams { pub engine: String, pub display_name: String, pub api_url: String, pub api_key: String, pub models: Vec, pub supports_streaming: Option, pub headers: Option>, pub requires_auth: bool, pub catalog_provider_id: Option, } #[derive(Debug, Clone)] pub struct UpdateCustomProviderParams { pub id: String, pub engine: String, pub display_name: String, pub api_url: String, pub api_key: String, pub models: Vec, pub supports_streaming: Option, pub headers: Option>, pub requires_auth: bool, pub catalog_provider_id: Option, } pub fn create_custom_provider( params: CreateCustomProviderParams, ) -> Result { let id = generate_id(¶ms.display_name); let api_key_env = if params.requires_auth { let api_key_name = generate_api_key_name(&id); let config = Config::global(); config.set_secret(&api_key_name, ¶ms.api_key)?; api_key_name } else { String::new() }; let model_infos: Vec = params .models .into_iter() .map(|name| ModelInfo::new(name, 128000)) .collect(); let provider_config = DeclarativeProviderConfig { name: id.clone(), engine: match params.engine.as_str() { "openai_compatible" => ProviderEngine::OpenAI, "anthropic_compatible" => ProviderEngine::Anthropic, "ollama_compatible" => ProviderEngine::Ollama, _ => return Err(anyhow::anyhow!("Invalid provider type: {}", params.engine)), }, display_name: params.display_name.clone(), description: Some(format!("Custom {} provider", params.display_name)), api_key_env, base_url: params.api_url, models: model_infos, headers: params.headers, timeout_seconds: None, supports_streaming: params.supports_streaming, requires_auth: params.requires_auth, catalog_provider_id: params.catalog_provider_id, }; let custom_providers_dir = custom_providers_dir(); std::fs::create_dir_all(&custom_providers_dir)?; let json_content = serde_json::to_string_pretty(&provider_config)?; let file_path = custom_providers_dir.join(format!("{}.json", id)); std::fs::write(file_path, json_content)?; Ok(provider_config) } pub fn update_custom_provider(params: UpdateCustomProviderParams) -> Result<()> { let loaded_provider = load_provider(¶ms.id)?; let existing_config = loaded_provider.config; let editable = loaded_provider.is_editable; let config = Config::global(); let api_key_env = if params.requires_auth { let api_key_name = if existing_config.api_key_env.is_empty() { generate_api_key_name(¶ms.id) } else { existing_config.api_key_env.clone() }; if !params.api_key.is_empty() { config.set_secret(&api_key_name, ¶ms.api_key)?; } api_key_name } else { String::new() }; if editable { let model_infos: Vec = params .models .into_iter() .map(|name| ModelInfo::new(name, 128000)) .collect(); let updated_config = DeclarativeProviderConfig { name: params.id.clone(), engine: match params.engine.as_str() { "openai_compatible" => ProviderEngine::OpenAI, "anthropic_compatible" => ProviderEngine::Anthropic, "ollama_compatible" => ProviderEngine::Ollama, _ => return Err(anyhow::anyhow!("Invalid provider type: {}", params.engine)), }, display_name: params.display_name, description: existing_config.description, api_key_env, base_url: params.api_url, models: model_infos, headers: match params.headers { Some(h) if h.is_empty() => None, Some(h) => Some(h), None => existing_config.headers, }, timeout_seconds: existing_config.timeout_seconds, supports_streaming: params.supports_streaming, requires_auth: params.requires_auth, catalog_provider_id: params.catalog_provider_id, }; let file_path = custom_providers_dir().join(format!("{}.json", updated_config.name)); let json_content = serde_json::to_string_pretty(&updated_config)?; std::fs::write(file_path, json_content)?; } Ok(()) } pub fn remove_custom_provider(id: &str) -> Result<()> { let config = Config::global(); let api_key_name = generate_api_key_name(id); let _ = config.delete_secret(&api_key_name); let custom_providers_dir = custom_providers_dir(); let file_path = custom_providers_dir.join(format!("{}.json", id)); if file_path.exists() { std::fs::remove_file(file_path)?; } Ok(()) } pub fn load_provider(id: &str) -> Result { let custom_file_path = custom_providers_dir().join(format!("{}.json", id)); if custom_file_path.exists() { let content = std::fs::read_to_string(&custom_file_path)?; let config: DeclarativeProviderConfig = serde_json::from_str(&content)?; return Ok(LoadedProvider { config, is_editable: true, }); } for file in FIXED_PROVIDERS.files() { if file.path().extension().and_then(|s| s.to_str()) != Some("json") { continue; } let content = file .contents_utf8() .ok_or_else(|| anyhow::anyhow!("Failed to read file as UTF-8: {:?}", file.path()))?; let config: DeclarativeProviderConfig = match serde_json::from_str(content) { Ok(config) => config, Err(_) => continue, }; if config.name == id { return Ok(LoadedProvider { config, is_editable: false, }); } } Err(anyhow::anyhow!("Provider not found: {}", id)) } pub fn load_custom_providers(dir: &Path) -> Result> { if !dir.exists() { return Ok(Vec::new()); } std::fs::read_dir(dir)? .filter_map(|entry| { let path = entry.ok()?.path(); (path.extension()? == "json").then_some(path) }) .map(|path| { let content = std::fs::read_to_string(&path)?; serde_json::from_str(&content) .map_err(|e| anyhow::anyhow!("Failed to parse {}: {}", path.display(), e)) }) .collect() } fn load_fixed_providers() -> Result> { let mut res = Vec::new(); for file in FIXED_PROVIDERS.files() { if file.path().extension().and_then(|s| s.to_str()) != Some("json") { continue; } let content = file .contents_utf8() .ok_or_else(|| anyhow::anyhow!("Failed to read file as UTF-8: {:?}", file.path()))?; match serde_json::from_str(content) { Ok(config) => res.push(config), Err(e) => { tracing::warn!( "Skipping invalid declarative provider {:?}: {}", file.path(), e ); } } } Ok(res) } pub fn register_declarative_providers( registry: &mut crate::providers::provider_registry::ProviderRegistry, ) -> Result<()> { let dir = custom_providers_dir(); let custom_providers = load_custom_providers(&dir)?; let fixed_providers = load_fixed_providers()?; for config in fixed_providers { register_declarative_provider(registry, config, ProviderType::Declarative); } for config in custom_providers { register_declarative_provider(registry, config, ProviderType::Custom); } Ok(()) } pub fn register_declarative_provider( registry: &mut crate::providers::provider_registry::ProviderRegistry, config: DeclarativeProviderConfig, provider_type: ProviderType, ) { let config_clone = config.clone(); match config.engine { ProviderEngine::OpenAI => { registry.register_with_name::( &config, provider_type, move |model| OpenAiProvider::from_custom_config(model, config_clone.clone()), ); } ProviderEngine::Ollama => { registry.register_with_name::( &config, provider_type, move |model| OllamaProvider::from_custom_config(model, config_clone.clone()), ); } ProviderEngine::Anthropic => { registry.register_with_name::( &config, provider_type, move |model| AnthropicProvider::from_custom_config(model, config_clone.clone()), ); } } }