mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 20:44:07 +08:00
feat: initialize server structure and database schema
- Added initial server setup with Cargo.toml defining dependencies and project structure. - Created build.rs for generating protobuf bindings and validating wire contracts. - Established database schema with initial migration files for conversations, messages, and runs. - Introduced tools and prompts for Cursor functionality, enhancing user interaction capabilities.
This commit is contained in:
@@ -0,0 +1,652 @@
|
||||
//! Verifies Conversation recovery and resumable checkpoint state.
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{collections::HashSet, sync::Arc};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
protocol::{connect, proto::agent::v1 as pb},
|
||||
TransportCommand, TransportRegistry,
|
||||
},
|
||||
model::ToolRoundId,
|
||||
provider::{FinishReason, ModelEvent},
|
||||
store::{BlobEdge, BlobId},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn applied_revision_schema_upgrades_to_checkpoints_without_losing_rows() {
|
||||
use std::borrow::Cow;
|
||||
|
||||
use sqlx::{
|
||||
migrate::Migrator,
|
||||
sqlite::{SqliteConnectOptions, SqlitePoolOptions},
|
||||
};
|
||||
|
||||
static ALL_MIGRATIONS: Migrator = sqlx::migrate!("./migrations");
|
||||
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let database_path = directory.path().join("upgrade.db");
|
||||
let database_url = format!("sqlite://{}", database_path.display());
|
||||
let pool = SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect_with(
|
||||
SqliteConnectOptions::new()
|
||||
.filename(&database_path)
|
||||
.create_if_missing(true)
|
||||
.foreign_keys(true),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let previous = Migrator {
|
||||
migrations: Cow::Owned(ALL_MIGRATIONS.iter().take(5).cloned().collect()),
|
||||
..Migrator::DEFAULT
|
||||
};
|
||||
previous.run(&pool).await.unwrap();
|
||||
|
||||
sqlx::query(
|
||||
"INSERT INTO conversations(conversation_id, updated_at_ms)
|
||||
VALUES ('upgrade-conversation', 1)",
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO conversation_revisions(
|
||||
conversation_id, parent_revision_id, state_digest, created_at_ms
|
||||
) VALUES ('upgrade-conversation', NULL, zeroblob(32), 2)",
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
let revision_id: i64 = sqlx::query_scalar("SELECT last_insert_rowid()")
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"UPDATE conversations SET current_revision_id = ? WHERE conversation_id = 'upgrade-conversation'",
|
||||
)
|
||||
.bind(revision_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO messages(
|
||||
conversation_id, message_id, role, origin, payload_json, created_at_ms
|
||||
) VALUES ('upgrade-conversation', 'message-1', 'user', 'user', '{}', 3)",
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO revision_messages(revision_id, ordinal, conversation_id, message_id)
|
||||
VALUES (?, 0, 'upgrade-conversation', 'message-1')",
|
||||
)
|
||||
.bind(revision_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO runs(
|
||||
run_id, conversation_id, base_revision_id, head_revision_id, run_kind, status,
|
||||
created_at_ms, updated_at_ms
|
||||
) VALUES ('upgrade-run', 'upgrade-conversation', ?, ?, 'root', 'completed', 4, 4)",
|
||||
)
|
||||
.bind(revision_id)
|
||||
.bind(revision_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO tool_rounds(
|
||||
round_id, run_id, base_revision_id, assistant_json, status, created_at_ms, updated_at_ms
|
||||
) VALUES ('upgrade-round', 'upgrade-run', ?, '{}', 'settled', 5, 5)",
|
||||
)
|
||||
.bind(revision_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO tool_round_calls(
|
||||
round_id, call_index, call_id, model_call_id, name, arguments_json, status,
|
||||
committed_revision_id
|
||||
) VALUES ('upgrade-round', 0, 'upgrade-call', 'model-call', 'Read', '{}', 'completed', ?)",
|
||||
)
|
||||
.bind(revision_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query(
|
||||
"INSERT INTO input_anchors(conversation_id, input_id, base_revision_id, created_at_ms)
|
||||
VALUES ('upgrade-conversation', 'input-1', ?, 6)",
|
||||
)
|
||||
.bind(revision_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
pool.close().await;
|
||||
|
||||
let upgraded = cursor_server::store::Store::connect(&database_url)
|
||||
.await
|
||||
.unwrap();
|
||||
let current: i64 = sqlx::query_scalar(
|
||||
"SELECT current_checkpoint_id FROM conversations
|
||||
WHERE conversation_id = 'upgrade-conversation'",
|
||||
)
|
||||
.fetch_one(upgraded.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
let linked_messages: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM checkpoint_messages WHERE checkpoint_id = ?")
|
||||
.bind(revision_id)
|
||||
.fetch_one(upgraded.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
let run_checkpoints: (i64, i64) = sqlx::query_as(
|
||||
"SELECT base_checkpoint_id, head_checkpoint_id FROM runs WHERE run_id = 'upgrade-run'",
|
||||
)
|
||||
.fetch_one(upgraded.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
let round_checkpoint: i64 = sqlx::query_scalar(
|
||||
"SELECT base_checkpoint_id FROM tool_rounds WHERE round_id = 'upgrade-round'",
|
||||
)
|
||||
.fetch_one(upgraded.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
let committed_checkpoint: i64 = sqlx::query_scalar(
|
||||
"SELECT committed_checkpoint_id FROM tool_round_calls WHERE call_id = 'upgrade-call'",
|
||||
)
|
||||
.fetch_one(upgraded.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
let anchor_checkpoint: i64 = sqlx::query_scalar(
|
||||
"SELECT base_checkpoint_id FROM input_anchors WHERE input_id = 'input-1'",
|
||||
)
|
||||
.fetch_one(upgraded.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(current, revision_id);
|
||||
assert_eq!(linked_messages, 1);
|
||||
assert_eq!(run_checkpoints, (revision_id, revision_id));
|
||||
assert_eq!(round_checkpoint, revision_id);
|
||||
assert_eq!(committed_checkpoint, revision_id);
|
||||
assert_eq!(anchor_checkpoint, revision_id);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn checkpoint_dependencies_are_content_addressed_without_a_persistent_stream_outbox() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let child = store.put_blob(b"message", &[]).await.unwrap();
|
||||
let root = store
|
||||
.put_blob(
|
||||
b"checkpoint",
|
||||
&[BlobEdge {
|
||||
child: child.clone(),
|
||||
field_name: "turns[0]".into(),
|
||||
}],
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(root, BlobId::digest(b"checkpoint"));
|
||||
assert_eq!(store.get_blob(&child).await.unwrap().unwrap(), b"message");
|
||||
let closure = store
|
||||
.blob_closure(std::slice::from_ref(&root))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(closure.contains(&root));
|
||||
assert!(closure.contains(&child));
|
||||
|
||||
let outbox: i64 = sqlx::query_scalar(
|
||||
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'outbox'",
|
||||
)
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(outbox, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn eligible_pending_checkpoint_resumes_tools_before_the_next_model_call() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-1".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "read-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: "model-2".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("resumed".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 = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
|
||||
let first = registry.get_or_create("first-run").await.unwrap();
|
||||
let mut first_output = first.subscribe();
|
||||
first
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(start_request()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut first_seqno = 1;
|
||||
let mut sent_blob_ids = HashSet::new();
|
||||
let staged = loop {
|
||||
let server = next_message(&mut first_output).await;
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(args)) = &kv.message {
|
||||
sent_blob_ids.insert(args.blob_id.clone());
|
||||
}
|
||||
acknowledge(&first, &mut first_seqno, kv.id).await;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state))
|
||||
if state.pending_tool_calls.len() == 1 =>
|
||||
{
|
||||
break state;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
assert!(
|
||||
!sent_blob_ids.contains(
|
||||
BlobId::digest(&staged.encode_to_vec())
|
||||
.as_bytes()
|
||||
.as_slice()
|
||||
),
|
||||
"ConversationStateStructure is inline and must not be sent as a Blob"
|
||||
);
|
||||
let staged_started_at_ms =
|
||||
serde_json::from_str::<serde_json::Value>(staged.pending_tool_calls.first().unwrap())
|
||||
.unwrap()["providerOptions"]["cursor"]["pendingToolCallStartedAtMs"]
|
||||
.as_u64()
|
||||
.unwrap();
|
||||
assert_eq!(provider.requests().len(), 1);
|
||||
first.disconnect().await;
|
||||
|
||||
let resumed = registry.get_or_create("resumed-run").await.unwrap();
|
||||
let mut resumed_output = resumed.subscribe();
|
||||
let mut resumed_checkpoints = Vec::new();
|
||||
let mut resumed_set_blob_ids = HashSet::new();
|
||||
resumed
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(resume_request(staged.clone())),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut resumed_seqno = 1;
|
||||
let exec_id = loop {
|
||||
let server = next_message(&mut resumed_output).await;
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(args)) = &kv.message {
|
||||
resumed_set_blob_ids.insert(args.blob_id.clone());
|
||||
}
|
||||
acknowledge(&resumed, &mut resumed_seqno, kv.id).await;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
resumed_checkpoints.push(state);
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => break exec.id,
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
assert_eq!(
|
||||
provider.requests().len(),
|
||||
1,
|
||||
"resume must execute the pending batch before calling the model"
|
||||
);
|
||||
let resumed_run_id = store
|
||||
.active_run_for_cursor_request("resumed-run")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let resumed_round = store
|
||||
.tool_round(&ToolRoundId::new(format!(
|
||||
"{}:round:resume",
|
||||
resumed_run_id.as_str()
|
||||
)))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(resumed_round.created_at_ms, staged_started_at_ms);
|
||||
resumed
|
||||
.command(TransportCommand::Append {
|
||||
seqno: resumed_seqno,
|
||||
message: Box::new(read_result(exec_id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
resumed_seqno += 1;
|
||||
|
||||
let mut saw_settled_barrier_blob = false;
|
||||
let mut saw_settled_checkpoint = false;
|
||||
loop {
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), resumed_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)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(args)) = &kv.message {
|
||||
resumed_set_blob_ids.insert(args.blob_id.clone());
|
||||
}
|
||||
if !saw_settled_barrier_blob {
|
||||
assert_eq!(
|
||||
provider.requests().len(),
|
||||
1,
|
||||
"the next model call must wait for the settled Blob barrier"
|
||||
);
|
||||
saw_settled_barrier_blob = true;
|
||||
}
|
||||
acknowledge(&resumed, &mut resumed_seqno, kv.id).await;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
if state.pending_tool_calls.is_empty() {
|
||||
saw_settled_checkpoint = true;
|
||||
}
|
||||
resumed_checkpoints.push(state);
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update))
|
||||
if matches!(
|
||||
update.message,
|
||||
Some(pb::interaction_update::Message::TextDelta(_))
|
||||
) =>
|
||||
{
|
||||
assert!(
|
||||
saw_settled_checkpoint,
|
||||
"the next model round must not become visible before settled checkpoint"
|
||||
);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
assert!(saw_settled_barrier_blob);
|
||||
assert!(saw_settled_checkpoint);
|
||||
assert!(staged
|
||||
.root_prompt_messages_json
|
||||
.iter()
|
||||
.all(|id| !resumed_set_blob_ids.contains(id)));
|
||||
assert_eq!(provider.requests().len(), 2);
|
||||
assert!(resumed_checkpoints
|
||||
.last()
|
||||
.unwrap()
|
||||
.read_paths
|
||||
.iter()
|
||||
.any(|path| path == "/tmp/a"));
|
||||
|
||||
let mut previous_steps = Vec::new();
|
||||
let mut saw_completed_read = false;
|
||||
for state in resumed_checkpoints {
|
||||
let Some(turn_id) = state.turns.last() else {
|
||||
continue;
|
||||
};
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
store
|
||||
.get_blob(&BlobId::from_bytes(turn_id).unwrap())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent turn");
|
||||
};
|
||||
assert!(turn.steps.len() >= previous_steps.len());
|
||||
assert_eq!(
|
||||
previous_steps,
|
||||
turn.steps[..previous_steps.len()],
|
||||
"published Step BlobIDs must be an immutable prefix"
|
||||
);
|
||||
previous_steps = turn.steps.clone();
|
||||
for step_id in &turn.steps {
|
||||
let step = pb::ConversationStep::decode(
|
||||
store
|
||||
.get_blob(&BlobId::from_bytes(step_id).unwrap())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
if let Some(pb::conversation_step::Message::ToolCall(call)) = step.message {
|
||||
if call.tool_call_id.as_deref() == Some("read-1") {
|
||||
assert!(call.started_at_ms.is_some());
|
||||
assert!(call.completed_at_ms.is_some());
|
||||
saw_completed_read = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
assert!(
|
||||
saw_completed_read,
|
||||
"settled Turn must keep the typed result"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn recovery_rejects_a_kv_get_payload_whose_hash_does_not_match_the_blob_id() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("bad-blob-run").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
let expected = BlobId::digest(b"expected");
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(resume_request(pb::ConversationStateStructure {
|
||||
root_prompt_messages_json: vec![expected.as_bytes().to_vec()],
|
||||
mode: Some(pb::AgentMode::Agent as i32),
|
||||
..Default::default()
|
||||
})),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let get_id = loop {
|
||||
let server = next_message(&mut output).await;
|
||||
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = server.message {
|
||||
if matches!(
|
||||
kv.message,
|
||||
Some(pb::kv_server_message::Message::GetBlobArgs(_))
|
||||
) {
|
||||
break kv.id;
|
||||
}
|
||||
}
|
||||
};
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 1,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::KvClientMessage(
|
||||
pb::KvClientMessage {
|
||||
id: get_id,
|
||||
message: Some(pb::kv_client_message::Message::GetBlobResult(
|
||||
pb::GetBlobResult {
|
||||
blob_data: Some(b"corrupt".to_vec()),
|
||||
error: None,
|
||||
},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let error = 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 serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
|
||||
}
|
||||
};
|
||||
assert_eq!(error["error"]["code"], "invalid_argument");
|
||||
assert!(error["error"]["message"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("Blob hash mismatch"));
|
||||
}
|
||||
|
||||
fn start_request() -> 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: "user-1".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("conversation".into()),
|
||||
run_id: Some("first-run".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn resume_request(state: pb::ConversationStateStructure) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::ResumeAction(
|
||||
pb::ResumeAction::default(),
|
||||
)),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state: Some(state),
|
||||
conversation_id: Some("conversation".into()),
|
||||
run_id: Some("resumed-run".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn read_result(id: u32) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id,
|
||||
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()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn next_message(
|
||||
output: &mut tokio::sync::mpsc::UnboundedReceiver<bytes::Bytes>,
|
||||
) -> pb::AgentServerMessage {
|
||||
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();
|
||||
assert_eq!(
|
||||
flags & connect::END_STREAM_FLAG,
|
||||
0,
|
||||
"unexpected EndStream: {}",
|
||||
String::from_utf8_lossy(&payload)
|
||||
);
|
||||
pb::AgentServerMessage::decode(payload).unwrap()
|
||||
}
|
||||
|
||||
async fn acknowledge(handle: &cursor_server::cursor::TransportHandle, seqno: &mut i64, id: u32) {
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: *seqno,
|
||||
message: Box::new(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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
*seqno += 1;
|
||||
}
|
||||
@@ -0,0 +1,371 @@
|
||||
//! Verifies explicit and automatic context compaction behavior.
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{collections::HashMap, sync::Arc, time::Duration};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::{
|
||||
protocol::{connect, proto::agent::v1 as pb},
|
||||
TransportCommand, TransportRegistry,
|
||||
},
|
||||
model::{
|
||||
ContentPart, ConversationId, MessageContent, ModelConfigInput, ModelType, Origin,
|
||||
ProjectedContent, Role, Usage, OPENAI_CHAT_ENDPOINT,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn summarize_replaces_model_history_and_preserves_cursor_history() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let model = store
|
||||
.create_model(&ModelConfigInput {
|
||||
sort_order: 0,
|
||||
display_name: "Test Model".into(),
|
||||
model_type: ModelType::OpenAi,
|
||||
base_url: "https://example.com/v1/chat/completions".into(),
|
||||
use_full_url: true,
|
||||
api_key: "test-key".into(),
|
||||
tooltip_data: "Test Model".into(),
|
||||
model_id: "test-model".into(),
|
||||
reasoning_effort: None,
|
||||
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
|
||||
openai_extra_params_enabled: false,
|
||||
openai_extra_params: serde_json::json!({}),
|
||||
custom_headers_enabled: false,
|
||||
custom_headers: serde_json::json!({}),
|
||||
anthropic_extra_params_enabled: false,
|
||||
anthropic_extra_params: serde_json::json!({}),
|
||||
context_window_tokens: None,
|
||||
max_completion_tokens: None,
|
||||
anthropic_max_tokens: None,
|
||||
anthropic_thinking_effort: None,
|
||||
thinking_budget_tokens: None,
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(text_response("old answer", 4_000, 12));
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "summary-call".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("Durable ".into()),
|
||||
ModelEvent::TextDelta("summary".into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(4_012),
|
||||
output_tokens: Some(9),
|
||||
total_tokens: Some(4_021),
|
||||
..Default::default()
|
||||
}),
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]);
|
||||
provider.push(text_response("new answer", 900, 5));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
|
||||
let first = run(
|
||||
®istry,
|
||||
"first",
|
||||
user_request(
|
||||
"conversation",
|
||||
"user-1",
|
||||
"remember alpha",
|
||||
&model.model_hash,
|
||||
None,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
let first_state = first.checkpoints.last().unwrap().clone();
|
||||
let old_turns = first_state.turns.clone();
|
||||
let old_roots = first_state.root_prompt_messages_json.clone();
|
||||
assert!(old_roots.len() >= 3);
|
||||
|
||||
let compacted = run(
|
||||
®istry,
|
||||
"compact",
|
||||
summary_request("conversation", &model.model_hash, first_state),
|
||||
)
|
||||
.await;
|
||||
assert_eq!(compacted.summary_started, 1);
|
||||
assert_eq!(compacted.summary, "Durable summary");
|
||||
assert_eq!(compacted.summary_completed, 1);
|
||||
assert_eq!(compacted.turn_ended, 1);
|
||||
assert_eq!(compacted.token_delta, 0);
|
||||
assert_eq!(compacted.checkpoints.len(), 3);
|
||||
assert!(compacted
|
||||
.checkpoints
|
||||
.windows(2)
|
||||
.all(|pair| pair[0] == pair[1]));
|
||||
|
||||
let compacted_state = compacted.checkpoints.last().unwrap();
|
||||
assert_eq!(compacted_state.root_prompt_messages_json.len(), 2);
|
||||
assert!(compacted_state.turns.starts_with(&old_turns));
|
||||
assert_eq!(compacted_state.turns.len(), old_turns.len() + 1);
|
||||
assert_eq!(compacted_state.self_summary_count, 1);
|
||||
let summary_id = compacted_state.summary.as_ref().unwrap();
|
||||
let summary = pb::ConversationSummary::decode(compacted.blobs[summary_id].as_slice()).unwrap();
|
||||
assert_eq!(summary.summary, "Durable summary");
|
||||
let archive_id = compacted_state.summary_archive.as_ref().unwrap();
|
||||
let archive =
|
||||
pb::ConversationSummaryArchive::decode(compacted.blobs[archive_id].as_slice()).unwrap();
|
||||
assert_eq!(archive.summary, "Durable summary");
|
||||
assert_eq!(archive.window_tail, 0);
|
||||
assert_eq!(archive.summarized_messages, old_roots[1..]);
|
||||
assert_eq!(
|
||||
archive.summary_message,
|
||||
*compacted_state.root_prompt_messages_json.last().unwrap()
|
||||
);
|
||||
|
||||
let stored = store
|
||||
.load_current_messages(&ConversationId::new("conversation"))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(stored.len(), 1);
|
||||
assert_eq!(stored[0].origin, Origin::Runtime);
|
||||
assert_eq!(stored[0].role, Role::User);
|
||||
assert!(matches!(
|
||||
&stored[0].content,
|
||||
MessageContent::Parts { parts }
|
||||
if matches!(parts.as_slice(), [ContentPart::Text { text }]
|
||||
if text == "<conversation_summary>\nDurable summary\n</conversation_summary>")
|
||||
));
|
||||
|
||||
let after = run(
|
||||
®istry,
|
||||
"after",
|
||||
user_request(
|
||||
"conversation",
|
||||
"user-2",
|
||||
"what remains?",
|
||||
&model.model_hash,
|
||||
Some(compacted_state.clone()),
|
||||
),
|
||||
)
|
||||
.await;
|
||||
assert!(after
|
||||
.checkpoints
|
||||
.last()
|
||||
.unwrap()
|
||||
.root_prompt_messages_json
|
||||
.starts_with(&compacted_state.root_prompt_messages_json));
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 3);
|
||||
assert!(requests[1].prompt.tools.is_empty());
|
||||
assert!(requests[1]
|
||||
.prompt
|
||||
.instructions
|
||||
.contains("compacting conversation history"));
|
||||
assert_eq!(requests[1].history.len(), 2);
|
||||
assert_eq!(requests[2].history.len(), 2);
|
||||
let ProjectedContent::Parts(summary_parts) = &requests[2].history[0].content else {
|
||||
panic!("first post-compaction message must be the summary")
|
||||
};
|
||||
assert!(
|
||||
matches!(summary_parts.as_slice(), [ContentPart::Text { text }]
|
||||
if text.contains("Durable summary"))
|
||||
);
|
||||
let ProjectedContent::Parts(new_user_parts) = &requests[2].history[1].content else {
|
||||
panic!("second post-compaction message must be the new runtime user")
|
||||
};
|
||||
assert!(
|
||||
matches!(new_user_parts.as_slice(), [ContentPart::Text { text }]
|
||||
if text.contains("what remains?") && !text.contains("remember alpha"))
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct Output {
|
||||
checkpoints: Vec<pb::ConversationStateStructure>,
|
||||
blobs: HashMap<Vec<u8>, Vec<u8>>,
|
||||
summary: String,
|
||||
summary_started: usize,
|
||||
summary_completed: usize,
|
||||
turn_ended: usize,
|
||||
token_delta: usize,
|
||||
}
|
||||
|
||||
async fn run(
|
||||
registry: &TransportRegistry,
|
||||
request_id: &str,
|
||||
request: pb::AgentClientMessage,
|
||||
) -> Output {
|
||||
let handle = registry.get_or_create(request_id).await.unwrap();
|
||||
let mut receiver = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(request),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut append_seqno = 1;
|
||||
let mut output = Output::default();
|
||||
loop {
|
||||
let frame = tokio::time::timeout(Duration::from_secs(5), receiver.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
|
||||
if flags & connect::END_STREAM_FLAG != 0 {
|
||||
return output;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = kv.message {
|
||||
output.blobs.insert(set.blob_id, set.blob_data);
|
||||
}
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
output.checkpoints.push(state)
|
||||
}
|
||||
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
|
||||
match update.message {
|
||||
Some(pb::interaction_update::Message::SummaryStarted(_)) => {
|
||||
output.summary_started += 1
|
||||
}
|
||||
Some(pb::interaction_update::Message::Summary(delta)) => {
|
||||
output.summary.push_str(&delta.summary)
|
||||
}
|
||||
Some(pb::interaction_update::Message::SummaryCompleted(_)) => {
|
||||
output.summary_completed += 1
|
||||
}
|
||||
Some(pb::interaction_update::Message::TurnEnded(_)) => output.turn_ended += 1,
|
||||
Some(pb::interaction_update::Message::TokenDelta(_)) => output.token_delta += 1,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn text_response(text: &str, input: u64, output: u64) -> Vec<ModelEvent> {
|
||||
vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: format!("call-{text}"),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta(text.into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Usage(Usage {
|
||||
input_tokens: Some(input),
|
||||
output_tokens: Some(output),
|
||||
total_tokens: Some(input + output),
|
||||
..Default::default()
|
||||
}),
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]
|
||||
}
|
||||
|
||||
fn user_request(
|
||||
conversation_id: &str,
|
||||
message_id: &str,
|
||||
text: &str,
|
||||
model_id: &str,
|
||||
state: Option<pb::ConversationStateStructure>,
|
||||
) -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: text.into(),
|
||||
message_id: message_id.into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
request(
|
||||
conversation_id,
|
||||
model_id,
|
||||
state,
|
||||
pb::conversation_action::Action::UserMessageAction(pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
request_context: Some(pb::RequestContext::default()),
|
||||
..Default::default()
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn summary_request(
|
||||
conversation_id: &str,
|
||||
model_id: &str,
|
||||
state: pb::ConversationStateStructure,
|
||||
) -> pb::AgentClientMessage {
|
||||
let user = pb::UserMessage {
|
||||
text: "/summarize".into(),
|
||||
message_id: "summary-command".into(),
|
||||
mode: pb::AgentMode::Agent as i32,
|
||||
..Default::default()
|
||||
};
|
||||
request(
|
||||
conversation_id,
|
||||
model_id,
|
||||
Some(state),
|
||||
pb::conversation_action::Action::UserMessageAction(pb::UserMessageAction {
|
||||
user_message: Some(user),
|
||||
request_context: Some(pb::RequestContext::default()),
|
||||
..Default::default()
|
||||
}),
|
||||
)
|
||||
}
|
||||
|
||||
fn request(
|
||||
conversation_id: &str,
|
||||
model_id: &str,
|
||||
state: Option<pb::ConversationStateStructure>,
|
||||
action: pb::conversation_action::Action,
|
||||
) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: model_id.into(),
|
||||
..Default::default()
|
||||
}),
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(action),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some(conversation_id.into()),
|
||||
conversation_state: state,
|
||||
run_id: Some("reusable-wire-run-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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
//! Verifies captured Cursor Connect framing and protobuf compatibility.
|
||||
#[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::{
|
||||
api::cursor,
|
||||
cursor::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::protocol::{
|
||||
connect,
|
||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||
},
|
||||
cursor::transport::TransportRegistry,
|
||||
};
|
||||
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/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
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 = cursor::router(registry)
|
||||
.unwrap()
|
||||
.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"));
|
||||
}
|
||||
@@ -0,0 +1,664 @@
|
||||
//! Verifies message delivery before, during, and after a Run.
|
||||
#[path = "support/fake_provider.rs"]
|
||||
mod fake_provider;
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
|
||||
use cursor_server::{
|
||||
cursor::{
|
||||
prompting::{PromptAssets, PromptCompiler},
|
||||
protocol::connect,
|
||||
protocol::proto::agent::v1 as pb,
|
||||
TransportCommand, TransportHandle, TransportRegistry,
|
||||
},
|
||||
model::{ContentPart, MessageContent, ProjectedContent, Role},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
const FOLLOW_UP: &str = "Perform any necessary follow-up actions in response to the subagent completion above. If no follow-up work is needed, no further action is required. If you mention an agent or subagent in your response, link it with the `[Name](id)` Don't use generic label such as `[agent]`, `[worker]`, or `[subagent]`.";
|
||||
const SHELL_FOLLOW_UP: &str = "Briefly inform the user about the task result and perform any follow-up actions (if needed). If there's no follow-ups needed, don't explicitly say that.";
|
||||
|
||||
#[tokio::test]
|
||||
async fn background_subagent_completion_starts_a_simulated_parent_turn() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(stop_response("model-call", "followed up"));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("completion-request").await.unwrap();
|
||||
let (checkpoint, blobs) = drive_completion(
|
||||
&handle,
|
||||
completion_run(
|
||||
"child-id",
|
||||
"reusable-parent-run",
|
||||
pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Multitask as i32),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
)
|
||||
.await;
|
||||
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 1);
|
||||
let [runtime] = requests[0].history.as_slice() else {
|
||||
panic!("completion Run must add exactly one runtime message")
|
||||
};
|
||||
assert_eq!(runtime.role, Role::User);
|
||||
let ProjectedContent::Parts(parts) = &runtime.content else {
|
||||
panic!("completion context must be text")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("completion context must have one text part")
|
||||
};
|
||||
assert!(text.contains("kind: subagent"));
|
||||
assert!(text.contains("agent_id: child-id"));
|
||||
assert!(text.contains("child result"));
|
||||
assert!(text.contains(FOLLOW_UP));
|
||||
|
||||
let messages = store
|
||||
.load_current_messages(&cursor_server::model::ConversationId::new(
|
||||
"parent-conversation",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(messages.iter().any(|message| {
|
||||
message.runtime_event_id.as_deref()
|
||||
== Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id:task-call")
|
||||
&& matches!(&message.content, MessageContent::Parts { parts } if !parts.is_empty())
|
||||
}));
|
||||
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
blobs
|
||||
.get(checkpoint.turns.last().expect("completion Turn"))
|
||||
.expect("completion Turn Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent conversation Turn")
|
||||
};
|
||||
let user = pb::UserMessage::decode(
|
||||
blobs
|
||||
.get(&turn.user_message)
|
||||
.expect("simulated UserMessage Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(user.text.contains(FOLLOW_UP));
|
||||
assert_eq!(user.is_simulated_msg, Some(true));
|
||||
assert_eq!(
|
||||
user.simulated_msg_reason,
|
||||
Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32)
|
||||
);
|
||||
assert_eq!(
|
||||
user.simulated_message_metadata.unwrap().task_id.as_deref(),
|
||||
Some("child-id")
|
||||
);
|
||||
|
||||
provider.push(stop_response("model-call-2", "followed up again"));
|
||||
let second = registry
|
||||
.get_or_create("completion-request-2")
|
||||
.await
|
||||
.unwrap();
|
||||
drive_completion(
|
||||
&second,
|
||||
completion_run("child-id-2", "reusable-parent-run-2", checkpoint),
|
||||
)
|
||||
.await;
|
||||
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
let runtime_ids = requests[1]
|
||||
.history
|
||||
.iter()
|
||||
.map(|message| message.message_id.as_str())
|
||||
.filter(|id| id.starts_with("runtime:"))
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
runtime_ids,
|
||||
[
|
||||
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id:task-call",
|
||||
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id-2:task-call"
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn background_completion_joins_the_active_run_instead_of_replacing_it() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
let first_ready = provider.push_gated(stop_response("model-call-1", "first response"));
|
||||
provider.push(stop_response("model-call-2", "processed both completions"));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let first = registry.get_or_create("active-completion-1").await.unwrap();
|
||||
let first_run = tokio::spawn(async move {
|
||||
drive_completion(
|
||||
&first,
|
||||
completion_run(
|
||||
"child-1",
|
||||
"parent-run-1",
|
||||
pb::ConversationStateStructure::default(),
|
||||
),
|
||||
)
|
||||
.await
|
||||
});
|
||||
while provider.requests().is_empty() {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
|
||||
let second = registry.get_or_create("active-completion-2").await.unwrap();
|
||||
let second_run = tokio::spawn(async move {
|
||||
drive_forwarded_completion(
|
||||
&second,
|
||||
completion_run(
|
||||
"child-2",
|
||||
"parent-run-2",
|
||||
pb::ConversationStateStructure::default(),
|
||||
),
|
||||
)
|
||||
.await
|
||||
});
|
||||
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
|
||||
first_ready.notify_one();
|
||||
|
||||
second_run.await.unwrap();
|
||||
first_run.await.unwrap();
|
||||
let requests = provider.requests();
|
||||
assert_eq!(requests.len(), 2);
|
||||
let history = serde_json::to_string(&requests[1].history).unwrap();
|
||||
assert!(history.contains("child-1"));
|
||||
assert!(history.contains("first response"));
|
||||
assert!(history.contains("child-2"));
|
||||
let statuses: Vec<String> = sqlx::query_scalar(
|
||||
"SELECT status FROM runs WHERE conversation_id = 'parent-conversation' ORDER BY created_at_ms",
|
||||
)
|
||||
.fetch_all(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(statuses, ["completed"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn retrying_one_background_completion_reuses_its_runtime_message() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(stop_response("model-call", "followed up"));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let first = registry.get_or_create("completion-retry-1").await.unwrap();
|
||||
let (checkpoint, _) = drive_completion(
|
||||
&first,
|
||||
completion_run(
|
||||
"retry-child",
|
||||
"completion-retry-run-1",
|
||||
pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Multitask as i32),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
)
|
||||
.await;
|
||||
|
||||
provider.push(stop_response("model-call-2", "followed up again"));
|
||||
let second = registry.get_or_create("completion-retry-2").await.unwrap();
|
||||
drive_completion(
|
||||
&second,
|
||||
completion_run("retry-child", "completion-retry-run-2", checkpoint),
|
||||
)
|
||||
.await;
|
||||
|
||||
let messages = store
|
||||
.load_current_messages(&cursor_server::model::ConversationId::new(
|
||||
"parent-conversation",
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
messages
|
||||
.iter()
|
||||
.filter(|message| {
|
||||
message.runtime_event_id.as_deref()
|
||||
== Some(
|
||||
"background-completed:BACKGROUND_TASK_KIND_SUBAGENT:retry-child:task-call",
|
||||
)
|
||||
})
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn background_shell_completion_wakes_the_parent_with_the_captured_notification() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(stop_response(
|
||||
"shell-wakeup",
|
||||
"The background server was stopped.",
|
||||
));
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("shell-completion-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let (checkpoint, blobs) = drive_completion(
|
||||
&handle,
|
||||
shell_completion_run(pb::ConversationStateStructure {
|
||||
mode: Some(pb::AgentMode::Agent as i32),
|
||||
..Default::default()
|
||||
}),
|
||||
)
|
||||
.await;
|
||||
|
||||
let requests = provider.requests();
|
||||
let [runtime] = requests[0].history.as_slice() else {
|
||||
panic!("Shell completion Run must add exactly one runtime message")
|
||||
};
|
||||
let ProjectedContent::Parts(parts) = &runtime.content else {
|
||||
panic!("Shell completion context must be text")
|
||||
};
|
||||
let [ContentPart::Text { text }] = parts.as_slice() else {
|
||||
panic!("Shell completion context must have one text part")
|
||||
};
|
||||
assert!(text.contains("<system_notification>"));
|
||||
assert!(text.contains("kind: shell"));
|
||||
assert!(text.contains("status: aborted"));
|
||||
assert!(text.contains("task_id: 977679"));
|
||||
assert!(text.contains("detail: terminated_by_user"));
|
||||
assert!(text.contains("output_path: /tmp/977679.txt"));
|
||||
assert!(text.contains(SHELL_FOLLOW_UP));
|
||||
assert!(text.starts_with("<timestamp>"));
|
||||
assert!(!text.contains("You are still in **Agent Mode**"));
|
||||
assert!(text.find("<system_notification>").unwrap() < text.find("<user_query>").unwrap());
|
||||
|
||||
let turn = pb::ConversationTurnStructure::decode(
|
||||
blobs
|
||||
.get(checkpoint.turns.last().expect("Shell completion Turn"))
|
||||
.expect("Shell completion Turn Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
|
||||
else {
|
||||
panic!("expected agent conversation Turn")
|
||||
};
|
||||
let user = pb::UserMessage::decode(
|
||||
blobs
|
||||
.get(&turn.user_message)
|
||||
.expect("simulated Shell UserMessage Blob")
|
||||
.as_slice(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(user.text, *text);
|
||||
assert_eq!(user.is_simulated_msg, Some(true));
|
||||
assert_eq!(
|
||||
user.simulated_msg_reason,
|
||||
Some(pb::SimulatedMsgReason::BackgroundTaskCompletion as i32)
|
||||
);
|
||||
let metadata = user.simulated_message_metadata.unwrap();
|
||||
assert_eq!(
|
||||
metadata.title.as_deref(),
|
||||
Some("Start Python HTTP server on 9000")
|
||||
);
|
||||
assert_eq!(metadata.task_id.as_deref(), Some("977679"));
|
||||
}
|
||||
|
||||
async fn drive_completion(
|
||||
handle: &TransportHandle,
|
||||
message: pb::AgentClientMessage,
|
||||
) -> (pb::ConversationStateStructure, HashMap<Vec<u8>, Vec<u8>>) {
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(message),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut blobs = HashMap::new();
|
||||
let mut final_checkpoint = None;
|
||||
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::ExecServerMessage(exec)) => {
|
||||
assert_eq!(exec.id, 0);
|
||||
assert!(matches!(
|
||||
exec.message,
|
||||
Some(pb::exec_server_message::Message::RequestContextArgs(_))
|
||||
));
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 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: 0 },
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
}),
|
||||
})
|
||||
.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(
|
||||
pb::ExecClientMessage {
|
||||
id: 0,
|
||||
message: Some(
|
||||
pb::exec_client_message::Message::RequestContextResult(
|
||||
pb::RequestContextResult {
|
||||
result: Some(
|
||||
pb::request_context_result::Result::Success(
|
||||
pb::RequestContextSuccess {
|
||||
request_context: Some(
|
||||
pb::RequestContext::default(),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
if let Some(pb::kv_server_message::Message::SetBlobArgs(set)) = &kv.message {
|
||||
blobs.insert(set.blob_id.clone(), set.blob_data.clone());
|
||||
}
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state))
|
||||
if state.pending_tool_calls.is_empty() =>
|
||||
{
|
||||
final_checkpoint = Some(state);
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
(
|
||||
final_checkpoint.expect("settled completion checkpoint"),
|
||||
blobs,
|
||||
)
|
||||
}
|
||||
|
||||
async fn drive_forwarded_completion(handle: &TransportHandle, message: pb::AgentClientMessage) {
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(message),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let mut append_seqno = 1;
|
||||
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!(payload.as_ref(), b"{}");
|
||||
return;
|
||||
}
|
||||
let server = pb::AgentServerMessage::decode(payload).unwrap();
|
||||
match server.message {
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
|
||||
assert_eq!(exec.id, 0);
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 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: 0 },
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
}),
|
||||
})
|
||||
.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(
|
||||
pb::ExecClientMessage {
|
||||
id: 0,
|
||||
message: Some(
|
||||
pb::exec_client_message::Message::RequestContextResult(
|
||||
pb::RequestContextResult {
|
||||
result: Some(
|
||||
pb::request_context_result::Result::Success(
|
||||
pb::RequestContextSuccess {
|
||||
request_context: Some(
|
||||
pb::RequestContext::default(),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
),
|
||||
),
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn completion_run(
|
||||
child_id: &str,
|
||||
run_id: &str,
|
||||
conversation_state: pb::ConversationStateStructure,
|
||||
) -> pb::AgentClientMessage {
|
||||
completion_run_with_detail(child_id, run_id, conversation_state, "child result")
|
||||
}
|
||||
|
||||
fn completion_run_with_detail(
|
||||
child_id: &str,
|
||||
run_id: &str,
|
||||
conversation_state: pb::ConversationStateStructure,
|
||||
detail: &str,
|
||||
) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(
|
||||
pb::conversation_action::Action::BackgroundTaskCompletionAction(
|
||||
pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![pb::BackgroundTaskCompletion {
|
||||
task_id: child_id.into(),
|
||||
kind: pb::BackgroundTaskKind::Subagent as i32,
|
||||
status: pb::BackgroundTaskStatus::Success as i32,
|
||||
title: "Inspect protocol".into(),
|
||||
detail: Some(detail.into()),
|
||||
output_path: Some("/tmp/child.jsonl".into()),
|
||||
reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32,
|
||||
subagent_id: Some(child_id.into()),
|
||||
tool_call_id: Some("task-call".into()),
|
||||
..Default::default()
|
||||
}],
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("parent-conversation".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state: Some(conversation_state),
|
||||
run_id: Some(run_id.into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn shell_completion_run(
|
||||
conversation_state: pb::ConversationStateStructure,
|
||||
) -> pb::AgentClientMessage {
|
||||
pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::RunRequest(
|
||||
pb::AgentRunRequest {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(
|
||||
pb::conversation_action::Action::BackgroundTaskCompletionAction(
|
||||
pb::BackgroundTaskCompletionAction {
|
||||
completions: vec![pb::BackgroundTaskCompletion {
|
||||
task_id: "977679".into(),
|
||||
kind: pb::BackgroundTaskKind::Shell as i32,
|
||||
status: pb::BackgroundTaskStatus::Aborted as i32,
|
||||
title: "Start Python HTTP server on 9000".into(),
|
||||
detail: Some("terminated_by_user".into()),
|
||||
output_path: Some("/tmp/977679.txt".into()),
|
||||
reason: pb::BackgroundTaskCompletionReason::TaskFinished as i32,
|
||||
tool_call_id: Some("shell-call".into()),
|
||||
..Default::default()
|
||||
}],
|
||||
},
|
||||
),
|
||||
),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_id: Some("parent-conversation".into()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
conversation_state: Some(conversation_state),
|
||||
run_id: Some("shell-parent-run".into()),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn stop_response(model_call_id: &str, text: &str) -> Vec<ModelEvent> {
|
||||
vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: model_call_id.into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta(text.into()),
|
||||
ModelEvent::TextEnd,
|
||||
ModelEvent::Done(FinishReason::Stop),
|
||||
]
|
||||
}
|
||||
|
||||
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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,641 @@
|
||||
//! Verifies unique persisted and streamed terminal outcomes.
|
||||
#[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::prompting::{PromptAssets, PromptCompiler},
|
||||
cursor::protocol::{
|
||||
connect,
|
||||
proto::{agent::v1 as pb, aiserver::v1 as ai},
|
||||
},
|
||||
cursor::{TransportCommand, TransportParent, TransportRegistry},
|
||||
model::{MessageContent, Role},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
Error,
|
||||
};
|
||||
use prost::Message;
|
||||
|
||||
#[tokio::test]
|
||||
async fn abort_command_cancels_the_run_and_closes_output() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let registry = TransportRegistry::new(
|
||||
store,
|
||||
Arc::new(fake_provider::FakeProvider::default()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("abort-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
|
||||
handle.command(TransportCommand::Disconnect).await.unwrap();
|
||||
|
||||
let frame = tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("Abort must emit a terminal frame");
|
||||
let (flags, payload) = connect::decode_frames(&frame).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 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_error(Error::Provider("provider failed".into()));
|
||||
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),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("failed-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(client_run()),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut append_seqno = 1;
|
||||
let mut checkpoints = Vec::new();
|
||||
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(TransportCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
append_seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
|
||||
checkpoints.push(state);
|
||||
}
|
||||
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_eq!(
|
||||
checkpoints.len(),
|
||||
1,
|
||||
"the initial user state is checkpointed"
|
||||
);
|
||||
assert!(checkpoints[0].pending_tool_calls.is_empty());
|
||||
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
|
||||
);
|
||||
assert_eq!(decoded.is_expected, Some(true));
|
||||
let custom = decoded.details.unwrap();
|
||||
assert_eq!(custom.title, "Provider 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_current_messages(&cursor_server::model::ConversationId::new(
|
||||
"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![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".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),
|
||||
]);
|
||||
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),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("protocol-failed-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(protocol_client_run("read it", "protocol-failed-user")),
|
||||
})
|
||||
.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(TransportCommand::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(TransportCommand::Append {
|
||||
seqno: append_seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: exec.id + 1_000,
|
||||
exec_id: String::new(),
|
||||
message: None,
|
||||
..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"],
|
||||
"unknown ExecClientMessage id: 1001"
|
||||
);
|
||||
assert_eq!(
|
||||
tokio::time::timeout(std::time::Duration::from_secs(1), output.recv())
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
|
||||
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(1);
|
||||
let (status, failure_summary) = loop {
|
||||
let row: (String, Option<String>) =
|
||||
sqlx::query_as("SELECT status, failure_summary FROM runs WHERE cursor_request_id = ?")
|
||||
.bind("protocol-failed-request")
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
if row.0 != "running" {
|
||||
break row;
|
||||
}
|
||||
assert!(
|
||||
tokio::time::Instant::now() < deadline,
|
||||
"Run remained running after the Cursor session failed"
|
||||
);
|
||||
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
|
||||
};
|
||||
assert_eq!(status, "failed");
|
||||
assert_eq!(
|
||||
failure_summary.as_deref(),
|
||||
Some("unknown ExecClientMessage id: 1001")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn newer_run_request_on_one_bidi_stream_replaces_the_active_run() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".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: "replacement-model-call".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("replacement completed".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 = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry
|
||||
.get_or_create("protocol-failed-request")
|
||||
.await
|
||||
.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(protocol_client_run("read it", "first-user")),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
let mut replacement_sent = false;
|
||||
let mut late_result_sent = false;
|
||||
let mut saw_abort = false;
|
||||
let mut cropped_state = None;
|
||||
let terminal_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(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(mut state)) => {
|
||||
state.root_prompt_messages_json.truncate(1);
|
||||
state.turns.clear();
|
||||
state.pending_tool_calls.clear();
|
||||
cropped_state = Some(state);
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec))
|
||||
if !replacement_sent =>
|
||||
{
|
||||
replacement_sent = true;
|
||||
let mut replacement =
|
||||
protocol_client_run("use the cropped history", "replacement-user");
|
||||
let Some(pb::agent_client_message::Message::RunRequest(request)) =
|
||||
replacement.message.as_mut()
|
||||
else {
|
||||
unreachable!()
|
||||
};
|
||||
request.conversation_state = Some(
|
||||
cropped_state
|
||||
.clone()
|
||||
.expect("first Run must publish a checkpoint before tools"),
|
||||
);
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(replacement),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(pb::AgentClientMessage {
|
||||
message: Some(pb::agent_client_message::Message::ExecClientMessage(
|
||||
pb::ExecClientMessage {
|
||||
id: exec.id,
|
||||
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||
pb::ReadResult::default(),
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
late_result_sent = true;
|
||||
}
|
||||
Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) => {
|
||||
if matches!(
|
||||
control.message,
|
||||
Some(pb::exec_server_control_message::Message::Abort(_))
|
||||
) {
|
||||
saw_abort = true;
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
};
|
||||
|
||||
assert!(replacement_sent);
|
||||
assert!(late_result_sent);
|
||||
assert!(saw_abort);
|
||||
assert!(terminal_json.get("error").is_none(), "{terminal_json}");
|
||||
assert_eq!(provider.requests().len(), 2);
|
||||
let replacement_history = serde_json::to_string(&provider.requests()[1].history).unwrap();
|
||||
assert!(replacement_history.contains("use the cropped history"));
|
||||
assert!(!replacement_history.contains("read it"));
|
||||
let statuses: Vec<String> = sqlx::query_scalar(
|
||||
"SELECT status FROM runs WHERE cursor_request_id = ? ORDER BY created_at_ms, run_id",
|
||||
)
|
||||
.bind("protocol-failed-request")
|
||||
.fetch_all(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(statuses, ["cancelled", "completed"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn parent_request_does_not_need_to_resolve_to_an_active_run() {
|
||||
assert_run_starts_without_parent_dependency(
|
||||
"finished-parent-request",
|
||||
Some(TransportParent {
|
||||
request_id: "already-finished-parent".into(),
|
||||
tool_call_id: "original-tool-call".into(),
|
||||
}),
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn subagent_type_does_not_require_parent_metadata() {
|
||||
assert_run_starts_without_parent_dependency(
|
||||
"parentless-subagent-request",
|
||||
None,
|
||||
Some("generalPurpose"),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
async fn assert_run_starts_without_parent_dependency(
|
||||
request_id: &str,
|
||||
parent: Option<TransportParent>,
|
||||
subagent_type_name: Option<&str>,
|
||||
) {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
let provider = fake_provider::FakeProvider::default();
|
||||
provider.push(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "independent-model-call".into(),
|
||||
},
|
||||
ModelEvent::TextStart,
|
||||
ModelEvent::TextDelta("continued independently".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 = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create(request_id).await.unwrap();
|
||||
if let Some(parent) = parent {
|
||||
handle.set_parent(parent).unwrap();
|
||||
}
|
||||
let mut output = handle.subscribe();
|
||||
let mut message = protocol_client_run("continue", "independent-user");
|
||||
let Some(pb::agent_client_message::Message::RunRequest(request)) = message.message.as_mut()
|
||||
else {
|
||||
unreachable!()
|
||||
};
|
||||
request.conversation_id = Some(format!("{request_id}-conversation"));
|
||||
request.subagent_type_name = subagent_type_name.map(str::to_owned);
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno: 0,
|
||||
message: Box::new(message),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let mut seqno = 1;
|
||||
let terminal_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();
|
||||
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = server.message {
|
||||
handle
|
||||
.command(TransportCommand::Append {
|
||||
seqno,
|
||||
message: Box::new(kv_ack(kv.id)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
seqno += 1;
|
||||
}
|
||||
};
|
||||
|
||||
assert!(terminal_json.get("error").is_none(), "{terminal_json}");
|
||||
assert_eq!(provider.requests().len(), 1);
|
||||
let row: (String, Option<String>, Option<String>) = sqlx::query_as(
|
||||
"SELECT run_kind, parent_run_id, parent_tool_call_id FROM runs WHERE cursor_request_id = ?",
|
||||
)
|
||||
.bind(request_id)
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(row, ("root".into(), None, 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()),
|
||||
requested_model: Some(pb::RequestedModel {
|
||||
model_id: "test-model".into(),
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn protocol_client_run(text: &str, message_id: &str) -> 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: text.into(),
|
||||
message_id: message_id.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()),
|
||||
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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,635 @@
|
||||
//! Verifies append-only provider history and stable prompt prefixes.
|
||||
#[path = "support/fixtures.rs"]
|
||||
mod fixtures;
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use cursor_server::{
|
||||
cursor::prompting::{Mode, PromptAssets, PromptCompiler},
|
||||
model::{project_messages, ProjectedContent},
|
||||
model::{
|
||||
CanonicalMessage, MessageContent, ModelSpec, Origin, Role, ToolCallContent, ToolDefinition,
|
||||
ToolResultContent,
|
||||
},
|
||||
};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
#[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 ProjectedContent::ToolResult(object_result) = &projected[0].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
let object_text = &object_result.content;
|
||||
assert_eq!(
|
||||
serde_json::from_str::<serde_json::Value>(object_text).unwrap(),
|
||||
object
|
||||
);
|
||||
let ProjectedContent::ToolResult(string_result) = &projected[1].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
assert_eq!(string_result.content, "plain text");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn projected_tool_result_prefixes_remain_stable() {
|
||||
let first = vec![named_tool_result("Grep", &"x".repeat(64 * 1024))];
|
||||
let mut second = first.clone();
|
||||
second.push(fixtures::user("u2", "continue"));
|
||||
|
||||
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 unbounded_tool_results_are_not_rewritten() {
|
||||
let original = "x".repeat(64 * 1024);
|
||||
let projected = project_messages(&[named_tool_result("Delete", &original)]).unwrap();
|
||||
let ProjectedContent::ToolResult(result) = &projected[0].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
assert_eq!(result.content, original);
|
||||
}
|
||||
|
||||
#[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(),
|
||||
tool_round_id: Some("round".into()),
|
||||
replay_state: None,
|
||||
tool_calls: Vec::new(),
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}];
|
||||
|
||||
let projected = project_messages(&messages).unwrap();
|
||||
let ProjectedContent::Assistant { text, thinking, .. } = &projected[0].content else {
|
||||
panic!("expected assistant")
|
||||
};
|
||||
assert_eq!(text, "visible answer");
|
||||
assert_eq!(thinking, "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, Role::Assistant);
|
||||
let ProjectedContent::Assistant {
|
||||
thinking, calls, ..
|
||||
} = &projected[0].content
|
||||
else {
|
||||
panic!("expected assistant")
|
||||
};
|
||||
assert_eq!(thinking, "complete reasoning");
|
||||
assert_eq!(calls[0].call_id, "call-first");
|
||||
assert_eq!(calls[1].call_id, "call-second");
|
||||
let ProjectedContent::ToolResult(second) = &projected[1].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
let ProjectedContent::ToolResult(first) = &projected[2].content else {
|
||||
panic!("expected tool result")
|
||||
};
|
||||
assert_eq!(second.call_id, "call-second");
|
||||
assert_eq!(first.call_id, "call-first");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(assets.mode(Mode::Agent).tools.len(), 21);
|
||||
assert_eq!(
|
||||
assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"Shell",
|
||||
"Grep",
|
||||
"Delete",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"GenerateImage",
|
||||
"EditNotebook",
|
||||
"TodoWrite",
|
||||
"StrReplace",
|
||||
"Write",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Glob",
|
||||
"AskQuestion",
|
||||
"Task",
|
||||
"GetMcpTools",
|
||||
"FetchMcpResource",
|
||||
"SwitchMode",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
]
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Ask,
|
||||
&[
|
||||
"AskQuestion",
|
||||
"CallMcpTool",
|
||||
"Delete",
|
||||
"FetchMcpResource",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Shell",
|
||||
"StrReplace",
|
||||
"Task",
|
||||
"TodoWrite",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Write",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"98bb57a9ade7f1a572c5c5fe77a905a129d28ecfd42b8d318250f6486b09e1ec",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Plan,
|
||||
&[
|
||||
"Shell",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"TodoWrite",
|
||||
"ReadLints",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"AskQuestion",
|
||||
"CreatePlan",
|
||||
"Task",
|
||||
"FetchMcpResource",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"9a7e0f9e0bd8ef0af01032fa311686f72c42ec260e3057f6fae5e68f5ed36fb8",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Debug,
|
||||
&[
|
||||
"AskQuestion",
|
||||
"CallMcpTool",
|
||||
"Delete",
|
||||
"FetchMcpResource",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Shell",
|
||||
"StrReplace",
|
||||
"Task",
|
||||
"TodoWrite",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Write",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"98bb57a9ade7f1a572c5c5fe77a905a129d28ecfd42b8d318250f6486b09e1ec",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Multitask,
|
||||
&[
|
||||
"AskQuestion",
|
||||
"CallMcpTool",
|
||||
"Delete",
|
||||
"FetchMcpResource",
|
||||
"Glob",
|
||||
"Grep",
|
||||
"Read",
|
||||
"ReadLints",
|
||||
"Shell",
|
||||
"StrReplace",
|
||||
"SwitchMode",
|
||||
"Task",
|
||||
"TodoWrite",
|
||||
"WebFetch",
|
||||
"WebSearch",
|
||||
"Write",
|
||||
"GenerateImage",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"976b309dd91e314d4916439ebb9da8995751d011532e39934a1da7593dc78ccb",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Subagent,
|
||||
&[
|
||||
"Shell",
|
||||
"Grep",
|
||||
"Delete",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"GenerateImage",
|
||||
"ReadLints",
|
||||
"EditNotebook",
|
||||
"TodoWrite",
|
||||
"StrReplace",
|
||||
"Write",
|
||||
"Read",
|
||||
"Glob",
|
||||
"GetMcpTools",
|
||||
"FetchMcpResource",
|
||||
"SwitchMode",
|
||||
"UpdateCurrentStep",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"6de1ee86a131ca093c7143f54fffcba2fc14b32ff45fd6f5e0df1347058ad744",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
Mode::Compaction,
|
||||
&[],
|
||||
"4f53cda18c2baa0c0354bb5f9a3ecbe5ed12ab4d8e11ba873c2f11161202b945",
|
||||
);
|
||||
assert_eq!(
|
||||
schema_digest(&assets.mode(Mode::Agent).tools),
|
||||
"282a1dff7957090d0a75eac4a46474ac7cffa1b0937bdf97354544e729bb15c2"
|
||||
);
|
||||
let task = assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
.iter()
|
||||
.find(|tool| tool.name == "Task")
|
||||
.unwrap();
|
||||
assert!(task.description.contains(
|
||||
"When the user does not specify a number, launch at most three subagents in a single response. If the user explicitly requests more, you may launch the requested number."
|
||||
));
|
||||
assert!(task.description.contains(
|
||||
"If the user explicitly requests parallel subagents, follow the number requested by the user."
|
||||
));
|
||||
assert!(!task
|
||||
.description
|
||||
.chars()
|
||||
.any(|character| ('\u{4e00}'..='\u{9fff}').contains(&character)));
|
||||
let shell = assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
.iter()
|
||||
.find(|tool| tool.name == "Shell")
|
||||
.unwrap();
|
||||
assert!(
|
||||
shell.parameters["properties"]["block_until_ms"]["description"]
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("do not combine it with `nohup`, `&`, `disown`")
|
||||
);
|
||||
for mode in [
|
||||
Mode::Agent,
|
||||
Mode::Ask,
|
||||
Mode::Debug,
|
||||
Mode::Multitask,
|
||||
Mode::Subagent,
|
||||
Mode::Compaction,
|
||||
] {
|
||||
assert!(!assets
|
||||
.mode(mode)
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "CreatePlan" || tool.name == "PatchEdit"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn every_captured_mode_owns_and_renders_its_runtime_template() {
|
||||
let compiler = PromptCompiler::new(
|
||||
PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap(),
|
||||
);
|
||||
let values = BTreeMap::from([
|
||||
("OPEN_FILES", String::new()),
|
||||
("SELECTED_CONTEXT", String::new()),
|
||||
("ACTION_CONTEXT", String::new()),
|
||||
("TIMESTAMP", "Sunday, Aug 16, 2026, 11:31 PM (UTC+8)".into()),
|
||||
("USER_QUERY", "question".into()),
|
||||
("DEBUG_SERVER_ENDPOINT", "http://debug".into()),
|
||||
("DEBUG_LOG_PATH", "/tmp/debug.log".into()),
|
||||
("DEBUG_SESSION_ID", "session".into()),
|
||||
]);
|
||||
for (mode, marker) in [
|
||||
(Mode::Agent, "You are still in **Agent Mode**"),
|
||||
(Mode::Ask, "Ask mode is active."),
|
||||
(Mode::Plan, "Plan mode is active."),
|
||||
(Mode::Debug, "You are now in **DEBUG MODE**"),
|
||||
(Mode::Multitask, "The user has engaged **Multitask Mode**"),
|
||||
] {
|
||||
let rendered = compiler.runtime_message(mode, &values).unwrap();
|
||||
assert!(rendered.contains(marker), "missing {mode:?} marker");
|
||||
assert!(rendered.contains("<user_query>\nquestion\n</user_query>"));
|
||||
assert_eq!(rendered.matches("<user_query>").count(), 1);
|
||||
}
|
||||
}
|
||||
|
||||
fn assert_mode(assets: &PromptAssets, mode: Mode, expected: &[&str], digest: &str) {
|
||||
assert_eq!(
|
||||
assets
|
||||
.mode(mode)
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
expected
|
||||
);
|
||||
assert_eq!(schema_digest(&assets.mode(mode).tools), digest);
|
||||
}
|
||||
|
||||
fn schema_digest(tools: &[ToolDefinition]) -> String {
|
||||
hex::encode(Sha256::digest(serde_json::to_vec(tools).unwrap()))
|
||||
}
|
||||
|
||||
#[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/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let base = compiler
|
||||
.prompt_spec(Mode::Agent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
let dynamic = compiler
|
||||
.prompt_spec(
|
||||
Mode::Agent,
|
||||
&ModelSpec::new("model"),
|
||||
&[ToolDefinition {
|
||||
name: "mcp_repo_lookup".into(),
|
||||
description: "lookup".into(),
|
||||
parameters: serde_json::json!({"type": "object"}),
|
||||
}],
|
||||
false,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(base.tools, dynamic.tools[..base.tools.len()]);
|
||||
assert_eq!(dynamic.tools.last().unwrap().name, "mcp_repo_lookup");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_tool_cannot_replace_a_mode_tool() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let error = compiler
|
||||
.prompt_spec(
|
||||
Mode::Agent,
|
||||
&ModelSpec::new("model"),
|
||||
&[ToolDefinition {
|
||||
name: "Read".into(),
|
||||
description: "replacement".into(),
|
||||
parameters: serde_json::json!({"type": "object"}),
|
||||
}],
|
||||
false,
|
||||
)
|
||||
.unwrap_err();
|
||||
assert!(error
|
||||
.to_string()
|
||||
.contains("dynamic MCP tool conflicts with a mode tool: Read"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn image_generation_capability_controls_only_the_generate_image_definition() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let without = compiler
|
||||
.prompt_spec(Mode::Agent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
let mut model = ModelSpec::new("model");
|
||||
model.supports_image_generation = true;
|
||||
let with = compiler
|
||||
.prompt_spec(Mode::Agent, &model, &[], false)
|
||||
.unwrap();
|
||||
|
||||
assert!(!without
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "GenerateImage"));
|
||||
assert!(with.tools.iter().any(|tool| tool.name == "GenerateImage"));
|
||||
assert_eq!(with.tools.len(), without.tools.len() + 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn agent_system_prompt_is_static_and_substitutes_the_model_name() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let mut model = ModelSpec::new("test-model-hash");
|
||||
model.display_name = Some("Test Model".into());
|
||||
let request = compiler
|
||||
.prompt_spec(Mode::Agent, &model, &[], false)
|
||||
.unwrap();
|
||||
let prompt = &request.instructions;
|
||||
assert!(prompt.contains("powered by Test Model"));
|
||||
assert!(!prompt.contains("test-model-hash"));
|
||||
assert!(!prompt.contains("{{FAKE_MODEL_NAME}}"));
|
||||
assert!(!prompt.contains("<user_info>"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn subagent_uses_the_agent_prompt_and_only_the_captured_tool_delta() {
|
||||
let assets = PromptAssets::load(
|
||||
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("prompt/cursor")
|
||||
.as_path(),
|
||||
)
|
||||
.unwrap();
|
||||
let compiler = PromptCompiler::new(assets);
|
||||
let agent_prompt = compiler
|
||||
.prompt_spec(Mode::Agent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
let subagent_prompt = compiler
|
||||
.prompt_spec(Mode::Subagent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
assert_eq!(agent_prompt.instructions, subagent_prompt.instructions);
|
||||
|
||||
let request = compiler
|
||||
.prompt_spec(Mode::Subagent, &ModelSpec::new("model"), &[], false)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
request
|
||||
.tools
|
||||
.iter()
|
||||
.map(|tool| tool.name.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec![
|
||||
"Shell",
|
||||
"Grep",
|
||||
"Delete",
|
||||
"WebSearch",
|
||||
"WebFetch",
|
||||
"ReadLints",
|
||||
"EditNotebook",
|
||||
"TodoWrite",
|
||||
"StrReplace",
|
||||
"Write",
|
||||
"Read",
|
||||
"Glob",
|
||||
"GetMcpTools",
|
||||
"FetchMcpResource",
|
||||
"SwitchMode",
|
||||
"UpdateCurrentStep",
|
||||
"CallMcpTool",
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
]
|
||||
);
|
||||
assert!(!request.tools.iter().any(|tool| tool.name == "Task"));
|
||||
|
||||
let suppressed = compiler
|
||||
.prompt_spec(Mode::Subagent, &ModelSpec::new("model"), &[], true)
|
||||
.unwrap();
|
||||
assert!(!suppressed
|
||||
.tools
|
||||
.iter()
|
||||
.any(|tool| tool.name == "UpdateCurrentStep"));
|
||||
}
|
||||
|
||||
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 {
|
||||
let output = output.into();
|
||||
CanonicalMessage {
|
||||
message_id: id.into(),
|
||||
role: Role::Tool,
|
||||
origin: Origin::Tool,
|
||||
content: MessageContent::ToolResult(ToolResultContent {
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
content: output
|
||||
.as_str()
|
||||
.map(str::to_string)
|
||||
.unwrap_or_else(|| output.to_string()),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn named_tool_result(name: &str, output: &str) -> CanonicalMessage {
|
||||
CanonicalMessage {
|
||||
message_id: format!("result-{name}"),
|
||||
role: Role::Tool,
|
||||
origin: Origin::Tool,
|
||||
content: MessageContent::ToolResult(ToolResultContent {
|
||||
call_id: format!("call-{name}"),
|
||||
name: name.into(),
|
||||
content: output.into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn assistant_tool_pair(
|
||||
id: &str,
|
||||
tool_round_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(),
|
||||
tool_round_id: Some(tool_round_id.into()),
|
||||
replay_state: None,
|
||||
tool_calls: vec![ToolCallContent {
|
||||
index,
|
||||
call_id: call_id.into(),
|
||||
name: "Tool".into(),
|
||||
arguments: serde_json::json!({}),
|
||||
}],
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
//! Provides captured Cursor wire fixtures for integration tests.
|
||||
use bytes::Bytes;
|
||||
use cursor_server::{cursor::protocol::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,92 @@
|
||||
//! Provides deterministic provider streams for integration tests.
|
||||
#![allow(dead_code)]
|
||||
|
||||
use std::{
|
||||
collections::VecDeque,
|
||||
sync::{Arc, Mutex},
|
||||
};
|
||||
|
||||
use cursor_server::{
|
||||
model::{ModelInvocation, ModelRequest},
|
||||
provider::{ModelEvent, Provider, ProviderStream},
|
||||
Error,
|
||||
};
|
||||
use futures_util::{stream, StreamExt};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
enum FakeResponse {
|
||||
Events(Vec<Result<ModelEvent, Error>>),
|
||||
Gated {
|
||||
ready: Arc<tokio::sync::Notify>,
|
||||
events: Vec<Result<ModelEvent, Error>>,
|
||||
},
|
||||
Pending,
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct FakeProvider {
|
||||
responses: Arc<Mutex<VecDeque<FakeResponse>>>,
|
||||
requests: Arc<Mutex<Vec<ModelRequest>>>,
|
||||
}
|
||||
|
||||
impl FakeProvider {
|
||||
pub fn push(&self, events: Vec<ModelEvent>) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Events(events.into_iter().map(Ok).collect()));
|
||||
}
|
||||
pub fn push_error(&self, error: Error) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Events(vec![Err(error)]));
|
||||
}
|
||||
pub fn push_pending(&self) {
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Pending);
|
||||
}
|
||||
pub fn push_gated(&self, events: Vec<ModelEvent>) -> Arc<tokio::sync::Notify> {
|
||||
let ready = Arc::new(tokio::sync::Notify::new());
|
||||
self.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push_back(FakeResponse::Gated {
|
||||
ready: ready.clone(),
|
||||
events: events.into_iter().map(Ok).collect(),
|
||||
});
|
||||
ready
|
||||
}
|
||||
pub fn requests(&self) -> Vec<ModelRequest> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl Provider for FakeProvider {
|
||||
fn stream(
|
||||
&self,
|
||||
invocation: ModelInvocation,
|
||||
_cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
self.requests.lock().unwrap().push(invocation.request);
|
||||
let events = self
|
||||
.responses
|
||||
.lock()
|
||||
.unwrap()
|
||||
.pop_front()
|
||||
.expect("fake response configured");
|
||||
match events {
|
||||
FakeResponse::Events(events) => Box::pin(stream::iter(events)),
|
||||
FakeResponse::Gated { ready, events } => Box::pin(
|
||||
stream::once(async move {
|
||||
ready.notified().await;
|
||||
events
|
||||
})
|
||||
.flat_map(stream::iter),
|
||||
),
|
||||
FakeResponse::Pending => Box::pin(stream::pending()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
//! Provides isolated stores and canonical message fixtures for tests.
|
||||
#![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)
|
||||
}
|
||||
@@ -0,0 +1,980 @@
|
||||
//! Verifies Tool dispatch, completion gating, and result continuation.
|
||||
#[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::{
|
||||
protocol::{connect, proto::agent::v1 as pb},
|
||||
tools::{
|
||||
codec,
|
||||
runtime::{CursorToolRuntime, ExecContext},
|
||||
ClientToolEvent, ToolBatchState, ToolDispatcher,
|
||||
},
|
||||
},
|
||||
cursor::{TransportCommand, TransportRegistry},
|
||||
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(),
|
||||
default_subagent_model: "model".into(),
|
||||
subagent_model: None,
|
||||
terminals_folder: "/tmp/terminals".into(),
|
||||
admin_command_denylist: Vec::new(),
|
||||
allow_subagents: true,
|
||||
subagents_disabled: false,
|
||||
mcp_routes: std::collections::HashMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn mcp_context(server: &str, provider: &str, tool: &str) -> ExecContext {
|
||||
let mut context = exec_context();
|
||||
context.mcp_routes.insert(
|
||||
(server.into(), tool.into()),
|
||||
cursor_server::cursor::tools::runtime::McpRoute {
|
||||
name: format!("{server}-{tool}"),
|
||||
provider_identifier: provider.into(),
|
||||
tool_name: tool.into(),
|
||||
description: "fixture MCP tool".into(),
|
||||
},
|
||||
);
|
||||
context
|
||||
}
|
||||
|
||||
#[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 dynamic_mcp_uses_one_definition_for_stream_ui_exec_and_result() {
|
||||
let definition = pb::McpToolDefinition {
|
||||
name: "cursor-ide-browser-browser_navigate".into(),
|
||||
provider_identifier: "cursor-ide-browser".into(),
|
||||
tool_name: "browser_navigate".into(),
|
||||
description: "Navigate the browser".into(),
|
||||
..Default::default()
|
||||
};
|
||||
let definitions = BTreeMap::from([(definition.name.clone(), definition.clone())]);
|
||||
let event = cursor_server::provider::ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "browser-call".into(),
|
||||
name: definition.name.clone(),
|
||||
};
|
||||
let partial =
|
||||
cursor_server::cursor::protocol::events::response_event(&event, "model:0", &definitions)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = partial.message else {
|
||||
panic!("expected interaction update")
|
||||
};
|
||||
let Some(pb::interaction_update::Message::PartialToolCall(partial)) = update.message else {
|
||||
panic!("expected partial tool call")
|
||||
};
|
||||
let partial = partial.tool_call.unwrap();
|
||||
assert_eq!(partial.started_at_ms, None);
|
||||
let Some(pb::tool_call::Tool::McpToolCall(tool)) = partial.tool else {
|
||||
panic!("expected MCP placeholder")
|
||||
};
|
||||
assert_eq!(tool.args.unwrap().tool_name, "browser_navigate");
|
||||
|
||||
let runtime = CursorToolRuntime::default();
|
||||
let dispatcher = ToolDispatcher::new(runtime.clone());
|
||||
let mut invocation = call("browser-call", &definition.name);
|
||||
invocation.arguments = json!({"url": "https://example.com"});
|
||||
let dispatched = dispatcher
|
||||
.start_batch(
|
||||
&[invocation],
|
||||
ToolBatchState {
|
||||
completed: &HashSet::new(),
|
||||
started: &HashSet::new(),
|
||||
response_text: "",
|
||||
response_thinking: "",
|
||||
},
|
||||
&[],
|
||||
&definitions,
|
||||
&exec_context(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(dispatched[0].messages.len(), 2);
|
||||
let Some(pb::agent_server_message::Message::InteractionUpdate(update)) =
|
||||
dispatched[0].messages[0].message.as_ref()
|
||||
else {
|
||||
panic!("expected tool-start interaction")
|
||||
};
|
||||
let Some(pb::interaction_update::Message::ToolCallStarted(started)) = update.message.as_ref()
|
||||
else {
|
||||
panic!("expected tool-start message")
|
||||
};
|
||||
assert!(started.tool_call.as_ref().unwrap().started_at_ms.is_some());
|
||||
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) =
|
||||
dispatched[0].messages[1].message.as_ref()
|
||||
else {
|
||||
panic!("expected MCP Exec")
|
||||
};
|
||||
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: "navigated".into(),
|
||||
output_location: None,
|
||||
},
|
||||
)),
|
||||
}],
|
||||
is_error: false,
|
||||
structured_content: None,
|
||||
})),
|
||||
})),
|
||||
..Default::default()
|
||||
},
|
||||
&runtime,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let codec::ClientExecEvent::Completed(completion) = event else {
|
||||
panic!("expected completed MCP result")
|
||||
};
|
||||
let Some(pb::tool_call::Tool::McpToolCall(tool)) = &completion.tool_call().tool else {
|
||||
panic!("expected rendered MCP result")
|
||||
};
|
||||
assert_eq!(tool.args.as_ref().unwrap().name, definition.name);
|
||||
assert!(tool.result.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn call_mcp_tool_uses_the_request_descriptor_and_returns_client_errors_to_the_model() {
|
||||
let runtime = CursorToolRuntime::default();
|
||||
let dispatcher = ToolDispatcher::new(runtime.clone());
|
||||
let completed = HashSet::new();
|
||||
let started = HashSet::new();
|
||||
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(),
|
||||
&mcp_context(
|
||||
"plugin-browser-use-browser-use",
|
||||
"browser-use",
|
||||
"browser_exec",
|
||||
),
|
||||
)
|
||||
.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::Error(pb::McpError {
|
||||
error: "invalid browser arguments".into(),
|
||||
})),
|
||||
})),
|
||||
..Default::default()
|
||||
},
|
||||
&runtime,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let codec::ClientExecEvent::Completed(completion) = event else {
|
||||
panic!("expected MCP completion")
|
||||
};
|
||||
assert_eq!(completion.result().content, "invalid browser arguments");
|
||||
assert!(completion.result().is_error);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mcp_auth_uses_the_cursor_auth_interaction_without_a_tool_definition() {
|
||||
let runtime = CursorToolRuntime::default();
|
||||
let dispatcher = ToolDispatcher::new(runtime.clone());
|
||||
let completed = HashSet::new();
|
||||
let started = HashSet::new();
|
||||
let mut auth = call("auth-gmail", "CallMcpTool");
|
||||
auth.arguments = json!({
|
||||
"server": "plugin-gmail-gmail",
|
||||
"toolName": "mcp_auth",
|
||||
"arguments": {}
|
||||
});
|
||||
let request = dispatcher
|
||||
.start_batch(
|
||||
&[auth],
|
||||
ToolBatchState {
|
||||
completed: &completed,
|
||||
started: &started,
|
||||
response_text: "",
|
||||
response_thinking: "",
|
||||
},
|
||||
&[],
|
||||
&BTreeMap::new(),
|
||||
&exec_context(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let Some(pb::agent_server_message::Message::InteractionQuery(query)) =
|
||||
request[0].messages[1].message.as_ref()
|
||||
else {
|
||||
panic!("expected MCP auth interaction")
|
||||
};
|
||||
let Some(pb::interaction_query::Query::McpAuthRequestQuery(auth)) = query.query.as_ref() else {
|
||||
panic!("expected MCP auth query")
|
||||
};
|
||||
let args = auth.args.as_ref().unwrap();
|
||||
assert_eq!(args.server_identifier, "plugin-gmail-gmail");
|
||||
assert_eq!(args.tool_call_id, "auth-gmail");
|
||||
|
||||
let event = dispatcher
|
||||
.interaction_response(&pb::InteractionResponse {
|
||||
id: query.id,
|
||||
result: Some(pb::interaction_response::Result::McpAuthRequestResponse(
|
||||
pb::McpAuthRequestResponse {
|
||||
result: Some(pb::mcp_auth_request_response::Result::Approved(
|
||||
pb::mcp_auth_request_response::Approved {},
|
||||
)),
|
||||
},
|
||||
)),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
let ClientToolEvent::Completed(completion) = event else {
|
||||
panic!("expected MCP auth completion")
|
||||
};
|
||||
let Some(pb::tool_call::Tool::McpAuthToolCall(auth)) = &completion.tool_call().tool else {
|
||||
panic!("expected MCP auth tool call")
|
||||
};
|
||||
assert!(matches!(
|
||||
auth.result.as_ref().and_then(|result| result.result.as_ref()),
|
||||
Some(pb::mcp_auth_result::Result::Success(success))
|
||||
if success.server_identifier == "plugin-gmail-gmail"
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn unknown_mcp_descriptor_returns_a_tool_error_without_client_discovery() {
|
||||
let dispatcher = ToolDispatcher::new(CursorToolRuntime::default());
|
||||
let completed = HashSet::new();
|
||||
let started = HashSet::new();
|
||||
let mut invocation = call("call-fast-context", "CallMcpTool");
|
||||
invocation.arguments = json!({
|
||||
"server": "fast-context",
|
||||
"toolName": "fast_context_search",
|
||||
"arguments": {"query": "MCP dispatch"}
|
||||
});
|
||||
let dispatched = dispatcher
|
||||
.start_batch(
|
||||
&[invocation],
|
||||
ToolBatchState {
|
||||
completed: &completed,
|
||||
started: &started,
|
||||
response_text: "",
|
||||
response_thinking: "",
|
||||
},
|
||||
&[],
|
||||
&BTreeMap::new(),
|
||||
&exec_context(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(dispatched[0].messages.len(), 1);
|
||||
let completion = dispatched[0]
|
||||
.completion
|
||||
.as_ref()
|
||||
.expect("missing descriptor should complete as a tool error");
|
||||
assert!(completion.result().is_error);
|
||||
assert!(completion.result().content.contains("descriptor not found"));
|
||||
}
|
||||
|
||||
#[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));
|
||||
assert_eq!(args.simple_commands, ["python3 -m http.server 8000"]);
|
||||
let parsing = args.parsing_result.as_ref().unwrap();
|
||||
assert!(!parsing.parsing_failed);
|
||||
assert_eq!(parsing.executable_commands.len(), 1);
|
||||
let executable = &parsing.executable_commands[0];
|
||||
assert_eq!(executable.name, "python3");
|
||||
assert_eq!(executable.full_text, "python3 -m http.server 8000");
|
||||
assert_eq!(
|
||||
executable
|
||||
.args
|
||||
.iter()
|
||||
.map(|argument| (argument.r#type.as_str(), argument.value.as_str()))
|
||||
.collect::<Vec<_>>(),
|
||||
[("word", "-m"), ("word", "http.server"), ("word", "8000")]
|
||||
);
|
||||
|
||||
let rendered = cursor_server::cursor::tools::codec::render_tool_call(&shell, false).unwrap();
|
||||
let Some(pb::tool_call::Tool::ShellToolCall(rendered)) = rendered.tool else {
|
||||
panic!("expected rendered ShellToolCall")
|
||||
};
|
||||
assert_eq!(rendered.description.as_deref(), Some("Start HTTP server"));
|
||||
assert_eq!(
|
||||
rendered.args.and_then(|args| args.description),
|
||||
Some("Start HTTP server".into())
|
||||
);
|
||||
|
||||
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"))
|
||||
.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 exec_stream_close_without_a_terminal_result_becomes_a_tool_error() {
|
||||
let pending = CursorToolRuntime::default();
|
||||
let mut shell = call("call-1", "Shell");
|
||||
shell.arguments = json!({"command": "git status"});
|
||||
let id = pending.reserve_exec(&shell, &exec_context()).await.unwrap();
|
||||
|
||||
let completion = codec::stream_closed(id, &pending)
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("a running Exec should complete when its stream closes");
|
||||
|
||||
assert_eq!(completion.result().call_id, "call-1");
|
||||
assert!(completion.result().is_error);
|
||||
assert_eq!(
|
||||
completion.result().content,
|
||||
"Cursor Exec stream closed before returning a terminal result"
|
||||
);
|
||||
let Some(pb::tool_call::Tool::ShellToolCall(shell)) = &completion.tool_call().tool else {
|
||||
panic!("expected typed Shell completion")
|
||||
};
|
||||
assert!(matches!(
|
||||
shell.result.as_ref().and_then(|result| result.result.as_ref()),
|
||||
Some(pb::shell_result::Result::SpawnError(error))
|
||||
if error.error == "Cursor Exec stream closed before returning a terminal result"
|
||||
));
|
||||
assert!(pending.exec_call(id).await.is_none());
|
||||
assert!(codec::stream_closed(id, &pending).await.unwrap().is_none());
|
||||
}
|
||||
|
||||
#[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 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 = TransportRegistry::new(
|
||||
store.clone(),
|
||||
Arc::new(provider.clone()),
|
||||
PromptCompiler::new(assets),
|
||||
);
|
||||
let handle = registry.get_or_create("tool-request").await.unwrap();
|
||||
let mut output = handle.subscribe();
|
||||
handle
|
||||
.command(TransportCommand::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(TransportCommand::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(TransportCommand::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(TransportCommand::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 cursor_request_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 },
|
||||
)),
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user