feat(orchestration): add safe parallel tool execution with metrics
This commit is contained in:
parent
11954b532a
commit
65aa1b04f2
19 changed files with 1384 additions and 84 deletions
|
|
@ -280,10 +280,14 @@
|
||||||
"orchestration": {
|
"orchestration": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
"max_spawn_depth": 3,
|
"max_spawn_depth": 3,
|
||||||
"max_parallel_workers": 4,
|
"max_parallel_workers": 8,
|
||||||
"max_tasks_per_agent": 20,
|
"max_tasks_per_agent": 20,
|
||||||
"default_task_timeout_seconds": 180,
|
"default_task_timeout_seconds": 180,
|
||||||
"retry_limit_per_task": 2
|
"retry_limit_per_task": 2,
|
||||||
|
"tool_calls_parallel_enabled": true,
|
||||||
|
"max_tool_call_concurrency": 8,
|
||||||
|
"parallel_tools_mode": "read_only_only",
|
||||||
|
"tool_parallel_overrides": {}
|
||||||
},
|
},
|
||||||
"audit": {
|
"audit": {
|
||||||
"enabled": false,
|
"enabled": false,
|
||||||
|
|
|
||||||
|
|
@ -20,6 +20,20 @@ This document describes how PicoClaw's subagent orchestration and periodic audit
|
||||||
- When depth is exceeded, spawn is rejected with `max spawn depth reached`.
|
- When depth is exceeded, spawn is rejected with `max spawn depth reached`.
|
||||||
- `max_parallel_workers`
|
- `max_parallel_workers`
|
||||||
- Enforced as max concurrent running tasks per manager.
|
- Enforced as max concurrent running tasks per manager.
|
||||||
|
- `tool_calls_parallel_enabled`
|
||||||
|
- Enables parallel execution for eligible tool calls in one LLM turn.
|
||||||
|
- `max_tool_call_concurrency`
|
||||||
|
- Bounded worker count for one tool-call batch (`<=0` means no explicit cap).
|
||||||
|
- `parallel_tools_mode`
|
||||||
|
- `read_only_only` (default): only tools marked `parallel_read_only` are parallelized.
|
||||||
|
- `all`: all tool calls are eligible only when the tool instance is concurrent-safe.
|
||||||
|
- `tool_parallel_overrides`
|
||||||
|
- Optional per-tool override map.
|
||||||
|
- Values:
|
||||||
|
- `parallel_read_only`: force this tool to be parallel-eligible.
|
||||||
|
- `serial_only`: force this tool to execute serially.
|
||||||
|
- Overrides take precedence over built-in tool policy and mode defaults.
|
||||||
|
- Overrides do not bypass instance safety checks.
|
||||||
- `max_tasks_per_agent`
|
- `max_tasks_per_agent`
|
||||||
- Enforced as max active (non-terminal) tasks per manager.
|
- Enforced as max active (non-terminal) tasks per manager.
|
||||||
- `default_task_timeout_seconds`
|
- `default_task_timeout_seconds`
|
||||||
|
|
@ -109,4 +123,3 @@ Optional model checks:
|
||||||
- Existing tools and loop behavior remain unchanged when `audit.enabled=false`.
|
- Existing tools and loop behavior remain unchanged when `audit.enabled=false`.
|
||||||
- New fields are additive and optional.
|
- New fields are additive and optional.
|
||||||
- `spawn` and `subagent` retain previous parameter contract; `agent_id` is additive for `subagent`.
|
- `spawn` and `subagent` retain previous parameter contract; `agent_id` is additive for `subagent`.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -167,6 +167,12 @@ func registerSharedTools(
|
||||||
cfg.Orchestration.MaxTasksPerAgent,
|
cfg.Orchestration.MaxTasksPerAgent,
|
||||||
cfg.Orchestration.MaxSpawnDepth,
|
cfg.Orchestration.MaxSpawnDepth,
|
||||||
)
|
)
|
||||||
|
subagentManager.SetToolCallParallelism(
|
||||||
|
cfg.Orchestration.ToolCallsParallelEnabled,
|
||||||
|
cfg.Orchestration.MaxToolCallConcurrency,
|
||||||
|
cfg.Orchestration.ParallelToolsMode,
|
||||||
|
cfg.Orchestration.ToolParallelOverrides,
|
||||||
|
)
|
||||||
subagentManager.SetTools(agent.Tools)
|
subagentManager.SetTools(agent.Tools)
|
||||||
currentAgentID := agentID
|
currentAgentID := agentID
|
||||||
subagentManager.SetExecutionResolver(func(targetAgentID string) (tools.SubagentExecutionConfig, error) {
|
subagentManager.SetExecutionResolver(func(targetAgentID string) (tools.SubagentExecutionConfig, error) {
|
||||||
|
|
@ -879,44 +885,48 @@ func (al *AgentLoop) runLLMIteration(
|
||||||
// Save assistant message with tool calls to session
|
// Save assistant message with tool calls to session
|
||||||
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
|
agent.Sessions.AddFullMessage(opts.SessionKey, assistantMsg)
|
||||||
|
|
||||||
// Execute tool calls
|
parallelCfg := tools.ToolCallParallelConfig{
|
||||||
for _, tc := range normalizedToolCalls {
|
Enabled: al.cfg != nil && al.cfg.Orchestration.ToolCallsParallelEnabled,
|
||||||
argsJSON, _ := json.Marshal(tc.Arguments)
|
MaxConcurrency: 0,
|
||||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
Mode: "",
|
||||||
logger.InfoCF("agent", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
}
|
||||||
map[string]any{
|
if al.cfg != nil {
|
||||||
"agent_id": agent.ID,
|
parallelCfg.MaxConcurrency = al.cfg.Orchestration.MaxToolCallConcurrency
|
||||||
"tool": tc.Name,
|
parallelCfg.Mode = al.cfg.Orchestration.ParallelToolsMode
|
||||||
"iteration": iteration,
|
parallelCfg.ToolPolicyOverrides = al.cfg.Orchestration.ToolParallelOverrides
|
||||||
})
|
}
|
||||||
|
|
||||||
// Create async callback for tools that implement AsyncTool
|
toolExecutions := tools.ExecuteToolCalls(ctx, agent.Tools, normalizedToolCalls, tools.ToolCallExecutionOptions{
|
||||||
// NOTE: Following openclaw's design, async tools do NOT send results directly to users.
|
Channel: opts.Channel,
|
||||||
// Instead, they notify the agent via PublishInbound, and the agent decides
|
ChatID: opts.ChatID,
|
||||||
// whether to forward the result to the user (in processSystemMessage).
|
SenderID: opts.SenderID,
|
||||||
asyncCallback := func(callbackCtx context.Context, result *tools.ToolResult) {
|
Iteration: iteration,
|
||||||
// Log the async completion but don't send directly to user
|
LogScope: "agent",
|
||||||
// The agent will handle user notification via processSystemMessage
|
Parallel: parallelCfg,
|
||||||
|
// Create async callback for tools that implement AsyncTool.
|
||||||
|
// Following openclaw's design, async tools do not send results directly
|
||||||
|
// to users. The agent handles user notification via processSystemMessage.
|
||||||
|
AsyncCallbackForCall: func(call providers.ToolCall) tools.AsyncCallback {
|
||||||
|
return func(callbackCtx context.Context, result *tools.ToolResult) {
|
||||||
|
if result == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
if !result.Silent && result.ForUser != "" {
|
if !result.Silent && result.ForUser != "" {
|
||||||
logger.InfoCF("agent", "Async tool completed, agent will handle notification",
|
logger.InfoCF("agent", "Async tool completed, agent will handle notification",
|
||||||
map[string]any{
|
map[string]any{
|
||||||
"tool": tc.Name,
|
"tool": call.Name,
|
||||||
"content_len": len(result.ForUser),
|
"content_len": len(result.ForUser),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
toolResult := agent.Tools.ExecuteWithContext(
|
for _, executed := range toolExecutions {
|
||||||
ctx,
|
toolResult := executed.Result
|
||||||
tc.Name,
|
tc := executed.ToolCall
|
||||||
tc.Arguments,
|
|
||||||
opts.Channel,
|
|
||||||
opts.ChatID,
|
|
||||||
opts.SenderID,
|
|
||||||
asyncCallback,
|
|
||||||
)
|
|
||||||
|
|
||||||
// Send ForUser content to user immediately if not Silent
|
// Send ForUser content to user immediately if not Silent.
|
||||||
if !toolResult.Silent && toolResult.ForUser != "" && opts.SendResponse {
|
if !toolResult.Silent && toolResult.ForUser != "" && opts.SendResponse {
|
||||||
al.bus.PublishOutbound(bus.OutboundMessage{
|
al.bus.PublishOutbound(bus.OutboundMessage{
|
||||||
Channel: opts.Channel,
|
Channel: opts.Channel,
|
||||||
|
|
|
||||||
152
pkg/agent/loop_parallel_test.go
Normal file
152
pkg/agent/loop_parallel_test.go
Normal file
|
|
@ -0,0 +1,152 @@
|
||||||
|
package agent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/bus"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
|
)
|
||||||
|
|
||||||
|
type parallelLoopMockProvider struct {
|
||||||
|
callCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *parallelLoopMockProvider) Chat(
|
||||||
|
_ context.Context,
|
||||||
|
messages []providers.Message,
|
||||||
|
_ []providers.ToolDefinition,
|
||||||
|
_ string,
|
||||||
|
_ map[string]any,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
m.callCount++
|
||||||
|
if m.callCount == 1 {
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
ToolCalls: []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "slow_parallel", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "fast_parallel", Arguments: map[string]any{}},
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if m.callCount == 2 {
|
||||||
|
toolMessages := make([]providers.Message, 0, 2)
|
||||||
|
for _, msg := range messages {
|
||||||
|
if msg.Role == "tool" {
|
||||||
|
toolMessages = append(toolMessages, msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(toolMessages) != 2 {
|
||||||
|
return nil, fmt.Errorf("tool message count = %d, want 2", len(toolMessages))
|
||||||
|
}
|
||||||
|
if toolMessages[0].ToolCallID != "tc-1" || toolMessages[0].Content != "slow-ok" {
|
||||||
|
return nil, fmt.Errorf("first tool message = %+v, want tc-1/slow-ok", toolMessages[0])
|
||||||
|
}
|
||||||
|
if toolMessages[1].ToolCallID != "tc-2" || toolMessages[1].Content != "fast-ok" {
|
||||||
|
return nil, fmt.Errorf("second tool message = %+v, want tc-2/fast-ok", toolMessages[1])
|
||||||
|
}
|
||||||
|
return &providers.LLMResponse{Content: "final-from-provider"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return &providers.LLMResponse{Content: "unexpected-extra-call"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *parallelLoopMockProvider) GetDefaultModel() string {
|
||||||
|
return "parallel-loop-mock"
|
||||||
|
}
|
||||||
|
|
||||||
|
type parallelTestTool struct {
|
||||||
|
name string
|
||||||
|
result string
|
||||||
|
delay time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *parallelTestTool) Name() string {
|
||||||
|
return t.name
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *parallelTestTool) Description() string {
|
||||||
|
return "parallel test tool"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *parallelTestTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *parallelTestTool) ParallelPolicy() tools.ToolParallelPolicy {
|
||||||
|
return tools.ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *parallelTestTool) Execute(_ context.Context, _ map[string]any) *tools.ToolResult {
|
||||||
|
if t.delay > 0 {
|
||||||
|
time.Sleep(t.delay)
|
||||||
|
}
|
||||||
|
return tools.SilentResult(t.result)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAgentLoop_RunLLMIteration_ParallelToolCallsPreserveOrder(t *testing.T) {
|
||||||
|
tmpDir, err := os.MkdirTemp("", "agent-loop-parallel-*")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Failed to create temp dir: %v", err)
|
||||||
|
}
|
||||||
|
defer os.RemoveAll(tmpDir)
|
||||||
|
|
||||||
|
cfg := config.DefaultConfig()
|
||||||
|
cfg.Agents.Defaults.Workspace = tmpDir
|
||||||
|
cfg.Agents.Defaults.Model = "test-model"
|
||||||
|
cfg.Agents.Defaults.MaxToolIterations = 4
|
||||||
|
cfg.Orchestration.ToolCallsParallelEnabled = true
|
||||||
|
cfg.Orchestration.MaxToolCallConcurrency = 8
|
||||||
|
cfg.Orchestration.ParallelToolsMode = tools.ParallelToolsModeReadOnlyOnly
|
||||||
|
|
||||||
|
msgBus := bus.NewMessageBus()
|
||||||
|
provider := ¶llelLoopMockProvider{}
|
||||||
|
al := NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
al.RegisterTool(¶llelTestTool{
|
||||||
|
name: "slow_parallel",
|
||||||
|
result: "slow-ok",
|
||||||
|
delay: 50 * time.Millisecond,
|
||||||
|
})
|
||||||
|
al.RegisterTool(¶llelTestTool{
|
||||||
|
name: "fast_parallel",
|
||||||
|
result: "fast-ok",
|
||||||
|
delay: 5 * time.Millisecond,
|
||||||
|
})
|
||||||
|
|
||||||
|
agent := al.registry.GetDefaultAgent()
|
||||||
|
if agent == nil {
|
||||||
|
t.Fatal("default agent not found")
|
||||||
|
}
|
||||||
|
|
||||||
|
finalContent, iterations, err := al.runLLMIteration(
|
||||||
|
context.Background(),
|
||||||
|
agent,
|
||||||
|
[]providers.Message{
|
||||||
|
{Role: "system", Content: "you are a test assistant"},
|
||||||
|
{Role: "user", Content: "run parallel tools"},
|
||||||
|
},
|
||||||
|
processOptions{
|
||||||
|
SessionKey: "parallel-session",
|
||||||
|
Channel: "cli",
|
||||||
|
ChatID: "direct",
|
||||||
|
SenderID: "tester",
|
||||||
|
SendResponse: false,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("runLLMIteration() error = %v", err)
|
||||||
|
}
|
||||||
|
if iterations != 2 {
|
||||||
|
t.Fatalf("iterations = %d, want 2", iterations)
|
||||||
|
}
|
||||||
|
if finalContent != "final-from-provider" {
|
||||||
|
t.Fatalf("finalContent = %q, want %q", finalContent, "final-from-provider")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -363,6 +363,10 @@ type OrchestrationConfig struct {
|
||||||
MaxTasksPerAgent int `json:"max_tasks_per_agent" env:"PICOCLAW_ORCHESTRATION_MAX_TASKS_PER_AGENT"`
|
MaxTasksPerAgent int `json:"max_tasks_per_agent" env:"PICOCLAW_ORCHESTRATION_MAX_TASKS_PER_AGENT"`
|
||||||
DefaultTaskTimeoutSeconds int `json:"default_task_timeout_seconds" env:"PICOCLAW_ORCHESTRATION_DEFAULT_TASK_TIMEOUT_SECONDS"`
|
DefaultTaskTimeoutSeconds int `json:"default_task_timeout_seconds" env:"PICOCLAW_ORCHESTRATION_DEFAULT_TASK_TIMEOUT_SECONDS"`
|
||||||
RetryLimitPerTask int `json:"retry_limit_per_task" env:"PICOCLAW_ORCHESTRATION_RETRY_LIMIT_PER_TASK"`
|
RetryLimitPerTask int `json:"retry_limit_per_task" env:"PICOCLAW_ORCHESTRATION_RETRY_LIMIT_PER_TASK"`
|
||||||
|
ToolCallsParallelEnabled bool `json:"tool_calls_parallel_enabled" env:"PICOCLAW_ORCHESTRATION_TOOL_CALLS_PARALLEL_ENABLED"`
|
||||||
|
MaxToolCallConcurrency int `json:"max_tool_call_concurrency" env:"PICOCLAW_ORCHESTRATION_MAX_TOOL_CALL_CONCURRENCY"`
|
||||||
|
ParallelToolsMode string `json:"parallel_tools_mode" env:"PICOCLAW_ORCHESTRATION_PARALLEL_TOOLS_MODE"`
|
||||||
|
ToolParallelOverrides map[string]string `json:"tool_parallel_overrides,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type AuditConfig struct {
|
type AuditConfig struct {
|
||||||
|
|
|
||||||
|
|
@ -367,12 +367,27 @@ func TestDefaultConfig_OrchestrationAndAuditDefaults(t *testing.T) {
|
||||||
if cfg.Orchestration.MaxSpawnDepth != 3 {
|
if cfg.Orchestration.MaxSpawnDepth != 3 {
|
||||||
t.Fatalf("MaxSpawnDepth = %d, want 3", cfg.Orchestration.MaxSpawnDepth)
|
t.Fatalf("MaxSpawnDepth = %d, want 3", cfg.Orchestration.MaxSpawnDepth)
|
||||||
}
|
}
|
||||||
|
if cfg.Orchestration.MaxParallelWorkers != 8 {
|
||||||
|
t.Fatalf("MaxParallelWorkers = %d, want 8", cfg.Orchestration.MaxParallelWorkers)
|
||||||
|
}
|
||||||
if cfg.Orchestration.DefaultTaskTimeoutSeconds != 180 {
|
if cfg.Orchestration.DefaultTaskTimeoutSeconds != 180 {
|
||||||
t.Fatalf(
|
t.Fatalf(
|
||||||
"DefaultTaskTimeoutSeconds = %d, want 180",
|
"DefaultTaskTimeoutSeconds = %d, want 180",
|
||||||
cfg.Orchestration.DefaultTaskTimeoutSeconds,
|
cfg.Orchestration.DefaultTaskTimeoutSeconds,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
if !cfg.Orchestration.ToolCallsParallelEnabled {
|
||||||
|
t.Fatal("ToolCallsParallelEnabled should be true by default")
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.MaxToolCallConcurrency != 8 {
|
||||||
|
t.Fatalf("MaxToolCallConcurrency = %d, want 8", cfg.Orchestration.MaxToolCallConcurrency)
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.ParallelToolsMode != "read_only_only" {
|
||||||
|
t.Fatalf("ParallelToolsMode = %q, want %q", cfg.Orchestration.ParallelToolsMode, "read_only_only")
|
||||||
|
}
|
||||||
|
if len(cfg.Orchestration.ToolParallelOverrides) != 0 {
|
||||||
|
t.Fatalf("ToolParallelOverrides len = %d, want 0", len(cfg.Orchestration.ToolParallelOverrides))
|
||||||
|
}
|
||||||
if cfg.Audit.IntervalMinutes != 30 {
|
if cfg.Audit.IntervalMinutes != 30 {
|
||||||
t.Fatalf("Audit.IntervalMinutes = %d, want 30", cfg.Audit.IntervalMinutes)
|
t.Fatalf("Audit.IntervalMinutes = %d, want 30", cfg.Audit.IntervalMinutes)
|
||||||
}
|
}
|
||||||
|
|
@ -389,7 +404,14 @@ func TestConfig_UnmarshalAuditAndOrchestration(t *testing.T) {
|
||||||
"orchestration": {
|
"orchestration": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
"max_spawn_depth": 4,
|
"max_spawn_depth": 4,
|
||||||
"max_parallel_workers": 2
|
"max_parallel_workers": 2,
|
||||||
|
"tool_calls_parallel_enabled": false,
|
||||||
|
"max_tool_call_concurrency": 3,
|
||||||
|
"parallel_tools_mode": "all",
|
||||||
|
"tool_parallel_overrides": {
|
||||||
|
"write_file": "serial_only",
|
||||||
|
"exec": "parallel_read_only"
|
||||||
|
}
|
||||||
},
|
},
|
||||||
"audit": {
|
"audit": {
|
||||||
"enabled": true,
|
"enabled": true,
|
||||||
|
|
@ -418,6 +440,30 @@ func TestConfig_UnmarshalAuditAndOrchestration(t *testing.T) {
|
||||||
if cfg.Orchestration.MaxSpawnDepth != 4 {
|
if cfg.Orchestration.MaxSpawnDepth != 4 {
|
||||||
t.Fatalf("max_spawn_depth = %d, want 4", cfg.Orchestration.MaxSpawnDepth)
|
t.Fatalf("max_spawn_depth = %d, want 4", cfg.Orchestration.MaxSpawnDepth)
|
||||||
}
|
}
|
||||||
|
if cfg.Orchestration.MaxParallelWorkers != 2 {
|
||||||
|
t.Fatalf("max_parallel_workers = %d, want 2", cfg.Orchestration.MaxParallelWorkers)
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.ToolCallsParallelEnabled {
|
||||||
|
t.Fatal("tool_calls_parallel_enabled should be false")
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.MaxToolCallConcurrency != 3 {
|
||||||
|
t.Fatalf("max_tool_call_concurrency = %d, want 3", cfg.Orchestration.MaxToolCallConcurrency)
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.ParallelToolsMode != "all" {
|
||||||
|
t.Fatalf("parallel_tools_mode = %q, want %q", cfg.Orchestration.ParallelToolsMode, "all")
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.ToolParallelOverrides["write_file"] != "serial_only" {
|
||||||
|
t.Fatalf(
|
||||||
|
"tool_parallel_overrides.write_file = %q, want %q",
|
||||||
|
cfg.Orchestration.ToolParallelOverrides["write_file"], "serial_only",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if cfg.Orchestration.ToolParallelOverrides["exec"] != "parallel_read_only" {
|
||||||
|
t.Fatalf(
|
||||||
|
"tool_parallel_overrides.exec = %q, want %q",
|
||||||
|
cfg.Orchestration.ToolParallelOverrides["exec"], "parallel_read_only",
|
||||||
|
)
|
||||||
|
}
|
||||||
if !cfg.Audit.Enabled {
|
if !cfg.Audit.Enabled {
|
||||||
t.Fatal("audit.enabled should be true")
|
t.Fatal("audit.enabled should be true")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -354,10 +354,14 @@ func DefaultConfig() *Config {
|
||||||
Orchestration: OrchestrationConfig{
|
Orchestration: OrchestrationConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
MaxSpawnDepth: 3,
|
MaxSpawnDepth: 3,
|
||||||
MaxParallelWorkers: 4,
|
MaxParallelWorkers: 8,
|
||||||
MaxTasksPerAgent: 20,
|
MaxTasksPerAgent: 20,
|
||||||
DefaultTaskTimeoutSeconds: 180,
|
DefaultTaskTimeoutSeconds: 180,
|
||||||
RetryLimitPerTask: 2,
|
RetryLimitPerTask: 2,
|
||||||
|
ToolCallsParallelEnabled: true,
|
||||||
|
MaxToolCallConcurrency: 8,
|
||||||
|
ParallelToolsMode: "read_only_only",
|
||||||
|
ToolParallelOverrides: map[string]string{},
|
||||||
},
|
},
|
||||||
Audit: AuditConfig{
|
Audit: AuditConfig{
|
||||||
Enabled: false,
|
Enabled: false,
|
||||||
|
|
|
||||||
|
|
@ -69,6 +69,40 @@ type AsyncTool interface {
|
||||||
SetCallback(cb AsyncCallback)
|
SetCallback(cb AsyncCallback)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ToolParallelPolicy declares whether a tool is safe to run concurrently
|
||||||
|
// within a single LLM tool-call batch.
|
||||||
|
type ToolParallelPolicy string
|
||||||
|
|
||||||
|
const (
|
||||||
|
// ToolParallelSerialOnly is the safe default for tools with side effects.
|
||||||
|
ToolParallelSerialOnly ToolParallelPolicy = "serial_only"
|
||||||
|
// ToolParallelReadOnly marks tools that are safe to run in parallel.
|
||||||
|
ToolParallelReadOnly ToolParallelPolicy = "parallel_read_only"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
// ParallelToolsModeReadOnlyOnly allows parallel execution only for tools
|
||||||
|
// explicitly marked as ToolParallelReadOnly.
|
||||||
|
ParallelToolsModeReadOnlyOnly = "read_only_only"
|
||||||
|
// ParallelToolsModeAll allows all tools to run in parallel.
|
||||||
|
ParallelToolsModeAll = "all"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ParallelPolicyProvider is an optional interface that tools can implement
|
||||||
|
// to opt into parallel batch execution.
|
||||||
|
type ParallelPolicyProvider interface {
|
||||||
|
ParallelPolicy() ToolParallelPolicy
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConcurrentSafeTool is an optional interface for tool instances that can
|
||||||
|
// safely handle concurrent ExecuteWithContext calls on the same singleton object.
|
||||||
|
//
|
||||||
|
// Tools that rely on mutable per-call instance state (for example through
|
||||||
|
// SetContext or SetCallback) should return false.
|
||||||
|
type ConcurrentSafeTool interface {
|
||||||
|
SupportsConcurrentExecution() bool
|
||||||
|
}
|
||||||
|
|
||||||
func ToolToSchema(tool Tool) map[string]any {
|
func ToolToSchema(tool Tool) map[string]any {
|
||||||
return map[string]any{
|
return map[string]any{
|
||||||
"type": "function",
|
"type": "function",
|
||||||
|
|
|
||||||
|
|
@ -99,6 +99,10 @@ func (t *ReadFileTool) Name() string {
|
||||||
return "read_file"
|
return "read_file"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *ReadFileTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *ReadFileTool) Description() string {
|
func (t *ReadFileTool) Description() string {
|
||||||
return "Read the contents of a file"
|
return "Read the contents of a file"
|
||||||
}
|
}
|
||||||
|
|
@ -204,6 +208,10 @@ func (t *ListDirTool) Name() string {
|
||||||
return "list_dir"
|
return "list_dir"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *ListDirTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *ListDirTool) Description() string {
|
func (t *ListDirTool) Description() string {
|
||||||
return "List files and directories in a path"
|
return "List files and directories in a path"
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -35,6 +36,65 @@ func (r *ToolRegistry) Get(name string) (Tool, bool) {
|
||||||
return tool, ok
|
return tool, ok
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ParallelPolicy returns the configured parallel policy for a tool.
|
||||||
|
// Tools that do not implement ParallelPolicyProvider default to serial-only.
|
||||||
|
func (r *ToolRegistry) ParallelPolicy(name string) ToolParallelPolicy {
|
||||||
|
tool, ok := r.Get(name)
|
||||||
|
if !ok || tool == nil {
|
||||||
|
return ToolParallelSerialOnly
|
||||||
|
}
|
||||||
|
provider, ok := tool.(ParallelPolicyProvider)
|
||||||
|
if !ok {
|
||||||
|
return ToolParallelSerialOnly
|
||||||
|
}
|
||||||
|
policy := provider.ParallelPolicy()
|
||||||
|
if policy == "" {
|
||||||
|
return ToolParallelSerialOnly
|
||||||
|
}
|
||||||
|
return policy
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsParallelInstanceSafe reports whether one shared tool instance can be used
|
||||||
|
// concurrently across multiple tool calls.
|
||||||
|
func (r *ToolRegistry) IsParallelInstanceSafe(name string) bool {
|
||||||
|
tool, ok := r.Get(name)
|
||||||
|
if !ok || tool == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if safeTool, ok := tool.(ConcurrentSafeTool); ok {
|
||||||
|
return safeTool.SupportsConcurrentExecution()
|
||||||
|
}
|
||||||
|
// Conservative default: tools with mutable per-call hooks are not safe
|
||||||
|
// unless they explicitly opt in via ConcurrentSafeTool.
|
||||||
|
if _, ok := tool.(ContextualTool); ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if _, ok := tool.(AsyncTool); ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// CanRunToolCallInParallel reports whether a tool call may run in parallel
|
||||||
|
// under the given mode.
|
||||||
|
//
|
||||||
|
// Supported modes:
|
||||||
|
// - "read_only_only" (default): only tools marked parallel_read_only are parallelized.
|
||||||
|
// - "all": every tool is eligible.
|
||||||
|
func (r *ToolRegistry) CanRunToolCallInParallel(name, mode string) bool {
|
||||||
|
if !r.IsParallelInstanceSafe(name) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
switch strings.ToLower(strings.TrimSpace(mode)) {
|
||||||
|
case ParallelToolsModeAll:
|
||||||
|
return true
|
||||||
|
case "", ParallelToolsModeReadOnlyOnly:
|
||||||
|
return r.ParallelPolicy(name) == ToolParallelReadOnly
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string]any) *ToolResult {
|
func (r *ToolRegistry) Execute(ctx context.Context, name string, args map[string]any) *ToolResult {
|
||||||
return r.ExecuteWithContext(ctx, name, args, "", "", "", nil)
|
return r.ExecuteWithContext(ctx, name, args, "", "", "", nil)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -45,6 +45,24 @@ func (m *mockAsyncRegistryTool) SetCallback(cb AsyncCallback) {
|
||||||
m.cb = cb
|
m.cb = cb
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type mockConcurrentSafeTool struct {
|
||||||
|
mockRegistryTool
|
||||||
|
concurrentSafe bool
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockConcurrentSafeTool) SupportsConcurrentExecution() bool {
|
||||||
|
return m.concurrentSafe
|
||||||
|
}
|
||||||
|
|
||||||
|
type mockParallelRegistryTool struct {
|
||||||
|
mockRegistryTool
|
||||||
|
policy ToolParallelPolicy
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *mockParallelRegistryTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return m.policy
|
||||||
|
}
|
||||||
|
|
||||||
// --- helpers ---
|
// --- helpers ---
|
||||||
|
|
||||||
func newMockTool(name, desc string) *mockRegistryTool {
|
func newMockTool(name, desc string) *mockRegistryTool {
|
||||||
|
|
@ -348,3 +366,71 @@ func TestToolRegistry_ConcurrentAccess(t *testing.T) {
|
||||||
t.Error("expected tools to be registered after concurrent access")
|
t.Error("expected tools to be registered after concurrent access")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_ParallelPolicy_DefaultSerial(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
r.Register(newMockTool("default_serial", ""))
|
||||||
|
|
||||||
|
policy := r.ParallelPolicy("default_serial")
|
||||||
|
if policy != ToolParallelSerialOnly {
|
||||||
|
t.Fatalf("ParallelPolicy = %q, want %q", policy, ToolParallelSerialOnly)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_ParallelPolicy_ReadOnly(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
r.Register(&mockParallelRegistryTool{
|
||||||
|
mockRegistryTool: *newMockTool("read_only", ""),
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
})
|
||||||
|
|
||||||
|
policy := r.ParallelPolicy("read_only")
|
||||||
|
if policy != ToolParallelReadOnly {
|
||||||
|
t.Fatalf("ParallelPolicy = %q, want %q", policy, ToolParallelReadOnly)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_CanRunToolCallInParallel(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
r.Register(&mockParallelRegistryTool{
|
||||||
|
mockRegistryTool: *newMockTool("read_only", ""),
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
})
|
||||||
|
r.Register(newMockTool("serial", ""))
|
||||||
|
|
||||||
|
if !r.CanRunToolCallInParallel("read_only", ParallelToolsModeReadOnlyOnly) {
|
||||||
|
t.Fatal("read_only tool should be parallel-eligible in read_only_only mode")
|
||||||
|
}
|
||||||
|
if r.CanRunToolCallInParallel("serial", ParallelToolsModeReadOnlyOnly) {
|
||||||
|
t.Fatal("serial tool should not be parallel-eligible in read_only_only mode")
|
||||||
|
}
|
||||||
|
if !r.CanRunToolCallInParallel("serial", ParallelToolsModeAll) {
|
||||||
|
t.Fatal("serial tool should be parallel-eligible in all mode")
|
||||||
|
}
|
||||||
|
if r.CanRunToolCallInParallel("serial", "unknown_mode") {
|
||||||
|
t.Fatal("unknown mode should disable parallel eligibility")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_CanRunToolCallInParallel_ContextualToolIsNotSafeByDefault(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
r.Register(&mockCtxTool{
|
||||||
|
mockRegistryTool: *newMockTool("ctx_tool", "contextual"),
|
||||||
|
})
|
||||||
|
|
||||||
|
if r.CanRunToolCallInParallel("ctx_tool", ParallelToolsModeAll) {
|
||||||
|
t.Fatal("contextual tool should not be parallel-eligible by default")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolRegistry_CanRunToolCallInParallel_ConcurrentSafeOverride(t *testing.T) {
|
||||||
|
r := NewToolRegistry()
|
||||||
|
r.Register(&mockConcurrentSafeTool{
|
||||||
|
mockRegistryTool: *newMockTool("safe_tool", "safe"),
|
||||||
|
concurrentSafe: true,
|
||||||
|
})
|
||||||
|
|
||||||
|
if !r.CanRunToolCallInParallel("safe_tool", ParallelToolsModeAll) {
|
||||||
|
t.Fatal("concurrent-safe tool should be parallel-eligible in all mode")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -33,6 +33,10 @@ func (t *SessionsListTool) Name() string {
|
||||||
return "sessions_list"
|
return "sessions_list"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *SessionsListTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *SessionsListTool) Description() string {
|
func (t *SessionsListTool) Description() string {
|
||||||
return "List known conversation sessions with metadata. Useful for debugging, navigation, and context inspection."
|
return "List known conversation sessions with metadata. Useful for debugging, navigation, and context inspection."
|
||||||
}
|
}
|
||||||
|
|
@ -169,6 +173,10 @@ func (t *SessionsHistoryTool) Name() string {
|
||||||
return "sessions_history"
|
return "sessions_history"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *SessionsHistoryTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *SessionsHistoryTool) Description() string {
|
func (t *SessionsHistoryTool) Description() string {
|
||||||
return "Get full or partial message history for one session key."
|
return "Get full or partial message history for one session key."
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -28,6 +28,10 @@ func (t *FindSkillsTool) Name() string {
|
||||||
return "find_skills"
|
return "find_skills"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *FindSkillsTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *FindSkillsTool) Description() string {
|
func (t *FindSkillsTool) Description() string {
|
||||||
return "Search for installable skills from skill registries. Returns skill slugs, descriptions, versions, and relevance scores. Use this to discover skills before installing them with install_skill."
|
return "Search for installable skills from skill registries. Returns skill slugs, descriptions, versions, and relevance scores. Use this to discover skills before installing them with install_skill."
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -73,6 +73,11 @@ type SubagentManager struct {
|
||||||
maxTasks int
|
maxTasks int
|
||||||
maxDepth int
|
maxDepth int
|
||||||
taskCancels map[string]context.CancelFunc
|
taskCancels map[string]context.CancelFunc
|
||||||
|
|
||||||
|
toolCallsParallelEnabled bool
|
||||||
|
maxToolCallConcurrency int
|
||||||
|
parallelToolsMode string
|
||||||
|
toolPolicyOverrides map[string]string
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewSubagentManager(
|
func NewSubagentManager(
|
||||||
|
|
@ -145,6 +150,33 @@ func (sm *SubagentManager) SetLimits(maxConcurrent, maxTasks, maxDepth int) {
|
||||||
sm.maxDepth = maxDepth
|
sm.maxDepth = maxDepth
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetToolCallParallelism configures in-batch parallel tool execution for
|
||||||
|
// subagent tool loops.
|
||||||
|
func (sm *SubagentManager) SetToolCallParallelism(
|
||||||
|
enabled bool,
|
||||||
|
maxConcurrency int,
|
||||||
|
mode string,
|
||||||
|
toolPolicyOverrides map[string]string,
|
||||||
|
) {
|
||||||
|
sm.mu.Lock()
|
||||||
|
defer sm.mu.Unlock()
|
||||||
|
sm.toolCallsParallelEnabled = enabled
|
||||||
|
sm.maxToolCallConcurrency = maxConcurrency
|
||||||
|
sm.parallelToolsMode = mode
|
||||||
|
sm.toolPolicyOverrides = clonePolicyOverrides(toolPolicyOverrides)
|
||||||
|
}
|
||||||
|
|
||||||
|
func clonePolicyOverrides(src map[string]string) map[string]string {
|
||||||
|
if len(src) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
dst := make(map[string]string, len(src))
|
||||||
|
for k, v := range src {
|
||||||
|
dst[k] = v
|
||||||
|
}
|
||||||
|
return dst
|
||||||
|
}
|
||||||
|
|
||||||
func (sm *SubagentManager) Spawn(
|
func (sm *SubagentManager) Spawn(
|
||||||
ctx context.Context,
|
ctx context.Context,
|
||||||
task, label, agentID, originChannel, originChatID string,
|
task, label, agentID, originChannel, originChatID string,
|
||||||
|
|
@ -298,6 +330,10 @@ After completing the task, provide a clear summary of what was done.`
|
||||||
defaultTools := sm.tools
|
defaultTools := sm.tools
|
||||||
defaultProvider := sm.provider
|
defaultProvider := sm.provider
|
||||||
defaultModel := sm.defaultModel
|
defaultModel := sm.defaultModel
|
||||||
|
toolCallsParallelEnabled := sm.toolCallsParallelEnabled
|
||||||
|
maxToolCallConcurrency := sm.maxToolCallConcurrency
|
||||||
|
parallelToolsMode := sm.parallelToolsMode
|
||||||
|
toolPolicyOverrides := clonePolicyOverrides(sm.toolPolicyOverrides)
|
||||||
sm.mu.RUnlock()
|
sm.mu.RUnlock()
|
||||||
|
|
||||||
execution := SubagentExecutionConfig{
|
execution := SubagentExecutionConfig{
|
||||||
|
|
@ -364,6 +400,10 @@ After completing the task, provide a clear summary of what was done.`
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
LLMOptions: llmOptions,
|
LLMOptions: llmOptions,
|
||||||
SenderID: fmt.Sprintf("subagent:%s", task.ID),
|
SenderID: fmt.Sprintf("subagent:%s", task.ID),
|
||||||
|
ToolCallsParallelEnabled: toolCallsParallelEnabled,
|
||||||
|
MaxToolCallConcurrency: maxToolCallConcurrency,
|
||||||
|
ParallelToolsMode: parallelToolsMode,
|
||||||
|
ToolPolicyOverrides: toolPolicyOverrides,
|
||||||
}, messages, task.OriginChannel, task.OriginChatID)
|
}, messages, task.OriginChannel, task.OriginChatID)
|
||||||
|
|
||||||
var result *ToolResult
|
var result *ToolResult
|
||||||
|
|
@ -617,6 +657,10 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
hasMaxTokens := sm.hasMaxTokens
|
hasMaxTokens := sm.hasMaxTokens
|
||||||
hasTemperature := sm.hasTemperature
|
hasTemperature := sm.hasTemperature
|
||||||
resolver := sm.resolver
|
resolver := sm.resolver
|
||||||
|
toolCallsParallelEnabled := sm.toolCallsParallelEnabled
|
||||||
|
maxToolCallConcurrency := sm.maxToolCallConcurrency
|
||||||
|
parallelToolsMode := sm.parallelToolsMode
|
||||||
|
toolPolicyOverrides := clonePolicyOverrides(sm.toolPolicyOverrides)
|
||||||
execution := SubagentExecutionConfig{
|
execution := SubagentExecutionConfig{
|
||||||
Provider: sm.provider,
|
Provider: sm.provider,
|
||||||
Model: sm.defaultModel,
|
Model: sm.defaultModel,
|
||||||
|
|
@ -660,6 +704,10 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
Tools: execution.Tools,
|
Tools: execution.Tools,
|
||||||
MaxIterations: maxIter,
|
MaxIterations: maxIter,
|
||||||
LLMOptions: llmOptions,
|
LLMOptions: llmOptions,
|
||||||
|
ToolCallsParallelEnabled: toolCallsParallelEnabled,
|
||||||
|
MaxToolCallConcurrency: maxToolCallConcurrency,
|
||||||
|
ParallelToolsMode: parallelToolsMode,
|
||||||
|
ToolPolicyOverrides: toolPolicyOverrides,
|
||||||
}, messages, t.originChannel, t.originChatID)
|
}, messages, t.originChannel, t.originChatID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
return ErrorResult(fmt.Sprintf("Subagent execution failed: %v", err)).WithError(err)
|
||||||
|
|
|
||||||
304
pkg/tools/toolcall_executor.go
Normal file
304
pkg/tools/toolcall_executor.go
Normal file
|
|
@ -0,0 +1,304 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"math"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
// ToolCallParallelConfig configures in-batch parallel execution for tool calls.
|
||||||
|
type ToolCallParallelConfig struct {
|
||||||
|
Enabled bool
|
||||||
|
MaxConcurrency int
|
||||||
|
Mode string
|
||||||
|
// ToolPolicyOverrides allows per-tool policy overrides.
|
||||||
|
// Values: "serial_only" or "parallel_read_only".
|
||||||
|
ToolPolicyOverrides map[string]string
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolCallExecutionOptions controls how tool calls are executed.
|
||||||
|
type ToolCallExecutionOptions struct {
|
||||||
|
Channel string
|
||||||
|
ChatID string
|
||||||
|
SenderID string
|
||||||
|
|
||||||
|
Iteration int
|
||||||
|
LogScope string
|
||||||
|
|
||||||
|
Parallel ToolCallParallelConfig
|
||||||
|
|
||||||
|
// AsyncCallbackForCall creates a callback for async-capable tools.
|
||||||
|
// It may be nil when async callbacks are not needed.
|
||||||
|
AsyncCallbackForCall func(call providers.ToolCall) AsyncCallback
|
||||||
|
}
|
||||||
|
|
||||||
|
// ToolCallExecution captures one tool call execution result.
|
||||||
|
type ToolCallExecution struct {
|
||||||
|
ToolCall providers.ToolCall
|
||||||
|
Result *ToolResult
|
||||||
|
DurationMS int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// ExecuteToolCalls executes tool calls with optional bounded parallelism while
|
||||||
|
// preserving output order exactly as provided in the input slice.
|
||||||
|
func ExecuteToolCalls(
|
||||||
|
ctx context.Context,
|
||||||
|
registry *ToolRegistry,
|
||||||
|
toolCalls []providers.ToolCall,
|
||||||
|
opts ToolCallExecutionOptions,
|
||||||
|
) []ToolCallExecution {
|
||||||
|
if len(toolCalls) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
batchStart := time.Now()
|
||||||
|
|
||||||
|
scope := opts.LogScope
|
||||||
|
if scope == "" {
|
||||||
|
scope = "tool"
|
||||||
|
}
|
||||||
|
|
||||||
|
results := make([]ToolCallExecution, len(toolCalls))
|
||||||
|
parallelCount := 0
|
||||||
|
serialCount := 0
|
||||||
|
mode := normalizeParallelMode(opts.Parallel.Mode)
|
||||||
|
|
||||||
|
shouldParallelize := func(tc providers.ToolCall) bool {
|
||||||
|
if registry == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !opts.Parallel.Enabled {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if opts.Parallel.MaxConcurrency == 1 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if !registry.IsParallelInstanceSafe(tc.Name) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
if override, ok := getOverridePolicy(tc.Name, opts.Parallel.ToolPolicyOverrides); ok {
|
||||||
|
return override == ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
switch mode {
|
||||||
|
case ParallelToolsModeAll:
|
||||||
|
return true
|
||||||
|
case ParallelToolsModeReadOnlyOnly:
|
||||||
|
return registry.CanRunToolCallInParallel(tc.Name, ParallelToolsModeReadOnlyOnly)
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
runOne := func(idx int) {
|
||||||
|
tc := toolCalls[idx]
|
||||||
|
argsJSON, _ := json.Marshal(tc.Arguments)
|
||||||
|
argsPreview := utils.Truncate(string(argsJSON), 200)
|
||||||
|
logger.InfoCF(scope, fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
||||||
|
map[string]any{
|
||||||
|
"tool": tc.Name,
|
||||||
|
"iteration": opts.Iteration,
|
||||||
|
})
|
||||||
|
|
||||||
|
var asyncCallback AsyncCallback
|
||||||
|
if opts.AsyncCallbackForCall != nil {
|
||||||
|
asyncCallback = opts.AsyncCallbackForCall(tc)
|
||||||
|
}
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
var toolResult *ToolResult
|
||||||
|
if registry != nil {
|
||||||
|
toolResult = registry.ExecuteWithContext(
|
||||||
|
ctx,
|
||||||
|
tc.Name,
|
||||||
|
tc.Arguments,
|
||||||
|
opts.Channel,
|
||||||
|
opts.ChatID,
|
||||||
|
opts.SenderID,
|
||||||
|
asyncCallback,
|
||||||
|
)
|
||||||
|
} else {
|
||||||
|
toolResult = ErrorResult("No tools available")
|
||||||
|
}
|
||||||
|
|
||||||
|
if toolResult == nil {
|
||||||
|
toolResult = ErrorResult(fmt.Sprintf("tool %q returned nil result", tc.Name)).
|
||||||
|
WithError(fmt.Errorf("tool %q returned nil result", tc.Name))
|
||||||
|
}
|
||||||
|
|
||||||
|
results[idx] = ToolCallExecution{
|
||||||
|
ToolCall: tc,
|
||||||
|
Result: toolResult,
|
||||||
|
DurationMS: time.Since(start).Milliseconds(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
runParallelBatch := func(batch []int) {
|
||||||
|
if len(batch) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
maxConc := opts.Parallel.MaxConcurrency
|
||||||
|
if maxConc <= 0 || maxConc > len(batch) {
|
||||||
|
maxConc = len(batch)
|
||||||
|
}
|
||||||
|
if maxConc <= 1 {
|
||||||
|
for _, idx := range batch {
|
||||||
|
runOne(idx)
|
||||||
|
}
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
logger.DebugCF(scope, "Executing parallel tool batch", map[string]any{
|
||||||
|
"iteration": opts.Iteration,
|
||||||
|
"batch_size": len(batch),
|
||||||
|
"max_parallel": maxConc,
|
||||||
|
"parallel_mode": mode,
|
||||||
|
})
|
||||||
|
|
||||||
|
jobs := make(chan int)
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
for i := 0; i < maxConc; i++ {
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
for idx := range jobs {
|
||||||
|
runOne(idx)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, idx := range batch {
|
||||||
|
jobs <- idx
|
||||||
|
}
|
||||||
|
close(jobs)
|
||||||
|
wg.Wait()
|
||||||
|
}
|
||||||
|
|
||||||
|
parallelBatch := make([]int, 0, len(toolCalls))
|
||||||
|
flushParallelBatch := func() {
|
||||||
|
if len(parallelBatch) == 0 {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
runParallelBatch(parallelBatch)
|
||||||
|
parallelBatch = parallelBatch[:0]
|
||||||
|
}
|
||||||
|
|
||||||
|
for i, tc := range toolCalls {
|
||||||
|
if shouldParallelize(tc) {
|
||||||
|
parallelCount++
|
||||||
|
parallelBatch = append(parallelBatch, i)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
serialCount++
|
||||||
|
flushParallelBatch()
|
||||||
|
runOne(i)
|
||||||
|
}
|
||||||
|
flushParallelBatch()
|
||||||
|
|
||||||
|
errorCount := 0
|
||||||
|
durations := make([]int64, 0, len(results))
|
||||||
|
for _, executed := range results {
|
||||||
|
if executed.Result != nil && executed.Result.IsError {
|
||||||
|
errorCount++
|
||||||
|
}
|
||||||
|
durations = append(durations, executed.DurationMS)
|
||||||
|
}
|
||||||
|
p50, p95, avg, max := summarizeDurations(durations)
|
||||||
|
|
||||||
|
logger.InfoCF(scope, "Tool call batch summary", map[string]any{
|
||||||
|
"iteration": opts.Iteration,
|
||||||
|
"tool_parallel_enabled": opts.Parallel.Enabled,
|
||||||
|
"max_tool_concurrency": opts.Parallel.MaxConcurrency,
|
||||||
|
"parallel_tools_mode": mode,
|
||||||
|
"parallel_candidate_count": parallelCount,
|
||||||
|
"serial_count": serialCount,
|
||||||
|
"total": len(toolCalls),
|
||||||
|
"error_count": errorCount,
|
||||||
|
"batch_duration_ms": time.Since(batchStart).Milliseconds(),
|
||||||
|
"tool_call_duration_p50_ms": p50,
|
||||||
|
"tool_call_duration_p95_ms": p95,
|
||||||
|
"tool_call_duration_avg_ms": avg,
|
||||||
|
"tool_call_duration_max_ms": max,
|
||||||
|
})
|
||||||
|
|
||||||
|
return results
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeParallelMode(mode string) string {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(mode)) {
|
||||||
|
case "", ParallelToolsModeReadOnlyOnly:
|
||||||
|
return ParallelToolsModeReadOnlyOnly
|
||||||
|
case ParallelToolsModeAll:
|
||||||
|
return ParallelToolsModeAll
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func getOverridePolicy(toolName string, overrides map[string]string) (ToolParallelPolicy, bool) {
|
||||||
|
if len(overrides) == 0 {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
raw, ok := overrides[toolName]
|
||||||
|
if !ok {
|
||||||
|
raw, ok = overrides[strings.ToLower(strings.TrimSpace(toolName))]
|
||||||
|
}
|
||||||
|
if !ok {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
switch strings.ToLower(strings.TrimSpace(raw)) {
|
||||||
|
case string(ToolParallelSerialOnly):
|
||||||
|
return ToolParallelSerialOnly, true
|
||||||
|
case string(ToolParallelReadOnly):
|
||||||
|
return ToolParallelReadOnly, true
|
||||||
|
default:
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func summarizeDurations(durations []int64) (p50, p95, avg, max int64) {
|
||||||
|
if len(durations) == 0 {
|
||||||
|
return 0, 0, 0, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
sorted := append([]int64(nil), durations...)
|
||||||
|
sort.Slice(sorted, func(i, j int) bool { return sorted[i] < sorted[j] })
|
||||||
|
total := int64(0)
|
||||||
|
for _, d := range sorted {
|
||||||
|
total += d
|
||||||
|
}
|
||||||
|
avg = total / int64(len(sorted))
|
||||||
|
max = sorted[len(sorted)-1]
|
||||||
|
p50 = percentileInt64(sorted, 0.50)
|
||||||
|
p95 = percentileInt64(sorted, 0.95)
|
||||||
|
return p50, p95, avg, max
|
||||||
|
}
|
||||||
|
|
||||||
|
func percentileInt64(sorted []int64, p float64) int64 {
|
||||||
|
if len(sorted) == 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
if p <= 0 {
|
||||||
|
return sorted[0]
|
||||||
|
}
|
||||||
|
if p >= 1 {
|
||||||
|
return sorted[len(sorted)-1]
|
||||||
|
}
|
||||||
|
// Nearest-rank percentile: rank = ceil(p*n), index = rank-1.
|
||||||
|
idx := int(math.Ceil(p*float64(len(sorted)))) - 1
|
||||||
|
if idx < 0 {
|
||||||
|
idx = 0
|
||||||
|
}
|
||||||
|
if idx >= len(sorted) {
|
||||||
|
idx = len(sorted) - 1
|
||||||
|
}
|
||||||
|
return sorted[idx]
|
||||||
|
}
|
||||||
399
pkg/tools/toolcall_executor_test.go
Normal file
399
pkg/tools/toolcall_executor_test.go
Normal file
|
|
@ -0,0 +1,399 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"sync/atomic"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type executorMockTool struct {
|
||||||
|
name string
|
||||||
|
policy ToolParallelPolicy
|
||||||
|
|
||||||
|
delay time.Duration
|
||||||
|
|
||||||
|
result *ToolResult
|
||||||
|
errMsg string
|
||||||
|
|
||||||
|
running *atomic.Int32
|
||||||
|
maxRunning *atomic.Int32
|
||||||
|
onExecute func()
|
||||||
|
onComplete func()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *executorMockTool) Name() string {
|
||||||
|
return t.name
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *executorMockTool) Description() string {
|
||||||
|
return "executor mock tool"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *executorMockTool) Parameters() map[string]any {
|
||||||
|
return map[string]any{
|
||||||
|
"type": "object",
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *executorMockTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return t.policy
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *executorMockTool) Execute(_ context.Context, _ map[string]any) *ToolResult {
|
||||||
|
if t.onExecute != nil {
|
||||||
|
t.onExecute()
|
||||||
|
}
|
||||||
|
|
||||||
|
if t.running != nil && t.maxRunning != nil {
|
||||||
|
cur := t.running.Add(1)
|
||||||
|
for {
|
||||||
|
prev := t.maxRunning.Load()
|
||||||
|
if cur <= prev || t.maxRunning.CompareAndSwap(prev, cur) {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
defer t.running.Add(-1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if t.delay > 0 {
|
||||||
|
time.Sleep(t.delay)
|
||||||
|
}
|
||||||
|
if t.onComplete != nil {
|
||||||
|
t.onComplete()
|
||||||
|
}
|
||||||
|
|
||||||
|
if t.errMsg != "" {
|
||||||
|
return ErrorResult(t.errMsg).WithError(fmt.Errorf("%s", t.errMsg))
|
||||||
|
}
|
||||||
|
if t.result != nil {
|
||||||
|
return t.result
|
||||||
|
}
|
||||||
|
return SilentResult("ok")
|
||||||
|
}
|
||||||
|
|
||||||
|
type executorContextualTool struct {
|
||||||
|
executorMockTool
|
||||||
|
channel string
|
||||||
|
chatID string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (t *executorContextualTool) SetContext(channel, chatID string) {
|
||||||
|
t.channel = channel
|
||||||
|
t.chatID = chatID
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_PreservesOrderWithParallelBatch(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "slow",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 60 * time.Millisecond,
|
||||||
|
result: SilentResult("slow-result"),
|
||||||
|
})
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "fast",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 5 * time.Millisecond,
|
||||||
|
result: SilentResult("fast-result"),
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "slow", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "fast", Arguments: map[string]any{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 8,
|
||||||
|
Mode: ParallelToolsModeReadOnlyOnly,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Fatalf("len(results) = %d, want 2", len(results))
|
||||||
|
}
|
||||||
|
if results[0].ToolCall.ID != "tc-1" || results[0].Result.ForLLM != "slow-result" {
|
||||||
|
t.Fatalf("results[0] = %+v, want tc-1/slow-result", results[0])
|
||||||
|
}
|
||||||
|
if results[1].ToolCall.ID != "tc-2" || results[1].Result.ForLLM != "fast-result" {
|
||||||
|
t.Fatalf("results[1] = %+v, want tc-2/fast-result", results[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_RespectsConcurrencyLimit(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
var running atomic.Int32
|
||||||
|
var maxRunning atomic.Int32
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "io",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 25 * time.Millisecond,
|
||||||
|
result: SilentResult("ok"),
|
||||||
|
running: &running,
|
||||||
|
maxRunning: &maxRunning,
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := make([]providers.ToolCall, 0, 20)
|
||||||
|
for i := 0; i < 20; i++ {
|
||||||
|
calls = append(calls, providers.ToolCall{
|
||||||
|
ID: fmt.Sprintf("tc-%d", i),
|
||||||
|
Name: "io",
|
||||||
|
Arguments: map[string]any{"index": i},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 3,
|
||||||
|
Mode: ParallelToolsModeReadOnlyOnly,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != len(calls) {
|
||||||
|
t.Fatalf("len(results) = %d, want %d", len(results), len(calls))
|
||||||
|
}
|
||||||
|
if maxRunning.Load() > 3 {
|
||||||
|
t.Fatalf("maxRunning = %d, want <= 3", maxRunning.Load())
|
||||||
|
}
|
||||||
|
if maxRunning.Load() < 2 {
|
||||||
|
t.Fatalf("maxRunning = %d, want >= 2 to confirm parallel execution", maxRunning.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_SerialBoundaryBeforeParallelBatch(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
writeDone := make(chan struct{})
|
||||||
|
var readStartedBeforeWriteDone atomic.Bool
|
||||||
|
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "write_file",
|
||||||
|
policy: ToolParallelSerialOnly,
|
||||||
|
delay: 60 * time.Millisecond,
|
||||||
|
result: SilentResult("write-ok"),
|
||||||
|
onComplete: func() { close(writeDone) },
|
||||||
|
})
|
||||||
|
|
||||||
|
readTool := &executorMockTool{
|
||||||
|
name: "read_file",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 20 * time.Millisecond,
|
||||||
|
result: SilentResult("read-ok"),
|
||||||
|
onExecute: func() {
|
||||||
|
select {
|
||||||
|
case <-writeDone:
|
||||||
|
default:
|
||||||
|
readStartedBeforeWriteDone.Store(true)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
registry.Register(readTool)
|
||||||
|
|
||||||
|
calls := []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "write_file", Arguments: map[string]any{"path": "x"}},
|
||||||
|
{ID: "tc-2", Name: "read_file", Arguments: map[string]any{"path": "x"}},
|
||||||
|
{ID: "tc-3", Name: "read_file", Arguments: map[string]any{"path": "y"}},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 8,
|
||||||
|
Mode: ParallelToolsModeReadOnlyOnly,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != 3 {
|
||||||
|
t.Fatalf("len(results) = %d, want 3", len(results))
|
||||||
|
}
|
||||||
|
if readStartedBeforeWriteDone.Load() {
|
||||||
|
t.Fatal("read tool started before preceding serial write tool finished")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_CollectsFailuresWithoutShortCircuit(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "ok",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 20 * time.Millisecond,
|
||||||
|
result: SilentResult("ok-result"),
|
||||||
|
})
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "fail",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 5 * time.Millisecond,
|
||||||
|
errMsg: "boom",
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "ok", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "fail", Arguments: map[string]any{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 8,
|
||||||
|
Mode: ParallelToolsModeReadOnlyOnly,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Fatalf("len(results) = %d, want 2", len(results))
|
||||||
|
}
|
||||||
|
if results[0].Result.IsError {
|
||||||
|
t.Fatalf("results[0].Result.IsError = true, want false")
|
||||||
|
}
|
||||||
|
if !results[1].Result.IsError {
|
||||||
|
t.Fatalf("results[1].Result.IsError = false, want true")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_OverrideForcesParallelInReadOnlyMode(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
var running atomic.Int32
|
||||||
|
var maxRunning atomic.Int32
|
||||||
|
|
||||||
|
// Default policy is serial_only for tools without ParallelPolicyProvider.
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "custom_tool",
|
||||||
|
delay: 20 * time.Millisecond,
|
||||||
|
result: SilentResult("ok"),
|
||||||
|
running: &running,
|
||||||
|
maxRunning: &maxRunning,
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "custom_tool", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "custom_tool", Arguments: map[string]any{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 8,
|
||||||
|
Mode: ParallelToolsModeReadOnlyOnly,
|
||||||
|
ToolPolicyOverrides: map[string]string{
|
||||||
|
"custom_tool": string(ToolParallelReadOnly),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Fatalf("len(results) = %d, want 2", len(results))
|
||||||
|
}
|
||||||
|
if maxRunning.Load() < 2 {
|
||||||
|
t.Fatalf("maxRunning = %d, want >= 2 when override forces parallel", maxRunning.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_OverrideForcesSerialInAllMode(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
var running atomic.Int32
|
||||||
|
var maxRunning atomic.Int32
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "any_tool",
|
||||||
|
delay: 20 * time.Millisecond,
|
||||||
|
result: SilentResult("ok"),
|
||||||
|
running: &running,
|
||||||
|
maxRunning: &maxRunning,
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "any_tool", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "any_tool", Arguments: map[string]any{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 8,
|
||||||
|
Mode: ParallelToolsModeAll,
|
||||||
|
ToolPolicyOverrides: map[string]string{
|
||||||
|
"any_tool": string(ToolParallelSerialOnly),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Fatalf("len(results) = %d, want 2", len(results))
|
||||||
|
}
|
||||||
|
if maxRunning.Load() > 1 {
|
||||||
|
t.Fatalf("maxRunning = %d, want <= 1 when override forces serial", maxRunning.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExecuteToolCalls_OverrideCannotBypassInstanceSafety(t *testing.T) {
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
var running atomic.Int32
|
||||||
|
var maxRunning atomic.Int32
|
||||||
|
registry.Register(&executorContextualTool{
|
||||||
|
executorMockTool: executorMockTool{
|
||||||
|
name: "ctx_tool",
|
||||||
|
delay: 20 * time.Millisecond,
|
||||||
|
result: SilentResult("ok"),
|
||||||
|
running: &running,
|
||||||
|
maxRunning: &maxRunning,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
calls := []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "ctx_tool", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "ctx_tool", Arguments: map[string]any{}},
|
||||||
|
}
|
||||||
|
|
||||||
|
results := ExecuteToolCalls(context.Background(), registry, calls, ToolCallExecutionOptions{
|
||||||
|
Channel: "telegram",
|
||||||
|
ChatID: "chat-1",
|
||||||
|
Iteration: 1,
|
||||||
|
LogScope: "test",
|
||||||
|
Parallel: ToolCallParallelConfig{
|
||||||
|
Enabled: true,
|
||||||
|
MaxConcurrency: 8,
|
||||||
|
Mode: ParallelToolsModeAll,
|
||||||
|
ToolPolicyOverrides: map[string]string{
|
||||||
|
"ctx_tool": string(ToolParallelReadOnly),
|
||||||
|
},
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
if len(results) != 2 {
|
||||||
|
t.Fatalf("len(results) = %d, want 2", len(results))
|
||||||
|
}
|
||||||
|
if maxRunning.Load() > 1 {
|
||||||
|
t.Fatalf("maxRunning = %d, want <= 1 because instance safety should block parallel", maxRunning.Load())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPercentileInt64_NearestRank(t *testing.T) {
|
||||||
|
values := []int64{10, 100}
|
||||||
|
if got := percentileInt64(values, 0.95); got != 100 {
|
||||||
|
t.Fatalf("p95 with n=2 = %d, want 100 (nearest-rank)", got)
|
||||||
|
}
|
||||||
|
if got := percentileInt64(values, 0.50); got != 10 {
|
||||||
|
t.Fatalf("p50 with n=2 = %d, want 10 (nearest-rank)", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
values3 := []int64{10, 20, 30}
|
||||||
|
if got := percentileInt64(values3, 0.95); got != 30 {
|
||||||
|
t.Fatalf("p95 with n=3 = %d, want 30 (nearest-rank)", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -10,11 +10,9 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/logger"
|
"github.com/sipeed/picoclaw/pkg/logger"
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// ToolLoopConfig configures the tool execution loop.
|
// ToolLoopConfig configures the tool execution loop.
|
||||||
|
|
@ -25,6 +23,11 @@ type ToolLoopConfig struct {
|
||||||
MaxIterations int
|
MaxIterations int
|
||||||
LLMOptions map[string]any
|
LLMOptions map[string]any
|
||||||
SenderID string
|
SenderID string
|
||||||
|
|
||||||
|
ToolCallsParallelEnabled bool
|
||||||
|
MaxToolCallConcurrency int
|
||||||
|
ParallelToolsMode string
|
||||||
|
ToolPolicyOverrides map[string]string
|
||||||
}
|
}
|
||||||
|
|
||||||
// ToolLoopResult contains the result of running the tool loop.
|
// ToolLoopResult contains the result of running the tool loop.
|
||||||
|
|
@ -136,32 +139,23 @@ func RunToolLoop(
|
||||||
}
|
}
|
||||||
messages = append(messages, assistantMsg)
|
messages = append(messages, assistantMsg)
|
||||||
|
|
||||||
// 7. Execute tool calls
|
toolExecutions := ExecuteToolCalls(ctx, config.Tools, normalizedToolCalls, ToolCallExecutionOptions{
|
||||||
for _, tc := range normalizedToolCalls {
|
Channel: channel,
|
||||||
start := time.Now()
|
ChatID: chatID,
|
||||||
argsJSON, _ := json.Marshal(tc.Arguments)
|
SenderID: config.SenderID,
|
||||||
argsPreview := utils.Truncate(string(argsJSON), 200)
|
Iteration: iteration,
|
||||||
logger.InfoCF("toolloop", fmt.Sprintf("Tool call: %s(%s)", tc.Name, argsPreview),
|
LogScope: "toolloop",
|
||||||
map[string]any{
|
Parallel: ToolCallParallelConfig{
|
||||||
"tool": tc.Name,
|
Enabled: config.ToolCallsParallelEnabled,
|
||||||
"iteration": iteration,
|
MaxConcurrency: config.MaxToolCallConcurrency,
|
||||||
|
Mode: config.ParallelToolsMode,
|
||||||
|
ToolPolicyOverrides: config.ToolPolicyOverrides,
|
||||||
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
// Execute tool (no async callback for subagents - they run independently)
|
for _, executed := range toolExecutions {
|
||||||
var toolResult *ToolResult
|
toolResult := executed.Result
|
||||||
if config.Tools != nil {
|
tc := executed.ToolCall
|
||||||
toolResult = config.Tools.ExecuteWithContext(
|
|
||||||
ctx,
|
|
||||||
tc.Name,
|
|
||||||
tc.Arguments,
|
|
||||||
channel,
|
|
||||||
chatID,
|
|
||||||
config.SenderID,
|
|
||||||
nil,
|
|
||||||
)
|
|
||||||
} else {
|
|
||||||
toolResult = ErrorResult("No tools available")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Determine content for LLM
|
// Determine content for LLM
|
||||||
contentForLLM := toolResult.ForLLM
|
contentForLLM := toolResult.ForLLM
|
||||||
|
|
@ -175,7 +169,7 @@ func RunToolLoop(
|
||||||
Arguments: tc.Arguments,
|
Arguments: tc.Arguments,
|
||||||
Result: contentForLLM,
|
Result: contentForLLM,
|
||||||
IsError: toolResult.IsError,
|
IsError: toolResult.IsError,
|
||||||
DurationMS: time.Since(start).Milliseconds(),
|
DurationMS: executed.DurationMS,
|
||||||
ToolCallID: tc.ID,
|
ToolCallID: tc.ID,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
|
||||||
114
pkg/tools/toolloop_parallel_test.go
Normal file
114
pkg/tools/toolloop_parallel_test.go
Normal file
|
|
@ -0,0 +1,114 @@
|
||||||
|
package tools
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
)
|
||||||
|
|
||||||
|
type scriptedToolLoopProvider struct {
|
||||||
|
callCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *scriptedToolLoopProvider) Chat(
|
||||||
|
_ context.Context,
|
||||||
|
messages []providers.Message,
|
||||||
|
_ []providers.ToolDefinition,
|
||||||
|
_ string,
|
||||||
|
_ map[string]any,
|
||||||
|
) (*providers.LLMResponse, error) {
|
||||||
|
p.callCount++
|
||||||
|
if p.callCount == 1 {
|
||||||
|
return &providers.LLMResponse{
|
||||||
|
ToolCalls: []providers.ToolCall{
|
||||||
|
{ID: "tc-1", Name: "slow", Arguments: map[string]any{}},
|
||||||
|
{ID: "tc-2", Name: "fast", Arguments: map[string]any{}},
|
||||||
|
},
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if p.callCount == 2 {
|
||||||
|
toolMessages := make([]providers.Message, 0, 2)
|
||||||
|
for _, m := range messages {
|
||||||
|
if m.Role == "tool" {
|
||||||
|
toolMessages = append(toolMessages, m)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(toolMessages) != 2 {
|
||||||
|
return nil, fmt.Errorf("tool message count = %d, want 2", len(toolMessages))
|
||||||
|
}
|
||||||
|
if toolMessages[0].ToolCallID != "tc-1" || toolMessages[0].Content != "slow-result" {
|
||||||
|
return nil, fmt.Errorf("first tool message = %+v, want tc-1/slow-result", toolMessages[0])
|
||||||
|
}
|
||||||
|
if toolMessages[1].ToolCallID != "tc-2" || toolMessages[1].Content != "fast-result" {
|
||||||
|
return nil, fmt.Errorf("second tool message = %+v, want tc-2/fast-result", toolMessages[1])
|
||||||
|
}
|
||||||
|
return &providers.LLMResponse{Content: "final-answer"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
return &providers.LLMResponse{Content: "unexpected-extra-call"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *scriptedToolLoopProvider) GetDefaultModel() string {
|
||||||
|
return "toolloop-scripted"
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunToolLoop_ParallelToolCallsPreserveOrder(t *testing.T) {
|
||||||
|
provider := &scriptedToolLoopProvider{}
|
||||||
|
registry := NewToolRegistry()
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "slow",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 50 * time.Millisecond,
|
||||||
|
result: SilentResult("slow-result"),
|
||||||
|
})
|
||||||
|
registry.Register(&executorMockTool{
|
||||||
|
name: "fast",
|
||||||
|
policy: ToolParallelReadOnly,
|
||||||
|
delay: 5 * time.Millisecond,
|
||||||
|
result: SilentResult("fast-result"),
|
||||||
|
})
|
||||||
|
|
||||||
|
result, err := RunToolLoop(
|
||||||
|
context.Background(),
|
||||||
|
ToolLoopConfig{
|
||||||
|
Provider: provider,
|
||||||
|
Model: "mock-model",
|
||||||
|
Tools: registry,
|
||||||
|
MaxIterations: 4,
|
||||||
|
ToolCallsParallelEnabled: true,
|
||||||
|
MaxToolCallConcurrency: 8,
|
||||||
|
ParallelToolsMode: ParallelToolsModeReadOnlyOnly,
|
||||||
|
},
|
||||||
|
[]providers.Message{
|
||||||
|
{Role: "system", Content: "you are a test assistant"},
|
||||||
|
{Role: "user", Content: "run tools"},
|
||||||
|
},
|
||||||
|
"cli",
|
||||||
|
"direct",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("RunToolLoop() error = %v", err)
|
||||||
|
}
|
||||||
|
if result == nil {
|
||||||
|
t.Fatal("RunToolLoop() returned nil result")
|
||||||
|
}
|
||||||
|
if result.Content != "final-answer" {
|
||||||
|
t.Fatalf("result.Content = %q, want %q", result.Content, "final-answer")
|
||||||
|
}
|
||||||
|
if provider.callCount != 2 {
|
||||||
|
t.Fatalf("provider.callCount = %d, want 2", provider.callCount)
|
||||||
|
}
|
||||||
|
if len(result.Trace) != 2 {
|
||||||
|
t.Fatalf("len(result.Trace) = %d, want 2", len(result.Trace))
|
||||||
|
}
|
||||||
|
if result.Trace[0].ToolCallID != "tc-1" || result.Trace[0].Result != "slow-result" {
|
||||||
|
t.Fatalf("trace[0] = %+v, want tc-1/slow-result", result.Trace[0])
|
||||||
|
}
|
||||||
|
if result.Trace[1].ToolCallID != "tc-2" || result.Trace[1].Result != "fast-result" {
|
||||||
|
t.Fatalf("trace[1] = %+v, want tc-2/fast-result", result.Trace[1])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -538,6 +538,10 @@ func (t *WebSearchTool) Name() string {
|
||||||
return "web_search"
|
return "web_search"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *WebSearchTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *WebSearchTool) Description() string {
|
func (t *WebSearchTool) Description() string {
|
||||||
return "Search the web for current information. Returns titles, URLs, and snippets from search results."
|
return "Search the web for current information. Returns titles, URLs, and snippets from search results."
|
||||||
}
|
}
|
||||||
|
|
@ -613,6 +617,10 @@ func (t *WebFetchTool) Name() string {
|
||||||
return "web_fetch"
|
return "web_fetch"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (t *WebFetchTool) ParallelPolicy() ToolParallelPolicy {
|
||||||
|
return ToolParallelReadOnly
|
||||||
|
}
|
||||||
|
|
||||||
func (t *WebFetchTool) Description() string {
|
func (t *WebFetchTool) Description() string {
|
||||||
return "Fetch a URL and extract readable content (HTML to text). Use this to get weather info, news, articles, or any web content."
|
return "Fetch a URL and extract readable content (HTML to text). Use this to get weather info, news, articles, or any web content."
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue