mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-08 07:21:13 +08:00
fix: shell
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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(),
|
||||
|
||||
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user