/** * IPC MCP / Tasks / Memory 域测试(v0.7.2 覆盖补齐) * * 锁定三类安全相关契约: * 1. mcp:addServer 的传输方式/必填字段/headers 逐项校验矩阵(v0.7.2 P2-8) * 2. tasks 域的枚举校验与会话越权防护(WHERE session_id = ? 契约) * 3. memory 域的查询参数收敛(topK 钳制/type 枚举)与删除表映射 */ import { describe, it, expect, vi, beforeEach, type Mock } from 'vitest'; const ipcMainHandleMock = vi.fn(); vi.mock('electron', () => ({ ipcMain: { handle: (...args: unknown[]) => ipcMainHandleMock(...args) }, })); vi.mock('electron-log', () => ({ default: { info: vi.fn(), warn: vi.fn(), error: vi.fn(), debug: vi.fn() }, })); import { registerMCPHandlers } from '../mcp'; import { registerTaskHandlers } from '../tasks'; import { registerMemoryHandlers } from '../memory'; import type { IPCContext } from '../context'; function getHandler(channel: string): (...args: unknown[]) => Promise { const call = ipcMainHandleMock.mock.calls.find(([ch]) => ch === channel); if (!call) throw new Error(`IPC handler not registered: ${channel}`); return call[1] as (...args: unknown[]) => Promise; } beforeEach(() => { ipcMainHandleMock.mockClear(); }); // ===== MCP 域 ===== describe('mcp:addServer — 校验矩阵', () => { function makeCtx(): { ctx: IPCContext; addServer: Mock } { const addServer = vi.fn(async () => undefined); return { ctx: { mcpManager: { addServer } } as unknown as IPCContext, addServer }; } it.each([ [null, 'Invalid config'], [undefined, 'Invalid config'], [{ transport: 'stdio', command: 'npx' }, 'name is required'], [{ name: 'x', transport: 'ftp', command: 'npx' }, 'Invalid transport'], [{ name: 'x', transport: 'stdio' }, 'command is required'], [{ name: 'x', transport: 'stdio', command: ' ' }, 'command is required'], [{ name: 'x', transport: 'streamable-http' }, 'url is required'], [{ name: 'x', transport: 'sse', url: 'not a url' }, 'Invalid url format'], ])('非法载荷 %# 拒绝', async (payload, expectedError) => { const { ctx } = makeCtx(); registerMCPHandlers(ctx); const result = (await getHandler('mcp:addServer')(null, payload)) as { success: boolean; error?: string; }; expect(result.success).toBe(false); expect(result.error).toContain(expectedError); }); it('stdio 合法载荷透传(不含 headers)', async () => { const { ctx, addServer } = makeCtx(); registerMCPHandlers(ctx); const result = await getHandler('mcp:addServer')(null, { name: 'fs-server', transport: 'stdio', command: 'npx', args: ['-y', 'server'], }); expect(result).toEqual({ success: true }); expect(addServer).toHaveBeenCalledWith( expect.objectContaining({ name: 'fs-server', transport: 'stdio', command: 'npx', enabled: true, }), ); }); it('v0.7.2 P2-8: headers 必须为扁平 string→string 对象', async () => { const { ctx } = makeCtx(); registerMCPHandlers(ctx); const handler = getHandler('mcp:addServer'); // 非对象 expect( await handler(null, { name: 'x', transport: 'streamable-http', url: 'https://a.com/mcp', headers: 'Bearer x', }), ).toMatchObject({ success: false }); // 数组 expect( await handler(null, { name: 'x', transport: 'streamable-http', url: 'https://a.com/mcp', headers: [], }), ).toMatchObject({ success: false }); // 值非字符串 expect( await handler(null, { name: 'x', transport: 'streamable-http', url: 'https://a.com/mcp', headers: { Authorization: 123 }, }), ).toMatchObject({ success: false }); // 空键 expect( await handler(null, { name: 'x', transport: 'streamable-http', url: 'https://a.com/mcp', headers: { '': 'v' }, }), ).toMatchObject({ success: false }); // 超过 20 项 const tooMany = Object.fromEntries(Array.from({ length: 21 }, (_, i) => [`h${i}`, 'v'])); expect( await handler(null, { name: 'x', transport: 'streamable-http', url: 'https://a.com/mcp', headers: tooMany, }), ).toMatchObject({ success: false }); // 超长值 expect( await handler(null, { name: 'x', transport: 'streamable-http', url: 'https://a.com/mcp', headers: { Authorization: 'x'.repeat(5000) }, }), ).toMatchObject({ success: false }); }); it('headers 合法时透传给 manager;未提供时不携带该字段', async () => { const { ctx, addServer } = makeCtx(); registerMCPHandlers(ctx); const handler = getHandler('mcp:addServer'); await handler(null, { name: 'remote', transport: 'streamable-http', url: 'https://a.com/mcp', headers: { Authorization: 'Bearer tok' }, }); expect(addServer).toHaveBeenLastCalledWith( expect.objectContaining({ headers: { Authorization: 'Bearer tok' } }), ); await handler(null, { name: 'remote2', transport: 'streamable-http', url: 'https://b.com/mcp', }); const call = addServer.mock.calls[1][0] as Record; expect(call.headers).toBeUndefined(); }); it('manager 抛错 → success:false + 错误信息', async () => { const { ctx, addServer } = makeCtx(); addServer.mockRejectedValueOnce(new Error('connect timeout')); registerMCPHandlers(ctx); const result = (await getHandler('mcp:addServer')(null, { name: 'x', transport: 'stdio', command: 'node', })) as { success: boolean; error?: string; }; expect(result.success).toBe(false); expect(result.error).toBe('connect timeout'); }); }); describe('mcp:removeServer / toggleServer — 校验', () => { it('非法 name 拒绝;合法调用透传', async () => { const removeServer = vi.fn(async () => undefined); const toggleServer = vi.fn(async () => undefined); registerMCPHandlers({ mcpManager: { removeServer, toggleServer } } as unknown as IPCContext); expect(await getHandler('mcp:removeServer')(null, '')).toMatchObject({ success: false }); expect(await getHandler('mcp:toggleServer')(null, 'x', 'yes')).toMatchObject({ success: false, }); expect(removeServer).not.toHaveBeenCalled(); expect(toggleServer).not.toHaveBeenCalled(); await getHandler('mcp:removeServer')(null, 'srv'); expect(removeServer).toHaveBeenCalledWith('srv'); await getHandler('mcp:toggleServer')(null, 'srv', false); expect(toggleServer).toHaveBeenCalledWith('srv', false); }); }); // ===== Tasks 域 ===== describe('tasks 域 — 会话越权防护与校验', () => { function makeDb(): { db: Record; prepare: Mock } { const run = vi.fn(); const get = vi.fn(); const all = vi.fn(() => []); const prepare = vi.fn(() => ({ run, get, all })); return { db: { run, get, all }, prepare }; } function makeCtx(): { ctx: IPCContext; prepare: Mock; run: Mock; get: Mock; all: Mock } { const { db, prepare } = makeDb(); const sessionService = { getDB: () => ({ prepare }) }; return { ctx: { sessionService } as unknown as IPCContext, prepare, run: db.run, get: db.get, all: db.all, }; } it('tasks:create 校验 sessionId/title/priority/parentId', async () => { const { ctx } = makeCtx(); registerTaskHandlers(ctx); const handler = getHandler('tasks:create'); expect(await handler(null, null)).toMatchObject({ success: false }); expect(await handler(null, { sessionId: '', title: 't' })).toMatchObject({ success: false }); expect(await handler(null, { sessionId: 's', title: ' ' })).toMatchObject({ success: false }); expect(await handler(null, { sessionId: 's', title: 't', priority: 'urgent' })).toMatchObject({ success: false, }); expect(await handler(null, { sessionId: 's', title: 't', parentId: 42 })).toMatchObject({ success: false, }); }); it('tasks:create 合法路径:order_idx = 同组 MAX+1,ID 统一 nanoid 前缀', async () => { const { ctx, prepare, run } = makeCtx(); (prepare as unknown as Mock).mockImplementation(() => ({ run, get: vi.fn(() => ({ maxOrder: 4 })), all: vi.fn(() => []), })); registerTaskHandlers(ctx); const result = (await getHandler('tasks:create')(null, { sessionId: 's1', title: '新任务', priority: 'high', parentId: null, })) as { success: boolean; id: string }; expect(result.success).toBe(true); expect(result.id).toMatch(/^task_/); }); it('tasks:update 强制 WHERE session_id = ?(越权防护契约)', async () => { const { ctx, prepare, run } = makeCtx(); prepare.mockImplementation(() => ({ run, get: vi.fn(() => ({ id: 't1' })), all: vi.fn(() => []), })); registerTaskHandlers(ctx); await getHandler('tasks:update')(null, 't1', { status: 'completed' }, 'sess-owner'); const sql = (prepare.mock.calls.find((c) => String(c[0]).startsWith('UPDATE'))?.[0] ?? '') as string; expect(sql).toContain('AND session_id = ?'); expect(run).toHaveBeenCalled(); }); it('tasks:update 枚举校验(status/priority/title/description 类型)', async () => { const { ctx } = makeCtx(); registerTaskHandlers(ctx); const handler = getHandler('tasks:update'); expect(await handler(null, 't1', { status: 'done' }, 's')).toMatchObject({ success: false }); expect(await handler(null, 't1', { priority: 'critical!' }, 's')).toMatchObject({ success: false, }); expect(await handler(null, 't1', { title: 123 }, 's')).toMatchObject({ success: false }); expect(await handler(null, 't1', { description: {} }, 's')).toMatchObject({ success: false }); // 空更新为幂等 no-op(既有契约:fields 为空直接 success,不触发 UPDATE) expect(await handler(null, 't1', {}, 's')).toMatchObject({ success: true }); expect(await handler(null, 't1', { status: 'completed' }, '')).toMatchObject({ success: false, }); // 越权防护 }); it('tasks:delete 强制 WHERE session_id = ?', async () => { const { ctx, prepare, run } = makeCtx(); registerTaskHandlers(ctx); await getHandler('tasks:delete')(null, 't1', 'sess-owner'); const sql = (prepare.mock.calls[0][0] ?? '') as string; expect(sql).toContain('AND session_id = ?'); expect(run).toHaveBeenCalled(); }); it('tasks:list 校验 sessionId 可选参数', async () => { const { ctx, all } = makeCtx(); registerTaskHandlers(ctx); expect(await getHandler('tasks:list')(null, 42)).toMatchObject({ success: false }); await getHandler('tasks:list')(null); expect(all).toHaveBeenCalled(); }); }); // ===== Memory 域 ===== describe('memory 域 — 查询参数收敛与删除映射', () => { function makeCtx(): { ctx: IPCContext; prepare: Mock } { const prepare = vi.fn(() => ({ run: vi.fn(), get: vi.fn(), all: vi.fn(() => []) })); return { ctx: { memoryManager: { search: vi.fn(() => []) }, sessionService: { getDB: () => ({ prepare }) }, } as unknown as IPCContext, prepare, }; } it('db:searchMemories 空 query 返回空数组;topK 钳制 1-100', async () => { const { ctx } = makeCtx(); registerMemoryHandlers(ctx); const handler = getHandler('db:searchMemories'); const search = (ctx.memoryManager as unknown as { search: Mock }).search; expect(await handler(null, ' ')).toEqual([]); await handler(null, 'q', { topK: 500 }); expect(search).toHaveBeenCalledWith('q', expect.objectContaining({ topK: 10 })); // 非法回退默认 10 await handler(null, 'q', { topK: 3 }); expect(search).toHaveBeenLastCalledWith('q', expect.objectContaining({ topK: 3 })); }); it('db:searchMemories type 必须为合法 MemoryType(非法被剔除)', async () => { const { ctx } = makeCtx(); registerMemoryHandlers(ctx); const handler = getHandler('db:searchMemories'); const search = (ctx.memoryManager as unknown as { search: Mock }).search; await handler(null, 'q', { type: 'semantic' }); expect(search).toHaveBeenLastCalledWith('q', expect.objectContaining({ type: 'semantic' })); await handler(null, 'q', { type: 'hacked' }); expect(search).toHaveBeenLastCalledWith( 'q', expect.not.objectContaining({ type: expect.anything() }), ); }); it('memory:listAll 校验 type 枚举与 limit 范围(LIMIT -1 防护)', async () => { const { ctx } = makeCtx(); registerMemoryHandlers(ctx); const handler = getHandler('memory:listAll'); expect(await handler(null, { type: 'nope' })).toMatchObject({ success: false }); expect(await handler(null, { limit: -1 })).toMatchObject({ success: false }); expect(await handler(null, { limit: 5000 })).toMatchObject({ success: false }); }); it('memory:delete 按类型映射到正确的表(防止三元默认落到 working_memories)', async () => { const { ctx, prepare } = makeCtx(); registerMemoryHandlers(ctx); const handler = getHandler('memory:delete'); await handler(null, 'episodic', 'm1'); expect(String(prepare.mock.calls[0][0])).toContain('episodic_memories'); prepare.mockClear(); await handler(null, 'semantic', 'm2'); expect(String(prepare.mock.calls[0][0])).toContain('semantic_memories'); prepare.mockClear(); expect(await handler(null, 'unknown', 'm3')).toMatchObject({ success: false }); }); });