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
This commit is contained in:
Administrator 2026-03-12 18:28:30 +08:00
parent fe51853262
commit bac9a1c84b
5 changed files with 241 additions and 28 deletions

View file

@ -229,7 +229,7 @@ func registerSharedTools(
}) })
agent.Tools.Register(spawnTool) agent.Tools.Register(spawnTool)
teamTool := tools.NewTeamTool(subagentManager) teamTool := tools.NewTeamTool(subagentManager, cfg)
if cfg.Tools.IsToolEnabled("team") { if cfg.Tools.IsToolEnabled("team") {
agent.Tools.Register(teamTool) agent.Tools.Register(teamTool)
} }

View file

@ -75,21 +75,21 @@ func (f *FlexibleStringSlice) UnmarshalText(text []byte) error {
} }
type TeamModelConfig struct { type TeamModelConfig struct {
Name string `json:"name" yaml:"name"` Name string `json:"name"`
Tags []string `json:"tags,omitempty" yaml:"tags,omitempty"` Tags []string `json:"tags,omitempty"`
} }
type TeamToolsConfig struct { type TeamToolsConfig struct {
ToolConfig ` envPrefix:"PICOCLAW_TOOLS_TEAM_"` ToolConfig
MaxMembers int `json:"max_members" env:"PICOCLAW_TOOLS_TEAM_MAX_MEMBERS"` MaxMembers int `json:"max_members"`
MaxTeamTokens int `json:"max_team_tokens" env:"PICOCLAW_TOOLS_TEAM_MAX_TOKENS"` MaxTeamTokens int `json:"max_team_tokens"`
MaxEvaluatorLoops int `json:"max_evaluator_loops" env:"PICOCLAW_TOOLS_TEAM_MAX_EVALUATOR_LOOPS"` MaxEvaluatorLoops int `json:"max_evaluator_loops"`
MaxTimeoutMinutes int `json:"max_timeout_minutes" env:"PICOCLAW_TOOLS_TEAM_MAX_TIMEOUT_MINUTES"` MaxTimeoutMinutes int `json:"max_timeout_minutes"`
MaxContextRunes int `json:"max_context_runes" env:"PICOCLAW_TOOLS_TEAM_MAX_CONTEXT_RUNES"` MaxContextRunes int `json:"max_context_runes"`
DisableAutoReviewer bool `json:"disable_auto_reviewer" env:"PICOCLAW_TOOLS_TEAM_DISABLE_AUTO_REVIEWER"` DisableAutoReviewer bool `json:"disable_auto_reviewer"`
ReviewerModel string `json:"reviewer_model" env:"PICOCLAW_TOOLS_TEAM_REVIEWER_MODEL"` ReviewerModel string `json:"reviewer_model"`
AllowedStrategies []string `json:"allowed_strategies" env:"PICOCLAW_TOOLS_TEAM_ALLOWED_STRATEGIES"` AllowedStrategies []string `json:"allowed_strategies"`
AllowedModels []TeamModelConfig `json:"allowed_models" env:"-"` AllowedModels []TeamModelConfig `json:"allowed_models"`
} }
type Config struct { type Config struct {
@ -771,7 +771,6 @@ type ToolsConfig struct {
WriteFile ToolConfig `json:"write_file" envPrefix:"PICOCLAW_TOOLS_WRITE_FILE_"` WriteFile ToolConfig `json:"write_file" envPrefix:"PICOCLAW_TOOLS_WRITE_FILE_"`
} }
type SearchCacheConfig struct { type SearchCacheConfig struct {
MaxSize int `json:"max_size" env:"PICOCLAW_SKILLS_SEARCH_CACHE_MAX_SIZE"` MaxSize int `json:"max_size" env:"PICOCLAW_SKILLS_SEARCH_CACHE_MAX_SIZE"`
TTLSeconds int `json:"ttl_seconds" env:"PICOCLAW_SKILLS_SEARCH_CACHE_TTL_SECONDS"` TTLSeconds int `json:"ttl_seconds" env:"PICOCLAW_SKILLS_SEARCH_CACHE_TTL_SECONDS"`

View file

@ -8,12 +8,14 @@ import (
"sync/atomic" "sync/atomic"
"time" "time"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/providers"
) )
type TeamTool struct { type TeamTool struct {
manager *SubagentManager manager *SubagentManager
cfg *config.Config
originChannel string originChannel string
originChatID string originChatID string
} }
@ -27,9 +29,10 @@ type TeamMember struct {
Produces string // Auto-reviewer: declares artifact type ("code", "data", "document") 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{ return &TeamTool{
manager: manager, manager: manager,
cfg: cfg,
originChannel: "cli", originChannel: "cli",
originChatID: "direct", originChatID: "direct",
} }
@ -191,6 +194,26 @@ func (t *TeamTool) maybeRunAutoReviewer(
reviewerConfig.Model = teamConfig.ReviewerModel 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) loopResult, err := RunToolLoop(ctx, reviewerConfig, reviewerMessages, t.originChannel, t.originChatID)
if err != nil { if err != nil {
return fmt.Sprintf("[Auto-Reviewer] Failed to run: %v", err) 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) budget.Store(effectiveMaxTokens)
} }
var members []TeamMember var members []TeamMember
for i, mRaw := range membersRaw { for i, mRaw := range membersRaw {
mMap, ok := mRaw.(map[string]any) mMap, ok := mRaw.(map[string]any)
@ -392,16 +414,51 @@ func upgradeRegistryForConcurrency(original *ToolRegistry) *ToolRegistry {
// buildWorkerConfig creates a ToolLoopConfig for a specific team member, // buildWorkerConfig creates a ToolLoopConfig for a specific team member,
// potentially overriding the model based on the member's definition. // 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 := baseConfig
cfg.Tools = registry cfg.Tools = registry
// Heterogeneous Agents: Override model if this team member requested a specific one // Heterogeneous Agents: Override model if this team member requested a specific one
if m.Model != "" { 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) 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 return cfg, nil
} }
@ -423,7 +480,7 @@ func (t *TeamTool) executeSequential(ctx context.Context, baseConfig ToolLoopCon
{Role: "user", Content: actualTask}, {Role: "user", Content: actualTask},
} }
workerConfig, err := buildWorkerConfig(baseConfig, baseConfig.Tools, m, t.manager) workerConfig, err := t.buildWorkerConfig(baseConfig, baseConfig.Tools, m)
if err != nil { if err != nil {
errStr := fmt.Sprintf("Phase %d (Role: %s) configuration failed: %v", i+1, m.Role, err) errStr := fmt.Sprintf("Phase %d (Role: %s) configuration failed: %v", i+1, m.Role, err)
finalOutput.WriteString(errStr + "\n") finalOutput.WriteString(errStr + "\n")
@ -474,7 +531,7 @@ func (t *TeamTool) executeParallel(ctx context.Context, baseConfig ToolLoopConfi
{Role: "user", Content: member.Task}, {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 { if err != nil {
resultsChan <- workResult{index: index, role: member.Role, err: err} resultsChan <- workResult{index: index, role: member.Role, err: err}
return 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 { func (t *TeamTool) executeEvaluatorOptimizer(ctx context.Context, baseConfig ToolLoopConfig, members []TeamMember, contextLimit int) *ToolResult {
if len(members) != 2 { if len(members) != 2 {
return ErrorResult("The evaluator_optimizer strategy requires exactly two members: [0] Worker, [1] Evaluator.") 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. // 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 { if err != nil {
return ErrorResult(fmt.Sprintf("Worker configuration failed: %v", err)).WithError(err) 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 { if err != nil {
return ErrorResult(fmt.Sprintf("Evaluator configuration failed: %v", err)).WithError(err) 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}, {Role: "user", Content: actualTask},
} }
workerConfig, err := buildWorkerConfig(baseConfig, baseConfig.Tools, m, t.manager) workerConfig, err := t.buildWorkerConfig(baseConfig, baseConfig.Tools, m)
if err != nil { if err != nil {
masterErrMu.Lock() masterErrMu.Lock()
if masterErr == nil { if masterErr == nil {

View file

@ -1,8 +1,11 @@
package tools package tools
import ( import (
"context"
"testing" "testing"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/providers"
"github.com/stretchr/testify/assert" "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") 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")
}

View file

@ -10,8 +10,8 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"fmt" "fmt"
"sync/atomic"
"sync" "sync"
"sync/atomic"
"github.com/sipeed/picoclaw/pkg/logger" "github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/providers" "github.com/sipeed/picoclaw/pkg/providers"