feat(team): add DAG execution tests and improve mock provider for concurrency

This commit is contained in:
Administrator 2026-03-24 10:21:58 +08:00
parent 4a714a7534
commit 375e64bba5
2 changed files with 390 additions and 1 deletions

View file

@ -899,6 +899,7 @@ func (t *TeamTool) executeDAG(ctx context.Context, cancel context.CancelFunc, ba
// 3. Initialize queue with nodes having 0 in-degree // 3. Initialize queue with nodes having 0 in-degree
nodesToProcess := len(members) nodesToProcess := len(members)
completedNodes := 0
for id, deg := range inDegree { for id, deg := range inDegree {
if deg == 0 { if deg == 0 {
readyChan <- id readyChan <- id
@ -914,7 +915,8 @@ func (t *TeamTool) executeDAG(ctx context.Context, cancel context.CancelFunc, ba
var finalResultsMu sync.Mutex var finalResultsMu sync.Mutex
// 4. DAG Execution Loop // 4. DAG Execution Loop
for i := 0; i < nodesToProcess; i++ { // Continue until all nodes have completed
for completedNodes < nodesToProcess {
select { select {
case <-ctx.Done(): case <-ctx.Done():
return ErrorResult("DAG execution timed out or cancelled") return ErrorResult("DAG execution timed out or cancelled")
@ -987,6 +989,9 @@ func (t *TeamTool) executeDAG(ctx context.Context, cancel context.CancelFunc, ba
return ErrorResult(res.err.Error()) return ErrorResult(res.err.Error())
} }
// Mark this node as completed
completedNodes++
// Update dependents // Update dependents
for _, dependentID := range graph[res.id] { for _, dependentID := range graph[res.id] {
contextMu.Lock() contextMu.Lock()

View file

@ -2,6 +2,7 @@ package tools
import ( import (
"context" "context"
"sync"
"testing" "testing"
"github.com/sipeed/picoclaw/pkg/config" "github.com/sipeed/picoclaw/pkg/config"
@ -192,6 +193,34 @@ func (m *mockProvider) GetDefaultModel() string {
return "mock-model" return "mock-model"
} }
// mockProviderWithID returns responses based on the role in the system message
// This is needed for DAG tests where execution order is non-deterministic
type mockProviderWithID struct {
responses map[string]string
mu sync.Mutex
callCount int
}
func (m *mockProviderWithID) Chat(ctx context.Context, messages []providers.Message, tools []providers.ToolDefinition, model string, options map[string]any) (*providers.LLMResponse, error) {
m.mu.Lock()
defer m.mu.Unlock()
m.callCount++
// Extract role from system message to determine which response to return
for _, msg := range messages {
if msg.Role == "system" {
if resp, ok := m.responses[msg.Content]; ok {
return &providers.LLMResponse{Content: resp}, nil
}
}
}
return &providers.LLMResponse{Content: "Default response"}, nil
}
func (m *mockProviderWithID) GetDefaultModel() string {
return "mock-model"
}
func TestExecuteSequential(t *testing.T) { func TestExecuteSequential(t *testing.T) {
// 1. Setup mock provider to return specific outputs for each agent // 1. Setup mock provider to return specific outputs for each agent
mock := &mockProvider{ mock := &mockProvider{
@ -225,3 +254,358 @@ func TestExecuteSequential(t *testing.T) {
assert.Contains(t, result.ForLLM, "Derived result from Agent B") assert.Contains(t, result.ForLLM, "Derived result from Agent B")
assert.Equal(t, 2, mock.callCount, "Should have called mock provider exactly twice") assert.Equal(t, 2, mock.callCount, "Should have called mock provider exactly twice")
} }
func TestExecuteDAG(t *testing.T) {
t.Run("Simple DAG with dependencies", func(t *testing.T) {
// Setup: A -> C, B -> C (C depends on both A and B)
mock := &mockProviderWithID{
responses: map[string]string{
"Data Collector A": "Data from A",
"Data Collector B": "Data from B",
"Aggregator": "Combined result from C using A and B",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mock,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(), // Empty registry for test
}
members := []TeamMember{
{ID: "A", Role: "Data Collector A", Task: "Collect data A"},
{ID: "B", Role: "Data Collector B", Task: "Collect data B"},
{ID: "C", Role: "Aggregator", Task: "Combine data", DependsOn: []string{"A", "B"}},
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
result := tool.executeDAG(ctx, cancel, baseConfig, members, 1000)
assert.False(t, result.IsError, "Should not return error")
assert.Contains(t, result.ForLLM, "Data from A")
assert.Contains(t, result.ForLLM, "Data from B")
assert.Contains(t, result.ForLLM, "Combined result from C")
assert.Equal(t, 3, mock.callCount, "Should have called mock provider 3 times")
})
t.Run("Linear chain DAG", func(t *testing.T) {
// Setup: A -> B -> C (linear dependency chain)
mock := &mockProviderWithID{
responses: map[string]string{
"Step 1": "Step 1 output",
"Step 2": "Step 2 output",
"Step 3": "Step 3 output",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mock,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "A", Role: "Step 1", Task: "Do step 1"},
{ID: "B", Role: "Step 2", Task: "Do step 2", DependsOn: []string{"A"}},
{ID: "C", Role: "Step 3", Task: "Do step 3", DependsOn: []string{"B"}},
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
result := tool.executeDAG(ctx, cancel, baseConfig, members, 1000)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "Step 1 output")
assert.Contains(t, result.ForLLM, "Step 2 output")
assert.Contains(t, result.ForLLM, "Step 3 output")
})
t.Run("Detect circular dependency", func(t *testing.T) {
// Setup: A -> B -> C -> A (circular)
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: &mockProvider{},
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "A", Role: "Worker A", Task: "Task A", DependsOn: []string{"C"}},
{ID: "B", Role: "Worker B", Task: "Task B", DependsOn: []string{"A"}},
{ID: "C", Role: "Worker C", Task: "Task C", DependsOn: []string{"B"}},
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
result := tool.executeDAG(ctx, cancel, baseConfig, members, 1000)
assert.True(t, result.IsError, "Should detect circular dependency")
assert.Contains(t, result.ForLLM, "cycle")
})
t.Run("Detect undefined dependency", func(t *testing.T) {
// Setup: A depends on non-existent "X"
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: &mockProvider{},
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "A", Role: "Worker A", Task: "Task A", DependsOn: []string{"X"}},
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
result := tool.executeDAG(ctx, cancel, baseConfig, members, 1000)
assert.True(t, result.IsError, "Should detect undefined dependency")
assert.Contains(t, result.ForLLM, "undefined member")
})
t.Run("Complex DAG with multiple roots", func(t *testing.T) {
// Setup: A, B (roots) -> C, D -> E
mock := &mockProviderWithID{
responses: map[string]string{
"Root A": "Root A output",
"Root B": "Root B output",
"Worker C": "C output",
"Worker D": "D output",
"Final": "Final E output",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mock,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "A", Role: "Root A", Task: "Root task A"},
{ID: "B", Role: "Root B", Task: "Root task B"},
{ID: "C", Role: "Worker C", Task: "Task C", DependsOn: []string{"A"}},
{ID: "D", Role: "Worker D", Task: "Task D", DependsOn: []string{"B"}},
{ID: "E", Role: "Final", Task: "Final task", DependsOn: []string{"C", "D"}},
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
result := tool.executeDAG(ctx, cancel, baseConfig, members, 1000)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "Root A output")
assert.Contains(t, result.ForLLM, "Root B output")
assert.Contains(t, result.ForLLM, "Final E output")
assert.Equal(t, 5, mock.callCount)
})
}
func TestExecuteEvaluatorOptimizer(t *testing.T) {
t.Run("Pass on first attempt", func(t *testing.T) {
mock := &mockProvider{
responses: []string{
"Perfect code implementation",
"[PASS] The code is correct",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{
MaxEvaluatorLoops: 3,
}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mock,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "worker", Role: "Coder", Task: "Write a function"},
{ID: "evaluator", Role: "Code Reviewer", Task: "Review the code"},
}
result := tool.executeEvaluatorOptimizer(context.Background(), baseConfig, members, 1000)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "Perfect code implementation")
assert.Contains(t, result.ForLLM, "[PASS]")
assert.Contains(t, result.ForUser, "passed on attempt 1")
assert.Equal(t, 2, mock.callCount, "Should call worker once and evaluator once")
})
t.Run("Pass on second attempt after feedback", func(t *testing.T) {
mock := &mockProvider{
responses: []string{
"Initial code with bug",
"Missing error handling",
"Fixed code with error handling",
"[PASS] Now it's correct",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{
MaxEvaluatorLoops: 3,
}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mock,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "worker", Role: "Coder", Task: "Write a function"},
{ID: "evaluator", Role: "Code Reviewer", Task: "Review the code"},
}
result := tool.executeEvaluatorOptimizer(context.Background(), baseConfig, members, 1000)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "Initial code with bug")
assert.Contains(t, result.ForLLM, "Missing error handling")
assert.Contains(t, result.ForLLM, "Fixed code with error handling")
assert.Contains(t, result.ForLLM, "[PASS]")
assert.Contains(t, result.ForUser, "passed on attempt 2")
assert.Equal(t, 4, mock.callCount, "Should call worker twice and evaluator twice")
})
t.Run("Exhaust max loops without pass", func(t *testing.T) {
mock := &mockProvider{
responses: []string{
"Attempt 1",
"Still has issues",
"Attempt 2",
"Still not good",
"Attempt 3",
"Still failing",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{
MaxEvaluatorLoops: 3,
}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mock,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "worker", Role: "Coder", Task: "Write a function"},
{ID: "evaluator", Role: "Code Reviewer", Task: "Review the code"},
}
result := tool.executeEvaluatorOptimizer(context.Background(), baseConfig, members, 1000)
assert.False(t, result.IsError, "Should not error, just report exhaustion")
assert.Contains(t, result.ForLLM, "Maximum evaluation loops reached")
assert.Contains(t, result.ForUser, "exhausted 3 attempts")
assert.Equal(t, 6, mock.callCount, "Should call worker 3 times and evaluator 3 times")
})
t.Run("Require exactly two members", func(t *testing.T) {
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: &mockProvider{},
Model: "mock-model",
MaxIterations: 1,
}
// Test with 1 member
members := []TeamMember{
{ID: "worker", Role: "Coder", Task: "Write a function"},
}
result := tool.executeEvaluatorOptimizer(context.Background(), baseConfig, members, 1000)
assert.True(t, result.IsError)
assert.Contains(t, result.ForLLM, "exactly two members")
// Test with 3 members
members = []TeamMember{
{ID: "worker", Role: "Coder", Task: "Write a function"},
{ID: "evaluator", Role: "Reviewer", Task: "Review"},
{ID: "extra", Role: "Extra", Task: "Extra task"},
}
result = tool.executeEvaluatorOptimizer(context.Background(), baseConfig, members, 1000)
assert.True(t, result.IsError)
assert.Contains(t, result.ForLLM, "exactly two members")
})
t.Run("Stateful worker memory across iterations", func(t *testing.T) {
// This test verifies that the worker's message history is preserved
// across iterations, allowing it to "remember" previous feedback
mockWithMemory := &mockProvider{
responses: []string{
"First attempt",
"Needs improvement: add validation",
"Second attempt with validation",
"[PASS] Good now",
},
}
manager := NewSubagentManager(nil, "mock-model", nil, "", config.TeamToolsConfig{
MaxEvaluatorLoops: 3,
}, nil)
tool := NewTeamTool(manager, &config.Config{})
baseConfig := ToolLoopConfig{
Provider: mockWithMemory,
Model: "mock-model",
MaxIterations: 1,
Tools: NewToolRegistry(),
}
members := []TeamMember{
{ID: "worker", Role: "Coder", Task: "Write validation logic"},
{ID: "evaluator", Role: "Reviewer", Task: "Check validation"},
}
result := tool.executeEvaluatorOptimizer(context.Background(), baseConfig, members, 1000)
assert.False(t, result.IsError)
assert.Contains(t, result.ForLLM, "First attempt")
assert.Contains(t, result.ForLLM, "Needs improvement")
assert.Contains(t, result.ForLLM, "Second attempt with validation")
assert.Contains(t, result.ForLLM, "[PASS]")
// Verify the worker was called twice (once initially, once after feedback)
// and evaluator was called twice
assert.Equal(t, 4, mockWithMemory.callCount)
})
}