mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
fix: preempt root loop on context injection
This commit is contained in:
@@ -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 },
|
||||
|
||||
@@ -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 },
|
||||
|
||||
@@ -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!(
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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() => {
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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::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 {
|
||||
|
||||
Reference in New Issue
Block a user