From 5450fc76e25cf4d6393f715a28e9ba9996915f04 Mon Sep 17 00:00:00 2001 From: leokun Date: Wed, 26 Aug 2026 20:20:38 +0800 Subject: [PATCH] fix: shell --- server/src/client/command.rs | 11 +- server/src/cursor/actor.rs | 64 +++++-- server/src/cursor/request/background.rs | 35 +++- server/src/cursor/request/prepare.rs | 63 +++++-- server/src/cursor/request/runtime.rs | 61 +++++- server/src/cursor/run_sse.rs | 95 +++++++++- server/src/cursor/session.rs | 17 ++ server/src/cursor/tools/codec/mod.rs | 2 +- server/src/cursor/tools/codec/response.rs | 47 +++++ server/src/cursor/tools/runtime.rs | 12 ++ server/src/run/actor.rs | 9 +- server/src/run/engine.rs | 94 ++++++++-- server/src/run/registry.rs | 39 +++- server/src/run/tool_round.rs | 78 ++++---- server/src/store/revisions.rs | 32 ++++ server/tests/background_completion.rs | 174 ++++++++++++++++- server/tests/client_contract.rs | 81 ++++++++ server/tests/error_lifecycle.rs | 25 ++- server/tests/interrupt.rs | 217 ++++++++++++++++++++++ server/tests/revision_branch.rs | 92 +++++++++ server/tests/runtime_modes.rs | 151 ++++++++++++++- server/tests/selected_images.rs | 11 +- server/tests/support/fake_provider.rs | 24 ++- server/tests/text_turn.rs | 4 +- server/tests/tool_loop.rs | 30 +++ 25 files changed, 1367 insertions(+), 101 deletions(-) diff --git a/server/src/client/command.rs b/server/src/client/command.rs index 7b14610..d29067c 100644 --- a/server/src/client/command.rs +++ b/server/src/client/command.rs @@ -1,10 +1,19 @@ +use tokio::sync::oneshot; + use crate::model::{CanonicalMessage, RuntimeEvent, ToolResult}; -#[derive(Clone, Debug, PartialEq)] +#[derive(Debug)] +pub struct MessageInsertion { + pub messages: Vec, + pub delivered: oneshot::Sender<()>, +} + +#[derive(Debug)] pub enum ClientCommand { ToolResult(ToolResult), RuntimeMessage(CanonicalMessage), RuntimeEvent(RuntimeEvent), + InsertMessages(MessageInsertion), ClientClosed { error: String }, Cancel, } diff --git a/server/src/cursor/actor.rs b/server/src/cursor/actor.rs index 6ca1cdd..de789fb 100644 --- a/server/src/cursor/actor.rs +++ b/server/src/cursor/actor.rs @@ -151,15 +151,39 @@ impl CursorActor { context.dynamic_tools.keys().cloned().collect(), context.turn_user.clone(), ); + if context.background_completion + && dependencies + .run_registry + .insert_messages( + &prepared.conversation_id, + prepared.initial_messages.clone(), + ) + .await + { + crate::cursor::lifecycle::finish_success( + &handle, + ); + let _ = handle + .command(CursorCommand::Finished) + .await; + return; + } let cancellation = handle.cancellation(); let (port, core) = crate::client::session(256); + let core_commands = core.commands.clone(); let actor = RunActor::new( dependencies.store.clone(), dependencies.provider, dependencies.run_registry, ); - let core_run = - actor.spawn(prepared, port, cancellation).await; + let core_run = actor + .spawn( + prepared, + port, + core_commands, + cancellation, + ) + .await; let session = CursorSession::new( handle.clone(), dependencies.store, @@ -181,7 +205,6 @@ impl CursorActor { %error, "Cursor session failed" ); - handle.cancel(); let _ = crate::cursor::lifecycle::fail( &handle, &error, ); @@ -235,12 +258,14 @@ impl CursorActor { { continue; } - if tool_runtime.take_exec(close.id).await.is_some() + match codec::stream_closed(close.id, &tool_runtime) + .await { - results_tx.send_error(crate::Error::Protocol(format!( - "Exec stream closed before result for id: {}", - close.id - ))); + Ok(Some(completion)) => { + results_tx.send(completion) + } + Ok(None) => {} + Err(error) => results_tx.send_error(error), } } Some(Message::Throw(throw)) => { @@ -311,13 +336,12 @@ impl CursorActor { // // The remaining unimplemented Action variants are // ShellCommandAction, StartPlanAction, - // AsyncAskQuestionCompletionAction, CancelSubagentAction, - // BackgroundShellAction, BackgroundSubagentAction, + // AsyncAskQuestionCompletionAction, BackgroundShellAction, + // BackgroundSubagentAction, // SubscriptionNotificationAction and GoalContinuationAction. - // CancelSubagentAction must not start an LLM; variants whose wire - // behavior is not captured yet need evidence before assigning - // semantics. Every unsupported runtime Action must return an explicit - // Protocol Error rather than falling through silently. + // Variants whose wire behavior is not captured yet need evidence + // before assigning semantics. Every unsupported runtime Action must + // return an explicit Protocol Error rather than falling through silently. Some( pb::agent_client_message::Message::ConversationAction( action, @@ -341,6 +365,18 @@ impl CursorActor { )); } } + Some( + pb::conversation_action::Action::CancelSubagentAction( + action, + ), + ) => { + if let Some(id) = tool_runtime + .running_task_exec_id(&action.subagent_id) + .await + { + let _ = handle.emit(&codec::abort(id)); + } + } Some(action) => { results_tx.send_error(crate::Error::Protocol(format!( "unsupported runtime ConversationAction: {}", diff --git a/server/src/cursor/request/background.rs b/server/src/cursor/request/background.rs index 08fd638..12a0b65 100644 --- a/server/src/cursor/request/background.rs +++ b/server/src/cursor/request/background.rs @@ -87,8 +87,15 @@ pub(super) fn project( } pb::BackgroundTaskKind::Unspecified => unreachable!(), }; - let identity = agent_id.unwrap_or(&completion.task_id); - let identity = format!("{}:{identity}", kind.as_str_name()); + let tool_call_id = completion + .tool_call_id + .as_deref() + .filter(|id| !id.is_empty()) + .ok_or_else(|| { + Error::Protocol("background task completion has no tool_call_id".into()) + })?; + let task_identity = agent_id.unwrap_or(&completion.task_id); + let identity = format!("{}:{task_identity}:{tool_call_id}", kind.as_str_name()); let context = completion_context(completion, kind, agent_id)?; if completions .insert(identity.clone(), (completion, context)) @@ -307,6 +314,30 @@ mod tests { assert_eq!(forward.context, reversed.context); } + #[test] + fn resumed_subagent_completions_use_the_task_call_as_part_of_their_identity() { + let first = project( + &pb::BackgroundTaskCompletionAction { + completions: vec![completion()], + }, + pb::AgentMode::Multitask as i32, + ) + .unwrap(); + let mut resumed = completion(); + resumed.tool_call_id = Some("task-call-2".into()); + let second = project( + &pb::BackgroundTaskCompletionAction { + completions: vec![resumed], + }, + pb::AgentMode::Multitask as i32, + ) + .unwrap(); + + assert_ne!(first.turn_user.message_id, second.turn_user.message_id); + assert!(first.turn_user.message_id.ends_with(":task-call")); + assert!(second.turn_user.message_id.ends_with(":task-call-2")); + } + #[test] fn completion_requires_the_captured_subagent_identity_and_terminal_reason() { let mut value = completion(); diff --git a/server/src/cursor/request/prepare.rs b/server/src/cursor/request/prepare.rs index dd7883b..be1c10a 100644 --- a/server/src/cursor/request/prepare.rs +++ b/server/src/cursor/request/prepare.rs @@ -41,6 +41,7 @@ pub struct CursorRunContext { pub dynamic_tools: BTreeMap, pub checkpoint_prompt: PromptSpec, pub compacting: bool, + pub background_completion: bool, } pub(crate) struct PrepareDependencies<'a> { @@ -123,7 +124,7 @@ pub(crate) async fn prepare( mode: mode_number, mut turn_user, action_context, - event_id, + mut event_id, input_id, starts_turn, compacting, @@ -172,14 +173,43 @@ pub(crate) async fn prepare( } Some(_) | None => store.ensure_conversation(&conversation_id).await?, }; - let base_revision_id = match input_id { + let base_revision_id = match input_id.as_deref() { Some(input_id) => { store - .anchor_input(&conversation_id, &input_id, proposed_base_revision_id) + .anchor_input(&conversation_id, input_id, proposed_base_revision_id) .await? } None => proposed_base_revision_id, }; + let mut projected_user_context = if input_id.is_some() && !compacting && !background_completion + { + runtime::compile_request_context( + "identity", + &request_context, + base_messages.as_deref().unwrap_or_default(), + )? + } else { + None + }; + if event_id.is_none() { + if let (Some(input_id), Some(user)) = (input_id.as_deref(), turn_user.as_ref()) { + event_id = Some( + runtime::user_event_id( + input_id, + checkpoint_mode, + user, + &request_context, + &action_context, + projected_user_context + .as_ref() + .map(|message| &message.content), + compiler, + blob_sync, + ) + .await?, + ); + } + } let existing_runtime = match event_id.as_deref() { Some(event_id) => { store @@ -193,6 +223,10 @@ pub(crate) async fn prepare( let message_id = format!("request-context:{event_id}"); match store.message(&conversation_id, &message_id).await? { Some(message) => Some(message), + None if input_id.is_some() => projected_user_context.take().map(|mut message| { + message.message_id = message_id; + message + }), None => runtime::compile_request_context( event_id, &request_context, @@ -202,7 +236,7 @@ pub(crate) async fn prepare( } _ => None, }; - let initial_messages = if compacting { + let mut initial_messages = if compacting { Vec::new() } else { match (turn_user.clone(), event_id) { @@ -256,6 +290,10 @@ pub(crate) async fn prepare( } } }; + let (base_revision_id, reused) = store + .match_revision_prefix(&conversation_id, base_revision_id, &initial_messages) + .await?; + initial_messages.drain(..reused); let action = if compacting { RunAction::Compact } else if starts_turn { @@ -310,6 +348,7 @@ pub(crate) async fn prepare( .collect(), checkpoint_prompt, compacting, + background_completion, }, )) } @@ -434,13 +473,13 @@ fn action(request: &pb::AgentRunRequest) -> Result { .filter(|text| !text.is_empty()) .cloned(), ); - let event_id = format!("cursor:user:{}", user.message_id); + let input_id = format!("cursor:user:{}", user.message_id); Ok(ActionProjection { mode, turn_user: Some(user.clone()), action_context: context.join("\n\n"), - event_id: Some(event_id.clone()), - input_id: Some(event_id), + event_id: None, + input_id: Some(input_id), starts_turn: true, compacting: false, background_completion: false, @@ -703,7 +742,7 @@ mod tests { } #[test] - fn queued_messages_reusing_a_request_id_keep_distinct_runtime_identities() { + fn queued_messages_keep_distinct_input_anchors_until_runtime_identity_is_compiled() { let request = |message_id: &str| pb::AgentRunRequest { action: Some(pb::ConversationAction { action: Some(pb::conversation_action::Action::UserMessageAction( @@ -725,9 +764,11 @@ mod tests { let first = action(&request("message-one")).unwrap(); let second = action(&request("message-two")).unwrap(); - assert_eq!(first.event_id.as_deref(), Some("cursor:user:message-one")); - assert_eq!(second.event_id.as_deref(), Some("cursor:user:message-two")); - assert_ne!(first.event_id, second.event_id); + assert_eq!(first.event_id, None); + assert_eq!(second.event_id, None); + assert_eq!(first.input_id.as_deref(), Some("cursor:user:message-one")); + assert_eq!(second.input_id.as_deref(), Some("cursor:user:message-two")); + assert_ne!(first.input_id, second.input_id); } #[test] diff --git a/server/src/cursor/request/runtime.rs b/server/src/cursor/request/runtime.rs index b8ea14a..0d3583c 100644 --- a/server/src/cursor/request/runtime.rs +++ b/server/src/cursor/request/runtime.rs @@ -10,6 +10,7 @@ use crate::{ proto::agent::v1 as pb, }, model::{CanonicalMessage, MessageContent, Origin, Role}, + store::BlobId, Error, Result, }; @@ -80,12 +81,66 @@ pub async fn compile( compiler: &PromptCompiler, blobs: &BlobSynchronizer, ) -> Result { - let time = Time::now( + let timestamp = Time::now( request_context .env .as_ref() .map(|env| env.time_zone.as_str()), - )?; + )? + .timestamp; + compile_with_timestamp( + event_id, + mode, + user, + request_context, + action_context, + timestamp, + compiler, + blobs, + ) + .await +} + +#[allow(clippy::too_many_arguments)] +pub(super) async fn user_event_id( + input_id: &str, + mode: Mode, + user: &pb::UserMessage, + request_context: &pb::RequestContext, + action_context: &str, + projected_request_context: Option<&MessageContent>, + compiler: &PromptCompiler, + blobs: &BlobSynchronizer, +) -> Result { + let runtime = compile_with_timestamp( + "identity".into(), + mode, + user, + request_context, + action_context, + String::new(), + compiler, + blobs, + ) + .await?; + let semantic = serde_json::to_vec(&(projected_request_context, runtime.content))?; + Ok(format!( + "{input_id}:{}", + BlobId::digest(&semantic).to_base64() + )) +} + +#[allow(clippy::too_many_arguments)] +async fn compile_with_timestamp( + event_id: String, + mode: Mode, + user: &pb::UserMessage, + request_context: &pb::RequestContext, + action_context: &str, + timestamp: String, + compiler: &PromptCompiler, + blobs: &BlobSynchronizer, +) -> Result { let mut values = BTreeMap::from([ ("OPEN_FILES", section(open_files(user))), ( @@ -98,7 +153,7 @@ pub async fn compile( ), ), ("ACTION_CONTEXT", section(action_context.to_string())), - ("TIMESTAMP", time.timestamp), + ("TIMESTAMP", timestamp), ("USER_QUERY", user.text.clone()), ("DEBUG_SERVER_ENDPOINT", String::new()), ("DEBUG_LOG_PATH", String::new()), diff --git a/server/src/cursor/run_sse.rs b/server/src/cursor/run_sse.rs index 2abbda8..387ae94 100644 --- a/server/src/cursor/run_sse.rs +++ b/server/src/cursor/run_sse.rs @@ -9,7 +9,11 @@ use tokio_stream::StreamExt; use tokio_util::sync::CancellationToken; use crate::{ - cursor::{connect::END_STREAM_FLAG, observability::CursorTraceRecorder, CursorSessionRegistry}, + cursor::{ + connect::{self, END_STREAM_FLAG}, + observability::CursorTraceRecorder, + CursorSessionRegistry, + }, Result, }; @@ -49,7 +53,7 @@ fn local_body_stream( trace.chunk(&chunk); if terminal { guard.complete(); - trace.finish(None); + trace.finish(end_stream_error(&chunk)); } yield Ok::(chunk); if terminal { @@ -67,6 +71,30 @@ fn is_end_stream_frame(frame: &Bytes) -> bool { .is_some_and(|flags| flags & END_STREAM_FLAG != 0) } +fn end_stream_error(frame: &Bytes) -> Option { + connect::decode_frames(frame) + .ok()? + .into_iter() + .find_map(|(flags, payload)| { + if flags & END_STREAM_FLAG == 0 { + return None; + } + let value = serde_json::from_slice::(&payload).ok()?; + let error = value.get("error")?; + let code = error.get("code").and_then(serde_json::Value::as_str); + let message = error + .get("message") + .and_then(serde_json::Value::as_str) + .filter(|message| !message.is_empty()); + Some(match (code, message) { + (Some(code), Some(message)) => format!("{code}: {message}"), + (Some(code), None) => code.to_string(), + (None, Some(message)) => message.to_string(), + (None, None) => error.to_string(), + }) + }) +} + struct LocalRunGuard { cancellation: CancellationToken, completed: bool, @@ -233,4 +261,67 @@ mod tests { drop(stream); assert!(!cancellation.is_cancelled()); } + + #[test] + fn connect_error_end_stream_exposes_the_trace_error() { + let frame = connect::encode_error_end_stream(&connect::ConnectStreamError { + code: connect::ConnectCode::InvalidArgument, + message: "unsupported runtime action".into(), + details: Vec::new(), + }) + .unwrap(); + + assert_eq!( + end_stream_error(&frame).as_deref(), + Some("invalid_argument: unsupported runtime action") + ); + assert_eq!(end_stream_error(&connect::encode_end_stream()), None); + } + + #[tokio::test] + async fn connect_error_end_stream_marks_the_local_trace_as_error() { + let store = crate::store::Store::connect("sqlite::memory:") + .await + .unwrap(); + store.set_detailed_logging(true).await.unwrap(); + let trace = CursorTraceRecorder::begin( + store.clone(), + "error-trace", + Some("conversation"), + "local_byok", + Some("model"), + ) + .await + .unwrap(); + let (sender, receiver) = mpsc::unbounded_channel(); + let cancellation = CancellationToken::new(); + sender + .send( + connect::encode_error_end_stream(&connect::ConnectStreamError { + code: connect::ConnectCode::InvalidArgument, + message: "unsupported runtime action".into(), + details: Vec::new(), + }) + .unwrap(), + ) + .unwrap(); + let mut stream = Box::pin(local_body_stream(receiver, cancellation, Some(trace))); + + stream.next().await.unwrap().unwrap(); + + let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(1); + let trace = loop { + let trace = store.cursor_trace("error-trace").await.unwrap().unwrap(); + if trace.status != "running" { + break trace; + } + assert!(tokio::time::Instant::now() < deadline); + tokio::time::sleep(std::time::Duration::from_millis(10)).await; + }; + assert_eq!(trace.status, "error"); + assert_eq!( + trace.error_message.as_deref(), + Some("invalid_argument: unsupported runtime action") + ); + } } diff --git a/server/src/cursor/session.rs b/server/src/cursor/session.rs index 52f2999..2f83275 100644 --- a/server/src/cursor/session.rs +++ b/server/src/cursor/session.rs @@ -88,6 +88,23 @@ impl CursorSession { } pub async fn run(mut self) -> Result<()> { + let result = self.run_inner().await; + if let Err(error) = &result { + self.abort_execs().await; + let error = match error { + Error::Protocol(message) => message.clone(), + error => error.to_string(), + }; + let _ = self + .core + .commands + .send(ClientCommand::ClientClosed { error }) + .await; + } + result + } + + async fn run_inner(&mut self) -> Result<()> { if self.context.compacting { self.handle.emit(&interaction::summary_started())?; } diff --git a/server/src/cursor/tools/codec/mod.rs b/server/src/cursor/tools/codec/mod.rs index 50e4bfa..f06e835 100644 --- a/server/src/cursor/tools/codec/mod.rs +++ b/server/src/cursor/tools/codec/mod.rs @@ -5,4 +5,4 @@ pub use request::{abort, mcp_request, mcp_state_request, request}; pub(crate) use request::{ await_read_request, edit_read_request, json_object_to_prost, mcp_meta_request, }; -pub use response::{client_event, ClientExecEvent}; +pub use response::{client_event, stream_closed, ClientExecEvent}; diff --git a/server/src/cursor/tools/codec/response.rs b/server/src/cursor/tools/codec/response.rs index ebad85a..93d36f6 100644 --- a/server/src/cursor/tools/codec/response.rs +++ b/server/src/cursor/tools/codec/response.rs @@ -130,6 +130,53 @@ pub async fn client_event( Ok(event) } +pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result> { + let Some(entry) = pending.take_exec(id).await else { + return Ok(None); + }; + let error = "Cursor Exec stream closed before returning a terminal result"; + if entry.call.name.eq_ignore_ascii_case("Shell") { + let command = entry + .call + .arguments + .get("command") + .and_then(serde_json::Value::as_str) + .unwrap_or_default() + .to_string(); + let working_directory = entry + .call + .arguments + .get("working_directory") + .and_then(serde_json::Value::as_str) + .unwrap_or_default() + .to_string(); + return Ok(Some(result::from_exec( + entry, + &pb::exec_client_message::Message::ShellResult(pb::ShellResult { + result: Some(pb::shell_result::Result::SpawnError(pb::ShellSpawnError { + command, + working_directory, + error: error.into(), + })), + ..Default::default() + }), + )?)); + } + let rendered = match &entry.stage { + ExecStage::DynamicMcp(definition) => { + interaction::render_dynamic_mcp(&entry.call, definition, false) + } + _ => interaction::render_tool_call(&entry.call, false)?, + }; + Ok(Some(ToolCompletion::from_rendered( + &entry.call, + entry.started_at_ms, + error.into(), + true, + rendered, + )?)) +} + async fn advance_await( entry: PendingExec, result: &pb::exec_client_message::Message, diff --git a/server/src/cursor/tools/runtime.rs b/server/src/cursor/tools/runtime.rs index 6175f29..6e95a97 100644 --- a/server/src/cursor/tools/runtime.rs +++ b/server/src/cursor/tools/runtime.rs @@ -338,6 +338,18 @@ impl CursorToolRuntime { ids } + pub async fn running_task_exec_id(&self, call_id: &str) -> Option { + self.execs + .lock() + .await + .iter() + .filter_map(|(id, entry)| { + (entry.call.call_id == call_id && entry.call.name.eq_ignore_ascii_case("Task")) + .then_some(*id) + }) + .min() + } + fn next_id(&self) -> Result { self.next_id .fetch_add(1, Ordering::Relaxed) diff --git a/server/src/run/actor.rs b/server/src/run/actor.rs index 5927435..04ebe06 100644 --- a/server/src/run/actor.rs +++ b/server/src/run/actor.rs @@ -2,7 +2,12 @@ use std::sync::Arc; use tokio_util::sync::CancellationToken; -use crate::{client::ClientPort, model::PreparedRun, provider::Provider, store::Store}; +use crate::{ + client::{ClientCommand, ClientPort}, + model::PreparedRun, + provider::Provider, + store::Store, +}; use super::{RunEngine, RunOutcome, RunRegistry}; @@ -26,6 +31,7 @@ impl RunActor { &self, prepared: PreparedRun, client: ClientPort, + commands: tokio::sync::mpsc::Sender, cancellation: CancellationToken, ) -> tokio::task::JoinHandle { let run_id = prepared.run_id.clone(); @@ -35,6 +41,7 @@ impl RunActor { conversation_id.clone(), run_id.clone(), cancellation.clone(), + commands, ) .await; let actor = self.clone(); diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index 413ce30..296309f 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -4,7 +4,10 @@ use std::sync::Arc; use tokio_util::sync::CancellationToken; use crate::{ - client::{ClientEvent, ClientPort, CommitBarrier, CommitCause, StateCommitted}, + client::{ + ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion, + StateCommitted, + }, model::{ CanonicalMessage, MessageContent, Origin, PreparedRun, Role, RunAction, ToolRoundAssistant, ToolRoundId, Usage, @@ -158,6 +161,7 @@ impl RunEngine { calls: round.calls.clone(), recovered_started_at_ms: Some(round.started_at_ms), }, + Vec::new(), ) .await { @@ -242,21 +246,27 @@ impl RunEngine { &cycle_cancellation, ); tokio::pin!(cycle); - let cycle = tokio::select! { - result = &mut cycle => result, + let mut pending_insertions = Vec::new(); + let cycle = loop { + tokio::select! { + result = &mut cycle => break result, command = client.commands.recv() => { let message = match command { - Some(crate::client::ClientCommand::RuntimeMessage(message)) => message, - Some(crate::client::ClientCommand::RuntimeEvent(event)) => event.into_message(), - Some(crate::client::ClientCommand::Cancel) => { + Some(ClientCommand::InsertMessages(insertion)) => { + pending_insertions.push(insertion); + continue; + } + Some(ClientCommand::RuntimeMessage(message)) => message, + Some(ClientCommand::RuntimeEvent(event)) => event.into_message(), + Some(ClientCommand::Cancel) => { cycle_cancellation.cancel(); return (RunOutcome::Cancelled, usage); } - Some(crate::client::ClientCommand::ClientClosed { error }) => { + Some(ClientCommand::ClientClosed { error }) => { cycle_cancellation.cancel(); return (RunOutcome::Failed(RunFailure::Client(error)), usage); } - Some(crate::client::ClientCommand::ToolResult(_)) => { + Some(ClientCommand::ToolResult(_)) => { cycle_cancellation.cancel(); return ( RunOutcome::Failed(RunFailure::Protocol( @@ -284,6 +294,19 @@ impl RunEngine { } } } + revision = match append_insertions( + &self.store, + prepared, + client, + cancellation, + revision, + std::mem::take(&mut pending_insertions), + ) + .await + { + Ok((revision, _)) => revision, + Err(outcome) => return (outcome, usage), + }; revision = match append_runtime_message( &self.store, prepared, @@ -294,11 +317,12 @@ impl RunEngine { ) .await { - Ok(revision) => revision, + Ok((revision, _)) => revision, Err(outcome) => return (outcome, usage), }; continue 'model; } + } }; let cycle = match cycle { Ok(cycle) => cycle, @@ -413,6 +437,27 @@ impl RunEngine { Ok(revision) => revision, Err(error) => return (RunOutcome::Failed(error.into()), usage), }; + if !pending_insertions.is_empty() { + let inserted = match append_insertions( + &self.store, + prepared, + client, + cancellation, + revision, + pending_insertions, + ) + .await + { + Ok((next, inserted)) => { + revision = next; + inserted + } + Err(outcome) => return (outcome, usage), + }; + if inserted { + continue 'model; + } + } let (barrier, ready) = CommitBarrier::before_continue(); if emit( client, @@ -453,6 +498,7 @@ impl RunEngine { calls: cycle.calls, recovered_started_at_ms: None, }, + pending_insertions, ) .await { @@ -683,14 +729,36 @@ fn fallback_summary(messages: &[CanonicalMessage]) -> String { ) } -async fn append_runtime_message( +pub(super) async fn append_insertions( + store: &Store, + prepared: &PreparedRun, + client: &mut ClientPort, + cancellation: &CancellationToken, + mut revision: crate::model::RevisionId, + insertions: Vec, +) -> std::result::Result<(crate::model::RevisionId, bool), RunOutcome> { + let mut inserted_any = false; + for insertion in insertions { + for message in insertion.messages { + let (next, inserted) = + append_runtime_message(store, prepared, client, cancellation, revision, message) + .await?; + revision = next; + inserted_any |= inserted; + } + let _ = insertion.delivered.send(()); + } + Ok((revision, inserted_any)) +} + +pub(super) async fn append_runtime_message( store: &Store, prepared: &PreparedRun, client: &mut ClientPort, cancellation: &CancellationToken, revision: crate::model::RevisionId, message: CanonicalMessage, -) -> std::result::Result { +) -> std::result::Result<(crate::model::RevisionId, bool), RunOutcome> { let event_id = message.runtime_event_id.clone().ok_or_else(|| { RunOutcome::Failed(RunFailure::Protocol( "runtime message has no event identity".into(), @@ -706,7 +774,7 @@ async fn append_runtime_message( .await .map_err(|error| RunOutcome::Failed(error.into()))?; if !inserted { - return Ok(revision); + return Ok((revision, false)); } let (barrier, ready) = CommitBarrier::before_continue(); emit( @@ -721,7 +789,7 @@ async fn append_runtime_message( .await .map_err(|_| client_failure())?; wait_for_state_ready(ready, cancellation).await?; - Ok(revision) + Ok((revision, true)) } async fn hydrate_tool_images( diff --git a/server/src/run/registry.rs b/server/src/run/registry.rs index d1305f4..66b450c 100644 --- a/server/src/run/registry.rs +++ b/server/src/run/registry.rs @@ -3,7 +3,10 @@ use std::{collections::HashMap, sync::Arc}; use tokio::sync::Mutex; use tokio_util::sync::CancellationToken; -use crate::model::{ConversationId, RunId}; +use crate::{ + client::{ClientCommand, MessageInsertion}, + model::{CanonicalMessage, ConversationId, RunId}, +}; #[derive(Clone, Default)] pub struct RunRegistry { @@ -13,6 +16,7 @@ pub struct RunRegistry { struct ActiveRun { run_id: RunId, cancellation: CancellationToken, + commands: tokio::sync::mpsc::Sender, } impl RunRegistry { @@ -21,12 +25,14 @@ impl RunRegistry { conversation_id: ConversationId, run_id: RunId, cancellation: CancellationToken, + commands: tokio::sync::mpsc::Sender, ) { let previous = self.active.lock().await.insert( conversation_id, ActiveRun { run_id: run_id.clone(), cancellation, + commands, }, ); if let Some(previous) = previous.filter(|previous| previous.run_id != run_id) { @@ -34,6 +40,37 @@ impl RunRegistry { } } + pub async fn insert_messages( + &self, + conversation_id: &ConversationId, + messages: Vec, + ) -> bool { + if messages.is_empty() { + return true; + } + let commands = self + .active + .lock() + .await + .get(conversation_id) + .map(|run| run.commands.clone()); + let Some(commands) = commands else { + return false; + }; + let (delivered, delivery) = tokio::sync::oneshot::channel(); + if commands + .send(ClientCommand::InsertMessages(MessageInsertion { + messages, + delivered, + })) + .await + .is_err() + { + return false; + } + delivery.await.is_ok() + } + pub async fn release(&self, conversation_id: &ConversationId, run_id: &RunId) { let mut active = self.active.lock().await; if active diff --git a/server/src/run/tool_round.rs b/server/src/run/tool_round.rs index b7f2846..0ae1cf9 100644 --- a/server/src/run/tool_round.rs +++ b/server/src/run/tool_round.rs @@ -1,7 +1,10 @@ use tokio_util::sync::CancellationToken; use crate::{ - client::{ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, StateCommitted}, + client::{ + ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion, + StateCommitted, + }, model::{PreparedRun, RevisionId, ToolCall, ToolRoundAssistant, ToolRoundId}, store::Store, }; @@ -22,6 +25,7 @@ pub(super) async fn execute( cancellation: &CancellationToken, mut revision: RevisionId, round: ToolRound, + insertions: Vec, ) -> std::result::Result { let ToolRound { id: round_id, @@ -66,7 +70,10 @@ pub(super) async fn execute( .await?; let mut remaining = calls.len(); - let mut pending_runtime_messages = Vec::new(); + let mut pending_runtime_messages = insertions + .into_iter() + .map(PendingRuntimeMessage::Insertion) + .collect::>(); while remaining > 0 { let command = tokio::select! { _ = cancellation.cancelled() => return Err(RunOutcome::Cancelled), @@ -116,10 +123,13 @@ pub(super) async fn execute( } } Some(ClientCommand::RuntimeEvent(event)) => { - pending_runtime_messages.push(event.into_message()); + pending_runtime_messages.push(PendingRuntimeMessage::Message(event.into_message())); } Some(ClientCommand::RuntimeMessage(message)) => { - pending_runtime_messages.push(message); + pending_runtime_messages.push(PendingRuntimeMessage::Message(message)); + } + Some(ClientCommand::InsertMessages(insertion)) => { + pending_runtime_messages.push(PendingRuntimeMessage::Insertion(insertion)) } Some(ClientCommand::Cancel) => return Err(RunOutcome::Cancelled), Some(ClientCommand::ClientClosed { error }) => { @@ -128,40 +138,42 @@ pub(super) async fn execute( None => return Err(client_failure()), } } - for message in pending_runtime_messages { - let event_id = message.runtime_event_id.clone().ok_or_else(|| { - RunOutcome::Failed(RunFailure::Protocol( - "runtime message has no event identity".into(), - )) - })?; - let (next, inserted) = store - .append_message_once( - &prepared.conversation_id, - &prepared.run_id, - revision, - &message, - ) - .await - .map_err(failed)?; - revision = next; - if inserted { - let (barrier, ready) = CommitBarrier::before_continue(); - send( - client, - ClientEvent::StateCommitted(StateCommitted { - revision_id: revision, - tool_round_version: 0, - cause: CommitCause::RuntimeEvent { event_id }, - barrier, - }), - ) - .await?; - super::engine::wait_for_state_ready(ready, cancellation).await?; + for pending in pending_runtime_messages { + match pending { + PendingRuntimeMessage::Message(message) => { + revision = super::engine::append_runtime_message( + store, + prepared, + client, + cancellation, + revision, + message, + ) + .await? + .0; + } + PendingRuntimeMessage::Insertion(insertion) => { + revision = super::engine::append_insertions( + store, + prepared, + client, + cancellation, + revision, + vec![insertion], + ) + .await? + .0; + } } } Ok(revision) } +enum PendingRuntimeMessage { + Message(crate::model::CanonicalMessage), + Insertion(MessageInsertion), +} + async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> { client .events diff --git a/server/src/store/revisions.rs b/server/src/store/revisions.rs index 5e638fd..a7048f8 100644 --- a/server/src/store/revisions.rs +++ b/server/src/store/revisions.rs @@ -58,6 +58,38 @@ impl Store { self.load_revision_messages(RevisionId(revision_id)).await } + pub async fn match_revision_prefix( + &self, + conversation_id: &ConversationId, + base_revision_id: RevisionId, + additions: &[CanonicalMessage], + ) -> Result<(RevisionId, usize)> { + let mut revision = base_revision_id; + let mut messages = self.load_revision_messages(revision).await?; + for (index, addition) in additions.iter().enumerate() { + messages.push(addition.clone()); + let digest = message_digest(&messages)?; + let child = sqlx::query_scalar::<_, i64>( + "SELECT revision_id FROM conversation_revisions + WHERE conversation_id = ? AND parent_revision_id = ? AND state_digest = ?", + ) + .bind(conversation_id.as_str()) + .bind(revision.0) + .bind(digest.as_slice()) + .fetch_optional(&self.pool) + .await? + .map(RevisionId); + let Some(child) = child else { + return Ok((revision, index)); + }; + if self.load_revision_messages(child).await? != messages { + return Ok((revision, index)); + } + revision = child; + } + Ok((revision, additions.len())) + } + pub async fn import_revision( &self, conversation_id: &ConversationId, diff --git a/server/tests/background_completion.rs b/server/tests/background_completion.rs index b6e7efb..2c675e5 100644 --- a/server/tests/background_completion.rs +++ b/server/tests/background_completion.rs @@ -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 = 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, diff --git a/server/tests/client_contract.rs b/server/tests/client_contract.rs index 3d8866e..342d34f 100644 --- a/server/tests/client_contract.rs +++ b/server/tests/client_contract.rs @@ -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; diff --git a/server/tests/error_lifecycle.rs b/server/tests/error_lifecycle.rs index c69d98e..a469860 100644 --- a/server/tests/error_lifecycle.rs +++ b/server/tests/error_lifecycle.rs @@ -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) = + 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] diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index 034b2fc..e214bbe 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -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() + }, + )), + } +} diff --git a/server/tests/revision_branch.rs b/server/tests/revision_branch.rs index 930cdf7..1186d99 100644 --- a/server/tests/revision_branch.rs +++ b/server/tests/revision_branch.rs @@ -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 + ] + ); +} diff --git a/server/tests/runtime_modes.rs b/server/tests/runtime_modes.rs index 3ec6b01..25b42f3 100644 --- a/server/tests/runtime_modes.rs +++ b/server/tests/runtime_modes.rs @@ -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::(&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("\nexplain the edited version\n")); + 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("")); let ProjectedContent::Parts(parts) = &request.history[1].content else { diff --git a/server/tests/selected_images.rs b/server/tests/selected_images.rs index 358331a..4338cba 100644 --- a/server/tests/selected_images.rs +++ b/server/tests/selected_images.rs @@ -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; } diff --git a/server/tests/support/fake_provider.rs b/server/tests/support/fake_provider.rs index 6db0d31..f2c5da0 100644 --- a/server/tests/support/fake_provider.rs +++ b/server/tests/support/fake_provider.rs @@ -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>), + Gated { + ready: Arc, + events: Vec>, + }, Pending, } @@ -43,6 +47,17 @@ impl FakeProvider { .unwrap() .push_back(FakeResponse::Pending); } + pub fn push_gated(&self, events: Vec) -> Arc { + 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 { 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()), } } diff --git a/server/tests/text_turn.rs b/server/tests/text_turn.rs index c2d3243..c0a3143 100644 --- a/server/tests/text_turn.rs +++ b/server/tests/text_turn.rs @@ -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(), diff --git a/server/tests/tool_loop.rs b/server/tests/tool_loop.rs index 38a073c..4811ea4 100644 --- a/server/tests/tool_loop.rs +++ b/server/tests/tool_loop.rs @@ -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();