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};
#[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
View File
@@ -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: {}",
+33 -2
View File
@@ -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();
+52 -11
View File
@@ -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]
+58 -3
View File
@@ -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()),
+93 -2
View File
@@ -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")
);
}
}
+17
View File
@@ -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())?;
}
+1 -1
View File
@@ -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};
+47
View File
@@ -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,
+12
View File
@@ -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)
+8 -1
View File
@@ -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
View File
@@ -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(
+38 -1
View File
@@ -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
+45 -33
View File
@@ -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(
client,
ClientEvent::StateCommitted(StateCommitted {
revision_id: revision,
tool_round_version: 0,
cause: CommitCause::RuntimeEvent { event_id },
barrier,
}),
)
.await?;
super::engine::wait_for_state_ready(ready, cancellation).await?;
for pending in pending_runtime_messages {
match pending {
PendingRuntimeMessage::Message(message) => {
revision = super::engine::append_runtime_message(
store,
prepared,
client,
cancellation,
revision,
message,
)
.await?
.0;
}
PendingRuntimeMessage::Insertion(insertion) => {
revision = super::engine::append_insertions(
store,
prepared,
client,
cancellation,
revision,
vec![insertion],
)
.await?
.0;
}
}
}
Ok(revision)
}
enum PendingRuntimeMessage {
Message(crate::model::CanonicalMessage),
Insertion(MessageInsertion),
}
async fn send(client: &ClientPort, event: ClientEvent) -> std::result::Result<(), RunOutcome> {
client
.events
+32
View File
@@ -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,