fix: persist each provider retry as its own llm_calls row

Retry attempts now finish the failed call, start a new call id, and record the next request so call history shows intermediate failures.
This commit is contained in:
leokun
2026-08-28 21:09:45 +08:00
parent 6a3570a95a
commit 5a0bc2e0e9
5 changed files with 228 additions and 69 deletions
+7 -1
View File
@@ -82,8 +82,12 @@ impl Provider for AnthropicProvider {
} }
apply_model(&mut body, &request.model)?; apply_model(&mut body, &request.model)?;
merge_extra_params(&mut body, &request.model.extra_params)?; merge_extra_params(&mut body, &request.model.extra_params)?;
let request_headers = recorded_headers(
&config,
&[("content-type", "application/json"), ("anthropic-version", "2023-06-01")],
);
if let Some(recorder) = &recorder { if let Some(recorder) = &recorder {
recorder.request(recorded_headers(&config, &[("content-type", "application/json"), ("anthropic-version", "2023-06-01")]), &body).await?; recorder.request(request_headers.clone(), &body).await?;
} }
let attempt = send_with_retry( let attempt = send_with_retry(
"Anthropic", "Anthropic",
@@ -94,6 +98,8 @@ impl Provider for AnthropicProvider {
RetryPolicy::default(), RetryPolicy::default(),
&cancellation, &cancellation,
recorder.as_ref(), recorder.as_ref(),
request_headers,
&body,
).await?; ).await?;
let Attempt::Response(response) = attempt else { return }; let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id }; yield ModelEvent::Start { model_call_id: call_id };
+6 -2
View File
@@ -16,7 +16,8 @@ use crate::{
}; };
use super::{ use super::{
apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, apply_openai_prompt_cache_key, merge_extra_params,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy}, retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream, CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
}; };
@@ -84,8 +85,9 @@ impl Provider for OpenAiChatProvider {
apply_model(&mut body, &request.model, config.max_output_tokens)?; apply_model(&mut body, &request.model, config.max_output_tokens)?;
merge_extra_params(&mut body, &request.model.extra_params)?; merge_extra_params(&mut body, &request.model.extra_params)?;
apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?; apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?;
let request_headers = recorded_headers(&config, &[("content-type", "application/json")]);
if let Some(recorder) = &recorder { if let Some(recorder) = &recorder {
recorder.request(recorded_headers(&config, &[("content-type", "application/json")]), &body).await?; recorder.request(request_headers.clone(), &body).await?;
} }
let attempt = send_with_retry( let attempt = send_with_retry(
"OpenAI Chat", "OpenAI Chat",
@@ -94,6 +96,8 @@ impl Provider for OpenAiChatProvider {
RetryPolicy::default(), RetryPolicy::default(),
&cancellation, &cancellation,
recorder.as_ref(), recorder.as_ref(),
request_headers,
&body,
).await?; ).await?;
let Attempt::Response(response) = attempt else { return }; let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id }; yield ModelEvent::Start { model_call_id: call_id };
+6 -2
View File
@@ -13,7 +13,8 @@ use crate::{
}; };
use super::{ use super::{
apply_openai_prompt_cache_key, merge_extra_params, recorder::recorded_headers, apply_openai_prompt_cache_key, merge_extra_params,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy}, retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream, CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
}; };
@@ -81,8 +82,9 @@ impl Provider for OpenAiResponsesProvider {
apply_model(&mut body, &request.model, config.max_output_tokens)?; apply_model(&mut body, &request.model, config.max_output_tokens)?;
merge_extra_params(&mut body, &request.model.extra_params)?; merge_extra_params(&mut body, &request.model.extra_params)?;
apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?; apply_openai_prompt_cache_key(&mut body, &request.model.model_id)?;
let request_headers = recorded_headers(&config, &[("content-type", "application/json")]);
if let Some(recorder) = &recorder { if let Some(recorder) = &recorder {
recorder.request(recorded_headers(&config, &[("content-type", "application/json")]), &body).await?; recorder.request(request_headers.clone(), &body).await?;
} }
let attempt = send_with_retry( let attempt = send_with_retry(
"OpenAI Responses", "OpenAI Responses",
@@ -91,6 +93,8 @@ impl Provider for OpenAiResponsesProvider {
RetryPolicy::default(), RetryPolicy::default(),
&cancellation, &cancellation,
recorder.as_ref(), recorder.as_ref(),
request_headers,
&body,
).await?; ).await?;
let Attempt::Response(response) = attempt else { return }; let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id }; yield ModelEvent::Start { model_call_id: call_id };
+181 -58
View File
@@ -1,6 +1,6 @@
use std::{ use std::{
sync::{ sync::{
atomic::{AtomicBool, AtomicI64, Ordering}, atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering},
Arc, Arc,
}, },
time::Instant, time::Instant,
@@ -42,13 +42,32 @@ pub struct CallRecorder {
struct Inner { struct Inner {
store: Store, store: Store,
base_call: NewLlmCall,
detailed: bool,
attempt: Mutex<AttemptState>,
next_attempt: AtomicU32,
next_generation: AtomicU64,
finished: AtomicBool,
}
struct AttemptState {
call_id: String, call_id: String,
started: Instant, started: Instant,
detailed: bool,
next_chunk: AtomicI64, next_chunk: AtomicI64,
chunks: Mutex<ChunkBuffer>, chunks: ChunkBuffer,
first_text_recorded: AtomicBool, first_text_recorded: AtomicBool,
finished: AtomicBool, }
impl AttemptState {
fn new(call_id: String) -> Self {
Self {
call_id,
started: Instant::now(),
next_chunk: AtomicI64::new(0),
chunks: ChunkBuffer::default(),
first_text_recorded: AtomicBool::new(false),
}
}
} }
#[derive(Default)] #[derive(Default)]
@@ -70,12 +89,11 @@ impl CallRecorder {
Ok(Self { Ok(Self {
inner: Arc::new(Inner { inner: Arc::new(Inner {
store, store,
call_id: call.call_id, base_call: call.clone(),
started: Instant::now(),
detailed: call.detailed, detailed: call.detailed,
next_chunk: AtomicI64::new(0), attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
chunks: Mutex::new(ChunkBuffer::default()), next_attempt: AtomicU32::new(0),
first_text_recorded: AtomicBool::new(false), next_generation: AtomicU64::new(0),
finished: AtomicBool::new(false), finished: AtomicBool::new(false),
}), }),
}) })
@@ -94,55 +112,63 @@ impl CallRecorder {
headers: serde_json::Value, headers: serde_json::Value,
body: &serde_json::Value, body: &serde_json::Value,
) -> Result<()> { ) -> Result<()> {
let attempt = self.inner.attempt.lock().await;
self.inner self.inner
.store .store
.record_llm_request(&self.inner.call_id, &headers, body, self.inner.detailed) .record_llm_request(&attempt.call_id, &headers, body, self.inner.detailed)
.await?; .await?;
Ok(()) Ok(())
} }
pub async fn response_headers(&self, status: u16) -> Result<()> { pub async fn response_headers(&self, status: u16) -> Result<()> {
let attempt = self.inner.attempt.lock().await;
self.inner self.inner
.store .store
.record_llm_response_headers(&self.inner.call_id, self.elapsed_ms(), status) .record_llm_response_headers(&attempt.call_id, elapsed_ms(attempt.started), status)
.await .await
} }
pub async fn response_chunk(&self, data: &[u8]) -> Result<()> { pub async fn response_chunk(&self, data: &[u8]) -> Result<()> {
let mut buffer = self.inner.chunks.lock().await; let mut attempt = self.inner.attempt.lock().await;
if self.is_finished() { if self.is_finished() {
return Ok(()); return Ok(());
} }
let seq = self.inner.next_chunk.fetch_add(1, Ordering::Relaxed); let seq = attempt.next_chunk.fetch_add(1, Ordering::Relaxed);
let schedule_flush = if buffer.chunks.is_empty() { let schedule_flush = if attempt.chunks.chunks.is_empty() {
buffer.generation = buffer.generation.wrapping_add(1); attempt.chunks.generation = self
buffer.first_chunk_at = Some(Instant::now()); .inner
Some(buffer.generation) .next_generation
.fetch_add(1, Ordering::Relaxed)
.wrapping_add(1);
attempt.chunks.first_chunk_at = Some(Instant::now());
Some(attempt.chunks.generation)
} else { } else {
None None
}; };
buffer.bytes += data.len(); attempt.chunks.bytes += data.len();
buffer.chunks.push(if self.inner.detailed { let elapsed = elapsed_ms(attempt.started);
BufferedLlmChunk::new(seq, self.elapsed_ms(), data) attempt.chunks.chunks.push(if self.inner.detailed {
BufferedLlmChunk::new(seq, elapsed, data)
} else { } else {
BufferedLlmChunk::metrics(seq, self.elapsed_ms(), data.len()) BufferedLlmChunk::metrics(seq, elapsed, data.len())
}); });
let expired = buffer let expired = attempt
.chunks
.first_chunk_at .first_chunk_at
.is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE); .is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE);
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS if attempt.chunks.chunks.len() >= MAX_BUFFERED_CHUNKS
|| buffer.bytes >= MAX_BUFFERED_BYTES || attempt.chunks.bytes >= MAX_BUFFERED_BYTES
|| expired || expired
{ {
self.flush_locked(&mut buffer).await?; self.flush_locked(&mut attempt).await?;
} }
drop(buffer); drop(attempt);
if let Some(generation) = schedule_flush { if let Some(generation) = schedule_flush {
let recorder = self.clone(); let recorder = self.clone();
tokio::spawn(async move { tokio::spawn(async move {
tokio::time::sleep(MAX_BUFFER_AGE).await; tokio::time::sleep(MAX_BUFFER_AGE).await;
if let Err(error) = recorder.flush_generation(generation).await { if let Err(error) = recorder.flush_generation(generation).await {
tracing::warn!(call_id = recorder.inner.call_id, %error, "failed to flush LLM response chunks"); tracing::warn!(call_id = recorder.call_id(), %error, "failed to flush LLM response chunks");
} }
}); });
} }
@@ -152,8 +178,8 @@ impl CallRecorder {
pub async fn event(&self, event: &ModelEvent) -> Result<()> { pub async fn event(&self, event: &ModelEvent) -> Result<()> {
match event { match event {
ModelEvent::TextDelta(_) => { ModelEvent::TextDelta(_) => {
if self let attempt = self.inner.attempt.lock().await;
.inner if attempt
.first_text_recorded .first_text_recorded
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok() .is_ok()
@@ -161,12 +187,10 @@ impl CallRecorder {
if let Err(error) = self if let Err(error) = self
.inner .inner
.store .store
.record_llm_first_text(&self.inner.call_id, self.elapsed_ms()) .record_llm_first_text(&attempt.call_id, elapsed_ms(attempt.started))
.await .await
{ {
self.inner attempt.first_text_recorded.store(false, Ordering::Release);
.first_text_recorded
.store(false, Ordering::Release);
return Err(error); return Err(error);
} }
} }
@@ -179,9 +203,10 @@ impl CallRecorder {
} }
pub async fn usage(&self, usage: Usage) -> Result<()> { pub async fn usage(&self, usage: Usage) -> Result<()> {
let attempt = self.inner.attempt.lock().await;
self.inner self.inner
.store .store
.record_llm_usage(&self.inner.call_id, usage) .record_llm_usage(&attempt.call_id, usage)
.await .await
} }
@@ -204,6 +229,32 @@ impl CallRecorder {
self.finish("cancelled", None, None, None).await self.finish("cancelled", None, None, None).await
} }
pub async fn retry(
&self,
error: &crate::Error,
headers: serde_json::Value,
body: &serde_json::Value,
) -> Result<()> {
self.failed(error).await?;
let attempt_number = self.inner.next_attempt.fetch_add(1, Ordering::Relaxed) + 1;
let mut call = self.inner.base_call.clone();
call.call_id = format!("{}:retry-{attempt_number}", self.inner.base_call.call_id);
self.inner.store.start_llm_call(&call).await?;
{
let mut attempt = self.inner.attempt.lock().await;
*attempt = AttemptState::new(call.call_id);
self.inner.finished.store(false, Ordering::Release);
}
if let Err(error) = self.request(headers, body).await {
self.failed(&error).await?;
return Err(error);
}
Ok(())
}
async fn finish( async fn finish(
&self, &self,
status: &str, status: &str,
@@ -214,37 +265,34 @@ impl CallRecorder {
if self.inner.finished.swap(true, Ordering::AcqRel) { if self.inner.finished.swap(true, Ordering::AcqRel) {
return Ok(()); return Ok(());
} }
if let Err(error) = self.flush_chunks().await { let mut attempt = self.inner.attempt.lock().await;
if let Err(error) = self.flush_locked(&mut attempt).await {
self.inner.finished.store(false, Ordering::Release); self.inner.finished.store(false, Ordering::Release);
return Err(error); return Err(error);
} }
self.inner self.inner
.store .store
.finish_llm_call( .finish_llm_call(
&self.inner.call_id, &attempt.call_id,
status, status,
reason, reason,
self.elapsed_ms(), elapsed_ms(attempt.started),
error_kind, error_kind,
error_message, error_message,
) )
.await .await
} }
async fn flush_chunks(&self) -> Result<()> {
let mut buffer = self.inner.chunks.lock().await;
self.flush_locked(&mut buffer).await
}
async fn flush_generation(&self, generation: u64) -> Result<()> { async fn flush_generation(&self, generation: u64) -> Result<()> {
let mut buffer = self.inner.chunks.lock().await; let mut attempt = self.inner.attempt.lock().await;
if buffer.generation != generation { if attempt.chunks.generation != generation {
return Ok(()); return Ok(());
} }
self.flush_locked(&mut buffer).await self.flush_locked(&mut attempt).await
} }
async fn flush_locked(&self, buffer: &mut ChunkBuffer) -> Result<()> { async fn flush_locked(&self, attempt: &mut AttemptState) -> Result<()> {
let buffer = &mut attempt.chunks;
if buffer.chunks.is_empty() { if buffer.chunks.is_empty() {
return Ok(()); return Ok(());
} }
@@ -254,7 +302,7 @@ impl CallRecorder {
if let Err(error) = self if let Err(error) = self
.inner .inner
.store .store
.record_llm_chunks(&self.inner.call_id, &chunks, self.inner.detailed) .record_llm_chunks(&attempt.call_id, &chunks, self.inner.detailed)
.await .await
{ {
buffer.bytes = chunks.iter().map(|chunk| chunk.byte_count).sum(); buffer.bytes = chunks.iter().map(|chunk| chunk.byte_count).sum();
@@ -265,15 +313,19 @@ impl CallRecorder {
Ok(()) Ok(())
} }
fn elapsed_ms(&self) -> i64 { fn call_id(&self) -> String {
self.inner self.inner
.started .attempt
.elapsed() .try_lock()
.as_millis() .map(|attempt| attempt.call_id.clone())
.min(i64::MAX as u128) as i64 .unwrap_or_else(|_| self.inner.base_call.call_id.clone())
} }
} }
fn elapsed_ms(started: Instant) -> i64 {
started.elapsed().as_millis().min(i64::MAX as u128) as i64
}
fn finish_reason(reason: FinishReason) -> &'static str { fn finish_reason(reason: FinishReason) -> &'static str {
match reason { match reason {
FinishReason::Stop => "stop", FinishReason::Stop => "stop",
@@ -296,6 +348,34 @@ mod tests {
use super::*; use super::*;
async fn test_recorder(store: &Store, call_id: &str, detailed: bool) -> CallRecorder { async fn test_recorder(store: &Store, call_id: &str, detailed: bool) -> CallRecorder {
let call = NewLlmCall {
call_id: call_id.into(),
run_id: "run".into(),
conversation_id: "conversation".into(),
provider_call_index: 0,
model_hash: "hash".into(),
provider_type: crate::model::ProviderType::OpenAiChat,
provider_url: "https://example.com".into(),
request_type: crate::model::ProviderType::OpenAiChat,
request_url: "https://example.com".into(),
model_id: "model".into(),
display_name: "Model".into(),
reasoning_effort: None,
fast: false,
message_count: 0,
tool_count: 0,
detailed,
};
sqlx::query(
"INSERT INTO model_configs(
model_hash, display_name, model_type, base_url, api_key,
tooltip_data, model_id, created_at_ms, updated_at_ms
) VALUES ('hash', 'Model', 'openai', 'https://example.com',
'key', 'Model', 'model', 1, 1)",
)
.execute(store.pool())
.await
.unwrap();
sqlx::query( sqlx::query(
"INSERT INTO llm_calls( "INSERT INTO llm_calls(
call_id, run_id, conversation_id, provider_call_index, provider_type, call_id, run_id, conversation_id, provider_call_index, provider_type,
@@ -313,19 +393,18 @@ mod tests {
CallRecorder { CallRecorder {
inner: Arc::new(Inner { inner: Arc::new(Inner {
store: store.clone(), store: store.clone(),
call_id: call_id.into(), base_call: call.clone(),
started: Instant::now(),
detailed, detailed,
next_chunk: AtomicI64::new(0), attempt: Mutex::new(AttemptState::new(call_id.into())),
chunks: Mutex::new(ChunkBuffer::default()), next_attempt: AtomicU32::new(0),
first_text_recorded: AtomicBool::new(false), next_generation: AtomicU64::new(0),
finished: AtomicBool::new(false), finished: AtomicBool::new(false),
}), }),
} }
} }
#[tokio::test] #[tokio::test]
async fn a_partial_chunk_batch_flushes_after_the_deadline() { async fn a_partial_chunk_flushes_after_the_deadline() {
let store = Store::connect("sqlite::memory:").await.unwrap(); let store = Store::connect("sqlite::memory:").await.unwrap();
let recorder = test_recorder(&store, "timed-flush-call", true).await; let recorder = test_recorder(&store, "timed-flush-call", true).await;
@@ -390,4 +469,48 @@ mod tests {
.unwrap(); .unwrap();
assert_eq!(count, 1); assert_eq!(count, 1);
} }
#[tokio::test]
async fn retry_finishes_the_old_call_and_records_the_new_request() {
let store = Store::connect("sqlite::memory:").await.unwrap();
let recorder = test_recorder(&store, "retry-call", true).await;
let error = crate::Error::Provider("OpenAI Chat 429: rate limited".into());
let body = serde_json::json!({"model": "model", "stream": true});
recorder
.retry(
&error,
serde_json::json!({"content-type": "application/json"}),
&body,
)
.await
.unwrap();
let old = store.llm_call("retry-call").await.unwrap().unwrap();
assert_eq!(old.status, "error");
assert_eq!(
old.error_message.as_deref(),
Some(error.to_string().as_str())
);
let new = store.llm_call("retry-call:retry-1").await.unwrap().unwrap();
assert_eq!(new.status, "running");
assert_eq!(new.request_bytes, Some(31));
assert!(store
.llm_call_request("retry-call:retry-1")
.await
.unwrap()
.is_some());
recorder.completed(FinishReason::Stop).await.unwrap();
assert_eq!(
store
.llm_call("retry-call:retry-1")
.await
.unwrap()
.unwrap()
.status,
"completed"
);
}
} }
+28 -6
View File
@@ -33,6 +33,8 @@ pub(crate) async fn send_with_retry<F>(
policy: RetryPolicy, policy: RetryPolicy,
cancellation: &CancellationToken, cancellation: &CancellationToken,
recorder: Option<&CallRecorder>, recorder: Option<&CallRecorder>,
request_headers: serde_json::Value,
request_body: &serde_json::Value,
) -> Result<Attempt> ) -> Result<Attempt>
where where
F: Fn() -> reqwest::RequestBuilder, F: Fn() -> reqwest::RequestBuilder,
@@ -43,19 +45,24 @@ where
response = build().send() => response, response = build().send() => response,
}?; }?;
if let Some(recorder) = recorder { if let Some(recorder) = recorder {
recorder.response_headers(response.status().as_u16()).await?; recorder
.response_headers(response.status().as_u16())
.await?;
} }
if response.status().is_success() { if response.status().is_success() {
return Ok(Attempt::Response(response)); return Ok(Attempt::Response(response));
} }
let status = response.status(); let status = response.status();
let bytes = response.bytes().await?; let bytes = response.bytes().await?;
if let Some(recorder) = recorder { let error = Error::Provider(format!(
recorder.response_chunk(&bytes).await?; "{label} {status}: {}",
} String::from_utf8_lossy(&bytes)
));
if attempt == policy.retries { if attempt == policy.retries {
let text = String::from_utf8_lossy(&bytes); if let Some(recorder) = recorder {
return Err(Error::Provider(format!("{label} {status}: {text}"))); recorder.failed(&error).await?;
}
return Err(error);
} }
tracing::warn!( tracing::warn!(
provider = label, provider = label,
@@ -65,6 +72,11 @@ where
delay_ms = policy.delay.as_millis(), delay_ms = policy.delay.as_millis(),
"provider returned a non-success status, retrying" "provider returned a non-success status, retrying"
); );
if let Some(recorder) = recorder {
recorder
.retry(&error, request_headers.clone(), request_body)
.await?;
}
tokio::select! { tokio::select! {
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled), _ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
_ = tokio::time::sleep(policy.delay) => {} _ = tokio::time::sleep(policy.delay) => {}
@@ -138,6 +150,8 @@ mod tests {
fast(5), fast(5),
&CancellationToken::new(), &CancellationToken::new(),
None, None,
serde_json::json!({}),
&serde_json::json!({}),
) )
.await .await
.unwrap(); .unwrap();
@@ -159,6 +173,8 @@ mod tests {
fast(5), fast(5),
&CancellationToken::new(), &CancellationToken::new(),
None, None,
serde_json::json!({}),
&serde_json::json!({}),
) )
.await .await
.unwrap_err(); .unwrap_err();
@@ -184,6 +200,8 @@ mod tests {
}, },
&CancellationToken::new(), &CancellationToken::new(),
None, None,
serde_json::json!({}),
&serde_json::json!({}),
) )
.await .await
.unwrap(); .unwrap();
@@ -214,6 +232,8 @@ mod tests {
}, },
&cancellation, &cancellation,
None, None,
serde_json::json!({}),
&serde_json::json!({}),
) )
.await .await
.unwrap(); .unwrap();
@@ -232,6 +252,8 @@ mod tests {
fast(5), fast(5),
&CancellationToken::new(), &CancellationToken::new(),
None, None,
serde_json::json!({}),
&serde_json::json!({}),
) )
.await .await
.unwrap(); .unwrap();