From 429ce66a7d0cf6bbbd7ca53f1dd7e97813f109a8 Mon Sep 17 00:00:00 2001 From: Administrator <1280842908@qq.com> Date: Sat, 28 Feb 2026 14:16:56 +0800 Subject: [PATCH] feat(teams): implement heterogeneous agent model routing - Added 'model' property to 'team' and 'spawn_sub_agent' JSON schemas. - Modified 'buildWorkerConfig' to override 'baseConfig.Model' when a specific LLM model is requested by the coordinator. - Allows teams to dynamically mix and match specialized vision, coding, and logical models within the same execution loop. --- pkg/tools/spawn_sub_agent.go | 9 ++++++ pkg/tools/team.go | 56 ++++++++++++++++++++++++++---------- 2 files changed, 50 insertions(+), 15 deletions(-) diff --git a/pkg/tools/spawn_sub_agent.go b/pkg/tools/spawn_sub_agent.go index cc740f16f..7617f8a12 100644 --- a/pkg/tools/spawn_sub_agent.go +++ b/pkg/tools/spawn_sub_agent.go @@ -43,6 +43,10 @@ func (t *SpawnSubAgentTool) Parameters() map[string]any { "type": "string", "description": "The system prompt/role assignment for the sub-agent (e.g., 'You are an expert code reviewer').", }, + "model": map[string]any{ + "type": "string", + "description": "Optional specific LLM model ID to route this task to (e.g., 'gpt-4o' for vision, 'claude-3-5-sonnet' for logic). If omitted, inherits the parent's model.", + }, }, "required": []string{"task", "role"}, } @@ -83,6 +87,11 @@ func (t *SpawnSubAgentTool) Execute(ctx context.Context, args map[string]any) *T // 2. Base Configuration (Timeout & LLM constraints) config := t.manager.BuildBaseWorkerConfig(ctx) + // 2.1 Model Override (Heterogeneous Agents) + if modelParam, ok := args["model"].(string); ok && strings.TrimSpace(modelParam) != "" { + config.Model = strings.TrimSpace(modelParam) + } + // Note: For MVP, we pass the current ToolRegistry unmodified. // To enforce strict sandboxing later, we can construct a new ToolRegistry here based on args['allowed_tools']. diff --git a/pkg/tools/team.go b/pkg/tools/team.go index 8fb500b70..4290487b0 100644 --- a/pkg/tools/team.go +++ b/pkg/tools/team.go @@ -21,6 +21,7 @@ type TeamMember struct { ID string Role string Task string + Model string // Heterogeneous Agents: Optional specific model for this task DependsOn []string // List of member IDs this member depends on } @@ -71,6 +72,10 @@ func (t *TeamTool) Parameters() map[string]any { "type": "string", "description": "The specific task this member needs to accomplish.", }, + "model": map[string]any{ + "type": "string", + "description": "Optional specific LLM model ID to route this task to (e.g., 'gpt-4o' for vision, 'claude-3-5-sonnet' for logic). If omitted, inherits the parent's model.", + }, "depends_on": map[string]any{ "type": "array", "description": "List of 'id' strings this member depends on. Only applicable for 'dag' strategy.", @@ -132,6 +137,9 @@ func (t *TeamTool) Execute(ctx context.Context, args map[string]any) *ToolResult id = fmt.Sprintf("member_%d", i) } + modelStr, _ := mMap["model"].(string) + modelStr = strings.TrimSpace(modelStr) + var dependsOn []string if depRaw, dOk := mMap["depends_on"].([]any); dOk { for _, d := range depRaw { @@ -145,6 +153,7 @@ func (t *TeamTool) Execute(ctx context.Context, args map[string]any) *ToolResult ID: id, Role: role, Task: task, + Model: modelStr, DependsOn: dependsOn, }) } @@ -199,7 +208,19 @@ func upgradeRegistryForConcurrency(original *ToolRegistry) *ToolRegistry { return upgraded } -func (t *TeamTool) executeSequential(ctx context.Context, config ToolLoopConfig, members []TeamMember) *ToolResult { +// buildWorkerConfig creates a ToolLoopConfig for a specific team member, +// potentially overriding the model based on the member's definition. +func buildWorkerConfig(baseConfig ToolLoopConfig, registry *ToolRegistry, m TeamMember) ToolLoopConfig { + cfg := baseConfig + cfg.Tools = registry + // Heterogeneous Agents: Override model if this team member requested a specific one + if m.Model != "" { + cfg.Model = m.Model + } + return cfg +} + +func (t *TeamTool) executeSequential(ctx context.Context, baseConfig ToolLoopConfig, members []TeamMember) *ToolResult { var finalOutput strings.Builder finalOutput.WriteString("Team Execution Summary (Sequential):\n\n") @@ -217,7 +238,8 @@ func (t *TeamTool) executeSequential(ctx context.Context, config ToolLoopConfig, {Role: "user", Content: actualTask}, } - loopResult, err := RunToolLoop(ctx, config, messages, t.originChannel, t.originChatID) + workerConfig := buildWorkerConfig(baseConfig, baseConfig.Tools, m) + loopResult, err := RunToolLoop(ctx, workerConfig, messages, t.originChannel, t.originChatID) if err != nil { errStr := fmt.Sprintf("Phase %d (Role: %s) failed: %v", i+1, m.Role, err) finalOutput.WriteString(errStr + "\n") @@ -235,7 +257,7 @@ func (t *TeamTool) executeSequential(ctx context.Context, config ToolLoopConfig, } } -func (t *TeamTool) executeParallel(ctx context.Context, config ToolLoopConfig, members []TeamMember) *ToolResult { +func (t *TeamTool) executeParallel(ctx context.Context, baseConfig ToolLoopConfig, members []TeamMember) *ToolResult { var wg sync.WaitGroup type workResult struct { index int @@ -248,22 +270,23 @@ func (t *TeamTool) executeParallel(ctx context.Context, config ToolLoopConfig, m for i, m := range members { wg.Add(1) - go func(index int, role, task string) { + go func(index int, member TeamMember) { defer wg.Done() messages := []providers.Message{ - {Role: "system", Content: role}, - {Role: "user", Content: task}, + {Role: "system", Content: member.Role}, + {Role: "user", Content: member.Task}, } - loopResult, err := RunToolLoop(ctx, config, messages, t.originChannel, t.originChatID) + workerConfig := buildWorkerConfig(baseConfig, baseConfig.Tools, member) + loopResult, err := RunToolLoop(ctx, workerConfig, messages, t.originChannel, t.originChatID) if err != nil { - resultsChan <- workResult{index: index, role: role, err: err} + resultsChan <- workResult{index: index, role: member.Role, err: err} return } - resultsChan <- workResult{index: index, role: role, res: loopResult.Content} - }(i, m.Role, m.Task) + resultsChan <- workResult{index: index, role: member.Role, res: loopResult.Content} + }(i, m) } // Wait for all goroutines to finish @@ -299,7 +322,7 @@ func (t *TeamTool) executeParallel(ctx context.Context, config ToolLoopConfig, m } } -func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, config ToolLoopConfig, members []TeamMember) *ToolResult { +func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, baseConfig ToolLoopConfig, members []TeamMember) *ToolResult { if len(members) != 2 { return ErrorResult("The evaluator_optimizer strategy requires exactly two members: [0] Worker, [1] Evaluator.") } @@ -321,7 +344,8 @@ func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, config ToolLoo finalOutput.WriteString(fmt.Sprintf("## Attempt %d\n", attempt)) // 2. Trigger Worker (resumes from its exact previous state!) - workerResult, err := RunToolLoop(ctx, config, workerMessages, t.originChannel, t.originChatID) + workerConfig := buildWorkerConfig(baseConfig, baseConfig.Tools, worker) + workerResult, err := RunToolLoop(ctx, workerConfig, workerMessages, t.originChannel, t.originChatID) if err != nil { errStr := fmt.Sprintf("Worker failed on attempt %d: %v", attempt, err) finalOutput.WriteString(errStr + "\n") @@ -341,7 +365,8 @@ func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, config ToolLoo {Role: "user", Content: evalContext}, } - evalResult, err := RunToolLoop(ctx, config, evalMessages, t.originChannel, t.originChatID) + evalConfig := buildWorkerConfig(baseConfig, baseConfig.Tools, evaluator) + evalResult, err := RunToolLoop(ctx, evalConfig, evalMessages, t.originChannel, t.originChatID) if err != nil { errStr := fmt.Sprintf("Evaluator failed on attempt %d: %v", attempt, err) finalOutput.WriteString(errStr + "\n") @@ -376,7 +401,7 @@ func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, config ToolLoo } } -func (t *TeamTool) executeDAG(ctx context.Context, config ToolLoopConfig, members []TeamMember) *ToolResult { +func (t *TeamTool) executeDAG(ctx context.Context, baseConfig ToolLoopConfig, members []TeamMember) *ToolResult { // 1. Build and VALIDATE dependency graph memberMap := make(map[string]TeamMember) inDegree := make(map[string]int) @@ -489,7 +514,8 @@ func (t *TeamTool) executeDAG(ctx context.Context, config ToolLoopConfig, member {Role: "user", Content: actualTask}, } - loopResult, err := RunToolLoop(ctx, config, messages, t.originChannel, t.originChatID) + workerConfig := buildWorkerConfig(baseConfig, baseConfig.Tools, m) + loopResult, err := RunToolLoop(ctx, workerConfig, messages, t.originChannel, t.originChatID) if err != nil { masterErrMu.Lock()