diff --git a/server/migrations/0009_add_tool_call_argument_error.sql b/server/migrations/0009_add_tool_call_argument_error.sql new file mode 100644 index 0000000..5d81bc6 --- /dev/null +++ b/server/migrations/0009_add_tool_call_argument_error.sql @@ -0,0 +1 @@ +ALTER TABLE tool_round_calls ADD COLUMN argument_error TEXT; diff --git a/server/src/config.rs b/server/src/config.rs index ed8ce51..d2a9810 100644 --- a/server/src/config.rs +++ b/server/src/config.rs @@ -46,7 +46,6 @@ pub struct ProviderConfig { pub custom_headers: reqwest::header::HeaderMap, pub max_output_tokens: Option, pub request_timeout: Duration, - pub retry_count: u32, pub allowed_body_fields: Option>, } diff --git a/server/src/cursor/checkpoint/messages/decode.rs b/server/src/cursor/checkpoint/messages/decode.rs index 5b601ec..a601212 100644 --- a/server/src/cursor/checkpoint/messages/decode.rs +++ b/server/src/cursor/checkpoint/messages/decode.rs @@ -102,6 +102,12 @@ pub fn decode_pending(value: &str) -> Result { .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 { name: call.name, arguments_text: serde_json::to_string(&call.arguments)?, arguments: call.arguments, + argument_error, }) }) .collect::>>()?; diff --git a/server/src/cursor/checkpoint/messages/encode.rs b/server/src/cursor/checkpoint/messages/encode.rs index f2c5eae..e5bc74a 100644 --- a/server/src/cursor/checkpoint/messages/encode.rs +++ b/server/src/cursor/checkpoint/messages/encode.rs @@ -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, 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!({ - "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, - }), - ) + 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(), ), diff --git a/server/src/cursor/checkpoint/messages/tests.rs b/server/src/cursor/checkpoint/messages/tests.rs index e454382..663f292 100644 --- a/server/src/cursor/checkpoint/messages/tests.rs +++ b/server/src/cursor/checkpoint/messages/tests.rs @@ -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"); } diff --git a/server/src/cursor/checkpoint/steps.rs b/server/src/cursor/checkpoint/steps.rs index 8168433..f4e243b 100644 --- a/server/src/cursor/checkpoint/steps.rs +++ b/server/src/cursor/checkpoint/steps.rs @@ -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(); diff --git a/server/src/cursor/conversation/output.rs b/server/src/cursor/conversation/output.rs index f86a2a3..d66c54d 100644 --- a/server/src/cursor/conversation/output.rs +++ b/server/src/cursor/conversation/output.rs @@ -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,8 +362,27 @@ 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)? { - self.handle.emit(&message)?; + 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 } => { @@ -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 { + 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() { diff --git a/server/src/cursor/conversation/runtime.rs b/server/src/cursor/conversation/runtime.rs index 6b9f939..dc3f1ef 100644 --- a/server/src/cursor/conversation/runtime.rs +++ b/server/src/cursor/conversation/runtime.rs @@ -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!( - "Exec {} failed: {}", - pending.call.call_id, throw.error - )), + 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" ), } } diff --git a/server/src/cursor/tools/codec/request.rs b/server/src/cursor/tools/codec/request.rs index 60ec9f0..a2f5867 100644 --- a/server/src/cursor/tools/codec/request.rs +++ b/server/src/cursor/tools/codec/request.rs @@ -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 diff --git a/server/src/cursor/tools/codec/response.rs b/server/src/cursor/tools/codec/response.rs index 0cf889a..ea3b679 100644 --- a/server/src/cursor/tools/codec/response.rs +++ b/server/src/cursor/tools/codec/response.rs @@ -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 { - 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 { - 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( diff --git a/server/src/cursor/tools/compat.rs b/server/src/cursor/tools/compat.rs index be85453..90ab81f 100644 --- a/server/src/cursor/tools/compat.rs +++ b/server/src/cursor/tools/compat.rs @@ -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() diff --git a/server/src/cursor/tools/edit.rs b/server/src/cursor/tools/edit.rs index 09146d7..084af3e 100644 --- a/server/src/cursor/tools/edit.rs +++ b/server/src/cursor/tools/edit.rs @@ -265,6 +265,7 @@ mod tests { "old_string": old_string, "new_string": "replacement", }), + argument_error: None, } } diff --git a/server/src/cursor/tools/mod.rs b/server/src/cursor/tools/mod.rs index 5a59081..45899e9 100644 --- a/server/src/cursor/tools/mod.rs +++ b/server/src/cursor/tools/mod.rs @@ -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,15 +162,19 @@ impl ToolDispatcher { let Some(next) = next else { return Ok(None); }; - self.start( - &next.call, - next.message_index, - next.publish_started, - &BTreeMap::new(), - &next.context, - ) - .await - .map(Some) + match self + .start( + &next.call, + next.message_index, + next.publish_started, + &BTreeMap::new(), + &next.context, + ) + .await + { + Ok(started) => Ok(Some(started)), + Err(error) => recover_validation_failure(&next.call, error).map(Some), + } } pub async fn interrupt_for_message(&self) -> Vec { @@ -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( - &self.results, - &self.search, - &self.fetch, - pending, - response, - ) - .await? - { - tool_call_dispatch::InteractionContinuation::Completed(completion) => { - ClientToolEvent::Completed(completion) - } - tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending, - }, + let call = pending.call.clone(); + let continuation = match tool_call_dispatch::resume_interaction( + &self.results, + &self.search, + &self.fetch, + pending, + response, ) + .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 { + match error { + Error::Protocol(message) => Ok(validation_failure(call, message)), + Error::Json(error) => Ok(validation_failure(call, error.to_string())), + error => Err(error), } } diff --git a/server/src/cursor/tools/stream.rs b/server/src/cursor/tools/stream.rs index 831ddfa..aa04498 100644 --- a/server/src/cursor/tools/stream.rs +++ b/server/src/cursor/tools/stream.rs @@ -248,6 +248,7 @@ mod tests { name: "Task".into(), arguments_text: arguments_text.into(), arguments: serde_json::Value::Null, + argument_error: None, } } diff --git a/server/src/cursor/tools/tool_call_result/interaction.rs b/server/src/cursor/tools/tool_call_result/interaction.rs index 79601b1..5295d84 100644 --- a/server/src/cursor/tools/tool_call_result/interaction.rs +++ b/server/src/cursor/tools/tool_call_result/interaction.rs @@ -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, } diff --git a/server/src/model/observability.rs b/server/src/model/observability.rs index ec5f443..b8a656d 100644 --- a/server/src/model/observability.rs +++ b/server/src/model/observability.rs @@ -9,6 +9,7 @@ mod usage { #[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)] pub struct Usage { pub input_tokens: Option, + pub context_input_tokens: Option, pub output_tokens: Option, pub total_tokens: Option, pub cache_read_tokens: Option, @@ -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); diff --git a/server/src/model/tool.rs b/server/src/model/tool.rs index 41175ea..c522765 100644 --- a/server/src/model/tool.rs +++ b/server/src/model/tool.rs @@ -19,6 +19,7 @@ pub struct ToolCall { pub name: String, pub arguments_text: String, pub arguments: Value, + pub argument_error: Option, } #[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)] diff --git a/server/src/plugin/wire.rs b/server/src/plugin/wire.rs index 2b97c20..eee28a1 100644 --- a/server/src/plugin/wire.rs +++ b/server/src/plugin/wire.rs @@ -147,8 +147,10 @@ pub fn model_event(value: &serde_json::Value) -> Result { .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), diff --git a/server/src/provider/anthropic.rs b/server/src/provider/anthropic.rs index c4bfd6e..1f26cb2 100644 --- a/server/src/provider/anthropic.rs +++ b/server/src/provider/anthropic.rs @@ -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 { } 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)); + } +} diff --git a/server/src/provider/attempt.rs b/server/src/provider/attempt.rs new file mode 100644 index 0000000..a9f17af --- /dev/null +++ b/server/src/provider/attempt.rs @@ -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( + label: &str, + build: F, + cancellation: &CancellationToken, + recorder: Option<&CallRecorder>, +) -> Result +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(_))); + } +} diff --git a/server/src/provider/mod.rs b/server/src/provider/mod.rs index 4fb9848..830c793 100644 --- a/server/src/provider/mod.rs +++ b/server/src/provider/mod.rs @@ -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; diff --git a/server/src/provider/openai_chat.rs b/server/src/provider/openai_chat.rs index 32835a4..32b546c 100644 --- a/server/src/provider/openai_chat.rs +++ b/server/src/provider/openai_chat.rs @@ -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 diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index 993f936..99eb7a7 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -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 { } 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 diff --git a/server/src/provider/recorder.rs b/server/src/provider/recorder.rs index f161730..652fc4e 100644 --- a/server/src/provider/recorder.rs +++ b/server/src/provider/recorder.rs @@ -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, - 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, diff --git a/server/src/provider/retry.rs b/server/src/provider/retry.rs deleted file mode 100644 index c050f3a..0000000 --- a/server/src/provider/retry.rs +++ /dev/null @@ -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( - label: &str, - build: F, - policy: RetryPolicy, - cancellation: &CancellationToken, - recorder: Option<&CallRecorder>, - request_headers: serde_json::Value, - request_body: &serde_json::Value, -) -> Result -where - F: Fn() -> reqwest::RequestBuilder, -{ - for attempt in 0..=policy.retries { - let response = tokio::select! { - _ = cancellation.cancelled() => return Ok(Attempt::Cancelled), - response = build().send() => response, - }?; - if let Some(recorder) = recorder { - recorder - .response_headers(response.status().as_u16()) - .await?; - } - if response.status().is_success() { - return Ok(Attempt::Response(response)); - } - let status = response.status(); - let bytes = response.bytes().await?; - 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") -} diff --git a/server/src/provider/router.rs b/server/src/provider/router.rs index 75dda91..d16f122 100644 --- a/server/src/provider/router.rs +++ b/server/src/provider/router.rs @@ -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()?; diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index 1c75334..1195fd6 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -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,119 +210,213 @@ impl RunEngine { model: prepared.model.clone(), history, }; - let invocation = crate::model::ModelInvocation { - call_id: format!("{}:{provider_call_index}", prepared.run_id), - run_id: prepared.run_id.to_string(), - conversation_id: prepared.conversation_id.to_string(), - provider_call_index, - request, - }; - let cycle_cancellation = cancellation.child_token(); - let cycle_events = client.events.clone(); - let cycle = consume_model_cycle( - self.provider.stream(invocation, cycle_cancellation.clone()), - &cycle_events, - &cycle_cancellation, - ); - tokio::pin!(cycle); + let mut retries = 0_u32; let mut pending_insertions = Vec::new(); - let cycle = loop { - tokio::select! { - biased; - command = client.commands.recv() => { - let interruption = match command { - Some(RunCommand::InsertMessages(insertion)) => { - pending_insertions.push(insertion); - continue; + 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, + run_id: prepared.run_id.to_string(), + conversation_id: prepared.conversation_id.to_string(), + provider_call_index, + request: request.clone(), + }; + let cycle_cancellation = cancellation.child_token(); + let cycle_events = client.events.clone(); + let cycle = consume_model_cycle( + self.provider.stream(invocation, cycle_cancellation.clone()), + &cycle_events, + &cycle_cancellation, + ); + tokio::pin!(cycle); + let cycle = 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) => { + cycle_cancellation.cancel(); + let _ = cycle.await; + let _ = emit(client, RunEvent::CycleInterrupted).await; + return (RunOutcome::Cancelled, usage); + } + Some(RunCommand::ToolResult(_)) => { + cycle_cancellation.cancel(); + let _ = cycle.await; + let _ = emit(client, RunEvent::CycleInterrupted).await; + return ( + RunOutcome::Failed(RunFailure::Protocol( + "received a tool result while the model was running".into(), + )), + usage, + ); + } + None => { + cycle_cancellation.cancel(); + let _ = cycle.await; + let _ = emit(client, RunEvent::CycleInterrupted).await; + return (client_failure(), usage); + } + }; + cycle_cancellation.cancel(); + let interrupted = cycle.await; + match interrupted { + Ok(cycle) => { + if let Some(cycle_usage) = cycle.usage { + accumulate_usage(&mut usage, cycle_usage); + } + } + Err(failure) => { + if let Some(cycle_usage) = failure.usage { + accumulate_usage(&mut usage, cycle_usage); + } + } } - Some(RunCommand::BreakMessages(messages)) => messages, - Some(RunCommand::Cancel) => { - cycle_cancellation.cancel(); - let _ = cycle.await; - let _ = emit(client, RunEvent::CycleInterrupted).await; - return (RunOutcome::Cancelled, usage); - } - Some(RunCommand::ToolResult(_)) => { - cycle_cancellation.cancel(); - let _ = cycle.await; - let _ = emit(client, RunEvent::CycleInterrupted).await; - return ( - RunOutcome::Failed(RunFailure::Protocol( - "received a tool result while the model was running".into(), - )), - usage, - ); - } - None => { - cycle_cancellation.cancel(); - let _ = cycle.await; - let _ = emit(client, RunEvent::CycleInterrupted).await; + if emit(client, RunEvent::CycleInterrupted).await.is_err() { return (client_failure(), usage); } - }; - cycle_cancellation.cancel(); - let interrupted = cycle.await; - match interrupted { - Ok(cycle) => { - if let Some(cycle_usage) = cycle.usage { - accumulate_usage(&mut usage, cycle_usage); - } - } - Err(failure) => { - if let Some(cycle_usage) = failure.usage { - accumulate_usage(&mut usage, cycle_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; + }, + result = &mut cycle => break result, + } + }; + 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 emit(client, RunEvent::CycleInterrupted).await.is_err() { + if cancellation.is_cancelled() { + let _ = emit(client, RunEvent::CycleInterrupted).await; + return (RunOutcome::Cancelled, 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); } - 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; - }, - 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 { - accumulate_usage(&mut usage, cycle_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 cancellation.is_cancelled() { - let _ = emit(client, RunEvent::CycleInterrupted).await; - return (RunOutcome::Cancelled, usage); - } - return (RunOutcome::Failed(failure), usage); } }; if let Some(cycle_usage) = cycle.usage { diff --git a/server/src/run/event.rs b/server/src/run/event.rs index ba6da13..345be46 100644 --- a/server/src/run/event.rs +++ b/server/src/run/event.rs @@ -104,6 +104,10 @@ pub enum RunEvent { AutoCompactionStarted, AutoCompactionCompleted, CycleInterrupted, + ModelAttemptFailed { + attempt: u32, + message: String, + }, TextStart, TextDelta(String), TextEnd, diff --git a/server/src/run/mod.rs b/server/src/run/mod.rs index 5ace76f..d9462db 100644 --- a/server/src/run/mod.rs +++ b/server/src/run/mod.rs @@ -7,6 +7,7 @@ mod event; mod handle; mod messages; mod model_cycle; +mod model_retry; mod port; mod tool_round; diff --git a/server/src/run/model_cycle.rs b/server/src/run/model_cycle.rs index bd4ddc7..e6a24ba 100644 --- a/server/src/run/model_cycle.rs +++ b/server/src/run/model_cycle.rs @@ -29,6 +29,7 @@ pub struct ModelCycleFailure { pub partial_text: String, pub partial_reasoning: String, pub usage: Option, + 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,13 +220,26 @@ 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; - tool.ended = true; - send(client, RunEvent::ToolCallEnd { index }).await } - Err(_) => Err("provider ended a tool call with invalid JSON 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 } 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::>(); 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, ) -> 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, +) -> 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"); } } diff --git a/server/src/run/model_retry.rs b/server/src/run/model_retry.rs new file mode 100644 index 0000000..cccde21 --- /dev/null +++ b/server/src/run/model_retry.rs @@ -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)); + } +} diff --git a/server/src/store/migrations.rs b/server/src/store/migrations.rs index 1a06448..44d69f4 100644 --- a/server/src/store/migrations.rs +++ b/server/src/store/migrations.rs @@ -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); } } diff --git a/server/src/store/tool_rounds.rs b/server/src/store/tool_rounds.rs index a7f3fc5..ca961fd 100644 --- a/server/src/store/tool_rounds.rs +++ b/server/src/store/tool_rounds.rs @@ -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 { diff --git a/server/tests/compaction.rs b/server/tests/compaction.rs index 3a4ab8f..d778a66 100644 --- a/server/tests/compaction.rs +++ b/server/tests/compaction.rs @@ -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::TextEnd, ModelEvent::Usage(Usage { input_tokens: Some(input), + context_input_tokens: Some(input), output_tokens: Some(output), total_tokens: Some(input + output), ..Default::default() diff --git a/server/tests/error_lifecycle.rs b/server/tests/error_lifecycle.rs index 2897f67..250a937 100644 --- a/server/tests/error_lifecycle.rs +++ b/server/tests/error_lifecycle.rs @@ -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::(&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,24 +365,42 @@ 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( - pb::ExecClientMessage { - id: exec.id + 1_000, - exec_id: String::new(), - message: None, + // 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() - }, - )), - }), - }) - .await - .unwrap(); - append_seqno += 1; + })), + }, + )), + ..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!( @@ -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) = - 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") - ); + let (status, failure_summary): (String, Option) = + sqlx::query_as("SELECT status, failure_summary FROM runs WHERE cursor_request_id = ?") + .bind("protocol-failed-request") + .fetch_one(store.pool()) + .await + .unwrap(); + assert_eq!(status, "completed"); + assert_eq!(failure_summary, None); } #[tokio::test] diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index c2db150..1cfd882 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -1580,6 +1580,7 @@ fn text_response(text: &str) -> Vec { ModelEvent::TextEnd, ModelEvent::Usage(Usage { input_tokens: Some(1), + context_input_tokens: Some(1), output_tokens: Some(1), total_tokens: Some(2), ..Default::default() diff --git a/server/tests/support/fake_provider.rs b/server/tests/support/fake_provider.rs index 89dbead..8b2d3e1 100644 --- a/server/tests/support/fake_provider.rs +++ b/server/tests/support/fake_provider.rs @@ -31,10 +31,13 @@ pub struct FakeProvider { impl FakeProvider { pub fn push(&self, events: Vec) { + self.push_results(events.into_iter().map(Ok).collect()); + } + pub fn push_results(&self, events: Vec>) { 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 diff --git a/server/tests/tool_round.rs b/server/tests/tool_round.rs index bc8d1e2..309ed1d 100644 --- a/server/tests/tool_round.rs +++ b/server/tests/tool_round.rs @@ -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]