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:
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(_) => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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_"));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user