feat: add model parameter to spawn and subagent tools
Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
73575c7d7d
commit
1b96f79577
3 changed files with 109 additions and 4 deletions
|
|
@ -44,6 +44,10 @@ func (t *SpawnTool) Parameters() map[string]any {
|
|||
"type": "string",
|
||||
"description": "Optional target agent ID to delegate the task to",
|
||||
},
|
||||
"model": map[string]any{
|
||||
"type": "string",
|
||||
"description": "Optional model_name from model_list to use for the subagent",
|
||||
},
|
||||
},
|
||||
"required": []string{"task"},
|
||||
}
|
||||
|
|
@ -71,6 +75,7 @@ func (t *SpawnTool) execute(ctx context.Context, args map[string]any, cb AsyncCa
|
|||
|
||||
label, _ := args["label"].(string)
|
||||
agentID, _ := args["agent_id"].(string)
|
||||
model, _ := args["model"].(string)
|
||||
|
||||
// Check allowlist if targeting a specific agent
|
||||
if agentID != "" && t.allowlistCheck != nil {
|
||||
|
|
@ -96,7 +101,7 @@ func (t *SpawnTool) execute(ctx context.Context, args map[string]any, cb AsyncCa
|
|||
}
|
||||
|
||||
// Pass callback to manager for async completion notification
|
||||
result, err := t.manager.Spawn(ctx, task, label, agentID, channel, chatID, cb)
|
||||
result, err := t.manager.Spawn(ctx, task, label, agentID, model, channel, chatID, cb)
|
||||
if err != nil {
|
||||
return ErrorResult(fmt.Sprintf("failed to spawn subagent: %v", err))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -120,7 +120,7 @@ func (sm *SubagentManager) RegisterTool(tool Tool) {
|
|||
|
||||
func (sm *SubagentManager) Spawn(
|
||||
ctx context.Context,
|
||||
task, label, agentID, originChannel, originChatID string,
|
||||
task, label, agentID, model, originChannel, originChatID string,
|
||||
callback AsyncCallback,
|
||||
) (string, error) {
|
||||
sm.mu.Lock()
|
||||
|
|
@ -133,6 +133,7 @@ func (sm *SubagentManager) Spawn(
|
|||
ID: taskID,
|
||||
Task: task,
|
||||
Label: label,
|
||||
Model: model,
|
||||
AgentID: agentID,
|
||||
OriginChannel: originChannel,
|
||||
OriginChatID: originChatID,
|
||||
|
|
@ -181,6 +182,19 @@ After completing the task, provide a clear summary of what was done.`
|
|||
default:
|
||||
}
|
||||
|
||||
// Resolve the model for this task
|
||||
resolvedModel, err := sm.ResolveModel(task.Model)
|
||||
if err != nil {
|
||||
sm.mu.Lock()
|
||||
task.Status = "failed"
|
||||
task.Result = err.Error()
|
||||
sm.mu.Unlock()
|
||||
if callback != nil {
|
||||
callback(ctx, ErrorResult(err.Error()))
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Run tool loop with access to tools
|
||||
sm.mu.RLock()
|
||||
tools := sm.tools
|
||||
|
|
@ -204,7 +218,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: resolvedModel,
|
||||
Tools: tools,
|
||||
MaxIterations: maxIter,
|
||||
LLMOptions: llmOptions,
|
||||
|
|
@ -319,6 +333,10 @@ func (t *SubagentTool) Parameters() map[string]any {
|
|||
"type": "string",
|
||||
"description": "Optional short label for the task (for display)",
|
||||
},
|
||||
"model": map[string]any{
|
||||
"type": "string",
|
||||
"description": "Optional model_name from model_list to use for the subagent",
|
||||
},
|
||||
},
|
||||
"required": []string{"task"},
|
||||
}
|
||||
|
|
@ -331,11 +349,18 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
|||
}
|
||||
|
||||
label, _ := args["label"].(string)
|
||||
model, _ := args["model"].(string)
|
||||
|
||||
if t.manager == nil {
|
||||
return ErrorResult("Subagent manager not configured").WithError(fmt.Errorf("manager is nil"))
|
||||
}
|
||||
|
||||
// Resolve the model before proceeding
|
||||
resolvedModel, err := t.manager.ResolveModel(model)
|
||||
if err != nil {
|
||||
return ErrorResult(fmt.Sprintf("invalid model: %v", err)).WithError(err)
|
||||
}
|
||||
|
||||
// Build messages for subagent
|
||||
messages := []providers.Message{
|
||||
{
|
||||
|
|
@ -383,7 +408,7 @@ func (t *SubagentTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
|||
|
||||
loopResult, err := RunToolLoop(ctx, ToolLoopConfig{
|
||||
Provider: sm.provider,
|
||||
Model: sm.defaultModel,
|
||||
Model: resolvedModel,
|
||||
Tools: tools,
|
||||
MaxIterations: maxIter,
|
||||
LLMOptions: llmOptions,
|
||||
|
|
|
|||
|
|
@ -362,6 +362,81 @@ func TestSubagentManager_ResolveModel(t *testing.T) {
|
|||
})
|
||||
}
|
||||
|
||||
func TestSubagentTool_Execute_WithModel(t *testing.T) {
|
||||
provider := &MockLLMProvider{}
|
||||
msgBus := bus.NewMessageBus()
|
||||
manager := NewSubagentManager(provider, "parent-model", "/tmp/test", msgBus)
|
||||
manager.SetModelValidator(func(name string) bool {
|
||||
return name == "custom-model" || name == "parent-model"
|
||||
})
|
||||
tool := NewSubagentTool(manager)
|
||||
|
||||
ctx := WithToolContext(context.Background(), "cli", "direct")
|
||||
args := map[string]any{
|
||||
"task": "Test with custom model",
|
||||
"model": "custom-model",
|
||||
}
|
||||
|
||||
result := tool.Execute(ctx, args)
|
||||
if result.IsError {
|
||||
t.Errorf("Expected success, got error: %s", result.ForLLM)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSubagentTool_Execute_WithInvalidModel(t *testing.T) {
|
||||
provider := &MockLLMProvider{}
|
||||
msgBus := bus.NewMessageBus()
|
||||
manager := NewSubagentManager(provider, "parent-model", "/tmp/test", msgBus)
|
||||
manager.SetModelValidator(func(name string) bool {
|
||||
return name == "parent-model"
|
||||
})
|
||||
tool := NewSubagentTool(manager)
|
||||
|
||||
ctx := WithToolContext(context.Background(), "cli", "direct")
|
||||
args := map[string]any{
|
||||
"task": "Test with bad model",
|
||||
"model": "nonexistent-model",
|
||||
}
|
||||
|
||||
result := tool.Execute(ctx, args)
|
||||
if !result.IsError {
|
||||
t.Error("Expected error for invalid model")
|
||||
}
|
||||
if !strings.Contains(result.ForLLM, "not found in model_list") {
|
||||
t.Errorf("Error should mention model_list, got: %s", result.ForLLM)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSpawnTool_Parameters_IncludesModel(t *testing.T) {
|
||||
provider := &MockLLMProvider{}
|
||||
manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil)
|
||||
tool := NewSpawnTool(manager)
|
||||
params := tool.Parameters()
|
||||
props := params["properties"].(map[string]any)
|
||||
model, ok := props["model"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatal("model parameter should exist in spawn tool")
|
||||
}
|
||||
if model["type"] != "string" {
|
||||
t.Errorf("model type should be 'string', got: %v", model["type"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSubagentTool_Parameters_IncludesModel(t *testing.T) {
|
||||
provider := &MockLLMProvider{}
|
||||
manager := NewSubagentManager(provider, "test-model", "/tmp/test", nil)
|
||||
tool := NewSubagentTool(manager)
|
||||
params := tool.Parameters()
|
||||
props := params["properties"].(map[string]any)
|
||||
model, ok := props["model"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatal("model parameter should exist in subagent tool")
|
||||
}
|
||||
if model["type"] != "string" {
|
||||
t.Errorf("model type should be 'string', got: %v", model["type"])
|
||||
}
|
||||
}
|
||||
|
||||
// TestSubagentTool_ForUserTruncation verifies long content is truncated for user
|
||||
func TestSubagentTool_ForUserTruncation(t *testing.T) {
|
||||
// Create a mock provider that returns very long content
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue