This commit is contained in:
leookun
2026-08-29 22:05:04 +08:00
parent 279e6bb07c
commit d200b3791d
204 changed files with 710 additions and 440 deletions
@@ -0,0 +1,707 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{collections::HashMap, sync::Arc};
use cursor_server::{
cursor::{
bidi_append::{self, DecodedAppend},
connect,
prompting::{PromptAssets, PromptCompiler},
proto::agent::v1 as pb,
CursorCommand, CursorSessionHandle, CursorSessionRegistry,
},
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 cancelled_conversation_drops_late_background_completion() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let active = registry.get_or_create("cancelled-request").await.unwrap();
active.set_conversation_id("parent-conversation").unwrap();
active.mark_conversation_cancelled();
bidi_append::append(
&registry,
DecodedAppend {
request_id: "late-completion".into(),
seqno: 1,
message: completion_run(
"child-id",
"cancelled-parent-run",
pb::ConversationStateStructure::default(),
),
},
None,
)
.await
.unwrap();
assert!(provider.requests().is_empty());
}
#[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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
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 = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
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: &CursorSessionHandle,
message: pb::AgentClientMessage,
) -> (pb::ConversationStateStructure, HashMap<Vec<u8>, Vec<u8>>) {
let mut output = handle.subscribe();
handle
.command(CursorCommand::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(CursorCommand::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(CursorCommand::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(CursorCommand::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: &CursorSessionHandle, message: pb::AgentClientMessage) {
let mut output = handle.subscribe();
handle
.command(CursorCommand::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(CursorCommand::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(CursorCommand::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(CursorCommand::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 },
)),
},
)),
}
}
+497
View File
@@ -0,0 +1,497 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{collections::HashSet, sync::Arc};
use cursor_server::{
cursor::{
connect,
prompting::{PromptAssets, PromptCompiler},
proto::agent::v1 as pb,
CursorCommand, CursorSessionRegistry,
},
model::ToolRoundId,
provider::{FinishReason, ModelEvent},
store::{BlobEdge, BlobId},
};
use prost::Message;
#[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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let first = registry.get_or_create("first-run").await.unwrap();
let mut first_output = first.subscribe();
first
.command(CursorCommand::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.cancel();
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(CursorCommand::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(CursorCommand::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 = CursorSessionRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("bad-blob-run").await.unwrap();
let mut output = handle.subscribe();
let expected = BlobId::digest(b"expected");
handle
.command(CursorCommand::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(CursorCommand::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::CursorSessionHandle,
seqno: &mut i64,
id: u32,
) {
handle
.command(CursorCommand::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;
}
+307
View File
@@ -0,0 +1,307 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
model::{
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind,
ToolDefinition, ToolResult,
},
provider::{FinishReason, ModelEvent},
run::{session, ClientCommand, ClientEvent, CommitCause, RunEngine, RunOutcome},
};
use tokio::{sync::oneshot, time::Duration};
use tokio_util::sync::CancellationToken;
#[tokio::test]
async fn a_client_without_checkpoint_protocol_runs_the_same_text_loop() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-1".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("hello".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let prepared = prepared(&store).await;
let (port, mut client) = session(32);
let engine = RunEngine::new(store.clone(), Arc::new(provider));
let run =
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
let mut saw_final_commit = false;
while let Some(event) = client.events.recv().await {
match event {
ClientEvent::TextDelta(text) => assert_eq!(text, "hello"),
ClientEvent::StateCommitted(state) => {
saw_final_commit |= state.cause == CommitCause::FinalTurn;
state.barrier.complete(Ok(()));
}
ClientEvent::Ended(outcome) => {
assert_eq!(outcome, RunOutcome::Completed);
break;
}
_ => {}
}
}
assert!(saw_final_commit);
assert_eq!(run.await.unwrap(), RunOutcome::Completed);
}
#[tokio::test]
async fn inserted_messages_wait_for_the_next_model_call_without_interrupting_the_active_call() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
let first_ready = provider.push_gated(vec![
ModelEvent::Start {
model_call_id: "call-1".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("first answer".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-2".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("followed up".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let prepared = prepared(&store).await;
let (port, mut client) = session(32);
let commands = client.commands.clone();
let engine = RunEngine::new(store, Arc::new(provider.clone()));
let run =
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
while provider.requests().is_empty() {
if let Ok(Some(ClientEvent::StateCommitted(state))) =
tokio::time::timeout(Duration::from_millis(20), client.events.recv()).await
{
state.barrier.complete(Ok(()));
}
}
let (delivered, mut delivery) = oneshot::channel();
commands
.send(ClientCommand::InsertMessages(
cursor_server::run::MessageInsertion {
messages: vec![cursor_server::model::RuntimeEvent {
event_id: "background:finished".into(),
text: "background work finished".into(),
}
.into_message()],
delivered,
},
))
.await
.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(20), &mut delivery)
.await
.is_err()
);
first_ready.notify_one();
while let Some(event) = client.events.recv().await {
match event {
ClientEvent::StateCommitted(state) => state.barrier.complete(Ok(())),
ClientEvent::Ended(outcome) => {
assert_eq!(outcome, RunOutcome::Completed);
break;
}
_ => {}
}
}
assert_eq!(run.await.unwrap(), RunOutcome::Completed);
delivery.await.unwrap();
let requests = provider.requests();
assert_eq!(requests.len(), 2);
let history = &requests[1].history;
assert!(matches!(
history[1].role,
cursor_server::model::Role::Assistant
));
assert_eq!(history[2].message_id, "runtime:background:finished");
}
#[tokio::test]
async fn a_failed_claim_cannot_overwrite_the_existing_run() {
let (_directory, store) = fixtures::temp_store().await;
let prepared = prepared(&store).await;
store.claim_run(&prepared).await.unwrap();
let (port, mut client) = session(8);
let outcome = RunEngine::new(
store.clone(),
Arc::new(fake_provider::FakeProvider::default()),
)
.run(prepared, port, CancellationToken::new())
.await;
assert!(matches!(outcome, RunOutcome::Failed(_)));
assert!(matches!(
client.events.recv().await,
Some(ClientEvent::Ended(RunOutcome::Failed(_)))
));
let status: String = sqlx::query_scalar("SELECT status FROM runs WHERE run_id = 'run'")
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(status, "running");
}
#[tokio::test]
async fn required_client_state_failure_prevents_a_completed_run() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-1".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("hello".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let prepared = prepared(&store).await;
let (port, mut client) = session(32);
let engine = RunEngine::new(store, Arc::new(provider));
let run =
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
while let Some(event) = client.events.recv().await {
match event {
ClientEvent::StateCommitted(state) if state.cause == CommitCause::FinalTurn => {
state.barrier.complete(Err("snapshot failed".into()));
}
ClientEvent::StateCommitted(state) => state.barrier.complete(Ok(())),
ClientEvent::Ended(outcome) => {
assert_eq!(
outcome,
RunOutcome::Failed(cursor_server::run::RunFailure::Client(
"snapshot failed".into()
))
);
break;
}
_ => {}
}
}
assert_eq!(
run.await.unwrap(),
RunOutcome::Failed(cursor_server::run::RunFailure::Client(
"snapshot failed".into()
))
);
}
#[tokio::test]
async fn generic_engine_waits_for_every_tool_result_without_any_cursor_wire_id() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-1".into(),
},
tool_start(0, "A"),
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: "{}".into(),
},
ModelEvent::ToolCallEnd { index: 0 },
tool_start(1, "B"),
ModelEvent::ToolCallArgumentsDelta {
index: 1,
delta: "{}".into(),
},
ModelEvent::ToolCallEnd { index: 1 },
ModelEvent::Done(FinishReason::ToolUse),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-2".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("done".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let prepared = prepared(&store).await;
let (port, mut client) = session(64);
let commands = client.commands.clone();
let engine = RunEngine::new(store, Arc::new(provider));
let run =
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
while let Some(event) = client.events.recv().await {
match event {
ClientEvent::ExecuteToolRound { calls, .. } => {
assert_eq!(calls.len(), 2);
commands
.send(ClientCommand::ToolResult(ToolResult {
call_id: "B".into(),
content: "result-B".into(),
is_error: false,
image: None,
}))
.await
.unwrap();
commands
.send(ClientCommand::ToolResult(ToolResult {
call_id: "A".into(),
content: "result-A".into(),
is_error: false,
image: None,
}))
.await
.unwrap();
}
ClientEvent::StateCommitted(state) => state.barrier.complete(Ok(())),
ClientEvent::Ended(outcome) => {
assert_eq!(outcome, RunOutcome::Completed);
break;
}
_ => {}
}
}
assert_eq!(run.await.unwrap(), RunOutcome::Completed);
}
async fn prepared(store: &cursor_server::store::Store) -> PreparedRun {
let conversation_id = ConversationId::new("conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
PreparedRun {
run_id: RunId::new("run"),
cursor_request_id: None,
conversation_id,
kind: RunKind::Root,
model: ModelSpec::new("model"),
prompt: PromptSpec {
instructions: "system".into(),
tools: vec![ToolDefinition {
name: "Tool".into(),
description: "test".into(),
parameters: serde_json::json!({"type":"object"}),
}],
},
initial_messages: vec![fixtures::user("user", "hello")],
action: RunAction::Start,
base_revision_id: root,
}
}
fn tool_start(index: usize, call_id: &str) -> ModelEvent {
ModelEvent::ToolCallStart {
index,
call_id: call_id.into(),
name: "Tool".into(),
}
}
+368
View File
@@ -0,0 +1,368 @@
#[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::{connect, proto::agent::v1 as pb, CursorCommand, CursorSessionRegistry},
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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let first = run(
&registry,
"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(
&registry,
"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(
&registry,
"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: &CursorSessionRegistry,
request_id: &str,
request: pb::AgentClientMessage,
) -> Output {
let handle = registry.get_or_create(request_id).await.unwrap();
let mut receiver = handle.subscribe();
handle
.command(CursorCommand::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(CursorCommand::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 },
)),
},
)),
}
}
+137
View File
@@ -0,0 +1,137 @@
#[path = "support/fake_cursor.rs"]
mod fake_cursor;
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{io::Write, sync::Arc};
use axum::{
body::{to_bytes, Body},
http::{header, Request, StatusCode},
};
use base64::{engine::general_purpose::STANDARD_NO_PAD, Engine};
use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::CursorSessionRegistry,
cursor::{
connect, handlers,
proto::{agent::v1 as pb, aiserver::v1 as ai},
},
};
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 = CursorSessionRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
Default::default(),
);
let wire = ai::BidiAppendRequest {
request_id: Some(ai::BidiRequestId {
request_id: "gzip-request".into(),
}),
..Default::default()
}
.encode_to_vec();
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder.write_all(&wire).unwrap();
let compressed = encoder.finish().unwrap();
let response = handlers::router(registry)
.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"));
}
+471
View File
@@ -0,0 +1,471 @@
#[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::{
connect,
proto::{agent::v1 as pb, aiserver::v1 as ai},
},
cursor::{CursorCommand, CursorSessionRegistry},
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 = CursorSessionRegistry::new(
store,
Arc::new(fake_provider::FakeProvider::default()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("abort-request").await.unwrap();
let mut output = handle.subscribe();
handle.command(CursorCommand::Abort).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!(handle.cancellation().is_cancelled());
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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("failed-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::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(CursorCommand::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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry
.get_or_create("protocol-failed-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(protocol_client_run()),
})
.await
.unwrap();
let mut append_seqno = 1;
let mut saw_turn_ended = false;
let error_json = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Error EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
break serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::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(CursorCommand::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 duplicate_run_request_on_one_bidi_stream_is_a_protocol_error() {
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 = CursorSessionRegistry::new(
store,
Arc::new(provider),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry
.get_or_create("protocol-failed-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(protocol_client_run()),
})
.await
.unwrap();
let mut seqno = 1;
let mut duplicate_sent = 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(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(_)) if !duplicate_sent => {
duplicate_sent = true;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(protocol_client_run()),
})
.await
.unwrap();
seqno += 1;
}
_ => {}
}
};
assert!(duplicate_sent);
assert_eq!(error_json["error"]["code"], "invalid_argument");
assert_eq!(
error_json["error"]["message"],
"duplicate RunRequest for request_id: protocol-failed-request"
);
}
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() -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
action: Some(pb::ConversationAction {
action: Some(pb::conversation_action::Action::UserMessageAction(
pb::UserMessageAction {
user_message: Some(pb::UserMessage {
text: "read it".into(),
message_id: "protocol-failed-user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("protocol-failed-conversation".into()),
run_id: Some("protocol-failed-request".into()),
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
+198
View File
@@ -0,0 +1,198 @@
use cursor_server::{
model::{
ConversationId, ModelConfigInput, ModelSpec, ModelType, NewLlmCall, PreparedRun,
PromptSpec, ProviderType, RunAction, RunId, RunKind, Usage, OPENAI_CHAT_ENDPOINT,
},
store::{RunStatus, Store},
};
async fn 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)
}
fn model_input() -> ModelConfigInput {
ModelConfigInput {
sort_order: 0,
display_name: "Model A".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Model A".into(),
model_id: "model-a".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: true,
openai_extra_params: serde_json::json!({"temperature":0}),
custom_headers_enabled: true,
custom_headers: serde_json::json!({"x-route":"one"}),
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,
}
}
#[tokio::test]
async fn model_configuration_round_trips_and_hash_uses_v0049_identity() {
let (_directory, store) = store().await;
let model = store.create_model(&model_input()).await.unwrap();
assert_eq!(model.model_hash.len(), 16);
assert_eq!(model.api_key, "secret");
assert_eq!(model.custom_headers["x-route"], "one");
let original_hash = model.model_hash.clone();
let mut input = model_input();
input.base_url = "https://example.com/v1/chat/completions".into();
input.sort_order = 3;
input.tooltip_data = "Updated tooltip".into();
let updated = store.update_model(&original_hash, &input).await.unwrap();
assert_eq!(updated.model_hash, original_hash);
assert_eq!(updated.sort_order, 3);
assert_eq!(updated.tooltip_data, "Updated tooltip");
}
#[tokio::test]
async fn arbitrary_request_url_is_independent_from_openai_protocol() {
let (_directory, store) = store().await;
let mut input = model_input();
input.base_url = "https://proxy.example.com/arbitrary/generate?api-version=2026-01-01".into();
let chat = store.create_model(&input).await.unwrap();
assert_eq!(chat.provider_type(), ProviderType::OpenAiChat);
assert_eq!(chat.request_url().unwrap(), input.base_url);
input.openai_endpoint = "/v1/responses".into();
let responses = store.update_model(&chat.model_hash, &input).await.unwrap();
assert_eq!(responses.provider_type(), ProviderType::OpenAiResponses);
assert_eq!(responses.request_url().unwrap(), input.base_url);
}
#[tokio::test]
async fn standard_server_address_resolves_to_the_same_model_identity_as_a_complete_url() {
let (_directory, store) = store().await;
let complete = model_input();
let mut standard = complete.clone();
standard.base_url = "https://example.com/v1".into();
standard.use_full_url = false;
let model = store.create_model(&standard).await.unwrap();
assert!(!model.use_full_url);
assert_eq!(
model.request_url().unwrap(),
"https://example.com/v1/chat/completions"
);
assert_eq!(
model.model_hash,
cursor_server::model::model_hash(&complete).unwrap()
);
}
#[tokio::test]
async fn call_summary_is_always_stored_and_payloads_follow_detailed_setting() {
let (_directory, store) = store().await;
let model = store.create_model(&model_input()).await.unwrap();
let call = NewLlmCall {
call_id: "call-1".into(),
run_id: "run-1".into(),
conversation_id: "conversation-1".into(),
provider_call_index: 0,
model_hash: model.model_hash,
provider_type: ProviderType::OpenAiChat,
provider_url: model.base_url.clone(),
request_type: ProviderType::OpenAiChat,
request_url: "https://example.com/v1/chat/completions".into(),
model_id: model.model_id,
display_name: model.display_name,
reasoning_effort: Some("high".into()),
fast: true,
message_count: 2,
tool_count: 3,
detailed: false,
};
let conversation_id = ConversationId::new("conversation-1");
let base_revision_id = store.ensure_conversation(&conversation_id).await.unwrap();
store
.claim_run(&PreparedRun {
run_id: RunId::new("run-1"),
cursor_request_id: None,
conversation_id,
kind: RunKind::Root,
model: ModelSpec::new(call.model_hash.clone()),
prompt: PromptSpec {
instructions: String::new(),
tools: Vec::new(),
},
initial_messages: Vec::new(),
action: RunAction::Resume {
pending_tool_round: None,
},
base_revision_id,
})
.await
.unwrap();
store.start_llm_call(&call).await.unwrap();
store
.record_llm_request(
"call-1",
&serde_json::json!({}),
&serde_json::json!({"model":"model-a"}),
false,
)
.await
.unwrap();
store
.record_llm_chunk("call-1", 0, 4, b"data", false)
.await
.unwrap();
store
.record_llm_usage(
"call-1",
Usage {
input_tokens: Some(10),
output_tokens: Some(5),
total_tokens: Some(15),
..Default::default()
},
)
.await
.unwrap();
store
.finish_llm_call("call-1", "completed", Some("stop"), 9, None, None)
.await
.unwrap();
let summary = store.llm_call("call-1").await.unwrap().unwrap();
assert_eq!(summary.total_tokens, Some(15));
assert_eq!(summary.reasoning_effort.as_deref(), Some("high"));
assert_eq!(summary.fast, Some(true));
assert_eq!(summary.request_bytes, Some(19));
assert_eq!(summary.response_bytes, 4);
assert!(store.llm_call_request("call-1").await.unwrap().is_none());
assert!(store.llm_call_chunks("call-1").await.unwrap().is_empty());
let abandoned = NewLlmCall {
call_id: "call-2".into(),
provider_call_index: 1,
..call
};
store.start_llm_call(&abandoned).await.unwrap();
store
.finish_run(&RunId::new("run-1"), RunStatus::Cancelled, None, None)
.await
.unwrap();
let abandoned = store.llm_call("call-2").await.unwrap().unwrap();
assert_eq!(abandoned.status, "cancelled");
assert!(abandoned.finished_at_ms.is_some());
assert!(abandoned.duration_ms.is_some());
assert_eq!(
store.llm_call("call-1").await.unwrap().unwrap().status,
"completed"
);
}
+235
View File
@@ -0,0 +1,235 @@
use std::time::Duration;
use axum::{http::header, response::IntoResponse, routing::post, Router};
use cursor_server::{
model::{
ModelConfigInput, ModelInvocation, ModelRequest, ModelSpec, ModelType, PromptSpec,
OPENAI_CHAT_ENDPOINT,
},
provider::{ModelEvent, Provider, ProviderRouter},
store::Store,
};
use futures_util::StreamExt;
use tokio_util::sync::CancellationToken;
async fn test_store(name: &str) -> (tempfile::TempDir, Store) {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
directory.path().join(name).display()
))
.await
.unwrap();
(directory, store)
}
#[tokio::test]
async fn cursor_traces_are_absent_when_detailed_logging_is_disabled() {
let (_directory, store) = test_store("cursor-trace-disabled.db").await;
assert!(!store
.start_cursor_trace_if_detailed(
"request-disabled",
Some("conversation"),
"local_byok",
Some("model"),
)
.await
.unwrap());
assert!(store
.cursor_trace("request-disabled")
.await
.unwrap()
.is_none());
}
#[tokio::test]
async fn cursor_trace_links_detailed_artifacts_to_the_logical_run() {
let (_directory, store) = test_store("cursor-trace-enabled.db").await;
store.set_detailed_logging(true).await.unwrap();
assert!(store
.start_cursor_trace_if_detailed(
"request-enabled",
Some("conversation"),
"cursor_official",
Some("official-model"),
)
.await
.unwrap());
store
.append_cursor_trace_artifact(
"request-enabled",
"bidi_append_request",
"cursor_client",
b"request",
&serde_json::json!({"append_seqno": 1}),
)
.await
.unwrap();
store
.add_cursor_trace_request_bytes("request-enabled", 7)
.await
.unwrap();
store
.start_cursor_trace_response("request-enabled", 200)
.await
.unwrap();
store
.add_cursor_trace_response_chunk("request-enabled", "cursor_official", b"response")
.await
.unwrap();
store
.finish_cursor_trace("request-enabled", None)
.await
.unwrap();
let trace = store
.cursor_trace("request-enabled")
.await
.unwrap()
.unwrap();
assert_eq!(trace.route, "cursor_official");
assert_eq!(trace.status, "completed");
assert_eq!(trace.request_bytes, 7);
assert_eq!(trace.response_bytes, 8);
assert_eq!(trace.response_event_count, 1);
let artifacts = store
.cursor_trace_artifacts("request-enabled")
.await
.unwrap();
assert_eq!(artifacts.len(), 2);
assert_eq!(artifacts[0].artifact_type, "bidi_append_request");
assert_eq!(artifacts[1].artifact_type, "run_sse_chunk");
assert_eq!(artifacts[1].data, b"response");
}
#[tokio::test]
async fn cursor_trace_artifact_and_blob_are_written_atomically() {
let (_directory, store) = test_store("cursor-trace-atomic.db").await;
store.set_detailed_logging(true).await.unwrap();
store
.start_cursor_trace_if_detailed(
"request-atomic",
Some("conversation"),
"cursor_official",
Some("model"),
)
.await
.unwrap();
sqlx::query(
"CREATE TRIGGER reject_trace_artifact
BEFORE INSERT ON cursor_run_trace_artifacts
BEGIN
SELECT RAISE(ABORT, 'rejected artifact');
END",
)
.execute(store.pool())
.await
.unwrap();
assert!(store
.append_cursor_trace_artifact(
"request-atomic",
"run_sse_chunk",
"cursor_official",
b"must-rollback",
&serde_json::json!({}),
)
.await
.is_err());
let blob_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM blobs")
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(blob_count, 0);
}
#[tokio::test]
async fn records_one_summary_and_raw_payloads_for_one_provider_request() {
let app = Router::new().route(
"/proxy/generate",
post(|| async {
(
[(header::CONTENT_TYPE, "text/event-stream")],
concat!(
"data: {\"choices\":[{\"delta\":{\"content\":\"hi\"},\"finish_reason\":null}]}\n\n",
"data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":10,\"completion_tokens\":2,\"total_tokens\":12}}\n\n",
"data: [DONE]\n\n"
),
)
.into_response()
}),
);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
let (_directory, store) = test_store("observability.db").await;
store.set_detailed_logging(true).await.unwrap();
let model = store
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Display Model".into(),
model_type: ModelType::OpenAi,
base_url: format!("http://{address}/proxy/generate"),
use_full_url: true,
api_key: "not-recorded".into(),
tooltip_data: "Display Model".into(),
model_id: "actual-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: true,
custom_headers: serde_json::json!({"x-safe":"visible","authorization":"hidden"}),
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 = ProviderRouter::new(store.clone(), Duration::from_secs(5));
let events = provider
.stream(
ModelInvocation {
call_id: "call-1".into(),
run_id: "run-1".into(),
conversation_id: "conversation-1".into(),
provider_call_index: 0,
request: ModelRequest {
prompt: PromptSpec {
instructions: "system".into(),
tools: Vec::new(),
},
model: ModelSpec::new(model.model_hash),
history: Vec::new(),
},
},
CancellationToken::new(),
)
.collect::<Vec<_>>()
.await;
assert!(events.iter().all(Result::is_ok));
assert!(events
.iter()
.any(|event| matches!(event, Ok(ModelEvent::Done(_)))));
let call = store.llm_call("call-1").await.unwrap().unwrap();
assert_eq!(call.status, "completed");
assert_eq!(call.request_type, "openai-chat");
assert_eq!(call.request_url, format!("http://{address}/proxy/generate"));
assert_eq!(call.total_tokens, Some(12));
assert!(call.ttfb_ms.is_some());
assert!(call.ttfr_ms.is_some());
assert!(call.ttft_ms.is_some());
let request = store.llm_call_request("call-1").await.unwrap().unwrap();
assert_eq!(request.body["model"], "actual-model");
assert_eq!(request.headers["x-safe"], "visible");
assert!(request.headers.get("authorization").is_none());
assert!(!store.llm_call_chunks("call-1").await.unwrap().is_empty());
server.abort();
}
+634
View File
@@ -0,0 +1,634 @@
#[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,
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,79 @@
use std::collections::{BTreeMap, HashSet};
use cursor_server::{
cursor::tools::{
runtime::{CursorToolRuntime, ExecContext},
ToolBatchState, ToolDispatcher,
},
model::ToolCall,
};
fn tool(name: &str) -> ToolCall {
let arguments = serde_json::json!({
"shell_id": "runtime-shell",
"block_until_ms": 30_000
});
ToolCall {
index: 0,
call_id: "call-1".into(),
model_call_id: "model-call-1".into(),
name: name.into(),
arguments_text: arguments.to_string(),
arguments,
}
}
async fn dispatch(
name: &str,
) -> cursor_server::Result<cursor_server::cursor::tools::DispatchedTool> {
let dispatcher = ToolDispatcher::new(CursorToolRuntime::default());
let completed = HashSet::new();
let started = HashSet::new();
let call = tool(name);
let dispatched = dispatcher
.start_batch(
&[call],
ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
},
&[],
&BTreeMap::new(),
&ExecContext::default(),
)
.await?;
Ok(dispatched.into_iter().next().expect("one dispatched tool"))
}
#[tokio::test]
async fn await_shell_emitted_during_active_run_becomes_a_failed_tool_result() {
let dispatched = dispatch("AwaitShell").await.unwrap();
assert_eq!(
dispatched.messages.len(),
1,
"started card is still published"
);
let completion = dispatched.completion.expect("compatibility completion");
assert!(completion.result().is_error);
assert!(completion
.result()
.content
.contains("current advertised tool set"));
}
#[tokio::test]
async fn hallucinated_unknown_tool_becomes_a_failed_tool_result_instead_of_a_protocol_error() {
let dispatched = dispatch("OldTool").await.unwrap();
assert_eq!(
dispatched.messages.len(),
1,
"started card is still published"
);
let completion = dispatched.completion.expect("compatibility completion");
assert!(completion.result().is_error);
assert!(completion.result().content.contains("not available"));
}
+312
View File
@@ -0,0 +1,312 @@
#[path = "support/fixtures.rs"]
mod fixtures;
use cursor_server::model::{
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind,
};
fn prepared(
run_id: &str,
conversation_id: &ConversationId,
base_revision_id: cursor_server::model::RevisionId,
) -> PreparedRun {
PreparedRun {
run_id: RunId::new(run_id),
cursor_request_id: None,
conversation_id: conversation_id.clone(),
kind: RunKind::Root,
model: ModelSpec::new("test-model"),
prompt: PromptSpec {
instructions: "test".into(),
tools: Vec::new(),
},
initial_messages: Vec::new(),
action: RunAction::Resume {
pending_tool_round: None,
},
base_revision_id,
}
}
#[tokio::test]
async fn selecting_an_old_revision_creates_a_branch_without_old_suffixes() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let first = prepared("run-1", &conversation_id, root);
store.claim_run(&first).await.unwrap();
let a = fixtures::user("a", "A");
let revision_a = store
.append_revision(
&conversation_id,
&first.run_id,
root,
std::slice::from_ref(&a),
)
.await
.unwrap();
let b = fixtures::user("b", "B");
let revision_b = store
.append_revision(
&conversation_id,
&first.run_id,
revision_a,
std::slice::from_ref(&b),
)
.await
.unwrap();
let second = prepared("run-2", &conversation_id, revision_a);
let claimed = store.claim_run(&second).await.unwrap();
assert_eq!(claimed.replaced_run_id.as_ref(), Some(&first.run_id));
let c = fixtures::user("c", "C");
let revision_c = store
.append_revision(
&conversation_id,
&second.run_id,
revision_a,
std::slice::from_ref(&c),
)
.await
.unwrap();
assert_eq!(
store.load_revision_messages(revision_b).await.unwrap(),
vec![a.clone(), b]
);
assert_eq!(
store.load_revision_messages(revision_c).await.unwrap(),
vec![a, c]
);
assert!(store
.append_revision(
&conversation_id,
&first.run_id,
revision_b,
&[fixtures::user("late", "late")],
)
.await
.is_err());
}
#[tokio::test]
async fn reused_cursor_request_id_maps_to_the_current_distinct_execution() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("queued-conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let mut first = prepared("reused-request:11111111", &conversation_id, root);
first.cursor_request_id = Some("reused-request".into());
store.claim_run(&first).await.unwrap();
assert_eq!(
store
.active_run_for_cursor_request("reused-request")
.await
.unwrap(),
Some(first.run_id.clone())
);
let mut second = prepared("reused-request:22222222", &conversation_id, root);
second.cursor_request_id = Some("reused-request".into());
store.claim_run(&second).await.unwrap();
assert_eq!(
store
.active_run_for_cursor_request("reused-request")
.await
.unwrap(),
Some(second.run_id)
);
}
#[tokio::test]
async fn identical_runtime_event_is_exactly_once_and_conflicts_are_rejected() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("runtime");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let run = prepared("run", &conversation_id, root);
store.claim_run(&run).await.unwrap();
let event = cursor_server::model::RuntimeEvent {
event_id: "branch:changed:7".into(),
text: "runtime state changed".into(),
}
.into_message();
let (revision, inserted) = store
.append_message_once(&conversation_id, &run.run_id, root, &event)
.await
.unwrap();
assert!(inserted);
let (same, inserted) = store
.append_message_once(&conversation_id, &run.run_id, revision, &event)
.await
.unwrap();
assert_eq!(same, revision);
assert!(!inserted);
let conflict = cursor_server::model::RuntimeEvent {
event_id: "branch:changed:7".into(),
text: "different".into(),
}
.into_message();
assert!(store
.append_message_once(&conversation_id, &run.run_id, revision, &conflict)
.await
.is_err());
}
#[tokio::test]
async fn editing_a_logical_input_discards_its_active_suffix() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("edited-conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let input_id = "cursor:user:stable-id";
assert_eq!(
store
.anchor_input(&conversation_id, input_id, root)
.await
.unwrap(),
root
);
let first = prepared("first-run", &conversation_id, root);
store.claim_run(&first).await.unwrap();
let original = fixtures::user("original", "original text");
let original_revision = store
.append_revision(
&conversation_id,
&first.run_id,
root,
std::slice::from_ref(&original),
)
.await
.unwrap();
let suffix = fixtures::user("suffix", "old suffix");
let old_head = store
.append_revision(
&conversation_id,
&first.run_id,
original_revision,
std::slice::from_ref(&suffix),
)
.await
.unwrap();
let edit_base = store
.anchor_input(&conversation_id, input_id, old_head)
.await
.unwrap();
assert_eq!(edit_base, root);
let second = prepared("second-run", &conversation_id, edit_base);
store.claim_run(&second).await.unwrap();
let edited = fixtures::user("edited", "edited text");
let edited_head = store
.append_revision(
&conversation_id,
&second.run_id,
edit_base,
std::slice::from_ref(&edited),
)
.await
.unwrap();
assert_eq!(
store.load_revision_messages(edited_head).await.unwrap(),
vec![edited]
);
assert_eq!(
store.load_revision_messages(old_head).await.unwrap(),
vec![original, suffix]
);
}
#[tokio::test]
async fn retry_reuses_only_the_matching_initial_child_chain() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("retry-conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let first = prepared("first-run", &conversation_id, root);
store.claim_run(&first).await.unwrap();
let context = fixtures::user("request-context:event", "context");
let context_revision = store
.append_revision(
&conversation_id,
&first.run_id,
root,
std::slice::from_ref(&context),
)
.await
.unwrap();
let runtime = cursor_server::model::RuntimeEvent {
event_id: "cursor:user:stable-id:version".into(),
text: "query".into(),
}
.into_message();
let (partial_revision, partial_count) = store
.match_revision_prefix(&conversation_id, root, &[context.clone(), runtime.clone()])
.await
.unwrap();
assert_eq!(partial_revision, context_revision);
assert_eq!(partial_count, 1);
let runtime_revision = store
.append_revision(
&conversation_id,
&first.run_id,
context_revision,
std::slice::from_ref(&runtime),
)
.await
.unwrap();
let suffix = fixtures::user("old-answer", "old answer");
let old_head = store
.append_revision(
&conversation_id,
&first.run_id,
runtime_revision,
std::slice::from_ref(&suffix),
)
.await
.unwrap();
let (retry_base, reused) = store
.match_revision_prefix(&conversation_id, root, &[context.clone(), runtime.clone()])
.await
.unwrap();
assert_eq!(retry_base, runtime_revision);
assert_eq!(reused, 2);
assert_eq!(
store.load_revision_messages(retry_base).await.unwrap(),
vec![context.clone(), runtime.clone()]
);
let retry = prepared("retry-run", &conversation_id, retry_base);
let claimed = store.claim_run(&retry).await.unwrap();
assert_eq!(claimed.replaced_run_id.as_ref(), Some(&first.run_id));
let first_status: String = sqlx::query_scalar("SELECT status FROM runs WHERE run_id = ?")
.bind(first.run_id.as_str())
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(first_status, "cancelled");
let changed = cursor_server::model::RuntimeEvent {
event_id: "cursor:user:stable-id:changed-version".into(),
text: "edited query".into(),
}
.into_message();
let (changed_base, reused) = store
.match_revision_prefix(&conversation_id, root, &[context, changed])
.await
.unwrap();
assert_eq!(changed_base, context_revision);
assert_eq!(reused, 1);
assert_eq!(
store.load_revision_messages(old_head).await.unwrap(),
vec![
fixtures::user("request-context:event", "context"),
runtime,
suffix
]
);
}
+735
View File
@@ -0,0 +1,735 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{
connect,
prompting::{PromptAssets, PromptCompiler},
proto::agent::v1 as pb,
CursorCommand, CursorSessionRegistry,
},
model::{ContentPart, ProjectedContent},
provider::{FinishReason, ModelEvent},
store::{BlobId, Store},
};
use prost::Message;
#[tokio::test]
async fn retrying_the_same_edited_input_reuses_its_initial_branch() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
for (call, answer) in [
("model", "answer"),
("model-retry", "retry answer"),
("model-edit", "edited answer"),
("model-context-edit", "context edited answer"),
] {
provider.push(vec![
ModelEvent::Start {
model_call_id: call.into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta(answer.into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
}
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
for (request_id, text, visible_file) in [
("edited-input", "explain this", "/workspace/src/main.rs"),
(
"edited-input-retry",
"explain this",
"/workspace/src/main.rs",
),
(
"edited-input-changed",
"explain the edited version",
"/workspace/src/main.rs",
),
(
"edited-input-context-changed",
"explain this",
"/workspace/src/edited.rs",
),
] {
let handle = registry.get_or_create(request_id).await.unwrap();
let mut output = handle.subscribe();
let mut request = run_request(references(&store).await);
let Some(pb::agent_client_message::Message::RunRequest(run)) = request.message.as_mut()
else {
unreachable!("run_request always returns a RunRequest")
};
let Some(pb::conversation_action::Action::UserMessageAction(action)) = run
.action
.as_mut()
.and_then(|action| action.action.as_mut())
else {
unreachable!("run_request always contains a UserMessageAction")
};
action
.user_message
.as_mut()
.expect("run_request always contains a UserMessage")
.text = text.into();
let user = action
.user_message
.as_mut()
.expect("run_request always contains a UserMessage");
let Some(pb::invocation_context::Data::IdeState(ide)) = user
.selected_context
.as_mut()
.and_then(|selected| selected.invocation_context.as_mut())
.and_then(|invocation| invocation.data.as_mut())
else {
unreachable!("run_request always contains IDE state")
};
ide.visible_files[0].path = visible_file.into();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(request),
})
.await
.unwrap();
let mut seqno = 1;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("retry must finish without closing the stream early");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
let end = serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
assert_eq!(end, serde_json::json!({}));
break;
}
let message = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = message.message {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
}
}
let requests = provider.requests();
assert_eq!(requests.len(), 4);
assert_eq!(requests[1].history, requests[0].history);
assert_eq!(requests[2].history.len(), requests[0].history.len());
assert_ne!(
requests[2].history.last().unwrap().message_id,
requests[0].history.last().unwrap().message_id
);
let ProjectedContent::Parts(parts) = &requests[2].history.last().unwrap().content else {
panic!("edited runtime message must use typed parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("this fixture has no images")
};
assert!(text.contains("<user_query>\nexplain the edited version\n</user_query>"));
assert_ne!(
requests[3].history.last().unwrap().message_id,
requests[0].history.last().unwrap().message_id
);
let ProjectedContent::Parts(parts) = &requests[3].history.last().unwrap().content else {
panic!("context-edited runtime message must use typed parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("this fixture has no images")
};
assert!(text.contains("/workspace/src/edited.rs"));
}
#[tokio::test]
async fn unchanged_request_context_is_not_repeated_and_preserves_the_provider_prefix() {
let (_directory, store) = fixtures::temp_store().await;
let first_references = references(&store).await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "model".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("answer".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "model-2".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("answer again".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("ask-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(run_request(first_references)),
})
.await
.unwrap();
let mut seqno = 1;
let mut 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 message = pb::AgentServerMessage::decode(payload).unwrap();
match message.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
checkpoint = Some(state);
}
_ => {}
}
}
let requests = provider.requests();
let request = &requests[0];
assert!(request
.prompt
.tools
.iter()
.any(|tool| tool.name == "AskQuestion"));
assert!(!request
.prompt
.tools
.iter()
.any(|tool| tool.name == "GenerateImage"));
assert_eq!(request.history.len(), 2);
assert!(request.history[0]
.message_id
.starts_with("request-context:"));
let ProjectedContent::Parts(context_parts) = &request.history[0].content else {
panic!("request context message must use typed parts")
};
let [ContentPart::Text { text: context_text }] = context_parts.as_slice() else {
panic!("request context message must contain one text part")
};
assert!(request.history[1]
.message_id
.starts_with("runtime:cursor:user:wire-user:"));
assert!(!request.prompt.instructions.contains("workspace rule"));
assert!(!request.prompt.instructions.contains("<mcp_meta_tools>"));
let ProjectedContent::Parts(parts) = &request.history[1].content else {
panic!("runtime message must use typed parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("this fixture has no images")
};
for expected in [
"<user_rule>\nworkspace rule\n</user_rule>",
"<agent_skill fullPath=\"/skills/test/SKILL.md\">test skill</agent_skill>",
"<subagent name=\"reviewer\">review code</subagent>",
"<mcp_meta_tool_server name=\"test\" identifier=\"mcp-test\">",
"<mcp_tool name=\"lookup\">",
"<definition_path>/tmp/mcp-test/lookup.json</definition_path>",
"<input_schema>{&quot;properties&quot;:{&quot;query&quot;:{&quot;type&quot;:&quot;string&quot;}},&quot;type&quot;:&quot;object&quot;}</input_schema>",
"Call a listed tool directly with CallMcpTool without calling GetMcpTools first.",
] {
assert!(
context_text.contains(expected),
"missing request context section: {expected}"
);
}
assert!(!context_text.contains("complete skill body"));
assert!(!context_text.contains("complete MCP server instructions"));
for expected in [
"Ask mode is active.",
"<user_query>\nexplain this\n</user_query>",
] {
assert!(
text.contains(expected),
"missing runtime section: {expected}"
);
}
assert!(!text.contains("<rules>"));
assert!(!text.contains("<mcp_meta_tools>"));
assert!(text.contains("/workspace/src/main.rs"));
let second = registry.get_or_create("ask-request-2").await.unwrap();
let mut second_output = second.subscribe();
second
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(run_request_with_state(
references(&store).await,
checkpoint.expect("first Run must publish a checkpoint"),
)),
})
.await
.unwrap();
let mut second_seqno = 1;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), second_output.recv())
.await
.unwrap()
.unwrap();
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
break;
}
let message = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = message.message {
second
.command(CursorCommand::Append {
seqno: second_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
second_seqno += 1;
}
}
let requests = provider.requests();
assert_eq!(requests.len(), 2);
assert_eq!(
requests[1].prompt.instructions, requests[0].prompt.instructions,
"unchanged request context must not rewrite the system prompt"
);
assert_eq!(
requests[1].history[..requests[0].history.len()],
requests[0].history,
"the previous provider history must remain an exact prefix"
);
assert_eq!(
requests[1]
.history
.iter()
.filter(|message| message.message_id.starts_with("request-context:"))
.count(),
1,
"identical request context must not be appended again"
);
}
#[tokio::test]
async fn missing_context_parts_use_current_cursor_response_and_cache_its_content() {
let (_directory, store) = fixtures::temp_store().await;
let referenced_context = fixture_context();
let references = references_for(&referenced_context);
let mut current_context = referenced_context.clone();
current_context
.mcp_meta_tool_options
.as_mut()
.unwrap()
.mcp_descriptors
.push(pb::McpDescriptor {
server_identifier: "live-mcp".into(),
tools: vec![pb::McpToolDescriptor {
tool_name: "current-tool".into(),
..Default::default()
}],
..Default::default()
});
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "model".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("answer".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("context-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(run_request(references)),
})
.await
.unwrap();
let mut seqno = 1;
let mut requested_context = 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 {
break;
}
let message = pb::AgentServerMessage::decode(payload).unwrap();
match message.message {
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
assert_eq!(exec.id, 0);
let Some(pb::exec_server_message::Message::RequestContextArgs(args)) = exec.message
else {
panic!("missing context must use RequestContextArgs")
};
assert_eq!(args.notes_session_id.as_deref(), Some("mode-conversation"));
requested_context = true;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(context_stream_close()),
})
.await
.unwrap();
seqno += 1;
handle
.command(CursorCommand::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(
current_context.clone(),
),
..Default::default()
},
),
),
},
),
),
..Default::default()
},
)),
}),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
_ => {}
}
}
assert!(requested_context);
let requests = provider.requests();
assert_eq!(requests.len(), 1);
let ProjectedContent::Parts(parts) = &requests[0].history[0].content else {
panic!("request context message must use typed parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("request context message must contain one text part")
};
assert!(text.contains("<mcp_meta_tool_server name=\"live-mcp\" identifier=\"live-mcp\">"));
assert!(text.contains("<mcp_tool name=\"current-tool\">"));
let stale = references_for(&referenced_context);
let current = references_for(&current_context);
for (id, _) in [
current.rules,
current.skills,
current.subagents,
current.mcps,
] {
assert!(store.get_blob(&id).await.unwrap().is_some());
}
assert!(store.get_blob(&stale.mcps.0).await.unwrap().is_none());
}
struct References {
rules: (BlobId, u32),
skills: (BlobId, u32),
subagents: (BlobId, u32),
mcps: (BlobId, u32),
}
async fn references(store: &Store) -> References {
let context = fixture_context();
let references = references_for(&context);
for data in part_data(&context) {
store.put_blob(&data, &[]).await.unwrap();
}
references
}
fn fixture_context() -> pb::RequestContext {
pb::RequestContext {
rules: vec![pb::CursorRule {
full_path: "/workspace/AGENTS.md".into(),
content: "workspace rule".into(),
..Default::default()
}],
non_file_rules: vec![pb::CursorRule {
full_path: "/skills/test/SKILL.md".into(),
content: "complete skill body".into(),
..Default::default()
}],
agent_skills: vec![pb::AgentSkill {
full_path: "/skills/test/SKILL.md".into(),
description: "test skill".into(),
..Default::default()
}],
custom_subagents: vec![pb::CustomSubagent {
name: "reviewer".into(),
description: "review code".into(),
..Default::default()
}],
mcp_meta_tool_options: Some(pb::McpMetaToolOptions {
enabled: true,
mcp_descriptors: vec![pb::McpDescriptor {
server_name: "test".into(),
server_identifier: "mcp-test".into(),
server_use_instructions: Some("complete MCP server instructions".into()),
tools: vec![pb::McpToolDescriptor {
tool_name: "lookup".into(),
definition_path: Some("/tmp/mcp-test/lookup.json".into()),
description: Some("look up a value".into()),
input_schema_json: Some(
r#"{"type":"object","properties":{"query":{"type":"string"}}}"#.into(),
),
..Default::default()
}],
..Default::default()
}],
}),
..Default::default()
}
}
fn references_for(context: &pb::RequestContext) -> References {
let mut parts = part_data(context).into_iter();
References {
rules: reference(&parts.next().unwrap()),
skills: reference(&parts.next().unwrap()),
subagents: reference(&parts.next().unwrap()),
mcps: reference(&parts.next().unwrap()),
}
}
fn part_data(context: &pb::RequestContext) -> Vec<Vec<u8>> {
vec![
pb::RequestContextRulesPart {
rules: context.rules.clone(),
non_file_rules: context.non_file_rules.clone(),
cloud_rule: context.cloud_rule.clone(),
}
.encode_to_vec(),
pb::RequestContextSkillsPart {
agent_skills: context.agent_skills.clone(),
skill_options: context.skill_options.clone(),
}
.encode_to_vec(),
pb::RequestContextSubagentsPart {
custom_subagents: context.custom_subagents.clone(),
}
.encode_to_vec(),
pb::RequestContextMcpsPart {
tools: context.tools.clone(),
mcp_instructions: context.mcp_instructions.clone(),
mcp_file_system_options: context.mcp_file_system_options.clone(),
mcp_meta_tool_options: context.mcp_meta_tool_options.clone(),
}
.encode_to_vec(),
]
}
fn reference(data: &[u8]) -> (BlobId, u32) {
(BlobId::digest(data), data.len() as u32)
}
fn run_request(references: References) -> pb::AgentClientMessage {
let (rules, rules_byte_length) = references.rules;
let (skills, skills_byte_length) = references.skills;
let (subagents, subagents_byte_length) = references.subagents;
let (mcps, mcps_byte_length) = references.mcps;
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
conversation_state: Some(pb::ConversationStateStructure {
mode: Some(pb::AgentMode::Agent as i32),
..Default::default()
}),
action: Some(pb::ConversationAction {
request_context_parts: Some(pb::RequestContextPartReferences {
rules_blob_id: rules.as_bytes().to_vec(),
rules_byte_length,
skills_blob_id: skills.as_bytes().to_vec(),
skills_byte_length,
subagents_blob_id: subagents.as_bytes().to_vec(),
subagents_byte_length,
mcps_blob_id: mcps.as_bytes().to_vec(),
mcps_byte_length,
dynamic_context: Some(pb::RequestContext {
env: Some(pb::RequestContextEnv {
os_version: "darwin".into(),
workspace_paths: vec!["/workspace".into()],
shell: "zsh".into(),
time_zone: "UTC".into(),
..Default::default()
}),
..Default::default()
}),
}),
action: Some(pb::conversation_action::Action::UserMessageAction(
pb::UserMessageAction {
user_message: Some(pb::UserMessage {
text: "explain this".into(),
message_id: "wire-user".into(),
mode: pb::AgentMode::Ask as i32,
selected_context: Some(pb::SelectedContext {
invocation_context: Some(pb::InvocationContext {
data: Some(pb::invocation_context::Data::IdeState(
pb::invocation_context::IdeState {
visible_files: vec![
pb::invocation_context::ide_state::File {
path: "/workspace/src/main.rs".into(),
total_lines: 10,
..Default::default()
},
],
..Default::default()
},
)),
}),
..Default::default()
}),
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("mode-conversation".into()),
run_id: Some("wire-run".into()),
requested_model: Some(pb::RequestedModel {
model_id: "test-model".into(),
..Default::default()
}),
..Default::default()
},
)),
}
}
fn run_request_with_state(
references: References,
state: pb::ConversationStateStructure,
) -> pb::AgentClientMessage {
let mut message = run_request(references);
let Some(pb::agent_client_message::Message::RunRequest(request)) = message.message.as_mut()
else {
unreachable!("run_request always returns a RunRequest")
};
request.conversation_state = Some(state);
let Some(pb::conversation_action::Action::UserMessageAction(action)) = request
.action
.as_mut()
.and_then(|action| action.action.as_mut())
else {
unreachable!("run_request always contains a UserMessageAction")
};
action
.user_message
.as_mut()
.expect("run_request always contains a UserMessage")
.message_id = "wire-user-2".into();
message
}
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 },
)),
},
)),
}
}
fn context_stream_close() -> pb::AgentClientMessage {
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 },
)),
},
)),
}
}
+54
View File
@@ -0,0 +1,54 @@
#[path = "support/fixtures.rs"]
mod fixtures;
use cursor_server::model::{
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind, RuntimeEvent,
};
#[tokio::test]
async fn runtime_event_is_appended_exactly_once() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let run = PreparedRun {
run_id: RunId::new("run"),
cursor_request_id: None,
conversation_id: conversation_id.clone(),
kind: RunKind::Root,
model: ModelSpec::new("model"),
prompt: PromptSpec {
instructions: String::new(),
tools: Vec::new(),
},
initial_messages: Vec::new(),
action: RunAction::Resume {
pending_tool_round: None,
},
base_revision_id: root,
};
store.claim_run(&run).await.unwrap();
let message = RuntimeEvent {
event_id: "branch:changed:7".into(),
text: "runtime state changed".into(),
}
.into_message();
let (revision, inserted) = store
.append_message_once(&conversation_id, &run.run_id, root, &message)
.await
.unwrap();
assert!(inserted);
assert!(
!store
.append_message_once(&conversation_id, &run.run_id, revision, &message)
.await
.unwrap()
.1
);
let messages = store.load_revision_messages(revision).await.unwrap();
assert_eq!(messages.len(), 1);
assert_eq!(
messages[0].runtime_event_id.as_deref(),
Some("branch:changed:7")
);
}
+225
View File
@@ -0,0 +1,225 @@
use std::borrow::Cow;
use cursor_server::store::Store;
use sqlx::{migrate::Migrator, sqlite::SqliteConnectOptions, Row};
#[tokio::test]
async fn version_two_database_upgrades_with_cursor_request_mapping() {
let directory = tempfile::tempdir().unwrap();
let database = directory.path().join("upgrade.db");
let pool = sqlx::SqlitePool::connect_with(
SqliteConnectOptions::new()
.filename(&database)
.create_if_missing(true),
)
.await
.unwrap();
let all = sqlx::migrate!("./migrations");
let prior = Migrator {
migrations: Cow::Owned(
all.iter()
.filter(|migration| migration.version <= 2)
.cloned()
.collect(),
),
ignore_missing: false,
locking: true,
no_tx: false,
};
prior.run(&pool).await.unwrap();
drop(pool);
let store = Store::connect(&format!("sqlite://{}", database.display()))
.await
.unwrap();
let columns = sqlx::query("PRAGMA table_info(runs)")
.fetch_all(store.pool())
.await
.unwrap();
assert!(columns
.iter()
.any(|column| column.get::<String, _>("name") == "cursor_request_id"));
let llm_call_columns = sqlx::query("PRAGMA table_info(llm_calls)")
.fetch_all(store.pool())
.await
.unwrap();
assert!(llm_call_columns
.iter()
.any(|column| column.get::<String, _>("name") == "first_valid_response_at_ms"));
assert!(llm_call_columns
.iter()
.any(|column| column.get::<String, _>("name") == "ttfr_ms"));
}
#[tokio::test]
async fn provider_and_model_rows_upgrade_to_flat_model_configuration() {
let directory = tempfile::tempdir().unwrap();
let database = directory.path().join("flat-model-upgrade.db");
let pool = sqlx::SqlitePool::connect_with(
SqliteConnectOptions::new()
.filename(&database)
.create_if_missing(true)
.foreign_keys(true),
)
.await
.unwrap();
let all = sqlx::migrate!("./migrations");
let prior = Migrator {
migrations: Cow::Owned(
all.iter()
.filter(|migration| migration.version <= 3)
.cloned()
.collect(),
),
ignore_missing: false,
locking: true,
no_tx: false,
};
prior.run(&pool).await.unwrap();
sqlx::query(
r#"INSERT INTO provider_endpoints(
name, provider_type, base_url, api_key, custom_headers_json, extra_params_json,
created_at_ms, updated_at_ms
) VALUES ('Example', 'openai-chat', 'https://example.com/v1', 'secret',
'{"x-client":"cursor-byok"}', '{"service_tier":"priority"}', 10, 11)"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"INSERT INTO provider_models(
model_hash, provider_id, model_id, display_name, endpoint_type, request_url,
enabled, sort_order, reasoning_enabled, supports_image_generation,
created_at_ms, updated_at_ms
) VALUES
('rspns001', 1, 'model-b', 'Model B', 'openai-responses',
'https://proxy.example.com/arbitrary/generate?api-version=2026-01-01',
1, 5, 0, 0, 12, 13),
('anthr001', 1, 'model-c', 'Model C', 'anthropic', '/proxy/claude',
1, 6, 0, 0, 12, 13),
('anthstd1', 1, 'model-d', 'Model D', 'anthropic', '',
1, 7, 0, 0, 12, 13)"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"INSERT INTO provider_models(
model_hash, provider_id, model_id, display_name, endpoint_type, request_url,
enabled, sort_order, context_window_tokens, max_output_tokens,
reasoning_enabled, reasoning_effort, supports_image_generation,
created_at_ms, updated_at_ms
) VALUES ('12345678', 1, 'model-a', 'Model A', 'openai-chat', '',
1, 4, 200000, 8192, 1, 'high', 0, 12, 13)"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"INSERT INTO llm_calls(
call_id, run_id, conversation_id, provider_call_index, model_hash,
provider_type, provider_url, request_type, request_url, model_id, display_name,
status, created_at_ms, message_count, tool_count, detailed
) VALUES ('call-1', 'run-1', 'conversation-1', 0, '12345678',
'openai-chat', 'https://example.com/v1', 'openai-chat',
'https://example.com/v1/chat/completions', 'model-a', 'Model A',
'completed', 14, 1, 0, 0)"#,
)
.execute(&pool)
.await
.unwrap();
all.run(&pool).await.unwrap();
drop(pool);
let store = Store::connect(&format!("sqlite://{}", database.display()))
.await
.unwrap();
let row = sqlx::query(
r#"SELECT model_hash, sort_order, display_name, model_type, base_url, api_key,
tooltip_data, model_id, reasoning_effort, openai_endpoint, use_full_url,
openai_extra_params_enabled, openai_extra_params_json,
custom_headers_enabled, custom_headers_json, context_window_tokens,
max_completion_tokens
FROM model_configs WHERE model_hash = '12345678'"#,
)
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(row.get::<String, _>("model_hash"), "12345678");
assert_eq!(row.get::<i64, _>("sort_order"), 4);
assert_eq!(row.get::<String, _>("model_type"), "openai");
assert_eq!(row.get::<String, _>("base_url"), "https://example.com/v1");
assert_eq!(row.get::<String, _>("api_key"), "secret");
assert_eq!(row.get::<i64, _>("use_full_url"), 0);
assert_eq!(row.get::<String, _>("reasoning_effort"), "high");
assert_eq!(
row.get::<String, _>("openai_endpoint"),
"/v1/chat/completions"
);
assert_eq!(row.get::<i64, _>("openai_extra_params_enabled"), 1);
assert_eq!(
row.get::<String, _>("openai_extra_params_json"),
r#"{"service_tier":"priority"}"#
);
assert_eq!(row.get::<i64, _>("custom_headers_enabled"), 1);
assert_eq!(row.get::<i64, _>("context_window_tokens"), 200000);
assert_eq!(row.get::<i64, _>("max_completion_tokens"), 8192);
let migrated_rows = sqlx::query(
"SELECT model_hash, model_type, base_url, use_full_url, openai_endpoint FROM model_configs WHERE model_hash IN ('rspns001', 'anthr001', 'anthstd1') ORDER BY model_hash",
)
.fetch_all(store.pool())
.await
.unwrap();
assert_eq!(migrated_rows[0].get::<String, _>("model_hash"), "anthr001");
assert_eq!(migrated_rows[0].get::<String, _>("model_type"), "anthropic");
assert_eq!(
migrated_rows[0].get::<String, _>("base_url"),
"https://example.com/v1/proxy/claude"
);
assert_eq!(migrated_rows[0].get::<i64, _>("use_full_url"), 1);
assert_eq!(migrated_rows[0].get::<String, _>("openai_endpoint"), "");
assert_eq!(migrated_rows[1].get::<String, _>("model_hash"), "anthstd1");
assert_eq!(migrated_rows[1].get::<String, _>("model_type"), "anthropic");
assert_eq!(
migrated_rows[1].get::<String, _>("base_url"),
"https://example.com/v1"
);
assert_eq!(migrated_rows[1].get::<i64, _>("use_full_url"), 0);
assert_eq!(migrated_rows[1].get::<String, _>("openai_endpoint"), "");
assert_eq!(migrated_rows[2].get::<String, _>("model_hash"), "rspns001");
assert_eq!(migrated_rows[2].get::<String, _>("model_type"), "openai");
assert_eq!(
migrated_rows[2].get::<String, _>("base_url"),
"https://proxy.example.com/arbitrary/generate?api-version=2026-01-01"
);
assert_eq!(migrated_rows[2].get::<i64, _>("use_full_url"), 1);
assert_eq!(
migrated_rows[2].get::<String, _>("openai_endpoint"),
"/v1/responses"
);
assert_eq!(
sqlx::query_scalar::<_, String>(
"SELECT model_hash FROM llm_calls WHERE call_id = 'call-1'"
)
.fetch_one(store.pool())
.await
.unwrap(),
"12345678"
);
assert!(sqlx::query("SELECT 1 FROM provider_endpoints")
.fetch_one(store.pool())
.await
.is_err());
assert!(sqlx::query("SELECT 1 FROM provider_models")
.fetch_one(store.pool())
.await
.is_err());
assert!(sqlx::query("PRAGMA foreign_key_check")
.fetch_all(store.pool())
.await
.unwrap()
.is_empty());
}
+239
View File
@@ -0,0 +1,239 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::sync::Arc;
use cursor_server::{
cursor::{
connect,
prompting::{PromptAssets, PromptCompiler},
proto::agent::v1 as pb,
CursorCommand, CursorSessionRegistry,
},
model::{ContentPart, ProjectedContent},
provider::{FinishReason, ModelEvent},
store::BlobId,
};
use prost::Message;
#[tokio::test]
async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_checkpoint() {
let (_directory, store) = fixtures::temp_store().await;
let stored_image = vec![4, 5];
let stored_image_id = store.put_blob(&stored_image, &[]).await.unwrap();
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "model".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("seen".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("image-run").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(request(&stored_image_id)),
})
.await
.unwrap();
let mut seqno = 1;
let mut checkpoints = Vec::new();
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 message = pb::AgentServerMessage::decode(payload).unwrap();
match message.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
checkpoints.push(state)
}
_ => {}
}
}
let requests = provider.requests();
let user = requests[0]
.history
.iter()
.find(|message| {
message
.message_id
.starts_with("runtime:cursor:user:image-user:")
})
.unwrap();
let ProjectedContent::Parts(parts) = &user.content else {
panic!("runtime user message must retain typed parts")
};
assert!(matches!(
&parts[0],
ContentPart::Text { text } if text.contains("<user_query>\nwhat is this?\n</user_query>")
));
assert_eq!(
&parts[1..],
&[
ContentPart::Image {
mime_type: "image/png".into(),
data: vec![1, 2, 3],
},
ContentPart::Image {
mime_type: "image/webp".into(),
data: stored_image,
},
ContentPart::Image {
mime_type: "image/jpeg".into(),
data: vec![6, 7, 8],
},
]
);
let serialized_request = serde_json::to_string(&requests[0]).unwrap();
assert!(!serialized_request.contains("private-image-uuid"));
assert!(!serialized_request.contains("/private/image/path"));
let state = checkpoints
.iter()
.find(|state| state.root_prompt_messages_json.len() >= 2)
.unwrap();
let mut user_root = None;
for raw_id in &state.root_prompt_messages_json {
let id = BlobId::from_bytes(raw_id).unwrap();
let bytes = store.get_blob(&id).await.unwrap().unwrap();
let value: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
if value["id"]
.as_str()
.is_some_and(|id| id.starts_with("runtime:cursor:user:image-user:"))
{
user_root = Some(value);
break;
}
}
let user_root = user_root.unwrap();
assert_eq!(user_root["content"][1]["type"], "image");
assert_eq!(user_root["content"][1]["mimeType"], "image/png");
assert_eq!(user_root["content"][1]["image"], "AQID");
assert!(user_root["content"][1].get("data").is_none());
assert_eq!(user_root["content"][2]["mimeType"], "image/webp");
assert_eq!(user_root["content"][2]["image"], "BAU=");
assert!(user_root["content"][2].get("data").is_none());
assert_eq!(user_root["content"][3]["mimeType"], "image/jpeg");
assert_eq!(user_root["content"][3]["image"], "BgcI");
assert!(user_root["content"][3].get("data").is_none());
}
fn request(stored_image_id: &BlobId) -> pb::AgentClientMessage {
let data = vec![1, 2, 3];
let inline_blob_data = vec![6, 7, 8];
let inline_blob_id = BlobId::digest(&inline_blob_data);
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
requested_model: Some(pb::RequestedModel {
model_id: "test-model".into(),
..Default::default()
}),
action: Some(pb::ConversationAction {
action: Some(pb::conversation_action::Action::UserMessageAction(
pb::UserMessageAction {
user_message: Some(pb::UserMessage {
text: "what is this?".into(),
message_id: "image-user".into(),
selected_context: Some(pb::SelectedContext {
selected_images: vec![
pb::SelectedImage {
uuid: "private-image-uuid".into(),
path: "/private/image/path".into(),
mime_type: "image/png".into(),
data_or_blob_id: Some(
pb::selected_image::DataOrBlobId::Data(data),
),
..Default::default()
},
pb::SelectedImage {
uuid: "stored-image".into(),
path: "/private/stored/path".into(),
mime_type: "image/webp".into(),
data_or_blob_id: Some(
pb::selected_image::DataOrBlobId::BlobId(
stored_image_id.as_bytes().to_vec(),
),
),
..Default::default()
},
pb::SelectedImage {
uuid: "inline-blob-image".into(),
path: "/private/inline/path".into(),
mime_type: "image/jpeg".into(),
data_or_blob_id: Some(
pb::selected_image::DataOrBlobId::BlobIdWithData(
pb::selected_image::BlobIdWithData {
blob_id: inline_blob_id.as_bytes().to_vec(),
data: inline_blob_data,
},
),
),
..Default::default()
},
],
..Default::default()
}),
mode: pb::AgentMode::Agent as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("image-conversation".into()),
run_id: Some("image-run".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 },
)),
},
)),
}
}
+299
View File
@@ -0,0 +1,299 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{sync::Arc, time::Duration};
use cursor_server::{
cursor::{
connect,
prompting::{PromptAssets, PromptCompiler},
proto::agent::v1 as pb,
CursorCommand, CursorSessionHandle, CursorSessionRegistry,
},
provider::{FinishReason, ModelEvent},
};
use prost::Message;
#[tokio::test]
async fn every_bidi_run_resolves_and_persists_its_own_subagent_model_and_background_state() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
for suffix in ["a", "b"] {
provider.push(task_response(suffix));
provider.push(stop_response(suffix));
}
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider),
PromptCompiler::new(assets),
Default::default(),
);
let first = registry.get_or_create("subagent-run-a").await.unwrap();
let first_checkpoint = drive(
&first,
run_request("subagent-run-a", "user-a", "model-a", None),
"model-a",
"child-a",
)
.await;
let second = registry.get_or_create("subagent-run-b").await.unwrap();
drive(
&second,
run_request(
"subagent-run-b",
"user-b",
"model-b",
Some(first_checkpoint),
),
"model-b",
"child-b",
)
.await;
}
async fn drive(
handle: &CursorSessionHandle,
request: pb::AgentClientMessage,
expected_model: &str,
child_id: &str,
) -> pb::ConversationStateStructure {
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(request),
})
.await
.unwrap();
let mut seqno = 1;
let mut saw_started = false;
let mut saw_exec = false;
let mut saw_completed = false;
let mut checkpoint = None;
loop {
let frame = tokio::time::timeout(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;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
let Some(pb::exec_server_message::Message::SubagentArgs(args)) = exec.message
else {
continue;
};
assert_eq!(args.model_id, expected_model);
assert_eq!(args.run_in_background, Some(true));
saw_exec = true;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(subagent_result(exec.id, child_id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
match update.message {
Some(pb::interaction_update::Message::ToolCallStarted(started)) => {
let task = task(started.tool_call.as_ref().unwrap());
assert_eq!(
task.args.as_ref().unwrap().model.as_deref(),
Some(expected_model)
);
saw_started = true;
}
Some(pb::interaction_update::Message::ToolCallCompleted(completed)) => {
let task = task(completed.tool_call.as_ref().unwrap());
let Some(pb::task_result::Result::Success(success)) = task
.result
.as_ref()
.and_then(|result| result.result.as_ref())
else {
panic!("expected Task success")
};
assert_eq!(
task.args.as_ref().unwrap().model.as_deref(),
Some(expected_model)
);
assert!(success.is_background);
assert_eq!(success.agent_id.as_deref(), Some(child_id));
saw_completed = true;
}
_ => {}
}
}
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state))
if state.pending_tool_calls.is_empty() =>
{
checkpoint = Some(state);
}
_ => {}
}
}
assert!(saw_started && saw_exec && saw_completed);
let checkpoint = checkpoint.expect("settled checkpoint");
let state = checkpoint
.subagent_states
.get(child_id)
.expect("background subagent persisted state");
assert_eq!(state.model_id.as_deref(), Some(expected_model));
let run = checkpoint
.subagent_runs_by_parent_tool_call_id
.get(&format!("task-{child_id}"))
.expect("background subagent run state");
assert_eq!(run.status, pb::SubagentRunStatus::Backgrounded as i32);
checkpoint
}
fn task(call: &pb::ToolCall) -> &pb::TaskToolCall {
let Some(pb::tool_call::Tool::TaskToolCall(task)) = call.tool.as_ref() else {
panic!("expected TaskToolCall")
};
task
}
fn run_request(
request_id: &str,
user_id: &str,
subagent_model: &str,
conversation_state: Option<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::UserMessageAction(
pb::UserMessageAction {
user_message: Some(pb::UserMessage {
text: format!("start {request_id}"),
message_id: user_id.into(),
mode: pb::AgentMode::Multitask as i32,
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some("subagent-e2e-conversation".into()),
run_id: Some(request_id.into()),
requested_model: Some(pb::RequestedModel {
model_id: "parent-model".into(),
..Default::default()
}),
conversation_state,
subagent_model_overrides: vec![pb::SubagentModelOverride {
subagent_type: "generalPurpose".into(),
selection: Some(pb::subagent_model_override::Selection::Model(
pb::RequestedModel {
model_id: subagent_model.into(),
..Default::default()
},
)),
}],
..Default::default()
},
)),
}
}
fn task_response(suffix: &str) -> Vec<ModelEvent> {
let child_id = format!("child-{suffix}");
let arguments = serde_json::json!({
"description": format!("background {suffix}"),
"prompt": "inspect",
"subagent_type": "generalPurpose",
"run_in_background": true
})
.to_string();
vec![
ModelEvent::Start {
model_call_id: format!("model-call-{suffix}"),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: format!("task-{child_id}"),
name: "Task".into(),
},
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: arguments,
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]
}
fn stop_response(suffix: &str) -> Vec<ModelEvent> {
vec![
ModelEvent::Start {
model_call_id: format!("final-{suffix}"),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("background task started".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]
}
fn subagent_result(id: u32, child_id: &str) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
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(),
final_message: Some("running in background".into()),
background_reason: pb::SubagentBackgroundReason::AgentRequest as i32,
..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 },
)),
},
)),
}
}
+208
View File
@@ -0,0 +1,208 @@
use std::collections::{BTreeMap, HashMap, HashSet};
use cursor_server::{
cursor::{
interaction,
proto::agent::v1 as pb,
tools::{
codec,
runtime::{CursorToolRuntime, ExecContext, SubagentModel},
ToolBatchState, ToolDispatcher,
},
},
model::{CanonicalMessage, MessageContent, Origin, Role, ToolCall},
};
#[test]
fn task_keeps_wire_type_model_parent_and_background_fields() {
let mut context = context();
context.subagent_model = Some(SubagentModel::Model("guide-model".into()));
let call = task_call(serde_json::json!({
"description": "guide",
"prompt": "inspect",
"subagent_type": "cursor-guide",
"run_in_background": true,
"interrupt": false,
}));
let call = context.prepare_call(&call).unwrap();
let message = codec::request(7, &call, &context).unwrap();
let pb::agent_server_message::Message::ExecServerMessage(exec) = message.message.unwrap()
else {
panic!("expected ExecServerMessage")
};
assert_eq!(exec.id, 7);
assert_eq!(exec.exec_id, "task-call");
assert_eq!(exec.accept_hook_additional_contexts, Some(false));
let pb::exec_server_message::Message::SubagentArgs(args) = exec.message.unwrap() else {
panic!("expected SubagentArgs")
};
assert_eq!(args.subagent_type, "cursor-guide");
assert_eq!(args.model_id, "guide-model");
assert_eq!(args.parent_conversation_id.as_deref(), Some("child"));
assert_eq!(args.root_parent_conversation_id.as_deref(), Some("root"));
assert_eq!(args.run_in_background, Some(true));
assert_eq!(args.interrupt, Some(false));
let rendered = interaction::render_tool_call(&call, false).unwrap();
let Some(pb::tool_call::Tool::TaskToolCall(task)) = rendered.tool else {
panic!("expected TaskToolCall")
};
assert_eq!(task.args.unwrap().model.as_deref(), Some("guide-model"));
}
#[test]
fn task_uses_explicit_call_model_then_the_run_default() {
let mut explicit = task_call(serde_json::json!({
"description": "task",
"prompt": "inspect",
"subagent_type": "generalPurpose",
"model": "call-model",
}));
explicit = context().prepare_call(&explicit).unwrap();
assert_eq!(explicit.arguments["model"], "call-model");
let default = task_call(serde_json::json!({
"description": "task",
"prompt": "inspect",
"subagent_type": "generalPurpose",
}));
let default = context().prepare_call(&default).unwrap();
assert_eq!(default.arguments["model"], "parent-model");
}
#[test]
fn task_renders_general_typed_and_custom_subagent_types_without_aliases() {
let cases = [
("generalPurpose", "unspecified"),
("cursor-guide", "cursor-guide"),
("MyReviewer", "MyReviewer"),
];
for (name, expected) in cases {
let call = task_call(serde_json::json!({
"description": "task",
"prompt": "inspect",
"subagent_type": name,
}));
let rendered = interaction::render_tool_call(&call, false).unwrap();
let Some(pb::tool_call::Tool::TaskToolCall(tool)) = rendered.tool else {
panic!("expected TaskToolCall")
};
let subagent = tool.args.unwrap().subagent_type.unwrap().r#type.unwrap();
match (expected, subagent) {
("unspecified", pb::subagent_type::Type::Unspecified(_)) => {}
("cursor-guide", pb::subagent_type::Type::CursorGuide(_)) => {}
(custom, pb::subagent_type::Type::Custom(value)) => {
assert_eq!(value.name, custom)
}
_ => panic!("wrong subagent oneof for {name}"),
}
}
}
#[test]
fn disabled_task_model_is_left_for_the_model_visible_reminder() {
let mut context = context();
context.subagent_model = Some(SubagentModel::Disabled);
let call = task_call(serde_json::json!({
"description": "review",
"prompt": "inspect",
"subagent_type": "security-review",
}));
assert!(context.task_disabled(&call));
assert!(context
.prepare_call(&call)
.unwrap()
.arguments
.get("model")
.is_none());
}
#[tokio::test]
async fn update_current_step_uses_a_one_based_turn_message_index() {
let dispatcher = ToolDispatcher::new(CursorToolRuntime::default());
let call = ToolCall {
index: 0,
call_id: "update-call".into(),
model_call_id: "model-call".into(),
name: "UpdateCurrentStep".into(),
arguments_text: r#"{"current_step":"testing"}"#.into(),
arguments: serde_json::json!({"current_step":"testing"}),
};
let completed = HashSet::new();
let started = HashSet::new();
let messages = vec![
CanonicalMessage::text("old-runtime", Role::User, Origin::Runtime, "old turn"),
CanonicalMessage {
message_id: "old-assistant".into(),
role: Role::Assistant,
origin: Origin::Assistant,
content: MessageContent::Assistant {
text: "old response".into(),
thinking: String::new(),
tool_round_id: None,
replay_state: None,
tool_calls: Vec::new(),
},
runtime_event_id: None,
},
CanonicalMessage::text(
"current-runtime",
Role::User,
Origin::Runtime,
"current turn",
),
];
let dispatched = dispatcher
.start_batch(
&[call],
ToolBatchState {
completed: &completed,
started: &started,
response_text: "",
response_thinking: "",
},
&messages,
&BTreeMap::new(),
&context(),
)
.await
.unwrap();
let completion = dispatched[0].completion.as_ref().unwrap();
let Some(pb::tool_call::Tool::CommunicateUpdateToolCall(tool)) =
completion.tool_call().tool.as_ref()
else {
panic!("expected CommunicateUpdateToolCall")
};
let pb::communicate_update_result::Result::Success(success) =
tool.result.as_ref().unwrap().result.as_ref().unwrap()
else {
panic!("expected communicate update success")
};
assert_eq!(success.message_index, 1);
assert_eq!(success.current_step, "testing");
}
fn context() -> ExecContext {
ExecContext {
conversation_id: "child".into(),
root_conversation_id: "root".into(),
default_subagent_model: "parent-model".into(),
subagent_model: None,
allow_subagents: true,
subagents_disabled: false,
terminals_folder: "/tmp/terminals".into(),
admin_command_denylist: Vec::new(),
mcp_routes: HashMap::new(),
}
}
fn task_call(arguments: serde_json::Value) -> ToolCall {
ToolCall {
index: 0,
call_id: "task-call".into(),
model_call_id: "model-call".into(),
name: "Task".into(),
arguments_text: arguments.to_string(),
arguments,
}
}
@@ -0,0 +1,9 @@
use bytes::Bytes;
use cursor_server::{cursor::connect, Result};
use prost::Message;
pub fn decode_single<M: Message + Default>(frame: &Bytes) -> Result<M> {
let frames = connect::decode_frames(frame)?;
assert_eq!(frames.len(), 1);
Ok(M::decode(frames[0].1.clone())?)
}
@@ -0,0 +1,91 @@
#![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()),
}
}
}
+17
View File
@@ -0,0 +1,17 @@
#![allow(dead_code)]
use cursor_server::{
model::{CanonicalMessage, Origin, Role},
store::Store,
};
pub async fn temp_store() -> (tempfile::TempDir, Store) {
let directory = tempfile::tempdir().unwrap();
let url = format!("sqlite://{}", directory.path().join("test.db").display());
let store = Store::connect(&url).await.unwrap();
(directory, store)
}
pub fn user(id: &str, text: &str) -> CanonicalMessage {
CanonicalMessage::text(id, Role::User, Origin::User, text)
}
+372
View File
@@ -0,0 +1,372 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{
collections::{HashMap, HashSet},
sync::Arc,
};
use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::{connect, proto::agent::v1 as pb},
cursor::{CursorCommand, CursorSessionRegistry},
model::{ModelConfigInput, ModelType, ProjectedContent, Role, Usage, OPENAI_CHAT_ENDPOINT},
provider::{FinishReason, ModelEvent},
};
use prost::Message;
#[tokio::test]
async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
let (_directory, store) = fixtures::temp_store().await;
let configured_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(vec![
ModelEvent::Start {
model_call_id: "ignored".into(),
},
ModelEvent::ThinkingStart,
ModelEvent::ThinkingDelta("reason".into()),
ModelEvent::ThinkingEnd,
ModelEvent::TextStart,
ModelEvent::TextDelta("hello".into()),
ModelEvent::TextEnd,
ModelEvent::Usage(Usage {
input_tokens: Some(20_000),
output_tokens: Some(8),
total_tokens: Some(20_008),
cache_read_tokens: Some(80),
cache_write_tokens: None,
reasoning_tokens: Some(3),
}),
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run(
"conversation",
"hello",
&configured_model.model_hash,
)),
})
.await
.unwrap();
let mut append_seqno = 1;
let mut text = String::new();
let mut thinking = String::new();
let mut thinking_duration_ms = None;
let mut saw_turn_ended = false;
let mut token_deltas = Vec::new();
let mut checkpoints = 0;
let mut after_turn_checkpoints = Vec::new();
let mut blobs = HashMap::<Vec<u8>, Vec<u8>>::new();
let mut set_blob_ids = HashSet::new();
let mut withheld_final_ack = 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 {
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(set)) = &kv.message {
assert!(
set_blob_ids.insert(set.blob_id.clone()),
"an acknowledged content-addressed Blob must not be SET twice"
);
blobs.insert(set.blob_id.clone(), set.blob_data.clone());
}
if text == "hello" && !withheld_final_ack {
withheld_final_ack = true;
assert!(
tokio::time::timeout(std::time::Duration::from_millis(100), output.recv())
.await
.is_err(),
"TurnEnded/checkpoint must wait for the final Blob ACK"
);
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
match update.message {
Some(pb::interaction_update::Message::TextDelta(delta)) => {
text.push_str(&delta.text)
}
Some(pb::interaction_update::Message::ThinkingDelta(delta)) => {
assert_eq!(
delta.thinking_style,
Some(pb::ThinkingStyle::Default as i32)
);
thinking.push_str(&delta.text);
}
Some(pb::interaction_update::Message::ThinkingCompleted(completed)) => {
thinking_duration_ms = Some(completed.thinking_duration_ms)
}
Some(pb::interaction_update::Message::TurnEnded(usage)) => {
assert_eq!(usage.input_tokens, Some(20_000));
assert_eq!(usage.output_tokens, Some(8));
assert_eq!(usage.cache_read_tokens, Some(80));
assert_eq!(usage.reasoning_tokens, Some(3));
saw_turn_ended = true;
}
Some(pb::interaction_update::Message::TokenDelta(delta)) => {
token_deltas.push(delta.tokens)
}
_ => {}
}
}
Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(state)) => {
checkpoints += 1;
if saw_turn_ended {
after_turn_checkpoints.push(state);
}
}
_ => {}
}
}
assert_eq!(text, "hello");
assert_eq!(thinking, "reason");
assert!(thinking_duration_ms.is_some_and(|duration| duration >= 1));
assert!(saw_turn_ended);
assert_eq!(token_deltas, vec![8]);
assert!(withheld_final_ack);
assert!(
checkpoints >= 2,
"final checkpoint is intentionally repeated"
);
assert_eq!(after_turn_checkpoints.len(), 3);
assert_eq!(after_turn_checkpoints[0].pending_tool_calls.len(), 1);
assert!(after_turn_checkpoints[1].pending_tool_calls.is_empty());
assert_eq!(after_turn_checkpoints[1], after_turn_checkpoints[2]);
let token_details = after_turn_checkpoints[1].token_details.as_ref().unwrap();
assert_eq!(token_details.used_tokens, 20_008);
assert_eq!(token_details.max_tokens, 256_000);
let breakdown = token_details.breakdown.as_ref().unwrap();
assert_eq!(breakdown.total_used_tokens, 20_008);
assert_eq!(breakdown.max_tokens, 256_000);
assert_eq!(breakdown.categories.len(), 9);
assert_eq!(
breakdown
.categories
.iter()
.map(|category| category.estimated_tokens)
.sum::<u32>(),
20_009
);
assert_eq!(
after_turn_checkpoints[0].turns, after_turn_checkpoints[1].turns,
"staged and settled checkpoints reuse the same frozen Turn"
);
assert_eq!(
after_turn_checkpoints[0].root_prompt_messages_json.len() + 1,
after_turn_checkpoints[1].root_prompt_messages_json.len()
);
let pending: serde_json::Value =
serde_json::from_str(&after_turn_checkpoints[0].pending_tool_calls[0]).unwrap();
assert_eq!(pending["role"], "assistant");
assert!(pending["content"]
.as_array()
.unwrap()
.iter()
.any(|part| part["type"] == "text" && part["text"] == "hello"));
let final_root = after_turn_checkpoints[1]
.root_prompt_messages_json
.last()
.unwrap();
let stable: serde_json::Value = serde_json::from_slice(blobs.get(final_root).unwrap()).unwrap();
assert_eq!(stable["role"], "assistant");
assert!(
stable.get("origin").is_none(),
"wire root is not CanonicalMessage JSON"
);
let turn = pb::ConversationTurnStructure::decode(
blobs
.get(after_turn_checkpoints[0].turns.last().unwrap())
.unwrap()
.as_slice(),
)
.unwrap();
let pb::conversation_turn_structure::Turn::AgentConversationTurn(turn) = turn.turn.unwrap()
else {
panic!("expected agent turn")
};
assert_eq!(turn.steps.len(), 2, "thinking and text are frozen once");
let requests = provider.requests();
assert_eq!(requests.len(), 1);
assert!(requests[0]
.prompt
.instructions
.contains("powered by Test Model"));
let projected = &requests[0].history;
assert_eq!(projected[0].role, Role::User);
assert!(projected[0].message_id.starts_with("request-context:"));
let ProjectedContent::Parts(context) = &projected[0].content else {
panic!("request context must be text")
};
assert!(matches!(
context.as_slice(),
[cursor_server::model::ContentPart::Text { text }]
if text.contains("<user_info>")
));
let ProjectedContent::Parts(runtime) = &projected[1].content else {
panic!("runtime user message must be text")
};
assert!(matches!(
runtime.as_slice(),
[cursor_server::model::ContentPart::Text { text }]
if text.contains("<user_query>\nhello\n</user_query>")
&& !text.contains("<user_info>")
));
assert_eq!(
projected.len(),
2,
"the raw UserMessage is not projected twice"
);
let messages = store
.load_current_messages(&cursor_server::model::ConversationId::new("conversation"))
.await
.unwrap();
assert!(messages[0].message_id.starts_with("request-context:"));
assert_eq!(messages[0].role, Role::User);
assert!(messages[1]
.message_id
.starts_with("runtime:cursor:user:user:"));
assert_eq!(messages[1].role, Role::User);
assert_eq!(
messages.len(),
3,
"request context plus runtime user and final assistant"
);
let stored_runs: Vec<String> = sqlx::query_scalar("SELECT run_id FROM runs ORDER BY run_id")
.fetch_all(store.pool())
.await
.unwrap();
assert_eq!(stored_runs.len(), 1);
let execution_suffix = stored_runs[0]
.strip_prefix("request:")
.expect("the local Run keeps the Cursor request id as a readable prefix");
assert_eq!(execution_suffix.len(), 8);
assert!(execution_suffix
.bytes()
.all(|byte| byte.is_ascii_hexdigit()));
}
fn client_run(conversation_id: &str, text: &str, model_id: &str) -> pb::AgentClientMessage {
let user = pb::UserMessage {
text: text.into(),
message_id: "user".into(),
mode: pb::AgentMode::Agent as i32,
..Default::default()
};
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::RunRequest(
pb::AgentRunRequest {
requested_model: Some(pb::RequestedModel {
model_id: model_id.into(),
parameters: vec![pb::requested_model::ModelParameterValue {
id: "context".into(),
value: "256k".into(),
}],
..Default::default()
}),
action: Some(pb::ConversationAction {
action: Some(pb::conversation_action::Action::UserMessageAction(
pb::UserMessageAction {
user_message: Some(user),
request_context: Some(pb::RequestContext {
env: Some(pb::RequestContextEnv {
os_version: "darwin".into(),
workspace_paths: vec!["/workspace".into()],
shell: "zsh".into(),
terminals_folder: "/terminals".into(),
agent_transcripts_folder: "/transcripts".into(),
..Default::default()
}),
git_repos: vec![pb::GitRepoInfo {
path: "/workspace".into(),
status: "M src/main.rs".into(),
..Default::default()
}],
..Default::default()
}),
..Default::default()
},
)),
..Default::default()
}),
conversation_id: Some(conversation_id.into()),
conversation_state: None,
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 },
)),
},
)),
}
}
+981
View File
@@ -0,0 +1,981 @@
#[path = "support/fake_provider.rs"]
mod fake_provider;
#[path = "support/fixtures.rs"]
mod fixtures;
use std::{
collections::{BTreeMap, HashSet},
sync::Arc,
};
use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::{
connect,
proto::agent::v1 as pb,
tools::{
codec,
runtime::{CursorToolRuntime, ExecContext},
ClientToolEvent, ToolBatchState, ToolDispatcher,
},
},
cursor::{CursorCommand, CursorSessionRegistry},
model::{MessageContent, ToolCall},
provider::{FinishReason, ModelEvent},
};
use prost::Message;
use serde_json::json;
fn call(id: &str, name: &str) -> ToolCall {
ToolCall {
index: 0,
call_id: id.into(),
model_call_id: "model:0".into(),
name: name.into(),
arguments_text: "{}".into(),
arguments: json!({}),
}
}
fn exec_context() -> ExecContext {
ExecContext {
conversation_id: "conversation".into(),
root_conversation_id: "conversation".into(),
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::interaction::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::interaction::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 = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry.get_or_create("tool-request").await.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run()),
})
.await
.unwrap();
let mut seqno = 1;
let mut saw_exec = false;
let mut saw_typed_completion = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.unwrap();
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&payload).unwrap(),
json!({})
);
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
saw_exec = true;
let exec_id = exec.id;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id: exec_id,
exec_id: String::new(),
message: Some(pb::exec_client_message::Message::ReadResult(
pb::ReadResult {
result: Some(pb::read_result::Result::Success(
pb::ReadSuccess {
path: "/tmp/a".into(),
total_lines: 1,
file_size: 1,
output: Some(
pb::read_success::Output::Content(
"x".into(),
),
),
..Default::default()
},
)),
},
)),
..Default::default()
},
)),
}),
})
.await
.unwrap();
seqno += 1;
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(
pb::agent_client_message::Message::ExecClientControlMessage(
pb::ExecClientControlMessage {
message: Some(
pb::exec_client_control_message::Message::StreamClose(
pb::ExecClientStreamClose { id: exec_id },
),
),
},
),
),
}),
})
.await
.unwrap();
seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
if let Some(pb::interaction_update::Message::ToolCallCompleted(completed)) =
update.message
{
let tool_call = completed.tool_call.expect("completed ToolCall");
assert!(tool_call.started_at_ms.unwrap_or_default() > 1);
assert!(tool_call.completed_at_ms.unwrap_or_default() > 1);
assert!(tool_call.completed_at_ms >= tool_call.started_at_ms);
let Some(pb::tool_call::Tool::ReadToolCall(read)) = tool_call.tool else {
panic!("expected completed ReadToolCall")
};
let result = read.result.expect("typed ReadToolResult");
assert!(matches!(
result.result,
Some(pb::read_tool_result::Result::Success(_))
));
saw_typed_completion = true;
}
}
_ => {}
}
}
assert!(saw_exec);
assert!(saw_typed_completion);
assert_eq!(provider.requests().len(), 2);
let database = sqlx::SqlitePool::connect(&format!(
"sqlite://{}",
directory.path().join("test.db").display()
))
.await
.unwrap();
let provider_call_index: i64 =
sqlx::query_scalar("SELECT provider_call_index FROM runs WHERE 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 },
)),
},
)),
}
}
+118
View File
@@ -0,0 +1,118 @@
#[path = "support/fixtures.rs"]
mod fixtures;
use cursor_server::{
model::{
ConversationId, MessageContent, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId,
RunKind, ToolCall, ToolResult, ToolRoundAssistant, ToolRoundId,
},
store::ToolRoundStatus,
};
#[tokio::test]
async fn results_commit_adjacent_pairs_in_arrival_order() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let run = PreparedRun {
run_id: RunId::new("run"),
cursor_request_id: None,
conversation_id: conversation_id.clone(),
kind: RunKind::Root,
model: ModelSpec::new("model"),
prompt: PromptSpec {
instructions: String::new(),
tools: Vec::new(),
},
initial_messages: Vec::new(),
action: RunAction::Resume {
pending_tool_round: None,
},
base_revision_id: root,
};
store.claim_run(&run).await.unwrap();
let round_id = ToolRoundId::new("round");
let calls = [call(0, "A"), call(1, "B"), call(2, "C")];
store
.create_tool_round(
&round_id,
&run.run_id,
root,
&ToolRoundAssistant {
text: "answer prefix".into(),
thinking: "reasoning".into(),
model_call_id: "model-call".into(),
replay_state: None,
},
&calls,
None,
)
.await
.unwrap();
let b = store
.commit_tool_result(&conversation_id, &run.run_id, &round_id, &result("B"))
.await
.unwrap();
assert_eq!(b.completion_seq, 0);
assert!(!b.settled);
let a = store
.commit_tool_result(&conversation_id, &run.run_id, &round_id, &result("A"))
.await
.unwrap();
assert_eq!(a.completion_seq, 1);
let c = store
.commit_tool_result(&conversation_id, &run.run_id, &round_id, &result("C"))
.await
.unwrap();
assert!(c.settled);
let messages = store.load_revision_messages(c.revision_id).await.unwrap();
assert_eq!(messages.len(), 6);
let ids = messages
.chunks_exact(2)
.map(|pair| match (&pair[0].content, &pair[1].content) {
(MessageContent::Assistant { tool_calls, .. }, MessageContent::ToolResult(result)) => {
assert_eq!(tool_calls[0].call_id, result.call_id);
result.call_id.clone()
}
_ => panic!("tool result must be adjacent to its assistant call"),
})
.collect::<Vec<_>>();
assert_eq!(ids, ["B", "A", "C"]);
let MessageContent::Assistant { text, thinking, .. } = &messages[0].content else {
unreachable!()
};
assert_eq!(text, "answer prefix");
assert_eq!(thinking, "reasoning");
for message in [&messages[2], &messages[4]] {
let MessageContent::Assistant { text, thinking, .. } = &message.content else {
unreachable!()
};
assert!(text.is_empty());
assert!(thinking.is_empty());
}
let snapshot = store.tool_round(&round_id).await.unwrap().unwrap();
assert_eq!(snapshot.status, ToolRoundStatus::Settled);
assert_eq!(snapshot.completed_call_ids.len(), 3);
}
fn call(index: usize, call_id: &str) -> ToolCall {
ToolCall {
index,
call_id: call_id.into(),
model_call_id: "model-call".into(),
name: "Read".into(),
arguments_text: r#"{"path":"/tmp/a"}"#.into(),
arguments: serde_json::json!({"path":"/tmp/a"}),
}
}
fn result(call_id: &str) -> ToolResult {
ToolResult {
call_id: call_id.into(),
content: format!("result-{call_id}"),
is_error: false,
image: None,
}
}
+181
View File
@@ -0,0 +1,181 @@
use axum::{http::StatusCode, response::IntoResponse, routing::get, Router};
use cursor_server::search::{HtmlEngine, JsonEngine, WebSearch};
use tokio::net::TcpListener;
const RESULT_SELECTOR: &str = ".result";
const TITLE_SELECTOR: &str = ".title";
const LINK_SELECTOR: &str = "a.title";
const SNIPPET_SELECTOR: &str = ".snippet";
#[test]
fn built_in_catalog_covers_the_reference_search_engines() {
let search = WebSearch::built_in();
let ids = search.engine_ids();
for expected in [
"google",
"bing",
"brave",
"duckduckgo",
"startpage",
"yahoo",
"mojeek",
"qwant",
"ecosia",
"yandex",
"baidu",
"sogou",
"so360",
"naver",
"seznam",
"wikipedia",
"github",
"stackoverflow",
"crates_io",
"npm",
"pypi",
"arxiv",
"crossref",
] {
assert!(ids.contains(&expected), "missing engine {expected}");
}
}
#[tokio::test]
async fn federated_search_deduplicates_and_rrf_ranks_results() {
let base = spawn_search_fixture().await;
let search = WebSearch::with_engines(vec![
engine("first", format!("{base}/first?q={{query}}")),
engine("second", format!("{base}/second?q={{query}}")),
]);
let results = search.search("rust agent").await.unwrap();
assert_eq!(results.len(), 3);
assert_eq!(results[0].url, "https://example.com/shared");
assert_eq!(results[0].engines, vec!["first", "second"]);
assert!(results.iter().any(|result| result.title == "Alpha"));
assert!(results.iter().any(|result| result.title == "Beta"));
}
#[tokio::test]
async fn one_failed_engine_does_not_discard_other_engine_results() {
let base = spawn_search_fixture().await;
let search = WebSearch::with_engines(vec![
engine("failed", format!("{base}/failed?q={{query}}")),
engine("first", format!("{base}/first?q={{query}}")),
]);
let results = search.search("rust").await.unwrap();
assert_eq!(results.len(), 2);
}
#[tokio::test]
async fn all_failed_engines_return_a_search_error() {
let base = spawn_search_fixture().await;
let search =
WebSearch::with_engines(vec![engine("failed", format!("{base}/failed?q={{query}}"))]);
let error = search.search("rust").await.unwrap_err();
assert!(error.to_string().contains("failed"));
}
#[tokio::test]
async fn json_engines_use_declared_result_fields() {
let base = spawn_search_fixture().await;
let search = WebSearch::with_engines(vec![JsonEngine::new(
"json",
format!("{base}/json?q={{query}}"),
"/items",
"/name",
"/url",
"/description",
None,
)]);
let results = search.search("rust").await.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].title, "Structured result");
assert_eq!(results[0].url, "https://example.com/structured");
assert_eq!(results[0].chunk, "Parsed from JSON");
}
#[tokio::test]
#[ignore = "live public search smoke test"]
async fn built_in_search_returns_live_results() {
let _ = tracing_subscriber::fmt()
.with_env_filter("cursor_server::search=debug")
.try_init();
let results = WebSearch::built_in()
.search("Rust programming language")
.await
.unwrap();
assert!(!results.is_empty());
for result in results {
println!("{}\t{}\t{:?}", result.title, result.url, result.engines);
}
}
fn engine(id: &'static str, url: String) -> HtmlEngine {
HtmlEngine::new(
id,
url,
RESULT_SELECTOR,
TITLE_SELECTOR,
LINK_SELECTOR,
SNIPPET_SELECTOR,
)
}
async fn spawn_search_fixture() -> String {
async fn first() -> impl IntoResponse {
r#"
<article class="result">
<a class="title" href="https://example.com/shared?utm_source=one">Shared</a>
<p class="snippet">Shared from first.</p>
</article>
<article class="result">
<a class="title" href="https://alpha.example/path">Alpha</a>
<p class="snippet">Alpha result.</p>
</article>
"#
}
async fn second() -> impl IntoResponse {
r#"
<article class="result">
<a class="title" href="https://example.com/shared#section">Shared result</a>
<p class="snippet">Shared from second.</p>
</article>
<article class="result">
<a class="title" href="https://beta.example/">Beta</a>
<p class="snippet">Beta result.</p>
</article>
"#
}
let app = Router::new()
.route("/first", get(first))
.route("/second", get(second))
.route(
"/failed",
get(|| async { (StatusCode::TOO_MANY_REQUESTS, "limited") }),
)
.route(
"/json",
get(|| async {
axum::Json(serde_json::json!({
"items": [{
"name": "Structured result",
"url": "https://example.com/structured",
"description": "Parsed from JSON"
}]
}))
}),
);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
format!("http://{address}")
}