Declarative providers (#5084)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
Co-authored-by: Michael Neale <michael.neale@gmail.com>
This commit is contained in:
Douwe Osinga
2025-10-15 09:48:14 -04:00
committed by GitHub
parent 925a042cb1
commit 9251da4314
28 changed files with 1219 additions and 802 deletions
-212
View File
@@ -1,212 +0,0 @@
use crate::config::paths::Paths;
use crate::config::Config;
use crate::model::ModelConfig;
use crate::providers::anthropic::AnthropicProvider;
use crate::providers::base::ModelInfo;
use crate::providers::ollama::OllamaProvider;
use crate::providers::openai::OpenAiProvider;
use anyhow::Result;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
pub fn custom_providers_dir() -> std::path::PathBuf {
Paths::config_dir().join("custom_providers")
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ProviderEngine {
OpenAI,
Ollama,
Anthropic,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CustomProviderConfig {
pub name: String,
pub engine: ProviderEngine,
pub display_name: String,
pub description: Option<String>,
pub api_key_env: String,
pub base_url: String,
pub models: Vec<ModelInfo>,
pub headers: Option<HashMap<String, String>>,
pub timeout_seconds: Option<u64>,
pub supports_streaming: Option<bool>,
}
impl CustomProviderConfig {
pub fn id(&self) -> &str {
&self.name
}
pub fn display_name(&self) -> &str {
&self.display_name
}
pub fn models(&self) -> &[ModelInfo] {
&self.models
}
pub fn generate_id(display_name: &str) -> String {
format!("custom_{}", display_name.to_lowercase().replace(' ', "_"))
}
pub fn generate_api_key_name(id: &str) -> String {
format!("{}_API_KEY", id.to_uppercase())
}
pub fn create_and_save(
provider_type: &str,
display_name: String,
api_url: String,
api_key: String,
models: Vec<String>,
supports_streaming: Option<bool>,
) -> Result<Self> {
let id = Self::generate_id(&display_name);
let api_key_name = Self::generate_api_key_name(&id);
let config = Config::global();
config.set_secret(&api_key_name, serde_json::Value::String(api_key))?;
let model_infos: Vec<ModelInfo> = models
.into_iter()
.map(|name| ModelInfo::new(name, 128000))
.collect();
let provider_config = CustomProviderConfig {
name: id.clone(),
engine: match provider_type {
"openai_compatible" => ProviderEngine::OpenAI,
"anthropic_compatible" => ProviderEngine::Anthropic,
"ollama_compatible" => ProviderEngine::Ollama,
_ => return Err(anyhow::anyhow!("Invalid provider type: {}", provider_type)),
},
display_name: display_name.clone(),
description: Some(format!("Custom {} provider", display_name)),
api_key_env: api_key_name,
base_url: api_url,
models: model_infos,
headers: None,
timeout_seconds: None,
supports_streaming,
};
// save to JSON file
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 remove(id: &str) -> Result<()> {
let config = Config::global();
let api_key_name = Self::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_custom_providers(dir: &Path) -> Result<Vec<CustomProviderConfig>> {
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()
}
pub fn register_custom_providers(
registry: &mut crate::providers::provider_registry::ProviderRegistry,
dir: &Path,
) -> Result<()> {
let configs = load_custom_providers(dir)?;
for config in configs {
let config_clone = config.clone();
let description = config
.description
.clone()
.unwrap_or_else(|| format!("Custom {} provider", config.display_name));
let default_model = config
.models
.first()
.map(|m| m.name.clone())
.unwrap_or_default();
let known_models: Vec<ModelInfo> = config
.models
.iter()
.map(|m| ModelInfo {
name: m.name.clone(),
context_limit: m.context_limit,
input_token_cost: m.input_token_cost,
output_token_cost: m.output_token_cost,
currency: m.currency.clone(),
supports_cache_control: Some(m.supports_cache_control.unwrap_or(false)),
})
.collect();
match config.engine {
ProviderEngine::OpenAI => {
registry.register_with_name::<OpenAiProvider, _>(
config.name.clone(),
config.display_name.clone(),
description,
default_model,
known_models,
move |model: ModelConfig| {
OpenAiProvider::from_custom_config(model, config_clone.clone())
},
);
}
ProviderEngine::Ollama => {
registry.register_with_name::<OllamaProvider, _>(
config.name.clone(),
config.display_name.clone(),
description,
default_model,
known_models,
move |model: ModelConfig| {
OllamaProvider::from_custom_config(model, config_clone.clone())
},
);
}
ProviderEngine::Anthropic => {
registry.register_with_name::<AnthropicProvider, _>(
config.name.clone(),
config.display_name.clone(),
description,
default_model,
known_models,
move |model: ModelConfig| {
AnthropicProvider::from_custom_config(model, config_clone.clone())
},
);
}
}
}
Ok(())
}
@@ -0,0 +1,317 @@
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<String>,
pub api_key_env: String,
pub base_url: String,
pub models: Vec<ModelInfo>,
pub headers: Option<HashMap<String, String>>,
pub timeout_seconds: Option<u64>,
pub supports_streaming: Option<bool>,
}
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<Mutex<()>> = 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())
}
pub fn create_custom_provider(
engine: &str,
display_name: String,
api_url: String,
api_key: String,
models: Vec<String>,
supports_streaming: Option<bool>,
) -> Result<DeclarativeProviderConfig> {
let id = generate_id(&display_name);
let api_key_name = generate_api_key_name(&id);
let config = Config::global();
config.set_secret(&api_key_name, serde_json::Value::String(api_key))?;
let model_infos: Vec<ModelInfo> = models
.into_iter()
.map(|name| ModelInfo::new(name, 128000))
.collect();
let provider_config = DeclarativeProviderConfig {
name: id.clone(),
engine: match engine {
"openai_compatible" => ProviderEngine::OpenAI,
"anthropic_compatible" => ProviderEngine::Anthropic,
"ollama_compatible" => ProviderEngine::Ollama,
_ => return Err(anyhow::anyhow!("Invalid provider type: {}", engine)),
},
display_name: display_name.clone(),
description: Some(format!("Custom {} provider", display_name)),
api_key_env: api_key_name,
base_url: api_url,
models: model_infos,
headers: None,
timeout_seconds: None,
supports_streaming,
};
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(
id: &str,
provider_type: &str,
display_name: String,
api_url: String,
api_key: String,
models: Vec<String>,
supports_streaming: Option<bool>,
) -> Result<()> {
let loaded_provider = load_provider(id)?;
let existing_config = loaded_provider.config;
let editable = loaded_provider.is_editable;
let config = Config::global();
if !api_key.is_empty() {
config.set_secret(
&existing_config.api_key_env,
serde_json::Value::String(api_key),
)?;
}
if editable {
let model_infos: Vec<ModelInfo> = models
.into_iter()
.map(|name| ModelInfo::new(name, 128000))
.collect();
let updated_config = DeclarativeProviderConfig {
name: id.to_string(),
engine: match provider_type {
"openai_compatible" => ProviderEngine::OpenAI,
"anthropic_compatible" => ProviderEngine::Anthropic,
"ollama_compatible" => ProviderEngine::Ollama,
_ => return Err(anyhow::anyhow!("Invalid provider type: {}", provider_type)),
},
display_name,
description: existing_config.description,
api_key_env: existing_config.api_key_env,
base_url: api_url,
models: model_infos,
headers: existing_config.headers,
timeout_seconds: existing_config.timeout_seconds,
supports_streaming,
};
let file_path = custom_providers_dir().join(format!("{}.json", id));
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<LoadedProvider> {
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 = serde_json::from_str(content)?;
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<Vec<DeclarativeProviderConfig>> {
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<Vec<DeclarativeProviderConfig>> {
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()))?;
let config: DeclarativeProviderConfig = serde_json::from_str(content)?;
res.push(config)
}
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::<OpenAiProvider, _>(
&config,
provider_type,
move |model| OpenAiProvider::from_custom_config(model, config_clone.clone()),
);
}
ProviderEngine::Ollama => {
registry.register_with_name::<OllamaProvider, _>(
&config,
provider_type,
move |model| OllamaProvider::from_custom_config(model, config_clone.clone()),
);
}
ProviderEngine::Anthropic => {
registry.register_with_name::<AnthropicProvider, _>(
&config,
provider_type,
move |model| AnthropicProvider::from_custom_config(model, config_clone.clone()),
);
}
}
}
+2 -2
View File
@@ -1,5 +1,5 @@
pub mod base;
pub mod custom_providers;
pub mod declarative_providers;
mod experiments;
pub mod extensions;
pub mod paths;
@@ -9,7 +9,7 @@ pub mod signup_tetrate;
pub use crate::agents::ExtensionConfig;
pub use base::{Config, ConfigError};
pub use custom_providers::CustomProviderConfig;
pub use declarative_providers::DeclarativeProviderConfig;
pub use experiments::ExperimentManager;
pub use extensions::{
get_all_extension_names, get_all_extensions, get_enabled_extensions, get_extension_by_name,
+5 -2
View File
@@ -15,7 +15,7 @@ use super::formats::anthropic::{
create_request, get_usage, response_to_message, response_to_streaming_message,
};
use super::utils::{emit_debug_trace, get_model, map_http_error_to_provider_error};
use crate::config::custom_providers::CustomProviderConfig;
use crate::config::declarative_providers::DeclarativeProviderConfig;
use crate::conversation::message::Message;
use crate::model::ModelConfig;
use crate::providers::retry::ProviderRetry;
@@ -69,7 +69,10 @@ impl AnthropicProvider {
})
}
pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result<Self> {
pub fn from_custom_config(
model: ModelConfig,
config: DeclarativeProviderConfig,
) -> Result<Self> {
let global_config = crate::config::Config::global();
let api_key: String = global_config
.get_secret(&config.api_key_env)
+8 -2
View File
@@ -81,6 +81,14 @@ impl ModelInfo {
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, ToSchema)]
pub enum ProviderType {
Preferred,
Builtin,
Declarative,
Custom,
}
/// Metadata about a provider's configuration requirements and capabilities
#[derive(Debug, Clone, Serialize, Deserialize, ToSchema)]
pub struct ProviderMetadata {
@@ -93,7 +101,6 @@ pub struct ProviderMetadata {
/// The default/recommended model for this provider
pub default_model: String,
/// A list of currently known models with their capabilities
/// TODO: eventually query the apis directly
pub known_models: Vec<ModelInfo>,
/// Link to the docs where models can be found
pub model_doc_link: String,
@@ -132,7 +139,6 @@ impl ProviderMetadata {
}
}
/// Create a new ProviderMetadata with ModelInfo objects that include cost data
pub fn with_models(
name: &str,
display_name: &str,
@@ -0,0 +1,29 @@
{
"name": "custom_deepseek",
"engine": "openai",
"display_name": "DeepSeek",
"description": "Custom DeepSeek provider",
"api_key_env": "DEEPSEEK_API_KEY",
"base_url": "https://api.deepseek.com",
"models": [
{
"name": "deepseek-chat",
"context_limit": 128000,
"input_token_cost": null,
"output_token_cost": null,
"currency": null,
"supports_cache_control": null
},
{
"name": "deepseek-reasoner",
"context_limit": 128000,
"input_token_cost": null,
"output_token_cost": null,
"currency": null,
"supports_cache_control": null
}
],
"headers": null,
"timeout_seconds": null,
"supports_streaming": true
}
@@ -0,0 +1,31 @@
{
"name": "groq",
"engine": "openai",
"display_name": "Groq (d)",
"description": "Fast inference with Groq hardware",
"api_key_env": "GROQ_API_KEY",
"base_url": "https://api.groq.com/openai/v1/chat/completions",
"models": [
{
"name": "openai/gpt-oss-120b",
"context_limit": 131072
},
{
"name": "llama-3.1-8b-instant",
"context_limit": 131072
},
{
"name": "llama-3.3-70b-versatile",
"context_limit": 131072
},
{
"name": "meta-llama/llama-guard-4-12b",
"context_limit": 131072
},
{
"name": "openai/gpt-oss-20b",
"context_limit": 131072
}
],
"supports_streaming": true
}
+44 -26
View File
@@ -12,7 +12,6 @@ use super::{
gemini_cli::GeminiCliProvider,
githubcopilot::GithubCopilotProvider,
google::GoogleProvider,
groq::GroqProvider,
lead_worker::LeadWorkerProvider,
litellm::LiteLLMProvider,
ollama::OllamaProvider,
@@ -25,8 +24,9 @@ use super::{
venice::VeniceProvider,
xai::XaiProvider,
};
use crate::config::custom_providers::{custom_providers_dir, register_custom_providers};
use crate::config::declarative_providers::register_declarative_providers;
use crate::model::ModelConfig;
use crate::providers::base::ProviderType;
use anyhow::Result;
use tokio::sync::OnceCell;
@@ -38,28 +38,43 @@ static REGISTRY: OnceCell<RwLock<ProviderRegistry>> = OnceCell::const_new();
async fn init_registry() -> RwLock<ProviderRegistry> {
let mut registry = ProviderRegistry::new().with_providers(|registry| {
registry.register::<AnthropicProvider, _>(|m| Box::pin(AnthropicProvider::from_env(m)));
registry.register::<AzureProvider, _>(|m| Box::pin(AzureProvider::from_env(m)));
registry.register::<BedrockProvider, _>(|m| Box::pin(BedrockProvider::from_env(m)));
registry.register::<ClaudeCodeProvider, _>(|m| Box::pin(ClaudeCodeProvider::from_env(m)));
registry.register::<CursorAgentProvider, _>(|m| Box::pin(CursorAgentProvider::from_env(m)));
registry.register::<DatabricksProvider, _>(|m| Box::pin(DatabricksProvider::from_env(m)));
registry.register::<GcpVertexAIProvider, _>(|m| Box::pin(GcpVertexAIProvider::from_env(m)));
registry.register::<GeminiCliProvider, _>(|m| Box::pin(GeminiCliProvider::from_env(m)));
registry
.register::<GithubCopilotProvider, _>(|m| Box::pin(GithubCopilotProvider::from_env(m)));
registry.register::<GoogleProvider, _>(|m| Box::pin(GoogleProvider::from_env(m)));
registry.register::<GroqProvider, _>(|m| Box::pin(GroqProvider::from_env(m)));
registry.register::<LiteLLMProvider, _>(|m| Box::pin(LiteLLMProvider::from_env(m)));
registry.register::<OllamaProvider, _>(|m| Box::pin(OllamaProvider::from_env(m)));
registry.register::<OpenAiProvider, _>(|m| Box::pin(OpenAiProvider::from_env(m)));
registry.register::<OpenRouterProvider, _>(|m| Box::pin(OpenRouterProvider::from_env(m)));
.register::<AnthropicProvider, _>(|m| Box::pin(AnthropicProvider::from_env(m)), true);
registry.register::<AzureProvider, _>(|m| Box::pin(AzureProvider::from_env(m)), false);
registry.register::<BedrockProvider, _>(|m| Box::pin(BedrockProvider::from_env(m)), false);
registry
.register::<SageMakerTgiProvider, _>(|m| Box::pin(SageMakerTgiProvider::from_env(m)));
registry.register::<SnowflakeProvider, _>(|m| Box::pin(SnowflakeProvider::from_env(m)));
registry.register::<TetrateProvider, _>(|m| Box::pin(TetrateProvider::from_env(m)));
registry.register::<VeniceProvider, _>(|m| Box::pin(VeniceProvider::from_env(m)));
registry.register::<XaiProvider, _>(|m| Box::pin(XaiProvider::from_env(m)));
.register::<ClaudeCodeProvider, _>(|m| Box::pin(ClaudeCodeProvider::from_env(m)), true);
registry.register::<CursorAgentProvider, _>(
|m| Box::pin(CursorAgentProvider::from_env(m)),
false,
);
registry
.register::<DatabricksProvider, _>(|m| Box::pin(DatabricksProvider::from_env(m)), true);
registry.register::<GcpVertexAIProvider, _>(
|m| Box::pin(GcpVertexAIProvider::from_env(m)),
false,
);
registry
.register::<GeminiCliProvider, _>(|m| Box::pin(GeminiCliProvider::from_env(m)), false);
registry.register::<GithubCopilotProvider, _>(
|m| Box::pin(GithubCopilotProvider::from_env(m)),
false,
);
registry.register::<GoogleProvider, _>(|m| Box::pin(GoogleProvider::from_env(m)), true);
registry.register::<LiteLLMProvider, _>(|m| Box::pin(LiteLLMProvider::from_env(m)), false);
registry.register::<OllamaProvider, _>(|m| Box::pin(OllamaProvider::from_env(m)), true);
registry.register::<OpenAiProvider, _>(|m| Box::pin(OpenAiProvider::from_env(m)), true);
registry
.register::<OpenRouterProvider, _>(|m| Box::pin(OpenRouterProvider::from_env(m)), true);
registry.register::<SageMakerTgiProvider, _>(
|m| Box::pin(SageMakerTgiProvider::from_env(m)),
false,
);
registry
.register::<SnowflakeProvider, _>(|m| Box::pin(SnowflakeProvider::from_env(m)), false);
registry.register::<TetrateProvider, _>(|m| Box::pin(TetrateProvider::from_env(m)), true);
registry.register::<VeniceProvider, _>(|m| Box::pin(VeniceProvider::from_env(m)), false);
registry.register::<XaiProvider, _>(|m| Box::pin(XaiProvider::from_env(m)), false);
});
if let Err(e) = load_custom_providers_into_registry(&mut registry) {
tracing::warn!("Failed to load custom providers: {}", e);
@@ -68,16 +83,19 @@ async fn init_registry() -> RwLock<ProviderRegistry> {
}
fn load_custom_providers_into_registry(registry: &mut ProviderRegistry) -> Result<()> {
let config_dir = custom_providers_dir();
register_custom_providers(registry, &config_dir)
register_declarative_providers(registry)
}
async fn get_registry() -> &'static RwLock<ProviderRegistry> {
REGISTRY.get_or_init(init_registry).await
}
pub async fn providers() -> Vec<ProviderMetadata> {
get_registry().await.read().unwrap().all_metadata()
pub async fn providers() -> Vec<(ProviderMetadata, ProviderType)> {
get_registry()
.await
.read()
.unwrap()
.all_metadata_with_types()
}
pub async fn refresh_custom_providers() -> Result<()> {
-131
View File
@@ -1,131 +0,0 @@
use super::api_client::{ApiClient, AuthMethod};
use super::errors::ProviderError;
use super::retry::ProviderRetry;
use super::utils::{get_model, handle_response_openai_compat};
use crate::conversation::message::Message;
use crate::model::ModelConfig;
use crate::providers::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
use anyhow::Result;
use async_trait::async_trait;
use rmcp::model::Tool;
use serde_json::Value;
pub const GROQ_API_HOST: &str = "https://api.groq.com";
pub const GROQ_DEFAULT_MODEL: &str = "moonshotai/kimi-k2-instruct";
pub const GROQ_KNOWN_MODELS: &[&str] = &[
"gemma2-9b-it",
"llama-3.3-70b-versatile",
"moonshotai/kimi-k2-instruct",
"qwen/qwen3-32b",
];
pub const GROQ_DOC_URL: &str = "https://console.groq.com/docs/models";
#[derive(serde::Serialize)]
pub struct GroqProvider {
#[serde(skip)]
api_client: ApiClient,
model: ModelConfig,
}
impl GroqProvider {
pub async fn from_env(model: ModelConfig) -> Result<Self> {
let config = crate::config::Config::global();
let api_key: String = config.get_secret("GROQ_API_KEY")?;
let host: String = config
.get_param("GROQ_HOST")
.unwrap_or_else(|_| GROQ_API_HOST.to_string());
let auth = AuthMethod::BearerToken(api_key);
let api_client = ApiClient::new(host, auth)?;
Ok(Self { api_client, model })
}
async fn post(&self, payload: Value) -> Result<Value, ProviderError> {
let response = self
.api_client
.response_post("openai/v1/chat/completions", &payload)
.await?;
handle_response_openai_compat(response).await
}
}
#[async_trait]
impl Provider for GroqProvider {
fn metadata() -> ProviderMetadata {
ProviderMetadata::new(
"groq",
"Groq",
"Fast inference with Groq hardware",
GROQ_DEFAULT_MODEL,
GROQ_KNOWN_MODELS.to_vec(),
GROQ_DOC_URL,
vec![
ConfigKey::new("GROQ_API_KEY", true, true, None),
ConfigKey::new("GROQ_HOST", false, false, Some(GROQ_API_HOST)),
],
)
}
fn get_model_config(&self) -> ModelConfig {
self.model.clone()
}
#[tracing::instrument(
skip(self, model_config, system, messages, tools),
fields(model_config, input, output, input_tokens, output_tokens, total_tokens)
)]
async fn complete_with_model(
&self,
model_config: &ModelConfig,
system: &str,
messages: &[Message],
tools: &[Tool],
) -> Result<(Message, ProviderUsage), ProviderError> {
let payload = create_request(
model_config,
system,
messages,
tools,
&super::utils::ImageFormat::OpenAi,
)?;
let response = self.with_retry(|| self.post(payload.clone())).await?;
let message = response_to_message(&response)?;
let usage = response.get("usage").map(get_usage).unwrap_or_else(|| {
tracing::debug!("Failed to get usage data");
Usage::default()
});
let response_model = get_model(&response);
super::utils::emit_debug_trace(model_config, &payload, &response, &usage);
Ok((message, ProviderUsage::new(response_model, usage)))
}
/// Fetch supported models from Groq; returns Err on failure, Ok(None) if no models found
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let response = self
.api_client
.request("openai/v1/models")
.header("Content-Type", "application/json")?
.response_get()
.await?;
let response = handle_response_openai_compat(response).await?;
let data = response
.get("data")
.and_then(|v| v.as_array())
.ok_or_else(|| {
ProviderError::UsageError("Missing or invalid `data` field in response".into())
})?;
let mut model_names: Vec<String> = data
.iter()
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(String::from))
.collect();
model_names.sort();
Ok(Some(model_names))
}
}
-1
View File
@@ -16,7 +16,6 @@ pub mod gcpvertexai;
pub mod gemini_cli;
pub mod githubcopilot;
pub mod google;
pub mod groq;
pub mod lead_worker;
pub mod litellm;
pub mod oauth;
+5 -2
View File
@@ -3,7 +3,7 @@ use super::base::{ConfigKey, MessageStream, Provider, ProviderMetadata, Provider
use super::errors::ProviderError;
use super::retry::ProviderRetry;
use super::utils::{get_model, handle_response_openai_compat, handle_status_openai_compat};
use crate::config::custom_providers::CustomProviderConfig;
use crate::config::declarative_providers::DeclarativeProviderConfig;
use crate::conversation::message::Message;
use crate::conversation::Conversation;
@@ -93,7 +93,10 @@ impl OllamaProvider {
})
}
pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result<Self> {
pub fn from_custom_config(
model: ModelConfig,
config: DeclarativeProviderConfig,
) -> Result<Self> {
let timeout = Duration::from_secs(config.timeout_seconds.unwrap_or(OLLAMA_TIMEOUT));
// Parse and normalize the custom URL
+5 -2
View File
@@ -20,7 +20,7 @@ use super::utils::{
emit_debug_trace, get_model, handle_response_openai_compat, handle_status_openai_compat,
ImageFormat,
};
use crate::config::custom_providers::CustomProviderConfig;
use crate::config::declarative_providers::DeclarativeProviderConfig;
use crate::conversation::message::Message;
use crate::model::ModelConfig;
@@ -110,7 +110,10 @@ impl OpenAiProvider {
})
}
pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result<Self> {
pub fn from_custom_config(
model: ModelConfig,
config: DeclarativeProviderConfig,
) -> Result<Self> {
let global_config = crate::config::Config::global();
let api_key: String = global_config
.get_secret(&config.api_key_env)
+42 -12
View File
@@ -1,4 +1,5 @@
use super::base::{Provider, ProviderMetadata};
use super::base::{ModelInfo, Provider, ProviderMetadata, ProviderType};
use crate::config::DeclarativeProviderConfig;
use crate::model::ModelConfig;
use anyhow::Result;
use futures::future::BoxFuture;
@@ -11,6 +12,7 @@ type ProviderConstructor =
pub struct ProviderEntry {
metadata: ProviderMetadata,
pub(crate) constructor: ProviderConstructor,
provider_type: ProviderType,
}
#[derive(Default)]
@@ -25,7 +27,7 @@ impl ProviderRegistry {
}
}
pub fn register<P, F>(&mut self, constructor: F)
pub fn register<P, F>(&mut self, constructor: F, preferred: bool)
where
P: Provider + 'static,
F: Fn(ModelConfig) -> BoxFuture<'static, Result<P>> + Send + Sync + 'static,
@@ -44,26 +46,50 @@ impl ProviderRegistry {
Ok(Arc::new(provider) as Arc<dyn Provider>)
})
}),
provider_type: if preferred {
ProviderType::Preferred
} else {
ProviderType::Builtin
},
},
);
}
pub fn register_with_name<P, F>(
&mut self,
custom_name: String,
display_name: String,
description: String,
default_model: String,
known_models: Vec<super::base::ModelInfo>,
config: &DeclarativeProviderConfig,
provider_type: ProviderType,
constructor: F,
) where
P: Provider + 'static,
F: Fn(ModelConfig) -> Result<P> + Send + Sync + 'static,
{
let base_metadata = P::metadata();
let description = config
.description
.clone()
.unwrap_or_else(|| format!("Custom {} provider", config.display_name));
let default_model = config
.models
.first()
.map(|m| m.name.clone())
.unwrap_or_default();
let known_models: Vec<ModelInfo> = config
.models
.iter()
.map(|m| ModelInfo {
name: m.name.clone(),
context_limit: m.context_limit,
input_token_cost: m.input_token_cost,
output_token_cost: m.output_token_cost,
currency: m.currency.clone(),
supports_cache_control: Some(m.supports_cache_control.unwrap_or(false)),
})
.collect();
let custom_metadata = ProviderMetadata {
name: custom_name.clone(),
display_name,
name: config.name.clone(),
display_name: config.display_name.clone(),
description,
default_model,
known_models,
@@ -72,7 +98,7 @@ impl ProviderRegistry {
};
self.entries.insert(
custom_name,
config.name.clone(),
ProviderEntry {
metadata: custom_metadata,
constructor: Arc::new(move |model| {
@@ -82,6 +108,7 @@ impl ProviderRegistry {
Ok(Arc::new(provider) as Arc<dyn Provider>)
})
}),
provider_type,
},
);
}
@@ -103,8 +130,11 @@ impl ProviderRegistry {
(entry.constructor)(model).await
}
pub fn all_metadata(&self) -> Vec<ProviderMetadata> {
self.entries.values().map(|e| e.metadata.clone()).collect()
pub fn all_metadata_with_types(&self) -> Vec<(ProviderMetadata, ProviderType)> {
self.entries
.values()
.map(|e| (e.metadata.clone(), e.provider_type))
.collect()
}
pub fn remove_custom_providers(&mut self) {
+11 -3
View File
@@ -80,6 +80,14 @@ pub fn map_http_error_to_provider_error(
);
ProviderError::Authentication(message)
}
StatusCode::PAYLOAD_TOO_LARGE => {
let payload_str = if let Some(payload) = &payload {
payload.to_string()
} else {
"Payload is too large.".to_string()
};
ProviderError::ContextLengthExceeded(payload_str)
}
StatusCode::BAD_REQUEST => {
let mut error_msg = "Unknown error".to_string();
if let Some(payload) = &payload {
@@ -929,12 +937,12 @@ mod tests {
"The model 'gpt-5' does not exist (code: model_not_found, type: invalid_request_error) (status 404)".to_string(),
)),
),
// Non-JSON body error (tests parse failure path)
// Non-JSON body error (tests 413 PAYLOAD_TOO_LARGE -> ContextLengthExceeded)
(
413,
Some(Value::String("Payload Too Large".to_string())),
Err(ProviderError::RequestFailed(
"Request failed with status: 413 Payload Too Large".to_string(),
Err(ProviderError::ContextLengthExceeded(
"Payload is too large.".to_string(),
)),
),
];