Custom providers update (#4099)

Co-authored-by: developerayo <shodipovi@gmail.com>
Co-authored-by: Douwe Osinga <douwe@squareup.com>
Co-authored-by: Claude <noreply@anthropic.com>
Co-authored-by: Zane Staggs <zane@squareup.com>
This commit is contained in:
Douwe Osinga
2025-08-19 18:03:10 -04:00
committed by GitHub
parent 1d93a59f31
commit 942ef5b0a3
31 changed files with 1704 additions and 567 deletions
+215
View File
@@ -0,0 +1,215 @@
use crate::config::{Config, APP_STRATEGY};
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 etcetera::{choose_app_strategy, AppStrategy};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::path::Path;
pub fn custom_providers_dir() -> std::path::PathBuf {
choose_app_strategy(APP_STRATEGY.clone())
.expect("goose requires a home dir")
.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(())
}
+2
View File
@@ -1,4 +1,5 @@
pub mod base;
pub mod custom_providers;
mod experiments;
pub mod extensions;
pub mod permission;
@@ -6,6 +7,7 @@ pub mod signup_openrouter;
pub use crate::agents::ExtensionConfig;
pub use base::{Config, ConfigError, APP_STRATEGY};
pub use custom_providers::CustomProviderConfig;
pub use experiments::ExperimentManager;
pub use extensions::{ExtensionConfigManager, ExtensionEntry};
pub use permission::PermissionManager;
+29 -2
View File
@@ -15,6 +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::conversation::message::Message;
use crate::impl_provider_default;
use crate::model::ModelConfig;
@@ -42,6 +43,7 @@ pub struct AnthropicProvider {
#[serde(skip)]
api_client: ApiClient,
model: ModelConfig,
supports_streaming: bool,
}
impl_provider_default!(AnthropicProvider);
@@ -62,7 +64,32 @@ impl AnthropicProvider {
let api_client =
ApiClient::new(host, auth)?.with_header("anthropic-version", ANTHROPIC_API_VERSION)?;
Ok(Self { api_client, model })
Ok(Self {
api_client,
model,
supports_streaming: true,
})
}
pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result<Self> {
let global_config = crate::config::Config::global();
let api_key: String = global_config
.get_secret(&config.api_key_env)
.map_err(|_| anyhow::anyhow!("Missing API key: {}", config.api_key_env))?;
let auth = AuthMethod::ApiKey {
header_name: "x-api-key".to_string(),
key: api_key,
};
let api_client = ApiClient::new(config.base_url, auth)?
.with_header("anthropic-version", ANTHROPIC_API_VERSION)?;
Ok(Self {
api_client,
model,
supports_streaming: config.supports_streaming.unwrap_or(true),
})
}
fn get_conditional_headers(&self) -> Vec<(&str, &str)> {
@@ -260,6 +287,6 @@ impl Provider for AnthropicProvider {
}
fn supports_streaming(&self) -> bool {
true
self.supports_streaming
}
}
+160 -249
View File
@@ -1,4 +1,4 @@
use std::sync::Arc;
use std::sync::{Arc, RwLock};
use super::{
anthropic::AnthropicProvider,
@@ -17,66 +17,87 @@ use super::{
ollama::OllamaProvider,
openai::OpenAiProvider,
openrouter::OpenRouterProvider,
provider_registry::ProviderRegistry,
sagemaker_tgi::SageMakerTgiProvider,
snowflake::SnowflakeProvider,
venice::VeniceProvider,
xai::XaiProvider,
};
use crate::config::custom_providers::{custom_providers_dir, register_custom_providers};
use crate::model::ModelConfig;
use anyhow::Result;
use once_cell::sync::Lazy;
#[cfg(test)]
use super::errors::ProviderError;
#[cfg(test)]
use rmcp::model::Tool;
fn default_lead_turns() -> usize {
3
}
fn default_failure_threshold() -> usize {
2
}
fn default_fallback_turns() -> usize {
2
const DEFAULT_LEAD_TURNS: usize = 3;
const DEFAULT_FAILURE_THRESHOLD: usize = 2;
const DEFAULT_FALLBACK_TURNS: usize = 2;
static REGISTRY: Lazy<RwLock<ProviderRegistry>> = Lazy::new(|| {
let registry = ProviderRegistry::new().with_providers(|registry| {
registry.register::<AnthropicProvider, _>(AnthropicProvider::from_env);
registry.register::<AzureProvider, _>(AzureProvider::from_env);
registry.register::<BedrockProvider, _>(BedrockProvider::from_env);
registry.register::<ClaudeCodeProvider, _>(ClaudeCodeProvider::from_env);
registry.register::<CursorAgentProvider, _>(CursorAgentProvider::from_env);
registry.register::<DatabricksProvider, _>(DatabricksProvider::from_env);
registry.register::<GcpVertexAIProvider, _>(GcpVertexAIProvider::from_env);
registry.register::<GeminiCliProvider, _>(GeminiCliProvider::from_env);
registry.register::<GoogleProvider, _>(GoogleProvider::from_env);
registry.register::<GroqProvider, _>(GroqProvider::from_env);
registry.register::<LiteLLMProvider, _>(LiteLLMProvider::from_env);
registry.register::<OllamaProvider, _>(OllamaProvider::from_env);
registry.register::<OpenAiProvider, _>(OpenAiProvider::from_env);
registry.register::<OpenRouterProvider, _>(OpenRouterProvider::from_env);
registry.register::<SageMakerTgiProvider, _>(SageMakerTgiProvider::from_env);
registry.register::<SnowflakeProvider, _>(SnowflakeProvider::from_env);
registry.register::<VeniceProvider, _>(VeniceProvider::from_env);
registry.register::<XaiProvider, _>(XaiProvider::from_env);
if let Err(e) = load_custom_providers_into_registry(registry) {
tracing::warn!("Failed to load custom providers: {}", e);
}
});
RwLock::new(registry)
});
fn load_custom_providers_into_registry(registry: &mut ProviderRegistry) -> Result<()> {
let config_dir = custom_providers_dir();
register_custom_providers(registry, &config_dir)
}
pub fn providers() -> Vec<ProviderMetadata> {
vec![
AnthropicProvider::metadata(),
AzureProvider::metadata(),
BedrockProvider::metadata(),
ClaudeCodeProvider::metadata(),
CursorAgentProvider::metadata(),
DatabricksProvider::metadata(),
GcpVertexAIProvider::metadata(),
GeminiCliProvider::metadata(),
// GithubCopilotProvider::metadata(),
GoogleProvider::metadata(),
GroqProvider::metadata(),
LiteLLMProvider::metadata(),
OllamaProvider::metadata(),
OpenAiProvider::metadata(),
OpenRouterProvider::metadata(),
SageMakerTgiProvider::metadata(),
VeniceProvider::metadata(),
SnowflakeProvider::metadata(),
XaiProvider::metadata(),
]
REGISTRY.read().unwrap().all_metadata()
}
pub fn refresh_custom_providers() -> Result<()> {
let mut registry = REGISTRY.write().unwrap();
registry.remove_custom_providers();
if let Err(e) = load_custom_providers_into_registry(&mut registry) {
tracing::warn!("Failed to refresh custom providers: {}", e);
return Err(e);
}
tracing::info!("Custom providers refreshed");
Ok(())
}
pub fn create(name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
let config = crate::config::Config::global();
// Check for lead model environment variables
if let Ok(lead_model_name) = config.get_param::<String>("GOOSE_LEAD_MODEL") {
tracing::info!("Creating lead/worker provider from environment variables");
return create_lead_worker_from_env(name, &model, &lead_model_name);
}
create_provider(name, model)
REGISTRY.read().unwrap().create(name, model)
}
/// Create a lead/worker provider from environment variables
fn create_lead_worker_from_env(
default_provider_name: &str,
default_model: &ModelConfig,
@@ -84,61 +105,36 @@ fn create_lead_worker_from_env(
) -> Result<Arc<dyn Provider>> {
let config = crate::config::Config::global();
// Get lead provider (optional, defaults to main provider)
let lead_provider_name = config
.get_param::<String>("GOOSE_LEAD_PROVIDER")
.unwrap_or_else(|_| default_provider_name.to_string());
// Get configuration parameters with defaults
let lead_turns = config
.get_param::<usize>("GOOSE_LEAD_TURNS")
.unwrap_or(default_lead_turns());
.unwrap_or(DEFAULT_LEAD_TURNS);
let failure_threshold = config
.get_param::<usize>("GOOSE_LEAD_FAILURE_THRESHOLD")
.unwrap_or(default_failure_threshold());
.unwrap_or(DEFAULT_FAILURE_THRESHOLD);
let fallback_turns = config
.get_param::<usize>("GOOSE_LEAD_FALLBACK_TURNS")
.unwrap_or(default_fallback_turns());
.unwrap_or(DEFAULT_FALLBACK_TURNS);
let lead_model_config = ModelConfig::new_with_context_env(
lead_model_name.to_string(),
Some("GOOSE_LEAD_CONTEXT_LIMIT"),
)?;
// For worker model, preserve the original context_limit from config (highest precedence)
// while still allowing environment variable overrides
let worker_model_config = {
// Start with a clone of the original model to preserve user-specified settings
let mut worker_config = ModelConfig::new_or_fail(default_model.model_name.as_str())
.with_context_limit(default_model.context_limit)
.with_temperature(default_model.temperature)
.with_max_tokens(default_model.max_tokens)
.with_toolshim(default_model.toolshim)
.with_toolshim_model(default_model.toolshim_model.clone());
let worker_model_config = create_worker_model_config(default_model)?;
// Apply environment variable overrides with proper precedence
let global_config = crate::config::Config::global();
let lead_provider = REGISTRY
.read()
.unwrap()
.create(&lead_provider_name, lead_model_config)?;
let worker_provider = REGISTRY
.read()
.unwrap()
.create(default_provider_name, worker_model_config)?;
// Check for worker-specific context limit
if let Ok(limit_str) = global_config.get_param::<String>("GOOSE_WORKER_CONTEXT_LIMIT") {
if let Ok(limit) = limit_str.parse::<usize>() {
worker_config = worker_config.with_context_limit(Some(limit));
}
} else if let Ok(limit_str) = global_config.get_param::<String>("GOOSE_CONTEXT_LIMIT") {
// Check for general context limit if worker-specific is not set
if let Ok(limit) = limit_str.parse::<usize>() {
worker_config = worker_config.with_context_limit(Some(limit));
}
}
worker_config
};
// Create the providers
let lead_provider = create_provider(&lead_provider_name, lead_model_config)?;
let worker_provider = create_provider(default_provider_name, worker_model_config)?;
// Create the lead/worker provider with configured settings
Ok(Arc::new(LeadWorkerProvider::new_with_settings(
lead_provider,
worker_provider,
@@ -148,30 +144,27 @@ fn create_lead_worker_from_env(
)))
}
fn create_provider(name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
// We use Arc instead of Box to be able to clone for multiple async tasks
match name {
"anthropic" => Ok(Arc::new(AnthropicProvider::from_env(model)?)),
"aws_bedrock" => Ok(Arc::new(BedrockProvider::from_env(model)?)),
"azure_openai" => Ok(Arc::new(AzureProvider::from_env(model)?)),
"claude-code" => Ok(Arc::new(ClaudeCodeProvider::from_env(model)?)),
"cursor-agent" => Ok(Arc::new(CursorAgentProvider::from_env(model)?)),
"databricks" => Ok(Arc::new(DatabricksProvider::from_env(model)?)),
"gcp_vertex_ai" => Ok(Arc::new(GcpVertexAIProvider::from_env(model)?)),
"gemini-cli" => Ok(Arc::new(GeminiCliProvider::from_env(model)?)),
// "github_copilot" => Ok(Arc::new(GithubCopilotProvider::from_env(model)?)),
"google" => Ok(Arc::new(GoogleProvider::from_env(model)?)),
"groq" => Ok(Arc::new(GroqProvider::from_env(model)?)),
"litellm" => Ok(Arc::new(LiteLLMProvider::from_env(model)?)),
"ollama" => Ok(Arc::new(OllamaProvider::from_env(model)?)),
"openai" => Ok(Arc::new(OpenAiProvider::from_env(model)?)),
"openrouter" => Ok(Arc::new(OpenRouterProvider::from_env(model)?)),
"sagemaker_tgi" => Ok(Arc::new(SageMakerTgiProvider::from_env(model)?)),
"snowflake" => Ok(Arc::new(SnowflakeProvider::from_env(model)?)),
"venice" => Ok(Arc::new(VeniceProvider::from_env(model)?)),
"xai" => Ok(Arc::new(XaiProvider::from_env(model)?)),
_ => Err(anyhow::anyhow!("Unknown provider: {}", name)),
fn create_worker_model_config(default_model: &ModelConfig) -> Result<ModelConfig> {
let mut worker_config = ModelConfig::new_or_fail(&default_model.model_name)
.with_context_limit(default_model.context_limit)
.with_temperature(default_model.temperature)
.with_max_tokens(default_model.max_tokens)
.with_toolshim(default_model.toolshim)
.with_toolshim_model(default_model.toolshim_model.clone());
let global_config = crate::config::Config::global();
if let Ok(limit_str) = global_config.get_param::<String>("GOOSE_WORKER_CONTEXT_LIMIT") {
if let Ok(limit) = limit_str.parse::<usize>() {
worker_config = worker_config.with_context_limit(Some(limit));
}
} else if let Ok(limit_str) = global_config.get_param::<String>("GOOSE_CONTEXT_LIMIT") {
if let Ok(limit) = limit_str.parse::<usize>() {
worker_config = worker_config.with_context_limit(Some(limit));
}
}
Ok(worker_config)
}
#[cfg(test)]
@@ -183,7 +176,6 @@ mod tests {
use rmcp::model::{AnnotateAble, RawTextContent, Role};
use std::env;
#[allow(dead_code)]
#[derive(Clone)]
struct MockTestProvider {
name: String,
@@ -233,222 +225,141 @@ mod tests {
}
}
struct EnvVarGuard {
vars: Vec<(String, Option<String>)>,
}
impl EnvVarGuard {
fn new(vars: &[&str]) -> Self {
let saved_vars = vars
.iter()
.map(|&var| (var.to_string(), env::var(var).ok()))
.collect();
for &var in vars {
env::remove_var(var);
}
Self { vars: saved_vars }
}
fn set(&self, key: &str, value: &str) {
env::set_var(key, value);
}
}
impl Drop for EnvVarGuard {
fn drop(&mut self) {
for (key, value) in &self.vars {
match value {
Some(val) => env::set_var(key, val),
None => env::remove_var(key),
}
}
}
}
#[test]
fn test_create_lead_worker_provider() {
// Save current env vars
let saved_lead = env::var("GOOSE_LEAD_MODEL").ok();
let saved_provider = env::var("GOOSE_LEAD_PROVIDER").ok();
let saved_turns = env::var("GOOSE_LEAD_TURNS").ok();
let _guard = EnvVarGuard::new(&[
"GOOSE_LEAD_MODEL",
"GOOSE_LEAD_PROVIDER",
"GOOSE_LEAD_TURNS",
]);
// Test with basic lead model configuration
env::set_var("GOOSE_LEAD_MODEL", "gpt-4o");
_guard.set("GOOSE_LEAD_MODEL", "gpt-4o");
// This will try to create a lead/worker provider
let gpt4mini_config = ModelConfig::new_or_fail("gpt-4o-mini");
let result = create("openai", gpt4mini_config.clone());
// The creation might succeed or fail depending on API keys, but we can verify the logic path
match result {
Ok(_) => {
// If it succeeds, it means we created a lead/worker provider successfully
// This would happen if API keys are available in the test environment
}
Ok(_) => {}
Err(error) => {
// If it fails, it should be due to missing API keys, confirming we tried to create providers
let error_msg = error.to_string();
assert!(error_msg.contains("OPENAI_API_KEY") || error_msg.contains("secret"));
}
}
// Test with different lead provider
env::set_var("GOOSE_LEAD_PROVIDER", "anthropic");
env::set_var("GOOSE_LEAD_TURNS", "5");
_guard.set("GOOSE_LEAD_PROVIDER", "anthropic");
_guard.set("GOOSE_LEAD_TURNS", "5");
let _result = create("openai", gpt4mini_config);
// Similar validation as above - will fail due to missing API keys but confirms the logic
// Restore env vars
match saved_lead {
Some(val) => env::set_var("GOOSE_LEAD_MODEL", val),
None => env::remove_var("GOOSE_LEAD_MODEL"),
}
match saved_provider {
Some(val) => env::set_var("GOOSE_LEAD_PROVIDER", val),
None => env::remove_var("GOOSE_LEAD_PROVIDER"),
}
match saved_turns {
Some(val) => env::set_var("GOOSE_LEAD_TURNS", val),
None => env::remove_var("GOOSE_LEAD_TURNS"),
}
}
#[test]
fn test_lead_model_env_vars_with_defaults() {
// Save current env vars
let saved_vars = [
("GOOSE_LEAD_MODEL", env::var("GOOSE_LEAD_MODEL").ok()),
("GOOSE_LEAD_PROVIDER", env::var("GOOSE_LEAD_PROVIDER").ok()),
("GOOSE_LEAD_TURNS", env::var("GOOSE_LEAD_TURNS").ok()),
(
"GOOSE_LEAD_FAILURE_THRESHOLD",
env::var("GOOSE_LEAD_FAILURE_THRESHOLD").ok(),
),
(
"GOOSE_LEAD_FALLBACK_TURNS",
env::var("GOOSE_LEAD_FALLBACK_TURNS").ok(),
),
];
let _guard = EnvVarGuard::new(&[
"GOOSE_LEAD_MODEL",
"GOOSE_LEAD_PROVIDER",
"GOOSE_LEAD_TURNS",
"GOOSE_LEAD_FAILURE_THRESHOLD",
"GOOSE_LEAD_FALLBACK_TURNS",
]);
// Clear all lead env vars
for (key, _) in &saved_vars {
env::remove_var(key);
}
_guard.set("GOOSE_LEAD_MODEL", "grok-3");
// Set only the required lead model
env::set_var("GOOSE_LEAD_MODEL", "grok-3");
// This should use defaults for all other values
let result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini"));
// Should attempt to create lead/worker provider (will fail due to missing API keys but confirms logic)
match result {
Ok(_) => {
// Success means we have API keys and created the provider
}
Ok(_) => {}
Err(error) => {
// Should fail due to missing API keys, confirming we tried to create providers
let error_msg = error.to_string();
assert!(error_msg.contains("OPENAI_API_KEY") || error_msg.contains("secret"));
}
}
// Test with custom values
env::set_var("GOOSE_LEAD_TURNS", "7");
env::set_var("GOOSE_LEAD_FAILURE_THRESHOLD", "4");
env::set_var("GOOSE_LEAD_FALLBACK_TURNS", "3");
_guard.set("GOOSE_LEAD_TURNS", "7");
_guard.set("GOOSE_LEAD_FAILURE_THRESHOLD", "4");
_guard.set("GOOSE_LEAD_FALLBACK_TURNS", "3");
let _result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini"));
// Should still attempt to create lead/worker provider with custom settings
// Restore all env vars
for (key, value) in saved_vars {
match value {
Some(val) => env::set_var(key, val),
None => env::remove_var(key),
}
}
}
#[test]
fn test_create_regular_provider_without_lead_config() {
// Save current env vars
let saved_lead = env::var("GOOSE_LEAD_MODEL").ok();
let saved_provider = env::var("GOOSE_LEAD_PROVIDER").ok();
let saved_turns = env::var("GOOSE_LEAD_TURNS").ok();
let saved_threshold = env::var("GOOSE_LEAD_FAILURE_THRESHOLD").ok();
let saved_fallback = env::var("GOOSE_LEAD_FALLBACK_TURNS").ok();
let _guard = EnvVarGuard::new(&[
"GOOSE_LEAD_MODEL",
"GOOSE_LEAD_PROVIDER",
"GOOSE_LEAD_TURNS",
"GOOSE_LEAD_FAILURE_THRESHOLD",
"GOOSE_LEAD_FALLBACK_TURNS",
]);
// Ensure all GOOSE_LEAD_* variables are not set
env::remove_var("GOOSE_LEAD_MODEL");
env::remove_var("GOOSE_LEAD_PROVIDER");
env::remove_var("GOOSE_LEAD_TURNS");
env::remove_var("GOOSE_LEAD_FAILURE_THRESHOLD");
env::remove_var("GOOSE_LEAD_FALLBACK_TURNS");
// This should try to create a regular provider
let result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini"));
// The creation might succeed or fail depending on API keys
match result {
Ok(_) => {
// If it succeeds, it means we created a regular provider successfully
// This would happen if API keys are available in the test environment
}
Ok(_) => {}
Err(error) => {
// If it fails, it should be due to missing API keys
let error_msg = error.to_string();
assert!(error_msg.contains("OPENAI_API_KEY") || error_msg.contains("secret"));
}
}
if let Some(val) = saved_lead {
env::set_var("GOOSE_LEAD_MODEL", val);
}
if let Some(val) = saved_provider {
env::set_var("GOOSE_LEAD_PROVIDER", val);
}
if let Some(val) = saved_turns {
env::set_var("GOOSE_LEAD_TURNS", val);
}
if let Some(val) = saved_threshold {
env::set_var("GOOSE_LEAD_FAILURE_THRESHOLD", val);
}
if let Some(val) = saved_fallback {
env::set_var("GOOSE_LEAD_FALLBACK_TURNS", val);
}
}
#[test]
fn test_worker_model_preserves_original_context_limit() {
use std::env;
let _guard = EnvVarGuard::new(&[
"GOOSE_LEAD_MODEL",
"GOOSE_WORKER_CONTEXT_LIMIT",
"GOOSE_CONTEXT_LIMIT",
]);
// Save current env vars
let saved_vars = [
("GOOSE_LEAD_MODEL", env::var("GOOSE_LEAD_MODEL").ok()),
(
"GOOSE_WORKER_CONTEXT_LIMIT",
env::var("GOOSE_WORKER_CONTEXT_LIMIT").ok(),
),
("GOOSE_CONTEXT_LIMIT", env::var("GOOSE_CONTEXT_LIMIT").ok()),
];
_guard.set("GOOSE_LEAD_MODEL", "gpt-4o");
// Clear env vars to ensure clean test
for (key, _) in &saved_vars {
env::remove_var(key);
}
// Set up lead model to trigger lead/worker mode
env::set_var("GOOSE_LEAD_MODEL", "gpt-4o");
// Create a default model with explicit context_limit
let default_model =
ModelConfig::new_or_fail("gpt-3.5-turbo").with_context_limit(Some(16_000));
// Test case 1: No environment variables - should preserve original context_limit
let result = create_lead_worker_from_env("openai", &default_model, "gpt-4o");
// Test case 2: With GOOSE_WORKER_CONTEXT_LIMIT - should override original
env::set_var("GOOSE_WORKER_CONTEXT_LIMIT", "32000");
_guard.set("GOOSE_WORKER_CONTEXT_LIMIT", "32000");
let _result = create_lead_worker_from_env("openai", &default_model, "gpt-4o");
env::remove_var("GOOSE_WORKER_CONTEXT_LIMIT");
// Test case 3: With GOOSE_CONTEXT_LIMIT - should override original
env::set_var("GOOSE_CONTEXT_LIMIT", "64000");
_guard.set("GOOSE_CONTEXT_LIMIT", "64000");
let _result = create_lead_worker_from_env("openai", &default_model, "gpt-4o");
env::remove_var("GOOSE_CONTEXT_LIMIT");
// Restore env vars
for (key, value) in saved_vars {
match value {
Some(val) => env::set_var(key, val),
None => env::remove_var(key),
}
}
// The main verification is that the function doesn't panic and handles
// the context limit preservation logic correctly. More detailed testing
// would require mocking the provider creation.
// The result could be Ok or Err depending on whether API keys are available
// in the test environment - both are acceptable for this test
match result {
Ok(_) => {
// Success means API keys are available and lead/worker provider was created
// This confirms our logic path is working
}
Err(_) => {
// Error is expected if API keys are not available
// This also confirms our logic path is working
}
Ok(_) => {}
Err(_) => {}
}
}
}
+2 -1
View File
@@ -24,6 +24,7 @@ pub mod ollama;
pub mod openai;
pub mod openrouter;
pub mod pricing;
pub mod provider_registry;
mod retry;
pub mod sagemaker_tgi;
pub mod snowflake;
@@ -35,4 +36,4 @@ pub mod utils_universal_openai_stream;
pub mod venice;
pub mod xai;
pub use factory::{create, providers};
pub use factory::{create, providers, refresh_custom_providers};
+47 -1
View File
@@ -3,6 +3,7 @@ use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
use super::errors::ProviderError;
use super::retry::ProviderRetry;
use super::utils::{get_model, handle_response_openai_compat};
use crate::config::custom_providers::CustomProviderConfig;
use crate::conversation::message::Message;
use crate::conversation::Conversation;
use crate::impl_provider_default;
@@ -30,6 +31,7 @@ pub struct OllamaProvider {
#[serde(skip)]
api_client: ApiClient,
model: ModelConfig,
supports_streaming: bool,
}
impl_provider_default!(OllamaProvider);
@@ -73,7 +75,47 @@ impl OllamaProvider {
let auth = AuthMethod::Custom(Box::new(NoAuth));
let api_client = ApiClient::with_timeout(base_url.to_string(), auth, timeout)?;
Ok(Self { api_client, model })
Ok(Self {
api_client,
model,
supports_streaming: true,
})
}
pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result<Self> {
let timeout = Duration::from_secs(config.timeout_seconds.unwrap_or(OLLAMA_TIMEOUT));
// Parse and normalize the custom URL
let base =
if config.base_url.starts_with("http://") || config.base_url.starts_with("https://") {
config.base_url.clone()
} else {
format!("http://{}", config.base_url)
};
let mut base_url = Url::parse(&base)
.map_err(|e| anyhow::anyhow!("Invalid base URL '{}': {}", config.base_url, e))?;
// Set default port if missing and not using standard ports
let explicit_default_port =
config.base_url.ends_with(":80") || config.base_url.ends_with(":443");
let is_https = base_url.scheme() == "https";
if base_url.port().is_none() && !explicit_default_port && !is_https {
base_url
.set_port(Some(OLLAMA_DEFAULT_PORT))
.map_err(|_| anyhow::anyhow!("Failed to set default port"))?;
}
// No authentication for Ollama
let auth = AuthMethod::Custom(Box::new(NoAuth));
let api_client = ApiClient::with_timeout(base_url.to_string(), auth, timeout)?;
Ok(Self {
api_client,
model,
supports_streaming: config.supports_streaming.unwrap_or(true),
})
}
async fn post(&self, payload: &Value) -> Result<Value, ProviderError> {
@@ -181,6 +223,10 @@ impl Provider for OllamaProvider {
Ok(safe_truncate(&description, 100))
}
fn supports_streaming(&self) -> bool {
self.supports_streaming
}
}
impl OllamaProvider {
+48 -1
View File
@@ -20,6 +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::conversation::message::Message;
use crate::impl_provider_default;
use crate::model::ModelConfig;
@@ -51,6 +52,7 @@ pub struct OpenAiProvider {
project: Option<String>,
model: ModelConfig,
custom_headers: Option<HashMap<String, String>>,
supports_streaming: bool,
}
impl_provider_default!(OpenAiProvider);
@@ -103,6 +105,51 @@ impl OpenAiProvider {
project,
model,
custom_headers,
supports_streaming: true,
})
}
pub fn from_custom_config(model: ModelConfig, config: CustomProviderConfig) -> Result<Self> {
let global_config = crate::config::Config::global();
let api_key: String = global_config
.get_secret(&config.api_key_env)
.map_err(|_e| anyhow::anyhow!("Missing API key: {}", config.api_key_env))?;
let url = url::Url::parse(&config.base_url)
.map_err(|e| anyhow::anyhow!("Invalid base URL '{}': {}", config.base_url, e))?;
let host = format!("{}://{}", url.scheme(), url.host_str().unwrap_or(""));
let base_path = url.path().trim_start_matches('/').to_string();
let base_path = if base_path.is_empty() {
"v1/chat/completions".to_string()
} else {
base_path
};
let timeout_secs = config.timeout_seconds.unwrap_or(600);
let auth = AuthMethod::BearerToken(api_key);
let mut api_client =
ApiClient::with_timeout(host, auth, std::time::Duration::from_secs(timeout_secs))?;
// Add custom headers if present
if let Some(headers) = &config.headers {
let mut header_map = reqwest::header::HeaderMap::new();
for (key, value) in headers {
let header_name = reqwest::header::HeaderName::from_bytes(key.as_bytes())?;
let header_value = reqwest::header::HeaderValue::from_str(value)?;
header_map.insert(header_name, header_value);
}
api_client = api_client.with_headers(header_map)?;
}
Ok(Self {
api_client,
base_path,
organization: None,
project: None,
model,
custom_headers: config.headers,
supports_streaming: config.supports_streaming.unwrap_or(true),
})
}
@@ -206,7 +253,7 @@ impl Provider for OpenAiProvider {
}
fn supports_streaming(&self) -> bool {
true
self.supports_streaming
}
async fn stream(
@@ -0,0 +1,102 @@
use super::base::{Provider, ProviderMetadata};
use crate::model::ModelConfig;
use anyhow::Result;
use std::collections::HashMap;
use std::sync::Arc;
type ProviderConstructor = Box<dyn Fn(ModelConfig) -> Result<Arc<dyn Provider>> + Send + Sync>;
struct ProviderEntry {
metadata: ProviderMetadata,
constructor: ProviderConstructor,
}
#[derive(Default)]
pub struct ProviderRegistry {
entries: HashMap<String, ProviderEntry>,
}
impl ProviderRegistry {
pub fn new() -> Self {
Self {
entries: HashMap::new(),
}
}
pub fn register<P, F>(&mut self, constructor: F)
where
P: Provider + 'static,
F: Fn(ModelConfig) -> Result<P> + Send + Sync + 'static,
{
let metadata = P::metadata();
let name = metadata.name.clone();
self.entries.insert(
name,
ProviderEntry {
metadata,
constructor: Box::new(move |model| Ok(Arc::new(constructor(model)?))),
},
);
}
/// create provider with custom name
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>,
constructor: F,
) where
P: Provider + 'static,
F: Fn(ModelConfig) -> Result<P> + Send + Sync + 'static,
{
let base_metadata = P::metadata();
let custom_metadata = ProviderMetadata {
name: custom_name.clone(),
display_name,
description,
default_model,
known_models,
model_doc_link: base_metadata.model_doc_link,
config_keys: base_metadata.config_keys,
};
self.entries.insert(
custom_name,
ProviderEntry {
metadata: custom_metadata,
constructor: Box::new(move |model| Ok(Arc::new(constructor(model)?))),
},
);
}
pub fn with_providers<F>(mut self, setup: F) -> Self
where
F: FnOnce(&mut Self),
{
setup(&mut self);
self
}
pub fn create(&self, name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
let _available_providers: Vec<_> = self.entries.keys().collect();
let entry = self
.entries
.get(name)
.ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", name))?;
(entry.constructor)(model)
}
pub fn all_metadata(&self) -> Vec<ProviderMetadata> {
self.entries.values().map(|e| e.metadata.clone()).collect()
}
pub fn remove_custom_providers(&mut self) {
self.entries.retain(|name, _| !name.starts_with("custom_"));
}
}