diff --git a/server/src/client/command.rs b/server/src/client/command.rs index d29067c..3ae1be4 100644 --- a/server/src/client/command.rs +++ b/server/src/client/command.rs @@ -11,7 +11,7 @@ pub struct MessageInsertion { #[derive(Debug)] pub enum ClientCommand { ToolResult(ToolResult), - RuntimeMessage(CanonicalMessage), + InterruptWithMessage(CanonicalMessage), RuntimeEvent(RuntimeEvent), InsertMessages(MessageInsertion), ClientClosed { error: String }, diff --git a/server/src/client/event.rs b/server/src/client/event.rs index 85470f2..b84e653 100644 --- a/server/src/client/event.rs +++ b/server/src/client/event.rs @@ -9,7 +9,7 @@ use crate::run::RunOutcome; pub enum CommitCause { InitialMessages, ToolRoundStarted(ToolRoundId), - ToolResult { call_id: String }, + ToolResult { call_id: String, interrupted: bool }, FinalTurn, Compaction { summary: String }, RuntimeEvent { event_id: String }, diff --git a/server/src/cursor/actor.rs b/server/src/cursor/actor.rs index 85c158a..679cc81 100644 --- a/server/src/cursor/actor.rs +++ b/server/src/cursor/actor.rs @@ -285,6 +285,10 @@ impl CursorActor { { continue; } + if tool_runtime.is_interrupted(throw.id).await { + tool_runtime.discard_exec(throw.id).await; + continue; + } match tool_runtime.take_exec(throw.id).await { Some(pending) => results_tx.send_error( crate::Error::Protocol(format!( diff --git a/server/src/cursor/interaction/mod.rs b/server/src/cursor/interaction/mod.rs index b2e6d7d..7861eeb 100644 --- a/server/src/cursor/interaction/mod.rs +++ b/server/src/cursor/interaction/mod.rs @@ -152,6 +152,19 @@ pub fn context_injection_queued(injection_id: String) -> pb::AgentServerMessage )) } +pub fn context_injection_rejected(injection_id: String, reason: String) -> pb::AgentServerMessage { + server_interaction(pb::interaction_update::Message::ContextInjectionState( + pb::ContextInjectionStateUpdate { + injection_id, + state: Some(pb::ContextInjectionState { + state: Some(pb::context_injection_state::State::Rejected( + pb::ContextInjectionRejected { reason }, + )), + }), + }, + )) +} + pub fn context_injection_delivered( injection_id: String, delivery_batch_id: String, diff --git a/server/src/cursor/session.rs b/server/src/cursor/session.rs index 2f83275..2999f75 100644 --- a/server/src/cursor/session.rs +++ b/server/src/cursor/session.rs @@ -122,6 +122,9 @@ impl CursorSession { let mut response_text = String::new(); let mut response_thinking = String::new(); let mut active_round = None::; + let mut active_tool_calls = HashSet::::new(); + let mut interrupted_rounds = HashSet::::new(); + let mut interrupted_tool_calls = HashSet::::new(); let mut final_checkpoint = None::; let mut compaction_checkpoint = None::; let mut turn_usage = None::; @@ -130,13 +133,16 @@ impl CursorSession { let mut presentation = Presentation::default(); loop { - let input = if let Some(completion) = ready.pop_front() { + let input = if let Ok(action) = self.runtime_actions.try_recv() { + Input::RuntimeAction(Some(Box::new(action))) + } else if let Some(completion) = ready.pop_front() { Input::Completion(completion) } else { tokio::select! { + biased; + action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)), event = self.core.events.recv() => Input::Event(event), completion = self.results.recv() => Input::CompletionResult(completion), - action = self.runtime_actions.recv() => Input::RuntimeAction(action.map(Box::new)), failure = worker.failures.recv(), if checkpoint_worker_open => Input::CheckpointFailure(failure), } }; @@ -147,15 +153,16 @@ impl CursorSession { } Input::Completion(completion) => { if let Some(completion) = self - .forward_completion(completion, &mut completions) + .forward_completion(completion, &mut completions, &interrupted_tool_calls) .await? { ready.push_back(completion); } } Input::CompletionResult(Some(result)) => { - if let Some(completion) = - self.forward_completion(result?, &mut completions).await? + if let Some(completion) = self + .forward_completion(result?, &mut completions, &interrupted_tool_calls) + .await? { ready.push_back(completion); } @@ -164,7 +171,15 @@ impl CursorSession { return Err(Error::Protocol("tool result channel closed".into())); } Input::RuntimeAction(Some(action)) => { - self.forward_injection(*action).await?; + self.forward_injection( + *action, + active_round.as_ref(), + &active_tool_calls, + &completions, + &mut interrupted_rounds, + &mut interrupted_tool_calls, + ) + .await?; } Input::RuntimeAction(None) => { return Err(Error::Protocol("runtime action channel closed".into())); @@ -283,7 +298,23 @@ impl CursorSession { round_id, calls: round_calls, } => { - active_round = Some(round_id); + active_round = Some(round_id.clone()); + active_tool_calls = round_calls + .iter() + .map(|call| call.call_id.clone()) + .collect(); + // Runtime actions are deliberately prioritized over core events. An + // injection can therefore be observed before the already-queued + // ToolRoundStarted event reaches this session. In that case the + // accepted injection is still pending delivery and this round must be + // detached without starting any root tools. + if interrupted_rounds.contains(&round_id) + || !self.pending_injections.is_empty() + { + interrupted_rounds.insert(round_id.clone()); + interrupted_tool_calls.extend(active_tool_calls.iter().cloned()); + continue; + } for dispatched in self .tools .start_batch( @@ -348,12 +379,11 @@ impl CursorSession { active_round = Some(round_id.clone()); } let mut tool_round_settled = false; - if let CommitCause::ToolResult { call_id } = &state.cause { - let completion = completions.remove(call_id).ok_or_else(|| { - Error::Protocol(format!( - "core committed a tool result without typed Cursor state: {call_id}" - )) - })?; + if let CommitCause::ToolResult { + call_id, + interrupted, + } = &state.cause + { let snapshot = self .store .tool_round(active_round.as_ref().ok_or_else(|| { @@ -372,9 +402,16 @@ impl CursorSession { "committed call is absent from tool round: {call_id}" )) })?; - self.handle - .emit(&interaction::tool_completed(call, &completion))?; - presentation.tool_completed(&completion); + if !interrupted { + let completion = completions.remove(call_id).ok_or_else(|| { + Error::Protocol(format!( + "core committed a tool result without typed Cursor state: {call_id}" + )) + })?; + self.handle + .emit(&interaction::tool_completed(call, &completion))?; + presentation.tool_completed(&completion); + } completed.insert(call_id.clone()); tool_round_settled = snapshot.status == ToolRoundStatus::Settled; } @@ -490,7 +527,10 @@ impl CursorSession { return Err(error); } } - active_round = None; + if let Some(round_id) = active_round.take() { + interrupted_rounds.remove(&round_id); + } + active_tool_calls.clear(); self.tool_runtime.clear_completed().await; } else if !matches!(&state.cause, CommitCause::ToolResult { .. }) && active_round.is_some() @@ -603,7 +643,11 @@ impl CursorSession { &self, mut completion: ToolCompletion, completions: &mut HashMap, + interrupted_tool_calls: &HashSet, ) -> Result> { + if interrupted_tool_calls.contains(&completion.result().call_id) { + return Ok(None); + } if let Some(image) = completion.take_read_image() { let blob_id = self.store.put_blob(&image.data, &[]).await?; completion.persist_read_image(&blob_id, &image)?; @@ -635,19 +679,33 @@ impl CursorSession { Ok(dispatched.completion) } - async fn forward_injection(&mut self, action: pb::InjectContextAction) -> Result<()> { + async fn forward_injection( + &mut self, + action: pb::InjectContextAction, + active_round: Option<&ToolRoundId>, + active_tool_calls: &HashSet, + completions: &HashMap, + interrupted_rounds: &mut HashSet, + interrupted_tool_calls: &mut HashSet, + ) -> Result<()> { if action.injection_id.is_empty() { return Err(Error::Protocol( "InjectContextAction has no injection_id".into(), )); } + if self.injection_ids.contains(&action.injection_id) { + return Ok(()); + } if action.expected_run_id != self.context.request_id { - return Err(Error::Protocol(format!( + let reason = format!( "InjectContextAction expected run {}, active run is {}", action.expected_run_id, self.context.request_id - ))); - } - if self.injection_ids.contains(&action.injection_id) { + ); + self.handle.emit(&interaction::context_injection_rejected( + action.injection_id.clone(), + reason, + ))?; + self.injection_ids.insert(action.injection_id); return Ok(()); } let user_message = match action.payload.as_ref() { @@ -675,25 +733,31 @@ impl CursorSession { ); self.handle .emit(&interaction::context_injection_queued(injection_id.clone()))?; + interrupted_tool_calls.extend( + active_tool_calls + .iter() + .filter(|call_id| !completions.contains_key(*call_id)) + .cloned(), + ); + if let Some(round_id) = active_round { + interrupted_rounds.insert(round_id.clone()); + } + self.interrupt_execs().await; if self .core .commands - .send(ClientCommand::RuntimeMessage(message)) + .send(ClientCommand::InterruptWithMessage(message)) .await .is_err() { self.pending_injections.remove(&injection_id); return Err(Error::RunNotFound(self.context.request_id.clone())); } - self.interrupt_execs().await; Ok(()) } async fn interrupt_execs(&self) { - // Keep runtime entries until Cursor returns the aborted result. The core tool - // round needs that terminal result before it can append the injected context - // after the complete assistant/tool pair and continue the same Run. - for id in self.tool_runtime.running_exec_ids().await { + for id in self.tools.interrupt_for_message().await { let _ = self.handle.emit(&codec::abort(id)); } } diff --git a/server/src/cursor/tools/codec/response.rs b/server/src/cursor/tools/codec/response.rs index 93d36f6..e9d73b5 100644 --- a/server/src/cursor/tools/codec/response.rs +++ b/server/src/cursor/tools/codec/response.rs @@ -25,6 +25,12 @@ pub async fn client_event( message: &pb::ExecClientMessage, pending: &CursorToolRuntime, ) -> Result { + if pending.is_interrupted(message.id).await { + if message.message.as_ref().is_some_and(is_terminal) { + pending.discard_exec(message.id).await; + } + return Ok(ClientExecEvent::Pending); + } let call = match pending.exec_call(message.id).await { Some(call) => call, None if pending.completed_call(message.id).await.is_some() => { @@ -131,6 +137,10 @@ pub async fn client_event( } pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result> { + if pending.is_interrupted(id).await { + pending.discard_exec(id).await; + return Ok(None); + } let Some(entry) = pending.take_exec(id).await else { return Ok(None); }; @@ -177,6 +187,22 @@ pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result bool { + use pb::{exec_client_message::Message, shell_stream::Event}; + + match message { + Message::ShellStream(stream) => matches!( + stream.event.as_ref(), + Some(Event::Exit(_)) + | Some(Event::Backgrounded(_)) + | Some(Event::Rejected(_)) + | Some(Event::PermissionDenied(_)) + | Some(Event::SandboxUnsupported(_)) + ), + _ => true, + } +} + async fn advance_await( entry: PendingExec, result: &pb::exec_client_message::Message, diff --git a/server/src/cursor/tools/mod.rs b/server/src/cursor/tools/mod.rs index 2003766..251968f 100644 --- a/server/src/cursor/tools/mod.rs +++ b/server/src/cursor/tools/mod.rs @@ -155,6 +155,11 @@ impl ToolDispatcher { .map(Some) } + pub async fn interrupt_for_message(&self) -> Vec { + self.edit_schedule.lock().await.clear(); + self.runtime.interrupt_for_message().await + } + async fn start( &self, call: &ToolCall, @@ -193,6 +198,9 @@ impl ToolDispatcher { &self, response: &pb::InteractionResponse, ) -> Result { + if self.runtime.is_interrupted(response.id).await { + return Ok(ClientToolEvent::Pending); + } let pending = match self.runtime.take_interaction(response.id).await { Some(pending) => pending, None if self.runtime.completed_call(response.id).await.is_some() => { diff --git a/server/src/cursor/tools/runtime.rs b/server/src/cursor/tools/runtime.rs index 6e95a97..798d01f 100644 --- a/server/src/cursor/tools/runtime.rs +++ b/server/src/cursor/tools/runtime.rs @@ -1,5 +1,5 @@ use std::{ - collections::HashMap, + collections::{HashMap, HashSet}, sync::{ atomic::{AtomicU32, Ordering}, Arc, @@ -19,6 +19,7 @@ pub struct CursorToolRuntime { execs: Arc>>, interactions: Arc>>, completed: Arc>>, + interrupted: Arc>>, } pub(crate) struct PendingExec { @@ -311,6 +312,10 @@ impl CursorToolRuntime { self.completed.lock().await.get(&id).cloned() } + pub async fn is_interrupted(&self, id: u32) -> bool { + self.interrupted.lock().await.contains(&id) + } + pub async fn clear_completed(&self) { self.completed.lock().await.clear(); } @@ -329,9 +334,39 @@ impl CursorToolRuntime { ids.sort_unstable(); self.interactions.lock().await.clear(); self.completed.lock().await.clear(); + self.interrupted.lock().await.clear(); ids } + pub async fn interrupt_for_message(&self) -> Vec { + let (abort_ids, interrupted_ids) = { + let mut entries = self.execs.lock().await; + let mut abort_ids = Vec::new(); + let mut interrupted_ids = Vec::new(); + entries.retain(|id, entry| { + interrupted_ids.push(*id); + let keep_running = entry.call.name.eq_ignore_ascii_case("Task"); + if !keep_running { + abort_ids.push(*id); + } + keep_running + }); + (abort_ids, interrupted_ids) + }; + let interaction_ids = { + let mut interactions = self.interactions.lock().await; + let ids = interactions.keys().copied().collect::>(); + interactions.clear(); + ids + }; + let mut interrupted = self.interrupted.lock().await; + interrupted.extend(interrupted_ids); + interrupted.extend(interaction_ids); + let mut abort_ids = abort_ids; + abort_ids.sort_unstable(); + abort_ids + } + pub async fn running_exec_ids(&self) -> Vec { let mut ids = self.execs.lock().await.keys().copied().collect::>(); ids.sort_unstable(); diff --git a/server/src/cursor/tools/schedule.rs b/server/src/cursor/tools/schedule.rs index 08ec794..7ee95c0 100644 --- a/server/src/cursor/tools/schedule.rs +++ b/server/src/cursor/tools/schedule.rs @@ -23,6 +23,11 @@ pub(super) struct DeferredEdit { } impl EditSchedule { + pub fn clear(&mut self) { + self.paths.clear(); + self.active_paths.clear(); + } + pub fn start_or_defer(&mut self, path: String, edit: DeferredEdit) -> Option { if let Some(queue) = self.paths.get_mut(&path) { queue.waiting.push_back(edit); diff --git a/server/src/run/engine.rs b/server/src/run/engine.rs index 296309f..6a35cf4 100644 --- a/server/src/run/engine.rs +++ b/server/src/run/engine.rs @@ -249,14 +249,14 @@ impl RunEngine { let mut pending_insertions = Vec::new(); let cycle = loop { tokio::select! { - result = &mut cycle => break result, + biased; command = client.commands.recv() => { let message = match command { Some(ClientCommand::InsertMessages(insertion)) => { pending_insertions.push(insertion); continue; } - Some(ClientCommand::RuntimeMessage(message)) => message, + Some(ClientCommand::InterruptWithMessage(message)) => message, Some(ClientCommand::RuntimeEvent(event)) => event.into_message(), Some(ClientCommand::Cancel) => { cycle_cancellation.cancel(); @@ -321,7 +321,8 @@ impl RunEngine { Err(outcome) => return (outcome, usage), }; continue 'model; - } + }, + result = &mut cycle => break result, } }; let cycle = match cycle { @@ -558,23 +559,68 @@ impl RunEngine { let cycle_cancellation = cancellation.child_token(); let (silent_events, mut discarded_events) = tokio::sync::mpsc::channel(256); let drain = tokio::spawn(async move { while discarded_events.recv().await.is_some() {} }); - let cycle = consume_model_cycle( - self.provider.stream(invocation, cycle_cancellation.clone()), - &silent_events, - &cycle_cancellation, - ) - .await; + let mut pending_insertions = Vec::new(); + let mut interrupted_message = None; + let cycle = { + let cycle = consume_model_cycle( + self.provider.stream(invocation, cycle_cancellation.clone()), + &silent_events, + &cycle_cancellation, + ); + tokio::pin!(cycle); + loop { + tokio::select! { + biased; + command = client.commands.recv() => match command { + Some(ClientCommand::InsertMessages(insertion)) => { + pending_insertions.push(insertion); + } + Some(ClientCommand::InterruptWithMessage(message)) => { + cycle_cancellation.cancel(); + interrupted_message = Some(message); + break cycle.await; + } + Some(ClientCommand::RuntimeEvent(event)) => { + cycle_cancellation.cancel(); + interrupted_message = Some(event.into_message()); + break cycle.await; + } + Some(ClientCommand::Cancel) => { + cycle_cancellation.cancel(); + return Err(RunOutcome::Cancelled); + } + Some(ClientCommand::ClientClosed { error }) => { + cycle_cancellation.cancel(); + return Err(RunOutcome::Failed(RunFailure::Client(error))); + } + Some(ClientCommand::ToolResult(_)) => { + cycle_cancellation.cancel(); + return Err(RunOutcome::Failed(RunFailure::Protocol( + "received a tool result while automatic compaction was running".into(), + ))); + } + None => { + cycle_cancellation.cancel(); + return Err(client_failure()); + } + }, + result = &mut cycle => break result, + } + } + }; drop(silent_events); let _ = drain.await; - let (summary, compaction_usage) = match cycle { - Ok(cycle) if cycle.calls.is_empty() && !cycle.text.trim().is_empty() => { + let (summary, compaction_usage) = match (interrupted_message.is_some(), cycle) { + (true, Ok(cycle)) => (fallback_summary(&compactable), cycle.usage), + (true, Err(failure)) => (fallback_summary(&compactable), failure.usage), + (false, Ok(cycle)) if cycle.calls.is_empty() && !cycle.text.trim().is_empty() => { (cycle.text.trim().to_string(), cycle.usage) } - Ok(cycle) => { + (false, Ok(cycle)) => { tracing::warn!("automatic compaction returned no usable summary; using fallback"); (fallback_summary(&compactable), cycle.usage) } - Err(failure) => { + (false, Err(failure)) => { tracing::warn!(error = ?failure.failure, "automatic compaction model failed; using fallback"); (fallback_summary(&compactable), failure.usage) } @@ -594,7 +640,7 @@ impl RunEngine { let mut replacement = retained_request_context.into_iter().collect::>(); replacement.push(summary_message); replacement.extend(prepared.initial_messages.iter().cloned()); - let revision = self + let mut revision = self .store .replace_revision( &prepared.conversation_id, @@ -620,6 +666,28 @@ impl RunEngine { emit(client, ClientEvent::AutoCompactionCompleted) .await .map_err(|_| client_failure())?; + revision = append_insertions( + &self.store, + prepared, + client, + cancellation, + revision, + pending_insertions, + ) + .await? + .0; + if let Some(message) = interrupted_message { + revision = append_runtime_message( + &self.store, + prepared, + client, + cancellation, + revision, + message, + ) + .await? + .0; + } Ok((revision, compaction_usage)) } } diff --git a/server/src/run/tool_round.rs b/server/src/run/tool_round.rs index 0ae1cf9..a48a1a8 100644 --- a/server/src/run/tool_round.rs +++ b/server/src/run/tool_round.rs @@ -1,3 +1,5 @@ +use std::collections::HashSet; + use tokio_util::sync::CancellationToken; use crate::{ @@ -5,7 +7,7 @@ use crate::{ ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion, StateCommitted, }, - model::{PreparedRun, RevisionId, ToolCall, ToolRoundAssistant, ToolRoundId}, + model::{PreparedRun, RevisionId, ToolCall, ToolResult, ToolRoundAssistant, ToolRoundId}, store::Store, }; @@ -70,6 +72,7 @@ pub(super) async fn execute( .await?; let mut remaining = calls.len(); + let mut completed_call_ids = HashSet::new(); let mut pending_runtime_messages = insertions .into_iter() .map(PendingRuntimeMessage::Insertion) @@ -92,6 +95,7 @@ pub(super) async fn execute( .await .map_err(failed)?; revision = committed.revision_id; + completed_call_ids.insert(call_id.clone()); tracing::info!( round_id = %round_id, call_id, @@ -113,7 +117,10 @@ pub(super) async fn execute( ClientEvent::StateCommitted(StateCommitted { revision_id: revision, tool_round_version: committed.tool_round_version, - cause: CommitCause::ToolResult { call_id }, + cause: CommitCause::ToolResult { + call_id, + interrupted: false, + }, barrier, }), ) @@ -125,8 +132,66 @@ pub(super) async fn execute( Some(ClientCommand::RuntimeEvent(event)) => { pending_runtime_messages.push(PendingRuntimeMessage::Message(event.into_message())); } - Some(ClientCommand::RuntimeMessage(message)) => { - pending_runtime_messages.push(PendingRuntimeMessage::Message(message)); + Some(ClientCommand::InterruptWithMessage(message)) => { + for call in calls + .iter() + .filter(|call| !completed_call_ids.contains(&call.call_id)) + { + let result = ToolResult { + call_id: call.call_id.clone(), + content: "Tool execution was interrupted by a newer user message.".into(), + is_error: true, + image: None, + }; + let committed = store + .commit_tool_result( + &prepared.conversation_id, + &prepared.run_id, + &round_id, + &result, + ) + .await + .map_err(failed)?; + revision = committed.revision_id; + let (barrier, ready) = if committed.settled { + let (barrier, ready) = CommitBarrier::before_continue(); + (barrier, Some(ready)) + } else { + (CommitBarrier::None, None) + }; + send( + client, + ClientEvent::StateCommitted(StateCommitted { + revision_id: revision, + tool_round_version: committed.tool_round_version, + cause: CommitCause::ToolResult { + call_id: call.call_id.clone(), + interrupted: true, + }, + barrier, + }), + ) + .await?; + if let Some(ready) = ready { + super::engine::wait_for_state_ready(ready, cancellation).await?; + } + } + for pending in pending_runtime_messages { + revision = + append_pending(store, prepared, client, cancellation, revision, pending) + .await?; + } + revision = super::engine::append_runtime_message( + store, + prepared, + client, + cancellation, + revision, + message, + ) + .await? + .0; + return Ok(revision); } Some(ClientCommand::InsertMessages(insertion)) => { pending_runtime_messages.push(PendingRuntimeMessage::Insertion(insertion)) @@ -139,32 +204,7 @@ pub(super) async fn execute( } } 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; - } - } + revision = append_pending(store, prepared, client, cancellation, revision, pending).await?; } Ok(revision) } @@ -174,6 +214,38 @@ enum PendingRuntimeMessage { Insertion(MessageInsertion), } +async fn append_pending( + store: &Store, + prepared: &PreparedRun, + client: &mut ClientPort, + cancellation: &CancellationToken, + revision: RevisionId, + pending: PendingRuntimeMessage, +) -> std::result::Result { + match pending { + PendingRuntimeMessage::Message(message) => Ok(super::engine::append_runtime_message( + store, + prepared, + client, + cancellation, + revision, + message, + ) + .await? + .0), + PendingRuntimeMessage::Insertion(insertion) => Ok(super::engine::append_insertions( + store, + prepared, + client, + cancellation, + revision, + vec![insertion], + ) + .await? + .0), + } +} + async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> { client .events diff --git a/server/tests/interrupt.rs b/server/tests/interrupt.rs index e214bbe..2e1a6b6 100644 --- a/server/tests/interrupt.rs +++ b/server/tests/interrupt.rs @@ -5,11 +5,15 @@ mod fixtures; use std::sync::Arc; +use bytes::Bytes; use cursor_server::{ cursor::prompting::{PromptAssets, PromptCompiler}, cursor::{connect, proto::agent::v1 as pb}, cursor::{CursorCommand, CursorSessionRegistry}, - model::{ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind}, + model::{ + ConversationId, ModelConfigInput, ModelSpec, ModelType, PreparedRun, PromptSpec, RunAction, + RunId, RunKind, Usage, OPENAI_CHAT_ENDPOINT, + }, provider::{FinishReason, ModelEvent}, run::RunRegistry, store::RunStatus, @@ -428,6 +432,472 @@ async fn injected_user_context_restarts_only_the_active_model_cycle() { ); } +#[tokio::test] +async fn injected_user_context_aborts_pending_tools_and_ignores_late_results() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + provider.push(tool_response("call-1", "Read", "{\"path\":\"/tmp/a\"}")); + let release = provider.push_gated(text_response("continued after tool interruption")); + 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("interrupt-tool-request") + .await + .unwrap(); + let mut output = handle.subscribe(); + handle + .command(CursorCommand::Append { + seqno: 0, + message: Box::new(client_run_for( + "interrupt-tool-request", + "interrupt-tool-conversation", + )), + }) + .await + .unwrap(); + + let mut append_seqno = 1; + let exec_id = wait_for_exec(&handle, &mut output, &mut append_seqno, "Read").await; + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(runtime_injection_for( + "tool-injection", + "interrupt-tool-request", + )), + }) + .await + .unwrap(); + append_seqno += 1; + + let mut saw_abort = false; + let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5); + while provider.requests().len() < 2 || !saw_abort { + assert!( + tokio::time::Instant::now() < deadline, + "root model did not restart after tool interruption" + ); + if let Ok(Some(frame)) = + tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await + { + let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + let server = pb::AgentServerMessage::decode(payload).unwrap(); + if let Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) = + server.message + { + if let Some(pb::exec_server_control_message::Message::Abort(abort)) = + control.message + { + assert_eq!(abort.id, exec_id); + saw_abort = true; + } + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + } + + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(read_success(exec_id)), + }) + .await + .unwrap(); + append_seqno += 1; + release.notify_one(); + + drain_successfully(&handle, &mut output, &mut append_seqno).await; + + let requests = provider.requests(); + assert_eq!( + requests[0].history, + requests[1].history[..requests[0].history.len()] + ); + let history = serde_json::to_string(&requests[1].history).unwrap(); + let interrupted = history + .find("Tool execution was interrupted by a newer user message.") + .expect("interrupted tool result missing from provider history"); + let injected = history + .find("injected follow-up") + .expect("injected message missing from provider history"); + assert!(interrupted < injected); +} + +#[tokio::test] +async fn injected_user_context_detaches_subagents_without_cancelling_them() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + provider.push(tool_response( + "task-call", + "Task", + &serde_json::json!({ + "description": "Inspect protocol", + "prompt": "Inspect the protocol", + "subagent_type": "generalPurpose", + "run_in_background": false + }) + .to_string(), + )); + let release = provider.push_gated(text_response("continued while subagent runs")); + 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("detach-subagent-request") + .await + .unwrap(); + let mut output = handle.subscribe(); + handle + .command(CursorCommand::Append { + seqno: 0, + message: Box::new(client_run_for( + "detach-subagent-request", + "detach-subagent-conversation", + )), + }) + .await + .unwrap(); + + let mut append_seqno = 1; + let exec_id = wait_for_exec(&handle, &mut output, &mut append_seqno, "Task").await; + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(runtime_injection_for( + "subagent-injection", + "detach-subagent-request", + )), + }) + .await + .unwrap(); + append_seqno += 1; + + let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5); + while provider.requests().len() < 2 { + assert!( + tokio::time::Instant::now() < deadline, + "root model did not restart while subagent remained active" + ); + if let Ok(Some(frame)) = + tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await + { + let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + let server = pb::AgentServerMessage::decode(payload).unwrap(); + if let Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) = + server.message + { + if let Some(pb::exec_server_control_message::Message::Abort(abort)) = + control.message + { + assert_ne!(abort.id, exec_id, "Task must not be aborted by injection"); + } + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + } + + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(subagent_success(exec_id)), + }) + .await + .unwrap(); + append_seqno += 1; + release.notify_one(); + + drain_successfully(&handle, &mut output, &mut append_seqno).await; + + let history = serde_json::to_string(&provider.requests()[1].history).unwrap(); + assert!(history.contains("Tool execution was interrupted by a newer user message.")); + assert!(history.contains("injected follow-up")); +} + +#[tokio::test] +async fn injected_user_context_interrupts_automatic_compaction() { + let (_directory, store) = fixtures::temp_store().await; + let model = store + .create_model(&ModelConfigInput { + sort_order: 0, + display_name: "Test Model".into(), + model_type: ModelType::OpenAi, + base_url: "https://example.com/v1/chat/completions".into(), + use_full_url: true, + api_key: "test-key".into(), + tooltip_data: "Test Model".into(), + model_id: "test-model".into(), + reasoning_effort: None, + openai_endpoint: OPENAI_CHAT_ENDPOINT.into(), + openai_extra_params_enabled: false, + openai_extra_params: serde_json::json!({}), + custom_headers_enabled: false, + custom_headers: serde_json::json!({}), + anthropic_extra_params_enabled: false, + anthropic_extra_params: serde_json::json!({}), + context_window_tokens: Some(10_001), + max_completion_tokens: None, + anthropic_max_tokens: None, + anthropic_thinking_effort: None, + thinking_budget_tokens: None, + }) + .await + .unwrap(); + let provider = fake_provider::FakeProvider::default(); + provider.push(text_response("seed answer")); + provider.push_pending(); + provider.push(text_response("continued after compacting injection")); + 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 seed_state = run_to_end( + ®istry, + "seed-request", + client_run_for_model( + "seed-request", + "compaction-injection-conversation", + &model.model_hash, + ), + ) + .await; + + let handle = registry + .get_or_create("inject-during-compaction") + .await + .unwrap(); + let mut output = handle.subscribe(); + let mut compacting_request = client_run_for_model_with_state( + "inject-during-compaction", + "compaction-injection-conversation", + &model.model_hash, + Some(seed_state), + ); + let Some(pb::agent_client_message::Message::RunRequest(request)) = + compacting_request.message.as_mut() + else { + panic!("expected RunRequest") + }; + request.requested_model.as_mut().unwrap().parameters.push( + pb::requested_model::ModelParameterValue { + id: "context".into(), + value: "10001".into(), + }, + ); + handle + .command(CursorCommand::Append { + seqno: 0, + message: Box::new(compacting_request), + }) + .await + .unwrap(); + + let mut append_seqno = 1; + let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5); + while provider.requests().len() < 2 { + assert!( + tokio::time::Instant::now() < deadline, + "automatic compaction did not start" + ); + if let Ok(Some(frame)) = + tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await + { + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + } + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(runtime_injection_for( + "compaction-injection", + "inject-during-compaction", + )), + }) + .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(); + if let Some(pb::agent_server_message::Message::InteractionUpdate(update)) = server.message { + if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message { + saw_continued |= delta.text.contains("continued after compacting injection"); + } + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + + let requests = provider.requests(); + assert_eq!(requests.len(), 3); + assert!(requests[1] + .prompt + .instructions + .starts_with("Summarize the conversation for the next model turn.")); + assert!(!serde_json::to_string(&requests[1].history) + .unwrap() + .contains("injected follow-up")); + assert!(serde_json::to_string(&requests[2].history) + .unwrap() + .contains("injected follow-up")); + assert!(saw_continued); +} + +#[tokio::test] +async fn stale_context_injection_is_rejected_without_failing_the_active_run() { + let (_directory, store) = fixtures::temp_store().await; + let provider = fake_provider::FakeProvider::default(); + let release = provider.push_gated(vec![ + ModelEvent::Start { + model_call_id: "active-cycle".into(), + }, + ModelEvent::TextStart, + ModelEvent::TextDelta("active run completed".into()), + ModelEvent::TextEnd, + ModelEvent::Done(FinishReason::Stop), + ]); + let assets = PromptAssets::load( + std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .join("prompt/cursor") + .as_path(), + ) + .unwrap(); + let registry = CursorSessionRegistry::new( + store, + Arc::new(provider.clone()), + PromptCompiler::new(assets), + Default::default(), + ); + let handle = registry.get_or_create("active-request").await.unwrap(); + let mut output = handle.subscribe(); + handle + .command(CursorCommand::Append { + seqno: 0, + message: Box::new(client_run_for( + "active-request", + "stale-injection-conversation", + )), + }) + .await + .unwrap(); + + let mut append_seqno = 1; + let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(5); + while provider.requests().is_empty() { + assert!( + tokio::time::Instant::now() < deadline, + "provider did not start" + ); + if let Ok(Some(frame)) = + tokio::time::timeout(std::time::Duration::from_millis(20), output.recv()).await + { + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } + } + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(runtime_injection_for("stale-injection", "replaced-request")), + }) + .await + .unwrap(); + append_seqno += 1; + handle + .command(CursorCommand::Append { + seqno: append_seqno, + message: Box::new(runtime_injection_for("stale-injection", "replaced-request")), + }) + .await + .unwrap(); + append_seqno += 1; + + let mut rejection_count = 0; + let mut released = 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(); + let rejected = match server.message { + Some(pb::agent_server_message::Message::InteractionUpdate(pb::InteractionUpdate { + message: + Some(pb::interaction_update::Message::ContextInjectionState( + pb::ContextInjectionStateUpdate { + injection_id, + state: + Some(pb::ContextInjectionState { + state: + Some(pb::context_injection_state::State::Rejected(rejected)), + }), + }, + )), + .. + })) if injection_id == "stale-injection" => { + assert_eq!( + rejected.reason, + "InjectContextAction expected run replaced-request, active run is active-request" + ); + true + } + _ => false, + }; + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + if rejected { + rejection_count += 1; + if !released { + released = true; + release.notify_one(); + } + } + } + + assert!(released, "stale injection was not rejected"); + assert_eq!(rejection_count, 1); + assert_eq!(provider.requests().len(), 1); +} + #[tokio::test] async fn cancel_subagent_action_aborts_the_target_task_and_keeps_the_parent_running() { let (_directory, store) = fixtures::temp_store().await; @@ -614,6 +1084,23 @@ fn client_run() -> pb::AgentClientMessage { } fn client_run_for(request_id: &str, conversation_id: &str) -> pb::AgentClientMessage { + client_run_for_model(request_id, conversation_id, "test-model") +} + +fn client_run_for_model( + request_id: &str, + conversation_id: &str, + model_id: &str, +) -> pb::AgentClientMessage { + client_run_for_model_with_state(request_id, conversation_id, model_id, None) +} + +fn client_run_for_model_with_state( + request_id: &str, + conversation_id: &str, + model_id: &str, + state: Option, +) -> pb::AgentClientMessage { pb::AgentClientMessage { message: Some(pb::agent_client_message::Message::RunRequest( pb::AgentRunRequest { @@ -634,15 +1121,177 @@ fn client_run_for(request_id: &str, conversation_id: &str) -> pb::AgentClientMes conversation_id: Some(conversation_id.into()), run_id: Some(request_id.into()), requested_model: Some(pb::RequestedModel { - model_id: "test-model".into(), + model_id: model_id.into(), ..Default::default() }), + conversation_state: state, ..Default::default() }, )), } } +fn text_response(text: &str) -> Vec { + vec![ + ModelEvent::Start { + model_call_id: format!("call-{text}"), + }, + ModelEvent::TextStart, + ModelEvent::TextDelta(text.into()), + ModelEvent::TextEnd, + ModelEvent::Usage(Usage { + input_tokens: Some(1), + output_tokens: Some(1), + total_tokens: Some(2), + ..Default::default() + }), + ModelEvent::Done(FinishReason::Stop), + ] +} + +fn tool_response(call_id: &str, name: &str, arguments: &str) -> Vec { + vec![ + ModelEvent::Start { + model_call_id: format!("call-{call_id}"), + }, + ModelEvent::ToolCallStart { + index: 0, + call_id: call_id.into(), + name: name.into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: arguments.into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Done(FinishReason::ToolUse), + ] +} + +async fn wait_for_exec( + handle: &cursor_server::cursor::CursorSessionHandle, + output: &mut tokio::sync::mpsc::UnboundedReceiver, + append_seqno: &mut i64, + tool: &str, +) -> u32 { + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) + .await + .unwrap() + .expect("RunSSE closed before 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(); + if let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = server.message { + let matches = match exec.message.as_ref() { + Some(pb::exec_server_message::Message::ReadArgs(_)) => tool == "Read", + Some(pb::exec_server_message::Message::SubagentArgs(_)) => tool == "Task", + _ => false, + }; + if matches { + return exec.id; + } + } + acknowledge_kv(handle, append_seqno, &frame).await; + } +} + +async fn drain_successfully( + handle: &cursor_server::cursor::CursorSessionHandle, + output: &mut tokio::sync::mpsc::UnboundedReceiver, + append_seqno: &mut i64, +) { + 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"{}"); + return; + } + acknowledge_kv(handle, append_seqno, &frame).await; + } +} + +fn read_success(id: u32) -> pb::AgentClientMessage { + pb::AgentClientMessage { + message: Some(pb::agent_client_message::Message::ExecClientMessage( + pb::ExecClientMessage { + id, + message: Some(pb::exec_client_message::Message::ReadResult( + pb::ReadResult { + result: Some(pb::read_result::Result::Success(pb::ReadSuccess { + path: "/tmp/a".into(), + total_lines: 1, + file_size: 1, + output: Some(pb::read_success::Output::Content("late".into())), + ..Default::default() + })), + }, + )), + ..Default::default() + }, + )), + } +} + +fn subagent_success(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::Success(pb::SubagentSuccess { + agent_id: "detached-child".into(), + ..Default::default() + })), + }, + )), + ..Default::default() + }, + )), + } +} + +async fn run_to_end( + registry: &CursorSessionRegistry, + request_id: &str, + request: pb::AgentClientMessage, +) -> pb::ConversationStateStructure { + let handle = registry.get_or_create(request_id).await.unwrap(); + let mut output = handle.subscribe(); + handle + .command(CursorCommand::Append { + seqno: 0, + message: Box::new(request), + }) + .await + .unwrap(); + let mut append_seqno = 1; + let mut state = None; + loop { + let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv()) + .await + .unwrap() + .expect("RunSSE closed before EndStream"); + let (flags, _) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + if flags & connect::END_STREAM_FLAG != 0 { + return state.expect("Run ended without a checkpoint"); + } + let (_, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap(); + let server = pb::AgentServerMessage::decode(payload).unwrap(); + if let Some(pb::agent_server_message::Message::ConversationCheckpointUpdate(update)) = + server.message + { + state = Some(update); + } + acknowledge_kv(&handle, &mut append_seqno, &frame).await; + } +} + async fn acknowledge_kv( handle: &cursor_server::cursor::CursorSessionHandle, append_seqno: &mut i64, @@ -700,13 +1349,17 @@ fn runtime_user_message() -> pb::AgentClientMessage { } fn runtime_injection() -> pb::AgentClientMessage { + runtime_injection_for("injection-1", "inject-request") +} + +fn runtime_injection_for(injection_id: &str, expected_run_id: &str) -> pb::AgentClientMessage { pb::AgentClientMessage { message: Some(pb::agent_client_message::Message::ConversationAction( pb::ConversationAction { action: Some(pb::conversation_action::Action::InjectContextAction( pb::InjectContextAction { - injection_id: "injection-1".into(), - expected_run_id: "inject-request".into(), + injection_id: injection_id.into(), + expected_run_id: expected_run_id.into(), payload: Some(pb::inject_context_action::Payload::UserContext( pb::UserContextInjection { user_message: Some(pb::UserMessage {