fix: adds ProviderRetry to openai provider (#5518)

Signed-off-by: Matt Yaple <matt@yaple.dev>
This commit is contained in:
Matt Yaple
2025-11-01 12:04:23 +00:00
committed by GitHub
parent 93f92e94f7
commit c7cd9b424c
+59 -24
View File
@@ -16,6 +16,7 @@ use super::base::{ConfigKey, ModelInfo, Provider, ProviderMetadata, ProviderUsag
use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse}; use super::embedding::{EmbeddingCapable, EmbeddingRequest, EmbeddingResponse};
use super::errors::ProviderError; use super::errors::ProviderError;
use super::formats::openai::{create_request, get_usage, response_to_message}; use super::formats::openai::{create_request, get_usage, response_to_message};
use super::retry::ProviderRetry;
use super::utils::{ use super::utils::{
get_model, handle_response_openai_compat, handle_status_openai_compat, ImageFormat, get_model, handle_response_openai_compat, handle_status_openai_compat, ImageFormat,
}; };
@@ -240,9 +241,15 @@ impl Provider for OpenAiProvider {
let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?; let payload = create_request(model_config, system, messages, tools, &ImageFormat::OpenAi)?;
let mut log = RequestLog::start(&self.model, &payload)?; let mut log = RequestLog::start(&self.model, &payload)?;
let json_response = self.post(&payload).await.inspect_err(|e| { let json_response = self
let _ = log.error(e); .with_retry(|| async {
})?; let payload_clone = payload.clone();
self.post(&payload_clone).await
})
.await
.inspect_err(|e| {
let _ = log.error(e);
})?;
let message = response_to_message(&json_response)?; let message = response_to_message(&json_response)?;
let usage = json_response let usage = json_response
@@ -260,19 +267,30 @@ 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.api_client.response_get(&models_path).await?; let response = self
let json = handle_response_openai_compat(response).await?; .with_retry(|| async {
if let Some(err_obj) = json.get("error") { let response = self.api_client.response_get(&models_path).await?;
let msg = err_obj let json = handle_response_openai_compat(response).await?;
.get("message") if let Some(err_obj) = json.get("error") {
.and_then(|v| v.as_str()) let msg = err_obj
.unwrap_or("unknown error"); .get("message")
return Err(ProviderError::Authentication(msg.to_string())); .and_then(|v| v.as_str())
} .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 = json.get("data").and_then(|v| v.as_array()).ok_or_else(|| { let data = response
ProviderError::UsageError("Missing data field in JSON response".into()) .get("data")
})?; .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))
@@ -310,17 +328,24 @@ impl Provider for OpenAiProvider {
let mut log = RequestLog::start(&self.model, &payload)?; let mut log = RequestLog::start(&self.model, &payload)?;
let response = self let response = self
.api_client .with_retry(|| async {
.response_post(&self.base_path, &payload) let resp = self
.await .api_client
.inspect_err(|e| { .response_post(&self.base_path, &payload)
let _ = log.error(e); .await?;
})?; let status = resp.status();
let response = handle_status_openai_compat(response) if !status.is_success() {
return Err(super::utils::map_http_error_to_provider_error(
status, None, // We'll let handle_status_openai_compat parse the error
));
}
Ok(resp)
})
.await .await
.inspect_err(|e| { .inspect_err(|e| {
let _ = log.error(e); let _ = log.error(e);
})?; })?;
let response = handle_status_openai_compat(response).await?;
let stream = response.bytes_stream().map_err(io::Error::other); let stream = response.bytes_stream().map_err(io::Error::other);
@@ -366,8 +391,18 @@ impl EmbeddingCapable for OpenAiProvider {
}; };
let response = self let response = self
.api_client .with_retry(|| async {
.api_post("v1/embeddings", &serde_json::to_value(request)?) let request_clone = EmbeddingRequest {
input: request.input.clone(),
model: request.model.clone(),
};
let request_value = serde_json::to_value(request_clone)
.map_err(|e| ProviderError::ExecutionError(e.to_string()))?;
self.api_client
.api_post("v1/embeddings", &request_value)
.await
.map_err(|e| ProviderError::ExecutionError(e.to_string()))
})
.await?; .await?;
if response.status != StatusCode::OK { if response.status != StatusCode::OK {