mirror of
https://github.com/NanmiCoder/cc-haha
synced 2026-07-25 14:53:32 +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')
|
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', () => {
|
it('flushes pending text before appending a thinking block', () => {
|
||||||
vi.useFakeTimers()
|
vi.useFakeTimers()
|
||||||
|
|
||||||
|
|||||||
@ -351,6 +351,13 @@ function clearPendingTaskToolUseIds(sessionId: string): void {
|
|||||||
pendingTaskToolUseIdsBySession.delete(sessionId)
|
pendingTaskToolUseIdsBySession.delete(sessionId)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function consumeAllPendingTaskToolUseIds(sessionId: string): boolean {
|
||||||
|
const hasPendingTaskTools =
|
||||||
|
(pendingTaskToolUseIdsBySession.get(sessionId)?.size ?? 0) > 0
|
||||||
|
pendingTaskToolUseIdsBySession.delete(sessionId)
|
||||||
|
return hasPendingTaskTools
|
||||||
|
}
|
||||||
|
|
||||||
function rememberPendingToolParentUseId(
|
function rememberPendingToolParentUseId(
|
||||||
sessionId: string,
|
sessionId: string,
|
||||||
toolUseId: string | null | undefined,
|
toolUseId: string | null | undefined,
|
||||||
@ -2423,6 +2430,12 @@ export const useChatStore = create<ChatStore>((set, get) => ({
|
|||||||
case 'message_complete': {
|
case 'message_complete': {
|
||||||
const session = get().sessions[sessionId]
|
const session = get().sessions[sessionId]
|
||||||
if (!session) break
|
if (!session) break
|
||||||
|
if (consumeAllPendingTaskToolUseIds(sessionId)) {
|
||||||
|
const cliTaskStore = useCLITaskStore.getState()
|
||||||
|
if (cliTaskStore.sessionId === sessionId) {
|
||||||
|
void cliTaskStore.refreshTasks(sessionId)
|
||||||
|
}
|
||||||
|
}
|
||||||
if (session.suppressNextTaskNotificationResponse) {
|
if (session.suppressNextTaskNotificationResponse) {
|
||||||
consumePendingDelta(sessionId)
|
consumePendingDelta(sessionId)
|
||||||
clearPendingToolInputDelta(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', () => {
|
it('marks completed tasks dismissed for the currently tracked session by default', () => {
|
||||||
useCLITaskStore.setState({
|
useCLITaskStore.setState({
|
||||||
sessionId: 'session-1',
|
sessionId: 'session-1',
|
||||||
|
|||||||
@ -39,6 +39,36 @@ type CLITaskStore = {
|
|||||||
toggleExpanded: () => void
|
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 {
|
function buildCompletedTaskKey(tasks: CLITask[]): string | null {
|
||||||
if (tasks.length === 0 || tasks.some((task) => task.status !== 'completed')) return 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) => {
|
fetchSessionTasks: async (sessionId) => {
|
||||||
if (get().sessionId !== sessionId) {
|
if (get().sessionId !== sessionId) {
|
||||||
|
invalidateTaskRequests()
|
||||||
set({
|
set({
|
||||||
sessionId,
|
sessionId,
|
||||||
tasks: [],
|
tasks: [],
|
||||||
@ -99,29 +130,37 @@ export const useCLITaskStore = create<CLITaskStore>((set, get) => ({
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
const request = beginTaskRequest()
|
||||||
try {
|
try {
|
||||||
const { tasks } = await cliTasksApi.getTasksForList(sessionId)
|
const { tasks } = await cliTasksApi.getTasksForList(sessionId)
|
||||||
// Only update if still tracking the same session
|
if (
|
||||||
if (get().sessionId === sessionId && !get().resetting) {
|
canApplyTaskResponse(sessionId, request)
|
||||||
|
&& get().sessionId === sessionId
|
||||||
|
&& !get().resetting
|
||||||
|
) {
|
||||||
|
markTaskResponseApplied(sessionId, request)
|
||||||
set((state) => ({
|
set((state) => ({
|
||||||
tasks,
|
tasks,
|
||||||
...resolveDismissState(tasks, state.dismissedCompletionKey),
|
...resolveDismissState(tasks, state.dismissedCompletionKey),
|
||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
} catch {
|
} catch {
|
||||||
// No tasks for this session — that's fine
|
// Preserve the last known task state across transient polling failures.
|
||||||
if (get().sessionId === sessionId && !get().resetting) {
|
|
||||||
set({ tasks: [], completedAndDismissed: false, dismissedCompletionKey: null, expanded: false })
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
refreshTasks: async (targetSessionId) => {
|
refreshTasks: async (targetSessionId) => {
|
||||||
const sessionId = targetSessionId ?? get().sessionId
|
const sessionId = targetSessionId ?? get().sessionId
|
||||||
if (!sessionId) return
|
if (!sessionId) return
|
||||||
|
const request = beginTaskRequest()
|
||||||
try {
|
try {
|
||||||
const { tasks } = await cliTasksApi.getTasksForList(sessionId)
|
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) => ({
|
set((state) => ({
|
||||||
tasks,
|
tasks,
|
||||||
...resolveDismissState(tasks, state.dismissedCompletionKey),
|
...resolveDismissState(tasks, state.dismissedCompletionKey),
|
||||||
@ -135,6 +174,7 @@ export const useCLITaskStore = create<CLITaskStore>((set, get) => ({
|
|||||||
setTasksFromTodos: (todos, targetSessionId) => {
|
setTasksFromTodos: (todos, targetSessionId) => {
|
||||||
const sessionId = targetSessionId ?? get().sessionId
|
const sessionId = targetSessionId ?? get().sessionId
|
||||||
if (!sessionId || get().sessionId !== sessionId) return
|
if (!sessionId || get().sessionId !== sessionId) return
|
||||||
|
invalidateTaskRequests()
|
||||||
const tasks = mapTodosToTasks(todos, sessionId)
|
const tasks = mapTodosToTasks(todos, sessionId)
|
||||||
set((state) => ({
|
set((state) => ({
|
||||||
tasks,
|
tasks,
|
||||||
@ -162,6 +202,7 @@ export const useCLITaskStore = create<CLITaskStore>((set, get) => ({
|
|||||||
const completionKey = buildCompletedTaskKey(tasks)
|
const completionKey = buildCompletedTaskKey(tasks)
|
||||||
if (!completionKey) return
|
if (!completionKey) return
|
||||||
|
|
||||||
|
invalidateTaskRequests()
|
||||||
set({
|
set({
|
||||||
tasks: [],
|
tasks: [],
|
||||||
resetting: true,
|
resetting: true,
|
||||||
@ -181,6 +222,7 @@ export const useCLITaskStore = create<CLITaskStore>((set, get) => ({
|
|||||||
|
|
||||||
clearTasks: (targetSessionId) => {
|
clearTasks: (targetSessionId) => {
|
||||||
if (targetSessionId && get().sessionId !== targetSessionId) return
|
if (targetSessionId && get().sessionId !== targetSessionId) return
|
||||||
|
invalidateTaskRequests()
|
||||||
set({
|
set({
|
||||||
sessionId: null,
|
sessionId: null,
|
||||||
tasks: [],
|
tasks: [],
|
||||||
|
|||||||
@ -64,7 +64,7 @@ export const TaskCreateTool = buildTool({
|
|||||||
userFacingName() {
|
userFacingName() {
|
||||||
return 'TaskCreate'
|
return 'TaskCreate'
|
||||||
},
|
},
|
||||||
shouldDefer: true,
|
alwaysLoad: true,
|
||||||
isEnabled() {
|
isEnabled() {
|
||||||
return isTodoV2Enabled()
|
return isTodoV2Enabled()
|
||||||
},
|
},
|
||||||
|
|||||||
@ -54,12 +54,12 @@ export const TaskGetTool = buildTool({
|
|||||||
userFacingName() {
|
userFacingName() {
|
||||||
return 'TaskGet'
|
return 'TaskGet'
|
||||||
},
|
},
|
||||||
shouldDefer: true,
|
alwaysLoad: true,
|
||||||
isEnabled() {
|
isEnabled() {
|
||||||
return isTodoV2Enabled()
|
return isTodoV2Enabled()
|
||||||
},
|
},
|
||||||
isConcurrencySafe() {
|
isConcurrencySafe() {
|
||||||
return true
|
return false
|
||||||
},
|
},
|
||||||
isReadOnly() {
|
isReadOnly() {
|
||||||
return true
|
return true
|
||||||
|
|||||||
@ -49,12 +49,12 @@ export const TaskListTool = buildTool({
|
|||||||
userFacingName() {
|
userFacingName() {
|
||||||
return 'TaskList'
|
return 'TaskList'
|
||||||
},
|
},
|
||||||
shouldDefer: true,
|
alwaysLoad: true,
|
||||||
isEnabled() {
|
isEnabled() {
|
||||||
return isTodoV2Enabled()
|
return isTodoV2Enabled()
|
||||||
},
|
},
|
||||||
isConcurrencySafe() {
|
isConcurrencySafe() {
|
||||||
return true
|
return false
|
||||||
},
|
},
|
||||||
isReadOnly() {
|
isReadOnly() {
|
||||||
return true
|
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() {
|
userFacingName() {
|
||||||
return 'TaskUpdate'
|
return 'TaskUpdate'
|
||||||
},
|
},
|
||||||
shouldDefer: true,
|
alwaysLoad: true,
|
||||||
isEnabled() {
|
isEnabled() {
|
||||||
return isTodoV2Enabled()
|
return isTodoV2Enabled()
|
||||||
},
|
},
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user