feat(providers): add retry for model fetching (#6347)

Signed-off-by: rabi <ramishra@redhat.com>
This commit is contained in:
Rabi Mishra
2026-01-06 22:44:56 +05:30
committed by GitHub
parent bd60c4496a
commit 0788f88929
6 changed files with 90 additions and 39 deletions
+5 -3
View File
@@ -20,7 +20,7 @@ use goose::config::{
use goose::conversation::message::Message; use goose::conversation::message::Message;
use goose::model::ModelConfig; use goose::model::ModelConfig;
use goose::providers::provider_test::test_provider_configuration; use goose::providers::provider_test::test_provider_configuration;
use goose::providers::{create, providers}; use goose::providers::{create, providers, retry_operation, RetryConfig};
use goose::session::{SessionManager, SessionType}; use goose::session::{SessionManager, SessionType};
use serde_json::Value; use serde_json::Value;
use std::collections::HashMap; use std::collections::HashMap;
@@ -576,13 +576,15 @@ pub async fn configure_provider_dialog() -> anyhow::Result<bool> {
} }
} }
// Attempt to fetch supported models for this provider
let spin = spinner(); let spin = spinner();
spin.start("Attempting to fetch supported models..."); spin.start("Attempting to fetch supported models...");
let models_res = { let models_res = {
let temp_model_config = ModelConfig::new(&provider_meta.default_model)?; let temp_model_config = ModelConfig::new(&provider_meta.default_model)?;
let temp_provider = create(provider_name, temp_model_config).await?; let temp_provider = create(provider_name, temp_model_config).await?;
temp_provider.fetch_recommended_models().await retry_operation(&RetryConfig::default(), || async {
temp_provider.fetch_recommended_models().await
})
.await
}; };
spin.stop(style("Model fetch complete").green()); spin.stop(style("Model fetch complete").green());
@@ -15,7 +15,9 @@ use goose::providers::auto_detect::detect_provider_from_api_key;
use goose::providers::base::{ProviderMetadata, ProviderType}; use goose::providers::base::{ProviderMetadata, ProviderType};
use goose::providers::canonical::maybe_get_canonical_model; use goose::providers::canonical::maybe_get_canonical_model;
use goose::providers::create_with_default_model; use goose::providers::create_with_default_model;
use goose::providers::errors::ProviderError;
use goose::providers::providers as get_providers; use goose::providers::providers as get_providers;
use goose::providers::{retry_operation, RetryConfig};
use goose::{ use goose::{
agents::execute_commands, agents::ExtensionConfig, config::permission::PermissionLevel, agents::execute_commands, agents::ExtensionConfig, config::permission::PermissionLevel,
slash_commands, slash_commands,
@@ -399,13 +401,15 @@ pub async fn get_provider_models(
.await .await
.map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?; .map_err(|_| StatusCode::INTERNAL_SERVER_ERROR)?;
let models_result = provider.fetch_recommended_models().await; let models_result = retry_operation(&RetryConfig::default(), || async {
provider.fetch_recommended_models().await
})
.await;
match models_result { match models_result {
Ok(Some(models)) => Ok(Json(models)), Ok(Some(models)) => Ok(Json(models)),
Ok(None) => Ok(Json(Vec::new())), Ok(None) => Ok(Json(Vec::new())),
Err(provider_error) => { Err(provider_error) => {
use goose::providers::errors::ProviderError;
let status_code = match provider_error { let status_code = match provider_error {
// Permanent misconfigurations - client should fix configuration // Permanent misconfigurations - client should fix configuration
ProviderError::Authentication(_) => StatusCode::BAD_REQUEST, ProviderError::Authentication(_) => StatusCode::BAD_REQUEST,
+11 -4
View File
@@ -1,4 +1,5 @@
use crate::model::ModelConfig; use crate::model::ModelConfig;
use crate::providers::retry::{retry_operation, RetryConfig};
pub async fn detect_provider_from_api_key(api_key: &str) -> Option<(String, Vec<String>)> { pub async fn detect_provider_from_api_key(api_key: &str) -> Option<(String, Vec<String>)> {
let provider_tests = vec![ let provider_tests = vec![
@@ -24,10 +25,16 @@ pub async fn detect_provider_from_api_key(api_key: &str) -> Option<(String, Vec<
) )
.await .await
{ {
Ok(provider) => match provider.fetch_supported_models().await { Ok(provider) => {
Ok(Some(models)) => Some((provider_name.to_string(), models)), match retry_operation(&RetryConfig::default(), || async {
_ => None, provider.fetch_supported_models().await
}, })
.await
{
Ok(Some(models)) => Some((provider_name.to_string(), models)),
_ => None,
}
}
Err(_) => None, Err(_) => None,
}; };
+1
View File
@@ -41,3 +41,4 @@ pub mod xai;
pub use factory::{ pub use factory::{
create, create_with_default_model, create_with_named_model, providers, refresh_custom_providers, create, create_with_default_model, create_with_named_model, providers, refresh_custom_providers,
}; };
pub use retry::{retry_operation, RetryConfig};
+12 -23
View File
@@ -320,30 +320,19 @@ impl Provider for OpenAiProvider {
async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> { async fn fetch_supported_models(&self) -> Result<Option<Vec<String>>, ProviderError> {
let models_path = self.base_path.replace("v1/chat/completions", "v1/models"); let models_path = self.base_path.replace("v1/chat/completions", "v1/models");
let response = self let response = self.api_client.response_get(&models_path).await?;
.with_retry(|| async { let json = handle_response_openai_compat(response).await?;
let response = self.api_client.response_get(&models_path).await?; if let Some(err_obj) = json.get("error") {
let json = handle_response_openai_compat(response).await?; let msg = err_obj
if let Some(err_obj) = json.get("error") { .get("message")
let msg = err_obj .and_then(|v| v.as_str())
.get("message") .unwrap_or("unknown error");
.and_then(|v| v.as_str()) return Err(ProviderError::Authentication(msg.to_string()));
.unwrap_or("unknown error"); }
return Err(ProviderError::Authentication(msg.to_string()));
}
Ok(json)
})
.await
.inspect_err(|e| {
tracing::warn!("Failed to fetch supported models from OpenAI: {:?}", e);
})?;
let data = response let data = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| {
.get("data") ProviderError::UsageError("Missing data field in JSON response".into())
.and_then(|v| v.as_array()) })?;
.ok_or_else(|| {
ProviderError::UsageError("Missing data field in JSON response".into())
})?;
let mut models: Vec<String> = data let mut models: Vec<String> = data
.iter() .iter()
.filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string)) .filter_map(|m| m.get("id").and_then(|v| v.as_str()).map(str::to_string))
+55 -7
View File
@@ -48,6 +48,10 @@ impl RetryConfig {
} }
} }
pub fn max_retries(&self) -> usize {
self.max_retries
}
pub fn delay_for_attempt(&self, attempt: usize) -> Duration { pub fn delay_for_attempt(&self, attempt: usize) -> Duration {
if attempt == 0 { if attempt == 0 {
return Duration::from_millis(0); return Duration::from_millis(0);
@@ -67,6 +71,56 @@ impl RetryConfig {
} }
} }
pub fn should_retry(error: &ProviderError) -> bool {
matches!(
error,
ProviderError::RateLimitExceeded { .. }
| ProviderError::ServerError(_)
| ProviderError::RequestFailed(_)
)
}
pub async fn retry_operation<F, Fut, T>(
config: &RetryConfig,
operation: F,
) -> Result<T, ProviderError>
where
F: Fn() -> Fut + Send,
Fut: Future<Output = Result<T, ProviderError>> + Send,
T: Send,
{
let mut attempts = 0;
loop {
match operation().await {
Ok(result) => return Ok(result),
Err(error) => {
if should_retry(&error) && attempts < config.max_retries {
attempts += 1;
tracing::warn!(
"Request failed, retrying ({}/{}): {:?}",
attempts,
config.max_retries,
error
);
let delay = match &error {
ProviderError::RateLimitExceeded {
retry_delay: Some(d),
..
} => *d,
_ => config.delay_for_attempt(attempts),
};
sleep(delay).await;
continue;
}
return Err(error);
}
}
}
}
/// Trait for retry functionality to keep Provider dyn-compatible /// Trait for retry functionality to keep Provider dyn-compatible
#[async_trait] #[async_trait]
pub trait ProviderRetry { pub trait ProviderRetry {
@@ -87,12 +141,7 @@ pub trait ProviderRetry {
return match operation().await { return match operation().await {
Ok(result) => Ok(result), Ok(result) => Ok(result),
Err(error) => { Err(error) => {
let should_retry = matches!( if should_retry(&error) && attempts < config.max_retries {
error,
ProviderError::RateLimitExceeded { .. } | ProviderError::ServerError(_)
);
if should_retry && attempts < config.max_retries {
attempts += 1; attempts += 1;
tracing::warn!( tracing::warn!(
"Request failed, retrying ({}/{}): {:?}", "Request failed, retrying ({}/{}): {:?}",
@@ -130,7 +179,6 @@ pub trait ProviderRetry {
} }
} }
// Let specific providers define their retry config if desired
impl<P: Provider> ProviderRetry for P { impl<P: Provider> ProviderRetry for P {
fn retry_config(&self) -> RetryConfig { fn retry_config(&self) -> RetryConfig {
Provider::retry_config(self) Provider::retry_config(self)