mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +08:00
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:
@@ -0,0 +1 @@
|
||||
ALTER TABLE tool_round_calls ADD COLUMN argument_error TEXT;
|
||||
@@ -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!({
|
||||
"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(),
|
||||
),
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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<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() {
|
||||
|
||||
@@ -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"
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -265,6 +265,7 @@ mod tests {
|
||||
"old_string": old_string,
|
||||
"new_string": "replacement",
|
||||
}),
|
||||
argument_error: None,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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<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(
|
||||
&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<DispatchedTool> {
|
||||
match error {
|
||||
Error::Protocol(message) => Ok(validation_failure(call, message)),
|
||||
Error::Json(error) => Ok(validation_failure(call, error.to_string())),
|
||||
error => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,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;
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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()?;
|
||||
|
||||
+203
-107
@@ -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 {
|
||||
|
||||
@@ -104,6 +104,10 @@ pub enum RunEvent {
|
||||
AutoCompactionStarted,
|
||||
AutoCompactionCompleted,
|
||||
CycleInterrupted,
|
||||
ModelAttemptFailed {
|
||||
attempt: u32,
|
||||
message: String,
|
||||
},
|
||||
TextStart,
|
||||
TextDelta(String),
|
||||
TextEnd,
|
||||
|
||||
@@ -7,6 +7,7 @@ mod event;
|
||||
mod handle;
|
||||
mod messages;
|
||||
mod model_cycle;
|
||||
mod model_retry;
|
||||
mod port;
|
||||
mod tool_round;
|
||||
|
||||
|
||||
+100
-15
@@ -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,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::<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");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
+184
-50
@@ -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,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<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")
|
||||
);
|
||||
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();
|
||||
assert_eq!(status, "completed");
|
||||
assert_eq!(failure_summary, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user