4 Commits
17 changed files with 1626 additions and 173 deletions
+443 -107
View File
@@ -1,9 +1,10 @@
use std::collections::HashMap;
use std::collections::{HashMap, HashSet};
use std::{marker::PhantomData, sync::Arc, time::Instant};
use futures::StreamExt;
use serde_json::{Value, json};
use tokio::sync::mpsc;
use tokio::time::{Duration, Instant as TokioInstant};
use tracing::{debug, info, trace, warn};
use crate::{
@@ -27,11 +28,14 @@ use crate::{
timeline::{TextBlockCollector, ThinkingBlockCollector, Timeline, ToolCallCollector},
tool::{
ToolCall, ToolDefinition as EngineToolDefinition, ToolError, ToolExecutionContext,
ToolOutputLimits, ToolResult, truncate_content,
ToolOutputLimits, ToolResult, ToolResultDisposition, truncate_content,
},
tool_server::{ToolServer, ToolServerHandle},
};
const TOOL_CANCEL_SIGNAL_TIMEOUT: Duration = Duration::from_millis(100);
const TOOL_CANCEL_GRACE_PERIOD: Duration = Duration::from_millis(500);
/// Engine errors
#[derive(Debug, thiserror::Error)]
pub enum EngineError {
@@ -53,6 +57,9 @@ pub enum EngineError {
/// A durable-history observer rejected an item before it entered history.
#[error("History append failed: {0}")]
HistoryAppend(String),
/// Tool terminalization lost its execution-attempt compare-and-set fence.
#[error("Tool execution attempt fence failed: {0}")]
ToolAttemptFence(String),
}
/// Tool registration error
@@ -70,6 +77,59 @@ pub struct EngineConfig {
_private: (),
}
/// Project terminal tool outputs into the assistant's original ToolCall order.
///
/// Runtime history intentionally retains completion order so every result can
/// be committed without waiting for slower siblings. The provider projection
/// is deterministic within each contiguous result batch and does not rewrite
/// the committed transcript.
struct ProviderHistoryProjection {
items: Vec<Item>,
original_to_projected_index: Vec<usize>,
}
fn materialize_provider_history(items: &[Item]) -> ProviderHistoryProjection {
let mut materialized: Vec<_> = items.iter().cloned().enumerate().collect();
let mut call_order = HashMap::<String, usize>::new();
let mut next_call_order = 0usize;
let mut index = 0usize;
while index < materialized.len() {
match &materialized[index].1 {
Item::ToolCall { call_id, .. } => {
call_order.insert(call_id.clone(), next_call_order);
next_call_order += 1;
index += 1;
}
Item::ToolResult { .. } => {
let start = index;
while index < materialized.len()
&& matches!(materialized[index].1, Item::ToolResult { .. })
{
index += 1;
}
materialized[start..index].sort_by_key(|(_, item)| match item {
Item::ToolResult { call_id, .. } => {
call_order.get(call_id).copied().unwrap_or(usize::MAX)
}
_ => unreachable!("tool-result run contains only ToolResult items"),
});
}
_ => index += 1,
}
}
let mut original_to_projected_index = vec![0; materialized.len()];
for (projected_index, (original_index, _)) in materialized.iter().enumerate() {
original_to_projected_index[*original_index] = projected_index;
}
ProviderHistoryProjection {
items: materialized.into_iter().map(|(_, item)| item).collect(),
original_to_projected_index,
}
}
/// Legacy serializable outcome used by the Worker session-log compatibility boundary.
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
@@ -126,10 +186,68 @@ pub struct EngineRunOutput<C: LlmClient, A = ()> {
/// Internal: tool execution result
enum ToolExecutionResult {
Completed(Vec<ToolResult>),
Completed,
Paused,
}
#[derive(Debug, Clone)]
struct ToolExecutionAttempt {
attempt_id: String,
terminal: bool,
}
/// Per-batch compare-and-set fence for terminal ToolResult commits.
///
/// A completion may commit only when its attempt id still matches the active
/// execution for that call and no prior terminal output has won the fence.
#[derive(Debug, Default)]
struct ToolExecutionAttemptFence {
attempts: HashMap<String, ToolExecutionAttempt>,
}
impl ToolExecutionAttemptFence {
fn register(&mut self, call_id: String, attempt_id: String) {
self.attempts.insert(
call_id,
ToolExecutionAttempt {
attempt_id,
terminal: false,
},
);
}
fn can_commit(&self, call_id: &str, attempt_id: &str) -> bool {
matches!(
self.attempts.get(call_id),
Some(attempt) if attempt.attempt_id == attempt_id && !attempt.terminal
)
}
fn commit_terminal(&mut self, call_id: &str, attempt_id: &str) -> bool {
let Some(attempt) = self.attempts.get_mut(call_id) else {
return false;
};
if attempt.attempt_id != attempt_id || attempt.terminal {
return false;
}
attempt.terminal = true;
true
}
fn is_terminal(&self, call_id: &str) -> bool {
self.attempts
.get(call_id)
.is_some_and(|attempt| attempt.terminal)
}
#[cfg(test)]
fn attempt_id(&self, call_id: &str) -> Option<&str> {
self.attempts
.get(call_id)
.map(|attempt| attempt.attempt_id.as_str())
}
}
const MAX_STREAM_CONTINUATIONS: u32 = 3;
/// Central component for managing LLM interactions
@@ -892,8 +1010,14 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
request = request.system(system);
}
// Add items directly (Request now uses Items natively)
request = request.items(context.iter().cloned());
// History keeps terminal tool outputs in completion order so each
// result can be committed immediately. Providers, however, expect a
// deterministic projection matching the assistant's ToolCall order.
let projection = materialize_provider_history(context);
let projected_cache_anchor = self
.cache_anchor
.and_then(|anchor| projection.original_to_projected_index.get(anchor).copied());
request = request.items(projection.items);
// Add tool definitions
for tool_def in tool_definitions {
@@ -906,7 +1030,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
// Attach the cache prefix anchor (may be narrower than `context`
// if the prune projection trimmed items from the head — keep it
// in range).
request.cache_anchor = self.cache_anchor.filter(|&anchor| anchor < context.len());
request.cache_anchor = projected_cache_anchor;
request.cache_key = self.cache_key.clone();
request
@@ -978,9 +1102,11 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
/// executes approved tools in parallel and applies post_tool_call hooks to results.
async fn execute_tools(
&mut self,
history: &mut History<A>,
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
tool_calls: Vec<ToolCall>,
) -> Result<ToolExecutionResult, EngineError> {
use futures::future::join_all;
use futures::stream::{FuturesUnordered, StreamExt};
// Map from tool call ID to (ToolCall, Meta, Tool, Context)
// Retained because it's needed for PostToolCall hooks
@@ -1039,110 +1165,297 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
context.clone(),
),
);
approved_calls.push((tool_call, context));
approved_calls.push((tool_call, context, Some(info.tool)));
} else {
// Unknown tools go into approved list as-is (will error at execution)
let context = ToolExecutionContext::new(&tool_call.id, &batch_id, call_index);
approved_calls.push((tool_call, context));
approved_calls.push((tool_call, context, None));
}
}
// Phase 2: Execute approved tools in parallel (cancellable)
let futures: Vec<_> = approved_calls
// Phase 2: Execute approved tools in parallel. FuturesUnordered yields
// each terminal result as soon as that call completes instead of
// holding fast siblings behind the slowest call in the batch.
let started_calls: Vec<_> = approved_calls
.iter()
.map(|(tool_call, context, _)| (tool_call.id.clone(), context.batch_id.clone()))
.collect();
let mut attempt_fence = ToolExecutionAttemptFence::default();
for (call_id, attempt_id) in &started_calls {
attempt_fence.register(call_id.clone(), attempt_id.clone());
}
let futures: FuturesUnordered<_> = approved_calls
.into_iter()
.map(|(tool_call, context)| {
let tool_server = self.tool_server.clone();
async move {
let input_json = serde_json::to_string(&tool_call.input).unwrap_or_default();
match tool_server
.call_tool(&tool_call.name, &input_json, context)
.await
{
.map(|(tool_call, context, tool)| async move {
let attempt_id = context.batch_id.clone();
let input_json = serde_json::to_string(&tool_call.input).unwrap_or_default();
let result = match tool {
None => ToolResult::error(
&tool_call.id,
format!("Tool not found: {}", tool_call.name),
),
Some(tool) => match tool.execute(&input_json, context).await {
Ok(output) => ToolResult::from_output(&tool_call.id, output),
Err(e) => ToolResult::error(&tool_call.id, e.to_string()),
}
}
Err(ToolError::Cancelled(output)) => {
ToolResult::from_output_with_disposition(
&tool_call.id,
output,
ToolResultDisposition::Cancelled,
)
}
Err(ToolError::Interrupted(output)) => {
ToolResult::from_output_with_disposition(
&tool_call.id,
output,
ToolResultDisposition::Interrupted,
)
}
Err(error) => ToolResult::error(&tool_call.id, error.to_string()),
},
};
(attempt_id, result)
})
.collect();
// Make tool execution cancellable
let mut results = tokio::select! {
results = join_all(futures) => results,
cancel = self.cancel_rx.recv() => {
if cancel.is_some() {
info!("Tool execution cancelled");
// Synthetic results are already terminal and need no execution wait.
// Commit them before polling ordinary calls so they obey the same
// commit-before-publish boundary.
let mut terminal_call_ids = HashSet::new();
for result in synthetic_results {
self.finalize_and_commit_tool_result(
history,
annotate,
result,
None,
&call_info_map,
&mut attempt_fence,
&mut terminal_call_ids,
)
.await?;
}
let mut futures = futures;
while !futures.is_empty() {
tokio::select! {
// If cancellation and a completed result are both ready, drain
// the completed result first. This preserves every terminal
// output observed before the cancellation boundary.
biased;
result = futures.next() => {
let (attempt_id, result) =
result.expect("non-empty FuturesUnordered returns a result");
self.finalize_and_commit_tool_result(
history,
annotate,
result,
Some(&attempt_id),
&call_info_map,
&mut attempt_fence,
&mut terminal_call_ids,
).await?;
}
self.timeline.abort_current_block();
return Err(EngineError::Cancelled);
}
};
results.extend(synthetic_results);
// Phase 3: Apply post_tool_call interceptor
for tool_result in &mut results {
if let Some((tool_call, meta, tool, context)) =
call_info_map.get(&tool_result.tool_use_id)
{
let mut info = ToolResultInfo {
call: tool_call.clone(),
result: tool_result.clone(),
meta: meta.clone(),
tool: tool.clone(),
context: context.clone(),
};
match self.interceptor.post_tool_call(&mut info).await {
PostToolAction::Continue => {}
PostToolAction::Abort(reason) => {
return Err(EngineError::Aborted(reason));
cancel = self.cancel_rx.recv() => {
if cancel.is_some() {
info!("Tool execution cancellation requested");
}
}
// Reflect interceptor-modified results
*tool_result = info.result;
}
}
// Phase 4: Cap `content` byte-size before it enters history.
// Runs *after* post_tool_call so interceptors (audit, logging,
// classification) still observe the full content, and any
// content they inject is also truncated — closing the last gap
// before the data reaches the next LLM request.
if let Some(limits) = self.tool_output_limits.as_ref() {
for tool_result in &mut results {
let Some(content) = tool_result.content.as_mut() else {
continue;
};
let Some((tool_call, _, _, _)) = call_info_map.get(&tool_result.tool_use_id) else {
continue;
};
let limit = limits.limit_for(&tool_call.name);
let before = content.len();
truncate_content(content, limit);
if content.len() != before {
warn!(
tool = %tool_call.name,
before_bytes = before,
after_bytes = content.len(),
limit_bytes = limit,
"Tool output exceeded byte limit and was truncated"
);
self.emit_warning(&format!(
"tool `{}` output truncated from {} to {} bytes (limit {})",
tool_call.name,
before,
content.len(),
limit
));
let cancellation_requests = call_info_map
.iter()
.filter(|(call_id, _)| !terminal_call_ids.contains(*call_id))
.map(|(call_id, (_, _, tool, _))| {
let call_id = call_id.clone();
let tool = tool.clone();
async move { (call_id.clone(), tool.cancel(&call_id).await) }
});
let cancellation_requests: FuturesUnordered<_> =
cancellation_requests.collect();
match tokio::time::timeout(
TOOL_CANCEL_SIGNAL_TIMEOUT,
cancellation_requests.collect::<Vec<_>>(),
)
.await
{
Ok(results) => {
for (call_id, result) in results {
if let Err(error) = result {
warn!(
%call_id,
error = %error,
"Tool cooperative cancellation request failed"
);
}
}
}
Err(_) => warn!("Tool cooperative cancellation request timed out"),
}
// Keep polling the original execution futures for a bounded
// grace period so cooperative providers can return their
// confirmed terminal output, including bounded progress.
let deadline = TokioInstant::now() + TOOL_CANCEL_GRACE_PERIOD;
while !futures.is_empty() {
tokio::select! {
biased;
result = futures.next() => {
let (attempt_id, result) =
result.expect("non-empty FuturesUnordered returns a result");
self.finalize_and_commit_tool_result(
history,
annotate,
result,
Some(&attempt_id),
&call_info_map,
&mut attempt_fence,
&mut terminal_call_ids,
).await?;
}
_ = tokio::time::sleep_until(deadline) => break,
}
}
// Calls that did not confirm a terminal outcome inside the
// grace period are durably closed as OutcomeUnknown before
// Engine/Worker final status becomes observable.
for (call_id, attempt_id) in &started_calls {
if !attempt_fence.is_terminal(call_id) {
self.finalize_and_commit_tool_result(
history,
annotate,
ToolResult::outcome_unknown(call_id),
Some(attempt_id),
&call_info_map,
&mut attempt_fence,
&mut terminal_call_ids,
).await?;
}
}
self.timeline.abort_current_block();
return Err(EngineError::Cancelled);
}
}
}
// Emit per-result callbacks on the post-truncation payload.
for tool_result in &results {
self.emit_tool_result(tool_result);
Ok(ToolExecutionResult::Completed)
}
/// Apply post-execution policy, bound the model-visible payload, durably
/// append one terminal ToolResult, and only then publish it to observers.
async fn finalize_and_commit_tool_result(
&mut self,
history: &mut History<A>,
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
mut tool_result: ToolResult,
execution_attempt_id: Option<&str>,
call_info_map: &HashMap<
String,
(
ToolCall,
crate::tool::ToolMeta,
Arc<dyn crate::tool::Tool>,
ToolExecutionContext,
),
>,
attempt_fence: &mut ToolExecutionAttemptFence,
terminal_call_ids: &mut HashSet<String>,
) -> Result<bool, EngineError> {
let call_id = tool_result.tool_use_id.as_str();
let may_commit = match execution_attempt_id {
Some(attempt_id) => attempt_fence.can_commit(call_id, attempt_id),
None => !terminal_call_ids.contains(call_id),
};
if !may_commit {
warn!(
call_id,
execution_attempt_id,
disposition = ?tool_result.disposition,
"Ignoring stale or duplicate tool result after terminal output commit"
);
return Ok(false);
}
Ok(ToolExecutionResult::Completed(results))
let call_info = call_info_map.get(&tool_result.tool_use_id);
let mut abort_reason = None;
if let Some((tool_call, meta, tool, context)) = call_info {
let mut info = ToolResultInfo {
call: tool_call.clone(),
result: tool_result,
meta: meta.clone(),
tool: tool.clone(),
context: context.clone(),
};
match self.interceptor.post_tool_call(&mut info).await {
PostToolAction::Continue => {}
PostToolAction::Abort(reason) => {
abort_reason = Some(reason);
}
}
tool_result = info.result;
}
if tool_result.is_error && tool_result.disposition.is_success() {
tool_result.disposition = ToolResultDisposition::Error;
}
tool_result.is_error = !tool_result.disposition.is_success();
// Cap content only after post_tool_call so interceptors still observe
// the full payload and any content they inject is bounded too.
if let (Some(limits), Some((tool_call, _, _, _)), Some(content)) = (
self.tool_output_limits.as_ref(),
call_info,
tool_result.content.as_mut(),
) {
let limit = limits.limit_for(&tool_call.name);
let before = content.len();
truncate_content(content, limit);
if content.len() != before {
warn!(
tool = %tool_call.name,
before_bytes = before,
after_bytes = content.len(),
limit_bytes = limit,
"Tool output exceeded byte limit and was truncated"
);
self.emit_warning(&format!(
"tool `{}` output truncated from {} to {} bytes (limit {})",
tool_call.name,
before,
content.len(),
limit
));
}
}
let item = Item::tool_result_item_with_disposition_and_attachments(
&tool_result.tool_use_id,
&tool_result.summary,
tool_result.content.clone(),
tool_result.disposition,
tool_result.attachments.clone(),
);
self.append_history_items(history, std::iter::once(item), annotate)?;
if let Some(attempt_id) = execution_attempt_id
&& !attempt_fence.commit_terminal(&tool_result.tool_use_id, attempt_id)
{
return Err(EngineError::ToolAttemptFence(
"tool execution attempt fence changed during terminal commit".to_string(),
));
}
terminal_call_ids.insert(tool_result.tool_use_id.clone());
debug!(
tool = call_info
.map(|(call, _, _, _)| call.name.as_str())
.unwrap_or("unknown"),
call_id = %tool_result.tool_use_id,
execution_attempt_id,
disposition = ?tool_result.disposition,
"Tool execution terminalized"
);
self.emit_tool_result(&tool_result);
if let Some(reason) = abort_reason {
return Err(EngineError::Aborted(reason));
}
Ok(true)
}
/// Internal turn execution logic
@@ -1658,23 +1971,9 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
tool_calls: Vec<ToolCall>,
) -> Result<Option<EngineResult>, EngineError> {
match self.execute_tools(tool_calls).await {
match self.execute_tools(history, annotate, tool_calls).await {
Ok(ToolExecutionResult::Paused) => Ok(Some(EngineResult::Paused)),
Ok(ToolExecutionResult::Completed(results)) => {
// Route per-result pushes through the callback path so
// observers see each tool result as it lands.
let items = results.into_iter().map(|result| {
Item::tool_result_item_with_attachments(
&result.tool_use_id,
&result.summary,
result.content,
result.is_error,
result.attachments,
)
});
self.append_history_items(history, items, annotate)?;
Ok(None)
}
Ok(ToolExecutionResult::Completed) => Ok(None),
Err(err) => Err(err),
}
}
@@ -2296,6 +2595,43 @@ mod tests {
use crate::tool::{Attachment, ImageAttachment};
use std::time::Duration;
#[test]
fn tool_execution_attempt_fence_rejects_duplicate_and_stale_results() {
let mut fence = ToolExecutionAttemptFence::default();
fence.register("call".to_string(), "attempt-1".to_string());
assert!(fence.can_commit("call", "attempt-1"));
assert!(fence.commit_terminal("call", "attempt-1"));
assert!(!fence.commit_terminal("call", "attempt-1"));
fence.register("call".to_string(), "attempt-2".to_string());
assert_eq!(fence.attempt_id("call"), Some("attempt-2"));
assert!(!fence.can_commit("call", "attempt-1"));
assert!(!fence.commit_terminal("call", "attempt-1"));
assert!(fence.commit_terminal("call", "attempt-2"));
}
#[test]
fn provider_projection_reorders_results_and_remaps_cache_anchor() {
let items = vec![
Item::tool_call_json("call_slow", "slow", serde_json::json!({})),
Item::tool_call_json("call_fast", "fast", serde_json::json!({})),
Item::tool_result_item("call_fast", "fast result", None, false),
Item::tool_result_item("call_slow", "slow result", None, false),
];
let projection = materialize_provider_history(&items);
let result_order: Vec<_> = projection
.items
.iter()
.filter_map(|item| match item {
Item::ToolResult { call_id, .. } => Some(call_id.as_str()),
_ => None,
})
.collect();
assert_eq!(result_order, ["call_slow", "call_fast"]);
assert_eq!(projection.original_to_projected_index, [0, 1, 3, 2]);
}
#[test]
fn tool_attachment_round_trips_through_durable_history_json() {
let body: Arc<[u8]> = Arc::from(&b"image-body"[..]);
+3 -1
View File
@@ -28,7 +28,9 @@ pub use handler::ToolUseBlockStart;
pub use history::{History, HistoryEntry};
pub use interceptor::Interceptor;
pub use message::{ContentPart, Item, Message, Role};
pub use tool::{ToolCall, ToolExecutionContext, ToolOutputLimits, ToolResult};
pub use tool::{
ToolCall, ToolExecutionContext, ToolOutputLimits, ToolResult, ToolResultDisposition,
};
pub use usage_record::UsageRecord;
/// Implementation dependencies used by code generated from `agen` macros.
+37 -2
View File
@@ -9,7 +9,7 @@
use std::{fmt, sync::Arc};
use crate::tool::Attachment;
use crate::tool::{Attachment, ToolResultDisposition};
use base64::Engine as _;
use serde::{Deserialize, Serialize};
@@ -121,6 +121,9 @@ pub enum Item {
/// Detailed output (removed by pruning when old enough)
#[serde(default, skip_serializing_if = "Option::is_none")]
content: Option<String>,
/// Typed terminal state used for replay and recovery.
#[serde(default, skip_serializing_if = "ToolResultDisposition::is_success")]
disposition: ToolResultDisposition,
/// Whether the tool result represents an execution error.
#[serde(default, skip_serializing_if = "is_false")]
is_error: bool,
@@ -261,7 +264,17 @@ impl Item {
content: Option<String>,
is_error: bool,
) -> Self {
Self::tool_result_item_with_attachments(call_id, summary, content, is_error, Vec::new())
Self::tool_result_item_with_disposition_and_attachments(
call_id,
summary,
content,
if is_error {
ToolResultDisposition::Error
} else {
ToolResultDisposition::Success
},
Vec::new(),
)
}
/// Create a tool result item with durable, prunable structured attachments.
@@ -272,11 +285,33 @@ impl Item {
is_error: bool,
attachments: Vec<Attachment>,
) -> Self {
Self::tool_result_item_with_disposition_and_attachments(
call_id,
summary,
content,
if is_error {
ToolResultDisposition::Error
} else {
ToolResultDisposition::Success
},
attachments,
)
}
pub fn tool_result_item_with_disposition_and_attachments(
call_id: impl Into<String>,
summary: impl Into<String>,
content: Option<String>,
disposition: ToolResultDisposition,
attachments: Vec<Attachment>,
) -> Self {
let is_error = !disposition.is_success();
Self::ToolResult {
id: None,
call_id: call_id.into(),
summary: summary.into(),
content,
disposition,
is_error,
attachments,
}
+68 -1
View File
@@ -23,6 +23,12 @@ pub enum ToolError {
/// Internal error
#[error("Internal error: {0}")]
Internal(String),
/// Cooperative cancellation completed with bounded terminal output.
#[error("Tool execution cancelled")]
Cancelled(ToolOutput),
/// Execution was interrupted with a confirmed bounded terminal output.
#[error("Tool execution interrupted")]
Interrupted(ToolOutput),
}
// =============================================================================
@@ -158,6 +164,28 @@ pub enum Attachment {
Image(ImageAttachment),
}
/// Terminal disposition of one started tool call.
///
/// `Cancelled` means the tool confirmed cancellation. `OutcomeUnknown` means
/// execution stopped without confirmation, so neither completion nor side
/// effects may be inferred.
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "snake_case")]
pub enum ToolResultDisposition {
#[default]
Success,
Error,
Interrupted,
Cancelled,
OutcomeUnknown,
}
impl ToolResultDisposition {
pub const fn is_success(&self) -> bool {
matches!(self, Self::Success)
}
}
/// Tool execution result.
///
/// Every output has a mandatory `summary` (1-2 lines) that persists in
@@ -402,6 +430,17 @@ pub trait Tool: Send + Sync {
input_json: &str,
ctx: ToolExecutionContext,
) -> Result<ToolOutput, ToolError>;
/// Request cooperative cancellation for one started call.
///
/// Implementations that own cancellable provider operations should signal
/// the exact execution identified by `call_id`, then let `execute` return
/// the confirmed bounded terminal output. The Engine applies a bounded
/// grace period and falls back to `OutcomeUnknown` when confirmation never
/// arrives.
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
Ok(())
}
}
// =============================================================================
@@ -429,6 +468,9 @@ pub struct ToolCall {
pub struct ToolResult {
/// Corresponding tool call ID
pub tool_use_id: String,
/// Typed terminal state.
#[serde(default, skip_serializing_if = "ToolResultDisposition::is_success")]
pub disposition: ToolResultDisposition,
/// Short summary (always kept in history)
pub summary: String,
/// Detailed output (prunable)
@@ -445,11 +487,20 @@ pub struct ToolResult {
impl ToolResult {
/// Create a success result from a [`ToolOutput`].
pub fn from_output(tool_use_id: impl Into<String>, output: ToolOutput) -> Self {
Self::from_output_with_disposition(tool_use_id, output, ToolResultDisposition::Success)
}
pub fn from_output_with_disposition(
tool_use_id: impl Into<String>,
output: ToolOutput,
disposition: ToolResultDisposition,
) -> Self {
Self {
tool_use_id: tool_use_id.into(),
disposition,
summary: output.summary,
content: output.content,
is_error: false,
is_error: !disposition.is_success(),
attachments: output.attachments,
}
}
@@ -458,12 +509,28 @@ impl ToolResult {
pub fn error(tool_use_id: impl Into<String>, message: impl Into<String>) -> Self {
Self {
tool_use_id: tool_use_id.into(),
disposition: ToolResultDisposition::Error,
summary: message.into(),
content: None,
is_error: true,
attachments: Vec::new(),
}
}
/// Close an execution whose completion and side effects cannot be confirmed.
pub fn outcome_unknown(tool_use_id: impl Into<String>) -> Self {
Self {
tool_use_id: tool_use_id.into(),
disposition: ToolResultDisposition::OutcomeUnknown,
summary: "Tool execution outcome unknown".to_string(),
content: Some(
"Execution was interrupted before completion could be confirmed. Completion and side effects are unknown."
.to_string(),
),
is_error: true,
attachments: Vec::new(),
}
}
}
#[cfg(test)]
+8 -1
View File
@@ -19,6 +19,7 @@ use std::sync::atomic::{AtomicUsize, Ordering};
pub struct MockLlmClient {
responses: Arc<Vec<Vec<Event>>>,
call_count: Arc<AtomicUsize>,
requests: Arc<Mutex<Vec<Request>>>,
}
impl MockLlmClient {
@@ -30,6 +31,7 @@ impl MockLlmClient {
Self {
responses: Arc::new(responses),
call_count: Arc::new(AtomicUsize::new(0)),
requests: Arc::new(Mutex::new(Vec::new())),
}
}
@@ -41,6 +43,10 @@ impl MockLlmClient {
pub fn event_count(&self) -> usize {
self.responses.iter().map(|v| v.len()).sum()
}
pub fn requests(&self) -> Vec<Request> {
self.requests.lock().unwrap().clone()
}
}
#[async_trait]
@@ -51,8 +57,9 @@ impl LlmClient for MockLlmClient {
async fn stream(
&self,
_request: Request,
request: Request,
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError> {
self.requests.lock().unwrap().push(request);
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
if count >= self.responses.len() {
return Err(ClientError::Api {
+512 -1
View File
@@ -10,8 +10,9 @@ use agen::interceptor::{Interceptor, PostToolAction, PreToolAction, ToolCallInfo
use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
use agen::tool::{
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult,
ToolResultDisposition,
};
use agen::{Engine, History};
use agen::{Engine, History, Item};
use async_trait::async_trait;
mod common;
@@ -70,6 +71,95 @@ impl Tool for SlowTool {
}
}
#[derive(Clone)]
struct FirstAttemptHangsTool {
calls: Arc<AtomicUsize>,
}
impl FirstAttemptHangsTool {
fn new() -> Self {
Self {
calls: Arc::new(AtomicUsize::new(0)),
}
}
fn definition(&self) -> ToolDefinition {
let tool = self.clone();
Arc::new(move || {
let meta = ToolMeta::new("hang_once")
.description("Hangs on the first execution attempt")
.input_schema(serde_json::json!({"type": "object"}));
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
})
}
fn call_count(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
#[async_trait]
impl Tool for FirstAttemptHangsTool {
async fn execute(
&self,
_input_json: &str,
_ctx: ToolExecutionContext,
) -> Result<ToolOutput, ToolError> {
let attempt = self.calls.fetch_add(1, Ordering::SeqCst);
if attempt == 0 {
std::future::pending::<()>().await;
}
Ok("completed on retry".to_string().into())
}
}
#[derive(Clone)]
struct CooperativeCancelTool {
calls: Arc<AtomicUsize>,
cancelled: Arc<tokio::sync::Notify>,
}
impl CooperativeCancelTool {
fn new() -> Self {
Self {
calls: Arc::new(AtomicUsize::new(0)),
cancelled: Arc::new(tokio::sync::Notify::new()),
}
}
fn definition(&self) -> ToolDefinition {
let tool = self.clone();
Arc::new(move || {
let meta = ToolMeta::new("cooperative")
.description("Returns bounded progress after cancellation")
.input_schema(serde_json::json!({"type": "object"}));
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
})
}
}
#[async_trait]
impl Tool for CooperativeCancelTool {
async fn execute(
&self,
_input_json: &str,
_ctx: ToolExecutionContext,
) -> Result<ToolOutput, ToolError> {
self.calls.fetch_add(1, Ordering::SeqCst);
self.cancelled.notified().await;
Err(ToolError::Cancelled(ToolOutput {
summary: "cooperative command cancelled".to_string(),
content: Some("stdout before cancellation\nstderr before cancellation".to_string()),
attachments: Vec::new(),
}))
}
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
self.cancelled.notify_one();
Ok(())
}
}
#[derive(Clone)]
struct ContextRecordingTool {
name: String,
@@ -179,6 +269,354 @@ async fn test_parallel_tool_execution() {
println!("Parallel execution completed in {:?}", elapsed);
}
#[tokio::test]
async fn completed_results_commit_before_publish_without_waiting_for_siblings() {
let client = MockLlmClient::with_responses(vec![
vec![
Event::tool_use_start(0, "call_slow", "slow_first"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::tool_use_start(1, "call_fast", "fast_second"),
Event::tool_input_delta(1, r#"{}"#),
Event::tool_use_stop(1),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
],
vec![
Event::text_block_start(0),
Event::text_delta(0, "Done"),
Event::text_block_stop(0, None),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
],
]);
let client_probe = client.clone();
let mut engine = Engine::new(client);
engine.register_tool(SlowTool::new("slow_first", 100).definition());
engine.register_tool(SlowTool::new("fast_second", 5).definition());
let observed = Arc::new(Mutex::new(Vec::<String>::new()));
let published = observed.clone();
engine.on_tool_result(move |result| {
published
.lock()
.unwrap()
.push(format!("publish:{}", result.tool_use_id));
});
let committed = observed.clone();
let mut annotate = move |item: &Item| {
if let Item::ToolResult { call_id, .. } = item {
committed.lock().unwrap().push(format!("commit:{call_id}"));
}
Ok(())
};
let mut history = History::new();
let _ = engine
.run_with_annotation(&mut history, "run both", &mut annotate)
.await;
observed.lock().unwrap().push("run-returned".to_string());
assert_eq!(
observed.lock().unwrap().as_slice(),
[
"commit:call_fast",
"publish:call_fast",
"commit:call_slow",
"publish:call_slow",
"run-returned",
]
);
let committed_order: Vec<_> = history
.iter()
.filter_map(|entry| match &entry.item {
Item::ToolResult { call_id, .. } => Some(call_id.as_str()),
_ => None,
})
.collect();
assert_eq!(committed_order, ["call_fast", "call_slow"]);
let requests = client_probe.requests();
let projected_order: Vec<_> = requests[1]
.items
.iter()
.filter_map(|item| match item {
Item::ToolResult { call_id, .. } => Some(call_id.as_str()),
_ => None,
})
.collect();
assert_eq!(projected_order, ["call_slow", "call_fast"]);
}
#[tokio::test]
async fn cancellation_preserves_completed_results_and_resume_skips_them() {
let client = MockLlmClient::with_responses(vec![
vec![
Event::tool_use_start(0, "call_hang", "hang_once"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::tool_use_start(1, "call_fast_a", "fast_a"),
Event::tool_input_delta(1, r#"{}"#),
Event::tool_use_stop(1),
Event::tool_use_start(2, "call_fast_b", "fast_b"),
Event::tool_input_delta(2, r#"{}"#),
Event::tool_use_stop(2),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
],
vec![
Event::text_block_start(0),
Event::text_delta(0, "Recovered"),
Event::text_block_stop(0, None),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
],
]);
let mut engine = Engine::new(client);
let hanging = FirstAttemptHangsTool::new();
let fast_a = SlowTool::new("fast_a", 1);
let fast_b = SlowTool::new("fast_b", 2);
engine.register_tool(hanging.definition());
engine.register_tool(fast_a.definition());
engine.register_tool(fast_b.definition());
let cancel = engine.cancel_sender();
let cancel_task = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(30)).await;
cancel.send(()).await.unwrap();
});
let mut history = History::new();
let output = engine.run(&mut history, "start").await;
let mut engine = output.engine;
cancel_task.await.unwrap();
let completed_before_resume = history
.iter()
.filter(|entry| {
matches!(
&entry.item,
Item::ToolResult { call_id, .. }
if call_id == "call_fast_a" || call_id == "call_fast_b"
)
})
.count();
let unknown_before_resume = history
.iter()
.filter(|entry| {
matches!(
&entry.item,
Item::ToolResult {
call_id,
disposition: ToolResultDisposition::OutcomeUnknown,
..
} if call_id == "call_hang"
)
})
.count();
assert_eq!(completed_before_resume, 2);
assert_eq!(unknown_before_resume, 1);
assert_eq!(fast_a.call_count(), 1);
assert_eq!(fast_b.call_count(), 1);
assert_eq!(hanging.call_count(), 1);
let _ = engine.resume(&mut history).await;
assert_eq!(
fast_a.call_count(),
1,
"completed call must not be re-executed"
);
assert_eq!(
fast_b.call_count(),
1,
"completed call must not be re-executed"
);
assert_eq!(
hanging.call_count(),
1,
"OutcomeUnknown is terminal and must not be re-executed"
);
let completed_after_resume = history
.iter()
.filter(|entry| {
matches!(
&entry.item,
Item::ToolResult { call_id, .. }
if call_id == "call_fast_a" || call_id == "call_fast_b"
)
})
.count();
assert_eq!(completed_after_resume, 2);
assert_eq!(
history
.iter()
.filter(|entry| {
matches!(
&entry.item,
Item::ToolResult {
call_id,
disposition: ToolResultDisposition::OutcomeUnknown,
..
} if call_id == "call_hang"
)
})
.count(),
1
);
}
#[tokio::test]
async fn cooperative_cancellation_commits_bounded_terminal_output() {
let client = MockLlmClient::with_responses(vec![vec![
Event::tool_use_start(0, "call_cooperative", "cooperative"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]]);
let mut engine = Engine::new(client);
let tool = CooperativeCancelTool::new();
engine.register_tool(tool.definition());
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
let published = observed.clone();
engine.on_tool_result(move |_| published.lock().unwrap().push("published"));
let committed = observed.clone();
let mut annotate = move |item: &Item| {
if matches!(item, Item::ToolResult { .. }) {
committed.lock().unwrap().push("committed");
}
Ok(())
};
let cancel = engine.cancel_sender();
let cancel_task = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(30)).await;
cancel.send(()).await.unwrap();
});
let mut history = History::new();
let output = engine
.run_with_annotation(&mut history, "start", &mut annotate)
.await;
observed.lock().unwrap().push("run-returned");
cancel_task.await.unwrap();
assert_eq!(
observed.lock().unwrap().as_slice(),
["committed", "published", "run-returned"]
);
assert_eq!(tool.calls.load(Ordering::SeqCst), 1);
let terminal: Vec<_> = history
.iter()
.filter_map(|entry| match &entry.item {
Item::ToolResult {
call_id,
disposition,
content,
..
} if call_id == "call_cooperative" => Some((*disposition, content.as_deref())),
_ => None,
})
.collect();
assert_eq!(terminal.len(), 1);
assert_eq!(terminal[0].0, ToolResultDisposition::Cancelled);
assert_eq!(
terminal[0].1,
Some("stdout before cancellation\nstderr before cancellation")
);
assert!(matches!(
output.result,
agen::EngineRunExit::Interrupted(agen::StopReason::Cancelled)
));
}
#[tokio::test]
async fn cancellation_completion_race_commits_one_terminal_output() {
for iteration in 0..24u64 {
let client = MockLlmClient::with_responses(vec![
vec![
Event::tool_use_start(0, "call_racy", "racy"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
],
vec![Event::Status(StatusEvent {
status: ResponseStatus::Completed,
})],
]);
let mut engine = Engine::new(client);
let delay = 2 + iteration % 3;
let tool = SlowTool::new("racy", delay);
engine.register_tool(tool.definition());
let cancel = engine.cancel_sender();
let cancel_task = tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(delay)).await;
let _ = cancel.send(()).await;
});
let mut history = History::new();
let _ = engine.run(&mut history, "race").await;
cancel_task.await.unwrap();
let terminal_count = history
.iter()
.filter(|entry| {
matches!(
&entry.item,
Item::ToolResult { call_id, .. } if call_id == "call_racy"
)
})
.count();
assert_eq!(terminal_count, 1, "iteration {iteration}");
assert_eq!(tool.call_count(), 1, "iteration {iteration}");
}
}
#[tokio::test]
async fn tool_result_commit_failure_prevents_publication() {
let client = MockLlmClient::with_responses(vec![vec![
Event::tool_use_start(0, "call_fast", "fast"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]]);
let mut engine = Engine::new(client);
engine.register_tool(SlowTool::new("fast", 1).definition());
let published = Arc::new(AtomicUsize::new(0));
let published_probe = published.clone();
engine.on_tool_result(move |_| {
published_probe.fetch_add(1, Ordering::SeqCst);
});
let mut history = History::new();
let mut reject_tool_result = |item: &Item| {
if matches!(item, Item::ToolResult { .. }) {
Err("session log unavailable".to_string())
} else {
Ok(())
}
};
let _ = engine
.run_with_annotation(&mut history, "start", &mut reject_tool_result)
.await;
assert_eq!(published.load(Ordering::SeqCst), 0);
assert!(
history
.iter()
.all(|entry| !matches!(entry.item, Item::ToolResult { .. }))
);
}
#[tokio::test]
async fn test_tool_execution_context_order_and_batch_id() {
let client = MockLlmClient::with_responses(vec![
@@ -583,3 +1021,76 @@ async fn test_before_tool_call_synthetic_result_committed() {
} if call_id == "call_1" && summary == "permission denied"
)));
}
#[tokio::test]
async fn post_tool_abort_commits_confirmed_result_before_stopping_run() {
let client = MockLlmClient::new(vec![
Event::tool_use_start(0, "call_confirmed", "confirmed"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]);
let mut engine = Engine::new(client);
let tool = SlowTool::new("confirmed", 1);
engine.register_tool(tool.definition());
struct AbortAfterResult;
#[async_trait]
impl Interceptor for AbortAfterResult {
async fn post_tool_call(&self, _info: &mut ToolResultInfo) -> PostToolAction {
PostToolAction::Abort("policy stopped the run".to_string())
}
}
engine.set_interceptor(AbortAfterResult);
let observed = Arc::new(Mutex::new(Vec::<&'static str>::new()));
let published = observed.clone();
engine.on_tool_result(move |_| published.lock().unwrap().push("published"));
let committed = observed.clone();
let mut annotate = move |item: &Item| {
if matches!(item, Item::ToolResult { .. }) {
committed.lock().unwrap().push("committed");
}
Ok(())
};
let mut history = History::new();
let output = engine
.run_with_annotation(&mut history, "run confirmed tool", &mut annotate)
.await;
observed.lock().unwrap().push("run-returned");
assert_eq!(tool.call_count(), 1);
assert_eq!(
observed.lock().unwrap().as_slice(),
["committed", "published", "run-returned"]
);
assert!(matches!(
output.result,
agen::EngineRunExit::Interrupted(agen::StopReason::Unexpected(
agen::EngineError::Aborted(ref reason)
)) if reason == "policy stopped the run"
));
let terminal: Vec<_> = history
.iter()
.filter_map(|entry| match &entry.item {
Item::ToolResult {
call_id,
disposition,
..
} if call_id == "call_confirmed" => Some(*disposition),
_ => None,
})
.collect();
assert_eq!(terminal, [ToolResultDisposition::Success]);
assert!(!history.iter().any(|entry| matches!(
&entry.item,
Item::ToolResult {
call_id,
disposition: ToolResultDisposition::OutcomeUnknown,
..
} if call_id == "call_confirmed"
)));
}
+19
View File
@@ -352,6 +352,18 @@ pub struct InternalWorkerSnapshot {
pub internal_workers: Vec<InternalWorkerSnapshot>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
#[serde(rename_all = "snake_case")]
pub enum ToolResultDisposition {
#[default]
Success,
Error,
Interrupted,
Cancelled,
OutcomeUnknown,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[cfg_attr(feature = "typescript", derive(ts_rs::TS))]
#[serde(tag = "event", content = "data", rename_all = "snake_case")]
@@ -501,6 +513,8 @@ pub enum Event {
/// summary-only, or when the result was pruned.
#[serde(default, skip_serializing_if = "Option::is_none")]
output: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
disposition: Option<ToolResultDisposition>,
#[serde(default)]
is_error: bool,
},
@@ -1839,6 +1853,7 @@ mod tests {
id: "call_1".into(),
summary: "Read 128 bytes".into(),
output: Some("hello world".into()),
disposition: Some(ToolResultDisposition::Success),
is_error: false,
};
let json = serde_json::to_string(&event).unwrap();
@@ -1855,11 +1870,13 @@ mod tests {
id,
summary,
output,
disposition,
is_error,
} => {
assert_eq!(id, "call_1");
assert_eq!(summary, "Read 128 bytes");
assert_eq!(output.as_deref(), Some("hello world"));
assert_eq!(disposition, Some(ToolResultDisposition::Success));
assert!(!is_error);
}
other => panic!("expected ToolResult, got {other:?}"),
@@ -1872,6 +1889,7 @@ mod tests {
id: "call_2".into(),
summary: "ok".into(),
output: None,
disposition: Some(ToolResultDisposition::Success),
is_error: false,
};
let json = serde_json::to_string(&event).unwrap();
@@ -1887,6 +1905,7 @@ mod tests {
id: "call_3".into(),
summary: "invalid argument".into(),
output: None,
disposition: Some(ToolResultDisposition::Error),
is_error: true,
};
let json = serde_json::to_string(&event).unwrap();
+2 -1
View File
@@ -8,7 +8,7 @@ use crate::{
CompletionKind, ErrorCode, Event, Greeting, InFlightBlock, InFlightSnapshot,
InFlightToolCallState, InternalWorkerKind, InternalWorkerRef, InternalWorkerSnapshot,
InvokeKind, MemoryWorkerEvent, Method, Permission, RewindSummary, RewindTarget, RewindTargetId,
RunResult, ScopeRule, Segment, TurnResult, WorkerEvent, WorkerStatus,
RunResult, ScopeRule, Segment, ToolResultDisposition, TurnResult, WorkerEvent, WorkerStatus,
subscription::{
EventSubscriptionSelector, SubscriptionEvent, SubscriptionEventPayload, SubscriptionFrame,
SubscriptionFramePayload, SubscriptionId, SubscriptionRejectionCode, SubscriptionRequest,
@@ -45,6 +45,7 @@ pub fn generated_protocol_types() -> String {
push_decl::<TurnResult>(&cfg, &mut output);
push_decl::<InvokeKind>(&cfg, &mut output);
push_decl::<RunResult>(&cfg, &mut output);
push_decl::<ToolResultDisposition>(&cfg, &mut output);
push_decl::<ErrorCode>(&cfg, &mut output);
push_decl::<Permission>(&cfg, &mut output);
push_decl::<InFlightToolCallState>(&cfg, &mut output);
+58 -9
View File
@@ -14,7 +14,7 @@
use agen::{
llm_client::types::{ContentPart, Item, Role},
tool::{Attachment, ImageAttachment},
tool::{Attachment, ImageAttachment, ToolResultDisposition},
};
use base64::{Engine as _, engine::general_purpose::STANDARD};
use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error as _};
@@ -61,6 +61,8 @@ pub enum LoggedItem {
content: Option<String>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
attachments: Vec<LoggedAttachment>,
#[serde(default, skip_serializing_if = "ToolResultDisposition::is_success")]
disposition: ToolResultDisposition,
#[serde(default, skip_serializing_if = "is_false")]
is_error: bool,
},
@@ -128,6 +130,7 @@ impl From<&Item> for LoggedItem {
summary,
content,
attachments,
disposition,
is_error,
..
} => Self::ToolResult {
@@ -135,6 +138,7 @@ impl From<&Item> for LoggedItem {
summary: summary.clone(),
content: content.clone(),
attachments: attachments.iter().map(LoggedAttachment::from).collect(),
disposition: *disposition,
is_error: *is_error,
},
Item::Reasoning {
@@ -184,15 +188,24 @@ impl From<LoggedItem> for Item {
summary,
content,
attachments,
disposition,
is_error,
} => Item::ToolResult {
id: None,
call_id,
summary,
content,
is_error,
attachments: attachments.into_iter().map(Attachment::from).collect(),
},
} => {
let disposition = if is_error && disposition.is_success() {
ToolResultDisposition::Error
} else {
disposition
};
Item::ToolResult {
id: None,
call_id,
summary,
content,
disposition,
is_error,
attachments: attachments.into_iter().map(Attachment::from).collect(),
}
}
LoggedItem::Reasoning {
text,
summary,
@@ -430,6 +443,42 @@ mod tests {
}
}
#[test]
fn outcome_unknown_tool_result_round_trips_as_terminal() {
let original = Item::tool_result_item_with_disposition_and_attachments(
"call_unknown",
"outcome unknown",
Some("bounded progress".to_string()),
ToolResultDisposition::OutcomeUnknown,
Vec::new(),
);
let logged: LoggedItem = (&original).into();
let json = serde_json::to_string(&logged).unwrap();
assert!(json.contains(r#""disposition":"outcome_unknown""#));
match Item::from(serde_json::from_str::<LoggedItem>(&json).unwrap()) {
Item::ToolResult {
disposition,
is_error,
..
} => {
assert_eq!(disposition, ToolResultDisposition::OutcomeUnknown);
assert!(is_error);
}
other => panic!("unexpected variant: {other:?}"),
}
}
#[test]
fn legacy_error_tool_result_infers_error_disposition() {
let legacy = r#"{"kind":"tool_result","call_id":"call_old","summary":"failed","content":null,"is_error":true}"#;
match Item::from(serde_json::from_str::<LoggedItem>(legacy).unwrap()) {
Item::ToolResult { disposition, .. } => {
assert_eq!(disposition, ToolResultDisposition::Error)
}
other => panic!("unexpected variant: {other:?}"),
}
}
#[test]
fn tool_result_persistence_round_trips_binary_attachments() {
let original = Item::tool_result_item_with_attachments(
+68 -8
View File
@@ -1,5 +1,6 @@
use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::sync::Arc;
use std::sync::{Arc, Mutex};
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta, ToolOutput};
use async_trait::async_trait;
@@ -20,15 +21,28 @@ struct BashParams {
pub(crate) struct BashTool {
session: WorkdirSessionHandle,
state: Arc<Mutex<BashExecutionState>>,
}
#[derive(Default)]
struct BashExecutionState {
active: HashMap<String, CommandHandle>,
cancellation_requested: HashSet<String>,
}
struct CommandGuard {
session: WorkdirSessionHandle,
state: Arc<Mutex<BashExecutionState>>,
call_id: String,
handle: Option<CommandHandle>,
}
impl Drop for CommandGuard {
fn drop(&mut self) {
let mut state = self.state.lock().unwrap();
state.active.remove(&self.call_id);
state.cancellation_requested.remove(&self.call_id);
drop(state);
if let Some(handle) = self.handle.take() {
let workdir = self.session.clone();
tokio::spawn(async move {
@@ -52,20 +66,35 @@ impl Tool for BashTool {
.unwrap_or(DEFAULT_TIMEOUT_SECS)
.clamp(1, MAX_TIMEOUT_SECS);
let cmd_summary = truncate_for_summary(&params.command);
let call_id = ctx.call_id;
let mut guard = CommandGuard {
session: self.session.clone(),
state: self.state.clone(),
call_id: call_id.clone(),
handle: None,
};
let handle = self
.session
.start_command(CommandRequest {
command: params.command,
timeout_secs,
output_limit: INLINE_BYTE_BUDGET,
tool_call_id: Some(ctx.call_id),
tool_call_id: Some(call_id.clone()),
})
.await
.map_err(crate::ToolsError::from)?;
let mut guard = CommandGuard {
session: self.session.clone(),
handle: Some(handle.clone()),
let cancel_after_start = {
let mut state = self.state.lock().unwrap();
state.active.insert(call_id.clone(), handle.clone());
state.cancellation_requested.contains(&call_id)
};
guard.handle = Some(handle.clone());
if cancel_after_start {
self.session
.cancel_command(handle.clone())
.await
.map_err(crate::ToolsError::from)?;
}
let output = self
.session
.command_output(CommandOutputRequest {
@@ -76,9 +105,17 @@ impl Tool for BashTool {
})
.await
.map_err(crate::ToolsError::from)?;
let cancellation_requested = {
let mut state = self.state.lock().unwrap();
state.active.remove(&call_id);
state.cancellation_requested.remove(&call_id)
};
guard.handle = None;
let summary = if output.timed_out {
let timed_out = output.timed_out;
let summary = if cancellation_requested {
format!("$ {cmd_summary} (cancelled)")
} else if output.timed_out {
format!("$ {cmd_summary} (timed out after {timeout_secs}s)")
} else {
match output.exit_code {
@@ -97,11 +134,33 @@ impl Tool for BashTool {
} else {
Some(output.content)
};
Ok(ToolOutput {
let output = ToolOutput {
summary,
content,
attachments: Vec::new(),
})
};
if cancellation_requested {
Err(ToolError::Cancelled(output))
} else if timed_out {
Err(ToolError::Interrupted(output))
} else {
Ok(output)
}
}
async fn cancel(&self, call_id: &str) -> Result<(), ToolError> {
let handle = {
let mut state = self.state.lock().unwrap();
state.cancellation_requested.insert(call_id.to_string());
state.active.get(call_id).cloned()
};
if let Some(handle) = handle {
self.session
.cancel_command(handle)
.await
.map_err(crate::ToolsError::from)?;
}
Ok(())
}
}
@@ -123,6 +182,7 @@ pub fn bash_tool(session: WorkdirSessionHandle, _output_dir: PathBuf) -> ToolDef
.input_schema(serde_json::to_value(schema).expect("Bash schema serialization"));
let tool: Arc<dyn Tool> = Arc::new(BashTool {
session: session.clone(),
state: Arc::new(Mutex::new(BashExecutionState::default())),
});
(meta, tool)
})
+41 -1
View File
@@ -7,7 +7,7 @@
use std::path::Path;
use std::sync::Arc;
use agen::tool::{Tool, ToolDefinition, ToolMeta};
use agen::tool::{Tool, ToolDefinition, ToolError, ToolMeta};
use manifest::{Permission, Scope, ScopeConfig, ScopeRule};
use serde_json::json;
use tempfile::TempDir;
@@ -401,5 +401,45 @@ async fn bash_provider_output_does_not_expose_internal_paths() {
assert_eq!(std::fs::read_dir(spill.path()).unwrap().count(), 0);
}
#[tokio::test]
async fn bash_cancellation_returns_bounded_progress_as_terminal_output() {
let (_dir, _spill, reg) = setup();
let bash = reg.get("Bash");
let executing = bash.clone();
let execution = tokio::spawn(async move {
executing
.execute(
r#"{"command":"printf 'before\\n'; printf 'err-before\\n' >&2; sleep 5; printf 'after\\n'"}"#,
Default::default(),
)
.await
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
bash.cancel("direct").await.expect("signal cancellation");
let error = tokio::time::timeout(std::time::Duration::from_secs(2), execution)
.await
.expect("cancelled Bash should terminate inside the Engine grace budget")
.expect("Bash task join");
let ToolError::Cancelled(output) = error.expect_err("cancelled command is non-success") else {
panic!("expected typed cancellation result");
};
let content = output.content.expect("bounded progress output");
assert!(
content.contains("before"),
"missing pre-cancel stdout: {content}"
);
assert!(
content.contains("err-before"),
"missing pre-cancel stderr: {content}"
);
assert!(
!content.contains("after"),
"post-cancel output leaked: {content}"
);
assert!(content.len() <= 16 * 1024, "output must remain bounded");
}
// Sanity: unused Path import guard
const _: fn() -> &'static Path = || Path::new("/");
+1
View File
@@ -1244,6 +1244,7 @@ impl App {
id,
summary,
output,
disposition: _,
is_error,
} => {
self.latest_llm_wait_event = None;
+13
View File
@@ -763,6 +763,19 @@ pub(crate) fn wire_event_bridges_on_engine<C, St>(
id: result.tool_use_id.clone(),
summary: result.summary.clone(),
output: result.content.clone(),
disposition: Some(match result.disposition {
agen::ToolResultDisposition::Success => protocol::ToolResultDisposition::Success,
agen::ToolResultDisposition::Error => protocol::ToolResultDisposition::Error,
agen::ToolResultDisposition::Interrupted => {
protocol::ToolResultDisposition::Interrupted
}
agen::ToolResultDisposition::Cancelled => {
protocol::ToolResultDisposition::Cancelled
}
agen::ToolResultDisposition::OutcomeUnknown => {
protocol::ToolResultDisposition::OutcomeUnknown
}
}),
is_error: result.is_error,
});
});
+13 -2
View File
@@ -13,7 +13,7 @@
#[cfg(test)]
use crate::prompt::catalog::PromptCatalog;
use agen::Item;
use agen::{Item, ToolResultDisposition};
/// Build synthetic `Item::ToolResult` items for every unanswered
/// `Item::ToolCall` in `history`, preserving order.
@@ -28,7 +28,16 @@ pub(crate) fn orphan_tool_result_closures(history: &[Item], summary: &str) -> Ve
for item in history {
if let Item::ToolCall { call_id, .. } = item {
if !answered.contains(call_id.as_str()) {
out.push(Item::tool_result(call_id.clone(), summary));
out.push(Item::tool_result_item_with_disposition_and_attachments(
call_id.clone(),
summary,
Some(
"Execution ended before completion could be confirmed. Completion and side effects are unknown."
.to_string(),
),
ToolResultDisposition::OutcomeUnknown,
Vec::new(),
));
}
}
}
@@ -77,10 +86,12 @@ mod tests {
Item::ToolResult {
call_id,
summary: got,
disposition,
..
} => {
assert_eq!(call_id, "c1");
assert_eq!(got, &summary);
assert_eq!(*disposition, ToolResultDisposition::OutcomeUnknown);
}
other => panic!("expected ToolResult, got {other:?}"),
}
+326 -34
View File
@@ -2953,50 +2953,52 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
Ok(())
}
/// Stage the post-interruption cleanup at the front of worker
/// history: close every unanswered `Item::ToolCall` with a synthetic
/// `Item::ToolResult` (Anthropic wire-validity), then append a
/// system note so the LLM understands the prior turn was cut
/// short. Called from `Worker::run` when the worker's
/// `last_run_interrupted` flag is set (i.e. the Worker just transitioned
/// out of Paused via a new user input).
fn apply_interrupt_prep(&mut self) -> Result<(), WorkerError> {
/// Durably close every unanswered ToolCall before the interrupted run's
/// final lifecycle record/status is published.
fn terminalize_orphan_tool_calls(&mut self) -> Result<(), WorkerError> {
let tool_result_summary = self
.prompts()
.load_full()
.interrupt_tool_result_summary()
.map_err(WorkerError::from)?;
let history_items = self.history();
let closures = crate::interrupt_prep::orphan_tool_result_closures(
&history_items,
&tool_result_summary,
);
if closures.is_empty() {
return Ok(());
}
let subject = worker_subject(self.session.session_id());
for item in closures {
let entry = HistoryEntry::new(
item,
new_history_metadata(
WorkerHistoryProvenance::ToolOutput {
worker: subject.clone(),
},
None,
),
);
self.commit_entry(LogEntry::AnnotatedToolResult {
ts: segment_log::now_millis(),
entry: to_logged_history_entry(&entry),
})?;
self.session.history_mut().push_entry(entry);
self.session.note_mutation();
}
Ok(())
}
fn apply_interrupt_prep(&mut self) -> Result<(), WorkerError> {
self.terminalize_orphan_tool_calls()?;
let system_note = self
.prompts()
.load_full()
.interrupt_system_note()
.map_err(WorkerError::from)?;
let history_items = self.history();
let closures = crate::interrupt_prep::orphan_tool_result_closures(
&history_items,
&tool_result_summary,
);
if !closures.is_empty() {
let subject = worker_subject(self.session.session_id());
for item in closures {
let entry = HistoryEntry::new(
item,
new_history_metadata(
WorkerHistoryProvenance::ToolOutput {
worker: subject.clone(),
},
None,
),
);
self.commit_entry(LogEntry::AnnotatedToolResult {
ts: segment_log::now_millis(),
entry: to_logged_history_entry(&entry),
})?;
self.session.history_mut().push_entry(entry);
self.session.note_mutation();
}
}
let interrupt_prompt_provenance =
self.prompt_render_provenance("internal.interrupt_system_note");
let interrupt_metadata = new_history_metadata(
@@ -3262,6 +3264,9 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
where
St: Clone + 'static,
{
if matches!(&result, EngineRunExit::Interrupted(_)) {
self.terminalize_orphan_tool_calls()?;
}
self.persist_turn(history_before, &result).await?;
if matches!(result, EngineRunExit::Yielded) {
@@ -5928,7 +5933,8 @@ fn stop_reason_error_code(reason: &StopReason) -> ErrorCode {
EngineError::Aborted(_)
| EngineError::Cancelled
| EngineError::ConfigWarnings(_)
| EngineError::HistoryAppend(_),
| EngineError::HistoryAppend(_)
| EngineError::ToolAttemptFence(_),
) => ErrorCode::Internal,
}
}
@@ -7442,6 +7448,118 @@ mod build_summary_prompt_tests {
}
}
#[derive(Clone)]
struct PauseResumeClient {
calls: Arc<std::sync::atomic::AtomicUsize>,
}
impl PauseResumeClient {
fn new() -> Self {
Self {
calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
}
}
}
#[async_trait::async_trait]
impl agen::llm_client::LlmClient for PauseResumeClient {
async fn stream(
&self,
_request: agen::llm_client::Request,
) -> Result<
std::pin::Pin<
Box<
dyn futures::Stream<
Item = Result<agen::llm_client::Event, agen::llm_client::ClientError>,
> + Send,
>,
>,
agen::llm_client::ClientError,
> {
use agen::llm_client::{Event, ResponseStatus, StatusEvent};
let call = self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let events = if call == 0 {
vec![
Event::tool_use_start(0, "call_pending", "pending_once"),
Event::tool_input_delta(0, r#"{}"#),
Event::tool_use_stop(0),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]
} else {
vec![
Event::text_block_start(0),
Event::text_delta(0, "done"),
Event::text_block_stop(0, None),
Event::Status(StatusEvent {
status: ResponseStatus::Completed,
}),
]
};
Ok(Box::pin(futures::stream::iter(events.into_iter().map(Ok))))
}
fn clone_boxed(&self) -> Box<dyn agen::llm_client::LlmClient> {
Box::new(self.clone())
}
}
#[derive(Clone)]
struct CountingPendingTool {
calls: Arc<std::sync::atomic::AtomicUsize>,
}
#[async_trait::async_trait]
impl agen::tool::Tool for CountingPendingTool {
async fn execute(
&self,
_input_json: &str,
_ctx: agen::ToolExecutionContext,
) -> Result<agen::tool::ToolOutput, agen::tool::ToolError> {
self.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok("executed once".to_string().into())
}
}
fn counting_pending_tool(
calls: Arc<std::sync::atomic::AtomicUsize>,
) -> agen::tool::ToolDefinition {
Arc::new(move || {
let meta = agen::tool::ToolMeta::new("pending_once")
.description("Counts resumable pending execution")
.input_schema(serde_json::json!({"type": "object"}));
(
meta,
Arc::new(CountingPendingTool {
calls: calls.clone(),
}) as Arc<dyn agen::tool::Tool>,
)
})
}
#[derive(Clone)]
struct PauseOnceHook {
should_pause: Arc<std::sync::atomic::AtomicBool>,
}
#[async_trait::async_trait]
impl crate::hook::Hook<crate::hook::PreToolCall> for PauseOnceHook {
async fn call(
&self,
_input: &crate::hook::ToolCallSummary,
) -> crate::hook::HookPreToolAction {
if self
.should_pause
.swap(false, std::sync::atomic::Ordering::SeqCst)
{
crate::hook::HookPreToolAction::Pause
} else {
crate::hook::HookPreToolAction::Continue
}
}
}
#[derive(Clone)]
struct NoopClient;
@@ -7862,6 +7980,7 @@ mod build_summary_prompt_tests {
summary: "wrote a file".into(),
content: None,
attachments: Vec::new(),
disposition: Default::default(),
is_error: false,
},
},
@@ -7904,6 +8023,7 @@ mod build_summary_prompt_tests {
summary: "wrote a file".into(),
content: None,
attachments: Vec::new(),
disposition: Default::default(),
is_error: false,
},
},
@@ -7948,6 +8068,7 @@ mod build_summary_prompt_tests {
summary: "side effect".into(),
content: None,
attachments: Vec::new(),
disposition: Default::default(),
is_error: false,
},
metadata: new_history_metadata(
@@ -8047,6 +8168,177 @@ mod build_summary_prompt_tests {
assert!(err.contains("session head changed"));
}
#[tokio::test]
async fn hook_paused_pending_tool_resumes_and_executes_once() {
let dir = tempfile::tempdir().unwrap();
let store = session_store::FsStore::new(dir.path().join("sessions")).unwrap();
let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let mut engine =
Engine::<_, Mutable, SessionHistoryMetadata>::new_annotated(PauseResumeClient::new());
engine.register_tool(counting_pending_tool(calls.clone()));
let mut worker = Worker::new(
minimal_manifest(),
engine,
store,
WorkerWorkspaceContext::no_workspace(),
WorkerFilesystemAuthority::None,
Scope::empty(),
)
.await
.unwrap();
let should_pause = Arc::new(std::sync::atomic::AtomicBool::new(true));
worker.add_pre_tool_call_hook(PauseOnceHook { should_pause });
assert_eq!(
worker.run_text("start").await.unwrap(),
WorkerRunResult::Paused
);
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0);
assert!(worker.history().iter().any(|item| matches!(
item,
Item::ToolCall { call_id, .. } if call_id == "call_pending"
)));
assert!(!worker.history().iter().any(|item| matches!(
item,
Item::ToolResult { call_id, .. } if call_id == "call_pending"
)));
assert_eq!(worker.resume().await.unwrap(), WorkerRunResult::Finished);
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
assert_eq!(
worker
.history()
.iter()
.filter(|item| matches!(
item,
Item::ToolResult {
call_id,
disposition: agen::ToolResultDisposition::Success,
..
} if call_id == "call_pending"
))
.count(),
1
);
}
#[tokio::test]
async fn interrupted_result_terminalizes_orphan_before_run_completed() {
let dir = tempfile::tempdir().unwrap();
let manifest = minimal_manifest();
let store = session_store::FsStore::new(dir.path().join("sessions")).unwrap();
let cwd = dir.path().join("workspace");
std::fs::create_dir_all(&cwd).unwrap();
let scope = Scope::writable(&cwd).unwrap();
let authority = WorkerFilesystemAuthority::local(cwd.clone(), cwd.clone());
let mut worker = Worker::new(
manifest,
Engine::<_, Mutable, SessionHistoryMetadata>::new_annotated(NoopClient),
store,
WorkerWorkspaceContext::local_filesystem(None),
authority,
scope,
)
.await
.unwrap();
worker.ensure_segment_head().unwrap();
worker.wire_history_persistence();
worker.set_history_for_test(vec![
Item::tool_call("call-known", "Read", "{}"),
Item::tool_result_item_with_disposition_and_attachments(
"call-known",
"known result",
Some("confirmed output".to_string()),
agen::ToolResultDisposition::Success,
Vec::new(),
),
Item::tool_call("call-orphan", "Bash", "{}"),
]);
let _ = worker
.handle_worker_result(
EngineRunExit::Interrupted(StopReason::Cancelled),
worker.history().len(),
)
.await
.unwrap();
let history = worker.history();
assert_eq!(
history
.iter()
.filter(|item| matches!(
item,
Item::ToolResult {
call_id,
disposition: agen::ToolResultDisposition::Success,
..
} if call_id == "call-known"
))
.count(),
1
);
assert!(!history.iter().any(|item| matches!(
item,
Item::ToolResult {
call_id,
disposition: agen::ToolResultDisposition::OutcomeUnknown,
..
} if call_id == "call-known"
)));
assert_eq!(
history
.iter()
.filter(|item| matches!(
item,
Item::ToolResult {
call_id,
disposition: agen::ToolResultDisposition::OutcomeUnknown,
..
} if call_id == "call-orphan"
))
.count(),
1
);
let entries = worker
.store
.read_all(
worker.segment_state.session_id(),
worker.segment_state.segment_id(),
)
.unwrap();
let terminal_index = entries
.iter()
.position(|entry| {
matches!(
entry,
LogEntry::AnnotatedToolResult {
entry: session_store::LoggedHistoryEntry {
item: session_store::LoggedItem::ToolResult {
call_id,
disposition: agen::ToolResultDisposition::OutcomeUnknown,
..
},
..
},
..
} if call_id == "call-orphan"
)
})
.expect("durable OutcomeUnknown closure");
let final_index = entries
.iter()
.position(|entry| {
matches!(
entry,
LogEntry::RunCompleted { .. } | LogEntry::RunErrored { .. }
)
})
.expect("durable final run status");
assert!(terminal_index < final_index);
}
#[tokio::test]
async fn apply_interrupt_prep_appends_via_callback_and_logs_independent_entries() {
let dir = tempfile::tempdir().unwrap();
+11 -4
View File
@@ -2169,9 +2169,13 @@ async fn paused_then_run_closes_orphan_tool_use_for_next_request() {
for item in items {
match item {
agen::Item::ToolResult {
call_id, summary, ..
call_id,
summary,
disposition,
..
} if call_id == "call_orphan" => {
assert_eq!(summary, "[Interrupted by user]");
assert_eq!(summary, "Tool execution outcome unknown");
assert_eq!(*disposition, agen::ToolResultDisposition::OutcomeUnknown);
saw_synthetic_tool_result = true;
}
agen::Item::Message { role, content, .. } if *role == agen::Role::System => {
@@ -2362,8 +2366,11 @@ async fn paused_cancel_abandons_resume_and_next_input_is_fresh_run() {
assert!(
items.iter().any(|item| matches!(
item,
agen::Item::ToolResult { call_id, summary, .. }
if call_id == "call_cancelled" && summary == "[Interrupted by user]"
agen::Item::ToolResult {
call_id,
disposition: agen::ToolResultDisposition::OutcomeUnknown,
..
} if call_id == "call_cancelled"
)),
"paused cancel should close orphan tool_use before future requests: {items:?}"
);
+3 -1
View File
@@ -16,6 +16,8 @@ export type InvokeKind = "user_send" | "notify" | "worker_event" | "system_remin
export type RunResult = "finished" | "paused" | "limit_reached" | "rolled_back";
export type ToolResultDisposition = "success" | "error" | "interrupted" | "cancelled" | "outcome_unknown";
export type ErrorCode = "already_running" | "not_running" | "not_paused" | "provider_error" | "tool_error" | "invalid_request" | "internal";
export type Permission = "read" | "write";
@@ -191,7 +193,7 @@ summary: string,
* Full tool output. Absent when the tool chose to return
* summary-only, or when the result was pruned.
*/
output?: string | null, is_error: boolean, } } | { "event": "usage", "data": { input_tokens: number | null, output_tokens: number | null, cache_read_input_tokens?: number | null, } } | { "event": "run_end", "data": { result: RunResult, } } | { "event": "error", "data": { code: ErrorCode, message: string, } } | { "event": "snapshot", "data": { entries: Array<unknown>, greeting: Greeting, status: WorkerStatus,
output?: string | null, disposition?: ToolResultDisposition | null, is_error: boolean, } } | { "event": "usage", "data": { input_tokens: number | null, output_tokens: number | null, cache_read_input_tokens?: number | null, } } | { "event": "run_end", "data": { result: RunResult, } } | { "event": "error", "data": { code: ErrorCode, message: string, } } | { "event": "snapshot", "data": { entries: Array<unknown>, greeting: Greeting, status: WorkerStatus,
/**
* Unfinished model output that has already streamed in the current
* run but is not yet represented by committed snapshot entries.