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