feat(tools): enhance SpawnStatusTool with task ID validation and sorting by creation timestamp
This commit is contained in:
parent
af6ef94f11
commit
6d370011b5
3 changed files with 98 additions and 24 deletions
|
|
@ -56,18 +56,23 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too
|
||||||
callerChannel := ToolChannel(ctx)
|
callerChannel := ToolChannel(ctx)
|
||||||
callerChatID := ToolChatID(ctx)
|
callerChatID := ToolChatID(ctx)
|
||||||
|
|
||||||
taskID, _ := args["task_id"].(string)
|
var taskID string
|
||||||
taskID = strings.TrimSpace(taskID)
|
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 != "" {
|
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 {
|
if !ok {
|
||||||
return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID))
|
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.
|
// Restrict lookup to tasks that belong to this conversation.
|
||||||
if callerChannel != "" && taskCopy.OriginChannel != "" && taskCopy.OriginChannel != callerChannel {
|
if callerChannel != "" && taskCopy.OriginChannel != "" && taskCopy.OriginChannel != callerChannel {
|
||||||
return ErrorResult(fmt.Sprintf("No subagent found with task ID: %s", taskID))
|
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))
|
return NewToolResult(spawnStatusFormatTask(&taskCopy))
|
||||||
}
|
}
|
||||||
|
|
||||||
origTasks := t.manager.ListTasks()
|
// ListTaskCopies returns consistent snapshots under the manager lock.
|
||||||
|
origTasks := t.manager.ListTaskCopies()
|
||||||
if len(origTasks) == 0 {
|
if len(origTasks) == 0 {
|
||||||
return NewToolResult("No subagents have been spawned yet.")
|
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))
|
tasks := make([]*SubagentTask, 0, len(origTasks))
|
||||||
for _, task := range origTasks {
|
for i := range origTasks {
|
||||||
if task == nil {
|
cpy := &origTasks[i]
|
||||||
continue
|
|
||||||
}
|
|
||||||
cpy := *task
|
|
||||||
|
|
||||||
// Filter to tasks that originate from the current conversation only.
|
// Filter to tasks that originate from the current conversation only.
|
||||||
if callerChannel != "" && cpy.OriginChannel != "" && cpy.OriginChannel != callerChannel {
|
if callerChannel != "" && cpy.OriginChannel != "" && cpy.OriginChannel != callerChannel {
|
||||||
|
|
@ -101,26 +102,24 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
tasks = append(tasks, &cpy)
|
tasks = append(tasks, cpy)
|
||||||
}
|
}
|
||||||
|
|
||||||
if len(tasks) == 0 {
|
if len(tasks) == 0 {
|
||||||
return NewToolResult("No subagents have been spawned yet.")
|
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 {
|
sort.Slice(tasks, func(i, j int) bool {
|
||||||
if tasks[i] == nil || tasks[j] == nil {
|
if tasks[i].Created != tasks[j].Created {
|
||||||
return false
|
return tasks[i].Created < tasks[j].Created
|
||||||
}
|
}
|
||||||
return tasks[i].ID < tasks[j].ID
|
return tasks[i].ID < tasks[j].ID
|
||||||
})
|
})
|
||||||
|
|
||||||
counts := map[string]int{}
|
counts := map[string]int{}
|
||||||
for _, task := range tasks {
|
for _, task := range tasks {
|
||||||
if task == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
counts[task.Status]++
|
counts[task.Status]++
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -135,9 +134,6 @@ func (t *SpawnStatusTool) Execute(ctx context.Context, args map[string]any) *Too
|
||||||
sb.WriteString("\n")
|
sb.WriteString("\n")
|
||||||
|
|
||||||
for _, task := range tasks {
|
for _, task := range tasks {
|
||||||
if task == nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
sb.WriteString(spawnStatusFormatTask(task))
|
sb.WriteString(spawnStatusFormatTask(task))
|
||||||
sb.WriteString("\n\n")
|
sb.WriteString("\n\n")
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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) {
|
func TestSpawnStatusTool_ResultTruncation(t *testing.T) {
|
||||||
provider := &MockLLMProvider{}
|
provider := &MockLLMProvider{}
|
||||||
manager := NewSubagentManager(provider, "test-model", "/tmp/test")
|
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) {
|
func TestSpawnStatusTool_ChannelFiltering_ListAll(t *testing.T) {
|
||||||
provider := &MockLLMProvider{}
|
provider := &MockLLMProvider{}
|
||||||
manager := NewSubagentManager(provider, "test-model", "/tmp/test")
|
manager := NewSubagentManager(provider, "test-model", "/tmp/test")
|
||||||
|
|
|
||||||
|
|
@ -109,8 +109,10 @@ func (sm *SubagentManager) Spawn(
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sm *SubagentManager) runTask(ctx context.Context, task *SubagentTask, callback AsyncCallback) {
|
func (sm *SubagentManager) runTask(ctx context.Context, task *SubagentTask, callback AsyncCallback) {
|
||||||
|
sm.mu.Lock()
|
||||||
task.Status = "running"
|
task.Status = "running"
|
||||||
task.Created = time.Now().UnixMilli()
|
task.Created = time.Now().UnixMilli()
|
||||||
|
sm.mu.Unlock()
|
||||||
|
|
||||||
// Build system prompt for subagent
|
// Build system prompt for subagent
|
||||||
systemPrompt := `You are a subagent. Complete the given task independently and report the result.
|
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
|
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 {
|
func (sm *SubagentManager) ListTasks() []*SubagentTask {
|
||||||
sm.mu.RLock()
|
sm.mu.RLock()
|
||||||
defer sm.mu.RUnlock()
|
defer sm.mu.RUnlock()
|
||||||
|
|
@ -230,6 +244,19 @@ func (sm *SubagentManager) ListTasks() []*SubagentTask {
|
||||||
return tasks
|
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.
|
// SubagentTool executes a subagent task synchronously and returns the result.
|
||||||
// Unlike SpawnTool which runs tasks asynchronously, SubagentTool waits for completion
|
// Unlike SpawnTool which runs tasks asynchronously, SubagentTool waits for completion
|
||||||
// and returns the result directly in the ToolResult.
|
// and returns the result directly in the ToolResult.
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue