mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 19:31:28 +08:00
fix: shell
This commit is contained in:
@@ -1,10 +1,19 @@
|
||||
use tokio::sync::oneshot;
|
||||
|
||||
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 {
|
||||
ToolResult(ToolResult),
|
||||
RuntimeMessage(CanonicalMessage),
|
||||
RuntimeEvent(RuntimeEvent),
|
||||
InsertMessages(MessageInsertion),
|
||||
ClientClosed { error: String },
|
||||
Cancel,
|
||||
}
|
||||
|
||||
+50
-14
@@ -151,15 +151,39 @@ impl CursorActor {
|
||||
context.dynamic_tools.keys().cloned().collect(),
|
||||
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 (port, core) = crate::client::session(256);
|
||||
let core_commands = core.commands.clone();
|
||||
let actor = RunActor::new(
|
||||
dependencies.store.clone(),
|
||||
dependencies.provider,
|
||||
dependencies.run_registry,
|
||||
);
|
||||
let core_run =
|
||||
actor.spawn(prepared, port, cancellation).await;
|
||||
let core_run = actor
|
||||
.spawn(
|
||||
prepared,
|
||||
port,
|
||||
core_commands,
|
||||
cancellation,
|
||||
)
|
||||
.await;
|
||||
let session = CursorSession::new(
|
||||
handle.clone(),
|
||||
dependencies.store,
|
||||
@@ -181,7 +205,6 @@ impl CursorActor {
|
||||
%error,
|
||||
"Cursor session failed"
|
||||
);
|
||||
handle.cancel();
|
||||
let _ = crate::cursor::lifecycle::fail(
|
||||
&handle, &error,
|
||||
);
|
||||
@@ -235,12 +258,14 @@ impl CursorActor {
|
||||
{
|
||||
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!(
|
||||
"Exec stream closed before result for id: {}",
|
||||
close.id
|
||||
)));
|
||||
Ok(Some(completion)) => {
|
||||
results_tx.send(completion)
|
||||
}
|
||||
Ok(None) => {}
|
||||
Err(error) => results_tx.send_error(error),
|
||||
}
|
||||
}
|
||||
Some(Message::Throw(throw)) => {
|
||||
@@ -311,13 +336,12 @@ impl CursorActor {
|
||||
//
|
||||
// The remaining unimplemented Action variants are
|
||||
// ShellCommandAction, StartPlanAction,
|
||||
// AsyncAskQuestionCompletionAction, CancelSubagentAction,
|
||||
// BackgroundShellAction, BackgroundSubagentAction,
|
||||
// AsyncAskQuestionCompletionAction, BackgroundShellAction,
|
||||
// BackgroundSubagentAction,
|
||||
// SubscriptionNotificationAction and GoalContinuationAction.
|
||||
// CancelSubagentAction must not start an LLM; variants whose wire
|
||||
// behavior is not captured yet need evidence before assigning
|
||||
// semantics. Every unsupported runtime Action must return an explicit
|
||||
// Protocol Error rather than falling through silently.
|
||||
// Variants whose wire behavior is not captured yet need evidence
|
||||
// before assigning semantics. Every unsupported runtime Action must
|
||||
// return an explicit Protocol Error rather than falling through silently.
|
||||
Some(
|
||||
pb::agent_client_message::Message::ConversationAction(
|
||||
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) => {
|
||||
results_tx.send_error(crate::Error::Protocol(format!(
|
||||
"unsupported runtime ConversationAction: {}",
|
||||
|
||||
@@ -87,8 +87,15 @@ pub(super) fn project(
|
||||
}
|
||||
pb::BackgroundTaskKind::Unspecified => unreachable!(),
|
||||
};
|
||||
let identity = agent_id.unwrap_or(&completion.task_id);
|
||||
let identity = format!("{}:{identity}", kind.as_str_name());
|
||||
let tool_call_id = completion
|
||||
.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)?;
|
||||
if completions
|
||||
.insert(identity.clone(), (completion, context))
|
||||
@@ -307,6 +314,30 @@ mod tests {
|
||||
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]
|
||||
fn completion_requires_the_captured_subagent_identity_and_terminal_reason() {
|
||||
let mut value = completion();
|
||||
|
||||
@@ -41,6 +41,7 @@ pub struct CursorRunContext {
|
||||
pub dynamic_tools: BTreeMap<String, pb::McpToolDefinition>,
|
||||
pub checkpoint_prompt: PromptSpec,
|
||||
pub compacting: bool,
|
||||
pub background_completion: bool,
|
||||
}
|
||||
|
||||
pub(crate) struct PrepareDependencies<'a> {
|
||||
@@ -123,7 +124,7 @@ pub(crate) async fn prepare(
|
||||
mode: mode_number,
|
||||
mut turn_user,
|
||||
action_context,
|
||||
event_id,
|
||||
mut event_id,
|
||||
input_id,
|
||||
starts_turn,
|
||||
compacting,
|
||||
@@ -172,14 +173,43 @@ pub(crate) async fn prepare(
|
||||
}
|
||||
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) => {
|
||||
store
|
||||
.anchor_input(&conversation_id, &input_id, proposed_base_revision_id)
|
||||
.anchor_input(&conversation_id, input_id, proposed_base_revision_id)
|
||||
.await?
|
||||
}
|
||||
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() {
|
||||
Some(event_id) => {
|
||||
store
|
||||
@@ -193,6 +223,10 @@ pub(crate) async fn prepare(
|
||||
let message_id = format!("request-context:{event_id}");
|
||||
match store.message(&conversation_id, &message_id).await? {
|
||||
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(
|
||||
event_id,
|
||||
&request_context,
|
||||
@@ -202,7 +236,7 @@ pub(crate) async fn prepare(
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let initial_messages = if compacting {
|
||||
let mut initial_messages = if compacting {
|
||||
Vec::new()
|
||||
} else {
|
||||
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 {
|
||||
RunAction::Compact
|
||||
} else if starts_turn {
|
||||
@@ -310,6 +348,7 @@ pub(crate) async fn prepare(
|
||||
.collect(),
|
||||
checkpoint_prompt,
|
||||
compacting,
|
||||
background_completion,
|
||||
},
|
||||
))
|
||||
}
|
||||
@@ -434,13 +473,13 @@ fn action(request: &pb::AgentRunRequest) -> Result<ActionProjection> {
|
||||
.filter(|text| !text.is_empty())
|
||||
.cloned(),
|
||||
);
|
||||
let event_id = format!("cursor:user:{}", user.message_id);
|
||||
let input_id = format!("cursor:user:{}", user.message_id);
|
||||
Ok(ActionProjection {
|
||||
mode,
|
||||
turn_user: Some(user.clone()),
|
||||
action_context: context.join("\n\n"),
|
||||
event_id: Some(event_id.clone()),
|
||||
input_id: Some(event_id),
|
||||
event_id: None,
|
||||
input_id: Some(input_id),
|
||||
starts_turn: true,
|
||||
compacting: false,
|
||||
background_completion: false,
|
||||
@@ -703,7 +742,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[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 {
|
||||
action: Some(pb::ConversationAction {
|
||||
action: Some(pb::conversation_action::Action::UserMessageAction(
|
||||
@@ -725,9 +764,11 @@ mod tests {
|
||||
let first = action(&request("message-one")).unwrap();
|
||||
let second = action(&request("message-two")).unwrap();
|
||||
|
||||
assert_eq!(first.event_id.as_deref(), Some("cursor:user:message-one"));
|
||||
assert_eq!(second.event_id.as_deref(), Some("cursor:user:message-two"));
|
||||
assert_ne!(first.event_id, second.event_id);
|
||||
assert_eq!(first.event_id, None);
|
||||
assert_eq!(second.event_id, None);
|
||||
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]
|
||||
|
||||
@@ -10,6 +10,7 @@ use crate::{
|
||||
proto::agent::v1 as pb,
|
||||
},
|
||||
model::{CanonicalMessage, MessageContent, Origin, Role},
|
||||
store::BlobId,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
@@ -80,12 +81,66 @@ pub async fn compile(
|
||||
compiler: &PromptCompiler,
|
||||
blobs: &BlobSynchronizer,
|
||||
) -> Result<CanonicalMessage> {
|
||||
let time = Time::now(
|
||||
let timestamp = Time::now(
|
||||
request_context
|
||||
.env
|
||||
.as_ref()
|
||||
.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([
|
||||
("OPEN_FILES", section(open_files(user))),
|
||||
(
|
||||
@@ -98,7 +153,7 @@ pub async fn compile(
|
||||
),
|
||||
),
|
||||
("ACTION_CONTEXT", section(action_context.to_string())),
|
||||
("TIMESTAMP", time.timestamp),
|
||||
("TIMESTAMP", timestamp),
|
||||
("USER_QUERY", user.text.clone()),
|
||||
("DEBUG_SERVER_ENDPOINT", String::new()),
|
||||
("DEBUG_LOG_PATH", String::new()),
|
||||
|
||||
@@ -9,7 +9,11 @@ use tokio_stream::StreamExt;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
cursor::{connect::END_STREAM_FLAG, observability::CursorTraceRecorder, CursorSessionRegistry},
|
||||
cursor::{
|
||||
connect::{self, END_STREAM_FLAG},
|
||||
observability::CursorTraceRecorder,
|
||||
CursorSessionRegistry,
|
||||
},
|
||||
Result,
|
||||
};
|
||||
|
||||
@@ -49,7 +53,7 @@ fn local_body_stream(
|
||||
trace.chunk(&chunk);
|
||||
if terminal {
|
||||
guard.complete();
|
||||
trace.finish(None);
|
||||
trace.finish(end_stream_error(&chunk));
|
||||
}
|
||||
yield Ok::<Bytes, Infallible>(chunk);
|
||||
if terminal {
|
||||
@@ -67,6 +71,30 @@ fn is_end_stream_frame(frame: &Bytes) -> bool {
|
||||
.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 {
|
||||
cancellation: CancellationToken,
|
||||
completed: bool,
|
||||
@@ -233,4 +261,67 @@ mod tests {
|
||||
drop(stream);
|
||||
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")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -88,6 +88,23 @@ impl CursorSession {
|
||||
}
|
||||
|
||||
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 {
|
||||
self.handle.emit(&interaction::summary_started())?;
|
||||
}
|
||||
|
||||
@@ -5,4 +5,4 @@ pub use request::{abort, mcp_request, mcp_state_request, request};
|
||||
pub(crate) use 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};
|
||||
|
||||
@@ -130,6 +130,53 @@ pub async fn client_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(
|
||||
entry: PendingExec,
|
||||
result: &pb::exec_client_message::Message,
|
||||
|
||||
@@ -338,6 +338,18 @@ impl CursorToolRuntime {
|
||||
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> {
|
||||
self.next_id
|
||||
.fetch_add(1, Ordering::Relaxed)
|
||||
|
||||
@@ -2,7 +2,12 @@ use std::sync::Arc;
|
||||
|
||||
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};
|
||||
|
||||
@@ -26,6 +31,7 @@ impl RunActor {
|
||||
&self,
|
||||
prepared: PreparedRun,
|
||||
client: ClientPort,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
cancellation: CancellationToken,
|
||||
) -> tokio::task::JoinHandle<RunOutcome> {
|
||||
let run_id = prepared.run_id.clone();
|
||||
@@ -35,6 +41,7 @@ impl RunActor {
|
||||
conversation_id.clone(),
|
||||
run_id.clone(),
|
||||
cancellation.clone(),
|
||||
commands,
|
||||
)
|
||||
.await;
|
||||
let actor = self.clone();
|
||||
|
||||
+81
-13
@@ -4,7 +4,10 @@ use std::sync::Arc;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::{ClientEvent, ClientPort, CommitBarrier, CommitCause, StateCommitted},
|
||||
client::{
|
||||
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
|
||||
StateCommitted,
|
||||
},
|
||||
model::{
|
||||
CanonicalMessage, MessageContent, Origin, PreparedRun, Role, RunAction, ToolRoundAssistant,
|
||||
ToolRoundId, Usage,
|
||||
@@ -158,6 +161,7 @@ impl RunEngine {
|
||||
calls: round.calls.clone(),
|
||||
recovered_started_at_ms: Some(round.started_at_ms),
|
||||
},
|
||||
Vec::new(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
@@ -242,21 +246,27 @@ impl RunEngine {
|
||||
&cycle_cancellation,
|
||||
);
|
||||
tokio::pin!(cycle);
|
||||
let cycle = tokio::select! {
|
||||
result = &mut cycle => result,
|
||||
let mut pending_insertions = Vec::new();
|
||||
let cycle = loop {
|
||||
tokio::select! {
|
||||
result = &mut cycle => break result,
|
||||
command = client.commands.recv() => {
|
||||
let message = match command {
|
||||
Some(crate::client::ClientCommand::RuntimeMessage(message)) => message,
|
||||
Some(crate::client::ClientCommand::RuntimeEvent(event)) => event.into_message(),
|
||||
Some(crate::client::ClientCommand::Cancel) => {
|
||||
Some(ClientCommand::InsertMessages(insertion)) => {
|
||||
pending_insertions.push(insertion);
|
||||
continue;
|
||||
}
|
||||
Some(ClientCommand::RuntimeMessage(message)) => message,
|
||||
Some(ClientCommand::RuntimeEvent(event)) => event.into_message(),
|
||||
Some(ClientCommand::Cancel) => {
|
||||
cycle_cancellation.cancel();
|
||||
return (RunOutcome::Cancelled, usage);
|
||||
}
|
||||
Some(crate::client::ClientCommand::ClientClosed { error }) => {
|
||||
Some(ClientCommand::ClientClosed { error }) => {
|
||||
cycle_cancellation.cancel();
|
||||
return (RunOutcome::Failed(RunFailure::Client(error)), usage);
|
||||
}
|
||||
Some(crate::client::ClientCommand::ToolResult(_)) => {
|
||||
Some(ClientCommand::ToolResult(_)) => {
|
||||
cycle_cancellation.cancel();
|
||||
return (
|
||||
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(
|
||||
&self.store,
|
||||
prepared,
|
||||
@@ -294,11 +317,12 @@ impl RunEngine {
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(revision) => revision,
|
||||
Ok((revision, _)) => revision,
|
||||
Err(outcome) => return (outcome, usage),
|
||||
};
|
||||
continue 'model;
|
||||
}
|
||||
}
|
||||
};
|
||||
let cycle = match cycle {
|
||||
Ok(cycle) => cycle,
|
||||
@@ -413,6 +437,27 @@ impl RunEngine {
|
||||
Ok(revision) => revision,
|
||||
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();
|
||||
if emit(
|
||||
client,
|
||||
@@ -453,6 +498,7 @@ impl RunEngine {
|
||||
calls: cycle.calls,
|
||||
recovered_started_at_ms: None,
|
||||
},
|
||||
pending_insertions,
|
||||
)
|
||||
.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,
|
||||
prepared: &PreparedRun,
|
||||
client: &mut ClientPort,
|
||||
cancellation: &CancellationToken,
|
||||
revision: crate::model::RevisionId,
|
||||
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(|| {
|
||||
RunOutcome::Failed(RunFailure::Protocol(
|
||||
"runtime message has no event identity".into(),
|
||||
@@ -706,7 +774,7 @@ async fn append_runtime_message(
|
||||
.await
|
||||
.map_err(|error| RunOutcome::Failed(error.into()))?;
|
||||
if !inserted {
|
||||
return Ok(revision);
|
||||
return Ok((revision, false));
|
||||
}
|
||||
let (barrier, ready) = CommitBarrier::before_continue();
|
||||
emit(
|
||||
@@ -721,7 +789,7 @@ async fn append_runtime_message(
|
||||
.await
|
||||
.map_err(|_| client_failure())?;
|
||||
wait_for_state_ready(ready, cancellation).await?;
|
||||
Ok(revision)
|
||||
Ok((revision, true))
|
||||
}
|
||||
|
||||
async fn hydrate_tool_images(
|
||||
|
||||
@@ -3,7 +3,10 @@ use std::{collections::HashMap, sync::Arc};
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::model::{ConversationId, RunId};
|
||||
use crate::{
|
||||
client::{ClientCommand, MessageInsertion},
|
||||
model::{CanonicalMessage, ConversationId, RunId},
|
||||
};
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct RunRegistry {
|
||||
@@ -13,6 +16,7 @@ pub struct RunRegistry {
|
||||
struct ActiveRun {
|
||||
run_id: RunId,
|
||||
cancellation: CancellationToken,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
}
|
||||
|
||||
impl RunRegistry {
|
||||
@@ -21,12 +25,14 @@ impl RunRegistry {
|
||||
conversation_id: ConversationId,
|
||||
run_id: RunId,
|
||||
cancellation: CancellationToken,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
) {
|
||||
let previous = self.active.lock().await.insert(
|
||||
conversation_id,
|
||||
ActiveRun {
|
||||
run_id: run_id.clone(),
|
||||
cancellation,
|
||||
commands,
|
||||
},
|
||||
);
|
||||
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) {
|
||||
let mut active = self.active.lock().await;
|
||||
if active
|
||||
|
||||
@@ -1,7 +1,10 @@
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::{ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, StateCommitted},
|
||||
client::{
|
||||
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
|
||||
StateCommitted,
|
||||
},
|
||||
model::{PreparedRun, RevisionId, ToolCall, ToolRoundAssistant, ToolRoundId},
|
||||
store::Store,
|
||||
};
|
||||
@@ -22,6 +25,7 @@ pub(super) async fn execute(
|
||||
cancellation: &CancellationToken,
|
||||
mut revision: RevisionId,
|
||||
round: ToolRound,
|
||||
insertions: Vec<MessageInsertion>,
|
||||
) -> std::result::Result<RevisionId, RunOutcome> {
|
||||
let ToolRound {
|
||||
id: round_id,
|
||||
@@ -66,7 +70,10 @@ pub(super) async fn execute(
|
||||
.await?;
|
||||
|
||||
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 {
|
||||
let command = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(RunOutcome::Cancelled),
|
||||
@@ -116,10 +123,13 @@ pub(super) async fn execute(
|
||||
}
|
||||
}
|
||||
Some(ClientCommand::RuntimeEvent(event)) => {
|
||||
pending_runtime_messages.push(event.into_message());
|
||||
pending_runtime_messages.push(PendingRuntimeMessage::Message(event.into_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::ClientClosed { error }) => {
|
||||
@@ -128,40 +138,42 @@ pub(super) async fn execute(
|
||||
None => return Err(client_failure()),
|
||||
}
|
||||
}
|
||||
for message in pending_runtime_messages {
|
||||
let event_id = message.runtime_event_id.clone().ok_or_else(|| {
|
||||
RunOutcome::Failed(RunFailure::Protocol(
|
||||
"runtime message has no event identity".into(),
|
||||
))
|
||||
})?;
|
||||
let (next, inserted) = store
|
||||
.append_message_once(
|
||||
&prepared.conversation_id,
|
||||
&prepared.run_id,
|
||||
revision,
|
||||
&message,
|
||||
)
|
||||
.await
|
||||
.map_err(failed)?;
|
||||
revision = next;
|
||||
if inserted {
|
||||
let (barrier, ready) = CommitBarrier::before_continue();
|
||||
send(
|
||||
for pending in pending_runtime_messages {
|
||||
match pending {
|
||||
PendingRuntimeMessage::Message(message) => {
|
||||
revision = super::engine::append_runtime_message(
|
||||
store,
|
||||
prepared,
|
||||
client,
|
||||
ClientEvent::StateCommitted(StateCommitted {
|
||||
revision_id: revision,
|
||||
tool_round_version: 0,
|
||||
cause: CommitCause::RuntimeEvent { event_id },
|
||||
barrier,
|
||||
}),
|
||||
cancellation,
|
||||
revision,
|
||||
message,
|
||||
)
|
||||
.await?;
|
||||
super::engine::wait_for_state_ready(ready, cancellation).await?;
|
||||
.await?
|
||||
.0;
|
||||
}
|
||||
PendingRuntimeMessage::Insertion(insertion) => {
|
||||
revision = super::engine::append_insertions(
|
||||
store,
|
||||
prepared,
|
||||
client,
|
||||
cancellation,
|
||||
revision,
|
||||
vec![insertion],
|
||||
)
|
||||
.await?
|
||||
.0;
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(revision)
|
||||
}
|
||||
|
||||
enum PendingRuntimeMessage {
|
||||
Message(crate::model::CanonicalMessage),
|
||||
Insertion(MessageInsertion),
|
||||
}
|
||||
|
||||
async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> {
|
||||
client
|
||||
.events
|
||||
|
||||
@@ -58,6 +58,38 @@ impl Store {
|
||||
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(
|
||||
&self,
|
||||
conversation_id: &ConversationId,
|
||||
|
||||
@@ -76,7 +76,7 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() {
|
||||
.unwrap();
|
||||
assert!(messages.iter().any(|message| {
|
||||
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())
|
||||
}));
|
||||
|
||||
@@ -131,12 +131,78 @@ async fn background_subagent_completion_starts_a_simulated_parent_turn() {
|
||||
assert_eq!(
|
||||
runtime_ids,
|
||||
[
|
||||
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id",
|
||||
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id-2"
|
||||
"runtime:background-completed:BACKGROUND_TASK_KIND_SUBAGENT:child-id:task-call",
|
||||
"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]
|
||||
async fn retrying_one_background_completion_reuses_its_runtime_message() {
|
||||
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();
|
||||
drive_completion(
|
||||
&second,
|
||||
completion_run_with_detail(
|
||||
"retry-child",
|
||||
"completion-retry-run-2",
|
||||
checkpoint,
|
||||
"updated retry payload",
|
||||
),
|
||||
completion_run("retry-child", "completion-retry-run-2", checkpoint),
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -192,7 +253,9 @@ async fn retrying_one_background_completion_reuses_its_runtime_message() {
|
||||
.iter()
|
||||
.filter(|message| {
|
||||
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(),
|
||||
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(
|
||||
child_id: &str,
|
||||
run_id: &str,
|
||||
|
||||
@@ -14,6 +14,7 @@ use cursor_server::{
|
||||
provider::{FinishReason, ModelEvent},
|
||||
run::{RunEngine, RunOutcome},
|
||||
};
|
||||
use tokio::{sync::oneshot, time::Duration};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
#[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);
|
||||
}
|
||||
|
||||
#[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]
|
||||
async fn a_failed_claim_cannot_overwrite_the_existing_run() {
|
||||
let (_directory, store) = fixtures::temp_store().await;
|
||||
|
||||
@@ -160,7 +160,7 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
|
||||
)
|
||||
.unwrap();
|
||||
let registry = CursorSessionRegistry::new(
|
||||
store,
|
||||
store.clone(),
|
||||
Arc::new(provider),
|
||||
PromptCompiler::new(assets),
|
||||
Default::default(),
|
||||
@@ -249,6 +249,29 @@ async fn runtime_protocol_failure_returns_connect_error_end_stream_and_closes()
|
||||
.unwrap(),
|
||||
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]
|
||||
|
||||
@@ -28,6 +28,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
|
||||
conversation.clone(),
|
||||
cursor_server::model::RunId::new("first"),
|
||||
first.clone(),
|
||||
cursor_server::client::session(1).1.commands,
|
||||
)
|
||||
.await;
|
||||
registry
|
||||
@@ -35,6 +36,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
|
||||
conversation.clone(),
|
||||
cursor_server::model::RunId::new("second"),
|
||||
second.clone(),
|
||||
cursor_server::client::session(1).1.commands,
|
||||
)
|
||||
.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 {
|
||||
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()
|
||||
},
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -218,3 +218,95 @@ async fn editing_a_logical_input_discards_its_active_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
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
@@ -18,6 +18,150 @@ use cursor_server::{
|
||||
};
|
||||
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]
|
||||
async fn unchanged_request_context_is_not_repeated_and_preserves_the_provider_prefix() {
|
||||
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 {
|
||||
panic!("request context message must contain one text part")
|
||||
};
|
||||
assert_eq!(
|
||||
request.history[1].message_id,
|
||||
"runtime:cursor:user:wire-user"
|
||||
);
|
||||
assert!(request.history[1]
|
||||
.message_id
|
||||
.starts_with("runtime:cursor:user:wire-user:"));
|
||||
assert!(!request.prompt.instructions.contains("workspace rule"));
|
||||
assert!(!request.prompt.instructions.contains("<mcp_meta_tools>"));
|
||||
let ProjectedContent::Parts(parts) = &request.history[1].content else {
|
||||
|
||||
@@ -89,7 +89,11 @@ async fn selected_image_bytes_flow_from_run_request_to_history_providers_and_che
|
||||
let user = requests[0]
|
||||
.history
|
||||
.iter()
|
||||
.find(|message| message.message_id == "runtime:cursor:user:image-user")
|
||||
.find(|message| {
|
||||
message
|
||||
.message_id
|
||||
.starts_with("runtime:cursor:user:image-user:")
|
||||
})
|
||||
.unwrap();
|
||||
let ProjectedContent::Parts(parts) = &user.content else {
|
||||
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 bytes = store.get_blob(&id).await.unwrap().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);
|
||||
break;
|
||||
}
|
||||
|
||||
@@ -10,11 +10,15 @@ use cursor_server::{
|
||||
provider::{ModelEvent, Provider, ProviderStream},
|
||||
Error,
|
||||
};
|
||||
use futures_util::stream;
|
||||
use futures_util::{stream, StreamExt};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
enum FakeResponse {
|
||||
Events(Vec<Result<ModelEvent, Error>>),
|
||||
Gated {
|
||||
ready: Arc<tokio::sync::Notify>,
|
||||
events: Vec<Result<ModelEvent, Error>>,
|
||||
},
|
||||
Pending,
|
||||
}
|
||||
|
||||
@@ -43,6 +47,17 @@ impl FakeProvider {
|
||||
.unwrap()
|
||||
.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> {
|
||||
self.requests.lock().unwrap().clone()
|
||||
}
|
||||
@@ -63,6 +78,13 @@ impl Provider for FakeProvider {
|
||||
.expect("fake response configured");
|
||||
match 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()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -283,7 +283,9 @@ async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
|
||||
.unwrap();
|
||||
assert!(messages[0].message_id.starts_with("request-context:"));
|
||||
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.len(),
|
||||
|
||||
@@ -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]
|
||||
async fn tool_success_is_not_inferred_from_debug_text() {
|
||||
let pending = CursorToolRuntime::default();
|
||||
|
||||
Reference in New Issue
Block a user