From 26cd169f7eee7d2e03f895ae18b33307f9069078 Mon Sep 17 00:00:00 2001 From: Leandro Barbosa Date: Wed, 18 Feb 2026 15:48:40 -0300 Subject: [PATCH] test: add comprehensive tests for tool policy pipeline - groups_test.go: 6 tests for ResolveToolNames (group expansion, individual, mixed, dedup, unknown group, empty) - policy_test.go: 11 tests for ApplyPolicy, DepthDenyList, Clone, Remove, and policy composition pipeline - handoff_test.go: 2 depth policy tests (leaf loses spawn/handoff, mid-chain retains all tools) --- pkg/multiagent/handoff_test.go | 84 ++++++++++++++++ pkg/tools/groups_test.go | 68 +++++++++++++ pkg/tools/policy_test.go | 179 +++++++++++++++++++++++++++++++++ 3 files changed, 331 insertions(+) create mode 100644 pkg/tools/groups_test.go create mode 100644 pkg/tools/policy_test.go diff --git a/pkg/multiagent/handoff_test.go b/pkg/multiagent/handoff_test.go index eb3527a28..56082f9af 100644 --- a/pkg/multiagent/handoff_test.go +++ b/pkg/multiagent/handoff_test.go @@ -737,6 +737,90 @@ func TestHandoffTool_SetContext(t *testing.T) { } } +// TestHandoff_DepthPolicy_LeafNoSpawn verifies that at max depth, the target agent's +// tool registry clone has spawn/handoff/list_agents removed. +func TestHandoff_DepthPolicy_LeafNoSpawn(t *testing.T) { + provider := &mockProvider{response: "leaf result"} + targetRegistry := tools.NewToolRegistry() + targetRegistry.Register(&simpleTool{name: "read_file"}) + targetRegistry.Register(&simpleTool{name: "spawn"}) + targetRegistry.Register(&simpleTool{name: "handoff"}) + targetRegistry.Register(&simpleTool{name: "list_agents"}) + + resolver := newMockResolver(&AgentInfo{ + ID: "leaf", Name: "Leaf Agent", Model: "test", + Provider: provider, Tools: targetRegistry, MaxIter: 5, + }) + + bb := NewBlackboard() + // Depth 2, maxDepth 3: the target will run at depth 3 (req.Depth+1), + // which equals maxDepth, triggering depth deny. + result := ExecuteHandoff(context.Background(), resolver, bb, HandoffRequest{ + FromAgentID: "main", + ToAgentID: "leaf", + Task: "do something as a leaf", + Depth: 2, + MaxDepth: 3, + Visited: []string{"main", "middle"}, + }, "cli", "direct") + + if !result.Success { + t.Fatalf("expected success, got error: %s", result.Error) + } + + // Original registry should still have all 4 tools (clone was modified, not original) + if targetRegistry.Count() != 4 { + t.Errorf("original registry count = %d, want 4 (unmodified)", targetRegistry.Count()) + } +} + +// TestHandoff_DepthPolicy_MidChain verifies that mid-chain agents retain all tools. +func TestHandoff_DepthPolicy_MidChain(t *testing.T) { + provider := &mockProvider{response: "mid-chain result"} + targetRegistry := tools.NewToolRegistry() + targetRegistry.Register(&simpleTool{name: "read_file"}) + targetRegistry.Register(&simpleTool{name: "spawn"}) + targetRegistry.Register(&simpleTool{name: "handoff"}) + targetRegistry.Register(&simpleTool{name: "list_agents"}) + + resolver := newMockResolver(&AgentInfo{ + ID: "mid", Name: "Mid Agent", Model: "test", + Provider: provider, Tools: targetRegistry, MaxIter: 5, + }) + + bb := NewBlackboard() + // Depth 0, maxDepth 3: target runs at depth 1, well below max. + result := ExecuteHandoff(context.Background(), resolver, bb, HandoffRequest{ + FromAgentID: "main", + ToAgentID: "mid", + Task: "mid-chain task", + Depth: 0, + MaxDepth: 3, + Visited: []string{"main"}, + }, "cli", "direct") + + if !result.Success { + t.Fatalf("expected success, got error: %s", result.Error) + } + + // Original registry should still have all 4 tools + if targetRegistry.Count() != 4 { + t.Errorf("original registry count = %d, want 4", targetRegistry.Count()) + } +} + +// simpleTool is a minimal tool for depth policy tests. +type simpleTool struct { + name string +} + +func (s *simpleTool) Name() string { return s.name } +func (s *simpleTool) Description() string { return "test tool" } +func (s *simpleTool) Parameters() map[string]interface{} { return nil } +func (s *simpleTool) Execute(_ context.Context, _ map[string]interface{}) *tools.ToolResult { + return tools.NewToolResult("ok") +} + func TestBuildHandoffSystemPrompt(t *testing.T) { agent := &AgentInfo{ Name: "Code Agent", diff --git a/pkg/tools/groups_test.go b/pkg/tools/groups_test.go new file mode 100644 index 000000000..90fcda8b8 --- /dev/null +++ b/pkg/tools/groups_test.go @@ -0,0 +1,68 @@ +package tools + +import ( + "sort" + "testing" +) + +func TestResolveToolNames_GroupExpansion(t *testing.T) { + result := ResolveToolNames([]string{"group:fs"}) + expected := []string{"read_file", "write_file", "edit_file", "append_file", "list_dir"} + + if len(result) != len(expected) { + t.Fatalf("len = %d, want %d: %v", len(result), len(expected), result) + } + sort.Strings(result) + sort.Strings(expected) + for i := range expected { + if result[i] != expected[i] { + t.Errorf("result[%d] = %q, want %q", i, result[i], expected[i]) + } + } +} + +func TestResolveToolNames_IndividualTool(t *testing.T) { + result := ResolveToolNames([]string{"exec"}) + if len(result) != 1 || result[0] != "exec" { + t.Errorf("result = %v, want [exec]", result) + } +} + +func TestResolveToolNames_Mixed(t *testing.T) { + result := ResolveToolNames([]string{"group:web", "exec"}) + expected := map[string]bool{"web_search": true, "web_fetch": true, "exec": true} + if len(result) != len(expected) { + t.Fatalf("len = %d, want %d: %v", len(result), len(expected), result) + } + for _, name := range result { + if !expected[name] { + t.Errorf("unexpected tool: %q", name) + } + } +} + +func TestResolveToolNames_Dedup(t *testing.T) { + result := ResolveToolNames([]string{"group:exec", "exec"}) + if len(result) != 1 { + t.Errorf("expected 1 (deduped), got %d: %v", len(result), result) + } +} + +func TestResolveToolNames_UnknownGroup(t *testing.T) { + result := ResolveToolNames([]string{"group:nonexistent"}) + // Unknown group ref treated as a literal tool name + if len(result) != 1 || result[0] != "group:nonexistent" { + t.Errorf("result = %v, want [group:nonexistent]", result) + } +} + +func TestResolveToolNames_Empty(t *testing.T) { + result := ResolveToolNames(nil) + if len(result) != 0 { + t.Errorf("result = %v, want empty", result) + } + result = ResolveToolNames([]string{}) + if len(result) != 0 { + t.Errorf("result = %v, want empty", result) + } +} diff --git a/pkg/tools/policy_test.go b/pkg/tools/policy_test.go new file mode 100644 index 000000000..d5b05578d --- /dev/null +++ b/pkg/tools/policy_test.go @@ -0,0 +1,179 @@ +package tools + +import ( + "sort" + "testing" +) + +func setupTestRegistry(names ...string) *ToolRegistry { + reg := NewToolRegistry() + for _, name := range names { + reg.Register(&dummyTool{name: name}) + } + return reg +} + +func registryNames(reg *ToolRegistry) []string { + names := reg.List() + sort.Strings(names) + return names +} + +func TestApplyPolicy_AllowOnly(t *testing.T) { + reg := setupTestRegistry("read_file", "write_file", "exec", "web_search") + ApplyPolicy(reg, ToolPolicy{Allow: []string{"read_file", "exec"}}) + + names := registryNames(reg) + if len(names) != 2 { + t.Fatalf("count = %d, want 2: %v", len(names), names) + } + if names[0] != "exec" || names[1] != "read_file" { + t.Errorf("names = %v, want [exec, read_file]", names) + } +} + +func TestApplyPolicy_DenyOnly(t *testing.T) { + reg := setupTestRegistry("read_file", "write_file", "exec", "web_search") + ApplyPolicy(reg, ToolPolicy{Deny: []string{"exec", "web_search"}}) + + names := registryNames(reg) + if len(names) != 2 { + t.Fatalf("count = %d, want 2: %v", len(names), names) + } + if names[0] != "read_file" || names[1] != "write_file" { + t.Errorf("names = %v, want [read_file, write_file]", names) + } +} + +func TestApplyPolicy_AllowAndDeny(t *testing.T) { + reg := setupTestRegistry("read_file", "write_file", "exec", "web_search") + ApplyPolicy(reg, ToolPolicy{ + Allow: []string{"read_file", "write_file", "exec"}, + Deny: []string{"exec"}, + }) + + names := registryNames(reg) + if len(names) != 2 { + t.Fatalf("count = %d, want 2: %v", len(names), names) + } + if names[0] != "read_file" || names[1] != "write_file" { + t.Errorf("names = %v, want [read_file, write_file]", names) + } +} + +func TestApplyPolicy_EmptyPolicy(t *testing.T) { + reg := setupTestRegistry("read_file", "write_file", "exec") + ApplyPolicy(reg, ToolPolicy{}) + + if reg.Count() != 3 { + t.Errorf("count = %d, want 3 (no-op)", reg.Count()) + } +} + +func TestApplyPolicy_GroupRefs(t *testing.T) { + reg := setupTestRegistry("read_file", "write_file", "edit_file", "append_file", "list_dir", "web_search", "web_fetch", "exec") + ApplyPolicy(reg, ToolPolicy{Deny: []string{"group:web"}}) + + names := registryNames(reg) + for _, name := range names { + if name == "web_search" || name == "web_fetch" { + t.Errorf("web tool %q should have been denied", name) + } + } + if reg.Count() != 6 { + t.Errorf("count = %d, want 6", reg.Count()) + } +} + +func TestDepthDenyList_Zero(t *testing.T) { + result := DepthDenyList(0, 3) + if result != nil { + t.Errorf("depth 0 should return nil, got %v", result) + } +} + +func TestDepthDenyList_AtMax(t *testing.T) { + result := DepthDenyList(3, 3) + expected := []string{"spawn", "handoff", "list_agents"} + if len(result) != len(expected) { + t.Fatalf("len = %d, want %d", len(result), len(expected)) + } + for i, name := range expected { + if result[i] != name { + t.Errorf("result[%d] = %q, want %q", i, result[i], name) + } + } +} + +func TestDepthDenyList_BelowMax(t *testing.T) { + result := DepthDenyList(1, 3) + if result != nil { + t.Errorf("mid-chain should return nil, got %v", result) + } +} + +func TestRegistryClone(t *testing.T) { + reg := setupTestRegistry("tool_a", "tool_b", "tool_c") + cloned := reg.Clone() + + // Same tools + if cloned.Count() != 3 { + t.Fatalf("cloned count = %d, want 3", cloned.Count()) + } + + // Independent: removing from clone doesn't affect original + cloned.Remove("tool_b") + if cloned.Count() != 2 { + t.Errorf("cloned count after remove = %d, want 2", cloned.Count()) + } + if reg.Count() != 3 { + t.Errorf("original count after clone remove = %d, want 3", reg.Count()) + } +} + +func TestRegistryRemove(t *testing.T) { + reg := setupTestRegistry("tool_a", "tool_b") + reg.Remove("tool_a") + + if reg.Count() != 1 { + t.Fatalf("count = %d, want 1", reg.Count()) + } + if _, ok := reg.Get("tool_a"); ok { + t.Error("tool_a should have been removed") + } + if _, ok := reg.Get("tool_b"); !ok { + t.Error("tool_b should still exist") + } + + // Remove nonexistent tool — no panic + reg.Remove("nonexistent") + if reg.Count() != 1 { + t.Errorf("count = %d, want 1 after removing nonexistent", reg.Count()) + } +} + +func TestPolicyPipeline_Compose(t *testing.T) { + // Simulate: global allow → per-agent deny → depth deny + reg := setupTestRegistry( + "read_file", "write_file", "exec", + "web_search", "spawn", "handoff", "list_agents", + ) + + // Layer 1: per-agent policy (deny web) + ApplyPolicy(reg, ToolPolicy{Deny: []string{"web_search"}}) + + // Layer 2: depth policy (leaf: deny spawn/handoff/list_agents) + denyList := DepthDenyList(3, 3) // at max depth + ApplyPolicy(reg, ToolPolicy{Deny: denyList}) + + names := registryNames(reg) + expected := map[string]bool{"read_file": true, "write_file": true, "exec": true} + if len(names) != len(expected) { + t.Fatalf("count = %d, want %d: %v", len(names), len(expected), names) + } + for _, name := range names { + if !expected[name] { + t.Errorf("unexpected tool: %q", name) + } + } +}