P1 修复面收口: - 超时三态区分(aborted→USER_INTERRUPT / ETIMEDOUT→TIMEOUT / 其余→ERROR), 根治"真实网络超时被误报为用户中断" - 流空闲超时统一(SSE/Ollama/Anthropic 读循环 60s 无数据抛 504 进重试通道) - 同会话并发 sendMessage 防重入(isRunning 守卫)+ 会话存在性预检 + 前置调用移入 try(ERROR+DONE 双事件保证,根治 isStreaming 假死) - 清空审计后 resetChainCache(根治 verifyChain 误报 TAMPERED) - DONE 不再提前清理 TRACE(TERMINATED 统一收尾,补全最终迭代录制) - IME 合成回车不发送(普通 Enter + Cmd/Ctrl+Enter 双分支)+ handleSend 闭包修复 P2 安全纵深: - preload 移除原始 electronAPI 暴露(渲染层零使用,关掉 XSS invoke 任意通道单点风险) - CORS 同源回显根治(仅当前浏览页面 Origin,did-navigate 同步) - MEMORY.md 命令保护正则扩展(括号/$/反引号/< 重定向边界 + 前导路径) - write_file append TOCTOU 统一(open 后 realpath 校验,新文件分支补漏) - 敏感键归一化(authKey 驼峰/连字符命中)+ MCP headers 鉴权值加密落库 - ReDoS 检测共享化(search_files/file_editor 统一拦截) - run_tests/lint_code 升风险 + 需确认 + npx --no-install(执行边界对齐 run_command) - MCP/SearXNG/llm.baseURL/updateFeedUrl 配置类 URL 高危目标校验(IPv6 去括号 + 十六进制映射解析 + 尾点剥离) P3 架构还债: - temperature/maxTokens 热生效(引擎/编排器/SubAgent 三处接线)+ setBatch 单事务落盘 - SessionRecorder flush 竞态根治(flushPromise 等待 + 超限内联落盘 + stopRecording async) - 内存收口(lastConsolidationBySession LRU / subTraces 清理 / 会话删除 disposeEngine) - i18n 全量收口(28 组件 + 353 key 双字典,状态标签改渲染时函数) - 死代码清理(updateTraceStep/HEADER_HEIGHT/void preA/失实注释) - 斜杠菜单 MUI 化 + 删除逻辑收敛 resetSessionState + Blob URL 统一释放 + 用户消息"仅保存"落库(saveMessage 透传前端 id 修复 id 错位) P4 能力演进: - 死循环检测拆分(驻留前置 + 乒乓后置带进度信号,合法交替不误报) - run-lock 30s 超时强制 abort(旧 run 卡死不无限排队) - RETRY 双通道 stream_reset(前端按 run 归属精确清空,根治重试文本重复) - FTS5 trigram 中文子串搜索(迁移 9 版本化 SCHEMA_VERSION=2,≤2 字符 LIKE 回退) - getContextWindow 兜底 1M→128K(未知模型防 413) 测试: - 855 → 2406 用例(+1551,2.8 倍):服务层 +325(含 MemoryManager 51 新用例)、 工具实体 +483、IPC/适配器 +390(含 OpenAI/Anthropic/Ollama 独立套件)、 纯函数表格化 +330;引入 jsdom + @testing-library(14 组件测试文件 249 用例) - 修复 R1(saveMessage id 透传)/ R2(stream_reset 精确归属)两个回归缺陷 - 遗留低危项清零:git-tools 顺序耦合 / web-fetch 真实时间退避 / slo 内存断言 / mcp-security 多余 skipIf / deepseek-balance 命名误导 / 组件 mock 注入脆弱性 版本: 0.7.4; README 同步(工具风险表/版本徽章); 依赖: 移除 @electron-toolkit/preload, 新增 jsdom/@testing-library(devDependencies 不打包) 回归: typecheck 双端 0 错误; ESLint 0/0; Electron ABI 全量 2406/2406 零跳过; 系统 Node 2110 通过 296 跳过(better-sqlite3 ABI)
716 lines
26 KiB
TypeScript
716 lines
26 KiB
TypeScript
/**
|
||
* 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<unknown> {
|
||
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<unknown>;
|
||
}
|
||
|
||
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<string, unknown>;
|
||
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');
|
||
});
|
||
|
||
it('headers 键超长(>128 字符)拒绝', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerMCPHandlers(ctx);
|
||
const result = (await getHandler('mcp:addServer')(null, {
|
||
name: 'x',
|
||
transport: 'streamable-http',
|
||
url: 'https://a.com/mcp',
|
||
headers: { ['X'.repeat(200)]: 'v' },
|
||
})) as { success: boolean };
|
||
expect(result.success).toBe(false);
|
||
});
|
||
|
||
it('headers 值为空字符串合法(允许空值);非字符串值拒绝', async () => {
|
||
const { ctx, addServer } = makeCtx();
|
||
registerMCPHandlers(ctx);
|
||
const handler = getHandler('mcp:addServer');
|
||
|
||
// 空字符串值合法
|
||
await handler(null, {
|
||
name: 'x',
|
||
transport: 'streamable-http',
|
||
url: 'https://a.com/mcp',
|
||
headers: { 'X-Empty': '' },
|
||
});
|
||
expect(addServer).toHaveBeenCalled();
|
||
});
|
||
|
||
it('transport=sse 时 url 必填且必须可解析', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerMCPHandlers(ctx);
|
||
const handler = getHandler('mcp:addServer');
|
||
|
||
expect(await handler(null, { name: 'x', transport: 'sse' })).toMatchObject({ success: false });
|
||
expect(await handler(null, { name: 'x', transport: 'sse', url: '::bad' })).toMatchObject({
|
||
success: false,
|
||
});
|
||
});
|
||
|
||
it('url 指向云元数据 169.254.169.254 → 拒绝(SSRF 配置面防护)', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerMCPHandlers(ctx);
|
||
const result = (await getHandler('mcp:addServer')(null, {
|
||
name: 'evil',
|
||
transport: 'streamable-http',
|
||
url: 'http://169.254.169.254/latest/meta-data',
|
||
})) as { success: boolean; error?: string };
|
||
expect(result.success).toBe(false);
|
||
expect(result.error).toContain('cloud metadata');
|
||
});
|
||
|
||
it('url 指向 metadata.google.internal → 拒绝', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerMCPHandlers(ctx);
|
||
const result = (await getHandler('mcp:addServer')(null, {
|
||
name: 'evil',
|
||
transport: 'sse',
|
||
url: 'http://metadata.google.internal/',
|
||
})) as { success: boolean };
|
||
expect(result.success).toBe(false);
|
||
});
|
||
|
||
it('url 指向本地回环 127.0.0.1 → 放行(本地 MCP 服务器合法用例)', async () => {
|
||
const { ctx, addServer } = makeCtx();
|
||
registerMCPHandlers(ctx);
|
||
const result = (await getHandler('mcp:addServer')(null, {
|
||
name: 'local',
|
||
transport: 'streamable-http',
|
||
url: 'http://127.0.0.1:8080/mcp',
|
||
})) as { success: boolean };
|
||
expect(result.success).toBe(true);
|
||
expect(addServer).toHaveBeenCalled();
|
||
});
|
||
|
||
it('mcp:listServers 透传 manager 状态', async () => {
|
||
const getServerStates = vi.fn(() => [{ name: 'srv', status: 'running' }]);
|
||
registerMCPHandlers({ mcpManager: { getServerStates } } as unknown as IPCContext);
|
||
const result = (await getHandler('mcp:listServers')(null)) as unknown[];
|
||
expect(result).toHaveLength(1);
|
||
});
|
||
|
||
it('mcp:toggleServer manager 抛错 → 返回错误', async () => {
|
||
const toggleServer = vi.fn(async () => {
|
||
throw new Error('toggle failed');
|
||
});
|
||
registerMCPHandlers({ mcpManager: { toggleServer } } as unknown as IPCContext);
|
||
const result = (await getHandler('mcp:toggleServer')(null, 'srv', true)) as {
|
||
success: boolean;
|
||
error?: string;
|
||
};
|
||
expect(result.success).toBe(false);
|
||
expect(result.error).toBe('toggle failed');
|
||
});
|
||
});
|
||
|
||
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<string, Mock>; 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();
|
||
});
|
||
|
||
it('tasks:list 带 sessionId 过滤 SQL', async () => {
|
||
const { ctx, prepare, all } = makeCtx();
|
||
registerTaskHandlers(ctx);
|
||
|
||
await getHandler('tasks:list')(null, 's1');
|
||
const sql = String(prepare.mock.calls[0][0]);
|
||
expect(sql).toContain('WHERE session_id = ?');
|
||
expect(all).toHaveBeenCalledWith('s1');
|
||
});
|
||
|
||
it('tasks:create 合法路径带 parentId → order_idx 按父子分组自增', async () => {
|
||
const { ctx, prepare, run } = makeCtx();
|
||
const get = vi.fn(() => ({ maxOrder: 2 }));
|
||
(prepare as unknown as Mock).mockImplementation(() => ({
|
||
run,
|
||
get,
|
||
all: vi.fn(() => []),
|
||
}));
|
||
registerTaskHandlers(ctx);
|
||
|
||
const result = (await getHandler('tasks:create')(null, {
|
||
sessionId: 's1',
|
||
title: '子任务',
|
||
parentId: 'task_p1',
|
||
priority: 'low',
|
||
})) as { success: boolean; id: string };
|
||
expect(result.success).toBe(true);
|
||
expect(result.id).toMatch(/^task_/);
|
||
// 同 session + parent 分组查询(参数经 get() 传入)
|
||
const selectSql = String(prepare.mock.calls[0][0]);
|
||
expect(selectSql).toContain('parent_id = ?');
|
||
expect(get).toHaveBeenCalledWith('s1', 'task_p1');
|
||
});
|
||
|
||
it('tasks:create 描述非字符串 → 缺省空串落库', async () => {
|
||
const { ctx, run } = makeCtx();
|
||
(ctx.sessionService as unknown as { getDB: () => unknown }).getDB = () => ({
|
||
prepare: vi.fn(() => ({
|
||
run,
|
||
get: vi.fn(() => ({ maxOrder: -1 })),
|
||
all: vi.fn(() => []),
|
||
})),
|
||
});
|
||
registerTaskHandlers(ctx);
|
||
|
||
const result = (await getHandler('tasks:create')(null, {
|
||
sessionId: 's1',
|
||
title: 't',
|
||
description: 123,
|
||
})) as { success: boolean };
|
||
expect(result.success).toBe(true);
|
||
});
|
||
|
||
it('tasks:create DB 抛错 → 返回错误信息', async () => {
|
||
const { ctx } = makeCtx();
|
||
(ctx.sessionService as unknown as { getDB: () => unknown }).getDB = () => ({
|
||
prepare: vi.fn(() => {
|
||
throw new Error('UNIQUE constraint failed');
|
||
}),
|
||
});
|
||
registerTaskHandlers(ctx);
|
||
|
||
const result = (await getHandler('tasks:create')(null, {
|
||
sessionId: 's1',
|
||
title: 't',
|
||
})) as { success: boolean; error?: string };
|
||
expect(result.success).toBe(false);
|
||
expect(result.error).toBe('UNIQUE constraint failed');
|
||
});
|
||
|
||
it('tasks:update 跨会话伪造 parent_id → WHERE 包含 session_id 保护(越权防护)', async () => {
|
||
const { ctx, prepare, run } = makeCtx();
|
||
prepare.mockImplementation(() => ({
|
||
run,
|
||
get: vi.fn(() => ({ id: 't1' })),
|
||
all: vi.fn(() => []),
|
||
}));
|
||
registerTaskHandlers(ctx);
|
||
|
||
// 伪造:用另一会话的 sessionId 修改 t1 —— WHERE 同时带 id 与 session_id
|
||
await getHandler('tasks:update')(null, 't1', { title: 'hack' }, 'sess_attacker');
|
||
const sql = String(
|
||
prepare.mock.calls.find((c) => String(c[0]).startsWith('UPDATE'))?.[0] ?? '',
|
||
);
|
||
expect(sql).toContain('WHERE id = ? AND session_id = ?');
|
||
expect(run).toHaveBeenCalled();
|
||
});
|
||
|
||
it('tasks:update 空更新(无 fields)→ 幂等 success 不触发 UPDATE', async () => {
|
||
const { ctx, prepare, run } = makeCtx();
|
||
prepare.mockClear();
|
||
registerTaskHandlers(ctx);
|
||
|
||
const result = await getHandler('tasks:update')(null, 't1', {}, 's');
|
||
expect(result).toEqual({ success: true });
|
||
expect(run).not.toHaveBeenCalled();
|
||
});
|
||
|
||
it('tasks:update 完成状态附带 completed_at 写入', 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' }, 's');
|
||
const sql = String(
|
||
prepare.mock.calls.find((c) => String(c[0]).startsWith('UPDATE'))?.[0] ?? '',
|
||
);
|
||
expect(sql).toContain('completed_at = ?');
|
||
});
|
||
|
||
it('tasks:update 非法 id/sessionId/updates 拒绝', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerTaskHandlers(ctx);
|
||
const handler = getHandler('tasks:update');
|
||
|
||
expect(await handler(null, '', { status: 'completed' }, 's')).toMatchObject({ success: false });
|
||
expect(await handler(null, 't1', { status: 'completed' }, '')).toMatchObject({
|
||
success: false,
|
||
});
|
||
expect(await handler(null, 't1', null, 's')).toMatchObject({ success: false });
|
||
expect(await handler(null, 't1', { assignedTo: 42 }, 's')).toMatchObject({ success: false });
|
||
});
|
||
|
||
it('tasks:delete 跨会话删除 → WHERE 保护(他会话任务不受影响)', async () => {
|
||
const { ctx, prepare, run } = makeCtx();
|
||
registerTaskHandlers(ctx);
|
||
|
||
await getHandler('tasks:delete')(null, 't1', 'sess_attacker');
|
||
const sql = String(prepare.mock.calls[0][0]);
|
||
expect(sql).toContain('WHERE id = ? AND session_id = ?');
|
||
expect(run).toHaveBeenCalledWith('t1', 'sess_attacker');
|
||
});
|
||
|
||
it('tasks:delete 非法 id / sessionId 拒绝', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerTaskHandlers(ctx);
|
||
const handler = getHandler('tasks:delete');
|
||
|
||
expect(await handler(null, '', 's')).toMatchObject({ success: false });
|
||
expect(await handler(null, 't1', 42)).toMatchObject({ success: false });
|
||
});
|
||
});
|
||
|
||
// ===== 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 });
|
||
});
|
||
|
||
it('memory:delete working 类型映射到 working_memories 表', async () => {
|
||
const { ctx, prepare } = makeCtx();
|
||
registerMemoryHandlers(ctx);
|
||
const handler = getHandler('memory:delete');
|
||
|
||
await handler(null, 'working', 'w1');
|
||
expect(String(prepare.mock.calls[0][0])).toContain('working_memories');
|
||
});
|
||
|
||
it('memory:delete 非法 id(非字符串/空)拒绝', async () => {
|
||
const { ctx } = makeCtx();
|
||
registerMemoryHandlers(ctx);
|
||
const handler = getHandler('memory:delete');
|
||
|
||
expect(await handler(null, 'episodic', '')).toMatchObject({ success: false });
|
||
expect(await handler(null, 'episodic', 42)).toMatchObject({ success: false });
|
||
expect(await handler(null, 123, 'm1')).toMatchObject({ success: false });
|
||
});
|
||
|
||
it('memory:listAll 合法 type 过滤:仅查对应表', async () => {
|
||
const { ctx, prepare } = makeCtx();
|
||
registerMemoryHandlers(ctx);
|
||
const handler = getHandler('memory:listAll');
|
||
|
||
await handler(null, { type: 'semantic', limit: 50 });
|
||
// 只执行 semantic 表的 SELECT
|
||
const sqls = prepare.mock.calls.map((c) => String(c[0]));
|
||
expect(sqls).toHaveLength(1);
|
||
expect(sqls[0]).toContain('semantic_memories');
|
||
});
|
||
|
||
it('memory:listAll 未指定 type → 查询全部三张表', async () => {
|
||
const { ctx, prepare } = makeCtx();
|
||
registerMemoryHandlers(ctx);
|
||
const handler = getHandler('memory:listAll');
|
||
|
||
await handler(null);
|
||
const sqls = prepare.mock.calls.map((c) => String(c[0]));
|
||
expect(sqls.some((s) => s.includes('episodic_memories'))).toBe(true);
|
||
expect(sqls.some((s) => s.includes('semantic_memories'))).toBe(true);
|
||
expect(sqls.some((s) => s.includes('working_memories'))).toBe(true);
|
||
});
|
||
|
||
it('memory:listAll limit 小数向下取整', async () => {
|
||
const { ctx, prepare } = makeCtx();
|
||
registerMemoryHandlers(ctx);
|
||
const handler = getHandler('memory:listAll');
|
||
|
||
await handler(null, { type: 'episodic', limit: 42.7 });
|
||
const all = prepare.mock.results[0].value.all as Mock;
|
||
expect(all).toHaveBeenCalledWith(42);
|
||
});
|
||
|
||
it('db:searchMemories options 非对象(null/字符串)→ 使用默认搜索选项', 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', null);
|
||
expect(search).toHaveBeenLastCalledWith('q', {});
|
||
await handler(null, 'q', 'not-an-object');
|
||
expect(search).toHaveBeenLastCalledWith('q', {});
|
||
});
|
||
|
||
it('db:searchMemories minImportance 数值透传(非有限值剔除)', 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', { minImportance: 0.5 });
|
||
expect(search).toHaveBeenLastCalledWith('q', expect.objectContaining({ minImportance: 0.5 }));
|
||
await handler(null, 'q', { minImportance: Number.NaN });
|
||
expect(search).toHaveBeenLastCalledWith(
|
||
'q',
|
||
expect.not.objectContaining({ minImportance: expect.anything() }),
|
||
);
|
||
});
|
||
|
||
it('db:searchMemories 空 query(仅空白)→ 返回 [] 不触达 search', 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([]);
|
||
expect(search).not.toHaveBeenCalled();
|
||
});
|
||
});
|