@@ -17,12 +17,11 @@ use super::formats::anthropic::{
|
||||
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;
|
||||
use crate::providers::retry::ProviderRetry;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-0";
|
||||
pub const ANTHROPIC_DEFAULT_MODEL: &str = "claude-sonnet-4-0";
|
||||
const ANTHROPIC_DEFAULT_FAST_MODEL: &str = "claude-3-7-sonnet-latest";
|
||||
const ANTHROPIC_KNOWN_MODELS: &[&str] = &[
|
||||
"claude-sonnet-4-0",
|
||||
@@ -45,10 +44,8 @@ pub struct AnthropicProvider {
|
||||
supports_streaming: bool,
|
||||
}
|
||||
|
||||
impl_provider_default!(AnthropicProvider);
|
||||
|
||||
impl AnthropicProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let model = model.with_fast(ANTHROPIC_DEFAULT_FAST_MODEL.to_string());
|
||||
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
@@ -11,7 +11,6 @@ use super::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{emit_debug_trace, get_model, handle_response_openai_compat, ImageFormat};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
@@ -68,10 +67,8 @@ impl AuthProvider for AzureAuthProvider {
|
||||
}
|
||||
}
|
||||
|
||||
impl_provider_default!(AzureProvider);
|
||||
|
||||
impl AzureProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let endpoint: String = config.get_param("AZURE_OPENAI_ENDPOINT")?;
|
||||
let deployment_name: String = config.get_param("AZURE_OPENAI_DEPLOYMENT_NAME")?;
|
||||
|
||||
@@ -4,7 +4,6 @@ use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
|
||||
use super::errors::ProviderError;
|
||||
use super::retry::{ProviderRetry, RetryConfig};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::utils::emit_debug_trace;
|
||||
use anyhow::Result;
|
||||
@@ -46,7 +45,7 @@ pub struct BedrockProvider {
|
||||
}
|
||||
|
||||
impl BedrockProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
// Attempt to load config and secrets to get AWS_ prefixed keys
|
||||
@@ -63,15 +62,14 @@ impl BedrockProvider {
|
||||
set_aws_env_vars(config.load_values());
|
||||
set_aws_env_vars(config.load_secrets());
|
||||
|
||||
let sdk_config = futures::executor::block_on(aws_config::load_from_env());
|
||||
let sdk_config = aws_config::load_from_env().await;
|
||||
|
||||
// validate credentials or return error back up
|
||||
futures::executor::block_on(
|
||||
sdk_config
|
||||
.credentials_provider()
|
||||
.unwrap()
|
||||
.provide_credentials(),
|
||||
)?;
|
||||
sdk_config
|
||||
.credentials_provider()
|
||||
.unwrap()
|
||||
.provide_credentials()
|
||||
.await?;
|
||||
let client = Client::new(&sdk_config);
|
||||
|
||||
let retry_config = Self::load_retry_config(config);
|
||||
@@ -172,8 +170,6 @@ impl BedrockProvider {
|
||||
}
|
||||
}
|
||||
|
||||
impl_provider_default!(BedrockProvider);
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for BedrockProvider {
|
||||
fn metadata() -> ProviderMetadata {
|
||||
|
||||
@@ -12,7 +12,6 @@ use super::errors::ProviderError;
|
||||
use super::utils::emit_debug_trace;
|
||||
use crate::config::Config;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::impl_provider_default;
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
@@ -27,10 +26,8 @@ pub struct ClaudeCodeProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(ClaudeCodeProvider);
|
||||
|
||||
impl ClaudeCodeProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let command: String = config
|
||||
.get_param("CLAUDE_CODE_COMMAND")
|
||||
@@ -518,16 +515,6 @@ mod tests {
|
||||
use super::ModelConfig;
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_model_config() {
|
||||
let provider = ClaudeCodeProvider::default();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "claude-sonnet-4-20250514");
|
||||
// Context limit should be set by the ModelConfig
|
||||
assert!(config.context_limit() > 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_permission_mode_flag_construction() {
|
||||
// Test that in auto mode, the --permission-mode acceptEdits flag is added
|
||||
@@ -540,21 +527,21 @@ mod tests {
|
||||
std::env::remove_var("GOOSE_MODE");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_invalid_model_no_fallback() {
|
||||
#[tokio::test]
|
||||
async fn test_claude_code_invalid_model_no_fallback() {
|
||||
// Test that an invalid model is kept as-is (no fallback)
|
||||
let invalid_model = ModelConfig::new_or_fail("invalid-model");
|
||||
let provider = ClaudeCodeProvider::from_env(invalid_model).unwrap();
|
||||
let provider = ClaudeCodeProvider::from_env(invalid_model).await.unwrap();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "invalid-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_claude_code_valid_model() {
|
||||
#[tokio::test]
|
||||
async fn test_claude_code_valid_model() {
|
||||
// Test that a valid model is preserved
|
||||
let valid_model = ModelConfig::new_or_fail("sonnet");
|
||||
let provider = ClaudeCodeProvider::from_env(valid_model).unwrap();
|
||||
let provider = ClaudeCodeProvider::from_env(valid_model).await.unwrap();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "sonnet");
|
||||
|
||||
@@ -11,7 +11,6 @@ use super::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::errors::ProviderError;
|
||||
use super::utils::emit_debug_trace;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::impl_provider_default;
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
@@ -26,10 +25,8 @@ pub struct CursorAgentProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(CursorAgentProvider);
|
||||
|
||||
impl CursorAgentProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let command: String = config
|
||||
.get_param("CURSOR_AGENT_COMMAND")
|
||||
@@ -450,46 +447,13 @@ mod tests {
|
||||
use super::ModelConfig;
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_cursor_agent_model_config() {
|
||||
let provider = CursorAgentProvider::default();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "auto");
|
||||
// Context limit should be set by the ModelConfig
|
||||
assert!(config.context_limit() > 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cursor_agent_invalid_model_no_fallback() {
|
||||
// Test that an invalid model is kept as-is (no fallback)
|
||||
let invalid_model = ModelConfig::new_or_fail("invalid-model");
|
||||
let provider = CursorAgentProvider::from_env(invalid_model).unwrap();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "invalid-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_cursor_agent_valid_model() {
|
||||
#[tokio::test]
|
||||
async fn test_cursor_agent_valid_model() {
|
||||
// Test that a valid model is preserved
|
||||
let valid_model = ModelConfig::new_or_fail("gpt-5");
|
||||
let provider = CursorAgentProvider::from_env(valid_model).unwrap();
|
||||
let provider = CursorAgentProvider::from_env(valid_model).await.unwrap();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "gpt-5");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_filter_extensions_from_system_prompt() {
|
||||
let provider = CursorAgentProvider::default();
|
||||
|
||||
let system_with_extensions = "Some system prompt\n\n# Extensions\nSome extension info\n\n# Next Section\nMore content";
|
||||
let filtered = provider.filter_extensions_from_system_prompt(system_with_extensions);
|
||||
assert_eq!(filtered, "Some system prompt\n# Next Section\nMore content");
|
||||
|
||||
let system_without_extensions = "Some system prompt\n\n# Other Section\nContent";
|
||||
let filtered = provider.filter_extensions_from_system_prompt(system_without_extensions);
|
||||
assert_eq!(filtered, system_without_extensions);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -21,7 +21,6 @@ use super::utils::{
|
||||
};
|
||||
use crate::config::ConfigError;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{get_usage, response_to_streaming_message};
|
||||
use crate::providers::retry::{
|
||||
@@ -109,10 +108,8 @@ pub struct DatabricksProvider {
|
||||
retry_config: RetryConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(DatabricksProvider);
|
||||
|
||||
impl DatabricksProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
let mut host: Result<String, ConfigError> = config.get_param("DATABRICKS_HOST");
|
||||
|
||||
@@ -28,56 +28,63 @@ use super::{
|
||||
use crate::config::custom_providers::{custom_providers_dir, register_custom_providers};
|
||||
use crate::model::ModelConfig;
|
||||
use anyhow::Result;
|
||||
use once_cell::sync::Lazy;
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
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::<GithubCopilotProvider, _>(GithubCopilotProvider::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::<TetrateProvider, _>(TetrateProvider::from_env);
|
||||
registry.register::<VeniceProvider, _>(VeniceProvider::from_env);
|
||||
registry.register::<XaiProvider, _>(XaiProvider::from_env);
|
||||
static REGISTRY: OnceCell<RwLock<ProviderRegistry>> = OnceCell::const_new();
|
||||
|
||||
if let Err(e) = load_custom_providers_into_registry(registry) {
|
||||
tracing::warn!("Failed to load custom providers: {}", e);
|
||||
}
|
||||
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)));
|
||||
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)));
|
||||
});
|
||||
if let Err(e) = load_custom_providers_into_registry(&mut 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> {
|
||||
REGISTRY.read().unwrap().all_metadata()
|
||||
async fn get_registry() -> &'static RwLock<ProviderRegistry> {
|
||||
REGISTRY.get_or_init(init_registry).await
|
||||
}
|
||||
|
||||
pub fn refresh_custom_providers() -> Result<()> {
|
||||
let mut registry = REGISTRY.write().unwrap();
|
||||
registry.remove_custom_providers();
|
||||
pub async fn providers() -> Vec<ProviderMetadata> {
|
||||
get_registry().await.read().unwrap().all_metadata()
|
||||
}
|
||||
|
||||
if let Err(e) = load_custom_providers_into_registry(&mut registry) {
|
||||
pub async fn refresh_custom_providers() -> Result<()> {
|
||||
let registry = get_registry().await;
|
||||
registry.write().unwrap().remove_custom_providers();
|
||||
|
||||
if let Err(e) = load_custom_providers_into_registry(&mut registry.write().unwrap()) {
|
||||
tracing::warn!("Failed to refresh custom providers: {}", e);
|
||||
return Err(e);
|
||||
}
|
||||
@@ -86,18 +93,36 @@ pub fn refresh_custom_providers() -> Result<()> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn create(name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
|
||||
pub async fn create(name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
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);
|
||||
return create_lead_worker_from_env(name, &model, &lead_model_name).await;
|
||||
}
|
||||
|
||||
REGISTRY.read().unwrap().create(name, model)
|
||||
let registry = get_registry().await;
|
||||
let constructor = {
|
||||
let guard = registry.read().unwrap();
|
||||
guard
|
||||
.entries
|
||||
.get(name)
|
||||
.ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", name))?
|
||||
.constructor
|
||||
.clone()
|
||||
};
|
||||
constructor(model).await
|
||||
}
|
||||
|
||||
fn create_lead_worker_from_env(
|
||||
pub async fn create_with_named_model(
|
||||
provider_name: &str,
|
||||
model_name: &str,
|
||||
) -> Result<Arc<dyn Provider>> {
|
||||
let config = ModelConfig::new(model_name)?;
|
||||
create(provider_name, config).await
|
||||
}
|
||||
|
||||
async fn create_lead_worker_from_env(
|
||||
default_provider_name: &str,
|
||||
default_model: &ModelConfig,
|
||||
lead_model_name: &str,
|
||||
@@ -125,14 +150,30 @@ fn create_lead_worker_from_env(
|
||||
|
||||
let worker_model_config = create_worker_model_config(default_model)?;
|
||||
|
||||
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)?;
|
||||
let registry = get_registry().await;
|
||||
|
||||
let lead_constructor = {
|
||||
let guard = registry.read().unwrap();
|
||||
guard
|
||||
.entries
|
||||
.get(&lead_provider_name)
|
||||
.ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", lead_provider_name))?
|
||||
.constructor
|
||||
.clone()
|
||||
};
|
||||
|
||||
let worker_constructor = {
|
||||
let guard = registry.read().unwrap();
|
||||
guard
|
||||
.entries
|
||||
.get(default_provider_name)
|
||||
.ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", default_provider_name))?
|
||||
.constructor
|
||||
.clone()
|
||||
};
|
||||
|
||||
let lead_provider = lead_constructor(lead_model_config).await?;
|
||||
let worker_provider = worker_constructor(worker_model_config).await?;
|
||||
|
||||
Ok(Arc::new(LeadWorkerProvider::new_with_settings(
|
||||
lead_provider,
|
||||
@@ -205,8 +246,8 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_lead_worker_provider() {
|
||||
#[tokio::test]
|
||||
async fn test_create_lead_worker_provider() {
|
||||
let _guard = EnvVarGuard::new(&[
|
||||
"GOOSE_LEAD_MODEL",
|
||||
"GOOSE_LEAD_PROVIDER",
|
||||
@@ -216,7 +257,7 @@ mod tests {
|
||||
_guard.set("GOOSE_LEAD_MODEL", "gpt-4o");
|
||||
|
||||
let gpt4mini_config = ModelConfig::new_or_fail("gpt-4o-mini");
|
||||
let result = create("openai", gpt4mini_config.clone());
|
||||
let result = create("openai", gpt4mini_config.clone()).await;
|
||||
|
||||
match result {
|
||||
Ok(_) => {}
|
||||
@@ -229,11 +270,11 @@ mod tests {
|
||||
_guard.set("GOOSE_LEAD_PROVIDER", "anthropic");
|
||||
_guard.set("GOOSE_LEAD_TURNS", "5");
|
||||
|
||||
let _result = create("openai", gpt4mini_config);
|
||||
let _result = create("openai", gpt4mini_config).await;
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_lead_model_env_vars_with_defaults() {
|
||||
#[tokio::test]
|
||||
async fn test_lead_model_env_vars_with_defaults() {
|
||||
let _guard = EnvVarGuard::new(&[
|
||||
"GOOSE_LEAD_MODEL",
|
||||
"GOOSE_LEAD_PROVIDER",
|
||||
@@ -244,7 +285,7 @@ mod tests {
|
||||
|
||||
_guard.set("GOOSE_LEAD_MODEL", "grok-3");
|
||||
|
||||
let result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini"));
|
||||
let result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini")).await;
|
||||
|
||||
match result {
|
||||
Ok(_) => {}
|
||||
@@ -261,8 +302,8 @@ mod tests {
|
||||
let _result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_regular_provider_without_lead_config() {
|
||||
#[tokio::test]
|
||||
async fn test_create_regular_provider_without_lead_config() {
|
||||
let _guard = EnvVarGuard::new(&[
|
||||
"GOOSE_LEAD_MODEL",
|
||||
"GOOSE_LEAD_PROVIDER",
|
||||
@@ -271,7 +312,7 @@ mod tests {
|
||||
"GOOSE_LEAD_FALLBACK_TURNS",
|
||||
]);
|
||||
|
||||
let result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini"));
|
||||
let result = create("openai", ModelConfig::new_or_fail("gpt-4o-mini")).await;
|
||||
|
||||
match result {
|
||||
Ok(_) => {}
|
||||
|
||||
@@ -18,7 +18,6 @@ use crate::providers::formats::gcpvertexai::{
|
||||
ModelProvider, RequestContext,
|
||||
};
|
||||
|
||||
use crate::impl_provider_default;
|
||||
use crate::providers::formats::gcpvertexai::GcpLocation::Iowa;
|
||||
use crate::providers::gcpauth::GcpAuth;
|
||||
use crate::providers::retry::RetryConfig;
|
||||
@@ -87,23 +86,7 @@ impl GcpVertexAIProvider {
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model` - Configuration for the model to be used
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
Self::new(model)
|
||||
}
|
||||
|
||||
/// Creates a new provider instance with the specified model configuration.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model` - Configuration for the model to be used
|
||||
pub fn new(model: ModelConfig) -> Result<Self> {
|
||||
futures::executor::block_on(Self::new_async(model))
|
||||
}
|
||||
|
||||
/// Async implementation of new provider instance creation.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `model` - Configuration for the model to be used
|
||||
async fn new_async(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let project_id = config.get_param("GCP_PROJECT_ID")?;
|
||||
let location = Self::determine_location(config)?;
|
||||
@@ -445,8 +428,6 @@ impl GcpVertexAIProvider {
|
||||
}
|
||||
}
|
||||
|
||||
impl_provider_default!(GcpVertexAIProvider);
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for GcpVertexAIProvider {
|
||||
/// Returns metadata about the GCP Vertex AI provider.
|
||||
|
||||
@@ -10,7 +10,7 @@ use super::base::{Provider, ProviderMetadata, ProviderUsage, Usage};
|
||||
use super::errors::ProviderError;
|
||||
use super::utils::emit_debug_trace;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::Role;
|
||||
use rmcp::model::Tool;
|
||||
@@ -26,10 +26,8 @@ pub struct GeminiCliProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(GeminiCliProvider);
|
||||
|
||||
impl GeminiCliProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let command: String = config
|
||||
.get_param("GEMINI_CLI_COMMAND")
|
||||
@@ -364,31 +362,21 @@ impl Provider for GeminiCliProvider {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_gemini_cli_model_config() {
|
||||
let provider = GeminiCliProvider::default();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "gemini-2.5-pro");
|
||||
// Context limit should be set by the ModelConfig
|
||||
assert!(config.context_limit() > 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_gemini_cli_invalid_model_no_fallback() {
|
||||
#[tokio::test]
|
||||
async fn test_gemini_cli_invalid_model_no_fallback() {
|
||||
// Test that an invalid model is kept as-is (no fallback)
|
||||
let invalid_model = ModelConfig::new_or_fail("invalid-model");
|
||||
let provider = GeminiCliProvider::from_env(invalid_model).unwrap();
|
||||
let provider = GeminiCliProvider::from_env(invalid_model).await.unwrap();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, "invalid-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_gemini_cli_valid_model() {
|
||||
#[tokio::test]
|
||||
async fn test_gemini_cli_valid_model() {
|
||||
// Test that a valid model is preserved
|
||||
let valid_model = ModelConfig::new_or_fail(GEMINI_CLI_DEFAULT_MODEL);
|
||||
let provider = GeminiCliProvider::from_env(valid_model).unwrap();
|
||||
let provider = GeminiCliProvider::from_env(valid_model).await.unwrap();
|
||||
let config = provider.get_model_config();
|
||||
|
||||
assert_eq!(config.model_name, GEMINI_CLI_DEFAULT_MODEL);
|
||||
|
||||
@@ -19,7 +19,7 @@ use super::utils::{emit_debug_trace, get_model, handle_response_openai_compat, I
|
||||
|
||||
use crate::config::{Config, ConfigError};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::base::ConfigKey;
|
||||
use rmcp::model::Tool;
|
||||
@@ -115,10 +115,8 @@ pub struct GithubCopilotProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(GithubCopilotProvider);
|
||||
|
||||
impl GithubCopilotProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let client = Client::builder()
|
||||
.timeout(Duration::from_secs(600))
|
||||
.build()?;
|
||||
|
||||
@@ -3,7 +3,7 @@ use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{emit_debug_trace, handle_response_google_compat, unescape_json_values};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::base::{ConfigKey, Provider, ProviderMetadata, ProviderUsage};
|
||||
use crate::providers::formats::google::{create_request, get_usage, response_to_message};
|
||||
@@ -52,10 +52,8 @@ pub struct GoogleProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(GoogleProvider);
|
||||
|
||||
impl GoogleProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let model = model.with_fast(GOOGLE_DEFAULT_FAST_MODEL.to_string());
|
||||
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
@@ -3,7 +3,6 @@ use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{get_model, handle_response_openai_compat};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
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};
|
||||
@@ -30,10 +29,8 @@ pub struct GroqProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(GroqProvider);
|
||||
|
||||
impl GroqProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
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
|
||||
|
||||
@@ -10,7 +10,7 @@ use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{emit_debug_trace, get_model, handle_response_openai_compat, ImageFormat};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
@@ -25,10 +25,8 @@ pub struct LiteLLMProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(LiteLLMProvider);
|
||||
|
||||
impl LiteLLMProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let api_key: String = config
|
||||
.get_secret("LITELLM_API_KEY")
|
||||
|
||||
@@ -37,4 +37,4 @@ pub mod utils_universal_openai_stream;
|
||||
pub mod venice;
|
||||
pub mod xai;
|
||||
|
||||
pub use factory::{create, providers, refresh_custom_providers};
|
||||
pub use factory::{create, create_with_named_model, providers, refresh_custom_providers};
|
||||
|
||||
@@ -6,7 +6,7 @@ use super::utils::{get_model, handle_response_openai_compat, handle_status_opena
|
||||
use crate::config::custom_providers::CustomProviderConfig;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::conversation::Conversation;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{
|
||||
create_request, get_usage, response_to_message, response_to_streaming_message,
|
||||
@@ -43,10 +43,8 @@ pub struct OllamaProvider {
|
||||
supports_streaming: bool,
|
||||
}
|
||||
|
||||
impl_provider_default!(OllamaProvider);
|
||||
|
||||
impl OllamaProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let host: String = config
|
||||
.get_param("OLLAMA_HOST")
|
||||
|
||||
@@ -22,7 +22,7 @@ use super::utils::{
|
||||
};
|
||||
use crate::config::custom_providers::CustomProviderConfig;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::base::MessageStream;
|
||||
use crate::providers::formats::openai::response_to_streaming_message;
|
||||
@@ -56,10 +56,8 @@ pub struct OpenAiProvider {
|
||||
supports_streaming: bool,
|
||||
}
|
||||
|
||||
impl_provider_default!(OpenAiProvider);
|
||||
|
||||
impl OpenAiProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let model = model.with_fast(OPEN_AI_DEFAULT_FAST_MODEL.to_string());
|
||||
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
@@ -11,7 +11,7 @@ use super::utils::{
|
||||
is_google_model,
|
||||
};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use rmcp::model::Tool;
|
||||
@@ -42,10 +42,8 @@ pub struct OpenRouterProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(OpenRouterProvider);
|
||||
|
||||
impl OpenRouterProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let model = model.with_fast(OPENROUTER_DEFAULT_FAST_MODEL.to_string());
|
||||
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
@@ -1,19 +1,21 @@
|
||||
use super::base::{Provider, ProviderMetadata};
|
||||
use crate::model::ModelConfig;
|
||||
use anyhow::Result;
|
||||
use futures::future::BoxFuture;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
type ProviderConstructor = Box<dyn Fn(ModelConfig) -> Result<Arc<dyn Provider>> + Send + Sync>;
|
||||
type ProviderConstructor =
|
||||
Arc<dyn Fn(ModelConfig) -> BoxFuture<'static, Result<Arc<dyn Provider>>> + Send + Sync>;
|
||||
|
||||
struct ProviderEntry {
|
||||
pub struct ProviderEntry {
|
||||
metadata: ProviderMetadata,
|
||||
constructor: ProviderConstructor,
|
||||
pub(crate) constructor: ProviderConstructor,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct ProviderRegistry {
|
||||
entries: HashMap<String, ProviderEntry>,
|
||||
pub(crate) entries: HashMap<String, ProviderEntry>,
|
||||
}
|
||||
|
||||
impl ProviderRegistry {
|
||||
@@ -26,7 +28,7 @@ impl ProviderRegistry {
|
||||
pub fn register<P, F>(&mut self, constructor: F)
|
||||
where
|
||||
P: Provider + 'static,
|
||||
F: Fn(ModelConfig) -> Result<P> + Send + Sync + 'static,
|
||||
F: Fn(ModelConfig) -> BoxFuture<'static, Result<P>> + Send + Sync + 'static,
|
||||
{
|
||||
let metadata = P::metadata();
|
||||
let name = metadata.name.clone();
|
||||
@@ -35,12 +37,17 @@ impl ProviderRegistry {
|
||||
name,
|
||||
ProviderEntry {
|
||||
metadata,
|
||||
constructor: Box::new(move |model| Ok(Arc::new(constructor(model)?))),
|
||||
constructor: Arc::new(move |model| {
|
||||
let fut = constructor(model);
|
||||
Box::pin(async move {
|
||||
let provider = fut.await?;
|
||||
Ok(Arc::new(provider) as Arc<dyn Provider>)
|
||||
})
|
||||
}),
|
||||
},
|
||||
);
|
||||
}
|
||||
|
||||
/// create provider with custom name
|
||||
pub fn register_with_name<P, F>(
|
||||
&mut self,
|
||||
custom_name: String,
|
||||
@@ -68,7 +75,13 @@ impl ProviderRegistry {
|
||||
custom_name,
|
||||
ProviderEntry {
|
||||
metadata: custom_metadata,
|
||||
constructor: Box::new(move |model| Ok(Arc::new(constructor(model)?))),
|
||||
constructor: Arc::new(move |model| {
|
||||
let result = constructor(model);
|
||||
Box::pin(async move {
|
||||
let provider = result?;
|
||||
Ok(Arc::new(provider) as Arc<dyn Provider>)
|
||||
})
|
||||
}),
|
||||
},
|
||||
);
|
||||
}
|
||||
@@ -81,15 +94,13 @@ impl ProviderRegistry {
|
||||
self
|
||||
}
|
||||
|
||||
pub fn create(&self, name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
|
||||
let _available_providers: Vec<_> = self.entries.keys().collect();
|
||||
|
||||
pub async fn create(&self, name: &str, model: ModelConfig) -> Result<Arc<dyn Provider>> {
|
||||
let entry = self
|
||||
.entries
|
||||
.get(name)
|
||||
.ok_or_else(|| anyhow::anyhow!("Unknown provider: {}", name))?;
|
||||
|
||||
(entry.constructor)(model)
|
||||
(entry.constructor)(model).await
|
||||
}
|
||||
|
||||
pub fn all_metadata(&self) -> Vec<ProviderMetadata> {
|
||||
|
||||
@@ -14,7 +14,7 @@ use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::emit_debug_trace;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use chrono::Utc;
|
||||
use rmcp::model::Role;
|
||||
@@ -33,7 +33,7 @@ pub struct SageMakerTgiProvider {
|
||||
}
|
||||
|
||||
impl SageMakerTgiProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
|
||||
// Get SageMaker endpoint name (just the name, not full URL)
|
||||
@@ -54,15 +54,14 @@ impl SageMakerTgiProvider {
|
||||
set_aws_env_vars(config.load_values());
|
||||
set_aws_env_vars(config.load_secrets());
|
||||
|
||||
let aws_config = futures::executor::block_on(aws_config::load_from_env());
|
||||
let aws_config = aws_config::load_from_env().await;
|
||||
|
||||
// Validate credentials
|
||||
futures::executor::block_on(
|
||||
aws_config
|
||||
.credentials_provider()
|
||||
.unwrap()
|
||||
.provide_credentials(),
|
||||
)?;
|
||||
aws_config
|
||||
.credentials_provider()
|
||||
.unwrap()
|
||||
.provide_credentials()
|
||||
.await?;
|
||||
|
||||
// Create client with longer timeout for model initialization
|
||||
let timeout_config = aws_config::timeout::TimeoutConfig::builder()
|
||||
@@ -255,8 +254,6 @@ impl SageMakerTgiProvider {
|
||||
}
|
||||
}
|
||||
|
||||
impl_provider_default!(SageMakerTgiProvider);
|
||||
|
||||
#[async_trait]
|
||||
impl Provider for SageMakerTgiProvider {
|
||||
fn metadata() -> ProviderMetadata {
|
||||
|
||||
@@ -11,7 +11,7 @@ use super::retry::ProviderRetry;
|
||||
use super::utils::{get_model, map_http_error_to_provider_error, ImageFormat};
|
||||
use crate::config::ConfigError;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::Tool;
|
||||
|
||||
@@ -40,10 +40,8 @@ pub struct SnowflakeProvider {
|
||||
image_format: ImageFormat,
|
||||
}
|
||||
|
||||
impl_provider_default!(SnowflakeProvider);
|
||||
|
||||
impl SnowflakeProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let mut host: Result<String, ConfigError> = config.get_param("SNOWFLAKE_HOST");
|
||||
if host.is_err() {
|
||||
|
||||
@@ -20,7 +20,7 @@ use super::utils::{
|
||||
};
|
||||
use crate::config::signup_tetrate::TETRATE_DEFAULT_MODEL;
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::model::ModelConfig;
|
||||
use crate::providers::formats::openai::{create_request, get_usage, response_to_message};
|
||||
use rmcp::model::Tool;
|
||||
@@ -48,10 +48,8 @@ pub struct TetrateProvider {
|
||||
supports_streaming: bool,
|
||||
}
|
||||
|
||||
impl_provider_default!(TetrateProvider);
|
||||
|
||||
impl TetrateProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let api_key: String = config.get_secret("TETRATE_API_KEY")?;
|
||||
// API host for LLM endpoints (/v1/chat/completions, /v1/models)
|
||||
|
||||
@@ -10,7 +10,7 @@ use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::map_http_error_to_provider_error;
|
||||
use crate::conversation::message::{Message, MessageContent};
|
||||
use crate::impl_provider_default;
|
||||
|
||||
use crate::mcp_utils::ToolResult;
|
||||
use crate::model::ModelConfig;
|
||||
use rmcp::model::{object, CallToolRequestParam, Role, Tool};
|
||||
@@ -80,10 +80,8 @@ pub struct VeniceProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(VeniceProvider);
|
||||
|
||||
impl VeniceProvider {
|
||||
pub fn from_env(mut model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(mut model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let api_key: String = config.get_secret("VENICE_API_KEY")?;
|
||||
let host: String = config
|
||||
|
||||
@@ -3,7 +3,7 @@ use super::errors::ProviderError;
|
||||
use super::retry::ProviderRetry;
|
||||
use super::utils::{get_model, handle_response_openai_compat};
|
||||
use crate::conversation::message::Message;
|
||||
use crate::impl_provider_default;
|
||||
|
||||
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};
|
||||
@@ -44,10 +44,8 @@ pub struct XaiProvider {
|
||||
model: ModelConfig,
|
||||
}
|
||||
|
||||
impl_provider_default!(XaiProvider);
|
||||
|
||||
impl XaiProvider {
|
||||
pub fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
pub async fn from_env(model: ModelConfig) -> Result<Self> {
|
||||
let config = crate::config::Config::global();
|
||||
let api_key: String = config.get_secret("XAI_API_KEY")?;
|
||||
let host: String = config
|
||||
|
||||
Reference in New Issue
Block a user