fix(tasks): keep task state reliable across providers #1075

This commit is contained in:
程序员阿江(Relakkes) 2026-07-20 23:53:35 +08:00
parent 727bcea077
commit 1270ced4b4
9 changed files with 232 additions and 13 deletions

View File

@ -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()

View File

@ -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<ChatStore>((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)

View File

@ -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',

View File

@ -39,6 +39,36 @@ type CLITaskStore = {
toggleExpanded: () => void
}
let taskRequestSequence = 0
let taskRequestGeneration = 0
const latestAppliedTaskRequestBySession = new Map<string, number>()
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<CLITaskStore>((set, get) => ({
fetchSessionTasks: async (sessionId) => {
if (get().sessionId !== sessionId) {
invalidateTaskRequests()
set({
sessionId,
tasks: [],
@ -99,29 +130,37 @@ export const useCLITaskStore = create<CLITaskStore>((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<CLITaskStore>((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<CLITaskStore>((set, get) => ({
const completionKey = buildCompletedTaskKey(tasks)
if (!completionKey) return
invalidateTaskRequests()
set({
tasks: [],
resetting: true,
@ -181,6 +222,7 @@ export const useCLITaskStore = create<CLITaskStore>((set, get) => ({
clearTasks: (targetSessionId) => {
if (targetSessionId && get().sessionId !== targetSessionId) return
invalidateTaskRequests()
set({
sessionId: null,
tasks: [],

View File

@ -64,7 +64,7 @@ export const TaskCreateTool = buildTool({
userFacingName() {
return 'TaskCreate'
},
shouldDefer: true,
alwaysLoad: true,
isEnabled() {
return isTodoV2Enabled()
},

View File

@ -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

View File

@ -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

View File

@ -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)
})
})

View File

@ -104,7 +104,7 @@ export const TaskUpdateTool = buildTool({
userFacingName() {
return 'TaskUpdate'
},
shouldDefer: true,
alwaysLoad: true,
isEnabled() {
return isTodoV2Enabled()
},