| |
| |
| |
| |
| |
|
|
| import * as path from 'node:path'; |
| import type { |
| Task as SDKTask, |
| TaskStatusUpdateEvent, |
| SendStreamingMessageSuccessResponse, |
| } from '@a2a-js/sdk'; |
| import { |
| ApprovalMode, |
| DEFAULT_GEMINI_MODEL, |
| DEFAULT_TRUNCATE_TOOL_OUTPUT_THRESHOLD, |
| GeminiClient, |
| HookSystem, |
| type MessageBus, |
| PolicyDecision, |
| tmpdir, |
| type Config, |
| type Storage, |
| NoopSandboxManager, |
| type ToolRegistry, |
| type SandboxManager, |
| } from '@google/gemini-cli-core'; |
| import { createMockMessageBus } from '@google/gemini-cli-core/src/test-utils/mock-message-bus.js'; |
| import { expect, vi } from 'vitest'; |
|
|
| export function createMockConfig( |
| overrides: Partial<Config> = {}, |
| ): Partial<Config> { |
| const tmpDir = tmpdir(); |
| |
| const mockConfig = { |
| get config() { |
| |
| return this as unknown as Config; |
| }, |
| get toolRegistry() { |
| |
| const config = this as unknown as Config; |
| |
| return config.getToolRegistry?.() as unknown as ToolRegistry; |
| }, |
| get messageBus() { |
| return ( |
| |
| (this as unknown as Config).getMessageBus?.() as unknown as MessageBus |
| ); |
| }, |
| get geminiClient() { |
| |
| const config = this as unknown as Config; |
| |
| return config.getGeminiClient?.() as unknown as GeminiClient; |
| }, |
| getToolRegistry: vi.fn().mockReturnValue({ |
| getTool: vi.fn(), |
| getAllToolNames: vi.fn().mockReturnValue([]), |
| getAllTools: vi.fn().mockReturnValue([]), |
| getToolsByServer: vi.fn().mockReturnValue([]), |
| }), |
| getApprovalMode: vi.fn().mockReturnValue(ApprovalMode.DEFAULT), |
| getIdeMode: vi.fn().mockReturnValue(false), |
| isInteractive: () => true, |
| getAllowedTools: vi.fn().mockReturnValue([]), |
| getWorkspaceContext: vi.fn().mockReturnValue({ |
| isPathWithinWorkspace: () => true, |
| }), |
| getTargetDir: () => tmpDir, |
| getCheckpointingEnabled: vi.fn().mockReturnValue(false), |
| |
| storage: { |
| getProjectTempDir: () => tmpDir, |
| getProjectTempCheckpointsDir: () => path.join(tmpDir, 'checkpoints'), |
| } as Storage, |
| getTruncateToolOutputThreshold: () => |
| DEFAULT_TRUNCATE_TOOL_OUTPUT_THRESHOLD, |
| getActiveModel: vi.fn().mockReturnValue(DEFAULT_GEMINI_MODEL), |
| getDebugMode: vi.fn().mockReturnValue(false), |
| getContentGeneratorConfig: vi.fn().mockReturnValue({ model: 'gemini-pro' }), |
| getModel: vi.fn().mockReturnValue('gemini-pro'), |
| getUsageStatisticsEnabled: vi.fn().mockReturnValue(false), |
| setFallbackModelHandler: vi.fn(), |
| initialize: vi.fn().mockResolvedValue(undefined), |
| getProxy: vi.fn().mockReturnValue(undefined), |
| getHistory: vi.fn().mockReturnValue([]), |
| getEmbeddingModel: vi.fn().mockReturnValue('text-embedding-004'), |
| getSessionId: vi.fn().mockReturnValue('test-session-id'), |
| getUserTier: vi.fn(), |
| getMessageBus: vi.fn(), |
| getPolicyEngine: vi.fn(), |
| getEnableExtensionReloading: vi.fn().mockReturnValue(false), |
| getEnableHooks: vi.fn().mockReturnValue(false), |
| getMcpClientManager: vi.fn().mockReturnValue({ |
| getMcpServers: vi.fn().mockReturnValue({}), |
| }), |
| getTelemetryLogPromptsEnabled: vi.fn().mockReturnValue(false), |
| getTelemetryTracesEnabled: vi.fn().mockReturnValue(false), |
| getGitService: vi.fn(), |
| validatePathAccess: vi.fn().mockReturnValue(undefined), |
| getShellExecutionConfig: vi.fn().mockReturnValue({ |
| |
| sandboxManager: new NoopSandboxManager() as unknown as SandboxManager, |
| sanitizationConfig: { |
| allowedEnvironmentVariables: [], |
| blockedEnvironmentVariables: [], |
| enableEnvironmentVariableRedaction: false, |
| }, |
| }), |
| isContextManagementEnabled: vi.fn().mockReturnValue(false), |
| getContextManagementConfig: vi.fn().mockReturnValue({ enabled: false }), |
| getExperimentalGemma: vi.fn().mockReturnValue(false), |
| ...overrides, |
| } as unknown as Config; |
|
|
| |
| (mockConfig as unknown as { config: Config; promptId: string }).promptId = |
| 'test-prompt-id'; |
|
|
| mockConfig.getMessageBus = vi.fn().mockReturnValue(createMockMessageBus()); |
| mockConfig.getHookSystem = vi |
| .fn() |
| .mockReturnValue(new HookSystem(mockConfig)); |
|
|
| mockConfig.getGeminiClient = vi |
| .fn() |
| .mockReturnValue(new GeminiClient(mockConfig)); |
|
|
| mockConfig.getPolicyEngine = vi.fn().mockReturnValue({ |
| check: async () => { |
| const mode = mockConfig.getApprovalMode(); |
| if (mode === ApprovalMode.YOLO) { |
| return { decision: PolicyDecision.ALLOW }; |
| } |
| return { decision: PolicyDecision.ASK_USER }; |
| }, |
| }); |
|
|
| return mockConfig; |
| } |
|
|
| export function createStreamMessageRequest( |
| text: string, |
| messageId: string, |
| taskId?: string, |
| ) { |
| const request: { |
| jsonrpc: string; |
| id: string; |
| method: string; |
| params: { |
| message: { |
| kind: string; |
| role: string; |
| parts: [{ kind: string; text: string }]; |
| messageId: string; |
| }; |
| metadata: { |
| coderAgent: { |
| kind: string; |
| workspacePath: string; |
| }; |
| }; |
| taskId?: string; |
| }; |
| } = { |
| jsonrpc: '2.0', |
| id: '1', |
| method: 'message/stream', |
| params: { |
| message: { |
| kind: 'message', |
| role: 'user', |
| parts: [{ kind: 'text', text }], |
| messageId, |
| }, |
| metadata: { |
| coderAgent: { |
| kind: 'agent-settings', |
| workspacePath: '/tmp', |
| }, |
| }, |
| }, |
| }; |
|
|
| if (taskId) { |
| request.params.taskId = taskId; |
| } |
|
|
| return request; |
| } |
|
|
| export function assertUniqueFinalEventIsLast( |
| events: SendStreamingMessageSuccessResponse[], |
| ) { |
| |
| |
| const finalEvent = events[events.length - 1].result as TaskStatusUpdateEvent; |
| expect(finalEvent.metadata?.['coderAgent']).toMatchObject({ |
| kind: 'state-change', |
| }); |
| expect(finalEvent.status?.state).toBe('input-required'); |
| expect(finalEvent.final).toBe(true); |
|
|
| |
| expect( |
| |
| events.filter((e) => (e.result as TaskStatusUpdateEvent).final).length, |
| ).toBe(1); |
| expect( |
| |
| events.findIndex((e) => (e.result as TaskStatusUpdateEvent).final), |
| ).toBe(events.length - 1); |
| } |
|
|
| export function assertTaskCreationAndWorkingStatus( |
| events: SendStreamingMessageSuccessResponse[], |
| ) { |
| |
| |
| const taskEvent = events[0].result as SDKTask; |
| expect(taskEvent.kind).toBe('task'); |
| expect(taskEvent.status.state).toBe('submitted'); |
|
|
| |
| |
| const workingEvent = events[1].result as TaskStatusUpdateEvent; |
| expect(workingEvent.kind).toBe('status-update'); |
| expect(workingEvent.status.state).toBe('working'); |
| } |
|
|