fix: preempt root loop on context injection

This commit is contained in:
leokun
2026-08-27 17:01:03 +08:00
parent aa68205735
commit ee915ee760
12 changed files with 1027 additions and 79 deletions
+1 -1
View File
@@ -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 },
+1 -1
View File
@@ -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 },
+4
View File
@@ -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!(
+13
View File
@@ -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,
+92 -28
View File
@@ -122,6 +122,9 @@ impl CursorSession {
let mut response_text = String::new();
let mut response_thinking = String::new();
let mut active_round = None::<ToolRoundId>;
let mut active_tool_calls = HashSet::<String>::new();
let mut interrupted_rounds = HashSet::<ToolRoundId>::new();
let mut interrupted_tool_calls = HashSet::<String>::new();
let mut final_checkpoint = None::<FinalCheckpoints>;
let mut compaction_checkpoint = None::<pb::ConversationStateStructure>;
let mut turn_usage = None::<Usage>;
@@ -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<String, ToolCompletion>,
interrupted_tool_calls: &HashSet<String>,
) -> Result<Option<ToolCompletion>> {
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<String>,
completions: &HashMap<String, ToolCompletion>,
interrupted_rounds: &mut HashSet<ToolRoundId>,
interrupted_tool_calls: &mut HashSet<String>,
) -> 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));
}
}
+26
View File
@@ -25,6 +25,12 @@ pub async fn client_event(
message: &pb::ExecClientMessage,
pending: &CursorToolRuntime,
) -> Result<ClientExecEvent> {
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<Option<ToolCompletion>> {
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<Optio
)?))
}
fn is_terminal(message: &pb::exec_client_message::Message) -> 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,
+8
View File
@@ -155,6 +155,11 @@ impl ToolDispatcher {
.map(Some)
}
pub async fn interrupt_for_message(&self) -> Vec<u32> {
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<ClientToolEvent> {
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() => {
+36 -1
View File
@@ -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<Mutex<HashMap<u32, PendingExec>>>,
interactions: Arc<Mutex<HashMap<u32, PendingInteraction>>>,
completed: Arc<Mutex<HashMap<u32, String>>>,
interrupted: Arc<Mutex<HashSet<u32>>>,
}
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<u32> {
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::<Vec<_>>();
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<u32> {
let mut ids = self.execs.lock().await.keys().copied().collect::<Vec<_>>();
ids.sort_unstable();
+5
View File
@@ -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<DeferredEdit> {
if let Some(queue) = self.paths.get_mut(&path) {
queue.waiting.push_back(edit);
+82 -14
View File
@@ -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::<Vec<_>>();
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))
}
}
+102 -30
View File
@@ -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<RevisionId, RunOutcome> {
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
+657 -4
View File
@@ -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(
&registry,
"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::ConversationStateStructure>,
) -> 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<ModelEvent> {
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<ModelEvent> {
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<Bytes>,
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<Bytes>,
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 {