diff --git a/server/src/provider/anthropic.rs b/server/src/provider/anthropic.rs index 1a50d22..1b85d09 100644 --- a/server/src/provider/anthropic.rs +++ b/server/src/provider/anthropic.rs @@ -11,8 +11,10 @@ use crate::{ }; use super::{ - merge_extra_params, recorder::recorded_headers, CallRecorder, FinishReason, ModelEvent, - Provider, ProviderStream, + merge_extra_params, + recorder::recorded_headers, + retry::{send_with_retry, Attempt, RetryPolicy}, + CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream, }; const DEFAULT_MAX_OUTPUT_TOKENS: u64 = 65_000; @@ -83,25 +85,17 @@ impl Provider for AnthropicProvider { if let Some(recorder) = &recorder { recorder.request(recorded_headers(&config, &[("content-type", "application/json"), ("anthropic-version", "2023-06-01")]), &body).await?; } - let request = client.post(&config.request_url) - .header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01") - .headers(config.custom_headers.clone()) - .json(&body).send(); - let response = tokio::select! { - _ = cancellation.cancelled() => return, - response = request => response, - }; - let response = response?; - if let Some(recorder) = &recorder { - recorder.response_headers(response.status().as_u16()).await?; - } - if !response.status().is_success() { - let status = response.status(); let bytes = response.bytes().await?; - if let Some(recorder) = &recorder { recorder.response_chunk(&bytes).await?; } - let text = String::from_utf8_lossy(&bytes); - Err(Error::Provider(format!("Anthropic {status}: {text}")))?; - return; - } + let attempt = send_with_retry( + "Anthropic", + || client.post(&config.request_url) + .header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01") + .headers(config.custom_headers.clone()) + .json(&body), + RetryPolicy::default(), + &cancellation, + recorder.as_ref(), + ).await?; + let Attempt::Response(response) = attempt else { return }; yield ModelEvent::Start { model_call_id: call_id }; let chunk_recorder = recorder.clone(); let chunks = response.bytes_stream() diff --git a/server/src/provider/mod.rs b/server/src/provider/mod.rs index b9fc7f7..ea07298 100644 --- a/server/src/provider/mod.rs +++ b/server/src/provider/mod.rs @@ -4,6 +4,7 @@ mod normalize; mod openai_chat; mod openai_responses; mod recorder; +mod retry; mod router; use std::pin::Pin; diff --git a/server/src/provider/openai_chat.rs b/server/src/provider/openai_chat.rs index b5da1f9..d5d86a7 100644 --- a/server/src/provider/openai_chat.rs +++ b/server/src/provider/openai_chat.rs @@ -16,8 +16,9 @@ use crate::{ }; use super::{ - apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, CallRecorder, - FinishReason, ModelEvent, Provider, ProviderStream, + apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, + retry::{send_with_retry, Attempt, RetryPolicy}, + CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream, }; #[derive(Default)] @@ -86,33 +87,15 @@ impl Provider for OpenAiChatProvider { if let Some(recorder) = &recorder { recorder.request(recorded_headers(&config, &[("content-type", "application/json")]), &body).await?; } - let request = client.post(&config.request_url) - .bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body).send(); - let response = tokio::select! { - _ = cancellation.cancelled() => return, - response = request => response, - }; - let response = match response { - Ok(r) => { - tracing::debug!(status = r.status().as_u16(), "OpenAI Chat HTTP response received"); - r - } - Err(e) => { - tracing::debug!(error = %e, "OpenAI Chat HTTP request failed"); - Err(Error::from(e))? - } - }; - if let Some(recorder) = &recorder { - recorder.response_headers(response.status().as_u16()).await?; - } - if !response.status().is_success() { - let status = response.status(); - let bytes = response.bytes().await?; - if let Some(recorder) = &recorder { recorder.response_chunk(&bytes).await?; } - let text = String::from_utf8_lossy(&bytes); - Err(Error::Provider(format!("OpenAI Chat {status}: {text}")))?; - return; - } + let attempt = send_with_retry( + "OpenAI Chat", + || client.post(&config.request_url) + .bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body), + RetryPolicy::default(), + &cancellation, + recorder.as_ref(), + ).await?; + let Attempt::Response(response) = attempt else { return }; yield ModelEvent::Start { model_call_id: call_id }; let chunk_recorder = recorder.clone(); let chunks = response.bytes_stream() diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index ab6ce8d..883cd01 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -13,8 +13,9 @@ use crate::{ }; use super::{ - apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, CallRecorder, - FinishReason, ModelEvent, Provider, ProviderStream, + apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, + retry::{send_with_retry, Attempt, RetryPolicy}, + CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream, }; #[derive(Default)] @@ -83,23 +84,15 @@ impl Provider for OpenAiResponsesProvider { if let Some(recorder) = &recorder { recorder.request(recorded_headers(&config, &[("content-type", "application/json")]), &body).await?; } - let request = client.post(&config.request_url) - .bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body).send(); - let response = tokio::select! { - _ = cancellation.cancelled() => return, - response = request => response, - }; - let response = response?; - if let Some(recorder) = &recorder { - recorder.response_headers(response.status().as_u16()).await?; - } - if !response.status().is_success() { - let status = response.status(); let bytes = response.bytes().await?; - if let Some(recorder) = &recorder { recorder.response_chunk(&bytes).await?; } - let text = String::from_utf8_lossy(&bytes); - Err(Error::Provider(format!("OpenAI Responses {status}: {text}")))?; - return; - } + let attempt = send_with_retry( + "OpenAI Responses", + || client.post(&config.request_url) + .bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body), + RetryPolicy::default(), + &cancellation, + recorder.as_ref(), + ).await?; + let Attempt::Response(response) = attempt else { return }; yield ModelEvent::Start { model_call_id: call_id }; let chunk_recorder = recorder.clone(); let chunks = response.bytes_stream() diff --git a/server/src/provider/retry.rs b/server/src/provider/retry.rs new file mode 100644 index 0000000..91dbf9a --- /dev/null +++ b/server/src/provider/retry.rs @@ -0,0 +1,242 @@ +use std::time::Duration; + +use tokio_util::sync::CancellationToken; + +use crate::{Error, Result}; + +use super::CallRecorder; + +#[derive(Clone, Copy, Debug)] +pub(crate) struct RetryPolicy { + pub retries: u32, + pub delay: Duration, +} + +impl Default for RetryPolicy { + fn default() -> Self { + Self { + retries: 5, + delay: Duration::from_secs(5), + } + } +} + +#[derive(Debug)] +pub(crate) enum Attempt { + Response(reqwest::Response), + Cancelled, +} + +pub(crate) async fn send_with_retry( + label: &str, + build: F, + policy: RetryPolicy, + cancellation: &CancellationToken, + recorder: Option<&CallRecorder>, +) -> Result +where + F: Fn() -> reqwest::RequestBuilder, +{ + for attempt in 0..=policy.retries { + let response = tokio::select! { + _ = cancellation.cancelled() => return Ok(Attempt::Cancelled), + response = build().send() => response, + }?; + if let Some(recorder) = recorder { + recorder.response_headers(response.status().as_u16()).await?; + } + if response.status().is_success() { + return Ok(Attempt::Response(response)); + } + let status = response.status(); + let bytes = response.bytes().await?; + if let Some(recorder) = recorder { + recorder.response_chunk(&bytes).await?; + } + if attempt == policy.retries { + let text = String::from_utf8_lossy(&bytes); + return Err(Error::Provider(format!("{label} {status}: {text}"))); + } + tracing::warn!( + provider = label, + status = status.as_u16(), + attempt = attempt + 1, + retries = policy.retries, + delay_ms = policy.delay.as_millis(), + "provider returned a non-success status, retrying" + ); + tokio::select! { + _ = cancellation.cancelled() => return Ok(Attempt::Cancelled), + _ = tokio::time::sleep(policy.delay) => {} + } + } + unreachable!("the retry loop returns on the final attempt") +} + +#[cfg(test)] +mod tests { + use super::*; + + use std::sync::{ + atomic::{AtomicU32, Ordering}, + Arc, + }; + + use axum::{extract::State, http::StatusCode, routing::post, Router}; + + fn fast(retries: u32) -> RetryPolicy { + RetryPolicy { + retries, + delay: Duration::from_millis(20), + } + } + + async fn status_server(statuses: Vec) -> (String, Arc) { + async fn endpoint( + State((statuses, calls)): State<(Arc>, Arc)>, + ) -> (StatusCode, String) { + let index = calls.fetch_add(1, Ordering::SeqCst) as usize; + let status = statuses + .get(index) + .copied() + .unwrap_or_else(|| *statuses.last().unwrap()); + ( + StatusCode::from_u16(status).unwrap(), + format!("body for attempt {index}"), + ) + } + + let calls = Arc::new(AtomicU32::new(0)); + let app = Router::new() + .route("/responses", post(endpoint)) + .with_state((Arc::new(statuses), calls.clone())); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + (format!("http://{address}/responses"), calls) + } + + fn sender(url: String) -> impl Fn() -> reqwest::RequestBuilder { + let client = reqwest::Client::new(); + move || client.post(&url).json(&serde_json::json!({"stream": true})) + } + + #[test] + fn the_default_policy_retries_five_times_every_five_seconds() { + let policy = RetryPolicy::default(); + assert_eq!(policy.retries, 5); + assert_eq!(policy.delay, Duration::from_secs(5)); + } + + #[tokio::test] + async fn a_non_success_response_is_retried_until_it_succeeds() { + let (url, calls) = status_server(vec![429, 500, 200]).await; + + let attempt = send_with_retry( + "Test", + sender(url), + fast(5), + &CancellationToken::new(), + None, + ) + .await + .unwrap(); + + let Attempt::Response(response) = attempt else { + panic!("expected a response"); + }; + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(calls.load(Ordering::SeqCst), 3); + } + + #[tokio::test] + async fn the_last_non_success_response_fails_after_the_retry_budget() { + let (url, calls) = status_server(vec![429]).await; + + let error = send_with_retry( + "Test", + sender(url), + fast(5), + &CancellationToken::new(), + None, + ) + .await + .unwrap_err(); + + assert!( + matches!(&error, Error::Provider(message) if message.contains("Test 429")), + "unexpected error: {error}" + ); + assert_eq!(calls.load(Ordering::SeqCst), 6); + } + + #[tokio::test] + async fn every_retry_waits_for_the_configured_delay() { + let (url, _) = status_server(vec![429, 429, 200]).await; + let started = std::time::Instant::now(); + + send_with_retry( + "Test", + sender(url), + RetryPolicy { + retries: 5, + delay: Duration::from_millis(150), + }, + &CancellationToken::new(), + None, + ) + .await + .unwrap(); + + assert!( + started.elapsed() >= Duration::from_millis(300), + "retries did not wait: {:?}", + started.elapsed() + ); + } + + #[tokio::test] + async fn cancellation_during_the_retry_delay_stops_the_attempts() { + let (url, calls) = status_server(vec![429]).await; + let cancellation = CancellationToken::new(); + let deadline = cancellation.clone(); + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(100)).await; + deadline.cancel(); + }); + + let attempt = send_with_retry( + "Test", + sender(url), + RetryPolicy { + retries: 5, + delay: Duration::from_millis(500), + }, + &cancellation, + None, + ) + .await + .unwrap(); + + assert!(matches!(attempt, Attempt::Cancelled)); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn a_success_response_is_returned_without_any_retry() { + let (url, calls) = status_server(vec![200]).await; + + let attempt = send_with_retry( + "Test", + sender(url), + fast(5), + &CancellationToken::new(), + None, + ) + .await + .unwrap(); + + assert!(matches!(attempt, Attempt::Response(_))); + assert_eq!(calls.load(Ordering::SeqCst), 1); + } +}