refactor: remove retry_count from ProviderConfig and enhance error handling in tool execution

- Removed the `retry_count` field from `ProviderConfig` as it is no longer needed.
- Introduced `argument_error` field in `ToolCall` to capture errors related to tool arguments.
- Updated various components to handle argument errors more gracefully, including in the `ToolDispatcher` and `ConversationOutput`.
- Enhanced tests to validate the new error handling and ensure proper functionality of tool calls.
This commit is contained in:
leokun
2026-09-01 10:14:53 +08:00
parent 29fde7d7c7
commit d83e14af9a
38 changed files with 1105 additions and 444 deletions
@@ -0,0 +1 @@
ALTER TABLE tool_round_calls ADD COLUMN argument_error TEXT;
-1
View File
@@ -46,7 +46,6 @@ pub struct ProviderConfig {
pub custom_headers: reqwest::header::HeaderMap,
pub max_output_tokens: Option<u64>,
pub request_timeout: Duration,
pub retry_count: u32,
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
}
@@ -102,6 +102,12 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
.into_iter()
.enumerate()
.map(|(index, call)| {
let argument_error = wire
.pointer("/providerOptions/cursor/pendingToolExecutionContracts")
.and_then(|contracts| contracts.get(&call.call_id))
.and_then(|contract| contract.get("argumentError"))
.and_then(Value::as_str)
.map(str::to_string);
Ok(ToolCall {
index,
call_id: call.call_id,
@@ -109,6 +115,7 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
name: call.name,
arguments_text: serde_json::to_string(&call.arguments)?,
arguments: call.arguments,
argument_error,
})
})
.collect::<Result<Vec<_>>>()?;
@@ -71,6 +71,7 @@ pub fn staged_tool_round(
allowed_tools,
dynamic_tools,
started_at_ms,
tool_calls: Some(calls),
}),
)?)?)
}
@@ -96,6 +97,7 @@ pub fn staged_final(
allowed_tools,
dynamic_tools,
started_at_ms,
tool_calls: None,
}),
)?)?)
}
@@ -105,6 +107,7 @@ pub(super) struct PendingContext<'a> {
allowed_tools: &'a [String],
dynamic_tools: &'a HashSet<String>,
started_at_ms: u64,
tool_calls: Option<&'a [ToolCall]>,
}
pub(super) fn wire_message(
@@ -132,16 +135,23 @@ pub(super) fn wire_message(
calls
.iter()
.map(|call| {
(
call.call_id.clone(),
json!({
let mut contract = json!({
"toolCallId": call.call_id,
"outerToolName": call.name,
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
"isDynamic": pending.dynamic_tools.contains(&call.name),
"allowedToolNames": pending.allowed_tools,
}),
)
});
if let Some(error) = pending
.tool_calls
.and_then(|calls| {
calls.iter().find(|candidate| candidate.call_id == call.call_id)
})
.and_then(|call| call.argument_error.as_deref())
{
contract["argumentError"] = Value::String(error.into());
}
(call.call_id.clone(), contract)
})
.collect(),
),
@@ -30,6 +30,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
name: "Read".into(),
arguments_text: r#"{"path":"/a"}"#.into(),
arguments: json!({"path":"/a"}),
argument_error: Some("Read arguments are not valid JSON".into()),
},
ToolCall {
index: 1,
@@ -38,6 +39,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
name: "Grep".into(),
arguments_text: r#"{"pattern":"x"}"#.into(),
arguments: json!({"pattern":"x"}),
argument_error: None,
},
];
let pending = staged_tool_round(
@@ -55,6 +57,10 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["toolIdentifier"],
"READ"
);
assert_eq!(
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["argumentError"],
"Read arguments are not valid JSON"
);
assert_eq!(wire["role"], "assistant");
assert_eq!(
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]
@@ -77,6 +83,10 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
assert_eq!(recovered.assistant.replay_state, Some(replay_state));
assert_eq!(recovered.calls.len(), 2);
assert_eq!(recovered.calls[0].call_id, "a");
assert_eq!(
recovered.calls[0].argument_error.as_deref(),
Some("Read arguments are not valid JSON")
);
assert_eq!(recovered.calls[1].call_id, "b");
}
+27
View File
@@ -75,6 +75,11 @@ impl StepBuffer {
});
}
pub fn finish_model_attempt(&mut self) {
self.finish_text();
self.finish_thinking(Duration::ZERO);
}
pub fn discard_model_output(&mut self) {
self.text.clear();
self.thinking.clear();
@@ -101,6 +106,28 @@ impl StepBuffer {
mod tests {
use super::*;
#[test]
fn failed_attempt_output_is_retained_for_the_next_checkpoint() {
let mut buffer = StepBuffer::default();
buffer.text_delta("partial answer");
buffer.thinking_delta("partial reasoning");
buffer.finish_model_attempt();
let steps = buffer.take().steps;
assert_eq!(steps.len(), 2);
assert!(matches!(
&steps[0].message,
Some(pb::conversation_step::Message::AssistantMessage(message))
if message.text == "partial answer"
));
assert!(matches!(
&steps[1].message,
Some(pb::conversation_step::Message::ThinkingMessage(message))
if message.text == "partial reasoning"
));
}
#[test]
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
let mut buffer = StepBuffer::default();
+78 -9
View File
@@ -21,7 +21,7 @@ use crate::{
protocol::proto::agent::v1 as pb,
services::blob_sync::BlobSynchronizer,
tools::{
codec,
codec, compat,
runtime::CursorToolRuntime,
stream::ToolCallStream,
tool_call_result::{ToolCompletion, ToolResultReceiver},
@@ -264,6 +264,32 @@ impl ConversationOutput {
streams.clear();
presentation.discard_model_output();
}
RunEvent::ModelAttemptFailed { attempt, message } => {
tracing::warn!(
run_id = %self.run.run_id(),
attempt,
%message,
"retrying model call from current checkpoint"
);
presentation.finish_model_attempt();
for call in calls.values_mut() {
if call.arguments.is_null() {
call.arguments = serde_json::from_str(&call.arguments_text)
.unwrap_or_else(|_| serde_json::json!({}));
}
let completion = compat::failure_with_message(
call,
format!("Model attempt failed before tool completion: {message}"),
);
self.handle
.emit(&codec::tool_completed(call, &completion))?;
presentation.tool_completed(&completion);
}
response_text.clear();
response_thinking.clear();
calls.clear();
streams.clear();
}
RunEvent::TextStart => {}
RunEvent::TextEnd => {
if !self.context.compacting {
@@ -312,6 +338,7 @@ impl ConversationOutput {
name: name.clone(),
arguments_text: String::new(),
arguments: serde_json::Value::Null,
argument_error: None,
};
self.emit_model_event(
crate::provider::ModelEvent::ToolCallStart {
@@ -335,10 +362,29 @@ impl ConversationOutput {
let stream = streams.get_mut(&index).ok_or_else(|| {
Error::Protocol(format!("missing Cursor tool stream: {index}"))
})?;
for message in stream.arguments_delta(call, &delta)? {
match stream.arguments_delta(call, &delta) {
Ok(messages) => {
for message in messages {
self.handle.emit(&message)?;
}
}
Err(Error::Protocol(message)) => {
tracing::warn!(
call_id = %call.call_id,
%message,
"ignoring invalid streaming tool arguments until completion"
);
}
Err(Error::Json(error)) => {
tracing::warn!(
call_id = %call.call_id,
%error,
"ignoring invalid streaming tool arguments until completion"
);
}
Err(error) => return Err(error),
}
}
RunEvent::ToolCallEnd { index } => {
let call = calls.get_mut(&index).ok_or_else(|| {
Error::Protocol(format!("unknown completed tool index: {index}"))
@@ -349,7 +395,8 @@ impl ConversationOutput {
call.arguments = if call.arguments_text.trim().is_empty() {
serde_json::json!({})
} else {
serde_json::from_str(&call.arguments_text)?
serde_json::from_str(&call.arguments_text)
.unwrap_or_else(|_| serde_json::json!({}))
};
}
RunEvent::Usage(usage) => {
@@ -359,10 +406,7 @@ impl ConversationOutput {
}
}
if !self.context.compacting {
context_tokens = usage
.input_tokens
.zip(usage.output_tokens)
.and_then(|(input, output)| input.checked_add(output));
context_tokens = usage.context_input_tokens;
}
match &mut turn_usage {
Some(total) => *total += usage,
@@ -523,6 +567,7 @@ impl ConversationOutput {
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
{
Ok(checkpoint) => {
context_tokens = checkpoint_context_tokens(&checkpoint);
compaction_checkpoint = Some(checkpoint);
state.barrier.complete(Ok(()));
}
@@ -1029,10 +1074,34 @@ pub(crate) fn finish_cancelled(handle: &TransportHandle) -> Result<()> {
Ok(())
}
fn checkpoint_context_tokens(checkpoint: &pb::ConversationStateStructure) -> Option<u64> {
checkpoint
.token_details
.as_ref()
.map(|details| u64::from(details.used_tokens))
}
#[cfg(test)]
mod tests {
use super::accept_tool_completion;
use crate::{run::CommandResult, Error};
use super::{accept_tool_completion, checkpoint_context_tokens};
use crate::{cursor::protocol::proto::agent::v1 as pb, run::CommandResult, Error};
#[test]
fn compacted_checkpoint_replaces_the_in_memory_context_usage() {
let compacted = pb::ConversationStateStructure {
token_details: Some(pb::ConversationTokenDetails {
used_tokens: 20_000,
..Default::default()
}),
..Default::default()
};
assert_eq!(checkpoint_context_tokens(&compacted), Some(20_000));
assert_eq!(
checkpoint_context_tokens(&pb::ConversationStateStructure::default()),
None
);
}
#[test]
fn closing_and_ended_runs_ignore_known_tool_completions() {
+10 -9
View File
@@ -11,7 +11,7 @@ use crate::{
protocol::proto::agent::v1 as pb,
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
tools::{
codec,
codec, compat,
runtime::CursorToolRuntime,
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
ClientToolEvent, ToolDispatcher,
@@ -262,17 +262,18 @@ impl ConversationRuntime {
.take_exec(throw.id)
.await
{
Some(pending) => generation.results.send_error(
crate::Error::Protocol(format!(
Some(pending) => generation.results.send(
compat::failure_with_message(
&pending.call,
format!(
"Exec {} failed: {}",
pending.call.call_id, throw.error
)),
),
None => generation.results.send_error(
crate::Error::Protocol(format!(
"unknown ExecClientThrow id: {}",
throw.id
)),
),
),
None => tracing::warn!(
id = throw.id,
"ignoring failure for unknown tool execution"
),
}
}
+1
View File
@@ -542,6 +542,7 @@ mod tests {
name: "Bash".into(),
arguments_text: String::new(),
arguments: json!({ "command": "ls -la" }),
argument_error: None,
};
let message = request(1, &call, &ExecContext::default()).unwrap();
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message
+35 -18
View File
@@ -3,7 +3,7 @@ use crate::{
cursor::{
protocol::{events, proto::agent::v1 as pb},
tools::{
edit,
compat, edit,
runtime::{CursorToolRuntime, ExecStage, PendingExec},
tool_call_result::{self as result, ToolCompletion},
},
@@ -34,16 +34,15 @@ pub async fn client_event(
let call = match pending.exec_call(message.id).await {
Some(call) => call,
None if pending.completed_call(message.id).await.is_some() => {
return Err(Error::Protocol(format!(
"duplicate terminal ExecClientMessage id: {}",
message.id
)))
tracing::warn!(id = message.id, "ignoring duplicate terminal tool response");
return Ok(ClientExecEvent::Pending);
}
None => {
return Err(Error::Protocol(format!(
"unknown ExecClientMessage id: {}",
message.id
)))
tracing::warn!(
id = message.id,
"ignoring response for unknown tool execution"
);
return Ok(ClientExecEvent::Pending);
}
};
let Some(wire_result) = &message.message else {
@@ -174,9 +173,22 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Optio
}
let rendered = match &entry.stage {
ExecStage::DynamicMcp(definition) => {
super::render_dynamic_mcp(&entry.call, definition, false)
Ok(super::render_dynamic_mcp(&entry.call, definition, false))
}
_ => super::render_tool_call(&entry.call, false)?,
_ => super::render_tool_call(&entry.call, false),
};
let rendered = match rendered {
Ok(rendered) => rendered,
Err(Error::Protocol(message)) => {
return Ok(Some(compat::failure_with_message(&entry.call, message)));
}
Err(Error::Json(error)) => {
return Ok(Some(compat::failure_with_message(
&entry.call,
error.to_string(),
)));
}
Err(error) => return Err(error),
};
Ok(Some(ToolCompletion::from_rendered(
&entry.call,
@@ -212,10 +224,10 @@ async fn advance_edit(
pb::exec_client_message::Message::ReadResult(result)
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
_ => {
return Err(Error::Protocol(format!(
"expected ReadResult for edit tool {}",
entry.call.name
)))
let message = format!("expected ReadResult for edit tool {}", entry.call.name);
return Ok(ClientExecEvent::Completed(Box::new(
compat::failure_with_message(&entry.call, message),
)));
}
};
let write = match edit::after_read(&entry.call, read) {
@@ -260,9 +272,14 @@ fn completed(
pending: PendingExec,
result: pb::exec_client_message::Message,
) -> Result<ClientExecEvent> {
Ok(ClientExecEvent::Completed(Box::new(result::from_exec(
pending, &result,
)?)))
let call = pending.call.clone();
let completion = match result::from_exec(pending, &result) {
Ok(completion) => completion,
Err(Error::Protocol(message)) => compat::failure_with_message(&call, message),
Err(Error::Json(error)) => compat::failure_with_message(&call, error.to_string()),
Err(error) => return Err(error),
};
Ok(ClientExecEvent::Completed(Box::new(completion)))
}
fn shell_exit_result(
+4 -1
View File
@@ -49,7 +49,10 @@ pub(crate) fn render(call: &ToolCall, completed: bool) -> pb::ToolCall {
}
pub(crate) fn failure(call: &ToolCall) -> ToolCompletion {
let error = failure_message(&call.name);
failure_with_message(call, failure_message(&call.name))
}
pub(crate) fn failure_with_message(call: &ToolCall, error: String) -> ToolCompletion {
let arguments = call
.arguments
.as_object()
+1
View File
@@ -265,6 +265,7 @@ mod tests {
"old_string": old_string,
"new_string": "replacement",
}),
argument_error: None,
}
}
+74 -24
View File
@@ -103,10 +103,20 @@ impl ToolDispatcher {
}
let message_index = first_tool_index + position;
let publish_started = !state.started.contains(&call.call_id);
if let Some(error) = &call.argument_error {
dispatched.push(validation_failure(call, error.clone()));
continue;
}
let edit_path = if dynamic_mcp.contains_key(&call.name) {
None
} else {
edit::execution_path(call)?
match edit::execution_path(call) {
Ok(path) => path,
Err(error) => {
dispatched.push(recover_validation_failure(call, error)?);
continue;
}
}
};
if let Some(path) = edit_path {
let next = self.edit_schedule.lock().await.start_or_defer(
@@ -121,22 +131,28 @@ impl ToolDispatcher {
let Some(next) = next else {
continue;
};
dispatched.push(
self.start(
let started = self
.start(
&next.call,
next.message_index,
next.publish_started,
dynamic_mcp,
&next.context,
)
.await?,
);
.await;
dispatched.push(match started {
Ok(started) => started,
Err(error) => recover_validation_failure(&next.call, error)?,
});
continue;
}
dispatched.push(
self.start(call, message_index, publish_started, dynamic_mcp, context)
.await?,
);
let started = self
.start(call, message_index, publish_started, dynamic_mcp, context)
.await;
dispatched.push(match started {
Ok(started) => started,
Err(error) => recover_validation_failure(call, error)?,
});
}
Ok(dispatched)
}
@@ -146,7 +162,8 @@ impl ToolDispatcher {
let Some(next) = next else {
return Ok(None);
};
self.start(
match self
.start(
&next.call,
next.message_index,
next.publish_started,
@@ -154,7 +171,10 @@ impl ToolDispatcher {
&next.context,
)
.await
.map(Some)
{
Ok(started) => Ok(Some(started)),
Err(error) => recover_validation_failure(&next.call, error).map(Some),
}
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
@@ -203,34 +223,64 @@ impl ToolDispatcher {
let pending = match self.runtime.take_interaction(response.id).await {
Some(pending) => pending,
None if self.runtime.completed_call(response.id).await.is_some() => {
return Err(Error::Protocol(format!(
"duplicate terminal InteractionResponse id: {}",
response.id
)));
tracing::warn!(
id = response.id,
"ignoring duplicate terminal interaction response"
);
return Ok(ClientToolEvent::Pending);
}
None => {
return Err(Error::Protocol(format!(
"unknown InteractionResponse id: {}",
response.id
)));
tracing::warn!(
id = response.id,
"ignoring response for unknown interaction"
);
return Ok(ClientToolEvent::Pending);
}
};
Ok(
match tool_call_dispatch::resume_interaction(
let call = pending.call.clone();
let continuation = match tool_call_dispatch::resume_interaction(
&self.results,
&self.search,
&self.fetch,
pending,
response,
)
.await?
.await
{
Ok(continuation) => continuation,
Err(Error::Protocol(message)) => {
return Ok(ClientToolEvent::Completed(Box::new(
compat::failure_with_message(&call, message),
)));
}
Err(Error::Json(error)) => {
return Ok(ClientToolEvent::Completed(Box::new(
compat::failure_with_message(&call, error.to_string()),
)));
}
Err(error) => return Err(error),
};
Ok(match continuation {
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
ClientToolEvent::Completed(completion)
}
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
},
)
})
}
}
fn validation_failure(call: &ToolCall, message: String) -> DispatchedTool {
DispatchedTool {
messages: Vec::new(),
completion: Some(compat::failure_with_message(call, message)),
}
}
fn recover_validation_failure(call: &ToolCall, error: Error) -> Result<DispatchedTool> {
match error {
Error::Protocol(message) => Ok(validation_failure(call, message)),
Error::Json(error) => Ok(validation_failure(call, error.to_string())),
error => Err(error),
}
}
+1
View File
@@ -248,6 +248,7 @@ mod tests {
name: "Task".into(),
arguments_text: arguments_text.into(),
arguments: serde_json::Value::Null,
argument_error: None,
}
}
@@ -406,6 +406,7 @@ mod tests {
name: "WebFetch".into(),
arguments_text: r#"{"url":"https://example.com"}"#.into(),
arguments: json!({"url": "https://example.com"}),
argument_error: None,
},
started_at_ms: 1,
}
+2
View File
@@ -9,6 +9,7 @@ mod usage {
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
pub struct Usage {
pub input_tokens: Option<u64>,
pub context_input_tokens: Option<u64>,
pub output_tokens: Option<u64>,
pub total_tokens: Option<u64>,
pub cache_read_tokens: Option<u64>,
@@ -19,6 +20,7 @@ mod usage {
impl AddAssign for Usage {
fn add_assign(&mut self, rhs: Self) {
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
self.context_input_tokens = sum(self.context_input_tokens, rhs.context_input_tokens);
self.output_tokens = sum(self.output_tokens, rhs.output_tokens);
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
+1
View File
@@ -19,6 +19,7 @@ pub struct ToolCall {
pub name: String,
pub arguments_text: String,
pub arguments: Value,
pub argument_error: Option<String>,
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
+4 -1
View File
@@ -147,8 +147,10 @@ pub fn model_event(value: &serde_json::Value) -> Result<ModelEvent> {
.get("usage")
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
let tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
let input_tokens = tokens("inputTokens");
ModelEvent::Usage(Usage {
input_tokens: tokens("inputTokens"),
input_tokens,
context_input_tokens: input_tokens,
output_tokens: tokens("outputTokens"),
total_tokens: tokens("totalTokens"),
cache_read_tokens: tokens("cacheReadTokens"),
@@ -285,6 +287,7 @@ mod tests {
usage,
ModelEvent::Usage(Usage {
input_tokens: Some(10),
context_input_tokens: Some(10),
output_tokens: Some(2),
total_tokens: None,
cache_read_tokens: Some(4),
+36 -10
View File
@@ -12,9 +12,9 @@ use crate::{
};
use super::{
attempt::{send_once, Attempt},
map_sse_error, merge_extra_params, provider_event_error,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -90,17 +90,14 @@ impl Provider for AnthropicProvider {
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
let attempt = send_once(
"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 { retries: config.retry_count, ..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 };
@@ -321,6 +318,7 @@ fn apply_model(body: &mut Value, model: &crate::model::ModelSpec) -> Result<()>
fn merge_usage(total: &mut Usage, update: Usage) {
merge_usage_field(&mut total.input_tokens, update.input_tokens);
merge_usage_field(&mut total.context_input_tokens, update.context_input_tokens);
merge_usage_field(&mut total.output_tokens, update.output_tokens);
merge_usage_field(&mut total.cache_read_tokens, update.cache_read_tokens);
merge_usage_field(&mut total.cache_write_tokens, update.cache_write_tokens);
@@ -475,14 +473,42 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
}
fn anthropic_usage(value: &Value) -> Usage {
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
let cache_read_tokens = value.get("cache_read_input_tokens").and_then(Value::as_u64);
let cache_write_tokens = value
.get("cache_creation_input_tokens")
.and_then(Value::as_u64);
let context_input_tokens = input_tokens.map(|input| {
input
.saturating_add(cache_read_tokens.unwrap_or_default())
.saturating_add(cache_write_tokens.unwrap_or_default())
});
Usage {
input_tokens: value.get("input_tokens").and_then(Value::as_u64),
input_tokens,
context_input_tokens,
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
cache_read_tokens: value.get("cache_read_input_tokens").and_then(Value::as_u64),
cache_write_tokens: value
.get("cache_creation_input_tokens")
.and_then(Value::as_u64),
cache_read_tokens,
cache_write_tokens,
reasoning_tokens: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cached_tokens_are_included_once_in_anthropic_context_input() {
let usage = anthropic_usage(&serde_json::json!({
"input_tokens": 10,
"output_tokens": 5,
"cache_read_input_tokens": 20,
"cache_creation_input_tokens": 30
}));
assert_eq!(usage.input_tokens, Some(10));
assert_eq!(usage.context_input_tokens, Some(60));
assert_eq!(usage.output_tokens, Some(5));
}
}
+100
View File
@@ -0,0 +1,100 @@
//! Sends one Provider HTTP attempt without applying retry policy.
use tokio_util::sync::CancellationToken;
use crate::{Error, Result};
use super::CallRecorder;
#[derive(Debug)]
pub(crate) enum Attempt {
Response(reqwest::Response),
Cancelled,
}
pub(crate) async fn send_once<F>(
label: &str,
build: F,
cancellation: &CancellationToken,
recorder: Option<&CallRecorder>,
) -> Result<Attempt>
where
F: FnOnce() -> reqwest::RequestBuilder,
{
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 = tokio::select! {
_ = cancellation.cancelled() => return Ok(Attempt::Cancelled),
bytes = response.bytes() => bytes,
}?;
Err(Error::Provider(format!(
"{label} {status}: {}",
String::from_utf8_lossy(&bytes)
)))
}
#[cfg(test)]
mod tests {
use super::*;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
async fn server(response: &'static [u8]) -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 1024];
let _ = socket.read(&mut request).await;
socket.write_all(response).await.unwrap();
});
format!("http://{address}")
}
#[tokio::test]
async fn non_success_status_is_one_failed_attempt() {
let url =
server(b"HTTP/1.1 503 Service Unavailable\r\nContent-Length: 4\r\n\r\ndown").await;
let client = reqwest::Client::new();
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
.await
.unwrap_err();
assert!(
matches!(error, Error::Provider(message) if message.contains("503") && message.contains("down"))
);
}
#[tokio::test]
async fn response_body_transport_failure_is_one_failed_attempt() {
let url =
server(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 100\r\n\r\nshort").await;
let client = reqwest::Client::new();
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
.await
.unwrap_err();
assert!(matches!(error, Error::Http(_)));
}
#[tokio::test]
async fn request_transport_failure_is_one_failed_attempt() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
drop(listener);
let url = format!("http://{address}");
let client = reqwest::Client::new();
let error = send_once("test", || client.get(&url), &CancellationToken::new(), None)
.await
.unwrap_err();
assert!(matches!(error, Error::Http(_)));
}
}
+1 -1
View File
@@ -1,11 +1,11 @@
//! Defines the provider interface and exports provider implementations.
mod anthropic;
mod attempt;
mod event;
mod normalize;
mod openai_chat;
mod openai_responses;
mod recorder;
mod retry;
mod router;
use std::pin::Pin;
+7 -8
View File
@@ -17,10 +17,10 @@ use crate::{
};
use super::{
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
provider_event_error,
apply_body_allowlist, apply_openai_prompt_cache_key,
attempt::{send_once, Attempt},
map_sse_error, merge_extra_params, provider_event_error,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -92,15 +92,12 @@ impl Provider for OpenAiChatProvider {
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
let attempt = send_once(
"OpenAI Chat",
|| client.post(&config.request_url)
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
RetryPolicy { retries: config.retry_count, ..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 };
@@ -421,8 +418,10 @@ fn merge_chat_fragment(target: &mut String, fragment: &str) {
}
pub(crate) fn openai_usage(value: &Value) -> Usage {
let input_tokens = value.get("prompt_tokens").and_then(Value::as_u64);
Usage {
input_tokens: value.get("prompt_tokens").and_then(Value::as_u64),
input_tokens,
context_input_tokens: input_tokens,
output_tokens: value.get("completion_tokens").and_then(Value::as_u64),
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
cache_read_tokens: value
+7 -8
View File
@@ -14,10 +14,10 @@ use crate::{
};
use super::{
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
provider_event_error,
apply_body_allowlist, apply_openai_prompt_cache_key,
attempt::{send_once, Attempt},
map_sse_error, merge_extra_params, provider_event_error,
recorder::recorded_headers,
retry::{send_with_retry, Attempt, RetryPolicy},
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
};
@@ -89,15 +89,12 @@ impl Provider for OpenAiResponsesProvider {
if let Some(recorder) = &recorder {
recorder.request(request_headers.clone(), &body).await?;
}
let attempt = send_with_retry(
let attempt = send_once(
"OpenAI Responses",
|| client.post(&config.request_url)
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
RetryPolicy { retries: config.retry_count, ..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 };
@@ -505,8 +502,10 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
}
fn responses_usage(value: &Value) -> Usage {
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
Usage {
input_tokens: value.get("input_tokens").and_then(Value::as_u64),
input_tokens,
context_input_tokens: input_tokens,
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
cache_read_tokens: value
+1 -31
View File
@@ -1,7 +1,7 @@
//! Records provider requests, responses, usage, and timing.
use std::{
sync::{
atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering},
atomic::{AtomicBool, AtomicI64, AtomicU64, Ordering},
Arc,
},
time::Instant,
@@ -71,7 +71,6 @@ struct Inner {
base_call: NewLlmCall,
detailed: bool,
attempt: Mutex<AttemptState>,
next_attempt: AtomicU32,
next_generation: AtomicU64,
finished: AtomicBool,
}
@@ -120,7 +119,6 @@ impl CallRecorder {
base_call: call.clone(),
detailed: call.detailed,
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
next_attempt: AtomicU32::new(0),
next_generation: AtomicU64::new(0),
finished: AtomicBool::new(false),
}),
@@ -284,34 +282,6 @@ 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);
{
let mut attempt = self.inner.attempt.lock().await;
*attempt = AttemptState::new(call.call_id.clone());
self.inner.finished.store(false, Ordering::Release);
if let Err(error) = self.inner.store.start_llm_call(&call).await {
self.inner.finished.store(true, Ordering::Release);
return Err(error);
}
}
if let Err(error) = self.request(headers, body).await {
self.failed(&error).await?;
return Err(error);
}
Ok(())
}
async fn finish(
&self,
status: &str,
-84
View File
@@ -1,84 +0,0 @@
//! Applies provider retry and backoff behavior.
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<F>(
label: &str,
build: 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,
{
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?;
let error = Error::Provider(format!(
"{label} {status}: {}",
String::from_utf8_lossy(&bytes)
));
if attempt == policy.retries {
return Err(error);
}
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"
);
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) => {}
}
}
unreachable!("the retry loop returns on the final attempt")
}
-3
View File
@@ -18,8 +18,6 @@ use super::{
OpenAiChatProvider, OpenAiResponsesProvider, Provider, ProviderStream,
};
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
pub struct ProviderRouter {
store: Store,
plugins: PluginRegistry,
@@ -91,7 +89,6 @@ impl Provider for ProviderRouter {
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
max_output_tokens: model.max_output_tokens(),
request_timeout,
retry_count: BUILTIN_PROVIDER_RETRIES,
allowed_body_fields: None,
};
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
+110 -14
View File
@@ -14,8 +14,10 @@ use crate::{
};
use super::{
consume_model_cycle, CommitBarrier, CommitCause, MessagesCommitted, ModelCycleFailure,
RunCommand, RunEvent, RunFailure, RunOutcome, RunPort,
consume_model_cycle,
model_retry::{should_retry, MODEL_RETRY_DELAY},
CommitBarrier, CommitCause, MessagesCommitted, RunCommand, RunEvent, RunFailure, RunOutcome,
RunPort,
};
pub struct RunEngine {
@@ -208,12 +210,20 @@ impl RunEngine {
model: prepared.model.clone(),
history,
};
let mut retries = 0_u32;
let mut pending_insertions = Vec::new();
let cycle = 'attempt: loop {
let call_id = if retries == 0 {
format!("{}:{provider_call_index}", prepared.run_id)
} else {
format!("{}:{provider_call_index}:retry-{retries}", prepared.run_id)
};
let invocation = crate::model::ModelInvocation {
call_id: format!("{}:{provider_call_index}", prepared.run_id),
call_id,
run_id: prepared.run_id.to_string(),
conversation_id: prepared.conversation_id.to_string(),
provider_call_index,
request,
request: request.clone(),
};
let cycle_cancellation = cancellation.child_token();
let cycle_events = client.events.clone();
@@ -223,7 +233,6 @@ impl RunEngine {
&cycle_cancellation,
);
tokio::pin!(cycle);
let mut pending_insertions = Vec::new();
let cycle = loop {
tokio::select! {
biased;
@@ -306,21 +315,108 @@ impl RunEngine {
result = &mut cycle => break result,
}
};
let cycle = match cycle {
Ok(cycle) => cycle,
Err(ModelCycleFailure {
failure,
usage: cycle_usage,
..
}) => {
if let Some(cycle_usage) = cycle_usage {
match cycle {
Ok(cycle) => break 'attempt cycle,
Err(cycle_failure) => {
if let Some(cycle_usage) = cycle_failure.usage {
accumulate_usage(&mut usage, cycle_usage);
}
if cancellation.is_cancelled() {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
return (RunOutcome::Failed(failure), usage);
if !should_retry(&cycle_failure, retries) {
return (RunOutcome::Failed(cycle_failure.failure), usage);
}
retries += 1;
let message = failure_message(&cycle_failure.failure);
tracing::warn!(
provider_call_index,
retries,
max_retries = super::model_retry::MAX_MODEL_RETRIES,
delay_ms = MODEL_RETRY_DELAY.as_millis() as u64,
%message,
checkpoint_id = checkpoint.0,
"model attempt failed; retrying from current checkpoint"
);
if emit(
client,
RunEvent::ModelAttemptFailed {
attempt: retries,
message,
},
)
.await
.is_err()
{
return (client_failure(), usage);
}
let delay = tokio::time::sleep(MODEL_RETRY_DELAY);
tokio::pin!(delay);
loop {
tokio::select! {
biased;
command = client.commands.recv() => {
let interruption = match command {
Some(RunCommand::InsertMessages(insertion)) => {
pending_insertions.push(insertion);
continue;
}
Some(RunCommand::BreakMessages(messages)) => messages,
Some(RunCommand::Cancel) => {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
Some(RunCommand::ToolResult(_)) => {
return (
RunOutcome::Failed(RunFailure::Protocol(
"received a tool result while waiting to retry the model".into(),
)),
usage,
);
}
None => return (client_failure(), usage),
};
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
return (client_failure(), usage);
}
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
std::mem::take(&mut pending_insertions),
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
checkpoint = match super::messages::append_batches(
&self.store,
prepared,
client,
cancellation,
checkpoint,
vec![interruption],
)
.await
{
Ok((checkpoint, _)) => checkpoint,
Err(outcome) => return (outcome, usage),
};
continue 'model;
}
_ = cancellation.cancelled() => {
let _ = emit(client, RunEvent::CycleInterrupted).await;
return (RunOutcome::Cancelled, usage);
}
_ = &mut delay => break,
}
}
}
}
};
if let Some(cycle_usage) = cycle.usage {
+4
View File
@@ -104,6 +104,10 @@ pub enum RunEvent {
AutoCompactionStarted,
AutoCompactionCompleted,
CycleInterrupted,
ModelAttemptFailed {
attempt: u32,
message: String,
},
TextStart,
TextDelta(String),
TextEnd,
+1
View File
@@ -7,6 +7,7 @@ mod event;
mod handle;
mod messages;
mod model_cycle;
mod model_retry;
mod port;
mod tool_round;
+100 -15
View File
@@ -29,6 +29,7 @@ pub struct ModelCycleFailure {
pub partial_text: String,
pub partial_reasoning: String,
pub usage: Option<Usage>,
pub retryable: bool,
}
struct OpenTool {
@@ -186,6 +187,7 @@ pub async fn consume_model_cycle(
name: name.clone(),
arguments_text: String::new(),
arguments: serde_json::Value::Null,
argument_error: None,
},
ended: false,
});
@@ -218,14 +220,27 @@ pub async fn consume_model_cycle(
serde_json::from_str(&tool.call.arguments_text)
};
match arguments {
Ok(arguments) => {
Ok(arguments) if arguments.is_object() => {
tool.call.arguments = arguments;
}
Ok(_) => {
tool.call.arguments = serde_json::json!({});
tool.call.argument_error = Some(format!(
"{} arguments must be a JSON object",
tool.call.name
));
}
Err(error) => {
tool.call.arguments = serde_json::json!({});
tool.call.argument_error = Some(format!(
"{} arguments are not valid JSON: {error}",
tool.call.name
));
}
}
tool.ended = true;
send(client, RunEvent::ToolCallEnd { index }).await
}
Err(_) => Err("provider ended a tool call with invalid JSON arguments"),
}
}
Some(_) => Err("provider emitted duplicate ToolCallEnd"),
None => Err("provider ended an unknown tool index"),
},
@@ -239,6 +254,13 @@ pub async fn consume_model_cycle(
ModelEvent::Usage(value) => {
if usage.replace(value).is_some() {
Err("provider emitted duplicate Usage")
} else if send(client, RunEvent::Usage(value)).await.is_err() {
return Err(failure(
RunFailure::Client("client event channel closed".into()),
text,
reasoning,
usage,
));
} else {
Ok(())
}
@@ -280,7 +302,7 @@ pub async fn consume_model_cycle(
.map(|tool| tool.call)
.collect::<Vec<_>>();
if finish_reason == FinishReason::Length {
return Err(failure(
return Err(terminal_failure(
RunFailure::Provider("model stopped before completing the response".into()),
text,
reasoning,
@@ -296,16 +318,6 @@ pub async fn consume_model_cycle(
usage,
));
}
if let Some(usage) = usage {
if send(client, RunEvent::Usage(usage)).await.is_err() {
return Err(failure(
RunFailure::Client("client event channel closed".into()),
text,
reasoning,
Some(usage),
));
}
}
let model_call_id = model_call_id.ok_or_else(|| {
failure(
RunFailure::Protocol("provider completed without Start".into()),
@@ -360,10 +372,83 @@ fn failure(
partial_reasoning: String,
usage: Option<Usage>,
) -> ModelCycleFailure {
let retryable = matches!(failure, RunFailure::Protocol(_) | RunFailure::Provider(_));
ModelCycleFailure {
failure,
partial_text,
partial_reasoning,
usage,
retryable,
}
}
fn terminal_failure(
failure: RunFailure,
partial_text: String,
partial_reasoning: String,
usage: Option<Usage>,
) -> ModelCycleFailure {
ModelCycleFailure {
failure,
partial_text,
partial_reasoning,
usage,
retryable: false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
model::Usage,
provider::{FinishReason, ModelEvent},
};
use tokio_stream::wrappers::ReceiverStream;
#[tokio::test]
async fn usage_is_forwarded_before_the_provider_call_finishes() {
let (provider_tx, provider_rx) = tokio::sync::mpsc::channel(4);
let stream = Box::pin(ReceiverStream::new(provider_rx));
let (event_tx, mut event_rx) = tokio::sync::mpsc::channel(4);
let cancellation = CancellationToken::new();
let cycle_cancellation = cancellation.clone();
let cycle = tokio::spawn(async move {
consume_model_cycle(stream, &event_tx, &cycle_cancellation).await
});
let usage = Usage {
input_tokens: Some(100),
context_input_tokens: Some(100),
output_tokens: Some(20),
total_tokens: Some(120),
..Default::default()
};
provider_tx
.send(Ok(ModelEvent::Start {
model_call_id: "call".into(),
}))
.await
.unwrap();
provider_tx
.send(Ok(ModelEvent::Usage(usage)))
.await
.unwrap();
let event = tokio::time::timeout(std::time::Duration::from_secs(1), event_rx.recv())
.await
.unwrap()
.unwrap();
assert!(matches!(event, RunEvent::Usage(value) if value == usage));
assert!(!cycle.is_finished(), "usage must arrive before Done");
provider_tx
.send(Ok(ModelEvent::Done(FinishReason::Stop)))
.await
.unwrap();
drop(provider_tx);
let result = cycle.await.unwrap().unwrap();
assert_eq!(result.usage, Some(usage));
assert!(event_rx.try_recv().is_err(), "usage must be forwarded once");
}
}
+42
View File
@@ -0,0 +1,42 @@
//! Defines retry policy for one logical model call.
use std::time::Duration;
use super::ModelCycleFailure;
pub(super) const MAX_MODEL_RETRIES: u32 = 8;
pub(super) const MODEL_RETRY_DELAY: Duration = Duration::from_secs(5);
pub(super) fn should_retry(failure: &ModelCycleFailure, retries: u32) -> bool {
failure.retryable && retries < MAX_MODEL_RETRIES
}
#[cfg(test)]
mod tests {
use super::*;
use crate::run::{ModelCycleFailure, RunFailure};
fn failure(retryable: bool) -> ModelCycleFailure {
ModelCycleFailure {
failure: RunFailure::Provider("failed".into()),
partial_text: String::new(),
partial_reasoning: String::new(),
usage: None,
retryable,
}
}
#[test]
fn permits_eight_retries_after_the_initial_attempt() {
let retryable = failure(true);
for retries in 0..MAX_MODEL_RETRIES {
assert!(should_retry(&retryable, retries));
}
assert!(!should_retry(&retryable, MAX_MODEL_RETRIES));
}
#[test]
fn terminal_failures_never_retry() {
assert!(!should_retry(&failure(false), 0));
}
}
+12 -1
View File
@@ -445,9 +445,20 @@ mod tests {
.await
.unwrap();
let argument_error_column_exists: i64 = sqlx::query_scalar(
"SELECT EXISTS(
SELECT 1 FROM pragma_table_info('tool_round_calls')
WHERE name = 'argument_error'
)",
)
.fetch_one(&pool)
.await
.unwrap();
assert_eq!(checksum_after, checksum_before);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7, 8]);
assert_eq!(versions, vec![1, 2, 3, 4, 5, 6, 7, 8, 9]);
assert_eq!(checkpoint_table_exists, 1);
assert_eq!(argument_error_column_exists, 1);
}
}
+7 -12
View File
@@ -93,22 +93,16 @@ impl Store {
for call in calls {
sqlx::query(
"INSERT INTO tool_round_calls
(round_id, call_index, call_id, model_call_id, name, arguments_json, status)
VALUES (?, ?, ?, ?, ?, ?, 'pending')",
(round_id, call_index, call_id, model_call_id, name, arguments_json, argument_error, status)
VALUES (?, ?, ?, ?, ?, ?, ?, 'pending')",
)
.bind(round_id.as_str())
.bind(call.index as i64)
.bind(&call.call_id)
.bind(&call.model_call_id)
.bind(&call.name)
// A no-argument tool call streams no argument text; persist it as an
// empty object so the `arguments_json` column always holds valid JSON
// and can be re-parsed on load.
.bind(if call.arguments_text.trim().is_empty() {
"{}"
} else {
call.arguments_text.as_str()
})
.bind(serde_json::to_string(&call.arguments)?)
.bind(call.argument_error.as_deref())
.execute(&mut *tx)
.await?;
}
@@ -276,7 +270,7 @@ impl Store {
return Ok(None);
};
let rows = sqlx::query(
"SELECT call_index, call_id, model_call_id, name, arguments_json, status
"SELECT call_index, call_id, model_call_id, name, arguments_json, argument_error, status
FROM tool_round_calls WHERE round_id = ? ORDER BY call_index",
)
.bind(round_id.as_str())
@@ -287,7 +281,7 @@ impl Store {
for row in rows {
let arguments_text: String = row.get(4);
let call_id: String = row.get(1);
if row.get::<&str, _>(5) == "completed" {
if row.get::<&str, _>(6) == "completed" {
completed.push(call_id.clone());
}
calls.push(ToolCall {
@@ -297,6 +291,7 @@ impl Store {
name: row.get(3),
arguments: serde_json::from_str(&arguments_text)?,
arguments_text,
argument_error: row.get(5),
});
}
Ok(Some(ToolRoundSnapshot {
+2
View File
@@ -62,6 +62,7 @@ async fn summarize_replaces_model_history_and_preserves_cursor_history() {
ModelEvent::TextEnd,
ModelEvent::Usage(Usage {
input_tokens: Some(4_012),
context_input_tokens: Some(4_012),
output_tokens: Some(9),
total_tokens: Some(4_021),
..Default::default()
@@ -442,6 +443,7 @@ fn text_response(text: &str, input: u64, output: u64) -> Vec<ModelEvent> {
ModelEvent::TextEnd,
ModelEvent::Usage(Usage {
input_tokens: Some(input),
context_input_tokens: Some(input),
output_tokens: Some(output),
total_tokens: Some(input + output),
..Default::default()
+168 -34
View File
@@ -51,10 +51,35 @@ async fn abort_command_cancels_the_run_and_closes_output() {
}
#[tokio::test]
async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_error() {
async fn provider_failure_retries_from_the_current_checkpoint_without_hiding_partial_output() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push_error(Error::Provider("provider failed".into()));
provider.push_results(vec![
Ok(ModelEvent::Start {
model_call_id: "attempt-0".into(),
}),
Ok(ModelEvent::TextStart),
Ok(ModelEvent::TextDelta("partial ".into())),
Ok(ModelEvent::ToolCallStart {
index: 0,
call_id: "failed-tool".into(),
name: "Read".into(),
}),
Ok(ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":".into(),
}),
Err(Error::Provider("stream disconnected".into())),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "attempt-1".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("completed".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
@@ -63,7 +88,106 @@ async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_e
.unwrap();
let registry = TransportRegistry::new(
store.clone(),
Arc::new(provider),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
);
let handle = registry.get_or_create("retry-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(TransportCommand::Append {
seqno: 0,
message: Box::new(protocol_client_run("retry", "retry-user")),
})
.await
.unwrap();
let mut seqno = 1;
let mut text = String::new();
let mut failed_tool_completed = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(60), output.recv())
.await
.unwrap()
.unwrap();
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&payload).unwrap(),
serde_json::json!({})
);
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(TransportCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
match update.message {
Some(pb::interaction_update::Message::TextDelta(delta)) => {
text.push_str(&delta.text);
}
Some(pb::interaction_update::Message::ToolCallCompleted(completed))
if completed.call_id == "failed-tool" =>
{
let tool = completed.tool_call.expect("failed tool completion");
failed_tool_completed = tool.completed_at_ms.is_some();
}
_ => {}
}
}
_ => {}
}
}
assert_eq!(text, "partial completed");
assert!(
failed_tool_completed,
"failed attempt must terminate its partial tool card"
);
let requests = provider.requests();
assert_eq!(requests.len(), 2);
assert_eq!(requests[0], requests[1]);
let messages = store
.load_current_messages(&cursor_server::model::ConversationId::new(
"protocol-failed-conversation",
))
.await
.unwrap();
assert!(messages.iter().any(|message| {
matches!(&message.content, MessageContent::Assistant { text, .. } if text == "completed")
}));
assert!(!messages.iter().any(|message| {
matches!(&message.content, MessageContent::Assistant { text, .. } if text.contains("partial"))
}));
}
#[tokio::test]
async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_error() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "length-limited".into(),
},
ModelEvent::Done(FinishReason::Length),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = TransportRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
);
let handle = registry.get_or_create("failed-request").await.unwrap();
@@ -148,6 +272,7 @@ async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_e
None
);
assert_eq!(provider.requests().len(), 1, "Length must not retry");
let messages = store
.load_current_messages(&cursor_server::model::ConversationId::new(
"failed-conversation",
@@ -164,7 +289,7 @@ async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_e
}
#[tokio::test]
async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes() {
async fn unknown_tool_response_id_is_ignored_and_the_run_continues() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
@@ -183,6 +308,15 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "model-call-2".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("done".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
@@ -209,7 +343,7 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
let mut append_seqno = 1;
let mut saw_turn_ended = false;
let error_json = loop {
let end_stream = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
@@ -231,25 +365,43 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
append_seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
// An unknown numeric bridge id is a runtime protocol error.
handle
.command(TransportCommand::Append {
seqno: append_seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
// Unknown bridge ids are ignored; the valid response still completes the tool.
for message in [
pb::ExecClientMessage {
id: exec.id + 1_000,
exec_id: String::new(),
message: None,
..Default::default()
},
pb::ExecClientMessage {
id: exec.id,
exec_id: String::new(),
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
path: "/tmp/a".into(),
output: Some(pb::read_success::Output::Content("value".into())),
..Default::default()
})),
},
)),
..Default::default()
},
] {
handle
.command(TransportCommand::Append {
seqno: append_seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(
pb::agent_client_message::Message::ExecClientMessage(message),
),
}),
})
.await
.unwrap();
append_seqno += 1;
}
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
if matches!(
update.message,
@@ -266,12 +418,8 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
}
};
assert!(!saw_turn_ended);
assert_eq!(error_json["error"]["code"], "invalid_argument");
assert_eq!(
error_json["error"]["message"],
"unknown ExecClientMessage id: 1001"
);
assert!(saw_turn_ended);
assert_eq!(end_stream, serde_json::json!({}));
assert_eq!(
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
.await
@@ -279,28 +427,14 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
None
);
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(1);
let (status, failure_summary) = loop {
let row: (String, Option<String>) =
let (status, failure_summary): (String, Option<String>) =
sqlx::query_as("SELECT status, failure_summary FROM runs WHERE cursor_request_id = ?")
.bind("protocol-failed-request")
.fetch_one(store.pool())
.await
.unwrap();
if row.0 != "running" {
break row;
}
assert!(
tokio::time::Instant::now() < deadline,
"Run remained running after the Cursor session failed"
);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
};
assert_eq!(status, "failed");
assert_eq!(
failure_summary.as_deref(),
Some("unknown ExecClientMessage id: 1001")
);
assert_eq!(status, "completed");
assert_eq!(failure_summary, None);
}
#[tokio::test]
+1
View File
@@ -1580,6 +1580,7 @@ fn text_response(text: &str) -> Vec<ModelEvent> {
ModelEvent::TextEnd,
ModelEvent::Usage(Usage {
input_tokens: Some(1),
context_input_tokens: Some(1),
output_tokens: Some(1),
total_tokens: Some(2),
..Default::default()
+4 -1
View File
@@ -31,10 +31,13 @@ pub struct FakeProvider {
impl FakeProvider {
pub fn push(&self, events: Vec<ModelEvent>) {
self.push_results(events.into_iter().map(Ok).collect());
}
pub fn push_results(&self, events: Vec<Result<ModelEvent, Error>>) {
self.responses
.lock()
.unwrap()
.push_back(FakeResponse::Events(events.into_iter().map(Ok).collect()));
.push_back(FakeResponse::Events(events));
}
pub fn push_error(&self, error: Error) {
self.responses
+96 -20
View File
@@ -25,9 +25,11 @@ use cursor_server::{
OPENAI_CHAT_ENDPOINT,
},
provider::{FinishReason, ModelEvent},
run::consume_model_cycle,
};
use prost::Message;
use serde_json::json;
use tokio_util::sync::CancellationToken;
fn call(id: &str, name: &str) -> ToolCall {
ToolCall {
@@ -37,6 +39,7 @@ fn call(id: &str, name: &str) -> ToolCall {
name: name.into(),
arguments_text: "{}".into(),
arguments: json!({}),
argument_error: None,
}
}
@@ -68,6 +71,37 @@ fn mcp_context(server: &str, provider: &str, tool: &str) -> ExecContext {
context
}
#[tokio::test]
async fn malformed_provider_tool_json_is_kept_as_a_tool_validation_error() {
let stream = Box::pin(futures_util::stream::iter(vec![
Ok(ModelEvent::Start {
model_call_id: "model-call".into(),
}),
Ok(ModelEvent::ToolCallStart {
index: 0,
call_id: "call-malformed".into(),
name: "Read".into(),
}),
Ok(ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":".into(),
}),
Ok(ModelEvent::ToolCallEnd { index: 0 }),
Ok(ModelEvent::Done(FinishReason::ToolUse)),
]));
let (events, _receiver) = tokio::sync::mpsc::channel(16);
let result = consume_model_cycle(stream, &events, &CancellationToken::new())
.await
.expect("malformed tool JSON must not fail the model cycle");
assert_eq!(result.calls.len(), 1);
assert!(result.calls[0]
.argument_error
.as_deref()
.is_some_and(|message| message.contains("not valid JSON")));
assert_eq!(result.calls[0].arguments, json!({}));
}
#[test]
fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
let call = ToolCall {
@@ -77,6 +111,7 @@ fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
name: "mcp_repo_lookup".into(),
arguments_text: "{\"query\":\"x\"}".into(),
arguments: json!({"query": "x"}),
argument_error: None,
};
let definition = pb::McpToolDefinition {
name: "mcp_repo_lookup".into(),
@@ -372,6 +407,52 @@ async fn unknown_mcp_descriptor_returns_a_tool_error_without_client_discovery()
assert!(completion.result().content.contains("descriptor not found"));
}
#[tokio::test]
async fn invalid_tool_arguments_complete_as_tool_errors() {
let dispatcher = ToolDispatcher::new(CursorToolRuntime::default());
let completed = HashSet::new();
let started = HashSet::new();
let state = || ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
};
let mut malformed = call("call-malformed", "Read");
malformed.argument_error = Some("Read arguments are not valid JSON".into());
let mut missing = call("call-missing", "Shell");
missing.arguments = json!({"description": "missing command"});
let mut wrong_type = call("call-type", "Shell");
wrong_type.arguments = json!({"command": 42});
let mut invalid_timeout = call("call-timeout", "Shell");
invalid_timeout.arguments = json!({"command": "pwd", "block_until_ms": -1});
for (invocation, expected) in [
(malformed, "not valid JSON"),
(missing, "missing command"),
(wrong_type, "missing command"),
(invalid_timeout, "out of range"),
] {
let dispatched = dispatcher
.start_batch(
&[invocation],
state(),
&[],
&BTreeMap::new(),
&exec_context(),
)
.await
.unwrap();
let completion = dispatched[0]
.completion
.as_ref()
.expect("invalid arguments must complete as a tool error");
assert!(completion.result().is_error);
assert!(completion.result().content.contains(expected));
}
}
#[tokio::test]
async fn shell_uses_background_timeout_and_preserves_stream_identity() {
let mut shell = call("call-shell", "Shell");
@@ -701,12 +782,15 @@ async fn an_exec_result_must_match_the_reserved_tool() {
},
&pending,
)
.await;
let Err(error) = result else {
panic!("mismatched result must fail")
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = result else {
panic!("mismatched result must complete as a tool error")
};
assert!(error
.to_string()
assert!(completion.result().is_error);
assert!(completion
.result()
.content
.contains("unexpected Exec result for tool Read"));
assert!(pending.exec_call(id).await.is_none());
assert_eq!(pending.completed_call(id).await.as_deref(), Some("call-1"));
@@ -718,15 +802,13 @@ async fn an_exec_result_must_match_the_reserved_tool() {
},
&pending,
)
.await;
let Err(duplicate) = duplicate else {
panic!("duplicate terminal result must fail")
};
assert!(duplicate.to_string().contains("duplicate terminal"));
.await
.unwrap();
assert!(matches!(duplicate, codec::ClientExecEvent::Pending));
}
#[tokio::test]
async fn unknown_exec_id_is_a_protocol_error() {
async fn unknown_exec_id_is_ignored() {
let result = codec::client_event(
&pb::ExecClientMessage {
id: 999,
@@ -737,15 +819,9 @@ async fn unknown_exec_id_is_a_protocol_error() {
},
&CursorToolRuntime::default(),
)
.await;
let Err(error) = result else {
panic!("unknown Exec id must fail")
};
assert!(matches!(
error,
cursor_server::Error::Protocol(message)
if message == "unknown ExecClientMessage id: 999"
));
.await
.unwrap();
assert!(matches!(result, codec::ClientExecEvent::Pending));
}
#[tokio::test]