fix: shell

This commit is contained in:
leokun
2026-08-26 20:20:38 +08:00
parent df053c3720
commit 5450fc76e2
25 changed files with 1367 additions and 101 deletions
+10 -1
View File
@@ -1,10 +1,19 @@
use tokio::sync::oneshot;
use crate::model::{CanonicalMessage, RuntimeEvent, ToolResult}; use crate::model::{CanonicalMessage, RuntimeEvent, ToolResult};
#[derive(Clone, Debug, PartialEq)] #[derive(Debug)]
pub struct MessageInsertion {
pub messages: Vec<CanonicalMessage>,
pub delivered: oneshot::Sender<()>,
}
#[derive(Debug)]
pub enum ClientCommand { pub enum ClientCommand {
ToolResult(ToolResult), ToolResult(ToolResult),
RuntimeMessage(CanonicalMessage), RuntimeMessage(CanonicalMessage),
RuntimeEvent(RuntimeEvent), RuntimeEvent(RuntimeEvent),
InsertMessages(MessageInsertion),
ClientClosed { error: String }, ClientClosed { error: String },
Cancel, Cancel,
} }
+50 -14
View File
@@ -151,15 +151,39 @@ impl CursorActor {
context.dynamic_tools.keys().cloned().collect(), context.dynamic_tools.keys().cloned().collect(),
context.turn_user.clone(), context.turn_user.clone(),
); );
if context.background_completion
&& dependencies
.run_registry
.insert_messages(
&prepared.conversation_id,
prepared.initial_messages.clone(),
)
.await
{
crate::cursor::lifecycle::finish_success(
&handle,
);
let _ = handle
.command(CursorCommand::Finished)
.await;
return;
}
let cancellation = handle.cancellation(); let cancellation = handle.cancellation();
let (port, core) = crate::client::session(256); let (port, core) = crate::client::session(256);
let core_commands = core.commands.clone();
let actor = RunActor::new( let actor = RunActor::new(
dependencies.store.clone(), dependencies.store.clone(),
dependencies.provider, dependencies.provider,
dependencies.run_registry, dependencies.run_registry,
); );
let core_run = let core_run = actor
actor.spawn(prepared, port, cancellation).await; .spawn(
prepared,
port,
core_commands,
cancellation,
)
.await;
let session = CursorSession::new( let session = CursorSession::new(
handle.clone(), handle.clone(),
dependencies.store, dependencies.store,
@@ -181,7 +205,6 @@ impl CursorActor {
%error, %error,
"Cursor session failed" "Cursor session failed"
); );
handle.cancel();
let _ = crate::cursor::lifecycle::fail( let _ = crate::cursor::lifecycle::fail(
&handle, &error, &handle, &error,
); );
@@ -235,12 +258,14 @@ impl CursorActor {
{ {
continue; continue;
} }
if tool_runtime.take_exec(close.id).await.is_some() match codec::stream_closed(close.id, &tool_runtime)
.await
{ {
results_tx.send_error(crate::Error::Protocol(format!( Ok(Some(completion)) => {
"Exec stream closed before result for id: {}", results_tx.send(completion)
close.id }
))); Ok(None) => {}
Err(error) => results_tx.send_error(error),
} }
} }
Some(Message::Throw(throw)) => { Some(Message::Throw(throw)) => {
@@ -311,13 +336,12 @@ impl CursorActor {
// //
// The remaining unimplemented Action variants are // The remaining unimplemented Action variants are
// ShellCommandAction, StartPlanAction, // ShellCommandAction, StartPlanAction,
// AsyncAskQuestionCompletionAction, CancelSubagentAction, // AsyncAskQuestionCompletionAction, BackgroundShellAction,
// BackgroundShellAction, BackgroundSubagentAction, // BackgroundSubagentAction,
// SubscriptionNotificationAction and GoalContinuationAction. // SubscriptionNotificationAction and GoalContinuationAction.
// CancelSubagentAction must not start an LLM; variants whose wire // Variants whose wire behavior is not captured yet need evidence
// behavior is not captured yet need evidence before assigning // before assigning semantics. Every unsupported runtime Action must
// semantics. Every unsupported runtime Action must return an explicit // return an explicit Protocol Error rather than falling through silently.
// Protocol Error rather than falling through silently.
Some( Some(
pb::agent_client_message::Message::ConversationAction( pb::agent_client_message::Message::ConversationAction(
action, action,
@@ -341,6 +365,18 @@ impl CursorActor {
)); ));
} }
} }
Some(
pb::conversation_action::Action::CancelSubagentAction(
action,
),
) => {
if let Some(id) = tool_runtime
.running_task_exec_id(&action.subagent_id)
.await
{
let _ = handle.emit(&codec::abort(id));
}
}
Some(action) => { Some(action) => {
results_tx.send_error(crate::Error::Protocol(format!( results_tx.send_error(crate::Error::Protocol(format!(
"unsupported runtime ConversationAction: {}", "unsupported runtime ConversationAction: {}",
+33 -2
View File
@@ -87,8 +87,15 @@ pub(super) fn project(
} }
pb::BackgroundTaskKind::Unspecified => unreachable!(), pb::BackgroundTaskKind::Unspecified => unreachable!(),
}; };
let identity = agent_id.unwrap_or(&completion.task_id); let tool_call_id = completion
let identity = format!("{}:{identity}", kind.as_str_name()); .tool_call_id
.as_deref()
.filter(|id| !id.is_empty())
.ok_or_else(|| {
Error::Protocol("background task completion has no tool_call_id".into())
})?;
let task_identity = agent_id.unwrap_or(&completion.task_id);
let identity = format!("{}:{task_identity}:{tool_call_id}", kind.as_str_name());
let context = completion_context(completion, kind, agent_id)?; let context = completion_context(completion, kind, agent_id)?;
if completions if completions
.insert(identity.clone(), (completion, context)) .insert(identity.clone(), (completion, context))
@@ -307,6 +314,30 @@ mod tests {
assert_eq!(forward.context, reversed.context); assert_eq!(forward.context, reversed.context);
} }
#[test]
fn resumed_subagent_completions_use_the_task_call_as_part_of_their_identity() {
let first = project(
&pb::BackgroundTaskCompletionAction {
completions: vec![completion()],
},
pb::AgentMode::Multitask as i32,
)
.unwrap();
let mut resumed = completion();
resumed.tool_call_id = Some("task-call-2".into());
let second = project(
&pb::BackgroundTaskCompletionAction {
completions: vec![resumed],
},
pb::AgentMode::Multitask as i32,
)
.unwrap();
assert_ne!(first.turn_user.message_id, second.turn_user.message_id);
assert!(first.turn_user.message_id.ends_with(":task-call"));
assert!(second.turn_user.message_id.ends_with(":task-call-2"));
}
#[test] #[test]
fn completion_requires_the_captured_subagent_identity_and_terminal_reason() { fn completion_requires_the_captured_subagent_identity_and_terminal_reason() {
let mut value = completion(); let mut value = completion();
+52 -11
View File
@@ -41,6 +41,7 @@ pub struct CursorRunContext {
pub dynamic_tools: BTreeMap<String, pb::McpToolDefinition>, pub dynamic_tools: BTreeMap<String, pb::McpToolDefinition>,
pub checkpoint_prompt: PromptSpec, pub checkpoint_prompt: PromptSpec,
pub compacting: bool, pub compacting: bool,
pub background_completion: bool,
} }
pub(crate) struct PrepareDependencies<'a> { pub(crate) struct PrepareDependencies<'a> {
@@ -123,7 +124,7 @@ pub(crate) async fn prepare(
mode: mode_number, mode: mode_number,
mut turn_user, mut turn_user,
action_context, action_context,
event_id, mut event_id,
input_id, input_id,
starts_turn, starts_turn,
compacting, compacting,
@@ -172,14 +173,43 @@ pub(crate) async fn prepare(
} }
Some(_) | None => store.ensure_conversation(&conversation_id).await?, Some(_) | None => store.ensure_conversation(&conversation_id).await?,
}; };
let base_revision_id = match input_id { let base_revision_id = match input_id.as_deref() {
Some(input_id) => { Some(input_id) => {
store store
.anchor_input(&conversation_id, &input_id, proposed_base_revision_id) .anchor_input(&conversation_id, input_id, proposed_base_revision_id)
.await? .await?
} }
None => proposed_base_revision_id, None => proposed_base_revision_id,
}; };
let mut projected_user_context = if input_id.is_some() && !compacting && !background_completion
{
runtime::compile_request_context(
"identity",
&request_context,
base_messages.as_deref().unwrap_or_default(),
)?
} else {
None
};
if event_id.is_none() {
if let (Some(input_id), Some(user)) = (input_id.as_deref(), turn_user.as_ref()) {
event_id = Some(
runtime::user_event_id(
input_id,
checkpoint_mode,
user,
&request_context,
&action_context,
projected_user_context
.as_ref()
.map(|message| &message.content),
compiler,
blob_sync,
)
.await?,
);
}
}
let existing_runtime = match event_id.as_deref() { let existing_runtime = match event_id.as_deref() {
Some(event_id) => { Some(event_id) => {
store store
@@ -193,6 +223,10 @@ pub(crate) async fn prepare(
let message_id = format!("request-context:{event_id}"); let message_id = format!("request-context:{event_id}");
match store.message(&conversation_id, &message_id).await? { match store.message(&conversation_id, &message_id).await? {
Some(message) => Some(message), Some(message) => Some(message),
None if input_id.is_some() => projected_user_context.take().map(|mut message| {
message.message_id = message_id;
message
}),
None => runtime::compile_request_context( None => runtime::compile_request_context(
event_id, event_id,
&request_context, &request_context,
@@ -202,7 +236,7 @@ pub(crate) async fn prepare(
} }
_ => None, _ => None,
}; };
let initial_messages = if compacting { let mut initial_messages = if compacting {
Vec::new() Vec::new()
} else { } else {
match (turn_user.clone(), event_id) { match (turn_user.clone(), event_id) {
@@ -256,6 +290,10 @@ pub(crate) async fn prepare(
} }
} }
}; };
let (base_revision_id, reused) = store
.match_revision_prefix(&conversation_id, base_revision_id, &initial_messages)
.await?;
initial_messages.drain(..reused);
let action = if compacting { let action = if compacting {
RunAction::Compact RunAction::Compact
} else if starts_turn { } else if starts_turn {
@@ -310,6 +348,7 @@ pub(crate) async fn prepare(
.collect(), .collect(),
checkpoint_prompt, checkpoint_prompt,
compacting, compacting,
background_completion,
}, },
)) ))
} }
@@ -434,13 +473,13 @@ fn action(request: &pb::AgentRunRequest) -> Result<ActionProjection> {
.filter(|text| !text.is_empty()) .filter(|text| !text.is_empty())
.cloned(), .cloned(),
); );
let event_id = format!("cursor:user:{}", user.message_id); let input_id = format!("cursor:user:{}", user.message_id);
Ok(ActionProjection { Ok(ActionProjection {
mode, mode,
turn_user: Some(user.clone()), turn_user: Some(user.clone()),
action_context: context.join("\n\n"), action_context: context.join("\n\n"),
event_id: Some(event_id.clone()), event_id: None,
input_id: Some(event_id), input_id: Some(input_id),
starts_turn: true, starts_turn: true,
compacting: false, compacting: false,
background_completion: false, background_completion: false,
@@ -703,7 +742,7 @@ mod tests {
} }
#[test] #[test]
fn queued_messages_reusing_a_request_id_keep_distinct_runtime_identities() { fn queued_messages_keep_distinct_input_anchors_until_runtime_identity_is_compiled() {
let request = |message_id: &str| pb::AgentRunRequest { let request = |message_id: &str| pb::AgentRunRequest {
action: Some(pb::ConversationAction { action: Some(pb::ConversationAction {
action: Some(pb::conversation_action::Action::UserMessageAction( action: Some(pb::conversation_action::Action::UserMessageAction(
@@ -725,9 +764,11 @@ mod tests {
let first = action(&request("message-one")).unwrap(); let first = action(&request("message-one")).unwrap();
let second = action(&request("message-two")).unwrap(); let second = action(&request("message-two")).unwrap();
assert_eq!(first.event_id.as_deref(), Some("cursor:user:message-one")); assert_eq!(first.event_id, None);
assert_eq!(second.event_id.as_deref(), Some("cursor:user:message-two")); assert_eq!(second.event_id, None);
assert_ne!(first.event_id, second.event_id); assert_eq!(first.input_id.as_deref(), Some("cursor:user:message-one"));
assert_eq!(second.input_id.as_deref(), Some("cursor:user:message-two"));
assert_ne!(first.input_id, second.input_id);
} }
#[test] #[test]
+58 -3
View File
@@ -10,6 +10,7 @@ use crate::{
proto::agent::v1 as pb, proto::agent::v1 as pb,
}, },
model::{CanonicalMessage, MessageContent, Origin, Role}, model::{CanonicalMessage, MessageContent, Origin, Role},
store::BlobId,
Error, Result, Error, Result,
}; };
@@ -80,12 +81,66 @@ pub async fn compile(
compiler: &PromptCompiler, compiler: &PromptCompiler,
blobs: &BlobSynchronizer, blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> { ) -> Result<CanonicalMessage> {
let time = Time::now( let timestamp = Time::now(
request_context request_context
.env .env
.as_ref() .as_ref()
.map(|env| env.time_zone.as_str()), .map(|env| env.time_zone.as_str()),
)?; )?
.timestamp;
compile_with_timestamp(
event_id,
mode,
user,
request_context,
action_context,
timestamp,
compiler,
blobs,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub(super) async fn user_event_id(
input_id: &str,
mode: Mode,
user: &pb::UserMessage,
request_context: &pb::RequestContext,
action_context: &str,
projected_request_context: Option<&MessageContent>,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<String> {
let runtime = compile_with_timestamp(
"identity".into(),
mode,
user,
request_context,
action_context,
String::new(),
compiler,
blobs,
)
.await?;
let semantic = serde_json::to_vec(&(projected_request_context, runtime.content))?;
Ok(format!(
"{input_id}:{}",
BlobId::digest(&semantic).to_base64()
))
}
#[allow(clippy::too_many_arguments)]
async fn compile_with_timestamp(
event_id: String,
mode: Mode,
user: &pb::UserMessage,
request_context: &pb::RequestContext,
action_context: &str,
timestamp: String,
compiler: &PromptCompiler,
blobs: &BlobSynchronizer,
) -> Result<CanonicalMessage> {
let mut values = BTreeMap::from([ let mut values = BTreeMap::from([
("OPEN_FILES", section(open_files(user))), ("OPEN_FILES", section(open_files(user))),
( (
@@ -98,7 +153,7 @@ pub async fn compile(
), ),
), ),
("ACTION_CONTEXT", section(action_context.to_string())), ("ACTION_CONTEXT", section(action_context.to_string())),
("TIMESTAMP", time.timestamp), ("TIMESTAMP", timestamp),
("USER_QUERY", user.text.clone()), ("USER_QUERY", user.text.clone()),
("DEBUG_SERVER_ENDPOINT", String::new()), ("DEBUG_SERVER_ENDPOINT", String::new()),
("DEBUG_LOG_PATH", String::new()), ("DEBUG_LOG_PATH", String::new()),
+93 -2
View File
@@ -9,7 +9,11 @@ use tokio_stream::StreamExt;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::{ use crate::{
cursor::{connect::END_STREAM_FLAG, observability::CursorTraceRecorder, CursorSessionRegistry}, cursor::{
connect::{self, END_STREAM_FLAG},
observability::CursorTraceRecorder,
CursorSessionRegistry,
},
Result, Result,
}; };
@@ -49,7 +53,7 @@ fn local_body_stream(
trace.chunk(&chunk); trace.chunk(&chunk);
if terminal { if terminal {
guard.complete(); guard.complete();
trace.finish(None); trace.finish(end_stream_error(&chunk));
} }
yield Ok::<Bytes, Infallible>(chunk); yield Ok::<Bytes, Infallible>(chunk);
if terminal { if terminal {
@@ -67,6 +71,30 @@ fn is_end_stream_frame(frame: &Bytes) -> bool {
.is_some_and(|flags| flags & END_STREAM_FLAG != 0) .is_some_and(|flags| flags & END_STREAM_FLAG != 0)
} }
fn end_stream_error(frame: &Bytes) -> Option<String> {
connect::decode_frames(frame)
.ok()?
.into_iter()
.find_map(|(flags, payload)| {
if flags & END_STREAM_FLAG == 0 {
return None;
}
let value = serde_json::from_slice::<serde_json::Value>(&payload).ok()?;
let error = value.get("error")?;
let code = error.get("code").and_then(serde_json::Value::as_str);
let message = error
.get("message")
.and_then(serde_json::Value::as_str)
.filter(|message| !message.is_empty());
Some(match (code, message) {
(Some(code), Some(message)) => format!("{code}: {message}"),
(Some(code), None) => code.to_string(),
(None, Some(message)) => message.to_string(),
(None, None) => error.to_string(),
})
})
}
struct LocalRunGuard { struct LocalRunGuard {
cancellation: CancellationToken, cancellation: CancellationToken,
completed: bool, completed: bool,
@@ -233,4 +261,67 @@ mod tests {
drop(stream); drop(stream);
assert!(!cancellation.is_cancelled()); assert!(!cancellation.is_cancelled());
} }
#[test]
fn connect_error_end_stream_exposes_the_trace_error() {
let frame = connect::encode_error_end_stream(&connect::ConnectStreamError {
code: connect::ConnectCode::InvalidArgument,
message: "unsupported runtime action".into(),
details: Vec::new(),
})
.unwrap();
assert_eq!(
end_stream_error(&frame).as_deref(),
Some("invalid_argument: unsupported runtime action")
);
assert_eq!(end_stream_error(&connect::encode_end_stream()), None);
}
#[tokio::test]
async fn connect_error_end_stream_marks_the_local_trace_as_error() {
let store = crate::store::Store::connect("sqlite::memory:")
.await
.unwrap();
store.set_detailed_logging(true).await.unwrap();
let trace = CursorTraceRecorder::begin(
store.clone(),
"error-trace",
Some("conversation"),
"local_byok",
Some("model"),
)
.await
.unwrap();
let (sender, receiver) = mpsc::unbounded_channel();
let cancellation = CancellationToken::new();
sender
.send(
connect::encode_error_end_stream(&connect::ConnectStreamError {
code: connect::ConnectCode::InvalidArgument,
message: "unsupported runtime action".into(),
details: Vec::new(),
})
.unwrap(),
)
.unwrap();
let mut stream = Box::pin(local_body_stream(receiver, cancellation, Some(trace)));
stream.next().await.unwrap().unwrap();
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(1);
let trace = loop {
let trace = store.cursor_trace("error-trace").await.unwrap().unwrap();
if trace.status != "running" {
break trace;
}
assert!(tokio::time::Instant::now() < deadline);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
};
assert_eq!(trace.status, "error");
assert_eq!(
trace.error_message.as_deref(),
Some("invalid_argument: unsupported runtime action")
);
}
} }
+17
View File
@@ -88,6 +88,23 @@ impl CursorSession {
} }
pub async fn run(mut self) -> Result<()> { pub async fn run(mut self) -> Result<()> {
let result = self.run_inner().await;
if let Err(error) = &result {
self.abort_execs().await;
let error = match error {
Error::Protocol(message) => message.clone(),
error => error.to_string(),
};
let _ = self
.core
.commands
.send(ClientCommand::ClientClosed { error })
.await;
}
result
}
async fn run_inner(&mut self) -> Result<()> {
if self.context.compacting { if self.context.compacting {
self.handle.emit(&interaction::summary_started())?; self.handle.emit(&interaction::summary_started())?;
} }
+1 -1
View File
@@ -5,4 +5,4 @@ pub use request::{abort, mcp_request, mcp_state_request, request};
pub(crate) use request::{ pub(crate) use request::{
await_read_request, edit_read_request, json_object_to_prost, mcp_meta_request, await_read_request, edit_read_request, json_object_to_prost, mcp_meta_request,
}; };
pub use response::{client_event, ClientExecEvent}; pub use response::{client_event, stream_closed, ClientExecEvent};
+47
View File
@@ -130,6 +130,53 @@ pub async fn client_event(
Ok(event) Ok(event)
} }
pub async fn stream_closed(id: u32, pending: &CursorToolRuntime) -> Result<Option<ToolCompletion>> {
let Some(entry) = pending.take_exec(id).await else {
return Ok(None);
};
let error = "Cursor Exec stream closed before returning a terminal result";
if entry.call.name.eq_ignore_ascii_case("Shell") {
let command = entry
.call
.arguments
.get("command")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.to_string();
let working_directory = entry
.call
.arguments
.get("working_directory")
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.to_string();
return Ok(Some(result::from_exec(
entry,
&pb::exec_client_message::Message::ShellResult(pb::ShellResult {
result: Some(pb::shell_result::Result::SpawnError(pb::ShellSpawnError {
command,
working_directory,
error: error.into(),
})),
..Default::default()
}),
)?));
}
let rendered = match &entry.stage {
ExecStage::DynamicMcp(definition) => {
interaction::render_dynamic_mcp(&entry.call, definition, false)
}
_ => interaction::render_tool_call(&entry.call, false)?,
};
Ok(Some(ToolCompletion::from_rendered(
&entry.call,
entry.started_at_ms,
error.into(),
true,
rendered,
)?))
}
async fn advance_await( async fn advance_await(
entry: PendingExec, entry: PendingExec,
result: &pb::exec_client_message::Message, result: &pb::exec_client_message::Message,
+12
View File
@@ -338,6 +338,18 @@ impl CursorToolRuntime {
ids ids
} }
pub async fn running_task_exec_id(&self, call_id: &str) -> Option<u32> {
self.execs
.lock()
.await
.iter()
.filter_map(|(id, entry)| {
(entry.call.call_id == call_id && entry.call.name.eq_ignore_ascii_case("Task"))
.then_some(*id)
})
.min()
}
fn next_id(&self) -> Result<u32> { fn next_id(&self) -> Result<u32> {
self.next_id self.next_id
.fetch_add(1, Ordering::Relaxed) .fetch_add(1, Ordering::Relaxed)
+8 -1
View File
@@ -2,7 +2,12 @@ use std::sync::Arc;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::{client::ClientPort, model::PreparedRun, provider::Provider, store::Store}; use crate::{
client::{ClientCommand, ClientPort},
model::PreparedRun,
provider::Provider,
store::Store,
};
use super::{RunEngine, RunOutcome, RunRegistry}; use super::{RunEngine, RunOutcome, RunRegistry};
@@ -26,6 +31,7 @@ impl RunActor {
&self, &self,
prepared: PreparedRun, prepared: PreparedRun,
client: ClientPort, client: ClientPort,
commands: tokio::sync::mpsc::Sender<ClientCommand>,
cancellation: CancellationToken, cancellation: CancellationToken,
) -> tokio::task::JoinHandle<RunOutcome> { ) -> tokio::task::JoinHandle<RunOutcome> {
let run_id = prepared.run_id.clone(); let run_id = prepared.run_id.clone();
@@ -35,6 +41,7 @@ impl RunActor {
conversation_id.clone(), conversation_id.clone(),
run_id.clone(), run_id.clone(),
cancellation.clone(), cancellation.clone(),
commands,
) )
.await; .await;
let actor = self.clone(); let actor = self.clone();
+81 -13
View File
@@ -4,7 +4,10 @@ use std::sync::Arc;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::{ use crate::{
client::{ClientEvent, ClientPort, CommitBarrier, CommitCause, StateCommitted}, client::{
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
StateCommitted,
},
model::{ model::{
CanonicalMessage, MessageContent, Origin, PreparedRun, Role, RunAction, ToolRoundAssistant, CanonicalMessage, MessageContent, Origin, PreparedRun, Role, RunAction, ToolRoundAssistant,
ToolRoundId, Usage, ToolRoundId, Usage,
@@ -158,6 +161,7 @@ impl RunEngine {
calls: round.calls.clone(), calls: round.calls.clone(),
recovered_started_at_ms: Some(round.started_at_ms), recovered_started_at_ms: Some(round.started_at_ms),
}, },
Vec::new(),
) )
.await .await
{ {
@@ -242,21 +246,27 @@ impl RunEngine {
&cycle_cancellation, &cycle_cancellation,
); );
tokio::pin!(cycle); tokio::pin!(cycle);
let cycle = tokio::select! { let mut pending_insertions = Vec::new();
result = &mut cycle => result, let cycle = loop {
tokio::select! {
result = &mut cycle => break result,
command = client.commands.recv() => { command = client.commands.recv() => {
let message = match command { let message = match command {
Some(crate::client::ClientCommand::RuntimeMessage(message)) => message, Some(ClientCommand::InsertMessages(insertion)) => {
Some(crate::client::ClientCommand::RuntimeEvent(event)) => event.into_message(), pending_insertions.push(insertion);
Some(crate::client::ClientCommand::Cancel) => { continue;
}
Some(ClientCommand::RuntimeMessage(message)) => message,
Some(ClientCommand::RuntimeEvent(event)) => event.into_message(),
Some(ClientCommand::Cancel) => {
cycle_cancellation.cancel(); cycle_cancellation.cancel();
return (RunOutcome::Cancelled, usage); return (RunOutcome::Cancelled, usage);
} }
Some(crate::client::ClientCommand::ClientClosed { error }) => { Some(ClientCommand::ClientClosed { error }) => {
cycle_cancellation.cancel(); cycle_cancellation.cancel();
return (RunOutcome::Failed(RunFailure::Client(error)), usage); return (RunOutcome::Failed(RunFailure::Client(error)), usage);
} }
Some(crate::client::ClientCommand::ToolResult(_)) => { Some(ClientCommand::ToolResult(_)) => {
cycle_cancellation.cancel(); cycle_cancellation.cancel();
return ( return (
RunOutcome::Failed(RunFailure::Protocol( RunOutcome::Failed(RunFailure::Protocol(
@@ -284,6 +294,19 @@ impl RunEngine {
} }
} }
} }
revision = match append_insertions(
&self.store,
prepared,
client,
cancellation,
revision,
std::mem::take(&mut pending_insertions),
)
.await
{
Ok((revision, _)) => revision,
Err(outcome) => return (outcome, usage),
};
revision = match append_runtime_message( revision = match append_runtime_message(
&self.store, &self.store,
prepared, prepared,
@@ -294,11 +317,12 @@ impl RunEngine {
) )
.await .await
{ {
Ok(revision) => revision, Ok((revision, _)) => revision,
Err(outcome) => return (outcome, usage), Err(outcome) => return (outcome, usage),
}; };
continue 'model; continue 'model;
} }
}
}; };
let cycle = match cycle { let cycle = match cycle {
Ok(cycle) => cycle, Ok(cycle) => cycle,
@@ -413,6 +437,27 @@ impl RunEngine {
Ok(revision) => revision, Ok(revision) => revision,
Err(error) => return (RunOutcome::Failed(error.into()), usage), Err(error) => return (RunOutcome::Failed(error.into()), usage),
}; };
if !pending_insertions.is_empty() {
let inserted = match append_insertions(
&self.store,
prepared,
client,
cancellation,
revision,
pending_insertions,
)
.await
{
Ok((next, inserted)) => {
revision = next;
inserted
}
Err(outcome) => return (outcome, usage),
};
if inserted {
continue 'model;
}
}
let (barrier, ready) = CommitBarrier::before_continue(); let (barrier, ready) = CommitBarrier::before_continue();
if emit( if emit(
client, client,
@@ -453,6 +498,7 @@ impl RunEngine {
calls: cycle.calls, calls: cycle.calls,
recovered_started_at_ms: None, recovered_started_at_ms: None,
}, },
pending_insertions,
) )
.await .await
{ {
@@ -683,14 +729,36 @@ fn fallback_summary(messages: &[CanonicalMessage]) -> String {
) )
} }
async fn append_runtime_message( pub(super) async fn append_insertions(
store: &Store,
prepared: &PreparedRun,
client: &mut ClientPort,
cancellation: &CancellationToken,
mut revision: crate::model::RevisionId,
insertions: Vec<MessageInsertion>,
) -> std::result::Result<(crate::model::RevisionId, bool), RunOutcome> {
let mut inserted_any = false;
for insertion in insertions {
for message in insertion.messages {
let (next, inserted) =
append_runtime_message(store, prepared, client, cancellation, revision, message)
.await?;
revision = next;
inserted_any |= inserted;
}
let _ = insertion.delivered.send(());
}
Ok((revision, inserted_any))
}
pub(super) async fn append_runtime_message(
store: &Store, store: &Store,
prepared: &PreparedRun, prepared: &PreparedRun,
client: &mut ClientPort, client: &mut ClientPort,
cancellation: &CancellationToken, cancellation: &CancellationToken,
revision: crate::model::RevisionId, revision: crate::model::RevisionId,
message: CanonicalMessage, message: CanonicalMessage,
) -> std::result::Result<crate::model::RevisionId, RunOutcome> { ) -> std::result::Result<(crate::model::RevisionId, bool), RunOutcome> {
let event_id = message.runtime_event_id.clone().ok_or_else(|| { let event_id = message.runtime_event_id.clone().ok_or_else(|| {
RunOutcome::Failed(RunFailure::Protocol( RunOutcome::Failed(RunFailure::Protocol(
"runtime message has no event identity".into(), "runtime message has no event identity".into(),
@@ -706,7 +774,7 @@ async fn append_runtime_message(
.await .await
.map_err(|error| RunOutcome::Failed(error.into()))?; .map_err(|error| RunOutcome::Failed(error.into()))?;
if !inserted { if !inserted {
return Ok(revision); return Ok((revision, false));
} }
let (barrier, ready) = CommitBarrier::before_continue(); let (barrier, ready) = CommitBarrier::before_continue();
emit( emit(
@@ -721,7 +789,7 @@ async fn append_runtime_message(
.await .await
.map_err(|_| client_failure())?; .map_err(|_| client_failure())?;
wait_for_state_ready(ready, cancellation).await?; wait_for_state_ready(ready, cancellation).await?;
Ok(revision) Ok((revision, true))
} }
async fn hydrate_tool_images( async fn hydrate_tool_images(
+38 -1
View File
@@ -3,7 +3,10 @@ use std::{collections::HashMap, sync::Arc};
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::model::{ConversationId, RunId}; use crate::{
client::{ClientCommand, MessageInsertion},
model::{CanonicalMessage, ConversationId, RunId},
};
#[derive(Clone, Default)] #[derive(Clone, Default)]
pub struct RunRegistry { pub struct RunRegistry {
@@ -13,6 +16,7 @@ pub struct RunRegistry {
struct ActiveRun { struct ActiveRun {
run_id: RunId, run_id: RunId,
cancellation: CancellationToken, cancellation: CancellationToken,
commands: tokio::sync::mpsc::Sender<ClientCommand>,
} }
impl RunRegistry { impl RunRegistry {
@@ -21,12 +25,14 @@ impl RunRegistry {
conversation_id: ConversationId, conversation_id: ConversationId,
run_id: RunId, run_id: RunId,
cancellation: CancellationToken, cancellation: CancellationToken,
commands: tokio::sync::mpsc::Sender<ClientCommand>,
) { ) {
let previous = self.active.lock().await.insert( let previous = self.active.lock().await.insert(
conversation_id, conversation_id,
ActiveRun { ActiveRun {
run_id: run_id.clone(), run_id: run_id.clone(),
cancellation, cancellation,
commands,
}, },
); );
if let Some(previous) = previous.filter(|previous| previous.run_id != run_id) { if let Some(previous) = previous.filter(|previous| previous.run_id != run_id) {
@@ -34,6 +40,37 @@ impl RunRegistry {
} }
} }
pub async fn insert_messages(
&self,
conversation_id: &ConversationId,
messages: Vec<CanonicalMessage>,
) -> bool {
if messages.is_empty() {
return true;
}
let commands = self
.active
.lock()
.await
.get(conversation_id)
.map(|run| run.commands.clone());
let Some(commands) = commands else {
return false;
};
let (delivered, delivery) = tokio::sync::oneshot::channel();
if commands
.send(ClientCommand::InsertMessages(MessageInsertion {
messages,
delivered,
}))
.await
.is_err()
{
return false;
}
delivery.await.is_ok()
}
pub async fn release(&self, conversation_id: &ConversationId, run_id: &RunId) { pub async fn release(&self, conversation_id: &ConversationId, run_id: &RunId) {
let mut active = self.active.lock().await; let mut active = self.active.lock().await;
if active if active
+45 -33
View File
@@ -1,7 +1,10 @@
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use crate::{ use crate::{
client::{ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, StateCommitted}, client::{
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
StateCommitted,
},
model::{PreparedRun, RevisionId, ToolCall, ToolRoundAssistant, ToolRoundId}, model::{PreparedRun, RevisionId, ToolCall, ToolRoundAssistant, ToolRoundId},
store::Store, store::Store,
}; };
@@ -22,6 +25,7 @@ pub(super) async fn execute(
cancellation: &CancellationToken, cancellation: &CancellationToken,
mut revision: RevisionId, mut revision: RevisionId,
round: ToolRound, round: ToolRound,
insertions: Vec<MessageInsertion>,
) -> std::result::Result<RevisionId, RunOutcome> { ) -> std::result::Result<RevisionId, RunOutcome> {
let ToolRound { let ToolRound {
id: round_id, id: round_id,
@@ -66,7 +70,10 @@ pub(super) async fn execute(
.await?; .await?;
let mut remaining = calls.len(); let mut remaining = calls.len();
let mut pending_runtime_messages = Vec::new(); let mut pending_runtime_messages = insertions
.into_iter()
.map(PendingRuntimeMessage::Insertion)
.collect::<Vec<_>>();
while remaining > 0 { while remaining > 0 {
let command = tokio::select! { let command = tokio::select! {
_ = cancellation.cancelled() => return Err(RunOutcome::Cancelled), _ = cancellation.cancelled() => return Err(RunOutcome::Cancelled),
@@ -116,10 +123,13 @@ pub(super) async fn execute(
} }
} }
Some(ClientCommand::RuntimeEvent(event)) => { Some(ClientCommand::RuntimeEvent(event)) => {
pending_runtime_messages.push(event.into_message()); pending_runtime_messages.push(PendingRuntimeMessage::Message(event.into_message()));
} }
Some(ClientCommand::RuntimeMessage(message)) => { Some(ClientCommand::RuntimeMessage(message)) => {
pending_runtime_messages.push(message); pending_runtime_messages.push(PendingRuntimeMessage::Message(message));
}
Some(ClientCommand::InsertMessages(insertion)) => {
pending_runtime_messages.push(PendingRuntimeMessage::Insertion(insertion))
} }
Some(ClientCommand::Cancel) => return Err(RunOutcome::Cancelled), Some(ClientCommand::Cancel) => return Err(RunOutcome::Cancelled),
Some(ClientCommand::ClientClosed { error }) => { Some(ClientCommand::ClientClosed { error }) => {
@@ -128,40 +138,42 @@ pub(super) async fn execute(
None => return Err(client_failure()), None => return Err(client_failure()),
} }
} }
for message in pending_runtime_messages { for pending in pending_runtime_messages {
let event_id = message.runtime_event_id.clone().ok_or_else(|| { match pending {
RunOutcome::Failed(RunFailure::Protocol( PendingRuntimeMessage::Message(message) => {
"runtime message has no event identity".into(), revision = super::engine::append_runtime_message(
)) store,
})?; prepared,
let (next, inserted) = store client,
.append_message_once( cancellation,
&prepared.conversation_id, revision,
&prepared.run_id, message,
revision, )
&message, .await?
) .0;
.await }
.map_err(failed)?; PendingRuntimeMessage::Insertion(insertion) => {
revision = next; revision = super::engine::append_insertions(
if inserted { store,
let (barrier, ready) = CommitBarrier::before_continue(); prepared,
send( client,
client, cancellation,
ClientEvent::StateCommitted(StateCommitted { revision,
revision_id: revision, vec![insertion],
tool_round_version: 0, )
cause: CommitCause::RuntimeEvent { event_id }, .await?
barrier, .0;
}), }
)
.await?;
super::engine::wait_for_state_ready(ready, cancellation).await?;
} }
} }
Ok(revision) Ok(revision)
} }
enum PendingRuntimeMessage {
Message(crate::model::CanonicalMessage),
Insertion(MessageInsertion),
}
async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> { async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> {
client client
.events .events
+32
View File
@@ -58,6 +58,38 @@ impl Store {
self.load_revision_messages(RevisionId(revision_id)).await self.load_revision_messages(RevisionId(revision_id)).await
} }
pub async fn match_revision_prefix(
&self,
conversation_id: &ConversationId,
base_revision_id: RevisionId,
additions: &[CanonicalMessage],
) -> Result<(RevisionId, usize)> {
let mut revision = base_revision_id;
let mut messages = self.load_revision_messages(revision).await?;
for (index, addition) in additions.iter().enumerate() {
messages.push(addition.clone());
let digest = message_digest(&messages)?;
let child = sqlx::query_scalar::<_, i64>(
"SELECT revision_id FROM conversation_revisions
WHERE conversation_id = ? AND parent_revision_id = ? AND state_digest = ?",
)
.bind(conversation_id.as_str())
.bind(revision.0)
.bind(digest.as_slice())
.fetch_optional(&self.pool)
.await?
.map(RevisionId);
let Some(child) = child else {
return Ok((revision, index));
};
if self.load_revision_messages(child).await? != messages {
return Ok((revision, index));
}
revision = child;
}
Ok((revision, additions.len()))
}
pub async fn import_revision( pub async fn import_revision(
&self, &self,
conversation_id: &ConversationId, conversation_id: &ConversationId,
+164 -10
View File
@@ -76,7 +76,7 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() {
.unwrap(); .unwrap();
assert!(messages.iter().any(|message| { assert!(messages.iter().any(|message| {
message.runtime_event_id.as_deref() message.runtime_event_id.as_deref()
== Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id") == Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id:task-call")
&& matches!(&message.content, MessageContent::Parts { parts } if !parts.is_empty()) && matches!(&message.content, MessageContent::Parts { parts } if !parts.is_empty())
})); }));
@@ -131,12 +131,78 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() {
assert_eq!( assert_eq!(
runtime_ids, runtime_ids,
[ [
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id", "runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id:task-call",
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id-2" "runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id-2:task-call"
] ]
); );
} }
#[tokio::test]
async fn background_completion_joins_the_active_run_instead_of_replacing_it() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
let first_ready = provider.push_gated(stop_response("model-call-1", "first response"));
provider.push(stop_response("model-call-2", "processed both completions"));
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let first = registry.get_or_create("active-completion-1").await.unwrap();
let first_run = tokio::spawn(async move {
drive_completion(
&first,
completion_run(
"child-1",
"parent-run-1",
pb::ConversationStateStructure::default(),
),
)
.await
});
while provider.requests().is_empty() {
tokio::task::yield_now().await;
}
let second = registry.get_or_create("active-completion-2").await.unwrap();
let second_run = tokio::spawn(async move {
drive_forwarded_completion(
&second,
completion_run(
"child-2",
"parent-run-2",
pb::ConversationStateStructure::default(),
),
)
.await
});
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
first_ready.notify_one();
second_run.await.unwrap();
first_run.await.unwrap();
let requests = provider.requests();
assert_eq!(requests.len(), 2);
let history = serde_json::to_string(&requests[1].history).unwrap();
assert!(history.contains("child-1"));
assert!(history.contains("first response"));
assert!(history.contains("child-2"));
let statuses: Vec<String> = sqlx::query_scalar(
"SELECT status FROM runs WHERE conversation_id = 'parent-conversation' ORDER BY created_at_ms",
)
.fetch_all(store.pool())
.await
.unwrap();
assert_eq!(statuses, ["completed"]);
}
#[tokio::test] #[tokio::test]
async fn retrying_one_background_completion_reuses_its_runtime_message() { async fn retrying_one_background_completion_reuses_its_runtime_message() {
let (_directory, store) = fixtures::temp_store().await; let (_directory, store) = fixtures::temp_store().await;
@@ -172,12 +238,7 @@ async fn retrying_one_background_completion_reuses_its_runtime_message() {
let second = registry.get_or_create("completion-retry-2").await.unwrap(); let second = registry.get_or_create("completion-retry-2").await.unwrap();
drive_completion( drive_completion(
&second, &second,
completion_run_with_detail( completion_run("retry-child", "completion-retry-run-2", checkpoint),
"retry-child",
"completion-retry-run-2",
checkpoint,
"updated retry payload",
),
) )
.await; .await;
@@ -192,7 +253,9 @@ async fn retrying_one_background_completion_reuses_its_runtime_message() {
.iter() .iter()
.filter(|message| { .filter(|message| {
message.runtime_event_id.as_deref() message.runtime_event_id.as_deref()
== Some("background-completed:BACKGROUND_TASK_KIND_SUBAGENT:retry-child") == Some(
"background-completed:BACKGROUND_TASK_KIND_SUBAGENT:retry-child:task-call",
)
}) })
.count(), .count(),
1 1
@@ -397,6 +460,97 @@ async fn drive_completion(
) )
} }
async fn drive_forwarded_completion(handle: &CursorSessionHandle, message: pb::AgentClientMessage) {
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(message),
})
.await
.unwrap();
let mut append_seqno = 1;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.unwrap();
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
return;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
assert_eq!(exec.id, 0);
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(
pb::agent_client_message::Message::ExecClientControlMessage(
pb::ExecClientControlMessage {
message: Some(
pb::exec_client_control_message::Message::StreamClose(
pb::ExecClientStreamClose { id: 0 },
),
),
},
),
),
}),
})
.await
.unwrap();
append_seqno += 1;
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id: 0,
message: Some(
pb::exec_client_message::Message::RequestContextResult(
pb::RequestContextResult {
result: Some(
pb::request_context_result::Result::Success(
pb::RequestContextSuccess {
request_context: Some(
pb::RequestContext::default(),
),
..Default::default()
},
),
),
},
),
),
..Default::default()
},
)),
}),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
_ => {}
}
}
}
fn completion_run( fn completion_run(
child_id: &str, child_id: &str,
run_id: &str, run_id: &str,
+81
View File
@@ -14,6 +14,7 @@ use cursor_server::{
provider::{FinishReason, ModelEvent}, provider::{FinishReason, ModelEvent},
run::{RunEngine, RunOutcome}, run::{RunEngine, RunOutcome},
}; };
use tokio::{sync::oneshot, time::Duration};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
#[tokio::test] #[tokio::test]
@@ -54,6 +55,86 @@ async fn a_client_without_checkpoint_protocol_runs_the_same_text_loop() {
assert_eq!(run.await.unwrap(), RunOutcome::Completed); assert_eq!(run.await.unwrap(), RunOutcome::Completed);
} }
#[tokio::test]
async fn inserted_messages_wait_for_the_next_model_call_without_interrupting_the_active_call() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
let first_ready = provider.push_gated(vec![
ModelEvent::Start {
model_call_id: "call-1".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("first answer".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "call-2".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("followed up".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let prepared = prepared(&store).await;
let (port, mut client) = session(32);
let commands = client.commands.clone();
let engine = RunEngine::new(store, Arc::new(provider.clone()));
let run =
tokio::spawn(async move { engine.run(prepared, port, CancellationToken::new()).await });
while provider.requests().is_empty() {
if let Ok(Some(ClientEvent::StateCommitted(state))) =
tokio::time::timeout(Duration::from_millis(20), client.events.recv()).await
{
state.barrier.complete(Ok(()));
}
}
let (delivered, mut delivery) = oneshot::channel();
commands
.send(ClientCommand::InsertMessages(
cursor_server::client::MessageInsertion {
messages: vec![cursor_server::model::RuntimeEvent {
event_id: "background:finished".into(),
text: "background work finished".into(),
}
.into_message()],
delivered,
},
))
.await
.unwrap();
assert!(
tokio::time::timeout(Duration::from_millis(20), &mut delivery)
.await
.is_err()
);
first_ready.notify_one();
while let Some(event) = client.events.recv().await {
match event {
ClientEvent::StateCommitted(state) => state.barrier.complete(Ok(())),
ClientEvent::Ended(outcome) => {
assert_eq!(outcome, RunOutcome::Completed);
break;
}
_ => {}
}
}
assert_eq!(run.await.unwrap(), RunOutcome::Completed);
delivery.await.unwrap();
let requests = provider.requests();
assert_eq!(requests.len(), 2);
let history = &requests[1].history;
assert!(matches!(
history[1].role,
cursor_server::model::Role::Assistant
));
assert_eq!(history[2].message_id, "runtime:background:finished");
}
#[tokio::test] #[tokio::test]
async fn a_failed_claim_cannot_overwrite_the_existing_run() { async fn a_failed_claim_cannot_overwrite_the_existing_run() {
let (_directory, store) = fixtures::temp_store().await; let (_directory, store) = fixtures::temp_store().await;
+24 -1
View File
@@ -160,7 +160,7 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
) )
.unwrap(); .unwrap();
let registry = CursorSessionRegistry::new( let registry = CursorSessionRegistry::new(
store, store.clone(),
Arc::new(provider), Arc::new(provider),
PromptCompiler::new(assets), PromptCompiler::new(assets),
Default::default(), Default::default(),
@@ -249,6 +249,29 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
.unwrap(), .unwrap(),
None None
); );
let deadline = tokio::time::Instant::now() + std::time::Duration::from_secs(1);
let (status, failure_summary) = loop {
let row: (String, Option<String>) =
sqlx::query_as("SELECT status, failure_summary FROM runs WHERE cursor_request_id = ?")
.bind("protocol-failed-request")
.fetch_one(store.pool())
.await
.unwrap();
if row.0 != "running" {
break row;
}
assert!(
tokio::time::Instant::now() < deadline,
"Run remained running after the Cursor session failed"
);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
};
assert_eq!(status, "failed");
assert_eq!(
failure_summary.as_deref(),
Some("unknown ExecClientMessage id: 1001")
);
} }
#[tokio::test] #[tokio::test]
+217
View File
@@ -28,6 +28,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
conversation.clone(), conversation.clone(),
cursor_server::model::RunId::new("first"), cursor_server::model::RunId::new("first"),
first.clone(), first.clone(),
cursor_server::client::session(1).1.commands,
) )
.await; .await;
registry registry
@@ -35,6 +36,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
conversation.clone(), conversation.clone(),
cursor_server::model::RunId::new("second"), cursor_server::model::RunId::new("second"),
second.clone(), second.clone(),
cursor_server::client::session(1).1.commands,
) )
.await; .await;
@@ -426,6 +428,187 @@ async fn injected_user_context_restarts_only_the_active_model_cycle() {
); );
} }
#[tokio::test]
async fn cancel_subagent_action_aborts_the_target_task_and_keeps_the_parent_running() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
provider.push(vec![
ModelEvent::Start {
model_call_id: "task-cycle".into(),
},
ModelEvent::ToolCallStart {
index: 0,
call_id: "task-call".into(),
name: "Task".into(),
},
ModelEvent::ToolCallArgumentsDelta {
index: 0,
delta: serde_json::json!({
"description": "Inspect protocol",
"prompt": "Inspect the protocol",
"subagent_type": "generalPurpose",
"run_in_background": false
})
.to_string(),
},
ModelEvent::ToolCallEnd { index: 0 },
ModelEvent::Done(FinishReason::ToolUse),
]);
provider.push(vec![
ModelEvent::Start {
model_call_id: "continued".into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta("continued after subagent cancellation".into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store,
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
let handle = registry
.get_or_create("cancel-subagent-request")
.await
.unwrap();
let mut output = handle.subscribe();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(client_run_for(
"cancel-subagent-request",
"cancel-subagent-conversation",
)),
})
.await
.unwrap();
let mut append_seqno = 1;
let exec_id = loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Task exec");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => {
let Some(pb::exec_server_message::Message::SubagentArgs(args)) = exec.message
else {
continue;
};
assert_eq!(args.tool_call_id, "task-call");
break exec.id;
}
_ => {}
}
};
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(runtime_cancel_subagent("task-call")),
})
.await
.unwrap();
append_seqno += 1;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before Task abort");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
assert_eq!(flags & connect::END_STREAM_FLAG, 0);
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::ExecServerControlMessage(control)) => {
let Some(pb::exec_server_control_message::Message::Abort(abort)) = control.message
else {
continue;
};
assert_eq!(abort.id, exec_id);
break;
}
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
_ => {}
}
}
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(subagent_aborted(exec_id)),
})
.await
.unwrap();
append_seqno += 1;
let mut saw_continued = false;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("RunSSE closed before successful EndStream");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
assert_eq!(payload.as_ref(), b"{}");
break;
}
let server = pb::AgentServerMessage::decode(payload).unwrap();
match server.message {
Some(pb::agent_server_message::Message::KvServerMessage(kv)) => {
handle
.command(CursorCommand::Append {
seqno: append_seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
append_seqno += 1;
}
Some(pb::agent_server_message::Message::InteractionUpdate(update)) => {
if let Some(pb::interaction_update::Message::TextDelta(delta)) = update.message {
saw_continued |= delta.text.contains("continued after subagent cancellation");
}
}
_ => {}
}
}
assert!(saw_continued);
assert!(!handle.cancellation().is_cancelled());
assert_eq!(provider.requests().len(), 2);
}
fn client_run() -> pb::AgentClientMessage { fn client_run() -> pb::AgentClientMessage {
client_run_for("cancel-request", "cancel-conversation") client_run_for("cancel-request", "cancel-conversation")
} }
@@ -541,3 +724,37 @@ fn runtime_injection() -> pb::AgentClientMessage {
)), )),
} }
} }
fn runtime_cancel_subagent(tool_call_id: &str) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ConversationAction(
pb::ConversationAction {
action: Some(pb::conversation_action::Action::CancelSubagentAction(
pb::CancelSubagentAction {
subagent_id: tool_call_id.into(),
},
)),
..Default::default()
},
)),
}
}
fn subagent_aborted(id: u32) -> pb::AgentClientMessage {
pb::AgentClientMessage {
message: Some(pb::agent_client_message::Message::ExecClientMessage(
pb::ExecClientMessage {
id,
message: Some(pb::exec_client_message::Message::SubagentResult(
pb::SubagentResult {
result: Some(pb::subagent_result::Result::Error(pb::SubagentError {
agent_id: None,
error: "Subagent was aborted by the user".into(),
})),
},
)),
..Default::default()
},
)),
}
}
+92
View File
@@ -218,3 +218,95 @@ async fn editing_a_logical_input_discards_its_active_suffix() {
vec![original, suffix] vec![original, suffix]
); );
} }
#[tokio::test]
async fn retry_reuses_only_the_matching_initial_child_chain() {
let (_directory, store) = fixtures::temp_store().await;
let conversation_id = ConversationId::new("retry-conversation");
let root = store.ensure_conversation(&conversation_id).await.unwrap();
let first = prepared("first-run", &conversation_id, root);
store.claim_run(&first).await.unwrap();
let context = fixtures::user("request-context:event", "context");
let context_revision = store
.append_revision(
&conversation_id,
&first.run_id,
root,
std::slice::from_ref(&context),
)
.await
.unwrap();
let runtime = cursor_server::model::RuntimeEvent {
event_id: "cursor:user:stable-id:version".into(),
text: "query".into(),
}
.into_message();
let (partial_revision, partial_count) = store
.match_revision_prefix(&conversation_id, root, &[context.clone(), runtime.clone()])
.await
.unwrap();
assert_eq!(partial_revision, context_revision);
assert_eq!(partial_count, 1);
let runtime_revision = store
.append_revision(
&conversation_id,
&first.run_id,
context_revision,
std::slice::from_ref(&runtime),
)
.await
.unwrap();
let suffix = fixtures::user("old-answer", "old answer");
let old_head = store
.append_revision(
&conversation_id,
&first.run_id,
runtime_revision,
std::slice::from_ref(&suffix),
)
.await
.unwrap();
let (retry_base, reused) = store
.match_revision_prefix(&conversation_id, root, &[context.clone(), runtime.clone()])
.await
.unwrap();
assert_eq!(retry_base, runtime_revision);
assert_eq!(reused, 2);
assert_eq!(
store.load_revision_messages(retry_base).await.unwrap(),
vec![context.clone(), runtime.clone()]
);
let retry = prepared("retry-run", &conversation_id, retry_base);
let claimed = store.claim_run(&retry).await.unwrap();
assert_eq!(claimed.replaced_run_id.as_ref(), Some(&first.run_id));
let first_status: String = sqlx::query_scalar("SELECT status FROM runs WHERE run_id = ?")
.bind(first.run_id.as_str())
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(first_status, "cancelled");
let changed = cursor_server::model::RuntimeEvent {
event_id: "cursor:user:stable-id:changed-version".into(),
text: "edited query".into(),
}
.into_message();
let (changed_base, reused) = store
.match_revision_prefix(&conversation_id, root, &[context, changed])
.await
.unwrap();
assert_eq!(changed_base, context_revision);
assert_eq!(reused, 1);
assert_eq!(
store.load_revision_messages(old_head).await.unwrap(),
vec![
fixtures::user("request-context:event", "context"),
runtime,
suffix
]
);
}
+147 -4
View File
@@ -18,6 +18,150 @@ use cursor_server::{
}; };
use prost::Message; use prost::Message;
#[tokio::test]
async fn retrying_the_same_edited_input_reuses_its_initial_branch() {
let (_directory, store) = fixtures::temp_store().await;
let provider = fake_provider::FakeProvider::default();
for (call, answer) in [
("model", "answer"),
("model-retry", "retry answer"),
("model-edit", "edited answer"),
("model-context-edit", "context edited answer"),
] {
provider.push(vec![
ModelEvent::Start {
model_call_id: call.into(),
},
ModelEvent::TextStart,
ModelEvent::TextDelta(answer.into()),
ModelEvent::TextEnd,
ModelEvent::Done(FinishReason::Stop),
]);
}
let assets = PromptAssets::load(
std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("prompt/cursor")
.as_path(),
)
.unwrap();
let registry = CursorSessionRegistry::new(
store.clone(),
Arc::new(provider.clone()),
PromptCompiler::new(assets),
Default::default(),
);
for (request_id, text, visible_file) in [
("edited-input", "explain this", "/workspace/src/main.rs"),
(
"edited-input-retry",
"explain this",
"/workspace/src/main.rs",
),
(
"edited-input-changed",
"explain the edited version",
"/workspace/src/main.rs",
),
(
"edited-input-context-changed",
"explain this",
"/workspace/src/edited.rs",
),
] {
let handle = registry.get_or_create(request_id).await.unwrap();
let mut output = handle.subscribe();
let mut request = run_request(references(&store).await);
let Some(pb::agent_client_message::Message::RunRequest(run)) = request.message.as_mut()
else {
unreachable!("run_request always returns a RunRequest")
};
let Some(pb::conversation_action::Action::UserMessageAction(action)) = run
.action
.as_mut()
.and_then(|action| action.action.as_mut())
else {
unreachable!("run_request always contains a UserMessageAction")
};
action
.user_message
.as_mut()
.expect("run_request always contains a UserMessage")
.text = text.into();
let user = action
.user_message
.as_mut()
.expect("run_request always contains a UserMessage");
let Some(pb::invocation_context::Data::IdeState(ide)) = user
.selected_context
.as_mut()
.and_then(|selected| selected.invocation_context.as_mut())
.and_then(|invocation| invocation.data.as_mut())
else {
unreachable!("run_request always contains IDE state")
};
ide.visible_files[0].path = visible_file.into();
handle
.command(CursorCommand::Append {
seqno: 0,
message: Box::new(request),
})
.await
.unwrap();
let mut seqno = 1;
loop {
let frame = tokio::time::timeout(std::time::Duration::from_secs(5), output.recv())
.await
.unwrap()
.expect("retry must finish without closing the stream early");
let (flags, payload) = connect::decode_frames(&frame).unwrap().pop().unwrap();
if flags & connect::END_STREAM_FLAG != 0 {
let end = serde_json::from_slice::<serde_json::Value>(&payload).unwrap();
assert_eq!(end, serde_json::json!({}));
break;
}
let message = pb::AgentServerMessage::decode(payload).unwrap();
if let Some(pb::agent_server_message::Message::KvServerMessage(kv)) = message.message {
handle
.command(CursorCommand::Append {
seqno,
message: Box::new(kv_ack(kv.id)),
})
.await
.unwrap();
seqno += 1;
}
}
}
let requests = provider.requests();
assert_eq!(requests.len(), 4);
assert_eq!(requests[1].history, requests[0].history);
assert_eq!(requests[2].history.len(), requests[0].history.len());
assert_ne!(
requests[2].history.last().unwrap().message_id,
requests[0].history.last().unwrap().message_id
);
let ProjectedContent::Parts(parts) = &requests[2].history.last().unwrap().content else {
panic!("edited runtime message must use typed parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("this fixture has no images")
};
assert!(text.contains("<user_query>\nexplain the edited version\n</user_query>"));
assert_ne!(
requests[3].history.last().unwrap().message_id,
requests[0].history.last().unwrap().message_id
);
let ProjectedContent::Parts(parts) = &requests[3].history.last().unwrap().content else {
panic!("context-edited runtime message must use typed parts")
};
let [ContentPart::Text { text }] = parts.as_slice() else {
panic!("this fixture has no images")
};
assert!(text.contains("/workspace/src/edited.rs"));
}
#[tokio::test] #[tokio::test]
async fn unchanged_request_context_is_not_repeated_and_preserves_the_provider_prefix() { async fn unchanged_request_context_is_not_repeated_and_preserves_the_provider_prefix() {
let (_directory, store) = fixtures::temp_store().await; let (_directory, store) = fixtures::temp_store().await;
@@ -115,10 +259,9 @@ async fn unchanged_request_context_is_not_repeated_and_preserves_the_provider_pr
let [ContentPart::Text { text: context_text }] = context_parts.as_slice() else { let [ContentPart::Text { text: context_text }] = context_parts.as_slice() else {
panic!("request context message must contain one text part") panic!("request context message must contain one text part")
}; };
assert_eq!( assert!(request.history[1]
request.history[1].message_id, .message_id
"runtime:cursor:user:wire-user" .starts_with("runtime:cursor:user:wire-user:"));
);
assert!(!request.prompt.instructions.contains("workspace rule")); assert!(!request.prompt.instructions.contains("workspace rule"));
assert!(!request.prompt.instructions.contains("<mcp_meta_tools>")); assert!(!request.prompt.instructions.contains("<mcp_meta_tools>"));
let ProjectedContent::Parts(parts) = &request.history[1].content else { let ProjectedContent::Parts(parts) = &request.history[1].content else {
+9 -2
View File
@@ -89,7 +89,11 @@ async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_che
let user = requests[0] let user = requests[0]
.history .history
.iter() .iter()
.find(|message| message.message_id == "runtime:cursor:user:image-user") .find(|message| {
message
.message_id
.starts_with("runtime:cursor:user:image-user:")
})
.unwrap(); .unwrap();
let ProjectedContent::Parts(parts) = &user.content else { let ProjectedContent::Parts(parts) = &user.content else {
panic!("runtime user message must retain typed parts") panic!("runtime user message must retain typed parts")
@@ -128,7 +132,10 @@ async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_che
let id = BlobId::from_bytes(raw_id).unwrap(); let id = BlobId::from_bytes(raw_id).unwrap();
let bytes = store.get_blob(&id).await.unwrap().unwrap(); let bytes = store.get_blob(&id).await.unwrap().unwrap();
let value: serde_json::Value = serde_json::from_slice(&bytes).unwrap(); let value: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
if value["id"] == "runtime:cursor:user:image-user" { if value["id"]
.as_str()
.is_some_and(|id| id.starts_with("runtime:cursor:user:image-user:"))
{
user_root = Some(value); user_root = Some(value);
break; break;
} }
+23 -1
View File
@@ -10,11 +10,15 @@ use cursor_server::{
provider::{ModelEvent, Provider, ProviderStream}, provider::{ModelEvent, Provider, ProviderStream},
Error, Error,
}; };
use futures_util::stream; use futures_util::{stream, StreamExt};
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
enum FakeResponse { enum FakeResponse {
Events(Vec<Result<ModelEvent, Error>>), Events(Vec<Result<ModelEvent, Error>>),
Gated {
ready: Arc<tokio::sync::Notify>,
events: Vec<Result<ModelEvent, Error>>,
},
Pending, Pending,
} }
@@ -43,6 +47,17 @@ impl FakeProvider {
.unwrap() .unwrap()
.push_back(FakeResponse::Pending); .push_back(FakeResponse::Pending);
} }
pub fn push_gated(&self, events: Vec<ModelEvent>) -> Arc<tokio::sync::Notify> {
let ready = Arc::new(tokio::sync::Notify::new());
self.responses
.lock()
.unwrap()
.push_back(FakeResponse::Gated {
ready: ready.clone(),
events: events.into_iter().map(Ok).collect(),
});
ready
}
pub fn requests(&self) -> Vec<ModelRequest> { pub fn requests(&self) -> Vec<ModelRequest> {
self.requests.lock().unwrap().clone() self.requests.lock().unwrap().clone()
} }
@@ -63,6 +78,13 @@ impl Provider for FakeProvider {
.expect("fake response configured"); .expect("fake response configured");
match events { match events {
FakeResponse::Events(events) => Box::pin(stream::iter(events)), FakeResponse::Events(events) => Box::pin(stream::iter(events)),
FakeResponse::Gated { ready, events } => Box::pin(
stream::once(async move {
ready.notified().await;
events
})
.flat_map(stream::iter),
),
FakeResponse::Pending => Box::pin(stream::pending()), FakeResponse::Pending => Box::pin(stream::pending()),
} }
} }
+3 -1
View File
@@ -283,7 +283,9 @@ async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
.unwrap(); .unwrap();
assert!(messages[0].message_id.starts_with("request-context:")); assert!(messages[0].message_id.starts_with("request-context:"));
assert_eq!(messages[0].role, Role::User); assert_eq!(messages[0].role, Role::User);
assert_eq!(messages[1].message_id, "runtime:cursor:user:user"); assert!(messages[1]
.message_id
.starts_with("runtime:cursor:user:user:"));
assert_eq!(messages[1].role, Role::User); assert_eq!(messages[1].role, Role::User);
assert_eq!( assert_eq!(
messages.len(), messages.len(),
+30
View File
@@ -563,6 +563,36 @@ async fn empty_exec_client_message_is_not_a_terminal_result() {
); );
} }
#[tokio::test]
async fn exec_stream_close_without_a_terminal_result_becomes_a_tool_error() {
let pending = CursorToolRuntime::default();
let mut shell = call("call-1", "Shell");
shell.arguments = json!({"command": "git status"});
let id = pending.reserve_exec(&shell, &exec_context()).await.unwrap();
let completion = codec::stream_closed(id, &pending)
.await
.unwrap()
.expect("a running Exec should complete when its stream closes");
assert_eq!(completion.result().call_id, "call-1");
assert!(completion.result().is_error);
assert_eq!(
completion.result().content,
"Cursor Exec stream closed before returning a terminal result"
);
let Some(pb::tool_call::Tool::ShellToolCall(shell)) = &completion.tool_call().tool else {
panic!("expected typed Shell completion")
};
assert!(matches!(
shell.result.as_ref().and_then(|result| result.result.as_ref()),
Some(pb::shell_result::Result::SpawnError(error))
if error.error == "Cursor Exec stream closed before returning a terminal result"
));
assert!(pending.exec_call(id).await.is_none());
assert!(codec::stream_closed(id, &pending).await.unwrap().is_none());
}
#[tokio::test] #[tokio::test]
async fn tool_success_is_not_inferred_from_debug_text() { async fn tool_success_is_not_inferred_from_debug_text() {
let pending = CursorToolRuntime::default(); let pending = CursorToolRuntime::default();