Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
62ada5eaa4 | ||
|
|
f5ff0b7c13 | ||
|
|
3337cafcdf | ||
|
|
e87784118b | ||
|
|
40fada28ea | ||
|
|
58cc94d4b7 | ||
|
|
ccabea59c9 | ||
|
|
183c37446e |
+586
-117
@@ -1,9 +1,10 @@
|
||||
use std::collections::HashMap;
|
||||
use std::{marker::PhantomData, sync::Arc, time::Instant};
|
||||
use std::collections::{HashMap, HashSet};
|
||||
use std::{future::Future, marker::PhantomData, pin::Pin, sync::Arc, time::Instant};
|
||||
|
||||
use futures::StreamExt;
|
||||
use serde_json::{Value, json};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::time::Instant as TokioInstant;
|
||||
use tracing::{debug, info, trace, warn};
|
||||
|
||||
use crate::{
|
||||
@@ -27,7 +28,8 @@ use crate::{
|
||||
timeline::{TextBlockCollector, ThinkingBlockCollector, Timeline, ToolCallCollector},
|
||||
tool::{
|
||||
ToolCall, ToolDefinition as EngineToolDefinition, ToolError, ToolExecutionContext,
|
||||
ToolOutputLimits, ToolResult, truncate_content,
|
||||
ToolExecutionHandle, ToolExecutionPolicy, ToolExecutionTerminal, ToolOutputLimits,
|
||||
ToolResult, ToolResultDisposition, truncate_content,
|
||||
},
|
||||
tool_server::{ToolServer, ToolServerHandle},
|
||||
};
|
||||
@@ -47,12 +49,18 @@ pub enum EngineError {
|
||||
/// Cancelled by CancellationToken
|
||||
#[error("Cancelled")]
|
||||
Cancelled,
|
||||
/// Paused by the caller at the next safe boundary.
|
||||
#[error("Paused")]
|
||||
PauseRequested,
|
||||
/// Config warnings (unsupported options)
|
||||
#[error("Config warnings: {}", .0.iter().map(|w| w.to_string()).collect::<Vec<_>>().join(", "))]
|
||||
ConfigWarnings(Vec<ConfigWarning>),
|
||||
/// 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 +78,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")]
|
||||
@@ -109,6 +170,7 @@ impl From<Result<EngineResult, EngineError>> for EngineRunExit {
|
||||
Self::Interrupted(StopReason::ContextWindowExceeded)
|
||||
}
|
||||
Err(EngineError::Cancelled) => Self::Interrupted(StopReason::Cancelled),
|
||||
Err(EngineError::PauseRequested) => Self::Paused,
|
||||
Err(error) => Self::Interrupted(StopReason::Unexpected(error)),
|
||||
}
|
||||
}
|
||||
@@ -126,10 +188,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
|
||||
@@ -228,6 +348,8 @@ pub struct Engine<C: LlmClient, S: EngineState = Mutable, A = ()> {
|
||||
tool_execution_batch_count: usize,
|
||||
/// Maximum number of AgentTurns (None = unlimited)
|
||||
max_turns: Option<u32>,
|
||||
/// Caller-selected policy for interrupting started provider operations.
|
||||
tool_execution_policy: ToolExecutionPolicy,
|
||||
/// AgentTurn-start callbacks (1:1 with LlmCall today)
|
||||
turn_start_cbs: Vec<Box<dyn Fn(usize) + Send + Sync>>,
|
||||
/// AgentTurn-end callbacks (1:1 with LlmCall today)
|
||||
@@ -266,6 +388,10 @@ pub struct Engine<C: LlmClient, S: EngineState = Mutable, A = ()> {
|
||||
/// Cancel notification channel (for interrupting execution)
|
||||
cancel_tx: mpsc::Sender<()>,
|
||||
cancel_rx: mpsc::Receiver<()>,
|
||||
/// Pause notification channel. Unlike cancellation, this waits for already
|
||||
/// started tools to reach provider-confirmed terminal results.
|
||||
pause_tx: mpsc::Sender<()>,
|
||||
pause_rx: mpsc::Receiver<()>,
|
||||
/// Byte-size caps applied to tool `content` before it reaches history.
|
||||
/// `None` disables truncation (tests and minimal setups).
|
||||
tool_output_limits: Option<ToolOutputLimits>,
|
||||
@@ -303,7 +429,10 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
|
||||
fn finish_logical_run(&mut self, result: &Result<EngineResult, EngineError>) {
|
||||
if !matches!(result, Ok(EngineResult::Paused) | Ok(EngineResult::Yielded)) {
|
||||
if !matches!(
|
||||
result,
|
||||
Ok(EngineResult::Paused | EngineResult::Yielded) | Err(EngineError::PauseRequested)
|
||||
) {
|
||||
self.active_run_turn_count = None;
|
||||
}
|
||||
}
|
||||
@@ -312,14 +441,19 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
while self.cancel_rx.try_recv().is_ok() {}
|
||||
}
|
||||
|
||||
/// Discard pending cancellation notifications while the engine is idle.
|
||||
fn drain_pause_queue(&mut self) {
|
||||
while self.pause_rx.try_recv().is_ok() {}
|
||||
}
|
||||
|
||||
/// Discard pending interruption notifications while the engine is idle.
|
||||
///
|
||||
/// Cancellation is a running-turn control signal. Callers that own a higher
|
||||
/// level run state can use this before starting a new turn so an old idle
|
||||
/// signal does not poison the next request, while cancellation queued after
|
||||
/// the run has been accepted remains observable by the turn loop.
|
||||
/// Cancellation and pause are running-turn control signals. Callers that own
|
||||
/// a higher level run state can use this before starting a new turn so an old
|
||||
/// idle signal does not poison the next request, while interruption queued
|
||||
/// after the run has been accepted remains observable by the turn loop.
|
||||
pub fn clear_pending_cancel(&mut self) {
|
||||
self.drain_cancel_queue();
|
||||
self.drain_pause_queue();
|
||||
}
|
||||
|
||||
fn try_cancelled(&mut self) -> bool {
|
||||
@@ -331,6 +465,14 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
}
|
||||
|
||||
fn try_paused(&mut self) -> bool {
|
||||
use tokio::sync::mpsc::error::TryRecvError;
|
||||
match self.pause_rx.try_recv() {
|
||||
Ok(()) => true,
|
||||
Err(TryRecvError::Empty | TryRecvError::Disconnected) => false,
|
||||
}
|
||||
}
|
||||
|
||||
/// Register a text block observer with scoped callbacks.
|
||||
///
|
||||
/// The setup closure is called once per text block. Inside it, register
|
||||
@@ -794,6 +936,22 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
self.cancel_tx.clone()
|
||||
}
|
||||
|
||||
/// Get the safe-boundary pause notification sender.
|
||||
pub fn pause_sender(&self) -> mpsc::Sender<()> {
|
||||
self.pause_tx.clone()
|
||||
}
|
||||
|
||||
/// Select the deadline policy applied to already-started provider operations.
|
||||
/// Worker/controller layers own this lifecycle policy; Agen owns only the
|
||||
/// mechanical terminalization of each call.
|
||||
pub fn set_tool_execution_policy(&mut self, policy: ToolExecutionPolicy) {
|
||||
self.tool_execution_policy = policy;
|
||||
}
|
||||
|
||||
pub fn tool_execution_policy(&self) -> ToolExecutionPolicy {
|
||||
self.tool_execution_policy
|
||||
}
|
||||
|
||||
/// Set request configuration at once
|
||||
pub fn set_request_config(&mut self, config: RequestConfig) {
|
||||
self.request_config = config;
|
||||
@@ -892,8 +1050,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 +1070,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 +1142,17 @@ 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};
|
||||
|
||||
// A pause observed before provider ownership starts leaves every call
|
||||
// NotStarted and therefore eligible for an explicit later retry.
|
||||
if self.try_paused() {
|
||||
return Ok(ToolExecutionResult::Paused);
|
||||
}
|
||||
|
||||
// Map from tool call ID to (ToolCall, Meta, Tool, Context)
|
||||
// Retained because it's needed for PostToolCall hooks
|
||||
@@ -1039,110 +1211,340 @@ 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
|
||||
.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
|
||||
{
|
||||
Ok(output) => ToolResult::from_output(&tool_call.id, output),
|
||||
Err(e) => ToolResult::error(&tool_call.id, e.to_string()),
|
||||
}
|
||||
}
|
||||
})
|
||||
// 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();
|
||||
|
||||
// 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");
|
||||
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<Pin<Box<dyn Future<Output = (String, ToolResult)> + Send>>> =
|
||||
FuturesUnordered::new();
|
||||
let mut execution_handles = HashMap::new();
|
||||
for (tool_call, context, tool) in approved_calls {
|
||||
let attempt_id = context.batch_id.clone();
|
||||
let input_json = serde_json::to_string(&tool_call.input).unwrap_or_default();
|
||||
let call_id = tool_call.id.clone();
|
||||
let future: Pin<Box<dyn Future<Output = (String, ToolResult)> + Send>> = match tool {
|
||||
None => {
|
||||
let result =
|
||||
ToolResult::error(&call_id, format!("Tool not found: {}", tool_call.name));
|
||||
Box::pin(async move { (attempt_id, result) })
|
||||
}
|
||||
self.timeline.abort_current_block();
|
||||
return Err(EngineError::Cancelled);
|
||||
}
|
||||
};
|
||||
results.extend(synthetic_results);
|
||||
Some(tool) => {
|
||||
let (handle, terminal) = ToolExecutionHandle::start(tool, input_json, context);
|
||||
execution_handles.insert(call_id.clone(), handle);
|
||||
Box::pin(async move {
|
||||
let result = match terminal.await {
|
||||
ToolExecutionTerminal::Confirmed(Ok(output)) => {
|
||||
ToolResult::from_output(&call_id, output)
|
||||
}
|
||||
ToolExecutionTerminal::Confirmed(Err(ToolError::Cancelled(output))) => {
|
||||
ToolResult::from_output_with_disposition(
|
||||
&call_id,
|
||||
output,
|
||||
ToolResultDisposition::Cancelled,
|
||||
)
|
||||
}
|
||||
ToolExecutionTerminal::Confirmed(Err(ToolError::Interrupted(
|
||||
output,
|
||||
))) => ToolResult::from_output_with_disposition(
|
||||
&call_id,
|
||||
output,
|
||||
ToolResultDisposition::Interrupted,
|
||||
),
|
||||
ToolExecutionTerminal::Confirmed(Err(error)) => {
|
||||
ToolResult::error(&call_id, error.to_string())
|
||||
}
|
||||
ToolExecutionTerminal::OutcomeUnknown => {
|
||||
ToolResult::outcome_unknown(&call_id)
|
||||
}
|
||||
};
|
||||
(attempt_id, result)
|
||||
})
|
||||
}
|
||||
};
|
||||
futures.push(future);
|
||||
}
|
||||
|
||||
// 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(),
|
||||
};
|
||||
// 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();
|
||||
let mut pause_requested = false;
|
||||
let mut pause_deadline = None;
|
||||
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?;
|
||||
}
|
||||
|
||||
match self.interceptor.post_tool_call(&mut info).await {
|
||||
PostToolAction::Continue => {}
|
||||
PostToolAction::Abort(reason) => {
|
||||
return Err(EngineError::Aborted(reason));
|
||||
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?;
|
||||
}
|
||||
pause = self.pause_rx.recv(), if !pause_requested => {
|
||||
if pause.is_some() {
|
||||
// Pause first waits for already-started tools to reach a
|
||||
// natural safe boundary. If they do not, Worker policy
|
||||
// escalates to the same explicit cancel-and-confirm path.
|
||||
pause_requested = true;
|
||||
pause_deadline = Some(
|
||||
TokioInstant::now()
|
||||
+ self.tool_execution_policy.pause_safe_boundary_timeout,
|
||||
);
|
||||
}
|
||||
}
|
||||
// Reflect interceptor-modified results
|
||||
*tool_result = info.result;
|
||||
}
|
||||
}
|
||||
_ = tokio::time::sleep_until(pause_deadline.unwrap_or_else(TokioInstant::now)), if pause_deadline.is_some() => {
|
||||
pause_deadline = None;
|
||||
let _ = self.cancel_tx.try_send(());
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
info!("Tool execution cancellation requested");
|
||||
}
|
||||
|
||||
// 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_request_deadline = TokioInstant::now()
|
||||
+ self.tool_execution_policy.cancellation_request_timeout;
|
||||
let cancellation_requests = execution_handles
|
||||
.iter()
|
||||
.filter(|(call_id, _)| !terminal_call_ids.contains(*call_id))
|
||||
.map(|(call_id, handle)| {
|
||||
let call_id = call_id.clone();
|
||||
let handle = handle.clone();
|
||||
async move {
|
||||
(
|
||||
call_id,
|
||||
handle.cancel_before(cancellation_request_deadline).await,
|
||||
)
|
||||
}
|
||||
});
|
||||
let cancellation_requests: FuturesUnordered<_> =
|
||||
cancellation_requests.collect();
|
||||
for (call_id, result) in cancellation_requests.collect::<Vec<_>>().await {
|
||||
if let Err(error) = result {
|
||||
warn!(
|
||||
%call_id,
|
||||
error = %error,
|
||||
"Tool cooperative cancellation request failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Keep polling the original execution handles to their
|
||||
// provider-confirmed terminal result until the caller-selected
|
||||
// deadline. The execution remains owned even if this polling
|
||||
// future is later dropped.
|
||||
let deadline = TokioInstant::now()
|
||||
+ self.tool_execution_policy.terminal_confirmation_timeout;
|
||||
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) {
|
||||
if let Some(handle) = execution_handles.get(call_id) {
|
||||
handle.force_close();
|
||||
}
|
||||
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();
|
||||
if pause_requested {
|
||||
return Ok(ToolExecutionResult::Paused);
|
||||
}
|
||||
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(if pause_requested {
|
||||
ToolExecutionResult::Paused
|
||||
} else {
|
||||
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
|
||||
@@ -1438,6 +1840,13 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
let stream_started = Instant::now();
|
||||
let stream_result = tokio::select! {
|
||||
stream_result = self.client.stream(request.clone()) => stream_result,
|
||||
pause = self.pause_rx.recv() => {
|
||||
if pause.is_some() {
|
||||
info!("Paused before stream started");
|
||||
}
|
||||
self.timeline.abort_current_block();
|
||||
return Err(EngineError::PauseRequested);
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
info!("Cancelled before stream started");
|
||||
@@ -1469,6 +1878,13 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
);
|
||||
let first_event_result = tokio::select! {
|
||||
first_event = wait_for_first_stream_event(stream, DEFAULT_FIRST_STREAM_EVENT_TIMEOUT) => first_event,
|
||||
pause = self.pause_rx.recv() => {
|
||||
if pause.is_some() {
|
||||
info!("Paused before first stream event");
|
||||
}
|
||||
self.timeline.abort_current_block();
|
||||
return Err(EngineError::PauseRequested);
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
info!("Cancelled before first stream event");
|
||||
@@ -1555,6 +1971,13 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
|
||||
tokio::select! {
|
||||
_ = tokio::time::sleep(wait) => {}
|
||||
pause = self.pause_rx.recv() => {
|
||||
if pause.is_some() {
|
||||
info!("Paused during LLM retry backoff");
|
||||
}
|
||||
self.timeline.abort_current_block();
|
||||
return Err(EngineError::PauseRequested);
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
info!("Cancelled during LLM retry backoff");
|
||||
@@ -1633,6 +2056,13 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
None => break,
|
||||
}
|
||||
}
|
||||
pause = self.pause_rx.recv() => {
|
||||
if pause.is_some() {
|
||||
info!("Paused during response stream");
|
||||
}
|
||||
self.timeline.abort_current_block();
|
||||
return Err(EngineError::PauseRequested);
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
info!("Stream cancelled");
|
||||
@@ -1658,23 +2088,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),
|
||||
}
|
||||
}
|
||||
@@ -1688,6 +2104,7 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
let thinking_block_collector = ThinkingBlockCollector::new();
|
||||
let mut timeline = Timeline::new();
|
||||
let (cancel_tx, cancel_rx) = mpsc::channel(1);
|
||||
let (pause_tx, pause_rx) = mpsc::channel(1);
|
||||
|
||||
// Register collectors with Timeline
|
||||
timeline.on_text_block(text_block_collector.clone());
|
||||
@@ -1710,6 +2127,7 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
llm_call_count: 0,
|
||||
tool_execution_batch_count: 0,
|
||||
max_turns: None,
|
||||
tool_execution_policy: ToolExecutionPolicy::default(),
|
||||
turn_start_cbs: Vec::new(),
|
||||
turn_end_cbs: Vec::new(),
|
||||
llm_call_start_cbs: Vec::new(),
|
||||
@@ -1724,6 +2142,8 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
request_config: RequestConfig::default(),
|
||||
cancel_tx,
|
||||
cancel_rx,
|
||||
pause_tx,
|
||||
pause_rx,
|
||||
tool_output_limits: None,
|
||||
prune_config: None,
|
||||
token_estimator: None,
|
||||
@@ -1982,6 +2402,7 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
llm_call_count: self.llm_call_count,
|
||||
tool_execution_batch_count: self.tool_execution_batch_count,
|
||||
max_turns: self.max_turns,
|
||||
tool_execution_policy: self.tool_execution_policy,
|
||||
turn_start_cbs: self.turn_start_cbs,
|
||||
turn_end_cbs: self.turn_end_cbs,
|
||||
llm_call_start_cbs: self.llm_call_start_cbs,
|
||||
@@ -1997,6 +2418,8 @@ impl<C: LlmClient, A> Engine<C, Mutable, A> {
|
||||
|
||||
cancel_tx: self.cancel_tx,
|
||||
cancel_rx: self.cancel_rx,
|
||||
pause_tx: self.pause_tx,
|
||||
pause_rx: self.pause_rx,
|
||||
tool_output_limits: self.tool_output_limits,
|
||||
prune_config: self.prune_config,
|
||||
token_estimator: self.token_estimator,
|
||||
@@ -2091,7 +2514,10 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
self.append_history_items(history, extras, annotate)?;
|
||||
}
|
||||
self.start_logical_run();
|
||||
let result = self.run_turn_loop(history, annotate).await;
|
||||
let result = match self.run_turn_loop(history, annotate).await {
|
||||
Err(EngineError::PauseRequested) => Ok(EngineResult::Paused),
|
||||
other => other,
|
||||
};
|
||||
let result = self.finalize_interruption(result).await;
|
||||
self.finish_logical_run(&result);
|
||||
result
|
||||
@@ -2114,7 +2540,10 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
) -> Result<EngineResult, EngineError> {
|
||||
self.ensure_logical_run();
|
||||
let result = self.run_turn_loop(history, annotate).await;
|
||||
let result = match self.run_turn_loop(history, annotate).await {
|
||||
Err(EngineError::PauseRequested) => Ok(EngineResult::Paused),
|
||||
other => other,
|
||||
};
|
||||
let result = self.finalize_interruption(result).await;
|
||||
self.finish_logical_run(&result);
|
||||
result
|
||||
@@ -2146,6 +2575,7 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
llm_call_count: self.llm_call_count,
|
||||
tool_execution_batch_count: self.tool_execution_batch_count,
|
||||
max_turns: self.max_turns,
|
||||
tool_execution_policy: self.tool_execution_policy,
|
||||
turn_start_cbs: self.turn_start_cbs,
|
||||
turn_end_cbs: self.turn_end_cbs,
|
||||
llm_call_start_cbs: self.llm_call_start_cbs,
|
||||
@@ -2161,6 +2591,8 @@ impl<C: LlmClient, A> Engine<C, Locked, A> {
|
||||
|
||||
cancel_tx: self.cancel_tx,
|
||||
cancel_rx: self.cancel_rx,
|
||||
pause_tx: self.pause_tx,
|
||||
pause_rx: self.pause_rx,
|
||||
tool_output_limits: self.tool_output_limits,
|
||||
prune_config: self.prune_config,
|
||||
token_estimator: self.token_estimator,
|
||||
@@ -2296,6 +2728,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"[..]);
|
||||
|
||||
@@ -28,7 +28,11 @@ 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, ToolExecutionHandle, ToolExecutionPolicy,
|
||||
ToolExecutionTerminal, ToolExecutionTerminalFuture, ToolOutputLimits, ToolResult,
|
||||
ToolResultDisposition,
|
||||
};
|
||||
pub use usage_record::UsageRecord;
|
||||
|
||||
/// Implementation dependencies used by code generated from `agen` macros.
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
+227
-2
@@ -3,7 +3,14 @@
|
||||
//! Traits for defining tools callable by LLM.
|
||||
//! Usually auto-implemented using the `#[tool]` macro.
|
||||
|
||||
use std::{collections::HashMap, fmt, sync::Arc};
|
||||
use std::{
|
||||
collections::HashMap,
|
||||
fmt,
|
||||
future::Future,
|
||||
pin::Pin,
|
||||
sync::Arc,
|
||||
task::{Context, Poll},
|
||||
};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use base64::{Engine as _, engine::general_purpose::STANDARD};
|
||||
@@ -23,6 +30,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 +171,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
|
||||
@@ -322,6 +357,12 @@ impl ToolExecutionContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// Identifies one live execution attempt without making the batch id a durable
|
||||
/// replay or idempotency authority.
|
||||
pub fn execution_id(&self) -> String {
|
||||
format!("{}:{}", self.batch_id, self.call_id)
|
||||
}
|
||||
|
||||
/// Context for direct, non-engine calls in unit tests and low-level callers.
|
||||
pub fn direct() -> Self {
|
||||
Self::new("direct", "direct", 0)
|
||||
@@ -334,6 +375,142 @@ impl Default for ToolExecutionContext {
|
||||
}
|
||||
}
|
||||
|
||||
/// The provider-confirmed terminal result of one started tool execution.
|
||||
///
|
||||
/// `OutcomeUnknown` is reserved for an execution task that had to be force-closed
|
||||
/// or failed before the provider could confirm its terminal result.
|
||||
#[derive(Debug)]
|
||||
pub enum ToolExecutionTerminal {
|
||||
Confirmed(Result<ToolOutput, ToolError>),
|
||||
OutcomeUnknown,
|
||||
}
|
||||
|
||||
/// The completion future paired with a [`ToolExecutionHandle`]. Dropping this
|
||||
/// future does not drop the provider execution: the spawned execution remains
|
||||
/// owned by its handle until it completes or is explicitly force-closed.
|
||||
pub struct ToolExecutionTerminalFuture {
|
||||
task: tokio::task::JoinHandle<Result<ToolOutput, ToolError>>,
|
||||
}
|
||||
|
||||
impl Future for ToolExecutionTerminalFuture {
|
||||
type Output = ToolExecutionTerminal;
|
||||
|
||||
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
|
||||
match Pin::new(&mut self.task).poll(cx) {
|
||||
Poll::Ready(Ok(result)) => Poll::Ready(ToolExecutionTerminal::Confirmed(result)),
|
||||
Poll::Ready(Err(_)) => Poll::Ready(ToolExecutionTerminal::OutcomeUnknown),
|
||||
Poll::Pending => Poll::Pending,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Live ownership and control for one started tool execution.
|
||||
///
|
||||
/// Execution, cancellation, and terminal confirmation remain provider-owned:
|
||||
/// this handle starts `Tool::execute`, delegates cooperative cancellation to
|
||||
/// `Tool::cancel_execution`, and treats execution-future completion as the
|
||||
/// provider's terminal confirmation. Agen may force-close only after its caller's
|
||||
/// deadline expires, at which point the outcome is necessarily unknown.
|
||||
#[derive(Clone)]
|
||||
pub struct ToolExecutionHandle {
|
||||
inner: Arc<ToolExecutionHandleInner>,
|
||||
}
|
||||
|
||||
struct ToolExecutionHandleInner {
|
||||
tool: Arc<dyn Tool>,
|
||||
context: ToolExecutionContext,
|
||||
abort: tokio::task::AbortHandle,
|
||||
}
|
||||
|
||||
impl Drop for ToolExecutionHandleInner {
|
||||
fn drop(&mut self) {
|
||||
// Losing the final live owner is an explicit forced close, never a
|
||||
// best-effort detached provider future.
|
||||
self.abort.abort();
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Debug for ToolExecutionHandle {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
f.debug_struct("ToolExecutionHandle")
|
||||
.field("call_id", &self.inner.context.call_id)
|
||||
.field("batch_id", &self.inner.context.batch_id)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl ToolExecutionHandle {
|
||||
pub fn start(
|
||||
tool: Arc<dyn Tool>,
|
||||
input_json: String,
|
||||
context: ToolExecutionContext,
|
||||
) -> (Self, ToolExecutionTerminalFuture) {
|
||||
let execution_tool = Arc::clone(&tool);
|
||||
let execution_context = context.clone();
|
||||
let task =
|
||||
tokio::spawn(
|
||||
async move { execution_tool.execute(&input_json, execution_context).await },
|
||||
);
|
||||
let abort = task.abort_handle();
|
||||
(
|
||||
Self {
|
||||
inner: Arc::new(ToolExecutionHandleInner {
|
||||
tool,
|
||||
context,
|
||||
abort,
|
||||
}),
|
||||
},
|
||||
ToolExecutionTerminalFuture { task },
|
||||
)
|
||||
}
|
||||
|
||||
pub fn context(&self) -> &ToolExecutionContext {
|
||||
&self.inner.context
|
||||
}
|
||||
|
||||
pub async fn cancel_before(&self, deadline: tokio::time::Instant) -> Result<(), ToolError> {
|
||||
match tokio::time::timeout_at(
|
||||
deadline,
|
||||
self.inner.tool.cancel_execution(&self.inner.context),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(_) => Err(ToolError::Internal(format!(
|
||||
"tool cancellation request exceeded its deadline for call {}",
|
||||
self.inner.context.call_id
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn force_close(&self) {
|
||||
self.inner.abort.abort();
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub struct ToolExecutionPolicy {
|
||||
/// Time a pause waits for already-started providers to reach a natural safe
|
||||
/// boundary before escalating to explicit cooperative cancellation.
|
||||
pub pause_safe_boundary_timeout: std::time::Duration,
|
||||
/// Maximum time allowed for a provider to accept one cooperative
|
||||
/// cancellation request.
|
||||
pub cancellation_request_timeout: std::time::Duration,
|
||||
/// Maximum time allowed for all providers to confirm terminal results after
|
||||
/// cancellation has been requested.
|
||||
pub terminal_confirmation_timeout: std::time::Duration,
|
||||
}
|
||||
|
||||
impl Default for ToolExecutionPolicy {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
pause_safe_boundary_timeout: std::time::Duration::from_millis(100),
|
||||
cancellation_request_timeout: std::time::Duration::from_millis(100),
|
||||
terminal_confirmation_timeout: std::time::Duration::from_millis(500),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Tool trait
|
||||
// =============================================================================
|
||||
@@ -402,6 +579,26 @@ 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
|
||||
/// every live execution identified by `call_id`, then let `execute` return
|
||||
/// the confirmed bounded terminal output. Direct callers may use this
|
||||
/// compatibility surface; Agen uses [`Tool::cancel_execution`] so providers
|
||||
/// can bind cancellation to one exact live attempt.
|
||||
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Request cooperative cancellation for one exact started execution.
|
||||
///
|
||||
/// The default preserves existing tools by delegating to `cancel(call_id)`.
|
||||
/// Providers with their own execution registry should override this method
|
||||
/// and key cancellation by [`ToolExecutionContext::execution_id`].
|
||||
async fn cancel_execution(&self, ctx: &ToolExecutionContext) -> Result<(), ToolError> {
|
||||
self.cancel(&ctx.call_id).await
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
@@ -429,6 +626,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 +645,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 +667,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)]
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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, ToolExecutionPolicy};
|
||||
use async_trait::async_trait;
|
||||
|
||||
mod common;
|
||||
@@ -70,6 +71,144 @@ 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 SafePauseTool {
|
||||
calls: Arc<AtomicUsize>,
|
||||
cancellations: Arc<AtomicUsize>,
|
||||
release: Arc<tokio::sync::Notify>,
|
||||
}
|
||||
|
||||
impl SafePauseTool {
|
||||
fn new() -> Self {
|
||||
Self {
|
||||
calls: Arc::new(AtomicUsize::new(0)),
|
||||
cancellations: Arc::new(AtomicUsize::new(0)),
|
||||
release: Arc::new(tokio::sync::Notify::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn definition(&self) -> ToolDefinition {
|
||||
let tool = self.clone();
|
||||
Arc::new(move || {
|
||||
let meta = ToolMeta::new("safe_pause")
|
||||
.description("Waits for a safe-boundary release")
|
||||
.input_schema(serde_json::json!({"type": "object"}));
|
||||
(meta, Arc::new(tool.clone()) as Arc<dyn Tool>)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Tool for SafePauseTool {
|
||||
async fn execute(
|
||||
&self,
|
||||
_input_json: &str,
|
||||
_ctx: ToolExecutionContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
self.release.notified().await;
|
||||
Ok(ToolOutput {
|
||||
summary: "safe-boundary complete".to_string(),
|
||||
content: Some("safe-boundary complete".to_string()),
|
||||
attachments: Vec::new(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn cancel(&self, _call_id: &str) -> Result<(), ToolError> {
|
||||
self.cancellations.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ContextRecordingTool {
|
||||
name: String,
|
||||
@@ -179,6 +318,450 @@ 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 pause_waits_for_started_tool_terminal_without_cancelling_provider() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::tool_use_start(0, "call_safe_pause", "safe_pause"),
|
||||
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 = SafePauseTool::new();
|
||||
engine.register_tool(tool.definition());
|
||||
|
||||
let pause = engine.pause_sender();
|
||||
let calls = Arc::clone(&tool.calls);
|
||||
let release = Arc::clone(&tool.release);
|
||||
let control = tokio::spawn(async move {
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while calls.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("tool execution starts");
|
||||
pause.send(()).await.unwrap();
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
release.notify_one();
|
||||
});
|
||||
|
||||
let started_at = std::time::Instant::now();
|
||||
let mut history = History::new();
|
||||
let output = engine.run(&mut history, "pause safely").await;
|
||||
control.await.unwrap();
|
||||
|
||||
assert!(started_at.elapsed() >= Duration::from_millis(50));
|
||||
assert_eq!(tool.calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(tool.cancellations.load(Ordering::SeqCst), 0);
|
||||
assert!(matches!(output.result, agen::EngineRunExit::Paused));
|
||||
assert!(history.iter().any(|entry| matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::Success,
|
||||
..
|
||||
} if call_id == "call_safe_pause"
|
||||
)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pause_escalates_to_explicit_cancel_and_confirm_after_safe_boundary_deadline() {
|
||||
let client = MockLlmClient::with_responses(vec![vec![
|
||||
Event::tool_use_start(0, "call_pause_cancel", "cooperative"),
|
||||
Event::tool_input_delta(0, r#"{}"#),
|
||||
Event::tool_use_stop(0),
|
||||
Event::Status(StatusEvent {
|
||||
status: ResponseStatus::Completed,
|
||||
}),
|
||||
]]);
|
||||
let mut engine = Engine::new(client);
|
||||
engine.set_tool_execution_policy(ToolExecutionPolicy {
|
||||
pause_safe_boundary_timeout: Duration::from_millis(20),
|
||||
cancellation_request_timeout: Duration::from_millis(50),
|
||||
terminal_confirmation_timeout: Duration::from_millis(100),
|
||||
});
|
||||
let tool = CooperativeCancelTool::new();
|
||||
engine.register_tool(tool.definition());
|
||||
|
||||
let pause = engine.pause_sender();
|
||||
let calls = Arc::clone(&tool.calls);
|
||||
let control = tokio::spawn(async move {
|
||||
tokio::time::timeout(Duration::from_secs(1), async {
|
||||
while calls.load(Ordering::SeqCst) == 0 {
|
||||
tokio::task::yield_now().await;
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("tool execution starts");
|
||||
pause.send(()).await.unwrap();
|
||||
});
|
||||
|
||||
let mut history = History::new();
|
||||
let output = engine.run(&mut history, "pause with escalation").await;
|
||||
control.await.unwrap();
|
||||
|
||||
assert!(matches!(output.result, agen::EngineRunExit::Paused));
|
||||
assert!(history.iter().any(|entry| matches!(
|
||||
&entry.item,
|
||||
Item::ToolResult {
|
||||
call_id,
|
||||
disposition: ToolResultDisposition::Cancelled,
|
||||
..
|
||||
} if call_id == "call_pause_cancel"
|
||||
)));
|
||||
}
|
||||
|
||||
#[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 +1166,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"
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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(
|
||||
|
||||
+159
-14
@@ -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,21 +21,65 @@ struct BashParams {
|
||||
|
||||
pub(crate) struct BashTool {
|
||||
session: WorkdirSessionHandle,
|
||||
state: Arc<Mutex<BashExecutionState>>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct ActiveCommand {
|
||||
call_id: String,
|
||||
execution_nonce: u64,
|
||||
handle: CommandHandle,
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct BashExecutionState {
|
||||
active: HashMap<String, ActiveCommand>,
|
||||
cancellation_requested: HashSet<String>,
|
||||
legacy_cancellation_requested: HashSet<String>,
|
||||
next_execution_nonce: u64,
|
||||
}
|
||||
|
||||
struct CommandGuard {
|
||||
session: WorkdirSessionHandle,
|
||||
state: Arc<Mutex<BashExecutionState>>,
|
||||
execution_id: String,
|
||||
execution_nonce: u64,
|
||||
handle: Option<CommandHandle>,
|
||||
}
|
||||
|
||||
impl Drop for CommandGuard {
|
||||
fn drop(&mut self) {
|
||||
if let Some(handle) = self.handle.take() {
|
||||
let workdir = self.session.clone();
|
||||
tokio::spawn(async move {
|
||||
let _ = workdir.cancel_command(handle).await;
|
||||
});
|
||||
}
|
||||
let Some(handle) = self.handle.take() else {
|
||||
return;
|
||||
};
|
||||
let workdir = self.session.clone();
|
||||
let state = Arc::clone(&self.state);
|
||||
let execution_id = self.execution_id.clone();
|
||||
let execution_nonce = self.execution_nonce;
|
||||
// A dropped provider future is not terminal confirmation. Keep the live
|
||||
// execution registered until cleanup has both requested cancellation and
|
||||
// observed terminal command output, so cancellation/session teardown
|
||||
// cannot race with an apparently empty registry.
|
||||
tokio::spawn(async move {
|
||||
let _ = workdir.cancel_command(handle.clone()).await;
|
||||
let _ = workdir
|
||||
.command_output(CommandOutputRequest {
|
||||
handle,
|
||||
cursor: 0,
|
||||
limit: INLINE_BYTE_BUDGET,
|
||||
wait: true,
|
||||
})
|
||||
.await;
|
||||
let mut state = state.lock().unwrap();
|
||||
if state
|
||||
.active
|
||||
.get(&execution_id)
|
||||
.is_some_and(|active| active.execution_nonce == execution_nonce)
|
||||
{
|
||||
state.active.remove(&execution_id);
|
||||
state.cancellation_requested.remove(&execution_id);
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,20 +97,50 @@ impl Tool for BashTool {
|
||||
.unwrap_or(DEFAULT_TIMEOUT_SECS)
|
||||
.clamp(1, MAX_TIMEOUT_SECS);
|
||||
let cmd_summary = truncate_for_summary(¶ms.command);
|
||||
let execution_id = ctx.execution_id();
|
||||
let call_id = ctx.call_id;
|
||||
let execution_nonce = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state.next_execution_nonce = state.next_execution_nonce.wrapping_add(1);
|
||||
state.next_execution_nonce
|
||||
};
|
||||
let mut guard = CommandGuard {
|
||||
session: self.session.clone(),
|
||||
state: self.state.clone(),
|
||||
execution_id: execution_id.clone(),
|
||||
execution_nonce,
|
||||
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(
|
||||
execution_id.clone(),
|
||||
ActiveCommand {
|
||||
call_id: call_id.clone(),
|
||||
execution_nonce,
|
||||
handle: handle.clone(),
|
||||
},
|
||||
);
|
||||
state.cancellation_requested.contains(&execution_id)
|
||||
|| state.legacy_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 +151,27 @@ impl Tool for BashTool {
|
||||
})
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
let cancellation_requested = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
let owns_registration = state
|
||||
.active
|
||||
.get(&execution_id)
|
||||
.is_some_and(|active| active.execution_nonce == execution_nonce);
|
||||
let exact = if owns_registration {
|
||||
state.active.remove(&execution_id);
|
||||
state.cancellation_requested.remove(&execution_id)
|
||||
} else {
|
||||
false
|
||||
};
|
||||
let legacy = state.legacy_cancellation_requested.remove(&call_id);
|
||||
exact || legacy
|
||||
};
|
||||
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 +190,62 @@ 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 handles = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state
|
||||
.legacy_cancellation_requested
|
||||
.insert(call_id.to_string());
|
||||
state
|
||||
.active
|
||||
.values()
|
||||
.filter(|active| active.call_id == call_id)
|
||||
.map(|active| active.handle.clone())
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
for handle in handles {
|
||||
self.session
|
||||
.cancel_command(handle)
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn cancel_execution(
|
||||
&self,
|
||||
ctx: &agen::tool::ToolExecutionContext,
|
||||
) -> Result<(), ToolError> {
|
||||
let execution_id = ctx.execution_id();
|
||||
let handle = {
|
||||
let mut state = self.state.lock().unwrap();
|
||||
state.cancellation_requested.insert(execution_id.clone());
|
||||
state
|
||||
.active
|
||||
.get(&execution_id)
|
||||
.map(|active| active.handle.clone())
|
||||
};
|
||||
if let Some(handle) = handle {
|
||||
self.session
|
||||
.cancel_command(handle)
|
||||
.await
|
||||
.map_err(crate::ToolsError::from)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -123,6 +267,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)
|
||||
})
|
||||
|
||||
@@ -7,7 +7,10 @@
|
||||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use agen::tool::{Tool, ToolDefinition, ToolMeta};
|
||||
use agen::tool::{
|
||||
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolExecutionHandle,
|
||||
ToolExecutionTerminal, ToolMeta,
|
||||
};
|
||||
use manifest::{Permission, Scope, ScopeConfig, ScopeRule};
|
||||
use serde_json::json;
|
||||
use tempfile::TempDir;
|
||||
@@ -401,5 +404,84 @@ 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 marker = dir.path().join("must-not-run-after-cancel");
|
||||
let command = format!(
|
||||
"printf 'before\\n'; printf 'err-before\\n' >&2; sleep 1; touch {}; printf 'after\\n'",
|
||||
marker.display()
|
||||
);
|
||||
let input = serde_json::to_string(&json!({ "command": command })).unwrap();
|
||||
let context = ToolExecutionContext::new("call-heavy", "attempt-heavy", 0);
|
||||
let bash = reg.get("Bash");
|
||||
let executing = bash.clone();
|
||||
let execution_context = context.clone();
|
||||
let execution = tokio::spawn(async move { executing.execute(&input, execution_context).await });
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
bash.cancel_execution(&context)
|
||||
.await
|
||||
.expect("signal exact execution 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");
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1_100)).await;
|
||||
assert!(
|
||||
!marker.exists(),
|
||||
"the cancelled command continued executing after terminal confirmation"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn bash_force_close_cleanup_stops_command_and_keeps_session_reusable() {
|
||||
let (dir, _spill, reg) = setup();
|
||||
let marker = dir.path().join("must-not-survive-force-close");
|
||||
let command = format!("sleep 1; touch {}", marker.display());
|
||||
let input = serde_json::to_string(&json!({ "command": command })).unwrap();
|
||||
let bash = reg.get("Bash");
|
||||
let context = ToolExecutionContext::new("call-force", "attempt-force", 0);
|
||||
let (handle, terminal) = ToolExecutionHandle::start(bash.clone(), input, context);
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
handle.force_close();
|
||||
assert!(matches!(
|
||||
terminal.await,
|
||||
ToolExecutionTerminal::OutcomeUnknown
|
||||
));
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(1_100)).await;
|
||||
assert!(
|
||||
!marker.exists(),
|
||||
"CommandGuard cleanup allowed a force-closed command to continue"
|
||||
);
|
||||
|
||||
let output = bash
|
||||
.execute(r#"{"command":"printf 'reused'"}"#, Default::default())
|
||||
.await
|
||||
.expect("workdir session remains reusable after cleanup");
|
||||
assert_eq!(output.content.as_deref(), Some("reused"));
|
||||
}
|
||||
|
||||
// Sanity: unused Path import guard
|
||||
const _: fn() -> &'static Path = || Path::new("/");
|
||||
|
||||
@@ -1244,6 +1244,7 @@ impl App {
|
||||
id,
|
||||
summary,
|
||||
output,
|
||||
disposition: _,
|
||||
is_error,
|
||||
} => {
|
||||
self.latest_llm_wait_event = None;
|
||||
|
||||
+151
-12
@@ -485,6 +485,7 @@ impl WorkerController {
|
||||
// into the controller task so the in-flight turn can be reached
|
||||
// via these handles while worker itself is borrowed by drive_turn.
|
||||
let cancel_tx = worker.engine_mut().cancel_sender();
|
||||
let pause_tx = worker.engine_mut().pause_sender();
|
||||
let notify_buffer = worker.notify_buffer_handle();
|
||||
|
||||
tokio::spawn(controller_loop(
|
||||
@@ -494,6 +495,7 @@ impl WorkerController {
|
||||
shared_state,
|
||||
runtime_dir,
|
||||
cancel_tx,
|
||||
pause_tx,
|
||||
notify_buffer,
|
||||
self_parent_socket,
|
||||
spawner_name,
|
||||
@@ -763,6 +765,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,
|
||||
});
|
||||
});
|
||||
@@ -1123,6 +1138,7 @@ async fn controller_loop<C, St>(
|
||||
shared_state: Arc<WorkerSharedState>,
|
||||
runtime_dir: Arc<RuntimeDir>,
|
||||
cancel_tx: mpsc::Sender<()>,
|
||||
pause_tx: mpsc::Sender<()>,
|
||||
notify_buffer: NotifyBuffer,
|
||||
self_parent_socket: Option<PathBuf>,
|
||||
spawner_name: String,
|
||||
@@ -1169,22 +1185,35 @@ async fn controller_loop<C, St>(
|
||||
// clear at run start prevents stale partial output left by an older
|
||||
// interrupted/error turn from being carried into the next snapshot.
|
||||
worker.clear_in_flight_events();
|
||||
set_controller_status(
|
||||
&shared_state,
|
||||
&runtime_dir,
|
||||
&event_tx,
|
||||
WorkerStatus::Running,
|
||||
)
|
||||
.await;
|
||||
let parent_originated = run.is_parent_originated();
|
||||
let user_input_run = matches!(&run, PendingRun::Run(_) | PendingRun::RunTracked { .. });
|
||||
if !user_input_run {
|
||||
set_controller_status(
|
||||
&shared_state,
|
||||
&runtime_dir,
|
||||
&event_tx,
|
||||
WorkerStatus::Running,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
let (mut new_status, shutdown) = match run {
|
||||
PendingRun::Run(input) => {
|
||||
let (input_commit_tx, input_commit_rx) = oneshot::channel();
|
||||
drive_turn(
|
||||
worker.run(input),
|
||||
worker.run_with_input_extensions_and_commit_hook(
|
||||
input,
|
||||
Vec::new(),
|
||||
move || {
|
||||
let _ = input_commit_tx.send(());
|
||||
},
|
||||
),
|
||||
&mut method_rx,
|
||||
&event_tx,
|
||||
&cancel_tx,
|
||||
&pause_tx,
|
||||
&shared_state,
|
||||
&runtime_dir,
|
||||
Some(input_commit_rx),
|
||||
¬ify_buffer,
|
||||
self_parent_socket.as_ref(),
|
||||
&spawner_name,
|
||||
@@ -1194,12 +1223,22 @@ async fn controller_loop<C, St>(
|
||||
.await
|
||||
}
|
||||
PendingRun::RunTracked { input, extension } => {
|
||||
let (input_commit_tx, input_commit_rx) = oneshot::channel();
|
||||
drive_turn(
|
||||
worker.run_with_input_extensions(input, vec![extension]),
|
||||
worker.run_with_input_extensions_and_commit_hook(
|
||||
input,
|
||||
vec![extension],
|
||||
move || {
|
||||
let _ = input_commit_tx.send(());
|
||||
},
|
||||
),
|
||||
&mut method_rx,
|
||||
&event_tx,
|
||||
&cancel_tx,
|
||||
&pause_tx,
|
||||
&shared_state,
|
||||
&runtime_dir,
|
||||
Some(input_commit_rx),
|
||||
¬ify_buffer,
|
||||
self_parent_socket.as_ref(),
|
||||
&spawner_name,
|
||||
@@ -1214,7 +1253,10 @@ async fn controller_loop<C, St>(
|
||||
&mut method_rx,
|
||||
&event_tx,
|
||||
&cancel_tx,
|
||||
&pause_tx,
|
||||
&shared_state,
|
||||
&runtime_dir,
|
||||
None,
|
||||
¬ify_buffer,
|
||||
self_parent_socket.as_ref(),
|
||||
&spawner_name,
|
||||
@@ -1229,7 +1271,10 @@ async fn controller_loop<C, St>(
|
||||
&mut method_rx,
|
||||
&event_tx,
|
||||
&cancel_tx,
|
||||
&pause_tx,
|
||||
&shared_state,
|
||||
&runtime_dir,
|
||||
None,
|
||||
¬ify_buffer,
|
||||
self_parent_socket.as_ref(),
|
||||
&spawner_name,
|
||||
@@ -1626,7 +1671,10 @@ async fn drive_turn<F>(
|
||||
method_rx: &mut mpsc::Receiver<Method>,
|
||||
event_tx: &broadcast::Sender<Event>,
|
||||
cancel_tx: &mpsc::Sender<()>,
|
||||
pause_tx: &mpsc::Sender<()>,
|
||||
shared_state: &Arc<WorkerSharedState>,
|
||||
runtime_dir: &RuntimeDir,
|
||||
mut input_commit_rx: Option<oneshot::Receiver<()>>,
|
||||
notify_buffer: &NotifyBuffer,
|
||||
parent_socket: Option<&PathBuf>,
|
||||
self_name: &str,
|
||||
@@ -1642,10 +1690,34 @@ where
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
// If input commit and provider completion become ready together, expose
|
||||
// Running only after processing the commit fence. This makes the
|
||||
// Running snapshot contract deterministic even for immediate clients.
|
||||
biased;
|
||||
committed = async {
|
||||
input_commit_rx
|
||||
.as_mut()
|
||||
.expect("input commit receiver guarded by select condition")
|
||||
.await
|
||||
}, if input_commit_rx.is_some() => {
|
||||
input_commit_rx = None;
|
||||
if committed.is_ok() {
|
||||
set_controller_status(
|
||||
shared_state,
|
||||
runtime_dir,
|
||||
event_tx,
|
||||
WorkerStatus::Running,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
result = &mut worker_future => {
|
||||
return match result {
|
||||
Ok(r) => {
|
||||
let (status, run_result) = match r {
|
||||
WorkerRunResult::Finished if pause_requested => {
|
||||
(WorkerStatus::Paused, RunResult::Paused)
|
||||
}
|
||||
WorkerRunResult::Finished => (WorkerStatus::Idle, RunResult::Finished),
|
||||
WorkerRunResult::Paused => (WorkerStatus::Paused, RunResult::Paused),
|
||||
WorkerRunResult::LimitReached => (WorkerStatus::Idle, RunResult::LimitReached),
|
||||
@@ -1718,7 +1790,7 @@ where
|
||||
}
|
||||
Some(Method::Pause) => {
|
||||
pause_requested = true;
|
||||
let _ = cancel_tx.try_send(());
|
||||
let _ = pause_tx.try_send(());
|
||||
}
|
||||
Some(Method::Shutdown) => {
|
||||
shutdown_requested = true;
|
||||
@@ -1970,11 +2042,13 @@ mod tests {
|
||||
event_tx: broadcast::Sender<Event>,
|
||||
cancel_tx: mpsc::Sender<()>,
|
||||
_cancel_rx: mpsc::Receiver<()>,
|
||||
pause_tx: mpsc::Sender<()>,
|
||||
_pause_rx: mpsc::Receiver<()>,
|
||||
shared_state: Arc<WorkerSharedState>,
|
||||
notify_buffer: NotifyBuffer,
|
||||
spawned_registry: Arc<SpawnedWorkerRegistry>,
|
||||
parent_socket_path: PathBuf,
|
||||
_runtime_dir: Arc<RuntimeDir>,
|
||||
runtime_dir: Arc<RuntimeDir>,
|
||||
_temp: TempDir,
|
||||
}
|
||||
|
||||
@@ -1988,6 +2062,7 @@ mod tests {
|
||||
let (method_tx, method_rx) = mpsc::channel::<Method>(16);
|
||||
let (event_tx, _) = broadcast::channel::<Event>(16);
|
||||
let (cancel_tx, cancel_rx) = mpsc::channel::<()>(1);
|
||||
let (pause_tx, pause_rx) = mpsc::channel::<()>(1);
|
||||
let shared_state = Arc::new(WorkerSharedState::new(
|
||||
"child-worker".to_string(),
|
||||
session_store::new_segment_id(),
|
||||
@@ -2013,11 +2088,13 @@ mod tests {
|
||||
event_tx,
|
||||
cancel_tx,
|
||||
_cancel_rx: cancel_rx,
|
||||
pause_tx,
|
||||
_pause_rx: pause_rx,
|
||||
shared_state,
|
||||
notify_buffer,
|
||||
spawned_registry,
|
||||
parent_socket_path,
|
||||
_runtime_dir: runtime_dir,
|
||||
runtime_dir,
|
||||
_temp: temp,
|
||||
}
|
||||
}
|
||||
@@ -2070,7 +2147,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"child-worker",
|
||||
@@ -2091,6 +2171,44 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pause_waits_for_run_boundary_and_uses_safe_pause_channel() {
|
||||
let mut env = make_env().await;
|
||||
let method_tx = env._method_tx.clone();
|
||||
tokio::spawn(async move {
|
||||
tokio::time::sleep(Duration::from_millis(10)).await;
|
||||
method_tx.send(Method::Pause).await.expect("send pause");
|
||||
});
|
||||
|
||||
let worker_future = async {
|
||||
tokio::time::sleep(Duration::from_millis(100)).await;
|
||||
Ok::<_, WorkerError>(WorkerRunResult::Finished)
|
||||
};
|
||||
let started_at = std::time::Instant::now();
|
||||
let (status, shutdown) = drive_turn(
|
||||
worker_future,
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
None,
|
||||
"child-worker",
|
||||
&env.spawned_registry,
|
||||
true,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(status, WorkerStatus::Paused);
|
||||
assert!(!shutdown);
|
||||
assert!(started_at.elapsed() >= Duration::from_millis(100));
|
||||
assert!(env._pause_rx.try_recv().is_ok());
|
||||
assert!(env._cancel_rx.try_recv().is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn non_parent_originated_finished_stays_silent() {
|
||||
let mut env = make_env().await;
|
||||
@@ -2102,7 +2220,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"child-worker",
|
||||
@@ -2137,7 +2258,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"child-worker",
|
||||
@@ -2178,7 +2302,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"child-worker",
|
||||
@@ -2217,7 +2344,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"parent",
|
||||
@@ -2253,7 +2383,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"parent",
|
||||
@@ -2287,7 +2420,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"parent",
|
||||
@@ -2320,7 +2456,10 @@ mod tests {
|
||||
&mut env.method_rx,
|
||||
&env.event_tx,
|
||||
&env.cancel_tx,
|
||||
&env.pause_tx,
|
||||
&env.shared_state,
|
||||
&env.runtime_dir,
|
||||
None,
|
||||
&env.notify_buffer,
|
||||
Some(&env.parent_socket_path),
|
||||
"child-worker",
|
||||
|
||||
@@ -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:?}"),
|
||||
}
|
||||
|
||||
+361
-37
@@ -11,7 +11,7 @@ use agen::llm_client::types::Role;
|
||||
use agen::state::Mutable;
|
||||
use agen::{
|
||||
Engine, EngineError, EngineResult, EngineRunExit, History, HistoryEntry, Item, StopReason,
|
||||
ToolOutputLimits, UsageRecord,
|
||||
ToolExecutionPolicy, ToolOutputLimits, UsageRecord,
|
||||
};
|
||||
use arc_swap::ArcSwap;
|
||||
use session_store::{
|
||||
@@ -2581,7 +2581,10 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
|
||||
result: &EngineRunExit,
|
||||
snapshot: &EmptyTurnRollbackSnapshot,
|
||||
) -> bool {
|
||||
if !matches!(result, EngineRunExit::Interrupted(StopReason::Cancelled)) {
|
||||
if !matches!(
|
||||
result,
|
||||
EngineRunExit::Paused | EngineRunExit::Interrupted(StopReason::Cancelled)
|
||||
) {
|
||||
return false;
|
||||
}
|
||||
if self.ai_activity_counter.load(Ordering::SeqCst) != snapshot.ai_activity_count {
|
||||
@@ -2752,10 +2755,28 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
|
||||
pub(crate) async fn run_with_input_extensions(
|
||||
&mut self,
|
||||
input: Vec<Segment>,
|
||||
mut input_extensions: Vec<SessionExtension>,
|
||||
input_extensions: Vec<SessionExtension>,
|
||||
) -> Result<WorkerRunResult, WorkerError>
|
||||
where
|
||||
St: Clone + 'static,
|
||||
{
|
||||
self.run_with_input_extensions_and_commit_hook(input, input_extensions, || {})
|
||||
.await
|
||||
}
|
||||
|
||||
/// Run user input and invoke `on_input_committed` only after the annotated
|
||||
/// input has crossed both the durable Store and live SegmentLogSink commit
|
||||
/// boundaries. The Controller uses this fence before exposing `Running`, so
|
||||
/// every in-flight snapshot for a user turn includes its committed input.
|
||||
pub(crate) async fn run_with_input_extensions_and_commit_hook<F>(
|
||||
&mut self,
|
||||
input: Vec<Segment>,
|
||||
mut input_extensions: Vec<SessionExtension>,
|
||||
on_input_committed: F,
|
||||
) -> Result<WorkerRunResult, WorkerError>
|
||||
where
|
||||
St: Clone + 'static,
|
||||
F: FnOnce(),
|
||||
{
|
||||
let (input, pending_flow_state, flow_projection) = self.prepare_flow_input(input)?;
|
||||
if let Some(state) = pending_flow_state.as_ref() {
|
||||
@@ -2810,6 +2831,7 @@ impl<C: LlmClient + 'static, St: Store> Worker<C, St> {
|
||||
.expect("flow_runtime_state poisoned") = Some(state);
|
||||
}
|
||||
self.user_segments.push(input.clone());
|
||||
on_input_committed();
|
||||
|
||||
// Resolve `@<path>` file refs to system messages stashed for the
|
||||
// WorkerInterceptor to attach right after the user message. Resolution
|
||||
@@ -2934,50 +2956,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(
|
||||
@@ -3243,6 +3267,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) {
|
||||
@@ -5793,6 +5820,15 @@ pub fn apply_worker_manifest<C: LlmClient + 'static, A>(
|
||||
) {
|
||||
worker.set_request_config(request_config_from_engine_manifest(wm));
|
||||
worker.set_max_turns(wm.max_turns.map(|n| n.get()));
|
||||
// Worker owns the lifecycle strategy for already-started tool operations.
|
||||
// The provider must first accept cooperative cancellation, then confirm a
|
||||
// terminal result before this bounded deadline; Agen handles only the
|
||||
// mechanical per-call terminalization.
|
||||
worker.set_tool_execution_policy(ToolExecutionPolicy {
|
||||
pause_safe_boundary_timeout: Duration::from_millis(100),
|
||||
cancellation_request_timeout: Duration::from_millis(250),
|
||||
terminal_confirmation_timeout: Duration::from_millis(500),
|
||||
});
|
||||
worker.set_tool_output_limits(Some(ToolOutputLimits {
|
||||
default_max_bytes: wm.tool_output.default_max_bytes,
|
||||
per_tool: wm.tool_output.per_tool.clone(),
|
||||
@@ -5908,8 +5944,10 @@ fn stop_reason_error_code(reason: &StopReason) -> ErrorCode {
|
||||
| StopReason::Unexpected(
|
||||
EngineError::Aborted(_)
|
||||
| EngineError::Cancelled
|
||||
| EngineError::PauseRequested
|
||||
| EngineError::ConfigWarnings(_)
|
||||
| EngineError::HistoryAppend(_),
|
||||
| EngineError::HistoryAppend(_)
|
||||
| EngineError::ToolAttemptFence(_),
|
||||
) => ErrorCode::Internal,
|
||||
}
|
||||
}
|
||||
@@ -7423,6 +7461,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;
|
||||
|
||||
@@ -7843,6 +7993,7 @@ mod build_summary_prompt_tests {
|
||||
summary: "wrote a file".into(),
|
||||
content: None,
|
||||
attachments: Vec::new(),
|
||||
disposition: Default::default(),
|
||||
is_error: false,
|
||||
},
|
||||
},
|
||||
@@ -7885,6 +8036,7 @@ mod build_summary_prompt_tests {
|
||||
summary: "wrote a file".into(),
|
||||
content: None,
|
||||
attachments: Vec::new(),
|
||||
disposition: Default::default(),
|
||||
is_error: false,
|
||||
},
|
||||
},
|
||||
@@ -7929,6 +8081,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(
|
||||
@@ -8028,6 +8181,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();
|
||||
|
||||
@@ -806,13 +806,30 @@ async fn snapshot_includes_user_input_for_in_flight_turn() {
|
||||
let client = MockClient::sequential(vec![MockResponse::Hang(simple_text_events())]);
|
||||
let worker = make_worker(client).await;
|
||||
let handle = spawn_controller(worker).await;
|
||||
let mut events = handle.subscribe();
|
||||
|
||||
handle
|
||||
.send(Method::run_text("hello in-flight"))
|
||||
.await
|
||||
.unwrap();
|
||||
wait_for_status(&handle, WorkerStatus::Running).await;
|
||||
tokio::time::timeout(std::time::Duration::from_secs(2), async {
|
||||
loop {
|
||||
if matches!(
|
||||
events.recv().await,
|
||||
Ok(Event::Status {
|
||||
status: WorkerStatus::Running,
|
||||
})
|
||||
) {
|
||||
break;
|
||||
}
|
||||
}
|
||||
})
|
||||
.await
|
||||
.expect("running status event");
|
||||
|
||||
// The Running event is the in-flight visibility fence: the committed
|
||||
// annotated input must already be available to an immediately attaching
|
||||
// subscriber rather than racing behind this status transition.
|
||||
let stream = tokio::net::UnixStream::connect(handle.runtime_dir.socket_path())
|
||||
.await
|
||||
.unwrap();
|
||||
@@ -2152,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 => {
|
||||
@@ -2345,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:?}"
|
||||
);
|
||||
|
||||
@@ -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.
|
||||
|
||||
Reference in New Issue
Block a user