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:
parent
fe51853262
commit
bac9a1c84b
5 changed files with 241 additions and 28 deletions
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"`
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue