From bac9a1c84b469f45fbb61d219c8dd0a5e2351f15 Mon Sep 17 00:00:00 2001 From: Administrator <1280842908@qq.com> Date: Thu, 12 Mar 2026 18:28:30 +0800 Subject: [PATCH] refactor(tools): refine team tool configuration and add unit tests - Update TeamTool to accept and use global configuration - Refactor buildWorkerConfig into a TeamTool method for better context handling - Improve error handling in DAG and sequential execution flows - Add comprehensive unit tests for TeamTool sequential execution - Clean up imports and formatting in toolloop.go and config.go --- pkg/agent/loop.go | 2 +- pkg/config/config.go | 29 ++++---- pkg/tools/team.go | 78 +++++++++++++++++--- pkg/tools/team_test.go | 158 +++++++++++++++++++++++++++++++++++++++++ pkg/tools/toolloop.go | 2 +- 5 files changed, 241 insertions(+), 28 deletions(-) diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index ce7f40800..01023e362 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -229,7 +229,7 @@ func registerSharedTools( }) agent.Tools.Register(spawnTool) - teamTool := tools.NewTeamTool(subagentManager) + teamTool := tools.NewTeamTool(subagentManager, cfg) if cfg.Tools.IsToolEnabled("team") { agent.Tools.Register(teamTool) } diff --git a/pkg/config/config.go b/pkg/config/config.go index 17c4f245a..b82b439e3 100644 --- a/pkg/config/config.go +++ b/pkg/config/config.go @@ -75,21 +75,21 @@ func (f *FlexibleStringSlice) UnmarshalText(text []byte) error { } type TeamModelConfig struct { - Name string `json:"name" yaml:"name"` - Tags []string `json:"tags,omitempty" yaml:"tags,omitempty"` + Name string `json:"name"` + Tags []string `json:"tags,omitempty"` } type TeamToolsConfig struct { - ToolConfig ` envPrefix:"PICOCLAW_TOOLS_TEAM_"` - MaxMembers int `json:"max_members" env:"PICOCLAW_TOOLS_TEAM_MAX_MEMBERS"` - MaxTeamTokens int `json:"max_team_tokens" env:"PICOCLAW_TOOLS_TEAM_MAX_TOKENS"` - MaxEvaluatorLoops int `json:"max_evaluator_loops" env:"PICOCLAW_TOOLS_TEAM_MAX_EVALUATOR_LOOPS"` - MaxTimeoutMinutes int `json:"max_timeout_minutes" env:"PICOCLAW_TOOLS_TEAM_MAX_TIMEOUT_MINUTES"` - MaxContextRunes int `json:"max_context_runes" env:"PICOCLAW_TOOLS_TEAM_MAX_CONTEXT_RUNES"` - DisableAutoReviewer bool `json:"disable_auto_reviewer" env:"PICOCLAW_TOOLS_TEAM_DISABLE_AUTO_REVIEWER"` - ReviewerModel string `json:"reviewer_model" env:"PICOCLAW_TOOLS_TEAM_REVIEWER_MODEL"` - AllowedStrategies []string `json:"allowed_strategies" env:"PICOCLAW_TOOLS_TEAM_ALLOWED_STRATEGIES"` - AllowedModels []TeamModelConfig `json:"allowed_models" env:"-"` + ToolConfig + MaxMembers int `json:"max_members"` + MaxTeamTokens int `json:"max_team_tokens"` + MaxEvaluatorLoops int `json:"max_evaluator_loops"` + MaxTimeoutMinutes int `json:"max_timeout_minutes"` + MaxContextRunes int `json:"max_context_runes"` + DisableAutoReviewer bool `json:"disable_auto_reviewer"` + ReviewerModel string `json:"reviewer_model"` + AllowedStrategies []string `json:"allowed_strategies"` + AllowedModels []TeamModelConfig `json:"allowed_models"` } type Config struct { @@ -607,8 +607,8 @@ type OpenAIProviderConfig struct { // Default protocol is "openai" if no prefix is specified. type ModelConfig struct { // Required fields - ModelName string `json:"model_name"` // User-facing alias for the model - Model string `json:"model"` // Protocol/model-identifier (e.g., "openai/gpt-4o", "anthropic/claude-sonnet-4.6") + ModelName string `json:"model_name"` // User-facing alias for the model + Model string `json:"model"` // Protocol/model-identifier (e.g., "openai/gpt-4o", "anthropic/claude-sonnet-4.6") // HTTP-based providers APIBase string `json:"api_base,omitempty"` // API endpoint URL @@ -771,7 +771,6 @@ type ToolsConfig struct { WriteFile ToolConfig `json:"write_file" envPrefix:"PICOCLAW_TOOLS_WRITE_FILE_"` } - type SearchCacheConfig struct { MaxSize int `json:"max_size" env:"PICOCLAW_SKILLS_SEARCH_CACHE_MAX_SIZE"` TTLSeconds int `json:"ttl_seconds" env:"PICOCLAW_SKILLS_SEARCH_CACHE_TTL_SECONDS"` diff --git a/pkg/tools/team.go b/pkg/tools/team.go index 0e0450749..8c2223d5e 100644 --- a/pkg/tools/team.go +++ b/pkg/tools/team.go @@ -8,12 +8,14 @@ import ( "sync/atomic" "time" + "github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/providers" ) type TeamTool struct { manager *SubagentManager + cfg *config.Config originChannel string originChatID string } @@ -27,9 +29,10 @@ type TeamMember struct { Produces string // Auto-reviewer: declares artifact type ("code", "data", "document") } -func NewTeamTool(manager *SubagentManager) *TeamTool { +func NewTeamTool(manager *SubagentManager, cfg *config.Config) *TeamTool { return &TeamTool{ manager: manager, + cfg: cfg, originChannel: "cli", originChatID: "direct", } @@ -191,6 +194,26 @@ func (t *TeamTool) maybeRunAutoReviewer( reviewerConfig.Model = teamConfig.ReviewerModel } + cnf, err := t.cfg.GetModelConfig(reviewerConfig.Model) + + if err == nil { + provider, model, err := providers.CreateProviderFromConfig(cnf) + + if err == nil { + reviewerConfig.Model = model + reviewerConfig.Provider = provider + } + } + + providerName := "unknown" + if reviewerConfig.Provider != nil { + providerName = reviewerConfig.Provider.GetDefaultModel() + } + + logger.InfoCF("team", fmt.Sprintf("reviewer use provider: [%s] and model: [%s]", providerName, reviewerConfig.Model), map[string]any{ + "model": teamConfig.ReviewerModel, + }) + loopResult, err := RunToolLoop(ctx, reviewerConfig, reviewerMessages, t.originChannel, t.originChatID) if err != nil { return fmt.Sprintf("[Auto-Reviewer] Failed to run: %v", err) @@ -268,7 +291,6 @@ func (t *TeamTool) Execute(ctx context.Context, args map[string]any) *ToolResult budget.Store(effectiveMaxTokens) } - var members []TeamMember for i, mRaw := range membersRaw { mMap, ok := mRaw.(map[string]any) @@ -392,16 +414,51 @@ func upgradeRegistryForConcurrency(original *ToolRegistry) *ToolRegistry { // 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, manager *SubagentManager) (ToolLoopConfig, error) { +func (t *TeamTool) buildWorkerConfig(baseConfig ToolLoopConfig, registry *ToolRegistry, m TeamMember) (ToolLoopConfig, error) { cfg := baseConfig cfg.Tools = registry // Heterogeneous Agents: Override model if this team member requested a specific one if m.Model != "" { - if !manager.IsModelAllowed(m.Model) { + if !t.manager.IsModelAllowed(m.Model) { return cfg, fmt.Errorf("requested model '%s' is not in the allowed fallback candidates list for this agent workspace", m.Model) } - cfg.Model = m.Model + // Resolve model name from model_list if it's an alias + //resolvedModel := m.Model + //if t.cfg != nil { + // for _, mc := range t.cfg.ModelList { + // if mc.ModelName == m.Model && mc.Model != "" { + // resolvedModel = mc.Model + // break + // } + // } + //} + + cnf, err := t.cfg.GetModelConfig(m.Model) + + if err != nil { + return cfg, err + } + + provider, model, err := providers.CreateProviderFromConfig(cnf) + + if err != nil { + return ToolLoopConfig{}, err + } + + cfg.Model = model + cfg.Provider = provider } + + providerName := "unknown" + if cfg.Provider != nil { + providerName = cfg.Provider.GetDefaultModel() + } + + logger.InfoCF("team", fmt.Sprintf("[%s] use provider: [%s] and model: [%s]", m.Role, providerName, cfg.Model), map[string]any{ + "member_index": m.ID, + "model": m.Model, + }) + return cfg, nil } @@ -423,7 +480,7 @@ func (t *TeamTool) executeSequential(ctx context.Context, baseConfig ToolLoopCon {Role: "user", Content: actualTask}, } - workerConfig, err := buildWorkerConfig(baseConfig, baseConfig.Tools, m, t.manager) + workerConfig, err := t.buildWorkerConfig(baseConfig, baseConfig.Tools, m) if err != nil { errStr := fmt.Sprintf("Phase %d (Role: %s) configuration failed: %v", i+1, m.Role, err) finalOutput.WriteString(errStr + "\n") @@ -474,7 +531,7 @@ func (t *TeamTool) executeParallel(ctx context.Context, baseConfig ToolLoopConfi {Role: "user", Content: member.Task}, } - workerConfig, err := buildWorkerConfig(baseConfig, baseConfig.Tools, member, t.manager) + workerConfig, err := t.buildWorkerConfig(baseConfig, baseConfig.Tools, member) if err != nil { resultsChan <- workResult{index: index, role: member.Role, err: err} return @@ -543,7 +600,6 @@ func (t *TeamTool) executeParallel(ctx context.Context, baseConfig ToolLoopConfi } } - func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, baseConfig ToolLoopConfig, members []TeamMember, contextLimit int) *ToolResult { if len(members) != 2 { return ErrorResult("The evaluator_optimizer strategy requires exactly two members: [0] Worker, [1] Evaluator.") @@ -572,11 +628,11 @@ func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, baseConfig Too } // Pre-compute both configs once — they don't change between loop iterations. - workerConfig, err := buildWorkerConfig(baseConfig, baseConfig.Tools, worker, t.manager) + workerConfig, err := t.buildWorkerConfig(baseConfig, baseConfig.Tools, worker) if err != nil { return ErrorResult(fmt.Sprintf("Worker configuration failed: %v", err)).WithError(err) } - evalConfig, err := buildWorkerConfig(baseConfig, NewToolRegistry(), evaluator, t.manager) + evalConfig, err := t.buildWorkerConfig(baseConfig, NewToolRegistry(), evaluator) if err != nil { return ErrorResult(fmt.Sprintf("Evaluator configuration failed: %v", err)).WithError(err) } @@ -767,7 +823,7 @@ func (t *TeamTool) executeDAG(ctx context.Context, cancel context.CancelFunc, ba {Role: "user", Content: actualTask}, } - workerConfig, err := buildWorkerConfig(baseConfig, baseConfig.Tools, m, t.manager) + workerConfig, err := t.buildWorkerConfig(baseConfig, baseConfig.Tools, m) if err != nil { masterErrMu.Lock() if masterErr == nil { diff --git a/pkg/tools/team_test.go b/pkg/tools/team_test.go index c49661249..986374b80 100644 --- a/pkg/tools/team_test.go +++ b/pkg/tools/team_test.go @@ -1,8 +1,11 @@ package tools import ( + "context" "testing" + "github.com/sipeed/picoclaw/pkg/config" + "github.com/sipeed/picoclaw/pkg/providers" "github.com/stretchr/testify/assert" ) @@ -60,3 +63,158 @@ func TestUpgradeRegistryForConcurrency(t *testing.T) { assert.False(t, isConcurrent, "Original registry components MUST REMAIN completely lock-free") } } + +func TestBuildWorkerConfig(t *testing.T) { + // 1. Setup global config with model aliases + cfg := &config.Config{ + ModelList: []config.ModelConfig{ + { + ModelName: "strong-model", + Model: "openai/gpt-4o", + APIKey: "sk-test", + APIBase: "https://api.openai.com/v1", + }, + { + ModelName: "fast-model", + Model: "anthropic/claude-3-haiku", + APIKey: "sk-test", + APIBase: "https://api.anthropic.com/v1", + }, + { + ModelName: "direct-id", + Model: "openai/gpt-3.5-turbo", + APIKey: "sk-test", + APIBase: "https://api.openai.com/v1", + }, + }, + } + + // 2. Setup SubagentManager (needed by TeamTool for IsModelAllowed check) + manager := NewSubagentManager(nil, "default-model", nil, "", config.TeamToolsConfig{ + AllowedModels: []config.TeamModelConfig{ + {Name: "fast-model", Tags: []string{"vision"}}, + {Name: "strong-model", Tags: []string{"coding"}}, + {Name: "direct-id", Tags: []string{"coding"}}, + }, + }, nil) + + // 3. Create TeamTool + tool := NewTeamTool(manager, cfg) + + baseConfig := ToolLoopConfig{ + Model: "base-model", + } + + tests := []struct { + name string + memberModel string + expectedModel string + expectError bool + }{ + { + name: "Resolve alias to actual ID", + memberModel: "strong-model", + expectedModel: "gpt-4o", + expectError: false, + }, + { + name: "Resolve another alias", + memberModel: "fast-model", + expectedModel: "claude-3-haiku", + expectError: false, + }, + { + name: "Resolve direct name if it matches an alias", + memberModel: "direct-id", + expectedModel: "gpt-3.5-turbo", + expectError: false, + }, + { + name: "Inherit base model if member model is empty", + memberModel: "", + expectedModel: "base-model", + expectError: false, + }, + { + name: "Error if model is not allowed", + memberModel: "forbidden-model", + expectedModel: "", + expectError: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + m := TeamMember{ + Model: tt.memberModel, + } + res, err := tool.buildWorkerConfig(baseConfig, nil, m) + + if tt.expectError { + assert.Error(t, err) + } else { + assert.NoError(t, err) + assert.Equal(t, tt.expectedModel, res.Model) + if tt.memberModel != "" { + assert.NotNil(t, res.Provider, "Provider should be set when model is specified") + } else { + assert.Nil(t, res.Provider, "Provider should be nil (inherited from baseConfig)") + } + } + }) + } +} + +type mockProvider struct { + responses []string + callCount int +} + +func (m *mockProvider) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, options map[string]any) (*providers.LLMResponse, error) { + if m.callCount >= len(m.responses) { + return &providers.LLMResponse{Content: "Default response"}, nil + } + resp := m.responses[m.callCount] + m.callCount++ + return &providers.LLMResponse{Content: resp}, nil +} + +func (m *mockProvider) GetDefaultModel() string { + return "mock-model" +} + +func TestExecuteSequential(t *testing.T) { + // 1. Setup mock provider to return specific outputs for each agent + mock := &mockProvider{ + responses: []string{ + "Result from Agent A", + "Derived result from Agent B", + }, + } + + // 2. Setup TeamTool + manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil) + tool := NewTeamTool(manager, &config.Config{}) + + baseConfig := ToolLoopConfig{ + Provider: mock, + Model: "mock-model", + MaxIterations: 1, + } + + members := []TeamMember{ + {ID: "worker-A", Role: "Researcher", Task: "Research topic X"}, + {ID: "worker-B", Role: "Writer", Task: "Write summary of researcher output"}, + } + + // 3. Run sequential execution + result := tool.executeSequential(context.Background(), baseConfig, members, 1000) + + // 4. Verify results + assert.False(t, result.IsError, "Should not return error") + assert.Contains(t, result.ForLLM, "Result from Agent A") + assert.Contains(t, result.ForLLM, "Derived result from Agent B") + assert.Equal(t, 2, mock.callCount, "Should have called mock provider exactly twice") +} + + diff --git a/pkg/tools/toolloop.go b/pkg/tools/toolloop.go index 5ec2b4a5d..673bb9158 100644 --- a/pkg/tools/toolloop.go +++ b/pkg/tools/toolloop.go @@ -10,8 +10,8 @@ import ( "context" "encoding/json" "fmt" - "sync/atomic" "sync" + "sync/atomic" "github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/providers"