Make async (#5126)

Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
Douwe Osinga
2025-10-10 19:31:50 -04:00
committed by GitHub
parent 69a7b7fe5a
commit 1015064367
43 changed files with 343 additions and 480 deletions
+2 -5
View File
@@ -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();
+1 -4
View File
@@ -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")?;
+7 -11
View File
@@ -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 {
+7 -20
View File
@@ -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");
+4 -40
View File
@@ -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);
}
}
+1 -4
View File
@@ -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");
+96 -55
View File
@@ -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(_) => {}
+1 -20
View File
@@ -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.
+8 -20
View File
@@ -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);
+2 -4
View File
@@ -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()?;
+2 -4
View File
@@ -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();
+1 -4
View File
@@ -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
+2 -4
View File
@@ -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")
+1 -1
View File
@@ -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};
+2 -4
View File
@@ -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")
+2 -4
View File
@@ -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();
+2 -4
View File
@@ -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();
+23 -12
View File
@@ -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> {
+8 -11
View File
@@ -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 {
+2 -4
View File
@@ -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() {
+2 -4
View File
@@ -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)
+2 -4
View File
@@ -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
+2 -4
View File
@@ -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