From 1270ced4b4406896609c7546f824f5177355f3f5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=A8=8B=E5=BA=8F=E5=91=98=E9=98=BF=E6=B1=9F=28Relakkes?= =?UTF-8?q?=29?= Date: Mon, 20 Jul 2026 23:53:35 +0800 Subject: [PATCH] fix(tasks): keep task state reliable across providers #1075 --- desktop/src/stores/chatStore.test.ts | 30 ++++++ desktop/src/stores/chatStore.ts | 13 +++ desktop/src/stores/cliTaskStore.test.ts | 107 +++++++++++++++++++++ desktop/src/stores/cliTaskStore.ts | 56 +++++++++-- src/tools/TaskCreateTool/TaskCreateTool.ts | 2 +- src/tools/TaskGetTool/TaskGetTool.ts | 4 +- src/tools/TaskListTool/TaskListTool.ts | 4 +- src/tools/TaskTools.eager.test.ts | 27 ++++++ src/tools/TaskUpdateTool/TaskUpdateTool.ts | 2 +- 9 files changed, 232 insertions(+), 13 deletions(-) create mode 100644 src/tools/TaskTools.eager.test.ts diff --git a/desktop/src/stores/chatStore.test.ts b/desktop/src/stores/chatStore.test.ts index 029b5d26..ca0f499d 100644 --- a/desktop/src/stores/chatStore.test.ts +++ b/desktop/src/stores/chatStore.test.ts @@ -5261,6 +5261,36 @@ describe('chatStore history mapping', () => { expect(updateTabStatusMock).toHaveBeenCalledWith(TEST_SESSION_ID, 'idle') }) + it('refreshes unfinished Task V2 tool state once when the message completes', () => { + cliTaskStoreSnapshot.sessionId = TEST_SESSION_ID + useChatStore.setState({ + sessions: { + [TEST_SESSION_ID]: makeSession({ chatState: 'tool_executing' }), + }, + }) + + useChatStore.getState().handleServerMessage(TEST_SESSION_ID, { + type: 'tool_use_complete', + toolName: 'TaskUpdate', + toolUseId: 'task-update-1', + input: { taskId: '1', status: 'completed' }, + }) + useChatStore.getState().handleServerMessage(TEST_SESSION_ID, { + type: 'message_complete', + usage: { input_tokens: 1, output_tokens: 2 }, + }) + + expect(refreshTasksMock).toHaveBeenCalledTimes(1) + expect(refreshTasksMock).toHaveBeenCalledWith(TEST_SESSION_ID) + + useChatStore.getState().handleServerMessage(TEST_SESSION_ID, { + type: 'message_complete', + usage: { input_tokens: 1, output_tokens: 2 }, + }) + + expect(refreshTasksMock).toHaveBeenCalledTimes(1) + }) + it('flushes pending text before appending a thinking block', () => { vi.useFakeTimers() diff --git a/desktop/src/stores/chatStore.ts b/desktop/src/stores/chatStore.ts index 7bcc8ac3..04b51772 100644 --- a/desktop/src/stores/chatStore.ts +++ b/desktop/src/stores/chatStore.ts @@ -351,6 +351,13 @@ function clearPendingTaskToolUseIds(sessionId: string): void { pendingTaskToolUseIdsBySession.delete(sessionId) } +function consumeAllPendingTaskToolUseIds(sessionId: string): boolean { + const hasPendingTaskTools = + (pendingTaskToolUseIdsBySession.get(sessionId)?.size ?? 0) > 0 + pendingTaskToolUseIdsBySession.delete(sessionId) + return hasPendingTaskTools +} + function rememberPendingToolParentUseId( sessionId: string, toolUseId: string | null | undefined, @@ -2423,6 +2430,12 @@ export const useChatStore = create((set, get) => ({ case 'message_complete': { const session = get().sessions[sessionId] if (!session) break + if (consumeAllPendingTaskToolUseIds(sessionId)) { + const cliTaskStore = useCLITaskStore.getState() + if (cliTaskStore.sessionId === sessionId) { + void cliTaskStore.refreshTasks(sessionId) + } + } if (session.suppressNextTaskNotificationResponse) { consumePendingDelta(sessionId) clearPendingToolInputDelta(sessionId) diff --git a/desktop/src/stores/cliTaskStore.test.ts b/desktop/src/stores/cliTaskStore.test.ts index 68bb32d8..88ddcef5 100644 --- a/desktop/src/stores/cliTaskStore.test.ts +++ b/desktop/src/stores/cliTaskStore.test.ts @@ -156,6 +156,113 @@ describe('cliTaskStore', () => { ]) }) + it('ignores an older task response that finishes after a newer refresh', async () => { + let resolveOlder: ((value: { tasks: CLITask[] }) => void) | null = null + let resolveNewer: ((value: { tasks: CLITask[] }) => void) | null = null + + vi.mocked(cliTasksApi.getTasksForList) + .mockImplementationOnce(() => new Promise((resolve) => { + resolveOlder = resolve + })) + .mockImplementationOnce(() => new Promise((resolve) => { + resolveNewer = resolve + })) + + useCLITaskStore.setState({ + sessionId: 'session-1', + tasks: [makeTask('session-1', 'pending')], + expanded: true, + completedAndDismissed: false, + dismissedCompletionKey: null, + }) + + const olderRequest = useCLITaskStore.getState().fetchSessionTasks('session-1') + const newerRequest = useCLITaskStore.getState().refreshTasks('session-1') + + resolveNewer!({ tasks: [makeTask('session-1', 'completed')] }) + await newerRequest + resolveOlder!({ tasks: [makeTask('session-1', 'in_progress')] }) + await olderRequest + + expect(useCLITaskStore.getState().tasks).toMatchObject([ + { taskListId: 'session-1', status: 'completed' }, + ]) + }) + + it('applies an older successful response when a newer refresh fails', async () => { + let resolveOlder: ((value: { tasks: CLITask[] }) => void) | null = null + let rejectNewer: ((reason: Error) => void) | null = null + + vi.mocked(cliTasksApi.getTasksForList) + .mockImplementationOnce(() => new Promise((resolve) => { + resolveOlder = resolve + })) + .mockImplementationOnce(() => new Promise((_, reject) => { + rejectNewer = reject + })) + + useCLITaskStore.setState({ + sessionId: 'session-1', + tasks: [], + expanded: false, + completedAndDismissed: false, + dismissedCompletionKey: null, + }) + + const olderRequest = useCLITaskStore.getState().fetchSessionTasks('session-1') + const newerRequest = useCLITaskStore.getState().refreshTasks('session-1') + + rejectNewer!(new Error('temporary failure')) + await newerRequest + resolveOlder!({ tasks: [makeTask('session-1', 'completed')] }) + await olderRequest + + expect(useCLITaskStore.getState().tasks).toMatchObject([ + { taskListId: 'session-1', status: 'completed' }, + ]) + }) + + it('preserves known tasks when polling fails transiently', async () => { + const knownTasks = [makeTask('session-1', 'in_progress')] + vi.mocked(cliTasksApi.getTasksForList).mockRejectedValueOnce(new Error('temporary failure')) + + useCLITaskStore.setState({ + sessionId: 'session-1', + tasks: knownTasks, + expanded: true, + completedAndDismissed: false, + dismissedCompletionKey: null, + }) + + await useCLITaskStore.getState().fetchSessionTasks('session-1') + + expect(useCLITaskStore.getState()).toMatchObject({ + sessionId: 'session-1', + tasks: knownTasks, + expanded: true, + }) + }) + + it('stays empty when the initial fetch for a new session fails', async () => { + vi.mocked(cliTasksApi.getTasksForList).mockRejectedValueOnce(new Error('temporary failure')) + + useCLITaskStore.setState({ + sessionId: 'session-1', + tasks: [makeTask('session-1', 'in_progress')], + expanded: true, + completedAndDismissed: false, + dismissedCompletionKey: null, + }) + + await useCLITaskStore.getState().fetchSessionTasks('session-2') + + expect(useCLITaskStore.getState()).toMatchObject({ + sessionId: 'session-2', + tasks: [], + expanded: false, + }) + }) + it('marks completed tasks dismissed for the currently tracked session by default', () => { useCLITaskStore.setState({ sessionId: 'session-1', diff --git a/desktop/src/stores/cliTaskStore.ts b/desktop/src/stores/cliTaskStore.ts index 8f7007ca..c23e3336 100644 --- a/desktop/src/stores/cliTaskStore.ts +++ b/desktop/src/stores/cliTaskStore.ts @@ -39,6 +39,36 @@ type CLITaskStore = { toggleExpanded: () => void } +let taskRequestSequence = 0 +let taskRequestGeneration = 0 +const latestAppliedTaskRequestBySession = new Map() + +type TaskRequest = { + requestId: number + generation: number +} + +function beginTaskRequest(): TaskRequest { + return { + requestId: ++taskRequestSequence, + generation: taskRequestGeneration, + } +} + +function canApplyTaskResponse(sessionId: string, request: TaskRequest): boolean { + return request.generation === taskRequestGeneration + && request.requestId > (latestAppliedTaskRequestBySession.get(sessionId) ?? 0) +} + +function markTaskResponseApplied(sessionId: string, request: TaskRequest): void { + latestAppliedTaskRequestBySession.set(sessionId, request.requestId) +} + +function invalidateTaskRequests(): void { + taskRequestGeneration += 1 + latestAppliedTaskRequestBySession.clear() +} + function buildCompletedTaskKey(tasks: CLITask[]): string | null { if (tasks.length === 0 || tasks.some((task) => task.status !== 'completed')) return null @@ -89,6 +119,7 @@ export const useCLITaskStore = create((set, get) => ({ fetchSessionTasks: async (sessionId) => { if (get().sessionId !== sessionId) { + invalidateTaskRequests() set({ sessionId, tasks: [], @@ -99,29 +130,37 @@ export const useCLITaskStore = create((set, get) => ({ }) } + const request = beginTaskRequest() try { const { tasks } = await cliTasksApi.getTasksForList(sessionId) - // Only update if still tracking the same session - if (get().sessionId === sessionId && !get().resetting) { + if ( + canApplyTaskResponse(sessionId, request) + && get().sessionId === sessionId + && !get().resetting + ) { + markTaskResponseApplied(sessionId, request) set((state) => ({ tasks, ...resolveDismissState(tasks, state.dismissedCompletionKey), })) } } catch { - // No tasks for this session — that's fine - if (get().sessionId === sessionId && !get().resetting) { - set({ tasks: [], completedAndDismissed: false, dismissedCompletionKey: null, expanded: false }) - } + // Preserve the last known task state across transient polling failures. } }, refreshTasks: async (targetSessionId) => { const sessionId = targetSessionId ?? get().sessionId if (!sessionId) return + const request = beginTaskRequest() try { const { tasks } = await cliTasksApi.getTasksForList(sessionId) - if (get().sessionId === sessionId && !get().resetting) { + if ( + canApplyTaskResponse(sessionId, request) + && get().sessionId === sessionId + && !get().resetting + ) { + markTaskResponseApplied(sessionId, request) set((state) => ({ tasks, ...resolveDismissState(tasks, state.dismissedCompletionKey), @@ -135,6 +174,7 @@ export const useCLITaskStore = create((set, get) => ({ setTasksFromTodos: (todos, targetSessionId) => { const sessionId = targetSessionId ?? get().sessionId if (!sessionId || get().sessionId !== sessionId) return + invalidateTaskRequests() const tasks = mapTodosToTasks(todos, sessionId) set((state) => ({ tasks, @@ -162,6 +202,7 @@ export const useCLITaskStore = create((set, get) => ({ const completionKey = buildCompletedTaskKey(tasks) if (!completionKey) return + invalidateTaskRequests() set({ tasks: [], resetting: true, @@ -181,6 +222,7 @@ export const useCLITaskStore = create((set, get) => ({ clearTasks: (targetSessionId) => { if (targetSessionId && get().sessionId !== targetSessionId) return + invalidateTaskRequests() set({ sessionId: null, tasks: [], diff --git a/src/tools/TaskCreateTool/TaskCreateTool.ts b/src/tools/TaskCreateTool/TaskCreateTool.ts index 2d82c8b7..cb892656 100644 --- a/src/tools/TaskCreateTool/TaskCreateTool.ts +++ b/src/tools/TaskCreateTool/TaskCreateTool.ts @@ -64,7 +64,7 @@ export const TaskCreateTool = buildTool({ userFacingName() { return 'TaskCreate' }, - shouldDefer: true, + alwaysLoad: true, isEnabled() { return isTodoV2Enabled() }, diff --git a/src/tools/TaskGetTool/TaskGetTool.ts b/src/tools/TaskGetTool/TaskGetTool.ts index ffcbfaff..ea9aac9f 100644 --- a/src/tools/TaskGetTool/TaskGetTool.ts +++ b/src/tools/TaskGetTool/TaskGetTool.ts @@ -54,12 +54,12 @@ export const TaskGetTool = buildTool({ userFacingName() { return 'TaskGet' }, - shouldDefer: true, + alwaysLoad: true, isEnabled() { return isTodoV2Enabled() }, isConcurrencySafe() { - return true + return false }, isReadOnly() { return true diff --git a/src/tools/TaskListTool/TaskListTool.ts b/src/tools/TaskListTool/TaskListTool.ts index 61e45589..a01ec63b 100644 --- a/src/tools/TaskListTool/TaskListTool.ts +++ b/src/tools/TaskListTool/TaskListTool.ts @@ -49,12 +49,12 @@ export const TaskListTool = buildTool({ userFacingName() { return 'TaskList' }, - shouldDefer: true, + alwaysLoad: true, isEnabled() { return isTodoV2Enabled() }, isConcurrencySafe() { - return true + return false }, isReadOnly() { return true diff --git a/src/tools/TaskTools.eager.test.ts b/src/tools/TaskTools.eager.test.ts new file mode 100644 index 00000000..d5fd6e88 --- /dev/null +++ b/src/tools/TaskTools.eager.test.ts @@ -0,0 +1,27 @@ +import { describe, expect, it } from 'bun:test' +import { TaskCreateTool } from './TaskCreateTool/TaskCreateTool.js' +import { TaskGetTool } from './TaskGetTool/TaskGetTool.js' +import { TaskListTool } from './TaskListTool/TaskListTool.js' +import { TaskUpdateTool } from './TaskUpdateTool/TaskUpdateTool.js' +import { isDeferredTool } from './ToolSearchTool/prompt.js' + +describe('Task tool discovery', () => { + it('keeps the complete task lifecycle available without ToolSearch', () => { + for (const tool of [ + TaskCreateTool, + TaskGetTool, + TaskListTool, + TaskUpdateTool, + ]) { + expect(tool.alwaysLoad).toBe(true) + expect(isDeferredTool(tool)).toBe(false) + } + }) +}) + +describe('Task tool execution ordering', () => { + it('serializes task reads against concurrent task mutations', () => { + expect(TaskGetTool.isConcurrencySafe({ taskId: '1' })).toBe(false) + expect(TaskListTool.isConcurrencySafe({})).toBe(false) + }) +}) diff --git a/src/tools/TaskUpdateTool/TaskUpdateTool.ts b/src/tools/TaskUpdateTool/TaskUpdateTool.ts index 427831bf..9be2d90f 100644 --- a/src/tools/TaskUpdateTool/TaskUpdateTool.ts +++ b/src/tools/TaskUpdateTool/TaskUpdateTool.ts @@ -104,7 +104,7 @@ export const TaskUpdateTool = buildTool({ userFacingName() { return 'TaskUpdate' }, - shouldDefer: true, + alwaysLoad: true, isEnabled() { return isTodoV2Enabled() },