fix: terminalize parallel tool outputs on completion
This commit is contained in:
+161
-57
@@ -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<Item>,
|
||||
original_to_projected_index: Vec<usize>,
|
||||
}
|
||||
|
||||
fn materialize_provider_history(items: &[Item]) -> ProviderHistoryProjection {
|
||||
let mut materialized: Vec<_> = items.iter().cloned().enumerate().collect();
|
||||
let mut call_order = HashMap::<String, usize>::new();
|
||||
let mut next_call_order = 0usize;
|
||||
let mut index = 0usize;
|
||||
|
||||
while index < materialized.len() {
|
||||
match &materialized[index].1 {
|
||||
Item::ToolCall { call_id, .. } => {
|
||||
call_order.insert(call_id.clone(), next_call_order);
|
||||
next_call_order += 1;
|
||||
index += 1;
|
||||
}
|
||||
Item::ToolResult { .. } => {
|
||||
let start = index;
|
||||
while index < materialized.len()
|
||||
&& matches!(materialized[index].1, Item::ToolResult { .. })
|
||||
{
|
||||
index += 1;
|
||||
}
|
||||
materialized[start..index].sort_by_key(|(_, item)| match item {
|
||||
Item::ToolResult { call_id, .. } => {
|
||||
call_order.get(call_id).copied().unwrap_or(usize::MAX)
|
||||
}
|
||||
_ => unreachable!("tool-result run contains only ToolResult items"),
|
||||
});
|
||||
}
|
||||
_ => index += 1,
|
||||
}
|
||||
}
|
||||
|
||||
let mut original_to_projected_index = vec![0; materialized.len()];
|
||||
for (projected_index, (original_index, _)) in materialized.iter().enumerate() {
|
||||
original_to_projected_index[*original_index] = projected_index;
|
||||
}
|
||||
|
||||
ProviderHistoryProjection {
|
||||
items: materialized.into_iter().map(|(_, item)| item).collect(),
|
||||
original_to_projected_index,
|
||||
}
|
||||
}
|
||||
|
||||
/// Legacy serializable outcome used by the Worker session-log compatibility boundary.
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
@@ -126,7 +179,7 @@ pub struct EngineRunOutput<C: LlmClient, A = ()> {
|
||||
|
||||
/// Internal: tool execution result
|
||||
enum ToolExecutionResult {
|
||||
Completed(Vec<ToolResult>),
|
||||
Completed,
|
||||
Paused,
|
||||
}
|
||||
|
||||
@@ -892,8 +945,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 +965,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 +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.
|
||||
async fn execute_tools(
|
||||
&mut self,
|
||||
history: &mut History<A>,
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
tool_calls: Vec<ToolCall>,
|
||||
) -> Result<ToolExecutionResult, EngineError> {
|
||||
use futures::future::join_all;
|
||||
use futures::stream::{FuturesUnordered, StreamExt};
|
||||
|
||||
// Map from tool call ID to (ToolCall, Meta, Tool, Context)
|
||||
// Retained because it's needed for PostToolCall hooks
|
||||
@@ -1047,8 +1108,10 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
}
|
||||
}
|
||||
|
||||
// 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,9 +1128,30 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Make tool execution cancellable
|
||||
let mut results = tokio::select! {
|
||||
results = join_all(futures) => results,
|
||||
// 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?;
|
||||
}
|
||||
cancel = self.cancel_rx.recv() => {
|
||||
if cancel.is_some() {
|
||||
info!("Tool execution cancelled");
|
||||
@@ -1075,17 +1159,34 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
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)
|
||||
{
|
||||
Ok(ToolExecutionResult::Completed)
|
||||
}
|
||||
|
||||
/// Apply post-execution policy, bound the model-visible payload, durably
|
||||
/// append one terminal ToolResult, and only then publish it to observers.
|
||||
async fn finalize_and_commit_tool_result(
|
||||
&mut self,
|
||||
history: &mut History<A>,
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
mut tool_result: ToolResult,
|
||||
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.clone(),
|
||||
result: tool_result,
|
||||
meta: meta.clone(),
|
||||
tool: tool.clone(),
|
||||
context: context.clone(),
|
||||
@@ -1097,24 +1198,16 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
return Err(EngineError::Aborted(reason));
|
||||
}
|
||||
}
|
||||
// Reflect interceptor-modified results
|
||||
*tool_result = info.result;
|
||||
}
|
||||
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;
|
||||
};
|
||||
// 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);
|
||||
@@ -1135,14 +1228,17 @@ impl<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Emit per-result callbacks on the post-truncation payload.
|
||||
for tool_result in &results {
|
||||
self.emit_tool_result(tool_result);
|
||||
}
|
||||
|
||||
Ok(ToolExecutionResult::Completed(results))
|
||||
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<C: LlmClient, S: EngineState, A> Engine<C, S, A> {
|
||||
annotate: &mut impl FnMut(&Item) -> Result<A, String>,
|
||||
tool_calls: Vec<ToolCall>,
|
||||
) -> Result<Option<EngineResult>, EngineError> {
|
||||
match self.execute_tools(tool_calls).await {
|
||||
match self.execute_tools(history, annotate, tool_calls).await {
|
||||
Ok(ToolExecutionResult::Paused) => Ok(Some(EngineResult::Paused)),
|
||||
Ok(ToolExecutionResult::Completed(results)) => {
|
||||
// Route per-result pushes through the callback path so
|
||||
// observers see each tool result as it lands.
|
||||
let items = results.into_iter().map(|result| {
|
||||
Item::tool_result_item_with_attachments(
|
||||
&result.tool_use_id,
|
||||
&result.summary,
|
||||
result.content,
|
||||
result.is_error,
|
||||
result.attachments,
|
||||
)
|
||||
});
|
||||
self.append_history_items(history, items, annotate)?;
|
||||
Ok(None)
|
||||
}
|
||||
Ok(ToolExecutionResult::Completed) => Ok(None),
|
||||
Err(err) => Err(err),
|
||||
}
|
||||
}
|
||||
@@ -2296,6 +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"[..]);
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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<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 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::<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]
|
||||
async fn test_tool_execution_context_order_and_batch_id() {
|
||||
let client = MockLlmClient::with_responses(vec![
|
||||
|
||||
Reference in New Issue
Block a user