mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-25 14:53:34 +00:00
* fix: parallelize tool call execution via JoinSet (#219) When the LLM returns multiple tool_calls in a single response, they were executed sequentially. This change makes both the worker and dispatcher paths concurrent using tokio::task::JoinSet, so N independent tool calls complete in ~max(latency) instead of sum(latency). Worker path: migrate execute_tools_parallel from join_all to JoinSet and route the respond_with_tools branch through the same parallel path. Dispatcher path: restructure the while-idx loop into three phases — preflight (sequential approval/hook checks), parallel execution via JoinSet, and sequential post-flight processing (session recording, auth detection, sanitization). Also fixes a pre-existing infinite loop bug where hook rejection used `continue` inside a `while idx` loop, skipping `idx += 1` and retrying the same rejected tool forever. Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address PR review — ordered results, deferred auth, dedup standalone fn - Fix auth early return skipping unrecorded tool results: defer auth response until after all results in the batch are recorded in session history and context_messages (both dispatcher and thread_ops paths) - Fix tool results appearing out of order: collect Phase 1 hook rejections indexed by original position, merge with Phase 2 execution results, and emit all in Phase 3 in original tool_calls order - Deduplicate execute_chat_tool: Agent method now delegates to the standalone function instead of duplicating 90 lines of logic - Fix benchmark compilation: add missing session_manager arg to Agent::new Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: rustfmt alignment for CI compatibility Co-Authored-By: Claude Opus 4.6 <[email protected]> * fix: address second round of PR review comments - Distinguish JoinError panic vs cancellation in log messages and error reasons across all 3 files (dispatcher, thread_ops, worker) - Simplify deferred_auth from Option<(String, String)> to Option<String> since only the instructions string is used - Add single-tool short-circuit in worker execute_tools_parallel to avoid JoinSet overhead for the common single-tool case Co-Authored-By: Claude Opus 4.6 <[email protected]> --------- Co-authored-by: Claude Opus 4.6 <[email protected]>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
9906190de7
commit
bfe393eb38
+472
-240
@@ -6,6 +6,7 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::task::JoinSet;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::Agent;
|
||||
@@ -254,23 +255,40 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// Execute each tool (with approval checking and hook interception)
|
||||
let mut idx = 0usize;
|
||||
while idx < tool_calls.len() {
|
||||
let mut tc = tool_calls[idx].clone();
|
||||
// === Phase 1: Preflight (sequential) ===
|
||||
// Walk tool_calls checking approval and hooks. Classify
|
||||
// each tool as Rejected (by hook) or Runnable. Stop at the
|
||||
// first tool that needs approval.
|
||||
//
|
||||
// Outcomes are indexed by original tool_calls position so
|
||||
// Phase 3 can emit results in the correct order.
|
||||
enum PreflightOutcome {
|
||||
/// Hook rejected/blocked this tool; contains the error message.
|
||||
Rejected(String),
|
||||
/// Tool passed preflight and will be executed.
|
||||
Runnable,
|
||||
}
|
||||
let mut preflight: Vec<(crate::llm::ToolCall, PreflightOutcome)> = Vec::new();
|
||||
let mut runnable: Vec<(usize, crate::llm::ToolCall)> = Vec::new();
|
||||
let mut approval_needed: Option<(
|
||||
usize,
|
||||
crate::llm::ToolCall,
|
||||
Arc<dyn crate::tools::Tool>,
|
||||
)> = None;
|
||||
|
||||
for (idx, original_tc) in tool_calls.iter().enumerate() {
|
||||
let mut tc = original_tc.clone();
|
||||
|
||||
// Check if tool requires approval
|
||||
if let Some(tool) = self.tools().get(&tc.name).await
|
||||
&& tool.requires_approval()
|
||||
{
|
||||
// Check if auto-approved for this session
|
||||
let mut is_auto_approved = {
|
||||
let sess = session.lock().await;
|
||||
sess.is_tool_auto_approved(&tc.name)
|
||||
};
|
||||
|
||||
// Override auto-approval for destructive parameters
|
||||
// (e.g. `rm -rf`, `git push --force` in shell commands).
|
||||
if is_auto_approved && tool.requires_approval_for(&tc.arguments) {
|
||||
tracing::info!(
|
||||
tool = %tc.name,
|
||||
@@ -280,175 +298,318 @@ impl Agent {
|
||||
}
|
||||
|
||||
if !is_auto_approved {
|
||||
// Need approval - store pending request and return.
|
||||
// Preserve remaining tool calls so they can be replayed
|
||||
// after approval.
|
||||
let pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
deferred_tool_calls: tool_calls[idx + 1..].to_vec(),
|
||||
};
|
||||
|
||||
return Ok(AgenticLoopResult::NeedApproval { pending });
|
||||
approval_needed = Some((idx, tc, tool));
|
||||
break; // remaining tools are deferred
|
||||
}
|
||||
}
|
||||
|
||||
// Hook: BeforeToolCall — allow hooks to modify or reject tool calls
|
||||
{
|
||||
let event = crate::hooks::HookEvent::ToolCall {
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
user_id: message.user_id.clone(),
|
||||
context: "chat".to_string(),
|
||||
};
|
||||
match self.hooks().run(&event).await {
|
||||
Err(crate::hooks::HookError::Rejected { reason }) => {
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
format!("Tool call rejected by hook: {}", reason),
|
||||
));
|
||||
continue;
|
||||
// Hook: BeforeToolCall
|
||||
let event = crate::hooks::HookEvent::ToolCall {
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
user_id: message.user_id.clone(),
|
||||
context: "chat".to_string(),
|
||||
};
|
||||
match self.hooks().run(&event).await {
|
||||
Err(crate::hooks::HookError::Rejected { reason }) => {
|
||||
preflight.push((
|
||||
tc,
|
||||
PreflightOutcome::Rejected(format!(
|
||||
"Tool call rejected by hook: {}",
|
||||
reason
|
||||
)),
|
||||
));
|
||||
continue; // skip to next tool (not infinite: using for loop)
|
||||
}
|
||||
Err(err) => {
|
||||
preflight.push((
|
||||
tc,
|
||||
PreflightOutcome::Rejected(format!(
|
||||
"Tool call blocked by hook policy: {}",
|
||||
err
|
||||
)),
|
||||
));
|
||||
continue;
|
||||
}
|
||||
Ok(crate::hooks::HookOutcome::Continue {
|
||||
modified: Some(new_params),
|
||||
}) => match serde_json::from_str(&new_params) {
|
||||
Ok(parsed) => tc.arguments = parsed,
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
tool = %tc.name,
|
||||
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
|
||||
e
|
||||
);
|
||||
}
|
||||
Err(err) => {
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
format!("Tool call blocked by hook policy: {}", err),
|
||||
));
|
||||
continue;
|
||||
},
|
||||
_ => {}
|
||||
}
|
||||
|
||||
let preflight_idx = preflight.len();
|
||||
preflight.push((tc.clone(), PreflightOutcome::Runnable));
|
||||
runnable.push((preflight_idx, tc));
|
||||
}
|
||||
|
||||
// === Phase 2: Parallel execution ===
|
||||
// Execute runnable tools and slot results back by preflight
|
||||
// index so Phase 3 can iterate in original order.
|
||||
let mut exec_results: Vec<Option<Result<String, Error>>> =
|
||||
(0..preflight.len()).map(|_| None).collect();
|
||||
|
||||
if runnable.len() <= 1 {
|
||||
// Single tool (or none): execute inline
|
||||
for (pf_idx, tc) in &runnable {
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
let result = self
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
exec_results[*pf_idx] = Some(result);
|
||||
}
|
||||
} else {
|
||||
// Multiple tools: execute in parallel via JoinSet
|
||||
let mut join_set = JoinSet::new();
|
||||
|
||||
for (pf_idx, tc) in &runnable {
|
||||
let pf_idx = *pf_idx;
|
||||
let tools = self.tools().clone();
|
||||
let safety = self.safety().clone();
|
||||
let channels = self.channels.clone();
|
||||
let job_ctx = job_ctx.clone();
|
||||
let tc = tc.clone();
|
||||
let channel = message.channel.clone();
|
||||
let metadata = message.metadata.clone();
|
||||
|
||||
join_set.spawn(async move {
|
||||
let _ = channels
|
||||
.send_status(
|
||||
&channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
let result = execute_chat_tool_standalone(
|
||||
&tools,
|
||||
&safety,
|
||||
&tc.name,
|
||||
&tc.arguments,
|
||||
&job_ctx,
|
||||
)
|
||||
.await;
|
||||
|
||||
let _ = channels
|
||||
.send_status(
|
||||
&channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
(pf_idx, result)
|
||||
});
|
||||
}
|
||||
|
||||
while let Some(join_result) = join_set.join_next().await {
|
||||
match join_result {
|
||||
Ok((pf_idx, result)) => {
|
||||
exec_results[pf_idx] = Some(result);
|
||||
}
|
||||
Ok(crate::hooks::HookOutcome::Continue {
|
||||
modified: Some(new_params),
|
||||
}) => match serde_json::from_str(&new_params) {
|
||||
Ok(parsed) => tc.arguments = parsed,
|
||||
Err(e) => {
|
||||
tracing::warn!(
|
||||
tool = %tc.name,
|
||||
"Hook returned non-JSON modification for ToolCall, ignoring: {}",
|
||||
Err(e) => {
|
||||
if e.is_panic() {
|
||||
tracing::error!("Chat tool execution task panicked: {}", e);
|
||||
} else {
|
||||
tracing::error!(
|
||||
"Chat tool execution task cancelled: {}",
|
||||
e
|
||||
);
|
||||
}
|
||||
},
|
||||
_ => {} // Continue, fail-open errors already logged
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
let tool_result = self
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: tool_result.is_ok(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
if let Ok(ref output) = tool_result
|
||||
&& !output.is_empty()
|
||||
{
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolResult {
|
||||
// Fill panicked slots with error results
|
||||
for (runnable_idx, (pf_idx, tc)) in runnable.iter().enumerate() {
|
||||
if exec_results[*pf_idx].is_none() {
|
||||
tracing::error!(
|
||||
tool = %tc.name,
|
||||
runnable_idx,
|
||||
"Filling failed task slot with error"
|
||||
);
|
||||
exec_results[*pf_idx] =
|
||||
Some(Err(crate::error::ToolError::ExecutionFailed {
|
||||
name: tc.name.clone(),
|
||||
preview: output.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
reason: "Task failed during execution".to_string(),
|
||||
}
|
||||
.into()));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Record result in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
match &tool_result {
|
||||
// === Phase 3: Post-flight (sequential, in original order) ===
|
||||
// Process all results — both hook rejections and execution
|
||||
// results — in the original tool_calls order. Auth intercept
|
||||
// is deferred until after every result is recorded.
|
||||
let mut deferred_auth: Option<String> = None;
|
||||
|
||||
for (pf_idx, (tc, outcome)) in preflight.into_iter().enumerate() {
|
||||
match outcome {
|
||||
PreflightOutcome::Rejected(error_msg) => {
|
||||
// Record hook rejection in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
turn.record_tool_error(error_msg.clone());
|
||||
}
|
||||
}
|
||||
context_messages
|
||||
.push(ChatMessage::tool_result(&tc.id, &tc.name, error_msg));
|
||||
}
|
||||
PreflightOutcome::Runnable => {
|
||||
// Retrieve the execution result for this slot
|
||||
let tool_result =
|
||||
exec_results[pf_idx].take().unwrap_or_else(|| {
|
||||
Err(crate::error::ToolError::ExecutionFailed {
|
||||
name: tc.name.clone(),
|
||||
reason: "No result available".to_string(),
|
||||
}
|
||||
.into())
|
||||
});
|
||||
|
||||
// Send ToolResult preview
|
||||
if let Ok(ref output) = tool_result
|
||||
&& !output.is_empty()
|
||||
{
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolResult {
|
||||
name: tc.name.clone(),
|
||||
preview: output.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
// Record result in thread
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id)
|
||||
&& let Some(turn) = thread.last_turn_mut()
|
||||
{
|
||||
match &tool_result {
|
||||
Ok(output) => {
|
||||
turn.record_tool_result(serde_json::json!(output));
|
||||
}
|
||||
Err(e) => {
|
||||
turn.record_tool_error(e.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Check for auth awaiting — defer the return
|
||||
// until all results are recorded.
|
||||
if deferred_auth.is_none()
|
||||
&& let Some((ext_name, instructions)) =
|
||||
check_auth_required(&tc.name, &tool_result)
|
||||
{
|
||||
let auth_data = parse_auth_result(&tool_result);
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
||||
thread.enter_auth_mode(ext_name.clone());
|
||||
}
|
||||
}
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::AuthRequired {
|
||||
extension_name: ext_name,
|
||||
instructions: Some(instructions.clone()),
|
||||
auth_url: auth_data.auth_url,
|
||||
setup_url: auth_data.setup_url,
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
deferred_auth = Some(instructions);
|
||||
}
|
||||
|
||||
// Sanitize and add tool result to context
|
||||
let result_content = match tool_result {
|
||||
Ok(output) => {
|
||||
turn.record_tool_result(serde_json::json!(output));
|
||||
let sanitized =
|
||||
self.safety().sanitize_tool_output(&tc.name, &output);
|
||||
self.safety().wrap_for_llm(
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
}
|
||||
Err(e) => {
|
||||
turn.record_tool_error(e.to_string());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
};
|
||||
|
||||
// If tool_auth returned awaiting_token, enter auth mode
|
||||
// and short-circuit: return the instructions directly so
|
||||
// the LLM doesn't get a chance to hallucinate tool calls.
|
||||
if let Some((ext_name, instructions)) =
|
||||
check_auth_required(&tc.name, &tool_result)
|
||||
{
|
||||
let auth_data = parse_auth_result(&tool_result);
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
||||
thread.enter_auth_mode(ext_name.clone());
|
||||
}
|
||||
}
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::AuthRequired {
|
||||
extension_name: ext_name,
|
||||
instructions: Some(instructions.clone()),
|
||||
auth_url: auth_data.auth_url,
|
||||
setup_url: auth_data.setup_url,
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
return Ok(AgenticLoopResult::Response(instructions));
|
||||
}
|
||||
|
||||
// Add tool result to context for next LLM call
|
||||
let result_content = match tool_result {
|
||||
Ok(output) => {
|
||||
// Sanitize output before showing to LLM
|
||||
let sanitized =
|
||||
self.safety().sanitize_tool_output(&tc.name, &output);
|
||||
self.safety().wrap_for_llm(
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
&sanitized.content,
|
||||
sanitized.was_modified,
|
||||
)
|
||||
result_content,
|
||||
));
|
||||
}
|
||||
Err(e) => format!("Error: {}", e),
|
||||
}
|
||||
}
|
||||
|
||||
// Return auth response after all results are recorded
|
||||
if let Some(instructions) = deferred_auth {
|
||||
return Ok(AgenticLoopResult::Response(instructions));
|
||||
}
|
||||
|
||||
// Handle approval if a tool needed it
|
||||
if let Some((approval_idx, tc, tool)) = approval_needed {
|
||||
let pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
deferred_tool_calls: tool_calls[approval_idx + 1..].to_vec(),
|
||||
};
|
||||
|
||||
context_messages.push(ChatMessage::tool_result(
|
||||
&tc.id,
|
||||
&tc.name,
|
||||
result_content,
|
||||
));
|
||||
|
||||
idx += 1;
|
||||
return Ok(AgenticLoopResult::NeedApproval { pending });
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -462,95 +623,108 @@ impl Agent {
|
||||
params: &serde_json::Value,
|
||||
job_ctx: &JobContext,
|
||||
) -> Result<String, Error> {
|
||||
let tool =
|
||||
self.tools()
|
||||
.get(tool_name)
|
||||
.await
|
||||
.ok_or_else(|| crate::error::ToolError::NotFound {
|
||||
name: tool_name.to_string(),
|
||||
})?;
|
||||
|
||||
// Validate tool parameters
|
||||
let validation = self.safety().validator().validate_tool_params(params);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
return Err(crate::error::ToolError::InvalidParameters {
|
||||
name: tool_name.to_string(),
|
||||
reason: format!("Invalid tool parameters: {}", details),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
params = %params,
|
||||
"Tool call started"
|
||||
);
|
||||
|
||||
// Execute with per-tool timeout
|
||||
let timeout = tool.execution_timeout();
|
||||
let start = std::time::Instant::now();
|
||||
let result = tokio::time::timeout(timeout, async {
|
||||
tool.execute(params.clone(), job_ctx).await
|
||||
})
|
||||
.await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
match &result {
|
||||
Ok(Ok(output)) => {
|
||||
let result_str = serde_json::to_string(&output.result)
|
||||
.unwrap_or_else(|_| "<serialize error>".to_string());
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
result = %result_str,
|
||||
"Tool call succeeded"
|
||||
);
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
error = %e,
|
||||
"Tool call failed"
|
||||
);
|
||||
}
|
||||
Err(_) => {
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
timeout_secs = timeout.as_secs(),
|
||||
"Tool call timed out"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let result = result
|
||||
.map_err(|_| crate::error::ToolError::Timeout {
|
||||
name: tool_name.to_string(),
|
||||
timeout,
|
||||
})?
|
||||
.map_err(|e| crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
// Convert result to string
|
||||
serde_json::to_string_pretty(&result.result).map_err(|e| {
|
||||
crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: format!("Failed to serialize result: {}", e),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
execute_chat_tool_standalone(self.tools(), self.safety(), tool_name, params, job_ctx).await
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute a chat tool without requiring `&Agent`.
|
||||
///
|
||||
/// This standalone function enables parallel invocation from spawned JoinSet
|
||||
/// tasks, which cannot borrow `&self`. It replicates the logic from
|
||||
/// `Agent::execute_chat_tool`.
|
||||
pub(super) async fn execute_chat_tool_standalone(
|
||||
tools: &crate::tools::ToolRegistry,
|
||||
safety: &crate::safety::SafetyLayer,
|
||||
tool_name: &str,
|
||||
params: &serde_json::Value,
|
||||
job_ctx: &crate::context::JobContext,
|
||||
) -> Result<String, Error> {
|
||||
let tool = tools
|
||||
.get(tool_name)
|
||||
.await
|
||||
.ok_or_else(|| crate::error::ToolError::NotFound {
|
||||
name: tool_name.to_string(),
|
||||
})?;
|
||||
|
||||
// Validate tool parameters
|
||||
let validation = safety.validator().validate_tool_params(params);
|
||||
if !validation.is_valid {
|
||||
let details = validation
|
||||
.errors
|
||||
.iter()
|
||||
.map(|e| format!("{}: {}", e.field, e.message))
|
||||
.collect::<Vec<_>>()
|
||||
.join("; ");
|
||||
return Err(crate::error::ToolError::InvalidParameters {
|
||||
name: tool_name.to_string(),
|
||||
reason: format!("Invalid tool parameters: {}", details),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
params = %params,
|
||||
"Tool call started"
|
||||
);
|
||||
|
||||
// Execute with per-tool timeout
|
||||
let timeout = tool.execution_timeout();
|
||||
let start = std::time::Instant::now();
|
||||
let result = tokio::time::timeout(timeout, async {
|
||||
tool.execute(params.clone(), job_ctx).await
|
||||
})
|
||||
.await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
match &result {
|
||||
Ok(Ok(output)) => {
|
||||
let result_str = serde_json::to_string(&output.result)
|
||||
.unwrap_or_else(|_| "<serialize error>".to_string());
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
result = %result_str,
|
||||
"Tool call succeeded"
|
||||
);
|
||||
}
|
||||
Ok(Err(e)) => {
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
error = %e,
|
||||
"Tool call failed"
|
||||
);
|
||||
}
|
||||
Err(_) => {
|
||||
tracing::debug!(
|
||||
tool = %tool_name,
|
||||
elapsed_ms = elapsed.as_millis() as u64,
|
||||
timeout_secs = timeout.as_secs(),
|
||||
"Tool call timed out"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
let result = result
|
||||
.map_err(|_| crate::error::ToolError::Timeout {
|
||||
name: tool_name.to_string(),
|
||||
timeout,
|
||||
})?
|
||||
.map_err(|e| crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
|
||||
serde_json::to_string_pretty(&result.result).map_err(|e| {
|
||||
crate::error::ToolError::ExecutionFailed {
|
||||
name: tool_name.to_string(),
|
||||
reason: format!("Failed to serialize result: {}", e),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
/// Parsed auth result fields for emitting StatusUpdate::AuthRequired.
|
||||
pub(super) struct ParsedAuthData {
|
||||
pub(super) auth_url: Option<String>,
|
||||
@@ -903,4 +1077,62 @@ mod tests {
|
||||
|
||||
assert!(check_auth_required("tool_activate", &result).is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_execute_chat_tool_standalone_success() {
|
||||
use crate::config::SafetyConfig;
|
||||
use crate::context::JobContext;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::ToolRegistry;
|
||||
use crate::tools::builtin::EchoTool;
|
||||
|
||||
let registry = ToolRegistry::new();
|
||||
registry.register(std::sync::Arc::new(EchoTool)).await;
|
||||
|
||||
let safety = SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
});
|
||||
|
||||
let job_ctx = JobContext::with_user("test", "chat", "test session");
|
||||
|
||||
let result = super::execute_chat_tool_standalone(
|
||||
®istry,
|
||||
&safety,
|
||||
"echo",
|
||||
&serde_json::json!({"message": "hello"}),
|
||||
&job_ctx,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(result.is_ok());
|
||||
let output = result.unwrap();
|
||||
assert!(output.contains("hello"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_execute_chat_tool_standalone_not_found() {
|
||||
use crate::config::SafetyConfig;
|
||||
use crate::context::JobContext;
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::ToolRegistry;
|
||||
|
||||
let registry = ToolRegistry::new();
|
||||
let safety = SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
});
|
||||
let job_ctx = JobContext::with_user("test", "chat", "test session");
|
||||
|
||||
let result = super::execute_chat_tool_standalone(
|
||||
®istry,
|
||||
&safety,
|
||||
"nonexistent",
|
||||
&serde_json::json!({}),
|
||||
&job_ctx,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(result.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -91,6 +91,7 @@ impl Session {
|
||||
None => self.create_thread(),
|
||||
Some(id) => {
|
||||
if self.threads.contains_key(&id) {
|
||||
// Safe: contains_key confirmed the entry exists.
|
||||
self.threads.get_mut(&id).unwrap()
|
||||
} else {
|
||||
// Stale active_thread ID: create a new thread, which
|
||||
|
||||
+196
-69
@@ -6,11 +6,14 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
use tokio::task::JoinSet;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::Agent;
|
||||
use crate::agent::compaction::ContextCompactor;
|
||||
use crate::agent::dispatcher::{AgenticLoopResult, check_auth_required, parse_auth_result};
|
||||
use crate::agent::dispatcher::{
|
||||
AgenticLoopResult, check_auth_required, execute_chat_tool_standalone, parse_auth_result,
|
||||
};
|
||||
use crate::agent::session::{PendingApproval, Session, ThreadState};
|
||||
use crate::agent::submission::SubmissionResult;
|
||||
use crate::channels::{IncomingMessage, StatusUpdate};
|
||||
@@ -785,9 +788,17 @@ impl Agent {
|
||||
.await;
|
||||
}
|
||||
|
||||
let mut deferred_queue = std::collections::VecDeque::from(deferred_tool_calls);
|
||||
while let Some(tc) = deferred_queue.pop_front() {
|
||||
// Re-check approval for each deferred tool call
|
||||
// === Phase 1: Preflight (sequential) ===
|
||||
// Walk deferred tools checking approval. Collect runnable
|
||||
// tools; stop at the first that needs approval.
|
||||
let mut runnable: Vec<crate::llm::ToolCall> = Vec::new();
|
||||
let mut approval_needed: Option<(
|
||||
usize,
|
||||
crate::llm::ToolCall,
|
||||
Arc<dyn crate::tools::Tool>,
|
||||
)> = None;
|
||||
|
||||
for (idx, tc) in deferred_tool_calls.iter().enumerate() {
|
||||
if let Some(tool) = self.tools().get(&tc.name).await
|
||||
&& tool.requires_approval()
|
||||
{
|
||||
@@ -801,73 +812,142 @@ impl Agent {
|
||||
};
|
||||
|
||||
if !is_auto_approved {
|
||||
let new_pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
deferred_tool_calls: deferred_queue.iter().cloned().collect(),
|
||||
};
|
||||
|
||||
let request_id = new_pending.request_id;
|
||||
let tool_name = new_pending.tool_name.clone();
|
||||
let description = new_pending.description.clone();
|
||||
let parameters = new_pending.parameters.clone();
|
||||
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
||||
thread.await_approval(new_pending);
|
||||
}
|
||||
}
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::Status("Awaiting approval".into()),
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
return Ok(SubmissionResult::NeedApproval {
|
||||
request_id,
|
||||
tool_name,
|
||||
description,
|
||||
parameters,
|
||||
});
|
||||
approval_needed = Some((idx, tc.clone(), tool));
|
||||
break; // remaining tools stay deferred
|
||||
}
|
||||
}
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
runnable.push(tc.clone());
|
||||
}
|
||||
|
||||
let deferred_result = self
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
// === Phase 2: Parallel execution ===
|
||||
let exec_results: Vec<(crate::llm::ToolCall, Result<String, Error>)> = if runnable.len()
|
||||
<= 1
|
||||
{
|
||||
// Single tool (or none): execute inline
|
||||
let mut results = Vec::new();
|
||||
for tc in &runnable {
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: deferred_result.is_ok(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
let result = self
|
||||
.execute_chat_tool(&tc.name, &tc.arguments, &job_ctx)
|
||||
.await;
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
results.push((tc.clone(), result));
|
||||
}
|
||||
results
|
||||
} else {
|
||||
// Multiple tools: execute in parallel via JoinSet
|
||||
let mut join_set = JoinSet::new();
|
||||
let runnable_count = runnable.len();
|
||||
|
||||
for (spawn_idx, tc) in runnable.iter().enumerate() {
|
||||
let tools = self.tools().clone();
|
||||
let safety = self.safety().clone();
|
||||
let channels = self.channels.clone();
|
||||
let job_ctx = job_ctx.clone();
|
||||
let tc = tc.clone();
|
||||
let channel = message.channel.clone();
|
||||
let metadata = message.metadata.clone();
|
||||
|
||||
join_set.spawn(async move {
|
||||
let _ = channels
|
||||
.send_status(
|
||||
&channel,
|
||||
StatusUpdate::ToolStarted {
|
||||
name: tc.name.clone(),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
let result = execute_chat_tool_standalone(
|
||||
&tools,
|
||||
&safety,
|
||||
&tc.name,
|
||||
&tc.arguments,
|
||||
&job_ctx,
|
||||
)
|
||||
.await;
|
||||
|
||||
let _ = channels
|
||||
.send_status(
|
||||
&channel,
|
||||
StatusUpdate::ToolCompleted {
|
||||
name: tc.name.clone(),
|
||||
success: result.is_ok(),
|
||||
},
|
||||
&metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
(spawn_idx, tc, result)
|
||||
});
|
||||
}
|
||||
|
||||
// Collect and reorder by original index
|
||||
let mut ordered: Vec<Option<(crate::llm::ToolCall, Result<String, Error>)>> =
|
||||
(0..runnable_count).map(|_| None).collect();
|
||||
while let Some(join_result) = join_set.join_next().await {
|
||||
match join_result {
|
||||
Ok((idx, tc, result)) => {
|
||||
ordered[idx] = Some((tc, result));
|
||||
}
|
||||
Err(e) => {
|
||||
if e.is_panic() {
|
||||
tracing::error!("Deferred tool execution task panicked: {}", e);
|
||||
} else {
|
||||
tracing::error!("Deferred tool execution task cancelled: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Fill panicked slots with error results
|
||||
ordered
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(i, opt)| {
|
||||
opt.unwrap_or_else(|| {
|
||||
let tc = runnable[i].clone();
|
||||
let err: Error = crate::error::ToolError::ExecutionFailed {
|
||||
name: tc.name.clone(),
|
||||
reason: "Task failed during execution".to_string(),
|
||||
}
|
||||
.into();
|
||||
(tc, Err(err))
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
};
|
||||
|
||||
// === Phase 3: Post-flight (sequential, in original order) ===
|
||||
// Process all results before any conditional return so every
|
||||
// tool result is recorded in the session audit trail.
|
||||
let mut deferred_auth: Option<String> = None;
|
||||
|
||||
for (tc, deferred_result) in exec_results {
|
||||
if let Ok(ref output) = deferred_result
|
||||
&& !output.is_empty()
|
||||
{
|
||||
@@ -897,9 +977,10 @@ impl Agent {
|
||||
}
|
||||
}
|
||||
|
||||
// Auth detection for deferred tools
|
||||
if let Some((ext_name, instructions)) =
|
||||
check_auth_required(&tc.name, &deferred_result)
|
||||
// Auth detection — defer return until all results are recorded
|
||||
if deferred_auth.is_none()
|
||||
&& let Some((ext_name, instructions)) =
|
||||
check_auth_required(&tc.name, &deferred_result)
|
||||
{
|
||||
self.handle_auth_intercept(
|
||||
&session,
|
||||
@@ -910,7 +991,7 @@ impl Agent {
|
||||
instructions.clone(),
|
||||
)
|
||||
.await;
|
||||
return Ok(SubmissionResult::response(instructions));
|
||||
deferred_auth = Some(instructions);
|
||||
}
|
||||
|
||||
let deferred_content = match deferred_result {
|
||||
@@ -928,6 +1009,52 @@ impl Agent {
|
||||
context_messages.push(ChatMessage::tool_result(&tc.id, &tc.name, deferred_content));
|
||||
}
|
||||
|
||||
// Return auth response after all results are recorded
|
||||
if let Some(instructions) = deferred_auth {
|
||||
return Ok(SubmissionResult::response(instructions));
|
||||
}
|
||||
|
||||
// Handle approval if a tool needed it
|
||||
if let Some((approval_idx, tc, tool)) = approval_needed {
|
||||
let new_pending = PendingApproval {
|
||||
request_id: Uuid::new_v4(),
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
description: tool.description().to_string(),
|
||||
tool_call_id: tc.id.clone(),
|
||||
context_messages: context_messages.clone(),
|
||||
deferred_tool_calls: deferred_tool_calls[approval_idx + 1..].to_vec(),
|
||||
};
|
||||
|
||||
let request_id = new_pending.request_id;
|
||||
let tool_name = new_pending.tool_name.clone();
|
||||
let description = new_pending.description.clone();
|
||||
let parameters = new_pending.parameters.clone();
|
||||
|
||||
{
|
||||
let mut sess = session.lock().await;
|
||||
if let Some(thread) = sess.threads.get_mut(&thread_id) {
|
||||
thread.await_approval(new_pending);
|
||||
}
|
||||
}
|
||||
|
||||
let _ = self
|
||||
.channels
|
||||
.send_status(
|
||||
&message.channel,
|
||||
StatusUpdate::Status("Awaiting approval".into()),
|
||||
&message.metadata,
|
||||
)
|
||||
.await;
|
||||
|
||||
return Ok(SubmissionResult::NeedApproval {
|
||||
request_id,
|
||||
tool_name,
|
||||
description,
|
||||
parameters,
|
||||
});
|
||||
}
|
||||
|
||||
// Continue the agentic loop (a tool was already executed this turn)
|
||||
let result = self
|
||||
.run_agentic_loop(message, session.clone(), thread_id, context_messages, true)
|
||||
|
||||
+282
-22
@@ -3,8 +3,8 @@
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures::future::join_all;
|
||||
use tokio::sync::mpsc;
|
||||
use tokio::task::JoinSet;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::scheduler::WorkerMessage;
|
||||
@@ -292,19 +292,21 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
tool_calls.clone(),
|
||||
));
|
||||
|
||||
for tc in tool_calls {
|
||||
let result = self.execute_tool(&tc.name, &tc.arguments).await;
|
||||
|
||||
// Create synthetic selection for process_tool_result
|
||||
let selection = ToolSelection {
|
||||
// Convert ToolCalls to ToolSelections and execute in parallel
|
||||
let selections: Vec<ToolSelection> = tool_calls
|
||||
.iter()
|
||||
.map(|tc| ToolSelection {
|
||||
tool_name: tc.name.clone(),
|
||||
parameters: tc.arguments.clone(),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: tc.id.clone(),
|
||||
};
|
||||
})
|
||||
.collect();
|
||||
|
||||
self.process_tool_result(reason_ctx, &selection, result)
|
||||
let results = self.execute_tools_parallel(&selections).await;
|
||||
for (selection, result) in selections.iter().zip(results) {
|
||||
self.process_tool_result(reason_ctx, selection, result.result)
|
||||
.await?;
|
||||
}
|
||||
}
|
||||
@@ -347,24 +349,71 @@ Report when the job is complete or if you encounter issues you cannot resolve."#
|
||||
}
|
||||
}
|
||||
|
||||
/// Execute multiple tools in parallel.
|
||||
/// Execute multiple tools in parallel using a JoinSet.
|
||||
///
|
||||
/// Each task is tagged with its original index so results are returned
|
||||
/// in the same order as `selections`, regardless of completion order.
|
||||
async fn execute_tools_parallel(&self, selections: &[ToolSelection]) -> Vec<ToolExecResult> {
|
||||
let futures: Vec<_> = selections
|
||||
.iter()
|
||||
.map(|selection| {
|
||||
let tool_name = selection.tool_name.clone();
|
||||
let params = selection.parameters.clone();
|
||||
let deps = self.deps.clone();
|
||||
let job_id = self.job_id;
|
||||
let count = selections.len();
|
||||
|
||||
async move {
|
||||
let result = Self::execute_tool_inner(&deps, job_id, &tool_name, ¶ms).await;
|
||||
ToolExecResult { result }
|
||||
// Short-circuit for single tool: execute directly without JoinSet overhead
|
||||
if count <= 1 {
|
||||
let mut results = Vec::with_capacity(count);
|
||||
for selection in selections {
|
||||
let result = Self::execute_tool_inner(
|
||||
&self.deps,
|
||||
self.job_id,
|
||||
&selection.tool_name,
|
||||
&selection.parameters,
|
||||
)
|
||||
.await;
|
||||
results.push(ToolExecResult { result });
|
||||
}
|
||||
return results;
|
||||
}
|
||||
|
||||
let mut join_set = JoinSet::new();
|
||||
|
||||
for (idx, selection) in selections.iter().enumerate() {
|
||||
let deps = self.deps.clone();
|
||||
let job_id = self.job_id;
|
||||
let tool_name = selection.tool_name.clone();
|
||||
let params = selection.parameters.clone();
|
||||
join_set.spawn(async move {
|
||||
let result = Self::execute_tool_inner(&deps, job_id, &tool_name, ¶ms).await;
|
||||
(idx, ToolExecResult { result })
|
||||
});
|
||||
}
|
||||
|
||||
// Collect and reorder by original index
|
||||
let mut results: Vec<Option<ToolExecResult>> = (0..count).map(|_| None).collect();
|
||||
while let Some(join_result) = join_set.join_next().await {
|
||||
match join_result {
|
||||
Ok((idx, exec_result)) => results[idx] = Some(exec_result),
|
||||
Err(e) => {
|
||||
if e.is_panic() {
|
||||
tracing::error!("Tool execution task panicked: {}", e);
|
||||
} else {
|
||||
tracing::error!("Tool execution task cancelled: {}", e);
|
||||
}
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
}
|
||||
}
|
||||
|
||||
join_all(futures).await
|
||||
// Fill any panicked slots with error results
|
||||
results
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(i, opt)| {
|
||||
opt.unwrap_or_else(|| ToolExecResult {
|
||||
result: Err(crate::error::ToolError::ExecutionFailed {
|
||||
name: selections[i].tool_name.clone(),
|
||||
reason: "Task failed during execution".to_string(),
|
||||
}
|
||||
.into()),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Inner tool execution logic that can be called from both single and parallel paths.
|
||||
@@ -823,6 +872,102 @@ mod tests {
|
||||
use crate::llm::ToolSelection;
|
||||
use crate::util::llm_signals_completion;
|
||||
|
||||
use super::*;
|
||||
use crate::config::SafetyConfig;
|
||||
use crate::context::JobContext;
|
||||
use crate::llm::{
|
||||
CompletionRequest, CompletionResponse, LlmProvider, ToolCompletionRequest,
|
||||
ToolCompletionResponse,
|
||||
};
|
||||
use crate::safety::SafetyLayer;
|
||||
use crate::tools::{Tool, ToolError, ToolOutput};
|
||||
|
||||
/// A test tool that sleeps for a configurable duration before returning.
|
||||
struct SlowTool {
|
||||
tool_name: String,
|
||||
delay: Duration,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl Tool for SlowTool {
|
||||
fn name(&self) -> &str {
|
||||
&self.tool_name
|
||||
}
|
||||
fn description(&self) -> &str {
|
||||
"Test tool with configurable delay"
|
||||
}
|
||||
fn parameters_schema(&self) -> serde_json::Value {
|
||||
serde_json::json!({"type": "object", "properties": {}})
|
||||
}
|
||||
async fn execute(
|
||||
&self,
|
||||
_params: serde_json::Value,
|
||||
_ctx: &JobContext,
|
||||
) -> Result<ToolOutput, ToolError> {
|
||||
let start = std::time::Instant::now();
|
||||
tokio::time::sleep(self.delay).await;
|
||||
Ok(ToolOutput::text(
|
||||
format!("done_{}", self.tool_name),
|
||||
start.elapsed(),
|
||||
))
|
||||
}
|
||||
fn requires_sanitization(&self) -> bool {
|
||||
false
|
||||
}
|
||||
}
|
||||
|
||||
/// Stub LLM provider (never called in these tests).
|
||||
struct StubLlm;
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl LlmProvider for StubLlm {
|
||||
fn model_name(&self) -> &str {
|
||||
"stub"
|
||||
}
|
||||
fn cost_per_token(&self) -> (rust_decimal::Decimal, rust_decimal::Decimal) {
|
||||
(rust_decimal::Decimal::ZERO, rust_decimal::Decimal::ZERO)
|
||||
}
|
||||
async fn complete(
|
||||
&self,
|
||||
_req: CompletionRequest,
|
||||
) -> Result<CompletionResponse, crate::error::LlmError> {
|
||||
unimplemented!("stub")
|
||||
}
|
||||
async fn complete_with_tools(
|
||||
&self,
|
||||
_req: ToolCompletionRequest,
|
||||
) -> Result<ToolCompletionResponse, crate::error::LlmError> {
|
||||
unimplemented!("stub")
|
||||
}
|
||||
}
|
||||
|
||||
/// Build a Worker wired to a ToolRegistry containing the given tools.
|
||||
async fn make_worker(tools: Vec<Arc<dyn Tool>>) -> Worker {
|
||||
let registry = ToolRegistry::new();
|
||||
for t in tools {
|
||||
registry.register(t).await;
|
||||
}
|
||||
|
||||
let cm = Arc::new(crate::context::ContextManager::new(5));
|
||||
let job_id = cm.create_job("test", "test job").await.unwrap();
|
||||
|
||||
let deps = WorkerDeps {
|
||||
context_manager: cm,
|
||||
llm: Arc::new(StubLlm),
|
||||
safety: Arc::new(SafetyLayer::new(&SafetyConfig {
|
||||
max_output_length: 100_000,
|
||||
injection_check_enabled: false,
|
||||
})),
|
||||
tools: Arc::new(registry),
|
||||
store: None,
|
||||
hooks: Arc::new(crate::hooks::HookRegistry::new()),
|
||||
timeout: Duration::from_secs(30),
|
||||
use_planning: false,
|
||||
};
|
||||
|
||||
Worker::new(job_id, deps)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_tool_selection_preserves_call_id() {
|
||||
let selection = ToolSelection {
|
||||
@@ -899,4 +1044,119 @@ mod tests {
|
||||
"The tool returned: TASK_COMPLETE signal"
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_parallel_speedup() {
|
||||
// 3 tools each sleeping 200ms should finish in roughly 200ms (parallel),
|
||||
// not ~600ms (sequential).
|
||||
let tools: Vec<Arc<dyn Tool>> = (0..3)
|
||||
.map(|i| {
|
||||
Arc::new(SlowTool {
|
||||
tool_name: format!("slow_{}", i),
|
||||
delay: Duration::from_millis(200),
|
||||
}) as Arc<dyn Tool>
|
||||
})
|
||||
.collect();
|
||||
|
||||
let worker = make_worker(tools).await;
|
||||
|
||||
let selections: Vec<ToolSelection> = (0..3)
|
||||
.map(|i| ToolSelection {
|
||||
tool_name: format!("slow_{}", i),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: format!("call_{}", i),
|
||||
})
|
||||
.collect();
|
||||
|
||||
let start = std::time::Instant::now();
|
||||
let results = worker.execute_tools_parallel(&selections).await;
|
||||
let elapsed = start.elapsed();
|
||||
|
||||
assert_eq!(results.len(), 3);
|
||||
for r in &results {
|
||||
assert!(r.result.is_ok(), "Tool should succeed");
|
||||
}
|
||||
// Parallel should complete well under the sequential 600ms threshold.
|
||||
assert!(
|
||||
elapsed < Duration::from_millis(500),
|
||||
"Parallel execution took {:?}, expected < 500ms",
|
||||
elapsed
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_result_ordering_preserved() {
|
||||
// Tools with different delays finish in different order.
|
||||
// Results must be returned in the original request order.
|
||||
let tools: Vec<Arc<dyn Tool>> = vec![
|
||||
Arc::new(SlowTool {
|
||||
tool_name: "tool_a".into(),
|
||||
delay: Duration::from_millis(300),
|
||||
}),
|
||||
Arc::new(SlowTool {
|
||||
tool_name: "tool_b".into(),
|
||||
delay: Duration::from_millis(100),
|
||||
}),
|
||||
Arc::new(SlowTool {
|
||||
tool_name: "tool_c".into(),
|
||||
delay: Duration::from_millis(200),
|
||||
}),
|
||||
];
|
||||
|
||||
let worker = make_worker(tools).await;
|
||||
|
||||
let selections = vec![
|
||||
ToolSelection {
|
||||
tool_name: "tool_a".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_a".into(),
|
||||
},
|
||||
ToolSelection {
|
||||
tool_name: "tool_b".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_b".into(),
|
||||
},
|
||||
ToolSelection {
|
||||
tool_name: "tool_c".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_c".into(),
|
||||
},
|
||||
];
|
||||
|
||||
let results = worker.execute_tools_parallel(&selections).await;
|
||||
|
||||
// Results must be in same order as selections, not completion order.
|
||||
assert!(results[0].result.as_ref().unwrap().contains("done_tool_a"));
|
||||
assert!(results[1].result.as_ref().unwrap().contains("done_tool_b"));
|
||||
assert!(results[2].result.as_ref().unwrap().contains("done_tool_c"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_missing_tool_produces_error_not_panic() {
|
||||
// If a tool doesn't exist, the result slot should contain an error.
|
||||
let worker = make_worker(vec![]).await;
|
||||
|
||||
let selections = vec![ToolSelection {
|
||||
tool_name: "nonexistent_tool".into(),
|
||||
parameters: serde_json::json!({}),
|
||||
reasoning: String::new(),
|
||||
alternatives: vec![],
|
||||
tool_call_id: "call_x".into(),
|
||||
}];
|
||||
|
||||
let results = worker.execute_tools_parallel(&selections).await;
|
||||
assert_eq!(results.len(), 1);
|
||||
assert!(
|
||||
results[0].result.is_err(),
|
||||
"Missing tool should produce an error, not a panic"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user