diff --git a/pkg/agent/loop.go b/pkg/agent/loop.go index f20a56b9c..48efd4f66 100644 --- a/pkg/agent/loop.go +++ b/pkg/agent/loop.go @@ -227,6 +227,17 @@ func registerSharedTools( if cfg.Tools.IsToolEnabled("subagent") { subagentManager := tools.NewSubagentManager(provider, agent.Model, agent.Workspace) subagentManager.SetLLMOptions(agent.MaxTokens, agent.Temperature) + subagentManager.SetAgentModelResolver(func(targetAgentID string) (string, bool) { + target, ok := registry.GetAgent(targetAgentID) + if !ok { + return "", false + } + model := strings.TrimSpace(target.Model) + if model == "" { + return "", false + } + return model, true + }) spawnTool := tools.NewSpawnTool(subagentManager) currentAgentID := agentID spawnTool.SetAllowlistChecker(func(targetAgentID string) bool { diff --git a/pkg/tools/spawn_test.go b/pkg/tools/spawn_test.go index 43223b8db..b0db8180c 100644 --- a/pkg/tools/spawn_test.go +++ b/pkg/tools/spawn_test.go @@ -4,6 +4,7 @@ import ( "context" "strings" "testing" + "time" ) func TestSpawnTool_Execute_EmptyTask(t *testing.T) { @@ -77,3 +78,45 @@ func TestSpawnTool_Execute_NilManager(t *testing.T) { t.Errorf("Error message should mention manager not configured, got: %s", result.ForLLM) } } + +func TestSpawnTool_ExecuteAsync_UsesTargetAgentModel(t *testing.T) { + provider := &MockLLMProvider{} + manager := NewSubagentManager(provider, "caller-model", "/tmp/test") + manager.SetAgentModelResolver(func(agentID string) (string, bool) { + if agentID == "analyst" { + return "target-model", true + } + return "", false + }) + tool := NewSpawnTool(manager) + + done := make(chan struct{}) + ctx := WithToolContext(context.Background(), "cli", "direct") + args := map[string]any{ + "task": "Write a haiku about coding", + "agent_id": "analyst", + } + + result := tool.ExecuteAsync(ctx, args, func(context.Context, *ToolResult) { + close(done) + }) + if result == nil { + t.Fatal("Result should not be nil") + } + if result.IsError { + t.Fatalf("Expected success for valid task, got error: %s", result.ForLLM) + } + if !result.Async { + t.Fatal("SpawnTool should return async result") + } + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("spawn callback was not invoked") + } + + if provider.lastModel != "target-model" { + t.Fatalf("lastModel = %q, want %q", provider.lastModel, "target-model") + } +} diff --git a/pkg/tools/subagent.go b/pkg/tools/subagent.go index e51cbaafa..4f22cf5e2 100644 --- a/pkg/tools/subagent.go +++ b/pkg/tools/subagent.go @@ -3,6 +3,7 @@ package tools import ( "context" "fmt" + "strings" "sync" "time" @@ -26,6 +27,7 @@ type SubagentManager struct { mu sync.RWMutex provider providers.LLMProvider defaultModel string + agentModelFor func(string) (string, bool) workspace string tools *ToolRegistry maxIterations int @@ -61,6 +63,13 @@ func (sm *SubagentManager) SetLLMOptions(maxTokens int, temperature float64) { sm.hasTemperature = true } +// SetAgentModelResolver resolves the effective model for a targeted subagent. +func (sm *SubagentManager) SetAgentModelResolver(resolve func(string) (string, bool)) { + sm.mu.Lock() + defer sm.mu.Unlock() + sm.agentModelFor = resolve +} + // SetTools sets the tool registry for subagent execution. // If not set, subagent will have access to the provided tools. func (sm *SubagentManager) SetTools(tools *ToolRegistry) { @@ -142,6 +151,8 @@ After completing the task, provide a clear summary of what was done.` // Run tool loop with access to tools sm.mu.RLock() tools := sm.tools + model := sm.defaultModel + agentModelFor := sm.agentModelFor maxIter := sm.maxIterations maxTokens := sm.maxTokens temperature := sm.temperature @@ -149,6 +160,12 @@ After completing the task, provide a clear summary of what was done.` hasTemperature := sm.hasTemperature sm.mu.RUnlock() + if task.AgentID != "" && agentModelFor != nil { + if resolvedModel, ok := agentModelFor(task.AgentID); ok && strings.TrimSpace(resolvedModel) != "" { + model = strings.TrimSpace(resolvedModel) + } + } + var llmOptions map[string]any if hasMaxTokens || hasTemperature { llmOptions = map[string]any{} @@ -162,7 +179,7 @@ After completing the task, provide a clear summary of what was done.` loopResult, err := RunToolLoop(ctx, ToolLoopConfig{ Provider: sm.provider, - Model: sm.defaultModel, + Model: model, Tools: tools, MaxIterations: maxIter, LLMOptions: llmOptions, diff --git a/pkg/tools/subagent_tool_test.go b/pkg/tools/subagent_tool_test.go index 4b6f130a5..66f2b73a3 100644 --- a/pkg/tools/subagent_tool_test.go +++ b/pkg/tools/subagent_tool_test.go @@ -11,6 +11,7 @@ import ( // MockLLMProvider is a test implementation of LLMProvider type MockLLMProvider struct { lastOptions map[string]any + lastModel string } func (m *MockLLMProvider) Chat( @@ -21,6 +22,7 @@ func (m *MockLLMProvider) Chat( options map[string]any, ) (*providers.LLMResponse, error) { m.lastOptions = options + m.lastModel = model // Find the last user message to generate a response for i := len(messages) - 1; i >= 0; i-- { if messages[i].Role == "user" { @@ -69,6 +71,34 @@ func TestSubagentManager_SetLLMOptions_AppliesToRunToolLoop(t *testing.T) { } } +func TestSubagentManager_RunTask_UsesResolvedTargetAgentModel(t *testing.T) { + provider := &MockLLMProvider{} + manager := NewSubagentManager(provider, "caller-model", "/tmp/test") + manager.SetAgentModelResolver(func(agentID string) (string, bool) { + if agentID == "analyst" { + return "target-model", true + } + return "", false + }) + + task := &SubagentTask{ + ID: "subagent-1", + Task: "Do something", + AgentID: "analyst", + OriginChannel: "cli", + OriginChatID: "direct", + } + + manager.runTask(context.Background(), task, nil) + + if provider.lastModel != "target-model" { + t.Fatalf("lastModel = %q, want %q", provider.lastModel, "target-model") + } + if task.Status != "completed" { + t.Fatalf("task.Status = %q, want %q", task.Status, "completed") + } +} + // TestSubagentTool_Name verifies tool name func TestSubagentTool_Name(t *testing.T) { provider := &MockLLMProvider{}