| |
| |
| |
| |
| |
|
|
| import { describe, it, expect, vi, beforeEach, type Mock } from 'vitest'; |
| import { CoderAgentExecutor } from './executor.js'; |
| import type { |
| ExecutionEventBus, |
| RequestContext, |
| TaskStore, |
| } from '@a2a-js/sdk/server'; |
| import { EventEmitter } from 'node:events'; |
| import { requestStorage } from '../http/requestStorage.js'; |
|
|
| vi.mock('../utils/path_utils.js', () => ({ |
| validateWorkspacePath: vi |
| .fn() |
| .mockImplementation(async (path?: string) => path || process.cwd()), |
| })); |
|
|
| |
| vi.mock('@google/gemini-cli-core', () => ({ |
| GeminiEventType: { |
| PRIMARY_TURN_STARTED: 'PRIMARY_TURN_STARTED', |
| SECONDARY_TURN_STARTED: 'SECONDARY_TURN_STARTED', |
| }, |
| SimpleExtensionLoader: vi.fn(), |
| checkPathTrust: vi.fn().mockReturnValue({ isTrusted: false }), |
| isHeadlessMode: vi.fn().mockReturnValue(true), |
| resolveToRealPath: vi.fn().mockImplementation((p) => p), |
| })); |
|
|
| vi.mock('../config/config.js', () => ({ |
| loadConfig: vi.fn().mockReturnValue({ |
| getSessionId: () => 'test-session', |
| getTargetDir: () => '/tmp', |
| getCheckpointingEnabled: () => false, |
| }), |
| loadEnvironment: vi.fn(), |
| setIsTrusted: vi.fn().mockReturnValue(false), |
| setTargetDir: vi.fn().mockReturnValue('/tmp'), |
| envStorage: { |
| run: (env: Record<string, string>, cb: () => unknown) => cb(), |
| }, |
| cwdSymbol: Symbol('cwd'), |
| })); |
|
|
| vi.mock('../config/settings.js', () => ({ |
| loadSettings: vi.fn().mockReturnValue({}), |
| })); |
|
|
| vi.mock('../config/extension.js', () => ({ |
| loadExtensions: vi.fn().mockReturnValue([]), |
| })); |
|
|
| vi.mock('../http/requestStorage.js', () => ({ |
| requestStorage: { |
| getStore: vi.fn(), |
| }, |
| })); |
|
|
| vi.mock('./task.js', () => { |
| const mockTaskInstance = (taskId: string, contextId: string) => ({ |
| id: taskId, |
| contextId, |
| taskState: 'working', |
| acceptUserMessage: vi |
| .fn() |
| .mockImplementation(async function* (context, aborted) { |
| const isConfirmation = ( |
| context.userMessage.parts as Array<{ kind: string }> |
| ).some((p) => p.kind === 'confirmation'); |
| |
| if (!isConfirmation && aborted) { |
| await new Promise((resolve) => { |
| aborted.addEventListener('abort', resolve, { once: true }); |
| }); |
| } |
| yield { type: 'content', value: 'hello' }; |
| }), |
| acceptAgentMessage: vi.fn().mockResolvedValue(undefined), |
| scheduleToolCalls: vi.fn().mockResolvedValue(undefined), |
| waitForPendingTools: vi.fn().mockResolvedValue(undefined), |
| getAndClearCompletedTools: vi.fn().mockReturnValue([]), |
| get hasPendingTools() { |
| return false; |
| }, |
| get pendingToolsCount() { |
| return 0; |
| }, |
| addToolResponsesToHistory: vi.fn(), |
| sendCompletedToolsToLlm: vi.fn().mockImplementation(async function* () {}), |
| cancelPendingTools: vi.fn(), |
| setTaskStateAndPublishUpdate: vi.fn(), |
| dispose: vi.fn(), |
| getMetadata: vi.fn().mockResolvedValue({}), |
| geminiClient: { |
| initialize: vi.fn().mockResolvedValue(undefined), |
| }, |
| toSDKTask: () => ({ |
| id: taskId, |
| contextId, |
| kind: 'task', |
| status: { state: 'working', timestamp: new Date().toISOString() }, |
| metadata: {}, |
| history: [], |
| artifacts: [], |
| }), |
| }); |
|
|
| const MockTask = vi.fn().mockImplementation(mockTaskInstance); |
| (MockTask as unknown as { create: Mock }).create = vi |
| .fn() |
| .mockImplementation(async (taskId: string, contextId: string) => |
| mockTaskInstance(taskId, contextId), |
| ); |
|
|
| return { Task: MockTask }; |
| }); |
|
|
| describe('CoderAgentExecutor', () => { |
| let executor: CoderAgentExecutor; |
| let mockTaskStore: TaskStore; |
| let mockEventBus: ExecutionEventBus; |
|
|
| beforeEach(() => { |
| vi.clearAllMocks(); |
| mockTaskStore = { |
| save: vi.fn().mockResolvedValue(undefined), |
| load: vi.fn().mockResolvedValue(undefined), |
| delete: vi.fn().mockResolvedValue(undefined), |
| list: vi.fn().mockResolvedValue([]), |
| } as unknown as TaskStore; |
|
|
| mockEventBus = new EventEmitter() as unknown as ExecutionEventBus; |
| mockEventBus.publish = vi.fn(); |
| mockEventBus.finished = vi.fn(); |
|
|
| executor = new CoderAgentExecutor(mockTaskStore); |
| }); |
|
|
| it('should distinguish between primary and secondary execution', async () => { |
| const taskId = 'test-task'; |
| const contextId = 'test-context'; |
|
|
| const mockSocket = new EventEmitter(); |
| const requestContext = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'hi' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket }, |
| }); |
|
|
| |
| const primaryPromise = executor.execute(requestContext, mockEventBus); |
|
|
| |
| await new Promise((resolve) => setTimeout(resolve, 50)); |
|
|
| expect( |
| ( |
| executor as unknown as { executingTasks: Set<string> } |
| ).executingTasks.has(taskId), |
| ).toBe(true); |
| const wrapper = executor.getTask(taskId); |
| expect(wrapper).toBeDefined(); |
|
|
| |
| const secondarySocket = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: secondarySocket }, |
| }); |
|
|
| const secondaryRequestContext = { |
| userMessage: { |
| messageId: 'msg-2', |
| taskId, |
| contextId, |
| parts: [{ kind: 'confirmation', callId: '1', outcome: 'proceed' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| const secondaryPromise = executor.execute( |
| secondaryRequestContext, |
| mockEventBus, |
| ); |
|
|
| |
| |
| await secondaryPromise; |
|
|
| |
| expect( |
| ( |
| executor as unknown as { executingTasks: Set<string> } |
| ).executingTasks.has(taskId), |
| ).toBe(true); |
| expect(wrapper?.task.dispose).not.toHaveBeenCalled(); |
|
|
| |
| secondarySocket.emit('end'); |
| expect( |
| ( |
| executor as unknown as { executingTasks: Set<string> } |
| ).executingTasks.has(taskId), |
| ).toBe(true); |
| expect(wrapper?.task.dispose).not.toHaveBeenCalled(); |
|
|
| |
| wrapper!.task.taskState = 'completed'; |
|
|
| |
| mockSocket.emit('end'); |
|
|
| await primaryPromise; |
|
|
| expect( |
| ( |
| executor as unknown as { executingTasks: Set<string> } |
| ).executingTasks.has(taskId), |
| ).toBe(false); |
| expect(wrapper?.task.dispose).toHaveBeenCalled(); |
| }); |
|
|
| it('should evict task from cache when it reaches terminal state', async () => { |
| const taskId = 'test-task-terminal'; |
| const contextId = 'test-context'; |
|
|
| const mockSocket = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket }, |
| }); |
|
|
| const requestContext = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'hi' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| const primaryPromise = executor.execute(requestContext, mockEventBus); |
| await new Promise((resolve) => setTimeout(resolve, 50)); |
|
|
| const wrapper = executor.getTask(taskId)!; |
| expect(wrapper).toBeDefined(); |
| |
| wrapper.task.taskState = 'completed'; |
|
|
| |
| mockSocket.emit('end'); |
| await primaryPromise; |
|
|
| expect(executor.getTask(taskId)).toBeUndefined(); |
| expect(wrapper.task.dispose).toHaveBeenCalled(); |
| }); |
|
|
| it('should yield the turn and transition to input-required if tools are pending', async () => { |
| const taskId = 'test-task-pending-tools'; |
| const contextId = 'test-context'; |
|
|
| const mockSocket = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket }, |
| }); |
|
|
| |
| const wrapper = await executor.createTask( |
| taskId, |
| contextId, |
| undefined, |
| mockEventBus, |
| ); |
| const hasPendingToolsSpy = vi |
| .spyOn(wrapper.task, 'hasPendingTools', 'get') |
| .mockReturnValue(true); |
| vi.spyOn(wrapper.task, 'pendingToolsCount', 'get').mockReturnValue(1); |
|
|
| const requestContext = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'confirmation', callId: '1', outcome: 'proceed' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| await executor.execute(requestContext, mockEventBus); |
|
|
| |
| expect(hasPendingToolsSpy).toHaveBeenCalled(); |
| expect(wrapper.task.getAndClearCompletedTools).not.toHaveBeenCalled(); |
| expect(wrapper.task.sendCompletedToolsToLlm).not.toHaveBeenCalled(); |
| expect(wrapper.task.setTaskStateAndPublishUpdate).toHaveBeenCalledWith( |
| 'input-required', |
| expect.any(Object), |
| undefined, |
| undefined, |
| true, |
| ); |
| }); |
|
|
| it('cancelTask should abort the active execution loop', async () => { |
| const abortSpy = vi.spyOn(AbortController.prototype, 'abort'); |
| const taskId = 'test-task-to-cancel'; |
| const contextId = 'test-context'; |
|
|
| const mockSocket = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket }, |
| }); |
|
|
| const requestContext = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'a long running prompt' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| |
| let primaryError: Error | null = null; |
| const primaryPromise = executor.execute(requestContext, mockEventBus); |
| primaryPromise.catch((err) => { |
| primaryError = err as Error; |
| }); |
|
|
| |
| let attempts = 0; |
| while (!executor.getTask(taskId)) { |
| if (primaryError) { |
| throw new Error(`Primary execution failed early: ${primaryError}`); |
| } |
| if (attempts++ > 100) { |
| |
| throw new Error('Timed out waiting for task to be registered'); |
| } |
| await new Promise((resolve) => setTimeout(resolve, 5)); |
| } |
|
|
| const wrapper = executor.getTask(taskId); |
| expect(wrapper).toBeDefined(); |
| const setTaskStateSpy = vi |
| .spyOn(wrapper!.task, 'setTaskStateAndPublishUpdate') |
| .mockImplementation((newState) => { |
| |
| wrapper!.task.taskState = newState; |
| }); |
|
|
| |
| await executor.cancelTask(taskId, mockEventBus); |
|
|
| |
| expect(abortSpy).toHaveBeenCalledOnce(); |
| expect(setTaskStateSpy).toHaveBeenCalledWith( |
| 'canceled', |
| expect.any(Object), |
| 'Task canceled by user request.', |
| undefined, |
| true, |
| ); |
|
|
| |
| |
| await primaryPromise; |
|
|
| |
| expect(executor.getTask(taskId)).toBeUndefined(); |
|
|
| abortSpy.mockRestore(); |
| }); |
|
|
| it('cancelTask should explicitly save task state to TaskStore and evict task during active aborts', async () => { |
| const taskId = 'test-task-active-abort-save'; |
| const contextId = 'test-context'; |
|
|
| const mockSocket = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket }, |
| }); |
|
|
| const requestContext = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'a long running prompt' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| const primaryPromise = executor.execute(requestContext, mockEventBus); |
|
|
| |
| let attempts = 0; |
| while (!executor.getTask(taskId)) { |
| if (attempts++ > 100) { |
| throw new Error('Timed out waiting for task to be registered'); |
| } |
| await new Promise((resolve) => setTimeout(resolve, 5)); |
| } |
|
|
| const wrapper = executor.getTask(taskId)!; |
| const saveSpy = vi.spyOn(mockTaskStore, 'save'); |
|
|
| |
| await executor.cancelTask(taskId, mockEventBus); |
|
|
| |
| expect(saveSpy).toHaveBeenCalled(); |
| expect(wrapper.task.dispose).toHaveBeenCalled(); |
| expect(executor.getTask(taskId)).toBeUndefined(); |
|
|
| |
| await primaryPromise; |
| }); |
|
|
| it('should allow executing a task that is in a terminal state by re-activating it', async () => { |
| const taskId = 'test-task-terminal-reactivate'; |
| const contextId = 'test-context'; |
|
|
| const mockSocket = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket }, |
| }); |
|
|
| const requestContext = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'hi' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| |
| const wrapper = await executor.createTask( |
| taskId, |
| contextId, |
| undefined, |
| mockEventBus, |
| ); |
| wrapper.task.taskState = 'canceled'; |
|
|
| |
| const primaryPromise = executor.execute(requestContext, mockEventBus); |
|
|
| |
| await new Promise((resolve) => setTimeout(resolve, 10)); |
|
|
| |
| const runningWrapper = executor.getTask(taskId); |
| expect(runningWrapper).toBeDefined(); |
| expect(runningWrapper!.task.taskState).not.toBe('canceled'); |
|
|
| |
| mockSocket.emit('end'); |
| await primaryPromise; |
| }); |
|
|
| it('should not evict or treat a new request as secondary when preceding request is canceled and winding down', async () => { |
| const taskId = 'test-repro-race-condition'; |
| const contextId = 'test-context'; |
|
|
| |
| vi.spyOn(mockTaskStore, 'save').mockImplementation(async () => { |
| await new Promise((resolve) => setTimeout(resolve, 50)); |
| }); |
|
|
| const mockSocket1 = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket1 }, |
| }); |
|
|
| const requestContext1 = { |
| userMessage: { |
| messageId: 'msg-1', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'prompt 1' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| const primaryPromise = executor.execute(requestContext1, mockEventBus); |
|
|
| |
| let attempts = 0; |
| while (!executor.getTask(taskId)) { |
| if (attempts++ > 100) { |
| throw new Error('Timed out waiting for task to be registered'); |
| } |
| await new Promise((resolve) => setTimeout(resolve, 5)); |
| } |
|
|
| const wrapper1 = executor.getTask(taskId)!; |
| expect(wrapper1).toBeDefined(); |
|
|
| |
| vi.spyOn(wrapper1.task, 'setTaskStateAndPublishUpdate').mockImplementation( |
| (newState) => { |
| wrapper1.task.taskState = newState; |
| }, |
| ); |
|
|
| |
| await executor.cancelTask(taskId, mockEventBus); |
|
|
| |
| const mockSocket2 = new EventEmitter(); |
| (requestStorage.getStore as Mock).mockReturnValue({ |
| req: { socket: mockSocket2 }, |
| }); |
|
|
| const requestContext2 = { |
| userMessage: { |
| messageId: 'msg-2', |
| taskId, |
| contextId, |
| parts: [{ kind: 'text', text: 'prompt 2' }], |
| metadata: { |
| coderAgent: { kind: 'agent-settings', workspacePath: '/tmp' }, |
| }, |
| }, |
| } as unknown as RequestContext; |
|
|
| const secondaryPromise = executor.execute(requestContext2, mockEventBus); |
|
|
| let secondaryResolved = false; |
| void secondaryPromise |
| .then(() => { |
| secondaryResolved = true; |
| }) |
| .catch(() => {}); |
|
|
| |
| await new Promise((resolve) => setTimeout(resolve, 20)); |
|
|
| |
| expect(secondaryResolved).toBe(false); |
|
|
| |
| await primaryPromise; |
|
|
| |
| const wrapper2 = executor.getTask(taskId); |
| expect(wrapper2).toBeDefined(); |
| expect(wrapper2).not.toBe(wrapper1); |
|
|
| |
| mockSocket2.emit('end'); |
| await secondaryPromise; |
| }); |
| }); |
|
|