mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +08:00
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:
@@ -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 };
|
||||||
|
|||||||
@@ -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 };
|
||||||
|
|||||||
@@ -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
@@ -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"
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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();
|
||||||
|
|||||||
Reference in New Issue
Block a user