diff --git a/src/agent/dispatcher.rs b/src/agent/dispatcher.rs index d2a60717..78d23c31 100644 --- a/src/agent/dispatcher.rs +++ b/src/agent/dispatcher.rs @@ -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, + )> = 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>> = + (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 = 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 { - 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::>() - .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(|_| "".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 { + 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::>() + .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(|_| "".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, @@ -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()); + } } diff --git a/src/agent/session.rs b/src/agent/session.rs index 364e6813..87a1e1e4 100644 --- a/src/agent/session.rs +++ b/src/agent/session.rs @@ -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 diff --git a/src/agent/thread_ops.rs b/src/agent/thread_ops.rs index e7db83ca..66f3723c 100644 --- a/src/agent/thread_ops.rs +++ b/src/agent/thread_ops.rs @@ -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 = Vec::new(); + let mut approval_needed: Option<( + usize, + crate::llm::ToolCall, + Arc, + )> = 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)> = 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)>> = + (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 = 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) diff --git a/src/agent/worker.rs b/src/agent/worker.rs index b2ba8a4e..50770af7 100644 --- a/src/agent/worker.rs +++ b/src/agent/worker.rs @@ -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 = 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 { - 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> = (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 { + 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 { + unimplemented!("stub") + } + async fn complete_with_tools( + &self, + _req: ToolCompletionRequest, + ) -> Result { + unimplemented!("stub") + } + } + + /// Build a Worker wired to a ToolRegistry containing the given tools. + async fn make_worker(tools: Vec>) -> 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> = (0..3) + .map(|i| { + Arc::new(SlowTool { + tool_name: format!("slow_{}", i), + delay: Duration::from_millis(200), + }) as Arc + }) + .collect(); + + let worker = make_worker(tools).await; + + let selections: Vec = (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> = 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" + ); + } }