mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-30 08:17:53 +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
+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