mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +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 custom_headers: reqwest::header::HeaderMap,
|
||||||
pub max_output_tokens: Option<u64>,
|
pub max_output_tokens: Option<u64>,
|
||||||
pub request_timeout: Duration,
|
pub request_timeout: Duration,
|
||||||
pub retry_count: u32,
|
|
||||||
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
|
pub allowed_body_fields: Option<std::collections::HashSet<String>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -102,6 +102,12 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
|
|||||||
.into_iter()
|
.into_iter()
|
||||||
.enumerate()
|
.enumerate()
|
||||||
.map(|(index, call)| {
|
.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 {
|
Ok(ToolCall {
|
||||||
index,
|
index,
|
||||||
call_id: call.call_id,
|
call_id: call.call_id,
|
||||||
@@ -109,6 +115,7 @@ pub fn decode_pending(value: &str) -> Result<RecoveredToolRound> {
|
|||||||
name: call.name,
|
name: call.name,
|
||||||
arguments_text: serde_json::to_string(&call.arguments)?,
|
arguments_text: serde_json::to_string(&call.arguments)?,
|
||||||
arguments: call.arguments,
|
arguments: call.arguments,
|
||||||
|
argument_error,
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
.collect::<Result<Vec<_>>>()?;
|
.collect::<Result<Vec<_>>>()?;
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ pub fn staged_tool_round(
|
|||||||
allowed_tools,
|
allowed_tools,
|
||||||
dynamic_tools,
|
dynamic_tools,
|
||||||
started_at_ms,
|
started_at_ms,
|
||||||
|
tool_calls: Some(calls),
|
||||||
}),
|
}),
|
||||||
)?)?)
|
)?)?)
|
||||||
}
|
}
|
||||||
@@ -96,6 +97,7 @@ pub fn staged_final(
|
|||||||
allowed_tools,
|
allowed_tools,
|
||||||
dynamic_tools,
|
dynamic_tools,
|
||||||
started_at_ms,
|
started_at_ms,
|
||||||
|
tool_calls: None,
|
||||||
}),
|
}),
|
||||||
)?)?)
|
)?)?)
|
||||||
}
|
}
|
||||||
@@ -105,6 +107,7 @@ pub(super) struct PendingContext<'a> {
|
|||||||
allowed_tools: &'a [String],
|
allowed_tools: &'a [String],
|
||||||
dynamic_tools: &'a HashSet<String>,
|
dynamic_tools: &'a HashSet<String>,
|
||||||
started_at_ms: u64,
|
started_at_ms: u64,
|
||||||
|
tool_calls: Option<&'a [ToolCall]>,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn wire_message(
|
pub(super) fn wire_message(
|
||||||
@@ -132,16 +135,23 @@ pub(super) fn wire_message(
|
|||||||
calls
|
calls
|
||||||
.iter()
|
.iter()
|
||||||
.map(|call| {
|
.map(|call| {
|
||||||
(
|
let mut contract = json!({
|
||||||
call.call_id.clone(),
|
"toolCallId": call.call_id,
|
||||||
json!({
|
"outerToolName": call.name,
|
||||||
"toolCallId": call.call_id,
|
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
|
||||||
"outerToolName": call.name,
|
"isDynamic": pending.dynamic_tools.contains(&call.name),
|
||||||
"toolIdentifier": tool_identifier(&call.name, pending.dynamic_tools),
|
"allowedToolNames": pending.allowed_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(),
|
.collect(),
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
|||||||
name: "Read".into(),
|
name: "Read".into(),
|
||||||
arguments_text: r#"{"path":"/a"}"#.into(),
|
arguments_text: r#"{"path":"/a"}"#.into(),
|
||||||
arguments: json!({"path":"/a"}),
|
arguments: json!({"path":"/a"}),
|
||||||
|
argument_error: Some("Read arguments are not valid JSON".into()),
|
||||||
},
|
},
|
||||||
ToolCall {
|
ToolCall {
|
||||||
index: 1,
|
index: 1,
|
||||||
@@ -38,6 +39,7 @@ fn pending_tool_round_is_one_complete_assistant_message_and_round_trips() {
|
|||||||
name: "Grep".into(),
|
name: "Grep".into(),
|
||||||
arguments_text: r#"{"pattern":"x"}"#.into(),
|
arguments_text: r#"{"pattern":"x"}"#.into(),
|
||||||
arguments: json!({"pattern":"x"}),
|
arguments: json!({"pattern":"x"}),
|
||||||
|
argument_error: None,
|
||||||
},
|
},
|
||||||
];
|
];
|
||||||
let pending = staged_tool_round(
|
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"],
|
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["toolIdentifier"],
|
||||||
"READ"
|
"READ"
|
||||||
);
|
);
|
||||||
|
assert_eq!(
|
||||||
|
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]["a"]["argumentError"],
|
||||||
|
"Read arguments are not valid JSON"
|
||||||
|
);
|
||||||
assert_eq!(wire["role"], "assistant");
|
assert_eq!(wire["role"], "assistant");
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
wire["providerOptions"]["cursor"]["pendingToolExecutionContracts"]
|
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.assistant.replay_state, Some(replay_state));
|
||||||
assert_eq!(recovered.calls.len(), 2);
|
assert_eq!(recovered.calls.len(), 2);
|
||||||
assert_eq!(recovered.calls[0].call_id, "a");
|
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");
|
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) {
|
pub fn discard_model_output(&mut self) {
|
||||||
self.text.clear();
|
self.text.clear();
|
||||||
self.thinking.clear();
|
self.thinking.clear();
|
||||||
@@ -101,6 +106,28 @@ impl StepBuffer {
|
|||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
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]
|
#[test]
|
||||||
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
|
fn interrupted_model_output_is_not_persisted_as_checkpoint_steps() {
|
||||||
let mut buffer = StepBuffer::default();
|
let mut buffer = StepBuffer::default();
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ use crate::{
|
|||||||
protocol::proto::agent::v1 as pb,
|
protocol::proto::agent::v1 as pb,
|
||||||
services::blob_sync::BlobSynchronizer,
|
services::blob_sync::BlobSynchronizer,
|
||||||
tools::{
|
tools::{
|
||||||
codec,
|
codec, compat,
|
||||||
runtime::CursorToolRuntime,
|
runtime::CursorToolRuntime,
|
||||||
stream::ToolCallStream,
|
stream::ToolCallStream,
|
||||||
tool_call_result::{ToolCompletion, ToolResultReceiver},
|
tool_call_result::{ToolCompletion, ToolResultReceiver},
|
||||||
@@ -264,6 +264,32 @@ impl ConversationOutput {
|
|||||||
streams.clear();
|
streams.clear();
|
||||||
presentation.discard_model_output();
|
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::TextStart => {}
|
||||||
RunEvent::TextEnd => {
|
RunEvent::TextEnd => {
|
||||||
if !self.context.compacting {
|
if !self.context.compacting {
|
||||||
@@ -312,6 +338,7 @@ impl ConversationOutput {
|
|||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
arguments_text: String::new(),
|
arguments_text: String::new(),
|
||||||
arguments: serde_json::Value::Null,
|
arguments: serde_json::Value::Null,
|
||||||
|
argument_error: None,
|
||||||
};
|
};
|
||||||
self.emit_model_event(
|
self.emit_model_event(
|
||||||
crate::provider::ModelEvent::ToolCallStart {
|
crate::provider::ModelEvent::ToolCallStart {
|
||||||
@@ -335,8 +362,27 @@ impl ConversationOutput {
|
|||||||
let stream = streams.get_mut(&index).ok_or_else(|| {
|
let stream = streams.get_mut(&index).ok_or_else(|| {
|
||||||
Error::Protocol(format!("missing Cursor tool stream: {index}"))
|
Error::Protocol(format!("missing Cursor tool stream: {index}"))
|
||||||
})?;
|
})?;
|
||||||
for message in stream.arguments_delta(call, &delta)? {
|
match stream.arguments_delta(call, &delta) {
|
||||||
self.handle.emit(&message)?;
|
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 } => {
|
RunEvent::ToolCallEnd { index } => {
|
||||||
@@ -349,7 +395,8 @@ impl ConversationOutput {
|
|||||||
call.arguments = if call.arguments_text.trim().is_empty() {
|
call.arguments = if call.arguments_text.trim().is_empty() {
|
||||||
serde_json::json!({})
|
serde_json::json!({})
|
||||||
} else {
|
} else {
|
||||||
serde_json::from_str(&call.arguments_text)?
|
serde_json::from_str(&call.arguments_text)
|
||||||
|
.unwrap_or_else(|_| serde_json::json!({}))
|
||||||
};
|
};
|
||||||
}
|
}
|
||||||
RunEvent::Usage(usage) => {
|
RunEvent::Usage(usage) => {
|
||||||
@@ -359,10 +406,7 @@ impl ConversationOutput {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if !self.context.compacting {
|
if !self.context.compacting {
|
||||||
context_tokens = usage
|
context_tokens = usage.context_input_tokens;
|
||||||
.input_tokens
|
|
||||||
.zip(usage.output_tokens)
|
|
||||||
.and_then(|(input, output)| input.checked_add(output));
|
|
||||||
}
|
}
|
||||||
match &mut turn_usage {
|
match &mut turn_usage {
|
||||||
Some(total) => *total += usage,
|
Some(total) => *total += usage,
|
||||||
@@ -523,6 +567,7 @@ impl ConversationOutput {
|
|||||||
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
.map_err(|_| Error::Protocol("checkpoint worker stopped".into()))?
|
||||||
{
|
{
|
||||||
Ok(checkpoint) => {
|
Ok(checkpoint) => {
|
||||||
|
context_tokens = checkpoint_context_tokens(&checkpoint);
|
||||||
compaction_checkpoint = Some(checkpoint);
|
compaction_checkpoint = Some(checkpoint);
|
||||||
state.barrier.complete(Ok(()));
|
state.barrier.complete(Ok(()));
|
||||||
}
|
}
|
||||||
@@ -1029,10 +1074,34 @@ pub(crate) fn finish_cancelled(handle: &TransportHandle) -> Result<()> {
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn checkpoint_context_tokens(checkpoint: &pb::ConversationStateStructure) -> Option<u64> {
|
||||||
|
checkpoint
|
||||||
|
.token_details
|
||||||
|
.as_ref()
|
||||||
|
.map(|details| u64::from(details.used_tokens))
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::accept_tool_completion;
|
use super::{accept_tool_completion, checkpoint_context_tokens};
|
||||||
use crate::{run::CommandResult, Error};
|
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]
|
#[test]
|
||||||
fn closing_and_ended_runs_ignore_known_tool_completions() {
|
fn closing_and_ended_runs_ignore_known_tool_completions() {
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ use crate::{
|
|||||||
protocol::proto::agent::v1 as pb,
|
protocol::proto::agent::v1 as pb,
|
||||||
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
|
services::{blob_sync::BlobSynchronizer, context_sync::RequestContextSynchronizer},
|
||||||
tools::{
|
tools::{
|
||||||
codec,
|
codec, compat,
|
||||||
runtime::CursorToolRuntime,
|
runtime::CursorToolRuntime,
|
||||||
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
|
tool_call_result::{tool_result_channel, ToolResultReceiver, ToolResultSender},
|
||||||
ClientToolEvent, ToolDispatcher,
|
ClientToolEvent, ToolDispatcher,
|
||||||
@@ -262,17 +262,18 @@ impl ConversationRuntime {
|
|||||||
.take_exec(throw.id)
|
.take_exec(throw.id)
|
||||||
.await
|
.await
|
||||||
{
|
{
|
||||||
Some(pending) => generation.results.send_error(
|
Some(pending) => generation.results.send(
|
||||||
crate::Error::Protocol(format!(
|
compat::failure_with_message(
|
||||||
"Exec {} failed: {}",
|
&pending.call,
|
||||||
pending.call.call_id, throw.error
|
format!(
|
||||||
)),
|
"Exec {} failed: {}",
|
||||||
|
pending.call.call_id, throw.error
|
||||||
|
),
|
||||||
|
),
|
||||||
),
|
),
|
||||||
None => generation.results.send_error(
|
None => tracing::warn!(
|
||||||
crate::Error::Protocol(format!(
|
id = throw.id,
|
||||||
"unknown ExecClientThrow id: {}",
|
"ignoring failure for unknown tool execution"
|
||||||
throw.id
|
|
||||||
)),
|
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -542,6 +542,7 @@ mod tests {
|
|||||||
name: "Bash".into(),
|
name: "Bash".into(),
|
||||||
arguments_text: String::new(),
|
arguments_text: String::new(),
|
||||||
arguments: json!({ "command": "ls -la" }),
|
arguments: json!({ "command": "ls -la" }),
|
||||||
|
argument_error: None,
|
||||||
};
|
};
|
||||||
let message = request(1, &call, &ExecContext::default()).unwrap();
|
let message = request(1, &call, &ExecContext::default()).unwrap();
|
||||||
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message
|
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ use crate::{
|
|||||||
cursor::{
|
cursor::{
|
||||||
protocol::{events, proto::agent::v1 as pb},
|
protocol::{events, proto::agent::v1 as pb},
|
||||||
tools::{
|
tools::{
|
||||||
edit,
|
compat, edit,
|
||||||
runtime::{CursorToolRuntime, ExecStage, PendingExec},
|
runtime::{CursorToolRuntime, ExecStage, PendingExec},
|
||||||
tool_call_result::{self as result, ToolCompletion},
|
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 {
|
let call = match pending.exec_call(message.id).await {
|
||||||
Some(call) => call,
|
Some(call) => call,
|
||||||
None if pending.completed_call(message.id).await.is_some() => {
|
None if pending.completed_call(message.id).await.is_some() => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(id = message.id, "ignoring duplicate terminal tool response");
|
||||||
"duplicate terminal ExecClientMessage id: {}",
|
return Ok(ClientExecEvent::Pending);
|
||||||
message.id
|
|
||||||
)))
|
|
||||||
}
|
}
|
||||||
None => {
|
None => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(
|
||||||
"unknown ExecClientMessage id: {}",
|
id = message.id,
|
||||||
message.id
|
"ignoring response for unknown tool execution"
|
||||||
)))
|
);
|
||||||
|
return Ok(ClientExecEvent::Pending);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let Some(wire_result) = &message.message else {
|
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 {
|
let rendered = match &entry.stage {
|
||||||
ExecStage::DynamicMcp(definition) => {
|
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(
|
Ok(Some(ToolCompletion::from_rendered(
|
||||||
&entry.call,
|
&entry.call,
|
||||||
@@ -212,10 +224,10 @@ async fn advance_edit(
|
|||||||
pb::exec_client_message::Message::ReadResult(result)
|
pb::exec_client_message::Message::ReadResult(result)
|
||||||
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
|
| pb::exec_client_message::Message::RedactedReadResult(result) => result,
|
||||||
_ => {
|
_ => {
|
||||||
return Err(Error::Protocol(format!(
|
let message = format!("expected ReadResult for edit tool {}", entry.call.name);
|
||||||
"expected ReadResult for edit tool {}",
|
return Ok(ClientExecEvent::Completed(Box::new(
|
||||||
entry.call.name
|
compat::failure_with_message(&entry.call, message),
|
||||||
)))
|
)));
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let write = match edit::after_read(&entry.call, read) {
|
let write = match edit::after_read(&entry.call, read) {
|
||||||
@@ -260,9 +272,14 @@ fn completed(
|
|||||||
pending: PendingExec,
|
pending: PendingExec,
|
||||||
result: pb::exec_client_message::Message,
|
result: pb::exec_client_message::Message,
|
||||||
) -> Result<ClientExecEvent> {
|
) -> Result<ClientExecEvent> {
|
||||||
Ok(ClientExecEvent::Completed(Box::new(result::from_exec(
|
let call = pending.call.clone();
|
||||||
pending, &result,
|
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(
|
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 {
|
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
|
let arguments = call
|
||||||
.arguments
|
.arguments
|
||||||
.as_object()
|
.as_object()
|
||||||
|
|||||||
@@ -265,6 +265,7 @@ mod tests {
|
|||||||
"old_string": old_string,
|
"old_string": old_string,
|
||||||
"new_string": "replacement",
|
"new_string": "replacement",
|
||||||
}),
|
}),
|
||||||
|
argument_error: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -103,10 +103,20 @@ impl ToolDispatcher {
|
|||||||
}
|
}
|
||||||
let message_index = first_tool_index + position;
|
let message_index = first_tool_index + position;
|
||||||
let publish_started = !state.started.contains(&call.call_id);
|
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) {
|
let edit_path = if dynamic_mcp.contains_key(&call.name) {
|
||||||
None
|
None
|
||||||
} else {
|
} 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 {
|
if let Some(path) = edit_path {
|
||||||
let next = self.edit_schedule.lock().await.start_or_defer(
|
let next = self.edit_schedule.lock().await.start_or_defer(
|
||||||
@@ -121,22 +131,28 @@ impl ToolDispatcher {
|
|||||||
let Some(next) = next else {
|
let Some(next) = next else {
|
||||||
continue;
|
continue;
|
||||||
};
|
};
|
||||||
dispatched.push(
|
let started = self
|
||||||
self.start(
|
.start(
|
||||||
&next.call,
|
&next.call,
|
||||||
next.message_index,
|
next.message_index,
|
||||||
next.publish_started,
|
next.publish_started,
|
||||||
dynamic_mcp,
|
dynamic_mcp,
|
||||||
&next.context,
|
&next.context,
|
||||||
)
|
)
|
||||||
.await?,
|
.await;
|
||||||
);
|
dispatched.push(match started {
|
||||||
|
Ok(started) => started,
|
||||||
|
Err(error) => recover_validation_failure(&next.call, error)?,
|
||||||
|
});
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
dispatched.push(
|
let started = self
|
||||||
self.start(call, message_index, publish_started, dynamic_mcp, context)
|
.start(call, message_index, publish_started, dynamic_mcp, context)
|
||||||
.await?,
|
.await;
|
||||||
);
|
dispatched.push(match started {
|
||||||
|
Ok(started) => started,
|
||||||
|
Err(error) => recover_validation_failure(call, error)?,
|
||||||
|
});
|
||||||
}
|
}
|
||||||
Ok(dispatched)
|
Ok(dispatched)
|
||||||
}
|
}
|
||||||
@@ -146,15 +162,19 @@ impl ToolDispatcher {
|
|||||||
let Some(next) = next else {
|
let Some(next) = next else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
self.start(
|
match self
|
||||||
&next.call,
|
.start(
|
||||||
next.message_index,
|
&next.call,
|
||||||
next.publish_started,
|
next.message_index,
|
||||||
&BTreeMap::new(),
|
next.publish_started,
|
||||||
&next.context,
|
&BTreeMap::new(),
|
||||||
)
|
&next.context,
|
||||||
.await
|
)
|
||||||
.map(Some)
|
.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> {
|
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 {
|
let pending = match self.runtime.take_interaction(response.id).await {
|
||||||
Some(pending) => pending,
|
Some(pending) => pending,
|
||||||
None if self.runtime.completed_call(response.id).await.is_some() => {
|
None if self.runtime.completed_call(response.id).await.is_some() => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(
|
||||||
"duplicate terminal InteractionResponse id: {}",
|
id = response.id,
|
||||||
response.id
|
"ignoring duplicate terminal interaction response"
|
||||||
)));
|
);
|
||||||
|
return Ok(ClientToolEvent::Pending);
|
||||||
}
|
}
|
||||||
None => {
|
None => {
|
||||||
return Err(Error::Protocol(format!(
|
tracing::warn!(
|
||||||
"unknown InteractionResponse id: {}",
|
id = response.id,
|
||||||
response.id
|
"ignoring response for unknown interaction"
|
||||||
)));
|
);
|
||||||
|
return Ok(ClientToolEvent::Pending);
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
Ok(
|
let call = pending.call.clone();
|
||||||
match tool_call_dispatch::resume_interaction(
|
let continuation = match tool_call_dispatch::resume_interaction(
|
||||||
&self.results,
|
&self.results,
|
||||||
&self.search,
|
&self.search,
|
||||||
&self.fetch,
|
&self.fetch,
|
||||||
pending,
|
pending,
|
||||||
response,
|
response,
|
||||||
)
|
|
||||||
.await?
|
|
||||||
{
|
|
||||||
tool_call_dispatch::InteractionContinuation::Completed(completion) => {
|
|
||||||
ClientToolEvent::Completed(completion)
|
|
||||||
}
|
|
||||||
tool_call_dispatch::InteractionContinuation::Pending => ClientToolEvent::Pending,
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
.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(),
|
name: "Task".into(),
|
||||||
arguments_text: arguments_text.into(),
|
arguments_text: arguments_text.into(),
|
||||||
arguments: serde_json::Value::Null,
|
arguments: serde_json::Value::Null,
|
||||||
|
argument_error: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -406,6 +406,7 @@ mod tests {
|
|||||||
name: "WebFetch".into(),
|
name: "WebFetch".into(),
|
||||||
arguments_text: r#"{"url":"https://example.com"}"#.into(),
|
arguments_text: r#"{"url":"https://example.com"}"#.into(),
|
||||||
arguments: json!({"url": "https://example.com"}),
|
arguments: json!({"url": "https://example.com"}),
|
||||||
|
argument_error: None,
|
||||||
},
|
},
|
||||||
started_at_ms: 1,
|
started_at_ms: 1,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ mod usage {
|
|||||||
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||||
pub struct Usage {
|
pub struct Usage {
|
||||||
pub input_tokens: Option<u64>,
|
pub input_tokens: Option<u64>,
|
||||||
|
pub context_input_tokens: Option<u64>,
|
||||||
pub output_tokens: Option<u64>,
|
pub output_tokens: Option<u64>,
|
||||||
pub total_tokens: Option<u64>,
|
pub total_tokens: Option<u64>,
|
||||||
pub cache_read_tokens: Option<u64>,
|
pub cache_read_tokens: Option<u64>,
|
||||||
@@ -19,6 +20,7 @@ mod usage {
|
|||||||
impl AddAssign for Usage {
|
impl AddAssign for Usage {
|
||||||
fn add_assign(&mut self, rhs: Self) {
|
fn add_assign(&mut self, rhs: Self) {
|
||||||
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
|
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.output_tokens = sum(self.output_tokens, rhs.output_tokens);
|
||||||
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
||||||
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_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 name: String,
|
||||||
pub arguments_text: String,
|
pub arguments_text: String,
|
||||||
pub arguments: Value,
|
pub arguments: Value,
|
||||||
|
pub argument_error: Option<String>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||||
|
|||||||
@@ -147,8 +147,10 @@ pub fn model_event(value: &serde_json::Value) -> Result<ModelEvent> {
|
|||||||
.get("usage")
|
.get("usage")
|
||||||
.ok_or_else(|| Error::Protocol("plugin usage event requires usage".into()))?;
|
.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 tokens = |name: &str| usage.get(name).and_then(serde_json::Value::as_u64);
|
||||||
|
let input_tokens = tokens("inputTokens");
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: tokens("inputTokens"),
|
input_tokens,
|
||||||
|
context_input_tokens: input_tokens,
|
||||||
output_tokens: tokens("outputTokens"),
|
output_tokens: tokens("outputTokens"),
|
||||||
total_tokens: tokens("totalTokens"),
|
total_tokens: tokens("totalTokens"),
|
||||||
cache_read_tokens: tokens("cacheReadTokens"),
|
cache_read_tokens: tokens("cacheReadTokens"),
|
||||||
@@ -285,6 +287,7 @@ mod tests {
|
|||||||
usage,
|
usage,
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: Some(10),
|
input_tokens: Some(10),
|
||||||
|
context_input_tokens: Some(10),
|
||||||
output_tokens: Some(2),
|
output_tokens: Some(2),
|
||||||
total_tokens: None,
|
total_tokens: None,
|
||||||
cache_read_tokens: Some(4),
|
cache_read_tokens: Some(4),
|
||||||
|
|||||||
@@ -12,9 +12,9 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
|
attempt::{send_once, Attempt},
|
||||||
map_sse_error, merge_extra_params, provider_event_error,
|
map_sse_error, merge_extra_params, provider_event_error,
|
||||||
recorder::recorded_headers,
|
recorder::recorded_headers,
|
||||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
|
||||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -90,17 +90,14 @@ impl Provider for AnthropicProvider {
|
|||||||
if let Some(recorder) = &recorder {
|
if let Some(recorder) = &recorder {
|
||||||
recorder.request(request_headers.clone(), &body).await?;
|
recorder.request(request_headers.clone(), &body).await?;
|
||||||
}
|
}
|
||||||
let attempt = send_with_retry(
|
let attempt = send_once(
|
||||||
"Anthropic",
|
"Anthropic",
|
||||||
|| client.post(&config.request_url)
|
|| client.post(&config.request_url)
|
||||||
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
|
.header("x-api-key", &config.api_key).header("anthropic-version", "2023-06-01")
|
||||||
.headers(config.custom_headers.clone())
|
.headers(config.custom_headers.clone())
|
||||||
.json(&body),
|
.json(&body),
|
||||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
|
||||||
&cancellation,
|
&cancellation,
|
||||||
recorder.as_ref(),
|
recorder.as_ref(),
|
||||||
request_headers,
|
|
||||||
&body,
|
|
||||||
).await?;
|
).await?;
|
||||||
let Attempt::Response(response) = attempt else { return };
|
let Attempt::Response(response) = attempt else { return };
|
||||||
yield ModelEvent::Start { model_call_id: call_id };
|
yield ModelEvent::Start { model_call_id: call_id };
|
||||||
@@ -321,6 +318,7 @@ fn apply_model(body: &mut Value, model: &crate::model::ModelSpec) -> Result<()>
|
|||||||
|
|
||||||
fn merge_usage(total: &mut Usage, update: Usage) {
|
fn merge_usage(total: &mut Usage, update: Usage) {
|
||||||
merge_usage_field(&mut total.input_tokens, update.input_tokens);
|
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.output_tokens, update.output_tokens);
|
||||||
merge_usage_field(&mut total.cache_read_tokens, update.cache_read_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);
|
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 {
|
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 {
|
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),
|
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
|
||||||
total_tokens: value.get("total_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_read_tokens,
|
||||||
cache_write_tokens: value
|
cache_write_tokens,
|
||||||
.get("cache_creation_input_tokens")
|
|
||||||
.and_then(Value::as_u64),
|
|
||||||
reasoning_tokens: None,
|
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.
|
//! Defines the provider interface and exports provider implementations.
|
||||||
mod anthropic;
|
mod anthropic;
|
||||||
|
mod attempt;
|
||||||
mod event;
|
mod event;
|
||||||
mod normalize;
|
mod normalize;
|
||||||
mod openai_chat;
|
mod openai_chat;
|
||||||
mod openai_responses;
|
mod openai_responses;
|
||||||
mod recorder;
|
mod recorder;
|
||||||
mod retry;
|
|
||||||
mod router;
|
mod router;
|
||||||
|
|
||||||
use std::pin::Pin;
|
use std::pin::Pin;
|
||||||
|
|||||||
@@ -17,10 +17,10 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
apply_body_allowlist, apply_openai_prompt_cache_key,
|
||||||
provider_event_error,
|
attempt::{send_once, Attempt},
|
||||||
|
map_sse_error, merge_extra_params, provider_event_error,
|
||||||
recorder::recorded_headers,
|
recorder::recorded_headers,
|
||||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
|
||||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -92,15 +92,12 @@ impl Provider for OpenAiChatProvider {
|
|||||||
if let Some(recorder) = &recorder {
|
if let Some(recorder) = &recorder {
|
||||||
recorder.request(request_headers.clone(), &body).await?;
|
recorder.request(request_headers.clone(), &body).await?;
|
||||||
}
|
}
|
||||||
let attempt = send_with_retry(
|
let attempt = send_once(
|
||||||
"OpenAI Chat",
|
"OpenAI Chat",
|
||||||
|| client.post(&config.request_url)
|
|| client.post(&config.request_url)
|
||||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
|
||||||
&cancellation,
|
&cancellation,
|
||||||
recorder.as_ref(),
|
recorder.as_ref(),
|
||||||
request_headers,
|
|
||||||
&body,
|
|
||||||
).await?;
|
).await?;
|
||||||
let Attempt::Response(response) = attempt else { return };
|
let Attempt::Response(response) = attempt else { return };
|
||||||
yield ModelEvent::Start { model_call_id: call_id };
|
yield ModelEvent::Start { model_call_id: call_id };
|
||||||
@@ -421,8 +418,10 @@ fn merge_chat_fragment(target: &mut String, fragment: &str) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn openai_usage(value: &Value) -> Usage {
|
pub(crate) fn openai_usage(value: &Value) -> Usage {
|
||||||
|
let input_tokens = value.get("prompt_tokens").and_then(Value::as_u64);
|
||||||
Usage {
|
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),
|
output_tokens: value.get("completion_tokens").and_then(Value::as_u64),
|
||||||
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
||||||
cache_read_tokens: value
|
cache_read_tokens: value
|
||||||
|
|||||||
@@ -14,10 +14,10 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
apply_body_allowlist, apply_openai_prompt_cache_key, map_sse_error, merge_extra_params,
|
apply_body_allowlist, apply_openai_prompt_cache_key,
|
||||||
provider_event_error,
|
attempt::{send_once, Attempt},
|
||||||
|
map_sse_error, merge_extra_params, provider_event_error,
|
||||||
recorder::recorded_headers,
|
recorder::recorded_headers,
|
||||||
retry::{send_with_retry, Attempt, RetryPolicy},
|
|
||||||
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
CallRecorder, FinishReason, ModelEvent, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -89,15 +89,12 @@ impl Provider for OpenAiResponsesProvider {
|
|||||||
if let Some(recorder) = &recorder {
|
if let Some(recorder) = &recorder {
|
||||||
recorder.request(request_headers.clone(), &body).await?;
|
recorder.request(request_headers.clone(), &body).await?;
|
||||||
}
|
}
|
||||||
let attempt = send_with_retry(
|
let attempt = send_once(
|
||||||
"OpenAI Responses",
|
"OpenAI Responses",
|
||||||
|| client.post(&config.request_url)
|
|| client.post(&config.request_url)
|
||||||
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
.bearer_auth(&config.api_key).headers(config.custom_headers.clone()).json(&body),
|
||||||
RetryPolicy { retries: config.retry_count, ..RetryPolicy::default() },
|
|
||||||
&cancellation,
|
&cancellation,
|
||||||
recorder.as_ref(),
|
recorder.as_ref(),
|
||||||
request_headers,
|
|
||||||
&body,
|
|
||||||
).await?;
|
).await?;
|
||||||
let Attempt::Response(response) = attempt else { return };
|
let Attempt::Response(response) = attempt else { return };
|
||||||
yield ModelEvent::Start { model_call_id: call_id };
|
yield ModelEvent::Start { model_call_id: call_id };
|
||||||
@@ -505,8 +502,10 @@ fn required_u64(value: &Value, name: &str) -> Result<u64> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn responses_usage(value: &Value) -> Usage {
|
fn responses_usage(value: &Value) -> Usage {
|
||||||
|
let input_tokens = value.get("input_tokens").and_then(Value::as_u64);
|
||||||
Usage {
|
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),
|
output_tokens: value.get("output_tokens").and_then(Value::as_u64),
|
||||||
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
total_tokens: value.get("total_tokens").and_then(Value::as_u64),
|
||||||
cache_read_tokens: value
|
cache_read_tokens: value
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
//! Records provider requests, responses, usage, and timing.
|
//! Records provider requests, responses, usage, and timing.
|
||||||
use std::{
|
use std::{
|
||||||
sync::{
|
sync::{
|
||||||
atomic::{AtomicBool, AtomicI64, AtomicU32, AtomicU64, Ordering},
|
atomic::{AtomicBool, AtomicI64, AtomicU64, Ordering},
|
||||||
Arc,
|
Arc,
|
||||||
},
|
},
|
||||||
time::Instant,
|
time::Instant,
|
||||||
@@ -71,7 +71,6 @@ struct Inner {
|
|||||||
base_call: NewLlmCall,
|
base_call: NewLlmCall,
|
||||||
detailed: bool,
|
detailed: bool,
|
||||||
attempt: Mutex<AttemptState>,
|
attempt: Mutex<AttemptState>,
|
||||||
next_attempt: AtomicU32,
|
|
||||||
next_generation: AtomicU64,
|
next_generation: AtomicU64,
|
||||||
finished: AtomicBool,
|
finished: AtomicBool,
|
||||||
}
|
}
|
||||||
@@ -120,7 +119,6 @@ impl CallRecorder {
|
|||||||
base_call: call.clone(),
|
base_call: call.clone(),
|
||||||
detailed: call.detailed,
|
detailed: call.detailed,
|
||||||
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
|
attempt: Mutex::new(AttemptState::new(call.call_id.clone())),
|
||||||
next_attempt: AtomicU32::new(0),
|
|
||||||
next_generation: AtomicU64::new(0),
|
next_generation: AtomicU64::new(0),
|
||||||
finished: AtomicBool::new(false),
|
finished: AtomicBool::new(false),
|
||||||
}),
|
}),
|
||||||
@@ -284,34 +282,6 @@ impl CallRecorder {
|
|||||||
self.finish("cancelled", None, None, None).await
|
self.finish("cancelled", None, None, None).await
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn retry(
|
|
||||||
&self,
|
|
||||||
error: &crate::Error,
|
|
||||||
headers: serde_json::Value,
|
|
||||||
body: &serde_json::Value,
|
|
||||||
) -> Result<()> {
|
|
||||||
self.failed(error).await?;
|
|
||||||
|
|
||||||
let attempt_number = self.inner.next_attempt.fetch_add(1, Ordering::Relaxed) + 1;
|
|
||||||
let mut call = self.inner.base_call.clone();
|
|
||||||
call.call_id = format!("{}:retry-{attempt_number}", self.inner.base_call.call_id);
|
|
||||||
{
|
|
||||||
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(
|
async fn finish(
|
||||||
&self,
|
&self,
|
||||||
status: &str,
|
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,
|
OpenAiChatProvider, OpenAiResponsesProvider, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
const BUILTIN_PROVIDER_RETRIES: u32 = 5;
|
|
||||||
|
|
||||||
pub struct ProviderRouter {
|
pub struct ProviderRouter {
|
||||||
store: Store,
|
store: Store,
|
||||||
plugins: PluginRegistry,
|
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() },
|
custom_headers: if model.custom_headers_enabled { custom_headers(&model.custom_headers)? } else { reqwest::header::HeaderMap::new() },
|
||||||
max_output_tokens: model.max_output_tokens(),
|
max_output_tokens: model.max_output_tokens(),
|
||||||
request_timeout,
|
request_timeout,
|
||||||
retry_count: BUILTIN_PROVIDER_RETRIES,
|
|
||||||
allowed_body_fields: None,
|
allowed_body_fields: None,
|
||||||
};
|
};
|
||||||
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
|
let client = crate::network::client_builder(&store).await?.timeout(request_timeout).build()?;
|
||||||
|
|||||||
+203
-107
@@ -14,8 +14,10 @@ use crate::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
consume_model_cycle, CommitBarrier, CommitCause, MessagesCommitted, ModelCycleFailure,
|
consume_model_cycle,
|
||||||
RunCommand, RunEvent, RunFailure, RunOutcome, RunPort,
|
model_retry::{should_retry, MODEL_RETRY_DELAY},
|
||||||
|
CommitBarrier, CommitCause, MessagesCommitted, RunCommand, RunEvent, RunFailure, RunOutcome,
|
||||||
|
RunPort,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct RunEngine {
|
pub struct RunEngine {
|
||||||
@@ -208,119 +210,213 @@ impl RunEngine {
|
|||||||
model: prepared.model.clone(),
|
model: prepared.model.clone(),
|
||||||
history,
|
history,
|
||||||
};
|
};
|
||||||
let invocation = crate::model::ModelInvocation {
|
let mut retries = 0_u32;
|
||||||
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 pending_insertions = Vec::new();
|
let mut pending_insertions = Vec::new();
|
||||||
let cycle = loop {
|
let cycle = 'attempt: loop {
|
||||||
tokio::select! {
|
let call_id = if retries == 0 {
|
||||||
biased;
|
format!("{}:{provider_call_index}", prepared.run_id)
|
||||||
command = client.commands.recv() => {
|
} else {
|
||||||
let interruption = match command {
|
format!("{}:{provider_call_index}:retry-{retries}", prepared.run_id)
|
||||||
Some(RunCommand::InsertMessages(insertion)) => {
|
};
|
||||||
pending_insertions.push(insertion);
|
let invocation = crate::model::ModelInvocation {
|
||||||
continue;
|
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,
|
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
|
||||||
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);
|
return (client_failure(), usage);
|
||||||
}
|
}
|
||||||
};
|
checkpoint = match super::messages::append_batches(
|
||||||
cycle_cancellation.cancel();
|
&self.store,
|
||||||
let interrupted = cycle.await;
|
prepared,
|
||||||
match interrupted {
|
client,
|
||||||
Ok(cycle) => {
|
cancellation,
|
||||||
if let Some(cycle_usage) = cycle.usage {
|
checkpoint,
|
||||||
accumulate_usage(&mut usage, cycle_usage);
|
std::mem::take(&mut pending_insertions),
|
||||||
}
|
)
|
||||||
}
|
.await
|
||||||
Err(failure) => {
|
{
|
||||||
if let Some(cycle_usage) = failure.usage {
|
Ok((checkpoint, _)) => checkpoint,
|
||||||
accumulate_usage(&mut usage, cycle_usage);
|
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);
|
return (client_failure(), usage);
|
||||||
}
|
}
|
||||||
checkpoint = match super::messages::append_batches(
|
|
||||||
&self.store,
|
let delay = tokio::time::sleep(MODEL_RETRY_DELAY);
|
||||||
prepared,
|
tokio::pin!(delay);
|
||||||
client,
|
loop {
|
||||||
cancellation,
|
tokio::select! {
|
||||||
checkpoint,
|
biased;
|
||||||
std::mem::take(&mut pending_insertions),
|
command = client.commands.recv() => {
|
||||||
)
|
let interruption = match command {
|
||||||
.await
|
Some(RunCommand::InsertMessages(insertion)) => {
|
||||||
{
|
pending_insertions.push(insertion);
|
||||||
Ok((checkpoint, _)) => checkpoint,
|
continue;
|
||||||
Err(outcome) => return (outcome, usage),
|
}
|
||||||
};
|
Some(RunCommand::BreakMessages(messages)) => messages,
|
||||||
checkpoint = match super::messages::append_batches(
|
Some(RunCommand::Cancel) => {
|
||||||
&self.store,
|
let _ = emit(client, RunEvent::CycleInterrupted).await;
|
||||||
prepared,
|
return (RunOutcome::Cancelled, usage);
|
||||||
client,
|
}
|
||||||
cancellation,
|
Some(RunCommand::ToolResult(_)) => {
|
||||||
checkpoint,
|
return (
|
||||||
vec![interruption],
|
RunOutcome::Failed(RunFailure::Protocol(
|
||||||
)
|
"received a tool result while waiting to retry the model".into(),
|
||||||
.await
|
)),
|
||||||
{
|
usage,
|
||||||
Ok((checkpoint, _)) => checkpoint,
|
);
|
||||||
Err(outcome) => return (outcome, usage),
|
}
|
||||||
};
|
None => return (client_failure(), usage),
|
||||||
continue 'model;
|
};
|
||||||
},
|
if emit(client, RunEvent::CycleInterrupted).await.is_err() {
|
||||||
result = &mut cycle => break result,
|
return (client_failure(), usage);
|
||||||
}
|
}
|
||||||
};
|
checkpoint = match super::messages::append_batches(
|
||||||
let cycle = match cycle {
|
&self.store,
|
||||||
Ok(cycle) => cycle,
|
prepared,
|
||||||
Err(ModelCycleFailure {
|
client,
|
||||||
failure,
|
cancellation,
|
||||||
usage: cycle_usage,
|
checkpoint,
|
||||||
..
|
std::mem::take(&mut pending_insertions),
|
||||||
}) => {
|
)
|
||||||
if let Some(cycle_usage) = cycle_usage {
|
.await
|
||||||
accumulate_usage(&mut usage, cycle_usage);
|
{
|
||||||
|
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 {
|
if let Some(cycle_usage) = cycle.usage {
|
||||||
|
|||||||
@@ -104,6 +104,10 @@ pub enum RunEvent {
|
|||||||
AutoCompactionStarted,
|
AutoCompactionStarted,
|
||||||
AutoCompactionCompleted,
|
AutoCompactionCompleted,
|
||||||
CycleInterrupted,
|
CycleInterrupted,
|
||||||
|
ModelAttemptFailed {
|
||||||
|
attempt: u32,
|
||||||
|
message: String,
|
||||||
|
},
|
||||||
TextStart,
|
TextStart,
|
||||||
TextDelta(String),
|
TextDelta(String),
|
||||||
TextEnd,
|
TextEnd,
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ mod event;
|
|||||||
mod handle;
|
mod handle;
|
||||||
mod messages;
|
mod messages;
|
||||||
mod model_cycle;
|
mod model_cycle;
|
||||||
|
mod model_retry;
|
||||||
mod port;
|
mod port;
|
||||||
mod tool_round;
|
mod tool_round;
|
||||||
|
|
||||||
|
|||||||
+100
-15
@@ -29,6 +29,7 @@ pub struct ModelCycleFailure {
|
|||||||
pub partial_text: String,
|
pub partial_text: String,
|
||||||
pub partial_reasoning: String,
|
pub partial_reasoning: String,
|
||||||
pub usage: Option<Usage>,
|
pub usage: Option<Usage>,
|
||||||
|
pub retryable: bool,
|
||||||
}
|
}
|
||||||
|
|
||||||
struct OpenTool {
|
struct OpenTool {
|
||||||
@@ -186,6 +187,7 @@ pub async fn consume_model_cycle(
|
|||||||
name: name.clone(),
|
name: name.clone(),
|
||||||
arguments_text: String::new(),
|
arguments_text: String::new(),
|
||||||
arguments: serde_json::Value::Null,
|
arguments: serde_json::Value::Null,
|
||||||
|
argument_error: None,
|
||||||
},
|
},
|
||||||
ended: false,
|
ended: false,
|
||||||
});
|
});
|
||||||
@@ -218,13 +220,26 @@ pub async fn consume_model_cycle(
|
|||||||
serde_json::from_str(&tool.call.arguments_text)
|
serde_json::from_str(&tool.call.arguments_text)
|
||||||
};
|
};
|
||||||
match arguments {
|
match arguments {
|
||||||
Ok(arguments) => {
|
Ok(arguments) if arguments.is_object() => {
|
||||||
tool.call.arguments = arguments;
|
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"),
|
Some(_) => Err("provider emitted duplicate ToolCallEnd"),
|
||||||
None => Err("provider ended an unknown tool index"),
|
None => Err("provider ended an unknown tool index"),
|
||||||
@@ -239,6 +254,13 @@ pub async fn consume_model_cycle(
|
|||||||
ModelEvent::Usage(value) => {
|
ModelEvent::Usage(value) => {
|
||||||
if usage.replace(value).is_some() {
|
if usage.replace(value).is_some() {
|
||||||
Err("provider emitted duplicate Usage")
|
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 {
|
} else {
|
||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
@@ -280,7 +302,7 @@ pub async fn consume_model_cycle(
|
|||||||
.map(|tool| tool.call)
|
.map(|tool| tool.call)
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
if finish_reason == FinishReason::Length {
|
if finish_reason == FinishReason::Length {
|
||||||
return Err(failure(
|
return Err(terminal_failure(
|
||||||
RunFailure::Provider("model stopped before completing the response".into()),
|
RunFailure::Provider("model stopped before completing the response".into()),
|
||||||
text,
|
text,
|
||||||
reasoning,
|
reasoning,
|
||||||
@@ -296,16 +318,6 @@ pub async fn consume_model_cycle(
|
|||||||
usage,
|
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(|| {
|
let model_call_id = model_call_id.ok_or_else(|| {
|
||||||
failure(
|
failure(
|
||||||
RunFailure::Protocol("provider completed without Start".into()),
|
RunFailure::Protocol("provider completed without Start".into()),
|
||||||
@@ -360,10 +372,83 @@ fn failure(
|
|||||||
partial_reasoning: String,
|
partial_reasoning: String,
|
||||||
usage: Option<Usage>,
|
usage: Option<Usage>,
|
||||||
) -> ModelCycleFailure {
|
) -> ModelCycleFailure {
|
||||||
|
let retryable = matches!(failure, RunFailure::Protocol(_) | RunFailure::Provider(_));
|
||||||
ModelCycleFailure {
|
ModelCycleFailure {
|
||||||
failure,
|
failure,
|
||||||
partial_text,
|
partial_text,
|
||||||
partial_reasoning,
|
partial_reasoning,
|
||||||
usage,
|
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
|
.await
|
||||||
.unwrap();
|
.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!(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!(checkpoint_table_exists, 1);
|
||||||
|
assert_eq!(argument_error_column_exists, 1);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -93,22 +93,16 @@ impl Store {
|
|||||||
for call in calls {
|
for call in calls {
|
||||||
sqlx::query(
|
sqlx::query(
|
||||||
"INSERT INTO tool_round_calls
|
"INSERT INTO tool_round_calls
|
||||||
(round_id, call_index, call_id, model_call_id, name, arguments_json, status)
|
(round_id, call_index, call_id, model_call_id, name, arguments_json, argument_error, status)
|
||||||
VALUES (?, ?, ?, ?, ?, ?, 'pending')",
|
VALUES (?, ?, ?, ?, ?, ?, ?, 'pending')",
|
||||||
)
|
)
|
||||||
.bind(round_id.as_str())
|
.bind(round_id.as_str())
|
||||||
.bind(call.index as i64)
|
.bind(call.index as i64)
|
||||||
.bind(&call.call_id)
|
.bind(&call.call_id)
|
||||||
.bind(&call.model_call_id)
|
.bind(&call.model_call_id)
|
||||||
.bind(&call.name)
|
.bind(&call.name)
|
||||||
// A no-argument tool call streams no argument text; persist it as an
|
.bind(serde_json::to_string(&call.arguments)?)
|
||||||
// empty object so the `arguments_json` column always holds valid JSON
|
.bind(call.argument_error.as_deref())
|
||||||
// and can be re-parsed on load.
|
|
||||||
.bind(if call.arguments_text.trim().is_empty() {
|
|
||||||
"{}"
|
|
||||||
} else {
|
|
||||||
call.arguments_text.as_str()
|
|
||||||
})
|
|
||||||
.execute(&mut *tx)
|
.execute(&mut *tx)
|
||||||
.await?;
|
.await?;
|
||||||
}
|
}
|
||||||
@@ -276,7 +270,7 @@ impl Store {
|
|||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let rows = sqlx::query(
|
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",
|
FROM tool_round_calls WHERE round_id = ? ORDER BY call_index",
|
||||||
)
|
)
|
||||||
.bind(round_id.as_str())
|
.bind(round_id.as_str())
|
||||||
@@ -287,7 +281,7 @@ impl Store {
|
|||||||
for row in rows {
|
for row in rows {
|
||||||
let arguments_text: String = row.get(4);
|
let arguments_text: String = row.get(4);
|
||||||
let call_id: String = row.get(1);
|
let call_id: String = row.get(1);
|
||||||
if row.get::<&str, _>(5) == "completed" {
|
if row.get::<&str, _>(6) == "completed" {
|
||||||
completed.push(call_id.clone());
|
completed.push(call_id.clone());
|
||||||
}
|
}
|
||||||
calls.push(ToolCall {
|
calls.push(ToolCall {
|
||||||
@@ -297,6 +291,7 @@ impl Store {
|
|||||||
name: row.get(3),
|
name: row.get(3),
|
||||||
arguments: serde_json::from_str(&arguments_text)?,
|
arguments: serde_json::from_str(&arguments_text)?,
|
||||||
arguments_text,
|
arguments_text,
|
||||||
|
argument_error: row.get(5),
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
Ok(Some(ToolRoundSnapshot {
|
Ok(Some(ToolRoundSnapshot {
|
||||||
|
|||||||
@@ -62,6 +62,7 @@ async fn summarize_replaces_model_history_and_preserves_cursor_history() {
|
|||||||
ModelEvent::TextEnd,
|
ModelEvent::TextEnd,
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: Some(4_012),
|
input_tokens: Some(4_012),
|
||||||
|
context_input_tokens: Some(4_012),
|
||||||
output_tokens: Some(9),
|
output_tokens: Some(9),
|
||||||
total_tokens: Some(4_021),
|
total_tokens: Some(4_021),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
@@ -442,6 +443,7 @@ fn text_response(text: &str, input: u64, output: u64) -> Vec<ModelEvent> {
|
|||||||
ModelEvent::TextEnd,
|
ModelEvent::TextEnd,
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: Some(input),
|
input_tokens: Some(input),
|
||||||
|
context_input_tokens: Some(input),
|
||||||
output_tokens: Some(output),
|
output_tokens: Some(output),
|
||||||
total_tokens: Some(input + output),
|
total_tokens: Some(input + output),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
|
|||||||
+184
-50
@@ -51,10 +51,35 @@ async fn abort_command_cancels_the_run_and_closes_output() {
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[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 (_directory, store) = fixtures::temp_store().await;
|
||||||
let provider = fake_provider::FakeProvider::default();
|
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(
|
let assets = PromptAssets::load(
|
||||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||||
.join("prompt/cursor")
|
.join("prompt/cursor")
|
||||||
@@ -63,7 +88,106 @@ async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_e
|
|||||||
.unwrap();
|
.unwrap();
|
||||||
let registry = TransportRegistry::new(
|
let registry = TransportRegistry::new(
|
||||||
store.clone(),
|
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),
|
PromptCompiler::new(assets),
|
||||||
);
|
);
|
||||||
let handle = registry.get_or_create("failed-request").await.unwrap();
|
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
|
None
|
||||||
);
|
);
|
||||||
|
|
||||||
|
assert_eq!(provider.requests().len(), 1, "Length must not retry");
|
||||||
let messages = store
|
let messages = store
|
||||||
.load_current_messages(&cursor_server::model::ConversationId::new(
|
.load_current_messages(&cursor_server::model::ConversationId::new(
|
||||||
"failed-conversation",
|
"failed-conversation",
|
||||||
@@ -164,7 +289,7 @@ async fn provider_failure_keeps_the_initial_checkpoint_then_returns_structured_e
|
|||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[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 (_directory, store) = fixtures::temp_store().await;
|
||||||
let provider = fake_provider::FakeProvider::default();
|
let provider = fake_provider::FakeProvider::default();
|
||||||
provider.push(vec![
|
provider.push(vec![
|
||||||
@@ -183,6 +308,15 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
|
|||||||
ModelEvent::ToolCallEnd { index: 0 },
|
ModelEvent::ToolCallEnd { index: 0 },
|
||||||
ModelEvent::Done(FinishReason::ToolUse),
|
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(
|
let assets = PromptAssets::load(
|
||||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||||
.join("prompt/cursor")
|
.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 append_seqno = 1;
|
||||||
let mut saw_turn_ended = false;
|
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())
|
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
|
||||||
.await
|
.await
|
||||||
.unwrap()
|
.unwrap()
|
||||||
@@ -231,24 +365,42 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
|
|||||||
append_seqno += 1;
|
append_seqno += 1;
|
||||||
}
|
}
|
||||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||||
// An unknown numeric bridge id is a runtime protocol error.
|
// Unknown bridge ids are ignored; the valid response still completes the tool.
|
||||||
handle
|
for message in [
|
||||||
.command(TransportCommand::Append {
|
pb::ExecClientMessage {
|
||||||
seqno: append_seqno,
|
id: exec.id + 1_000,
|
||||||
message: Box::new(pb::AgentClientMessage {
|
exec_id: String::new(),
|
||||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
message: None,
|
||||||
pb::ExecClientMessage {
|
..Default::default()
|
||||||
id: exec.id + 1_000,
|
},
|
||||||
exec_id: String::new(),
|
pb::ExecClientMessage {
|
||||||
message: None,
|
id: exec.id,
|
||||||
|
exec_id: String::new(),
|
||||||
|
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||||
|
pb::ReadResult {
|
||||||
|
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
|
||||||
|
path: "/tmp/a".into(),
|
||||||
|
output: Some(pb::read_success::Output::Content("value".into())),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
},
|
})),
|
||||||
)),
|
},
|
||||||
}),
|
)),
|
||||||
})
|
..Default::default()
|
||||||
.await
|
},
|
||||||
.unwrap();
|
] {
|
||||||
append_seqno += 1;
|
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)) => {
|
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||||
if matches!(
|
if matches!(
|
||||||
@@ -266,12 +418,8 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
assert!(!saw_turn_ended);
|
assert!(saw_turn_ended);
|
||||||
assert_eq!(error_json["error"]["code"], "invalid_argument");
|
assert_eq!(end_stream, serde_json::json!({}));
|
||||||
assert_eq!(
|
|
||||||
error_json["error"]["message"],
|
|
||||||
"unknown ExecClientMessage id: 1001"
|
|
||||||
);
|
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
||||||
.await
|
.await
|
||||||
@@ -279,28 +427,14 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
|
|||||||
None
|
None
|
||||||
);
|
);
|
||||||
|
|
||||||
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(1);
|
let (status, failure_summary): (String, Option<String>) =
|
||||||
let (status, failure_summary) = loop {
|
sqlx::query_as("SELECT status, failure_summary FROM runs WHERE cursor_request_id = ?")
|
||||||
let row: (String, Option<String>) =
|
.bind("protocol-failed-request")
|
||||||
sqlx::query_as("SELECT status, failure_summary FROM runs WHERE cursor_request_id = ?")
|
.fetch_one(store.pool())
|
||||||
.bind("protocol-failed-request")
|
.await
|
||||||
.fetch_one(store.pool())
|
.unwrap();
|
||||||
.await
|
assert_eq!(status, "completed");
|
||||||
.unwrap();
|
assert_eq!(failure_summary, None);
|
||||||
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")
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
@@ -1580,6 +1580,7 @@ fn text_response(text: &str) -> Vec<ModelEvent> {
|
|||||||
ModelEvent::TextEnd,
|
ModelEvent::TextEnd,
|
||||||
ModelEvent::Usage(Usage {
|
ModelEvent::Usage(Usage {
|
||||||
input_tokens: Some(1),
|
input_tokens: Some(1),
|
||||||
|
context_input_tokens: Some(1),
|
||||||
output_tokens: Some(1),
|
output_tokens: Some(1),
|
||||||
total_tokens: Some(2),
|
total_tokens: Some(2),
|
||||||
..Default::default()
|
..Default::default()
|
||||||
|
|||||||
@@ -31,10 +31,13 @@ pub struct FakeProvider {
|
|||||||
|
|
||||||
impl FakeProvider {
|
impl FakeProvider {
|
||||||
pub fn push(&self, events: Vec<ModelEvent>) {
|
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
|
self.responses
|
||||||
.lock()
|
.lock()
|
||||||
.unwrap()
|
.unwrap()
|
||||||
.push_back(FakeResponse::Events(events.into_iter().map(Ok).collect()));
|
.push_back(FakeResponse::Events(events));
|
||||||
}
|
}
|
||||||
pub fn push_error(&self, error: Error) {
|
pub fn push_error(&self, error: Error) {
|
||||||
self.responses
|
self.responses
|
||||||
|
|||||||
+96
-20
@@ -25,9 +25,11 @@ use cursor_server::{
|
|||||||
OPENAI_CHAT_ENDPOINT,
|
OPENAI_CHAT_ENDPOINT,
|
||||||
},
|
},
|
||||||
provider::{FinishReason, ModelEvent},
|
provider::{FinishReason, ModelEvent},
|
||||||
|
run::consume_model_cycle,
|
||||||
};
|
};
|
||||||
use prost::Message;
|
use prost::Message;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
use tokio_util::sync::CancellationToken;
|
||||||
|
|
||||||
fn call(id: &str, name: &str) -> ToolCall {
|
fn call(id: &str, name: &str) -> ToolCall {
|
||||||
ToolCall {
|
ToolCall {
|
||||||
@@ -37,6 +39,7 @@ fn call(id: &str, name: &str) -> ToolCall {
|
|||||||
name: name.into(),
|
name: name.into(),
|
||||||
arguments_text: "{}".into(),
|
arguments_text: "{}".into(),
|
||||||
arguments: json!({}),
|
arguments: json!({}),
|
||||||
|
argument_error: None,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -68,6 +71,37 @@ fn mcp_context(server: &str, provider: &str, tool: &str) -> ExecContext {
|
|||||||
context
|
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]
|
#[test]
|
||||||
fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
|
fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
|
||||||
let call = ToolCall {
|
let call = ToolCall {
|
||||||
@@ -77,6 +111,7 @@ fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
|
|||||||
name: "mcp_repo_lookup".into(),
|
name: "mcp_repo_lookup".into(),
|
||||||
arguments_text: "{\"query\":\"x\"}".into(),
|
arguments_text: "{\"query\":\"x\"}".into(),
|
||||||
arguments: json!({"query": "x"}),
|
arguments: json!({"query": "x"}),
|
||||||
|
argument_error: None,
|
||||||
};
|
};
|
||||||
let definition = pb::McpToolDefinition {
|
let definition = pb::McpToolDefinition {
|
||||||
name: "mcp_repo_lookup".into(),
|
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"));
|
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]
|
#[tokio::test]
|
||||||
async fn shell_uses_background_timeout_and_preserves_stream_identity() {
|
async fn shell_uses_background_timeout_and_preserves_stream_identity() {
|
||||||
let mut shell = call("call-shell", "Shell");
|
let mut shell = call("call-shell", "Shell");
|
||||||
@@ -701,12 +782,15 @@ async fn an_exec_result_must_match_the_reserved_tool() {
|
|||||||
},
|
},
|
||||||
&pending,
|
&pending,
|
||||||
)
|
)
|
||||||
.await;
|
.await
|
||||||
let Err(error) = result else {
|
.unwrap();
|
||||||
panic!("mismatched result must fail")
|
let codec::ClientExecEvent::Completed(completion) = result else {
|
||||||
|
panic!("mismatched result must complete as a tool error")
|
||||||
};
|
};
|
||||||
assert!(error
|
assert!(completion.result().is_error);
|
||||||
.to_string()
|
assert!(completion
|
||||||
|
.result()
|
||||||
|
.content
|
||||||
.contains("unexpected Exec result for tool Read"));
|
.contains("unexpected Exec result for tool Read"));
|
||||||
assert!(pending.exec_call(id).await.is_none());
|
assert!(pending.exec_call(id).await.is_none());
|
||||||
assert_eq!(pending.completed_call(id).await.as_deref(), Some("call-1"));
|
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,
|
&pending,
|
||||||
)
|
)
|
||||||
.await;
|
.await
|
||||||
let Err(duplicate) = duplicate else {
|
.unwrap();
|
||||||
panic!("duplicate terminal result must fail")
|
assert!(matches!(duplicate, codec::ClientExecEvent::Pending));
|
||||||
};
|
|
||||||
assert!(duplicate.to_string().contains("duplicate terminal"));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
async fn unknown_exec_id_is_a_protocol_error() {
|
async fn unknown_exec_id_is_ignored() {
|
||||||
let result = codec::client_event(
|
let result = codec::client_event(
|
||||||
&pb::ExecClientMessage {
|
&pb::ExecClientMessage {
|
||||||
id: 999,
|
id: 999,
|
||||||
@@ -737,15 +819,9 @@ async fn unknown_exec_id_is_a_protocol_error() {
|
|||||||
},
|
},
|
||||||
&CursorToolRuntime::default(),
|
&CursorToolRuntime::default(),
|
||||||
)
|
)
|
||||||
.await;
|
.await
|
||||||
let Err(error) = result else {
|
.unwrap();
|
||||||
panic!("unknown Exec id must fail")
|
assert!(matches!(result, codec::ClientExecEvent::Pending));
|
||||||
};
|
|
||||||
assert!(matches!(
|
|
||||||
error,
|
|
||||||
cursor_server::Error::Protocol(message)
|
|
||||||
if message == "unknown ExecClientMessage id: 999"
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[tokio::test]
|
#[tokio::test]
|
||||||
|
|||||||
Reference in New Issue
Block a user