mirror of
https://github.com/NanmiCoder/cc-haha
synced 2026-07-22 14:30:53 +08:00
fix(tasks): keep task state reliable across providers #1075
This commit is contained in:
parent
727bcea077
commit
1270ced4b4
@ -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()
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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',
|
||||
|
||||
@ -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: [],
|
||||
|
||||
@ -64,7 +64,7 @@ export const TaskCreateTool = buildTool({
|
||||
userFacingName() {
|
||||
return 'TaskCreate'
|
||||
},
|
||||
shouldDefer: true,
|
||||
alwaysLoad: true,
|
||||
isEnabled() {
|
||||
return isTodoV2Enabled()
|
||||
},
|
||||
|
||||
@ -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
|
||||
|
||||
@ -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
|
||||
|
||||
27
src/tools/TaskTools.eager.test.ts
Normal file
27
src/tools/TaskTools.eager.test.ts
Normal 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)
|
||||
})
|
||||
})
|
||||
@ -104,7 +104,7 @@ export const TaskUpdateTool = buildTool({
|
||||
userFacingName() {
|
||||
return 'TaskUpdate'
|
||||
},
|
||||
shouldDefer: true,
|
||||
alwaysLoad: true,
|
||||
isEnabled() {
|
||||
return isTodoV2Enabled()
|
||||
},
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user