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)?;
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 {
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(
"Anthropic",
@@ -94,6 +98,8 @@ impl Provider for AnthropicProvider {
RetryPolicy::default(),
&cancellation,
recorder.as_ref(),
request_headers,
&body,
).await?;
let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id };
+6 -2
View File
@@ -16,7 +16,8 @@ use crate::{
};
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},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -84,8 +85,9 @@ impl Provider for OpenAiChatProvider {
apply_model(&mut body, &request.model, config.max_output_tokens)?;
merge_extra_params(&mut body, &request.model.extra_params)?;
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 {
recorder.request(recorded_headers(&config, &[("content-type", "application/json")]), &body).await?;
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
"OpenAI Chat",
@@ -94,6 +96,8 @@ impl Provider for OpenAiChatProvider {
RetryPolicy::default(),
&cancellation,
recorder.as_ref(),
request_headers,
&body,
).await?;
let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id };
+6 -2
View File
@@ -13,7 +13,8 @@ use crate::{
};
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},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -81,8 +82,9 @@ impl Provider for OpenAiResponsesProvider {
apply_model(&mut body, &request.model, config.max_output_tokens)?;
merge_extra_params(&mut body, &request.model.extra_params)?;
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 {
recorder.request(recorded_headers(&config, &[("content-type", "application/json")]), &body).await?;
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
"OpenAI Responses",
@@ -91,6 +93,8 @@ impl Provider for OpenAiResponsesProvider {
RetryPolicy::default(),
&cancellation,
recorder.as_ref(),
request_headers,
&body,
).await?;
let Attempt::Response(response) = attempt else { return };
yield ModelEvent::Start { model_call_id: call_id };
+181 -58
View File
@@ -1,6 +1,6 @@
use std::{
sync::{
atomic::{AtomicBool, AtomicI64, Ordering},
atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering},
Arc,
},
time::Instant,
@@ -42,13 +42,32 @@ pub struct CallRecorder {
struct Inner {
store: Store,
base_call: NewLlmCall,
detailed: bool,
attempt: Mutex<AttemptState>,
next_attempt: AtomicU32,
next_generation: AtomicU64,
finished: AtomicBool,
}
struct AttemptState {
call_id: String,
started: Instant,
detailed: bool,
next_chunk: AtomicI64,
chunks: Mutex<ChunkBuffer>,
chunks: ChunkBuffer,
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)]
@@ -70,12 +89,11 @@ impl CallRecorder {
Ok(Self {
inner: Arc::new(Inner {
store,
call_id: call.call_id,
started: Instant::now(),
base_call: call.clone(),
detailed: call.detailed,
next_chunk: AtomicI64::new(0),
chunks: Mutex::new(ChunkBuffer::default()),
first_text_recorded: AtomicBool::new(false),
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
next_attempt: AtomicU32::new(0),
next_generation: AtomicU64::new(0),
finished: AtomicBool::new(false),
}),
})
@@ -94,55 +112,63 @@ impl CallRecorder {
headers: serde_json::Value,
body: &serde_json::Value,
) -> Result<()> {
let attempt = self.inner.attempt.lock().await;
self.inner
.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?;
Ok(())
}
pub async fn response_headers(&self, status: u16) -> Result<()> {
let attempt = self.inner.attempt.lock().await;
self.inner
.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
}
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() {
return Ok(());
}
let seq = self.inner.next_chunk.fetch_add(1, Ordering::Relaxed);
let schedule_flush = if buffer.chunks.is_empty() {
buffer.generation = buffer.generation.wrapping_add(1);
buffer.first_chunk_at = Some(Instant::now());
Some(buffer.generation)
let seq = attempt.next_chunk.fetch_add(1, Ordering::Relaxed);
let schedule_flush = if attempt.chunks.chunks.is_empty() {
attempt.chunks.generation = self
.inner
.next_generation
.fetch_add(1, Ordering::Relaxed)
.wrapping_add(1);
attempt.chunks.first_chunk_at = Some(Instant::now());
Some(attempt.chunks.generation)
} else {
None
};
buffer.bytes += data.len();
buffer.chunks.push(if self.inner.detailed {
BufferedLlmChunk::new(seq, self.elapsed_ms(), data)
attempt.chunks.bytes += data.len();
let elapsed = elapsed_ms(attempt.started);
attempt.chunks.chunks.push(if self.inner.detailed {
BufferedLlmChunk::new(seq, elapsed, data)
} 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
.is_some_and(|started| started.elapsed() >= MAX_BUFFER_AGE);
if buffer.chunks.len() >= MAX_BUFFERED_CHUNKS
|| buffer.bytes >= MAX_BUFFERED_BYTES
if attempt.chunks.chunks.len() >= MAX_BUFFERED_CHUNKS
|| attempt.chunks.bytes >= MAX_BUFFERED_BYTES
|| expired
{
self.flush_locked(&mut buffer).await?;
self.flush_locked(&mut attempt).await?;
}
drop(buffer);
drop(attempt);
if let Some(generation) = schedule_flush {
let recorder = self.clone();
tokio::spawn(async move {
tokio::time::sleep(MAX_BUFFER_AGE).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<()> {
match event {
ModelEvent::TextDelta(_) => {
if self
.inner
let attempt = self.inner.attempt.lock().await;
if attempt
.first_text_recorded
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
@@ -161,12 +187,10 @@ impl CallRecorder {
if let Err(error) = self
.inner
.store
.record_llm_first_text(&self.inner.call_id, self.elapsed_ms())
.record_llm_first_text(&attempt.call_id, elapsed_ms(attempt.started))
.await
{
self.inner
.first_text_recorded
.store(false, Ordering::Release);
attempt.first_text_recorded.store(false, Ordering::Release);
return Err(error);
}
}
@@ -179,9 +203,10 @@ impl CallRecorder {
}
pub async fn usage(&self, usage: Usage) -> Result<()> {
let attempt = self.inner.attempt.lock().await;
self.inner
.store
.record_llm_usage(&self.inner.call_id, usage)
.record_llm_usage(&attempt.call_id, usage)
.await
}
@@ -204,6 +229,32 @@ impl CallRecorder {
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(
&self,
status: &str,
@@ -214,37 +265,34 @@ impl CallRecorder {
if self.inner.finished.swap(true, Ordering::AcqRel) {
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);
return Err(error);
}
self.inner
.store
.finish_llm_call(
&self.inner.call_id,
&attempt.call_id,
status,
reason,
self.elapsed_ms(),
elapsed_ms(attempt.started),
error_kind,
error_message,
)
.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<()> {
let mut buffer = self.inner.chunks.lock().await;
if buffer.generation != generation {
let mut attempt = self.inner.attempt.lock().await;
if attempt.chunks.generation != generation {
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() {
return Ok(());
}
@@ -254,7 +302,7 @@ impl CallRecorder {
if let Err(error) = self
.inner
.store
.record_llm_chunks(&self.inner.call_id, &chunks, self.inner.detailed)
.record_llm_chunks(&attempt.call_id, &chunks, self.inner.detailed)
.await
{
buffer.bytes = chunks.iter().map(|chunk| chunk.byte_count).sum();
@@ -265,15 +313,19 @@ impl CallRecorder {
Ok(())
}
fn elapsed_ms(&self) -> i64 {
fn call_id(&self) -> String {
self.inner
.started
.elapsed()
.as_millis()
.min(i64::MAX as u128) as i64
.attempt
.try_lock()
.map(|attempt| attempt.call_id.clone())
.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 {
match reason {
FinishReason::Stop => "stop",
@@ -296,6 +348,34 @@ mod tests {
use super::*;
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(
"INSERT INTO llm_calls(
call_id, run_id, conversation_id, provider_call_index, provider_type,
@@ -313,19 +393,18 @@ mod tests {
CallRecorder {
inner: Arc::new(Inner {
store: store.clone(),
call_id: call_id.into(),
started: Instant::now(),
base_call: call.clone(),
detailed,
next_chunk: AtomicI64::new(0),
chunks: Mutex::new(ChunkBuffer::default()),
first_text_recorded: AtomicBool::new(false),
attempt: Mutex::new(AttemptState::new(call_id.into())),
next_attempt: AtomicU32::new(0),
next_generation: AtomicU64::new(0),
finished: AtomicBool::new(false),
}),
}
}
#[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 recorder = test_recorder(&store, "timed-flush-call", true).await;
@@ -390,4 +469,48 @@ mod tests {
.unwrap();
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,
cancellation: &CancellationToken,
recorder: Option<&CallRecorder>,
request_headers: serde_json::Value,
request_body: &serde_json::Value,
) -> Result<Attempt>
where
F: Fn() -> reqwest::RequestBuilder,
@@ -43,19 +45,24 @@ where
response = build().send() => response,
}?;
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() {
return Ok(Attempt::Response(response));
}
let status = response.status();
let bytes = response.bytes().await?;
if let Some(recorder) = recorder {
recorder.response_chunk(&bytes).await?;
}
let error = Error::Provider(format!(
"{label} {status}: {}",
String::from_utf8_lossy(&bytes)
));
if attempt == policy.retries {
let text = String::from_utf8_lossy(&bytes);
return Err(Error::Provider(format!("{label} {status}: {text}")));
if let Some(recorder) = recorder {
recorder.failed(&error).await?;
}
return Err(error);
}
tracing::warn!(
provider = label,
@@ -65,6 +72,11 @@ where
delay_ms = policy.delay.as_millis(),
"provider returned a non-success status, retrying"
);
if let Some(recorder) = recorder {
recorder
.retry(&error, request_headers.clone(), request_body)
.await?;
}
tokio::select! {
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
_ = tokio::time::sleep(policy.delay) => {}
@@ -138,6 +150,8 @@ mod tests {
fast(5),
&CancellationToken::new(),
None,
serde_json::json!({}),
&serde_json::json!({}),
)
.await
.unwrap();
@@ -159,6 +173,8 @@ mod tests {
fast(5),
&CancellationToken::new(),
None,
serde_json::json!({}),
&serde_json::json!({}),
)
.await
.unwrap_err();
@@ -184,6 +200,8 @@ mod tests {
},
&CancellationToken::new(),
None,
serde_json::json!({}),
&serde_json::json!({}),
)
.await
.unwrap();
@@ -214,6 +232,8 @@ mod tests {
},
&cancellation,
None,
serde_json::json!({}),
&serde_json::json!({}),
)
.await
.unwrap();
@@ -232,6 +252,8 @@ mod tests {
fast(5),
&CancellationToken::new(),
None,
serde_json::json!({}),
&serde_json::json!({}),
)
.await
.unwrap();