diff --git a/pkg/agent/agent_init.go b/pkg/agent/agent_init.go index 335fd8537..7a9f5866b 100644 --- a/pkg/agent/agent_init.go +++ b/pkg/agent/agent_init.go @@ -320,7 +320,15 @@ func registerSharedTools( subagentTool := tools.NewSubagentTool(subagentManager) subagentTool.SetSpawner(NewSubTurnSpawner(al)) agent.Tools.Register(subagentTool) + } + // Register the synchronous batched subagent tool + multiSubagentTool := tools.NewMultiSubagentTool(subagentManager) + multiSubagentTool.SetSpawner(NewSubTurnSpawner(al)) + multiSubagentTool.SetAllowlistChecker(func(targetAgentID string) bool { + return registry.CanSpawnSubagent(currentAgentID, targetAgentID) + }) + agent.Tools.Register(multiSubagentTool) if spawnStatusEnabled { agent.Tools.Register(tools.NewSpawnStatusTool(subagentManager)) } diff --git a/pkg/tools/multi_subagent.go b/pkg/tools/multi_subagent.go new file mode 100644 index 000000000..3df7d197f --- /dev/null +++ b/pkg/tools/multi_subagent.go @@ -0,0 +1,215 @@ +package tools + +import ( + "context" + "fmt" + "strings" + "sync" +) + +// MultiSubagentTool executes several subagent tasks synchronously and collects +// their results. Each call may optionally include an allowed target agent ID for +// routing metadata/reporting, while execution otherwise follows the same direct +// SubTurnSpawner path as the existing synchronous subagent tool. +type MultiSubagentTool struct { + spawner SubTurnSpawner + defaultModel string + maxTokens int + temperature float64 + allowlistCheck func(targetAgentID string) bool +} + +func NewMultiSubagentTool(manager *SubagentManager) *MultiSubagentTool { + if manager == nil { + return &MultiSubagentTool{} + } + return &MultiSubagentTool{ + defaultModel: manager.defaultModel, + maxTokens: manager.maxTokens, + temperature: manager.temperature, + } +} + +func (t *MultiSubagentTool) SetSpawner(spawner SubTurnSpawner) { + t.spawner = spawner +} + +func (t *MultiSubagentTool) SetAllowlistChecker(check func(targetAgentID string) bool) { + t.allowlistCheck = check +} + +func (t *MultiSubagentTool) Name() string { + return "multi_subagent" +} + +func (t *MultiSubagentTool) Description() string { + return "Execute several subagent tasks synchronously and return grouped results. Use this when parallel delegation materially helps the current turn. Optional agent_id values are validated against the allowlist and reported per call." +} + +func (t *MultiSubagentTool) Parameters() map[string]any { + return map[string]any{ + "type": "object", + "properties": map[string]any{ + "calls": map[string]any{ + "type": "array", + "description": "The subagent calls to execute in parallel", + "items": map[string]any{ + "type": "object", + "properties": map[string]any{ + "task": map[string]any{ + "type": "string", + "description": "The task for the subagent to complete", + }, + "label": map[string]any{ + "type": "string", + "description": "Optional short label for the task", + }, + "agent_id": map[string]any{ + "type": "string", + "description": "Optional target agent ID metadata for this call, such as mini or code", + }, + }, + "required": []string{"task"}, + }, + }, + }, + "required": []string{"calls"}, + } +} + +type multiSubagentCall struct { + Task string + Label string + AgentID string +} + +type multiSubagentCallResult struct { + Index int + Label string + AgentID string + Result *ToolResult + Err error +} + +func (t *MultiSubagentTool) Execute(ctx context.Context, args map[string]any) *ToolResult { + if t.spawner == nil { + return ErrorResult("Subagent manager not configured").WithError(fmt.Errorf("spawner not set")) + } + + rawCalls, ok := args["calls"].([]any) + if !ok || len(rawCalls) == 0 { + return ErrorResult("calls is required and must be a non-empty array") + } + + calls := make([]multiSubagentCall, 0, len(rawCalls)) + for i, raw := range rawCalls { + callMap, ok := raw.(map[string]any) + if !ok { + return ErrorResult(fmt.Sprintf("calls[%d] must be an object", i)) + } + task, ok := callMap["task"].(string) + if !ok || strings.TrimSpace(task) == "" { + return ErrorResult(fmt.Sprintf("calls[%d].task is required and must be a non-empty string", i)) + } + label, _ := callMap["label"].(string) + agentID, _ := callMap["agent_id"].(string) + if agentID != "" && t.allowlistCheck != nil && !t.allowlistCheck(agentID) { + return ErrorResult(fmt.Sprintf("calls[%d] is not allowed to target agent '%s'", i, agentID)) + } + calls = append(calls, multiSubagentCall{Task: task, Label: label, AgentID: agentID}) + } + + results := make([]multiSubagentCallResult, len(calls)) + var wg sync.WaitGroup + for i, call := range calls { + wg.Add(1) + go func(i int, call multiSubagentCall) { + defer wg.Done() + + systemPrompt := fmt.Sprintf( + `You are a subagent. Complete the given task independently and provide a clear, concise result. + +Task: %s`, + call.Task, + ) + if call.Label != "" { + systemPrompt = fmt.Sprintf( + `You are a subagent labeled "%s". Complete the given task independently and provide a clear, concise result. + +Task: %s`, + call.Label, + call.Task, + ) + } + + res, err := t.spawner.SpawnSubTurn(ctx, SubTurnConfig{ + Model: t.defaultModel, + Tools: nil, + SystemPrompt: systemPrompt, + ActualSystemPrompt: buildMultiSubagentActualSystemPrompt(call.Label), + MaxTokens: t.maxTokens, + Temperature: t.temperature, + Async: false, + }) + results[i] = multiSubagentCallResult{ + Index: i, + Label: call.Label, + AgentID: call.AgentID, + Result: res, + Err: err, + } + }(i, call) + } + wg.Wait() + + var llmParts []string + var userParts []string + hadError := false + for _, item := range results { + label := item.Label + if label == "" { + label = fmt.Sprintf("call-%d", item.Index+1) + } + agentSegment := "default" + if item.AgentID != "" { + agentSegment = item.AgentID + } + + if item.Err != nil { + hadError = true + llmParts = append(llmParts, fmt.Sprintf("[%s | agent=%s] ERROR: %v", label, agentSegment, item.Err)) + userParts = append(userParts, fmt.Sprintf("%s: error: %v", label, item.Err)) + continue + } + if item.Result == nil { + hadError = true + llmParts = append(llmParts, fmt.Sprintf("[%s | agent=%s] ERROR: no result returned", label, agentSegment)) + userParts = append(userParts, fmt.Sprintf("%s: error: no result returned", label)) + continue + } + if item.Result.IsError { + hadError = true + } + llmParts = append(llmParts, fmt.Sprintf("[%s | agent=%s]\n%s", label, agentSegment, item.Result.ForLLM)) + userText := item.Result.ForUser + if strings.TrimSpace(userText) == "" { + userText = item.Result.ForLLM + } + userParts = append(userParts, fmt.Sprintf("%s: %s", label, userText)) + } + + return &ToolResult{ + ForLLM: "Parallel subagent results:\n" + strings.Join(llmParts, "\n\n"), + ForUser: strings.Join(userParts, "\n"), + Silent: false, + IsError: hadError, + Async: false, + } +} + +func buildMultiSubagentActualSystemPrompt(label string) string { + if strings.TrimSpace(label) == "" { + return "You are a subagent. Complete the given task independently and provide a clear, concise result." + } + return fmt.Sprintf("You are a subagent labeled %q. Complete the given task independently and provide a clear, concise result.", label) +} diff --git a/pkg/tools/multi_subagent_test.go b/pkg/tools/multi_subagent_test.go new file mode 100644 index 000000000..299e67f65 --- /dev/null +++ b/pkg/tools/multi_subagent_test.go @@ -0,0 +1,174 @@ +package tools + +import ( + "context" + "errors" + "strings" + "sync" + "testing" +) + +type recordingMultiSubagentSpawner struct { + mu sync.Mutex + cfgs []SubTurnConfig + byTask map[string]*ToolResult + errTask map[string]error +} + +func (s *recordingMultiSubagentSpawner) SpawnSubTurn(ctx context.Context, cfg SubTurnConfig) (*ToolResult, error) { + s.mu.Lock() + defer s.mu.Unlock() + s.cfgs = append(s.cfgs, cfg) + res := (*ToolResult)(nil) + var err error + if s.byTask != nil { + res = s.byTask[cfg.SystemPrompt] + } + if s.errTask != nil { + err = s.errTask[cfg.SystemPrompt] + } + return res, err +} + +func TestMultiSubagentTool_RequiresSpawner(t *testing.T) { + tool := NewMultiSubagentTool(nil) + res := tool.Execute(context.Background(), map[string]any{ + "calls": []any{map[string]any{"task": "hello"}}, + }) + if !res.IsError { + t.Fatal("expected error when spawner is missing") + } + if !strings.Contains(res.ForLLM, "Subagent manager not configured") { + t.Fatalf("unexpected error: %s", res.ForLLM) + } +} + +func TestMultiSubagentTool_ValidatesCalls(t *testing.T) { + tool := NewMultiSubagentTool(nil) + tool.SetSpawner(&recordingMultiSubagentSpawner{}) + + cases := []map[string]any{ + {}, + {"calls": []any{}}, + {"calls": []any{"bad"}}, + {"calls": []any{map[string]any{"task": " "}}}, + } + for _, args := range cases { + res := tool.Execute(context.Background(), args) + if !res.IsError { + t.Fatalf("expected validation error for args: %#v", args) + } + } +} + +func TestMultiSubagentTool_AllowlistRejectsTargetedCall(t *testing.T) { + tool := NewMultiSubagentTool(nil) + tool.SetSpawner(&recordingMultiSubagentSpawner{}) + tool.SetAllowlistChecker(func(targetAgentID string) bool { + return targetAgentID == "mini" + }) + + res := tool.Execute(context.Background(), map[string]any{ + "calls": []any{ + map[string]any{"task": "one", "agent_id": "code"}, + }, + }) + if !res.IsError { + t.Fatal("expected allowlist rejection") + } + if !strings.Contains(res.ForLLM, "not allowed to target agent 'code'") { + t.Fatalf("unexpected allowlist error: %s", res.ForLLM) + } +} + +func TestMultiSubagentTool_UsesDefaultModelAndGroupsResults(t *testing.T) { + promptAlpha := "You are a subagent labeled \"alpha\". Complete the given task independently and provide a clear, concise result.\n\nTask: task a" + promptBeta := "You are a subagent. Complete the given task independently and provide a clear, concise result.\n\nTask: task b" + spawner := &recordingMultiSubagentSpawner{ + byTask: map[string]*ToolResult{ + promptAlpha: {ForLLM: "alpha llm", ForUser: "alpha user"}, + promptBeta: {ForLLM: "beta llm", ForUser: ""}, + }, + } + tool := &MultiSubagentTool{ + spawner: spawner, + defaultModel: "main-model", + maxTokens: 123, + temperature: 0.4, + } + + res := tool.Execute(context.Background(), map[string]any{ + "calls": []any{ + map[string]any{"task": "task a", "label": "alpha", "agent_id": "mini"}, + map[string]any{"task": "task b"}, + }, + }) + if res.IsError { + t.Fatalf("unexpected error result: %s", res.ForLLM) + } + if len(spawner.cfgs) != 2 { + t.Fatalf("expected 2 subturns, got %d", len(spawner.cfgs)) + } + for i, cfg := range spawner.cfgs { + if cfg.Model != "main-model" { + t.Fatalf("cfg[%d].Model = %q, want %q", i, cfg.Model, "main-model") + } + if cfg.Async { + t.Fatalf("cfg[%d].Async = true, want false", i) + } + if cfg.MaxTokens != 123 { + t.Fatalf("cfg[%d].MaxTokens = %d, want 123", i, cfg.MaxTokens) + } + if cfg.Temperature != 0.4 { + t.Fatalf("cfg[%d].Temperature = %v, want 0.4", i, cfg.Temperature) + } + } + if !strings.Contains(res.ForLLM, "[alpha | agent=mini]") { + t.Fatalf("expected labeled grouped LLM output, got: %s", res.ForLLM) + } + if !strings.Contains(res.ForLLM, "[call-2 | agent=default]") { + t.Fatalf("expected default grouped LLM output, got: %s", res.ForLLM) + } + if !strings.Contains(res.ForUser, "alpha: alpha user") { + t.Fatalf("expected ForUser to prefer subresult user content, got: %s", res.ForUser) + } + if !strings.Contains(res.ForUser, "call-2: beta llm") { + t.Fatalf("expected ForUser fallback to llm content, got: %s", res.ForUser) + } +} + +func TestMultiSubagentTool_PropagatesPerCallErrors(t *testing.T) { + promptOne := "You are a subagent labeled \"one\". Complete the given task independently and provide a clear, concise result.\n\nTask: task 1" + promptTwo := "You are a subagent labeled \"two\". Complete the given task independently and provide a clear, concise result.\n\nTask: task 2" + promptThree := "You are a subagent labeled \"three\". Complete the given task independently and provide a clear, concise result.\n\nTask: task 3" + spawner := &recordingMultiSubagentSpawner{ + byTask: map[string]*ToolResult{ + promptOne: {ForLLM: "ok-one", ForUser: "ok-one"}, + promptThree: {ForLLM: "suberror", ForUser: "suberror", IsError: true}, + }, + errTask: map[string]error{ + promptTwo: errors.New("boom"), + }, + } + tool := &MultiSubagentTool{ + spawner: spawner, + defaultModel: "main-model", + } + + res := tool.Execute(context.Background(), map[string]any{ + "calls": []any{ + map[string]any{"task": "task 1", "label": "one"}, + map[string]any{"task": "task 2", "label": "two"}, + map[string]any{"task": "task 3", "label": "three"}, + }, + }) + if !res.IsError { + t.Fatal("expected aggregated error state") + } + if !strings.Contains(res.ForLLM, "[two | agent=default] ERROR: boom") { + t.Fatalf("expected explicit per-call error in llm output, got: %s", res.ForLLM) + } + if !strings.Contains(res.ForLLM, "[three | agent=default]\nsuberror") { + t.Fatalf("expected subresult error content in llm output, got: %s", res.ForLLM) + } +}