diff --git a/crates/agen/src/engine.rs b/crates/agen/src/engine.rs index 4139d0b4..bc7cf043 100644 --- a/crates/agen/src/engine.rs +++ b/crates/agen/src/engine.rs @@ -70,6 +70,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, + original_to_projected_index: Vec, +} + +fn materialize_provider_history(items: &[Item]) -> ProviderHistoryProjection { + let mut materialized: Vec<_> = items.iter().cloned().enumerate().collect(); + let mut call_order = HashMap::::new(); + let mut next_call_order = 0usize; + let mut index = 0usize; + + while index < materialized.len() { + match &materialized[index].1 { + Item::ToolCall { call_id, .. } => { + call_order.insert(call_id.clone(), next_call_order); + next_call_order += 1; + index += 1; + } + Item::ToolResult { .. } => { + let start = index; + while index < materialized.len() + && matches!(materialized[index].1, Item::ToolResult { .. }) + { + index += 1; + } + materialized[start..index].sort_by_key(|(_, item)| match item { + Item::ToolResult { call_id, .. } => { + call_order.get(call_id).copied().unwrap_or(usize::MAX) + } + _ => unreachable!("tool-result run contains only ToolResult items"), + }); + } + _ => index += 1, + } + } + + let mut original_to_projected_index = vec![0; materialized.len()]; + for (projected_index, (original_index, _)) in materialized.iter().enumerate() { + original_to_projected_index[*original_index] = projected_index; + } + + ProviderHistoryProjection { + items: materialized.into_iter().map(|(_, item)| item).collect(), + original_to_projected_index, + } +} + /// Legacy serializable outcome used by the Worker session-log compatibility boundary. #[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)] #[serde(rename_all = "snake_case")] @@ -126,7 +179,7 @@ pub struct EngineRunOutput { /// Internal: tool execution result enum ToolExecutionResult { - Completed(Vec), + Completed, Paused, } @@ -892,8 +945,14 @@ impl Engine { 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 +965,7 @@ impl Engine { // 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 +1037,11 @@ impl Engine { /// executes approved tools in parallel and applies post_tool_call hooks to results. async fn execute_tools( &mut self, + history: &mut History, + annotate: &mut impl FnMut(&Item) -> Result, tool_calls: Vec, ) -> Result { - use futures::future::join_all; + use futures::stream::{FuturesUnordered, StreamExt}; // Map from tool call ID to (ToolCall, Meta, Tool, Context) // Retained because it's needed for PostToolCall hooks @@ -1047,8 +1108,10 @@ impl Engine { } } - // Phase 2: Execute approved tools in parallel (cancellable) - let futures: Vec<_> = approved_calls + // Phase 2: Execute approved tools in parallel. FuturesUnordered yields + // each terminal result as soon as that call completes instead of + // holding fast siblings behind the slowest call in the batch. + let futures: FuturesUnordered<_> = approved_calls .into_iter() .map(|(tool_call, context)| { let tool_server = self.tool_server.clone(); @@ -1065,84 +1128,117 @@ impl Engine { }) .collect(); - // Make tool execution cancellable - let mut results = tokio::select! { - results = join_all(futures) => results, - cancel = self.cancel_rx.recv() => { - if cancel.is_some() { - info!("Tool execution cancelled"); + // Synthetic results are already terminal and need no execution wait. + // Commit them before polling ordinary calls so they obey the same + // commit-before-publish boundary. + for result in synthetic_results { + self.finalize_and_commit_tool_result(history, annotate, result, &call_info_map) + .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(); - return Err(EngineError::Cancelled); - } - }; - results.extend(synthetic_results); - - // Phase 3: Apply post_tool_call interceptor - for tool_result in &mut results { - if let Some((tool_call, meta, tool, context)) = - call_info_map.get(&tool_result.tool_use_id) - { - let mut info = ToolResultInfo { - call: tool_call.clone(), - result: tool_result.clone(), - meta: meta.clone(), - tool: tool.clone(), - context: context.clone(), - }; - - match self.interceptor.post_tool_call(&mut info).await { - PostToolAction::Continue => {} - PostToolAction::Abort(reason) => { - return Err(EngineError::Aborted(reason)); + cancel = self.cancel_rx.recv() => { + if cancel.is_some() { + info!("Tool execution cancelled"); } - } - // Reflect interceptor-modified results - *tool_result = info.result; - } - } - - // Phase 4: Cap `content` byte-size before it enters history. - // Runs *after* post_tool_call so interceptors (audit, logging, - // classification) still observe the full content, and any - // content they inject is also truncated — closing the last gap - // before the data reaches the next LLM request. - if let Some(limits) = self.tool_output_limits.as_ref() { - for tool_result in &mut results { - let Some(content) = tool_result.content.as_mut() else { - continue; - }; - let Some((tool_call, _, _, _)) = call_info_map.get(&tool_result.tool_use_id) else { - continue; - }; - let limit = limits.limit_for(&tool_call.name); - let before = content.len(); - truncate_content(content, limit); - if content.len() != before { - warn!( - tool = %tool_call.name, - before_bytes = before, - after_bytes = content.len(), - limit_bytes = limit, - "Tool output exceeded byte limit and was truncated" - ); - self.emit_warning(&format!( - "tool `{}` output truncated from {} to {} bytes (limit {})", - tool_call.name, - before, - content.len(), - limit - )); + self.timeline.abort_current_block(); + return Err(EngineError::Cancelled); } } } - // Emit per-result callbacks on the post-truncation payload. - for tool_result in &results { - self.emit_tool_result(tool_result); + Ok(ToolExecutionResult::Completed) + } + + /// Apply post-execution policy, bound the model-visible payload, durably + /// append one terminal ToolResult, and only then publish it to observers. + async fn finalize_and_commit_tool_result( + &mut self, + history: &mut History, + annotate: &mut impl FnMut(&Item) -> Result, + mut tool_result: ToolResult, + call_info_map: &HashMap< + String, + ( + ToolCall, + crate::tool::ToolMeta, + Arc, + 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 @@ -1658,23 +1754,9 @@ impl Engine { annotate: &mut impl FnMut(&Item) -> Result, tool_calls: Vec, ) -> Result, EngineError> { - match self.execute_tools(tool_calls).await { + match self.execute_tools(history, annotate, tool_calls).await { Ok(ToolExecutionResult::Paused) => Ok(Some(EngineResult::Paused)), - Ok(ToolExecutionResult::Completed(results)) => { - // Route per-result pushes through the callback path so - // observers see each tool result as it lands. - let items = results.into_iter().map(|result| { - Item::tool_result_item_with_attachments( - &result.tool_use_id, - &result.summary, - result.content, - result.is_error, - result.attachments, - ) - }); - self.append_history_items(history, items, annotate)?; - Ok(None) - } + Ok(ToolExecutionResult::Completed) => Ok(None), Err(err) => Err(err), } } @@ -2296,6 +2378,28 @@ mod tests { use crate::tool::{Attachment, ImageAttachment}; 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] fn tool_attachment_round_trips_through_durable_history_json() { let body: Arc<[u8]> = Arc::from(&b"image-body"[..]); diff --git a/crates/agen/tests/common/mod.rs b/crates/agen/tests/common/mod.rs index 81902e78..5f4d4e94 100644 --- a/crates/agen/tests/common/mod.rs +++ b/crates/agen/tests/common/mod.rs @@ -19,6 +19,7 @@ use std::sync::atomic::{AtomicUsize, Ordering}; pub struct MockLlmClient { responses: Arc>>, call_count: Arc, + requests: Arc>>, } 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 { + self.requests.lock().unwrap().clone() + } } #[async_trait] @@ -51,8 +57,9 @@ impl LlmClient for MockLlmClient { async fn stream( &self, - _request: Request, + request: Request, ) -> Result> + 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 { diff --git a/crates/agen/tests/parallel_execution_test.rs b/crates/agen/tests/parallel_execution_test.rs index 7a68cbf4..2141b55c 100644 --- a/crates/agen/tests/parallel_execution_test.rs +++ b/crates/agen/tests/parallel_execution_test.rs @@ -11,7 +11,7 @@ use agen::llm_client::event::{Event, ResponseStatus, StatusEvent}; use agen::tool::{ Tool, ToolDefinition, ToolError, ToolExecutionContext, ToolMeta, ToolOutput, ToolResult, }; -use agen::{Engine, History}; +use agen::{Engine, History, Item}; use async_trait::async_trait; mod common; @@ -70,6 +70,48 @@ impl Tool for SlowTool { } } +#[derive(Clone)] +struct FirstAttemptHangsTool { + calls: Arc, +} + +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) + }) + } + + 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 { + 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 ContextRecordingTool { name: String, @@ -179,6 +221,201 @@ 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::::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] async fn test_tool_execution_context_order_and_batch_id() { let client = MockLlmClient::with_responses(vec![