all tools

This commit is contained in:
leookun
2026-08-16 17:29:29 +08:00
parent eafede22e3
commit 4db2061611
95 changed files with 19023 additions and 5637 deletions
@@ -0,0 +1,79 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{connect, proto::agent::v1 as pb},
prompting::{PromptAssets, PromptCompiler},
run::RunRegistry,
store::BlobId,
};
use prost::Message;
#[tokio::test]
async fn checkpoint_dependency_is_not_ready_until_blob_ack() {
let (_directory, store) = fixtures::temp_store().await;
store.create_pending_run("request").await.unwrap();
let id = BlobId::digest(b"tool result");
let key = format!("blob:{}", id.to_base64());
store
.enqueue_outbox("request", &key, "kv_set", b"tool result", &[])
.await
.unwrap();
assert!(!store
.dependencies_acked("request", std::slice::from_ref(&id))
.await
.unwrap());
assert!(store.ack_outbox("request", &key).await.unwrap());
assert!(store.dependencies_acked("request", &[id]).await.unwrap());
}
#[tokio::test]
async fn pending_blob_outbox_is_replayed_when_the_run_is_recreated() {
let (_directory, store) = fixtures::temp_store().await;
store.create_pending_run("recover-request").await.unwrap();
let data = b"durable tool result";
let blob_id = store.put_blob(data, &[]).await.unwrap();
store
.enqueue_outbox(
"recover-request",
&format!("blob:{}", blob_id.to_base64()),
"kv_set",
data,
&[],
)
.await
.unwrap();
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
"test".into(),
);
let handle = registry.get_or_create("recover-request").await.unwrap();
let mut output = handle.subscribe();
let frame = tokio::time::timeout(std::time::Duration::from_secs(2), output.recv())
.await
.unwrap()
.unwrap();
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
let message = pb::AgentServerMessage::decode(payload).unwrap();
let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = message.message else {
panic!("expected replayed KV SET")
};
let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = kv.message else {
panic!("expected SetBlobArgs")
};
assert_eq!(set.blob_id, blob_id.as_bytes());
assert_eq!(set.blob_data, data);
}
+136
View File
@@ -0,0 +1,136 @@
#[path = "support/fake_cursor.rs"]
mod fake_cursor;
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{io::Write, sync::Arc};
use axum::{
body::{to_bytes, Body},
http::{header, Request, StatusCode},
};
use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine};
use cursor_server::{
cursor::{
connect, handlers,
proto::{agent::v1 as pb, aiserver::v1 as ai},
},
prompting::{PromptAssets, PromptCompiler},
run::RunRegistry,
};
use flate2::{write::GzEncoder, Compression};
use prost::Message;
use tower::ServiceExt;
#[test]
fn connect_envelope_is_flag_plus_big_endian_length_plus_protobuf() {
let message = pb::BidiRequestId {
request_id: "abc".into(),
};
let frame = connect::encode_message(&message).unwrap();
assert_eq!(&frame[..5], &[0, 0, 0, 0, 5]);
let decoded: pb::BidiRequestId = fake_cursor::decode_single(&frame).unwrap();
assert_eq!(decoded.request_id, "abc");
}
#[test]
fn end_stream_matches_captured_connect_shape() {
assert_eq!(
connect::encode_end_stream().as_ref(),
&[2, 0, 0, 0, 2, b'{', b'}']
);
}
#[test]
fn error_end_stream_is_flagged_json_not_protobuf() {
let frame = connect::encode_error_end_stream(&connect::ConnectStreamError {
code: connect::ConnectCode::Unavailable,
message: "overloaded".into(),
details: vec![connect::ConnectErrorDetail {
type_name: "aiserver.v1.ErrorDetails".into(),
value: "AQ".into(),
}],
})
.unwrap();
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags, connect::END_STREAM_FLAG);
let json: serde_json::Value = serde_json::from_slice(&payload).unwrap();
assert_eq!(json["error"]["code"], "unavailable");
assert_eq!(json["error"]["message"], "overloaded");
assert_eq!(
json["error"]["details"][0]["type"],
"aiserver.v1.ErrorDetails"
);
}
#[test]
fn cursor_error_details_subset_decodes_captured_wire_value() {
let captured = "CAISVQoUQXV0aGVudGljYXRpb24gZXJyb3ISMklmIHlvdSBhcmUgbG9nZ2VkIGluLCB0cnkgbG9nZ2luZyBvdXQgYW5kIGJhY2sgaW4uIABSBwoFbG9naW4YAQ";
let bytes = STANDARD_NO_PAD.decode(captured).unwrap();
let details = ai::ErrorDetails::decode(bytes.as_slice()).unwrap();
assert_eq!(details.error, 2, "ERROR_NOT_LOGGED_IN");
assert_eq!(details.is_expected, Some(true));
let custom = details.details.unwrap();
assert_eq!(custom.title, "Authentication error");
assert_eq!(custom.is_retryable, Some(false));
}
#[test]
fn captured_kv_ack_hex_decodes_as_agent_client_message() {
let bytes = hex::decode("1a0408011a00").unwrap();
let message = pb::AgentClientMessage::decode(bytes.as_slice()).unwrap();
let Some(pb::agent_client_message::Message::KvClientMessage(kv)) = message.message else {
panic!("expected KV client message")
};
assert_eq!(kv.id, 1);
assert!(matches!(
kv.message,
Some(pb::kv_client_message::Message::SetBlobResult(_))
));
}
#[tokio::test]
async fn bidi_append_gzip_body_is_decompressed_before_protobuf_decode() {
let (_directory, store) = fixtures::temp_store().await;
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
"test-model".into(),
);
let wire = ai::BidiAppendRequest {
request_id: Some(ai::BidiRequestId {
request_id: "gzip-request".into(),
}),
..Default::default()
}
.encode_to_vec();
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&wire).unwrap();
let compressed = encoder.finish().unwrap();
let response = handlers::router(registry)
.oneshot(
Request::post("/aiserver.v1.BidiService/BidiAppend")
.header(header::CONTENT_TYPE, "application/proto")
.header(header::CONTENT_ENCODING, "gzip")
.body(Body::from(compressed))
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
let body = to_bytes(response.into_body(), 4096).await.unwrap();
let text = std::str::from_utf8(&body).unwrap();
assert!(text.contains("BidiAppend contains no AgentClientMessage"));
assert!(!text.contains("protobuf decode error"));
}
+306
View File
@@ -0,0 +1,306 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine};
use cursor_server::{
cursor::{
connect,
proto::{agent::v1 as pb, aiserver::v1 as ai},
},
model::{MessageContent, Role},
prompting::{PromptAssets, PromptCompiler},
provider::{FinishReason, ResponseEvent},
run::{RunCommand, RunRegistry},
Error,
};
use prost::Message;
#[tokio::test]
async fn provider_failure_checkpoints_then_returns_structured_error_and_closes() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push_error(Error::Provider("provider failed".into()));
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store.clone(),
Arc::new(provider),
PromptCompiler::new(assets),
"test-model".into(),
);
let handle = registry.get_or_create("failed-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(RunCommand::Append {
seqno: 0,
message: Box::new(client_run()),
})
.await
.unwrap();
let mut append_seqno = 1;
let mut saw_checkpoint = false;
let mut saw_turn_ended = false;
let error_json = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
break serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(RunCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(_)) => {
saw_checkpoint = true;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
if matches!(
update.message,
Some(pb::interaction_update::Message::TurnEnded(_))
) {
saw_turn_ended = true;
}
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
assert!(!delta.text.contains("Cursor server error"));
}
}
_ => {}
}
};
assert!(saw_checkpoint);
assert!(!saw_turn_ended);
assert_eq!(error_json["error"]["code"], "unavailable");
let detail = &error_json["error"]["details"][0];
assert_eq!(detail["type"], "aiserver.v1.ErrorDetails");
let encoded = detail["value"].as_str().unwrap();
assert!(!encoded.ends_with('='));
let decoded = STANDARD_NO_PAD.decode(encoded).unwrap();
let decoded = ai::ErrorDetails::decode(decoded.as_slice()).unwrap();
assert_eq!(
decoded.error,
ai::error_details::Error::ProviderError as i32
);
let custom = decoded.details.unwrap();
assert_eq!(custom.title, "Server Error");
assert_eq!(custom.is_retryable, Some(true));
assert_eq!(custom.should_show_immediate_error, Some(false));
assert_eq!(
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
.await
.unwrap(),
None
);
let messages = store.load_messages("failed-conversation").await.unwrap();
assert!(messages.iter().any(|message| message.role == Role::User));
assert!(!messages.iter().any(|message| {
matches!(
&message.content,
MessageContent::Assistant { text, .. } if text.contains("Cursor server error")
)
}));
}
#[tokio::test]
async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ResponseEvent::Start {
model_call_id: "model-call".into(),
},
ResponseEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
},
ResponseEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":\"/tmp/a\"}".into(),
},
ResponseEvent::ToolCallEnd { index: 0 },
ResponseEvent::Done(FinishReason::ToolUse),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store,
Arc::new(provider),
PromptCompiler::new(assets),
"test-model".into(),
);
let handle = registry
.get_or_create("protocol-failed-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(RunCommand::Append {
seqno: 0,
message: Box::new(protocol_client_run()),
})
.await
.unwrap();
let mut append_seqno = 1;
let mut saw_turn_ended = false;
let error_json = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Error EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
break serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(RunCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
// An unknown numeric bridge id is a runtime protocol error.
handle
.command(RunCommand::Append {
seqno: append_seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id: exec.id + 1_000,
exec_id: String::new(),
message: None,
..Default::default()
},
)),
}),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
if matches!(
update.message,
Some(pb::interaction_update::Message::TurnEnded(_))
) {
saw_turn_ended = true;
}
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
assert!(!delta.text.contains("unknown tool result"));
assert!(!delta.text.contains("protocol error"));
}
}
_ => {}
}
};
assert!(!saw_turn_ended);
assert_eq!(error_json["error"]["code"], "invalid_argument");
assert_eq!(
error_json["error"]["message"],
"protocol error: unknown ExecClientMessage id: 1001"
);
assert_eq!(
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
.await
.unwrap(),
None
);
}
fn client_run() -> pb::AgentClientMessage {
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(pb::UserMessage {
text: "hello".into(),
message_id: "failed-user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("failed-conversation".into()),
run_id: Some("failed-request".into()),
..Default::default()
},
)),
}
}
fn protocol_client_run() -> pb::AgentClientMessage {
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(pb::UserMessage {
text: "read it".into(),
message_id: "protocol-failed-user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("protocol-failed-conversation".into()),
run_id: Some("protocol-failed-request".into()),
..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 },
)),
},
)),
}
}
+190
View File
@@ -0,0 +1,190 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{connect, proto::agent::v1 as pb},
prompting::{PromptAssets, PromptCompiler},
provider::{FinishReason, ResponseEvent},
run::{RunCommand, RunRegistry},
};
use prost::Message;
#[tokio::test]
async fn a_new_revision_invalidates_late_events_from_the_old_run() {
let (_directory, store) = fixtures::temp_store().await;
let first = store.begin_revision("conversation").await.unwrap();
let second = store.begin_revision("conversation").await.unwrap();
assert!(second > first);
assert!(!store
.revision_is_current("conversation", first)
.await
.unwrap());
assert!(store
.revision_is_current("conversation", second)
.await
.unwrap());
}
#[tokio::test]
async fn registry_shutdown_cancels_runs_and_closes_run_sse_outputs() {
let (_directory, store) = fixtures::temp_store().await;
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
"test-model".into(),
);
let handle = registry.get_or_create("active-run").await.unwrap();
let mut output = handle.subscribe();
registry.shutdown().await;
assert!(handle.cancellation().is_cancelled());
let terminal = output.recv().await.expect("canceled EndStream");
let (flags, payload) = connect::decode_frames(&terminal).unwrap().pop().unwrap();
assert_eq!(flags, connect::END_STREAM_FLAG);
let payload: serde_json::Value = serde_json::from_slice(&payload).unwrap();
assert_eq!(payload["error"]["code"], "canceled");
assert_eq!(output.recv().await, None);
}
#[tokio::test]
async fn cancel_aborts_active_exec_before_canceled_end_stream() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ResponseEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
},
ResponseEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":\"/tmp/a\"}".into(),
},
ResponseEvent::ToolCallEnd { index: 0 },
ResponseEvent::Done(FinishReason::ToolUse),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store,
Arc::new(provider),
PromptCompiler::new(assets),
"test-model".into(),
);
let handle = registry.get_or_create("cancel-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(RunCommand::Append {
seqno: 0,
message: Box::new(client_run()),
})
.await
.unwrap();
let mut append_seqno = 1;
let exec_id = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.unwrap();
let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(RunCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => break exec.id,
_ => {}
}
};
handle.cancel();
let mut saw_abort = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before canceled EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
let json: serde_json::Value = serde_json::from_slice(&payload).unwrap();
assert_eq!(json["error"]["code"], "canceled");
assert!(saw_abort, "ExecServerAbort must precede canceled EndStream");
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) =
server.message
{
let Some(pb::exec_server_control_message::Message::Abort(abort)) = control.message
else {
panic!("expected ExecServerAbort")
};
assert_eq!(abort.id, exec_id);
saw_abort = true;
}
}
assert_eq!(output.recv().await, None);
}
fn client_run() -> pb::AgentClientMessage {
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(pb::UserMessage {
text: "read".into(),
message_id: "cancel-user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("cancel-conversation".into()),
run_id: Some("cancel-request".into()),
..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 },
)),
},
)),
}
}
+181
View File
@@ -0,0 +1,181 @@
#[path = "support/fixtures.rs"]
mod fixtures;
use cursor_server::{
model::{CanonicalMessage, MessageContent, Origin, Role, ToolCallContent, ToolResultContent},
prompting::{project_messages, Mode, PromptAssets, PromptCompiler, ToolDefinition},
};
#[test]
fn projecting_an_append_only_context_preserves_the_complete_prefix() {
let first = vec![fixtures::user("u1", "one")];
let mut second = first.clone();
second.push(fixtures::user("u2", "two"));
let projected_first = project_messages(&first).unwrap();
let projected_second = project_messages(&second).unwrap();
assert_eq!(projected_first, projected_second[..projected_first.len()]);
}
#[test]
fn every_tool_result_is_projected_as_string_content() {
let object = serde_json::json!({"merge": false, "todos": []});
let messages = vec![
tool_result("object", object.clone()),
tool_result("string", serde_json::Value::String("plain text".into())),
];
let projected = project_messages(&messages).unwrap();
let object_text = projected[0]
.content
.as_str()
.expect("object ToolResult must be JSON-encoded into a string");
assert_eq!(
serde_json::from_str::<serde_json::Value>(object_text).unwrap(),
object
);
assert_eq!(projected[1].content.as_str(), Some("plain text"));
}
#[test]
fn assistant_text_and_thinking_remain_separate_during_projection() {
let messages = vec![CanonicalMessage {
message_id: "assistant".into(),
role: Role::Assistant,
origin: Origin::Assistant,
content: MessageContent::Assistant {
text: "visible answer".into(),
thinking: "private reasoning".into(),
model_call_id: Some("model-call".into()),
tool_calls: Vec::new(),
},
runtime_event_id: None,
}];
let projected = project_messages(&messages).unwrap();
assert_eq!(projected[0].content.as_str(), Some("visible answer"));
assert_eq!(projected[0].thinking.as_deref(), Some("private reasoning"));
}
#[test]
fn split_tool_pairs_reconstruct_the_original_provider_assistant_message() {
let messages = vec![
assistant_tool_pair(
"assistant-second",
"model-call",
1,
"call-second",
"visible answer",
"complete reasoning",
),
tool_result_with_call("result-second", "call-second", "second"),
assistant_tool_pair("assistant-first", "model-call", 0, "call-first", "", ""),
tool_result_with_call("result-first", "call-first", "first"),
];
let projected = project_messages(&messages).unwrap();
assert_eq!(projected.len(), 3);
assert_eq!(projected[0].role, "assistant");
assert_eq!(projected[0].thinking.as_deref(), Some("complete reasoning"));
let calls = projected[0].tool_calls.as_ref().unwrap();
assert_eq!(calls[0]["id"], "call-first");
assert_eq!(calls[1]["id"], "call-second");
assert_eq!(projected[1].tool_call_id.as_deref(), Some("call-first"));
assert_eq!(projected[2].tool_call_id.as_deref(), Some("call-second"));
}
#[test]
fn every_prompt_mode_loads_the_captured_tool_set() {
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
assert_eq!(assets.mode(Mode::Agent).tools.len(), 20);
assert_eq!(assets.mode(Mode::Ask).tools.len(), 18);
assert_eq!(assets.mode(Mode::Plan).tools.len(), 16);
assert_eq!(assets.mode(Mode::Debug).tools.len(), 18);
assert_eq!(assets.mode(Mode::Multitask).tools.len(), 20);
assert_eq!(assets.mode(Mode::Subagent).tools.len(), 4);
}
#[test]
fn dynamic_mcp_tools_are_appended_after_the_stable_mode_tool_prefix() {
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let compiler = PromptCompiler::new(assets);
let messages = vec![fixtures::user("u1", "one")];
let base = compiler
.compile(Mode::Agent, "model", "call", &messages)
.unwrap();
let dynamic = compiler
.compile_with_dynamic_tools(
Mode::Agent,
"model",
"call",
&messages,
&[ToolDefinition {
name: "mcp_repo_lookup".into(),
description: "lookup".into(),
input_schema: serde_json::json!({"type": "object"}),
}],
)
.unwrap();
assert_eq!(base.tools, dynamic.tools[..base.tools.len()]);
assert_eq!(dynamic.tools.last().unwrap().name, "mcp_repo_lookup");
}
fn tool_result(id: &str, output: serde_json::Value) -> CanonicalMessage {
tool_result_with_call(id, &format!("call-{id}"), output)
}
fn tool_result_with_call(
id: &str,
call_id: &str,
output: impl Into<serde_json::Value>,
) -> CanonicalMessage {
CanonicalMessage {
message_id: id.into(),
role: Role::Tool,
origin: Origin::Tool,
content: MessageContent::ToolResult(ToolResultContent {
call_id: call_id.into(),
name: "Tool".into(),
output: output.into(),
is_error: false,
}),
runtime_event_id: None,
}
}
fn assistant_tool_pair(
id: &str,
model_call_id: &str,
index: usize,
call_id: &str,
text: &str,
thinking: &str,
) -> CanonicalMessage {
CanonicalMessage {
message_id: id.into(),
role: Role::Assistant,
origin: Origin::Assistant,
content: MessageContent::Assistant {
text: text.into(),
thinking: thinking.into(),
model_call_id: Some(model_call_id.into()),
tool_calls: vec![ToolCallContent {
index,
call_id: call_id.into(),
name: "Tool".into(),
arguments: serde_json::json!({}),
}],
},
runtime_event_id: None,
}
}
+27
View File
@@ -0,0 +1,27 @@
#[path = "support/fixtures.rs"]
mod fixtures;
use cursor_server::model::RuntimeEvent;
#[tokio::test]
async fn runtime_event_is_appended_exactly_once() {
let (_directory, store) = fixtures::temp_store().await;
let event = RuntimeEvent {
event_id: "branch:changed:7".into(),
text: "runtime state changed".into(),
};
assert!(store
.append_runtime_event_once("conversation", event.clone())
.await
.unwrap());
assert!(!store
.append_runtime_event_once("conversation", event)
.await
.unwrap());
let messages = store.load_messages("conversation").await.unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(
messages[0].runtime_event_id.as_deref(),
Some("branch:changed:7")
);
}
@@ -0,0 +1,9 @@
use bytes::Bytes;
use cursor_server::{cursor::connect, Result};
use prost::Message;
pub fn decode_single<M: Message + Default>(frame: &Bytes) -> Result<M> {
let frames = connect::decode_frames(frame)?;
assert_eq!(frames.len(), 1);
Ok(M::decode(frames[0].1.clone())?)
}
@@ -0,0 +1,50 @@
#![allow(dead_code)]
use std::{
collections::VecDeque,
sync::{Arc, Mutex},
};
use cursor_server::{
prompting::ModelRequest,
provider::{Provider, ProviderStream, ResponseEvent},
Error,
};
use futures_util::stream;
use tokio_util::sync::CancellationToken;
type FakeResponse = Vec<Result<ResponseEvent, Error>>;
#[derive(Clone, Default)]
pub struct FakeProvider {
responses: Arc<Mutex<VecDeque<FakeResponse>>>,
requests: Arc<Mutex<Vec<ModelRequest>>>,
}
impl FakeProvider {
pub fn push(&self, events: Vec<ResponseEvent>) {
self.responses
.lock()
.unwrap()
.push_back(events.into_iter().map(Ok).collect());
}
pub fn push_error(&self, error: Error) {
self.responses.lock().unwrap().push_back(vec![Err(error)]);
}
pub fn requests(&self) -> Vec<ModelRequest> {
self.requests.lock().unwrap().clone()
}
}
impl Provider for FakeProvider {
fn stream(&self, request: ModelRequest, _cancellation: CancellationToken) -> ProviderStream {
self.requests.lock().unwrap().push(request);
let events = self
.responses
.lock()
.unwrap()
.pop_front()
.expect("fake response configured");
Box::pin(stream::iter(events))
}
}
+17
View File
@@ -0,0 +1,17 @@
#![allow(dead_code)]
use cursor_server::{
model::{CanonicalMessage, Origin, Role},
store::Store,
};
pub async fn temp_store() -> (tempfile::TempDir, Store) {
let directory = tempfile::tempdir().unwrap();
let url = format!("sqlite://{}", directory.path().join("test.db").display());
let store = Store::connect(&url).await.unwrap();
(directory, store)
}
pub fn user(id: &str, text: &str) -> CanonicalMessage {
CanonicalMessage::text(id, Role::User, Origin::User, text)
}
+158
View File
@@ -0,0 +1,158 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{connect, proto::agent::v1 as pb},
prompting::{PromptAssets, PromptCompiler},
provider::{FinishReason, ResponseEvent},
run::{RunCommand, RunRegistry},
};
use prost::Message;
#[tokio::test]
async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ResponseEvent::Start {
model_call_id: "ignored".into(),
},
ResponseEvent::ThinkingStart,
ResponseEvent::ThinkingDelta("reason".into()),
ResponseEvent::ThinkingEnd,
ResponseEvent::TextStart,
ResponseEvent::TextDelta("hello".into()),
ResponseEvent::TextEnd,
ResponseEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
"test-model".into(),
);
let handle = registry.get_or_create("request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(RunCommand::Append {
seqno: 0,
message: Box::new(client_run("request", "conversation", "hello")),
})
.await
.unwrap();
let mut append_seqno = 1;
let mut text = String::new();
let mut thinking = String::new();
let mut thinking_duration_ms = None;
let mut saw_turn_ended = false;
let mut checkpoints = 0;
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 {
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(RunCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_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::ThinkingDelta(delta)) => {
assert_eq!(
delta.thinking_style,
Some(pb::ThinkingStyle::Default as i32)
);
thinking.push_str(&delta.text);
}
Some(pb::interaction_update::Message::ThinkingCompleted(completed)) => {
thinking_duration_ms = Some(completed.thinking_duration_ms)
}
Some(pb::interaction_update::Message::TurnEnded(_)) => saw_turn_ended = true,
_ => {}
}
}
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(_)) => {
checkpoints += 1
}
_ => {}
}
}
assert_eq!(text, "hello");
assert_eq!(thinking, "reason");
assert!(thinking_duration_ms.is_some_and(|duration| duration >= 1));
assert!(saw_turn_ended);
assert_eq!(checkpoints, 2, "final checkpoint is intentionally repeated");
assert_eq!(provider.requests().len(), 1);
assert!(store
.load_messages("conversation")
.await
.unwrap()
.iter()
.any(|message| message.message_id == "user"));
}
fn client_run(request_id: &str, conversation_id: &str, text: &str) -> pb::AgentClientMessage {
let user = pb::UserMessage {
text: text.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(conversation_id.into()),
run_id: Some(request_id.into()),
..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 },
)),
},
)),
}
}
+521
View File
@@ -0,0 +1,521 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{
connect, exec,
pending::{ExecContext, PendingExecRegistry},
proto::agent::v1 as pb,
},
model::{MessageContent, ToolCall},
prompting::{PromptAssets, PromptCompiler},
provider::{FinishReason, ResponseEvent},
run::{RunCommand, RunRegistry},
};
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(),
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 = cursor_server::cursor::exec::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 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 = exec::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 = PendingExecRegistry::default();
let id = pending.reserve(&shell, &context).await.unwrap();
let delta = exec::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 exec::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 = exec::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 exec::ClientExecEvent::Completed(completion) = completion else {
panic!("expected background completion")
};
assert_eq!(
completion.result().output.as_str(),
Some(
"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));
}
#[tokio::test]
async fn exec_ids_are_monotonic_and_released_ids_are_not_reused() {
let pending = PendingExecRegistry::default();
let first = pending
.reserve(&call("call-1", "Read"), &exec_context())
.await
.unwrap();
assert_eq!(first, 1);
assert_eq!(
pending.call(first).await.map(|call| call.call_id),
Some("call-1".into())
);
pending.discard(first).await;
assert!(pending.call(first).await.is_none());
let second = pending
.reserve(&call("call-2", "Read"), &exec_context())
.await
.unwrap();
assert_eq!(second, 2, "released Exec ids must not be reused in one Run");
}
#[tokio::test]
async fn empty_exec_client_message_is_not_a_terminal_result() {
let pending = PendingExecRegistry::default();
let id = pending
.reserve(&call("call-1", "Read"), &exec_context())
.await
.unwrap();
let event = exec::client_event(
&pb::ExecClientMessage {
id,
message: None,
..Default::default()
},
&pending,
)
.await
.unwrap();
assert!(matches!(event, exec::ClientExecEvent::Pending));
assert_eq!(
pending.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 = PendingExecRegistry::default();
let mut write = call("call-1", "Write");
write.arguments = json!({"path": "/tmp/a", "contents": "x"});
let id = pending.reserve(&write, &exec_context()).await.unwrap();
let event = exec::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 exec::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 an_exec_result_must_match_the_reserved_tool() {
let pending = PendingExecRegistry::default();
let id = pending
.reserve(&call("call-1", "Read"), &exec_context())
.await
.unwrap();
let result = exec::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.call(id).await.is_none());
}
#[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![
ResponseEvent::Start {
model_call_id: "ignored".into(),
},
ResponseEvent::ToolCallStart {
index: 0,
call_id: "call-1".into(),
name: "Read".into(),
},
ResponseEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{\"path\":\"/tmp/a\"}".into(),
},
ResponseEvent::ToolCallEnd { index: 0 },
ResponseEvent::Done(FinishReason::ToolUse),
]);
provider.push(vec![
ResponseEvent::Start {
model_call_id: "ignored".into(),
},
ResponseEvent::TextStart,
ResponseEvent::TextDelta("done".into()),
ResponseEvent::TextEnd,
ResponseEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../prompt")
.as_path(),
)
.unwrap();
let registry = RunRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
"test-model".into(),
);
let handle = registry.get_or_create("tool-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(RunCommand::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(RunCommand::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(RunCommand::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(RunCommand::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 request_id = ?")
.bind("tool-request")
.fetch_one(&database)
.await
.unwrap();
assert_eq!(provider_call_index, 1);
let messages = store.load_messages("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()),
..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 },
)),
},
)),
}
}