Files
cursor-byok/cursor-server/tests/tool_loop.rs
T

845 lines
30 KiB
Rust

#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{
collections::{BTreeMap, HashSet},
sync::Arc,
};
use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::{
connect,
proto::agent::v1 as pb,
tools::{
codec,
runtime::{CursorToolRuntime, ExecContext},
ToolBatchState, ToolDispatcher,
},
},
cursor::{CursorCommand, CursorSessionRegistry},
model::{MessageContent, ToolCall},
provider::{FinishReason, ModelEvent},
};
use prost::Message;
use serde_json::json;
fn call(id: &str, name: &str) -> ToolCall {
ToolCall {
index: 0,
call_id: id.into(),
model_call_id: "model:0".into(),
name: name.into(),
arguments_text: "{}".into(),
arguments: json!({}),
}
}
fn exec_context() -> ExecContext {
ExecContext {
conversation_id: "conversation".into(),
root_conversation_id: "conversation".into(),
model_id: "model".into(),
subagent_models: std::collections::HashMap::new(),
terminals_folder: "/tmp/terminals".into(),
admin_command_denylist: Vec::new(),
}
}
#[test]
fn dynamic_mcp_call_routes_to_the_captured_exec_message() {
let call = ToolCall {
index: 0,
call_id: "mcp-call".into(),
model_call_id: "model:0".into(),
name: "mcp_repo_lookup".into(),
arguments_text: "{\"query\":\"x\"}".into(),
arguments: json!({"query": "x"}),
};
let definition = pb::McpToolDefinition {
name: "mcp_repo_lookup".into(),
provider_identifier: "repo".into(),
tool_name: "lookup".into(),
description: "lookup".into(),
input_schema: None,
input_schema_json: None,
};
let message = codec::mcp_request(7, &call, &definition).unwrap();
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message else {
panic!("expected ExecServerMessage")
};
let Some(pb::exec_server_message::Message::McpArgs(args)) = exec.message else {
panic!("expected McpArgs")
};
assert_eq!(exec.exec_id, "mcp-call");
assert_eq!(args.provider_identifier, "repo");
assert_eq!(args.tool_name, "lookup");
assert_eq!(
args.args["query"].kind,
Some(prost_types::value::Kind::StringValue("x".into()))
);
}
#[tokio::test]
async fn call_mcp_tool_reuses_the_exact_definition_returned_by_mcp_state() {
let runtime = CursorToolRuntime::default();
let dispatcher = ToolDispatcher::new(runtime.clone());
let completed = HashSet::new();
let started = HashSet::new();
let state = ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
};
let discovery = call("get-tools", "GetMcpTools");
let request = dispatcher
.start_batch(&[discovery], state, &[], &BTreeMap::new(), &exec_context())
.await
.unwrap();
let Some(pb::agent_server_message::Message::ExecServerMessage(request)) =
request[0].messages[1].message.as_ref()
else {
panic!("expected MCP state Exec")
};
let event = codec::client_event(
&pb::ExecClientMessage {
id: request.id,
message: Some(pb::exec_client_message::Message::McpStateExecResult(
pb::McpStateExecResult {
result: Some(pb::mcp_state_exec_result::Result::Success(
pb::McpStateSuccess {
servers: vec![pb::McpStateServer {
server_name: "browser-use".into(),
server_identifier: "plugin-browser-use-browser-use".into(),
tools: vec![pb::McpToolDefinition {
name: "plugin-browser-use-browser-use-browser_exec".into(),
provider_identifier: "browser-use".into(),
tool_name: "browser_exec".into(),
description: "execute browser code".into(),
..Default::default()
}],
..Default::default()
}],
},
)),
},
)),
..Default::default()
},
&runtime,
)
.await
.unwrap();
assert!(matches!(event, codec::ClientExecEvent::Completed(_)));
let mut invocation = call("call-mcp", "CallMcpTool");
invocation.arguments = json!({
"server": "plugin-browser-use-browser-use",
"toolName": "browser_exec",
"description": "run browser code",
"arguments": {"code": "print('ok')"}
});
let requests = dispatcher
.start_batch(
&[invocation],
ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
},
&[],
&BTreeMap::new(),
&exec_context(),
)
.await
.unwrap();
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) =
requests[0].messages[1].message.as_ref()
else {
panic!("expected MCP Exec")
};
let Some(pb::exec_server_message::Message::McpArgs(args)) = exec.message.as_ref() else {
panic!("expected McpArgs")
};
assert_eq!(args.name, "plugin-browser-use-browser-use-browser_exec");
assert_eq!(args.provider_identifier, "browser-use");
assert_eq!(args.tool_name, "browser_exec");
assert_eq!(args.server_identifier, "plugin-browser-use-browser-use");
assert_eq!(
args.args["code"].kind,
Some(prost_types::value::Kind::StringValue("print('ok')".into()))
);
let event = codec::client_event(
&pb::ExecClientMessage {
id: exec.id,
message: Some(pb::exec_client_message::Message::McpResult(pb::McpResult {
result: Some(pb::mcp_result::Result::Success(pb::McpSuccess {
content: vec![pb::McpToolResultContentItem {
content: Some(pb::mcp_tool_result_content_item::Content::Text(
pb::McpTextContent {
text: "browser result".into(),
output_location: None,
},
)),
}],
is_error: false,
structured_content: None,
})),
})),
..Default::default()
},
&runtime,
)
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = event else {
panic!("expected MCP completion")
};
assert_eq!(completion.result().content, "browser result");
assert!(!completion.result().is_error);
}
#[tokio::test]
async fn shell_uses_background_timeout_and_preserves_stream_identity() {
let mut shell = call("call-shell", "Shell");
shell.arguments = json!({
"command": "python3 -m http.server 8000",
"working_directory": "/tmp/project",
"block_until_ms": 3000,
"description": "Start HTTP server"
});
let context = exec_context();
let request = codec::request(7, &shell, &context).unwrap();
let Some(pb::agent_server_message::Message::ExecServerMessage(request)) = request.message
else {
panic!("expected ExecServerMessage")
};
assert_eq!(request.accept_hook_additional_contexts, Some(true));
let Some(pb::exec_server_message::Message::ShellStreamArgs(args)) = request.message else {
panic!("expected ShellArgs")
};
assert_eq!(args.timeout, 3000);
assert_eq!(
args.timeout_behavior,
pb::TimeoutBehavior::Background as i32
);
assert_eq!(args.hard_timeout, Some(86_400_000));
assert_eq!(args.description.as_deref(), Some("Start HTTP server"));
assert!(args.close_stdin);
assert_eq!(args.conversation_id.as_deref(), Some("conversation"));
assert_eq!(args.file_output_threshold_bytes, Some(40_000));
let pending = CursorToolRuntime::default();
let id = pending.reserve_exec(&shell, &context).await.unwrap();
let delta = codec::client_event(
&pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::ShellStream(
pb::ShellStream {
event: Some(pb::shell_stream::Event::Stdout(pb::ShellStreamStdout {
data: "Serving HTTP on port 8000\n".into(),
})),
},
)),
..Default::default()
},
&pending,
)
.await
.unwrap();
let codec::ClientExecEvent::Delta(delta) = delta else {
panic!("expected Shell stdout delta")
};
let Some(pb::agent_server_message::Message::InteractionUpdate(delta)) = delta.message else {
panic!("expected InteractionUpdate")
};
let Some(pb::interaction_update::Message::ToolCallDelta(delta)) = delta.message else {
panic!("expected ToolCallDelta")
};
assert_eq!(delta.call_id, "call-shell");
assert_eq!(delta.model_call_id, "model:0");
let Some(pb::tool_call_delta::Delta::ShellToolCallDelta(shell_delta)) =
delta.tool_call_delta.and_then(|delta| delta.delta)
else {
panic!("expected ShellToolCallDelta")
};
let Some(pb::shell_tool_call_delta::Delta::Stdout(stdout)) = shell_delta.delta else {
panic!("expected stdout")
};
assert_eq!(stdout.content, "Serving HTTP on port 8000\n");
let completion = codec::client_event(
&pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::ShellStream(
pb::ShellStream {
event: Some(pb::shell_stream::Event::Backgrounded(
pb::ShellStreamBackgrounded {
shell_id: 42,
command: "python3 -m http.server 8000".into(),
working_directory: "/tmp/project".into(),
pid: Some(1234),
ms_to_wait: Some(3000),
reason: Some(pb::ShellBackgroundReason::Timeout as i32),
},
)),
},
)),
..Default::default()
},
&pending,
)
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = completion else {
panic!("expected background completion")
};
assert_eq!(
completion.result().content,
(
"shell running in background shell_id=42 pid=1234 terminals_folder=/tmp/terminals\nServing HTTP on port 8000\n"
)
);
let Some(pb::tool_call::Tool::ShellToolCall(tool)) = &completion.tool_call().tool else {
panic!("expected ShellToolCall")
};
let result = tool.result.as_ref().expect("background ShellResult");
assert_eq!(result.is_background, Some(true));
assert_eq!(result.terminals_folder.as_deref(), Some("/tmp/terminals"));
assert_eq!(result.pid, Some(1234));
assert!(
pending.drain_running().await.is_empty(),
"a backgrounded Shell is no longer an abortable Run Exec"
);
}
#[tokio::test]
async fn exec_ids_are_monotonic_and_released_ids_are_not_reused() {
let pending = CursorToolRuntime::default();
let first = pending
.reserve_exec(&call("call-1", "Read"), &exec_context())
.await
.unwrap();
assert_eq!(first, 1);
assert_eq!(
pending.exec_call(first).await.map(|call| call.call_id),
Some("call-1".into())
);
pending.discard_exec(first).await;
assert!(pending.exec_call(first).await.is_none());
let second = pending
.reserve_exec(&call("call-2", "Read"), &exec_context())
.await
.unwrap();
assert_eq!(second, 2, "released Exec ids must not be reused in one Run");
let interaction = pending
.reserve_interaction(&call("call-3", "AskQuestion"), &exec_context())
.await
.unwrap();
assert_eq!(
interaction, 3,
"Exec and Interaction share one wire-id space"
);
}
#[tokio::test]
async fn empty_exec_client_message_is_not_a_terminal_result() {
let pending = CursorToolRuntime::default();
let id = pending
.reserve_exec(&call("call-1", "Read"), &exec_context())
.await
.unwrap();
let event = codec::client_event(
&pb::ExecClientMessage {
id,
message: None,
..Default::default()
},
&pending,
)
.await
.unwrap();
assert!(matches!(event, codec::ClientExecEvent::Pending));
assert_eq!(
pending.exec_call(id).await.map(|call| call.call_id),
Some("call-1".into())
);
}
#[tokio::test]
async fn tool_success_is_not_inferred_from_debug_text() {
let pending = CursorToolRuntime::default();
let mut write = call("call-1", "Write");
write.arguments = json!({"path": "/tmp/a", "contents": "x"});
let id = pending.reserve_exec(&write, &exec_context()).await.unwrap();
let event = codec::client_event(
&pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::WriteResult(
pb::WriteResult {
result: Some(pb::write_result::Result::Success(pb::WriteSuccess {
path: "/tmp/a".into(),
file_content_after_write: Some("enum Error { Example }".into()),
..Default::default()
})),
},
)),
..Default::default()
},
&pending,
)
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = event else {
panic!("expected terminal write result")
};
assert!(!completion.result().is_error);
assert!(matches!(
completion.tool_call().tool,
Some(pb::tool_call::Tool::EditToolCall(_))
));
}
#[tokio::test]
async fn new_task_result_exposes_the_subagent_name_and_id_to_the_model() {
let pending = CursorToolRuntime::default();
let mut task = call("call-task", "Task");
task.arguments = json!({
"description": "Analyze game logic",
"prompt": "Inspect the game",
"run_in_background": true,
"subagent_type": "generalPurpose"
});
let id = pending.reserve_exec(&task, &exec_context()).await.unwrap();
let event = codec::client_event(
&pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::SubagentResult(
pb::SubagentResult {
result: Some(pb::subagent_result::Result::Success(pb::SubagentSuccess {
agent_id: "child-id".into(),
..Default::default()
})),
},
)),
..Default::default()
},
&pending,
)
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = event else {
panic!("expected terminal Task result")
};
assert_eq!(
completion.result().content,
"Subagent name: Analyze game logic\nSubagent ID: child-id"
);
let Some(pb::tool_call::Tool::TaskToolCall(tool)) = &completion.tool_call().tool else {
panic!("expected TaskToolCall")
};
let Some(pb::task_result::Result::Success(success)) = tool
.result
.as_ref()
.and_then(|result| result.result.as_ref())
else {
panic!("expected typed Task success")
};
assert_eq!(success.agent_id.as_deref(), Some("child-id"));
}
#[tokio::test]
async fn an_exec_result_must_match_the_reserved_tool() {
let pending = CursorToolRuntime::default();
let id = pending
.reserve_exec(&call("call-1", "Read"), &exec_context())
.await
.unwrap();
let result = codec::client_event(
&pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::WriteResult(
pb::WriteResult {
result: Some(pb::write_result::Result::Success(pb::WriteSuccess {
path: "/tmp/a".into(),
..Default::default()
})),
},
)),
..Default::default()
},
&pending,
)
.await;
let Err(error) = result else {
panic!("mismatched result must fail")
};
assert!(error
.to_string()
.contains("unexpected Exec result for tool Read"));
assert!(pending.exec_call(id).await.is_none());
assert_eq!(pending.completed_call(id).await.as_deref(), Some("call-1"));
let duplicate = codec::client_event(
&pb::ExecClientMessage {
id,
message: None,
..Default::default()
},
&pending,
)
.await;
let Err(duplicate) = duplicate else {
panic!("duplicate terminal result must fail")
};
assert!(duplicate.to_string().contains("duplicate terminal"));
}
#[tokio::test]
async fn unknown_exec_id_is_a_protocol_error() {
let result = codec::client_event(
&pb::ExecClientMessage {
id: 999,
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult::default(),
)),
..Default::default()
},
&CursorToolRuntime::default(),
)
.await;
let Err(error) = result else {
panic!("unknown Exec id must fail")
};
assert!(matches!(
error,
cursor_server::Error::Protocol(message)
if message == "unknown ExecClientMessage id: 999"
));
}
#[tokio::test]
async fn await_shell_consumes_the_background_output_file_terminal_state() {
let runtime = CursorToolRuntime::default();
let dispatcher = ToolDispatcher::new(runtime.clone());
let mut await_call = call("await-call", "AwaitShell");
await_call.arguments = json!({
"shell_id": "42",
"block_until_ms": 1000,
"pattern": "ready",
});
await_call.arguments_text = await_call.arguments.to_string();
let completed = HashSet::new();
let started = HashSet::new();
let dispatched = dispatcher
.start_batch(
&[await_call],
ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
},
&[],
&BTreeMap::new(),
&exec_context(),
)
.await
.unwrap();
let exec = dispatched[0]
.messages
.iter()
.find_map(|message| match message.message.as_ref() {
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => Some(exec),
_ => None,
})
.unwrap();
let Some(pb::exec_server_message::Message::ReadArgs(read)) = exec.message.as_ref() else {
panic!("expected AwaitShell ReadArgs")
};
assert_eq!(read.path, "/tmp/terminals/42.txt");
let event = codec::client_event(
&pb::ExecClientMessage {
id: exec.id,
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
output: Some(pb::read_success::Output::Content(
"server ready\nexit_code: 0\n".into(),
)),
..Default::default()
})),
},
)),
..Default::default()
},
&runtime,
)
.await
.unwrap();
let codec::ClientExecEvent::Completed(completion) = event else {
panic!("expected completed AwaitShell")
};
assert_eq!(completion.result().call_id, "await-call");
assert!(!completion.result().is_error);
let Some(pb::tool_call::Tool::AwaitToolCall(tool)) = completion.tool_call().tool.as_ref()
else {
panic!("expected AwaitToolCall")
};
let pb::await_result::Result::Success(success) =
tool.result.as_ref().unwrap().result.as_ref().unwrap()
else {
panic!("expected Await success")
};
let pb::await_success::AwaitResult::Complete(complete) = success.await_result.as_ref().unwrap()
else {
panic!("expected completed background task")
};
assert_eq!(complete.task_id, "42");
assert_eq!(complete.exit_code, Some(0));
assert_eq!(complete.regex_match.as_deref(), Some("ready"));
}
#[tokio::test]
async fn provider_tool_use_waits_for_client_result_then_calls_provider_again() {
let (directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "ignored".into(),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
},
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":\"/tmp/a\"}".into(),
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "ignored".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("done".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("tool-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run()),
})
.await
.unwrap();
let mut seqno = 1;
let mut saw_exec = false;
let mut saw_typed_completion = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), 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(),
json!({})
);
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
saw_exec = true;
let exec_id = exec.id;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id: exec_id,
exec_id: String::new(),
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(
pb::ReadSuccess {
path: "/tmp/a".into(),
total_lines: 1,
file_size: 1,
output: Some(
pb::read_success::Output::Content(
"x".into(),
),
),
..Default::default()
},
)),
},
)),
..Default::default()
},
)),
}),
})
.await
.unwrap();
seqno += 1;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(
pb::agent_client_message::Message::ExecClientControlMessage(
pb::ExecClientControlMessage {
message: Some(
pb::exec_client_control_message::Message::StreamClose(
pb::ExecClientStreamClose { id: exec_id },
),
),
},
),
),
}),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
if let Some(pb::interaction_update::Message::ToolCallCompleted(completed)) =
update.message
{
let tool_call = completed.tool_call.expect("completed ToolCall");
assert!(tool_call.started_at_ms.unwrap_or_default() > 1);
assert!(tool_call.completed_at_ms.unwrap_or_default() > 1);
assert!(tool_call.completed_at_ms >= tool_call.started_at_ms);
let Some(pb::tool_call::Tool::ReadToolCall(read)) = tool_call.tool else {
panic!("expected completed ReadToolCall")
};
let result = read.result.expect("typed ReadToolResult");
assert!(matches!(
result.result,
Some(pb::read_tool_result::Result::Success(_))
));
saw_typed_completion = true;
}
}
_ => {}
}
}
assert!(saw_exec);
assert!(saw_typed_completion);
assert_eq!(provider.requests().len(), 2);
let database = sqlx::SqlitePool::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
let provider_call_index: i64 =
sqlx::query_scalar("SELECT provider_call_index FROM runs WHERE run_id = ?")
.bind("tool-request")
.fetch_one(&database)
.await
.unwrap();
assert_eq!(provider_call_index, 1);
let messages = store
.load_current_messages(&cursor_server::model::ConversationId::new(
"tool-conversation",
))
.await
.unwrap();
let result_position = messages
.iter()
.position(|message| matches!(message.content, MessageContent::ToolResult(_)))
.expect("tool result persisted");
let MessageContent::Assistant { tool_calls, .. } = &messages[result_position - 1].content
else {
panic!("tool result must immediately follow its assistant tool call")
};
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0].call_id, "call-1");
}
fn client_run() -> pb::AgentClientMessage {
let user = pb::UserMessage {
text: "read it".into(),
message_id: "user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
};
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
action: Some(pb::ConversationAction {
action: Some(pb::conversation_action::Action::UserMessageAction(
pb::UserMessageAction {
user_message: Some(user),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("tool-conversation".into()),
run_id: Some("tool-request".into()),
requested_model: Some(pb::RequestedModel {
model_id: "test-model".into(),
..Default::default()
}),
..Default::default()
},
)),
}
}
fn kv_ack(id: u32) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::KvClientMessage(
pb::KvClientMessage {
id,
message: Some(pb::kv_client_message::Message::SetBlobResult(
pb::SetBlobResult { error: None },
)),
},
)),
}
}