fix: terminalize parallel tool outputs on completion
This commit is contained in:
+197
-93
@@ -70,6 +70,59 @@ pub struct EngineConfig {
|
|||||||
_private: (),
|
_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.
|
/// Legacy serializable outcome used by the Worker session-log compatibility boundary.
|
||||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
|
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
|
||||||
#[serde(rename_all = "snake_case")]
|
#[serde(rename_all = "snake_case")]
|
||||||
@@ -126,7 +179,7 @@ pub struct EngineRunOutput<C: LlmClient, A = ()> {
|
|||||||
|
|
||||||
/// Internal: tool execution result
|
/// Internal: tool execution result
|
||||||
enum ToolExecutionResult {
|
enum ToolExecutionResult {
|
||||||
Completed(Vec<ToolResult>),
|
Completed,
|
||||||
Paused,
|
Paused,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -892,8 +945,14 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
request = request.system(system);
|
request = request.system(system);
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add items directly (Request now uses Items natively)
|
// History keeps terminal tool outputs in completion order so each
|
||||||
request = request.items(context.iter().cloned());
|
// 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
|
// Add tool definitions
|
||||||
for tool_def in tool_definitions {
|
for tool_def in tool_definitions {
|
||||||
@@ -906,7 +965,7 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
// Attach the cache prefix anchor (may be narrower than `context`
|
// Attach the cache prefix anchor (may be narrower than `context`
|
||||||
// if the prune projection trimmed items from the head — keep it
|
// if the prune projection trimmed items from the head — keep it
|
||||||
// in range).
|
// 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.cache_key = self.cache_key.clone();
|
||||||
|
|
||||||
request
|
request
|
||||||
@@ -978,9 +1037,11 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
/// executes approved tools in parallel and applies post_tool_call hooks to results.
|
/// executes approved tools in parallel and applies post_tool_call hooks to results.
|
||||||
async fn execute_tools(
|
async fn execute_tools(
|
||||||
&mut self,
|
&mut self,
|
||||||
|
history: &mut History<A>,
|
||||||
|
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||||
tool_calls: Vec<ToolCall>,
|
tool_calls: Vec<ToolCall>,
|
||||||
) -> Result<ToolExecutionResult, EngineError> {
|
) -> Result<ToolExecutionResult, EngineError> {
|
||||||
use futures::future::join_all;
|
use futures::stream::{FuturesUnordered, StreamExt};
|
||||||
|
|
||||||
// Map from tool call ID to (ToolCall, Meta, Tool, Context)
|
// Map from tool call ID to (ToolCall, Meta, Tool, Context)
|
||||||
// Retained because it's needed for PostToolCall hooks
|
// Retained because it's needed for PostToolCall hooks
|
||||||
@@ -1047,8 +1108,10 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Phase 2: Execute approved tools in parallel (cancellable)
|
// Phase 2: Execute approved tools in parallel. FuturesUnordered yields
|
||||||
let futures: Vec<_> = approved_calls
|
// each terminal result as soon as that call completes instead of
|
||||||
|
// holding fast siblings behind the slowest call in the batch.
|
||||||
|
let futures: FuturesUnordered<_> = approved_calls
|
||||||
.into_iter()
|
.into_iter()
|
||||||
.map(|(tool_call, context)| {
|
.map(|(tool_call, context)| {
|
||||||
let tool_server = self.tool_server.clone();
|
let tool_server = self.tool_server.clone();
|
||||||
@@ -1065,84 +1128,117 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
})
|
})
|
||||||
.collect();
|
.collect();
|
||||||
|
|
||||||
// Make tool execution cancellable
|
// Synthetic results are already terminal and need no execution wait.
|
||||||
let mut results = tokio::select! {
|
// Commit them before polling ordinary calls so they obey the same
|
||||||
results = join_all(futures) => results,
|
// commit-before-publish boundary.
|
||||||
cancel = self.cancel_rx.recv() => {
|
for result in synthetic_results {
|
||||||
if cancel.is_some() {
|
self.finalize_and_commit_tool_result(history, annotate, result, &call_info_map)
|
||||||
info!("Tool execution cancelled");
|
.await?;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut futures = futures;
|
||||||
|
while !futures.is_empty() {
|
||||||
|
tokio::select! {
|
||||||
|
// If cancellation and a completed result are both ready, drain
|
||||||
|
// the completed result first. This preserves every terminal
|
||||||
|
// output observed before the cancellation boundary.
|
||||||
|
biased;
|
||||||
|
result = futures.next() => {
|
||||||
|
let result = result.expect("non-empty FuturesUnordered returns a result");
|
||||||
|
self.finalize_and_commit_tool_result(
|
||||||
|
history,
|
||||||
|
annotate,
|
||||||
|
result,
|
||||||
|
&call_info_map,
|
||||||
|
).await?;
|
||||||
}
|
}
|
||||||
self.timeline.abort_current_block();
|
cancel = self.cancel_rx.recv() => {
|
||||||
return Err(EngineError::Cancelled);
|
if cancel.is_some() {
|
||||||
}
|
info!("Tool execution cancelled");
|
||||||
};
|
|
||||||
results.extend(synthetic_results);
|
|
||||||
|
|
||||||
// Phase 3: Apply post_tool_call interceptor
|
|
||||||
for tool_result in &mut results {
|
|
||||||
if let Some((tool_call, meta, tool, context)) =
|
|
||||||
call_info_map.get(&tool_result.tool_use_id)
|
|
||||||
{
|
|
||||||
let mut info = ToolResultInfo {
|
|
||||||
call: tool_call.clone(),
|
|
||||||
result: tool_result.clone(),
|
|
||||||
meta: meta.clone(),
|
|
||||||
tool: tool.clone(),
|
|
||||||
context: context.clone(),
|
|
||||||
};
|
|
||||||
|
|
||||||
match self.interceptor.post_tool_call(&mut info).await {
|
|
||||||
PostToolAction::Continue => {}
|
|
||||||
PostToolAction::Abort(reason) => {
|
|
||||||
return Err(EngineError::Aborted(reason));
|
|
||||||
}
|
}
|
||||||
}
|
self.timeline.abort_current_block();
|
||||||
// Reflect interceptor-modified results
|
return Err(EngineError::Cancelled);
|
||||||
*tool_result = info.result;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Phase 4: Cap `content` byte-size before it enters history.
|
|
||||||
// Runs *after* post_tool_call so interceptors (audit, logging,
|
|
||||||
// classification) still observe the full content, and any
|
|
||||||
// content they inject is also truncated — closing the last gap
|
|
||||||
// before the data reaches the next LLM request.
|
|
||||||
if let Some(limits) = self.tool_output_limits.as_ref() {
|
|
||||||
for tool_result in &mut results {
|
|
||||||
let Some(content) = tool_result.content.as_mut() else {
|
|
||||||
continue;
|
|
||||||
};
|
|
||||||
let Some((tool_call, _, _, _)) = call_info_map.get(&tool_result.tool_use_id) else {
|
|
||||||
continue;
|
|
||||||
};
|
|
||||||
let limit = limits.limit_for(&tool_call.name);
|
|
||||||
let before = content.len();
|
|
||||||
truncate_content(content, limit);
|
|
||||||
if content.len() != before {
|
|
||||||
warn!(
|
|
||||||
tool = %tool_call.name,
|
|
||||||
before_bytes = before,
|
|
||||||
after_bytes = content.len(),
|
|
||||||
limit_bytes = limit,
|
|
||||||
"Tool output exceeded byte limit and was truncated"
|
|
||||||
);
|
|
||||||
self.emit_warning(&format!(
|
|
||||||
"tool `{}` output truncated from {} to {} bytes (limit {})",
|
|
||||||
tool_call.name,
|
|
||||||
before,
|
|
||||||
content.len(),
|
|
||||||
limit
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Emit per-result callbacks on the post-truncation payload.
|
Ok(ToolExecutionResult::Completed)
|
||||||
for tool_result in &results {
|
}
|
||||||
self.emit_tool_result(tool_result);
|
|
||||||
|
/// 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,
|
||||||
|
call_info_map: &HashMap<
|
||||||
|
String,
|
||||||
|
(
|
||||||
|
ToolCall,
|
||||||
|
crate::tool::ToolMeta,
|
||||||
|
Arc<dyn crate::tool::Tool>,
|
||||||
|
ToolExecutionContext,
|
||||||
|
),
|
||||||
|
>,
|
||||||
|
) -> Result<(), EngineError> {
|
||||||
|
let call_info = call_info_map.get(&tool_result.tool_use_id);
|
||||||
|
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) => {
|
||||||
|
return Err(EngineError::Aborted(reason));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
tool_result = info.result;
|
||||||
}
|
}
|
||||||
|
|
||||||
Ok(ToolExecutionResult::Completed(results))
|
// 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_attachments(
|
||||||
|
&tool_result.tool_use_id,
|
||||||
|
&tool_result.summary,
|
||||||
|
tool_result.content.clone(),
|
||||||
|
tool_result.is_error,
|
||||||
|
tool_result.attachments.clone(),
|
||||||
|
);
|
||||||
|
self.append_history_items(history, std::iter::once(item), annotate)?;
|
||||||
|
self.emit_tool_result(&tool_result);
|
||||||
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Internal turn execution logic
|
/// Internal turn execution logic
|
||||||
@@ -1658,23 +1754,9 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
|||||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||||
tool_calls: Vec<ToolCall>,
|
tool_calls: Vec<ToolCall>,
|
||||||
) -> Result<Option<EngineResult>, EngineError> {
|
) -> 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::Paused) => Ok(Some(EngineResult::Paused)),
|
||||||
Ok(ToolExecutionResult::Completed(results)) => {
|
Ok(ToolExecutionResult::Completed) => Ok(None),
|
||||||
// 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)
|
|
||||||
}
|
|
||||||
Err(err) => Err(err),
|
Err(err) => Err(err),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -2296,6 +2378,28 @@ mod tests {
|
|||||||
use crate::tool::{Attachment, ImageAttachment};
|
use crate::tool::{Attachment, ImageAttachment};
|
||||||
use std::time::Duration;
|
use std::time::Duration;
|
||||||
|
|
||||||
|
#[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]
|
#[test]
|
||||||
fn tool_attachment_round_trips_through_durable_history_json() {
|
fn tool_attachment_round_trips_through_durable_history_json() {
|
||||||
let body: Arc<[u8]> = Arc::from(&b"image-body"[..]);
|
let body: Arc<[u8]> = Arc::from(&b"image-body"[..]);
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ use std::sync::atomic::{AtomicUsize, Ordering};
|
|||||||
pub struct MockLlmClient {
|
pub struct MockLlmClient {
|
||||||
responses: Arc<Vec<Vec<Event>>>,
|
responses: Arc<Vec<Vec<Event>>>,
|
||||||
call_count: Arc<AtomicUsize>,
|
call_count: Arc<AtomicUsize>,
|
||||||
|
requests: Arc<Mutex<Vec<Request>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl MockLlmClient {
|
impl MockLlmClient {
|
||||||
@@ -30,6 +31,7 @@ impl MockLlmClient {
|
|||||||
Self {
|
Self {
|
||||||
responses: Arc::new(responses),
|
responses: Arc::new(responses),
|
||||||
call_count: Arc::new(AtomicUsize::new(0)),
|
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 {
|
pub fn event_count(&self) -> usize {
|
||||||
self.responses.iter().map(|v| v.len()).sum()
|
self.responses.iter().map(|v| v.len()).sum()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn requests(&self) -> Vec<Request> {
|
||||||
|
self.requests.lock().unwrap().clone()
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
@@ -51,8 +57,9 @@ impl LlmClient for MockLlmClient {
|
|||||||
|
|
||||||
async fn stream(
|
async fn stream(
|
||||||
&self,
|
&self,
|
||||||
_request: Request,
|
request: Request,
|
||||||
) -> Result<Pin<Box<dyn Stream<Item = Result<Event, ClientError>> + Send>>, ClientError> {
|
) -> 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);
|
let count = self.call_count.fetch_add(1, Ordering::SeqCst);
|
||||||
if count >= self.responses.len() {
|
if count >= self.responses.len() {
|
||||||
return Err(ClientError::Api {
|
return Err(ClientError::Api {
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ use agen::llm_client::event::{Event, ResponseStatus, StatusEvent};
|
|||||||
use agen::tool::{
|
use agen::tool::{
|
||||||
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult,
|
Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult,
|
||||||
};
|
};
|
||||||
use agen::{Engine, History};
|
use agen::{Engine, History, Item};
|
||||||
use async_trait::async_trait;
|
use async_trait::async_trait;
|
||||||
|
|
||||||
mod common;
|
mod common;
|
||||||
@@ -70,6 +70,48 @@ 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)]
|
#[derive(Clone)]
|
||||||
struct ContextRecordingTool {
|
struct ContextRecordingTool {
|
||||||
name: String,
|
name: String,
|
||||||
@@ -179,6 +221,201 @@ async fn test_parallel_tool_execution() {
|
|||||||
println!("Parallel execution completed in {:?}", elapsed);
|
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;
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
observed.lock().unwrap().as_slice(),
|
||||||
|
[
|
||||||
|
"commit:call_fast",
|
||||||
|
"publish:call_fast",
|
||||||
|
"commit:call_slow",
|
||||||
|
"publish:call_slow",
|
||||||
|
]
|
||||||
|
);
|
||||||
|
|
||||||
|
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", "fast"),
|
||||||
|
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, "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 = SlowTool::new("fast", 1);
|
||||||
|
engine.register_tool(hanging.definition());
|
||||||
|
engine.register_tool(fast.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"
|
||||||
|
)
|
||||||
|
})
|
||||||
|
.count();
|
||||||
|
assert_eq!(completed_before_resume, 1);
|
||||||
|
assert_eq!(fast.call_count(), 1);
|
||||||
|
assert_eq!(hanging.call_count(), 1);
|
||||||
|
|
||||||
|
let _ = engine.resume(&mut history).await;
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
fast.call_count(),
|
||||||
|
1,
|
||||||
|
"completed call must not be re-executed"
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
hanging.call_count(),
|
||||||
|
2,
|
||||||
|
"only the unresolved call is retried"
|
||||||
|
);
|
||||||
|
let completed_after_resume = history
|
||||||
|
.iter()
|
||||||
|
.filter(|entry| {
|
||||||
|
matches!(
|
||||||
|
&entry.item,
|
||||||
|
Item::ToolResult { call_id, .. } if call_id == "call_fast"
|
||||||
|
)
|
||||||
|
})
|
||||||
|
.count();
|
||||||
|
assert_eq!(completed_after_resume, 1);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[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]
|
#[tokio::test]
|
||||||
async fn test_tool_execution_context_order_and_batch_id() {
|
async fn test_tool_execution_context_order_and_batch_id() {
|
||||||
let client = MockLlmClient::with_responses(vec![
|
let client = MockLlmClient::with_responses(vec![
|
||||||
|
|||||||
Reference in New Issue
Block a user