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)
teamTool := tools.NewTeamTool(subagentManager)
teamTool := tools.NewTeamTool(subagentManager, cfg)
if cfg.Tools.IsToolEnabled("team") {
agent.Tools.Register(teamTool)
}

View file

@ -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"`

View file

@ -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 {

View file

@ -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")
}

View file

@ -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"