fix: shell

This commit is contained in:
leokun
2026-08-26 20:20:38 +08:00
parent df053c3720
commit 5450fc76e2
25 changed files with 1367 additions and 101 deletions
+164 -10
View File
@@ -76,7 +76,7 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() {
.unwrap();
assert!(messages.iter().any(|message| {
message.runtime_event_id.as_deref()
== Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id")
== Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id:task-call")
&& matches!(&message.content, MessageContent::Parts { parts } if !parts.is_empty())
}));
@@ -131,12 +131,78 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() {
assert_eq!(
runtime_ids,
[
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id",
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id-2"
"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;
@@ -172,12 +238,7 @@ async fn retrying_one_background_completion_reuses_its_runtime_message() {
let second = registry.get_or_create("completion-retry-2").await.unwrap();
drive_completion(
&second,
completion_run_with_detail(
"retry-child",
"completion-retry-run-2",
checkpoint,
"updated retry payload",
),
completion_run("retry-child", "completion-retry-run-2", checkpoint),
)
.await;
@@ -192,7 +253,9 @@ async fn retrying_one_background_completion_reuses_its_runtime_message() {
.iter()
.filter(|message| {
message.runtime_event_id.as_deref()
== Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:retry-child")
== Some(
"background-completed:BACKGROUND_TASK_KIND_SUBAGENT:retry-child:task-call",
)
})
.count(),
1
@@ -397,6 +460,97 @@ async fn drive_completion(
)
}
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,
+81
View File
@@ -14,6 +14,7 @@ use cursor_server::{
provider::{FinishReason, ModelEvent},
run::{RunEngine, RunOutcome},
};
use tokio::{sync::oneshot, time::Duration};
use tokio_util::sync::CancellationToken;
#[tokio::test]
@@ -54,6 +55,86 @@ async fn a_client_without_checkpoint_protocol_runs_the_same_text_loop() {
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::client::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;
+24 -1
View File
@@ -160,7 +160,7 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
store.clone(),
Arc::new(provider),
PromptCompiler::new(assets),
Default::default(),
@@ -249,6 +249,29 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
.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]
+217
View File
@@ -28,6 +28,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
conversation.clone(),
cursor_server::model::RunId::new("first"),
first.clone(),
cursor_server::client::session(1).1.commands,
)
.await;
registry
@@ -35,6 +36,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
conversation.clone(),
cursor_server::model::RunId::new("second"),
second.clone(),
cursor_server::client::session(1).1.commands,
)
.await;
@@ -426,6 +428,187 @@ async fn injected_user_context_restarts_only_the_active_model_cycle() {
);
}
#[tokio::test]
async fn cancel_subagent_action_aborts_the_target_task_and_keeps_the_parent_running() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "task-cycle".into(),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: "task-call".into(),
name: "Task".into(),
},
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: serde_json::json!({
"description": "Inspect protocol",
"prompt": "Inspect the protocol",
"subagent_type": "generalPurpose",
"run_in_background": false
})
.to_string(),
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "continued".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("continued after subagent cancellation".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,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry
.get_or_create("cancel-subagent-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run_for(
"cancel-subagent-request",
"cancel-subagent-conversation",
)),
})
.await
.unwrap();
let mut append_seqno = 1;
let exec_id = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Task exec");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
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)) => {
let Some(pb::exec_server_message::Message::SubagentArgs(args)) = exec.message
else {
continue;
};
assert_eq!(args.tool_call_id, "task-call");
break exec.id;
}
_ => {}
}
};
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_cancel_subagent("task-call")),
})
.await
.unwrap();
append_seqno += 1;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Task abort");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) => {
let Some(pb::exec_server_control_message::Message::Abort(abort)) = control.message
else {
continue;
};
assert_eq!(abort.id, exec_id);
break;
}
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;
}
_ => {}
}
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(subagent_aborted(exec_id)),
})
.await
.unwrap();
append_seqno += 1;
let mut saw_continued = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before successful EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
break;
}
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::InteractionUpdate(update)) => {
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
saw_continued |= delta.text.contains("continued after subagent cancellation");
}
}
_ => {}
}
}
assert!(saw_continued);
assert!(!handle.cancellation().is_cancelled());
assert_eq!(provider.requests().len(), 2);
}
fn client_run() -> pb::AgentClientMessage {
client_run_for("cancel-request", "cancel-conversation")
}
@@ -541,3 +724,37 @@ fn runtime_injection() -> pb::AgentClientMessage {
)),
}
}
fn runtime_cancel_subagent(tool_call_id: &str) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ConversationAction(
pb::ConversationAction {
action: Some(pb::conversation_action::Action::CancelSubagentAction(
pb::CancelSubagentAction {
subagent_id: tool_call_id.into(),
},
)),
..Default::default()
},
)),
}
}
fn subagent_aborted(id: u32) -> 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::Error(pb::SubagentError {
agent_id: None,
error: "Subagent was aborted by the user".into(),
})),
},
)),
..Default::default()
},
)),
}
}
+92
View File
@@ -218,3 +218,95 @@ async fn editing_a_logical_input_discards_its_active_suffix() {
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
]
);
}
+147 -4
View File
@@ -18,6 +18,150 @@ use cursor_server::{
};
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;
@@ -115,10 +259,9 @@ async fn unchanged_request_context_is_not_repeated_and_preserves_the_provider_pr
let [ContentPart::Text { text: context_text }] = context_parts.as_slice() else {
panic!("request context message must contain one text part")
};
assert_eq!(
request.history[1].message_id,
"runtime:cursor:user:wire-user"
);
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 {
+9 -2
View File
@@ -89,7 +89,11 @@ async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_che
let user = requests[0]
.history
.iter()
.find(|message| message.message_id == "runtime:cursor:user:image-user")
.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")
@@ -128,7 +132,10 @@ async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_che
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"] == "runtime:cursor:user:image-user" {
if value["id"]
.as_str()
.is_some_and(|id| id.starts_with("runtime:cursor:user:image-user:"))
{
user_root = Some(value);
break;
}
+23 -1
View File
@@ -10,11 +10,15 @@ use cursor_server::{
provider::{ModelEvent, Provider, ProviderStream},
Error,
};
use futures_util::stream;
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,
}
@@ -43,6 +47,17 @@ impl FakeProvider {
.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()
}
@@ -63,6 +78,13 @@ impl Provider for FakeProvider {
.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()),
}
}
+3 -1
View File
@@ -283,7 +283,9 @@ async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
.unwrap();
assert!(messages[0].message_id.starts_with("request-context:"));
assert_eq!(messages[0].role, Role::User);
assert_eq!(messages[1].message_id, "runtime:cursor:user:user");
assert!(messages[1]
.message_id
.starts_with("runtime:cursor:user:user:"));
assert_eq!(messages[1].role, Role::User);
assert_eq!(
messages.len(),
+30
View File
@@ -563,6 +563,36 @@ async fn empty_exec_client_message_is_not_a_terminal_result() {
);
}
#[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();