From 6d370011b5e2b00ee91ed00fe00b0a4e5df9dc02 Mon Sep 17 00:00:00 2001 From: SHINE-six Date: Sat, 14 Mar 2026 17:22:28 +0800 Subject: [PATCH] feat(tools): enhance SpawnStatusTool with task ID validation and sorting by creation timestamp --- pkg/tools/spawn_status.go | 44 +++++++++++++---------------- pkg/tools/spawn_status_test.go | 51 ++++++++++++++++++++++++++++++++++ pkg/tools/subagent.go | 27 ++++++++++++++++++ 3 files changed, 98 insertions(+), 24 deletions(-) diff --git a/pkg/tools/spawn_status.go b/pkg/tools/spawn_status.go index fd17e1abc..fec37bd99 100644 --- a/pkg/tools/spawn_status.go +++ b/pkg/tools/spawn_status.go @@ -56,18 +56,23 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too callerChannel := ToolChannel(ctx) callerChatID := ToolChatID(ctx) - taskID, _ := args["task_id"].(string) - taskID = strings.TrimSpace(taskID) + var taskID string + if rawTaskID, ok := args["task_id"]; ok && rawTaskID != nil { + taskIDStr, ok := rawTaskID.(string) + if !ok { + return ErrorResult("task_id must be a string") + } + taskID = strings.TrimSpace(taskIDStr) + } if taskID != "" { - task, ok := t.manager.GetTask(taskID) + // GetTaskCopy returns a consistent snapshot under the manager lock, + // eliminating any data race with the concurrent subagent goroutine. + taskCopy, ok := t.manager.GetTaskCopy(taskID) if !ok { return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID)) } - // Snapshot before formatting to avoid racing with the subagent goroutine. - taskCopy := *task - // Restrict lookup to tasks that belong to this conversation. if callerChannel != "" && taskCopy.OriginChannel != "" && taskCopy.OriginChannel != callerChannel { return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID)) @@ -79,19 +84,15 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too return NewToolResult(spawnStatusFormatTask(&taskCopy)) } - origTasks := t.manager.ListTasks() + // ListTaskCopies returns consistent snapshots under the manager lock. + origTasks := t.manager.ListTaskCopies() if len(origTasks) == 0 { return NewToolResult("No subagents have been spawned yet.") } - // Snapshot each task to avoid reading concurrently mutated state via shared - // pointers once the manager lock is released. tasks := make([]*SubagentTask, 0, len(origTasks)) - for _, task := range origTasks { - if task == nil { - continue - } - cpy := *task + for i := range origTasks { + cpy := &origTasks[i] // Filter to tasks that originate from the current conversation only. if callerChannel != "" && cpy.OriginChannel != "" && cpy.OriginChannel != callerChannel { @@ -101,26 +102,24 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too continue } - tasks = append(tasks, &cpy) + tasks = append(tasks, cpy) } if len(tasks) == 0 { return NewToolResult("No subagents have been spawned yet.") } - // Deterministic ordering: sort by ID string (e.g. "subagent-1" < "subagent-2"). + // Order by creation time (ascending) so spawning order is preserved. + // Fall back to ID string for tasks created in the same millisecond. sort.Slice(tasks, func(i, j int) bool { - if tasks[i] == nil || tasks[j] == nil { - return false + if tasks[i].Created != tasks[j].Created { + return tasks[i].Created < tasks[j].Created } return tasks[i].ID < tasks[j].ID }) counts := map[string]int{} for _, task := range tasks { - if task == nil { - continue - } counts[task.Status]++ } @@ -135,9 +134,6 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too sb.WriteString("\n") for _, task := range tasks { - if task == nil { - continue - } sb.WriteString(spawnStatusFormatTask(task)) sb.WriteString("\n\n") } diff --git a/pkg/tools/spawn_status_test.go b/pkg/tools/spawn_status_test.go index 4d77d0a54..89e6bfe3a 100644 --- a/pkg/tools/spawn_status_test.go +++ b/pkg/tools/spawn_status_test.go @@ -182,6 +182,22 @@ func TestSpawnStatusTool_GetByID_NotFound(t *testing.T) { } } +func TestSpawnStatusTool_TaskID_NonString(t *testing.T) { + provider := &MockLLMProvider{} + manager := NewSubagentManager(provider, "test-model", "/tmp/test") + tool := NewSpawnStatusTool(manager) + + for _, badVal := range []any{42, 3.14, true, map[string]any{"x": 1}, []string{"a"}} { + result := tool.Execute(context.Background(), map[string]any{"task_id": badVal}) + if !result.IsError { + t.Errorf("Expected error for task_id=%T(%v), got success: %s", badVal, badVal, result.ForLLM) + } + if !strings.Contains(result.ForLLM, "task_id must be a string") { + t.Errorf("Expected type-error message, got: %s", result.ForLLM) + } + } +} + func TestSpawnStatusTool_ResultTruncation(t *testing.T) { provider := &MockLLMProvider{} manager := NewSubagentManager(provider, "test-model", "/tmp/test") @@ -266,6 +282,41 @@ func TestSpawnStatusTool_StatusCounts(t *testing.T) { } } +func TestSpawnStatusTool_SortByCreatedTimestamp(t *testing.T) { + provider := &MockLLMProvider{} + manager := NewSubagentManager(provider, "test-model", "/tmp/test") + + now := time.Now().UnixMilli() + manager.mu.Lock() + // Intentionally insert with out-of-order IDs and timestamps that reflect + // true spawn order: subagent-2 was spawned first, subagent-10 second. + manager.tasks["subagent-10"] = &SubagentTask{ + ID: "subagent-10", Task: "second", Status: "running", + Created: now + 1, + } + manager.tasks["subagent-2"] = &SubagentTask{ + ID: "subagent-2", Task: "first", Status: "running", + Created: now, + } + manager.mu.Unlock() + + tool := NewSpawnStatusTool(manager) + result := tool.Execute(context.Background(), map[string]any{}) + + if result.IsError { + t.Fatalf("Unexpected error: %s", result.ForLLM) + } + + pos2 := strings.Index(result.ForLLM, "subagent-2") + pos10 := strings.Index(result.ForLLM, "subagent-10") + if pos2 < 0 || pos10 < 0 { + t.Fatalf("Both task IDs should appear in output:\n%s", result.ForLLM) + } + if pos2 > pos10 { + t.Errorf("Expected subagent-2 (created first) to appear before subagent-10, but got:\n%s", result.ForLLM) + } +} + func TestSpawnStatusTool_ChannelFiltering_ListAll(t *testing.T) { provider := &MockLLMProvider{} manager := NewSubagentManager(provider, "test-model", "/tmp/test") diff --git a/pkg/tools/subagent.go b/pkg/tools/subagent.go index e51cbaafa..fc13c0f09 100644 --- a/pkg/tools/subagent.go +++ b/pkg/tools/subagent.go @@ -109,8 +109,10 @@ func (sm *SubagentManager) Spawn( } func (sm *SubagentManager) runTask(ctx context.Context, task *SubagentTask, callback AsyncCallback) { + sm.mu.Lock() task.Status = "running" task.Created = time.Now().UnixMilli() + sm.mu.Unlock() // Build system prompt for subagent systemPrompt := `You are a subagent. Complete the given task independently and report the result. @@ -219,6 +221,18 @@ func (sm *SubagentManager) GetTask(taskID string) (*SubagentTask, bool) { return task, ok } +// GetTaskCopy returns a copy of the task with the given ID, taken under the +// read lock, so the caller receives a consistent snapshot with no data race. +func (sm *SubagentManager) GetTaskCopy(taskID string) (SubagentTask, bool) { + sm.mu.RLock() + defer sm.mu.RUnlock() + task, ok := sm.tasks[taskID] + if !ok { + return SubagentTask{}, false + } + return *task, true +} + func (sm *SubagentManager) ListTasks() []*SubagentTask { sm.mu.RLock() defer sm.mu.RUnlock() @@ -230,6 +244,19 @@ func (sm *SubagentManager) ListTasks() []*SubagentTask { return tasks } +// ListTaskCopies returns value copies of all tasks, taken under the read lock, +// so callers receive consistent snapshots with no data race. +func (sm *SubagentManager) ListTaskCopies() []SubagentTask { + sm.mu.RLock() + defer sm.mu.RUnlock() + + copies := make([]SubagentTask, 0, len(sm.tasks)) + for _, task := range sm.tasks { + copies = append(copies, *task) + } + return copies +} + // SubagentTool executes a subagent task synchronously and returns the result. // Unlike SpawnTool which runs tasks asynchronously, SubagentTool waits for completion // and returns the result directly in the ToolResult.