Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
163 changes: 163 additions & 0 deletions electron/main/ai/runtime-v2/agent-capabilities.test.mjs
Original file line number Diff line number Diff line change
@@ -0,0 +1,163 @@
import assert from 'node:assert/strict'
import test from 'node:test'
import { DatabaseSync } from 'node:sqlite'
import { readFileSync } from 'node:fs'
import { createServer } from 'node:http'

import { AgentMemoryStore } from './agent-memory-store.ts'
import { ControlledMcpStore } from './controlled-mcp-store.ts'
import { ControlledHttpMcpClient, validateControlledMcpUrl } from './controlled-mcp-client.ts'
import { createControlledMemoryTools } from './tools/memory-tools.ts'
import { createRuntimePlan } from './planner.ts'

test('运行规划默认注入当前项目的长期记忆上下文', () => {
const plan = createRuntimePlan({
surface: {
id: 'global-page',
title: '全局助手',
scope: 'project',
allowedTools: []
},
request: { userMessage: '继续分析人物弧光' }
})

assert.equal(plan.contextProviders.includes('agent-memories'), true)
})

test('长期记忆按项目隔离、去重并允许用户删除和调整重要度', () => {
const db = new DatabaseSync(':memory:')
const store = new AgentMemoryStore(db)
const first = store.create({ projectId: 'p1', content: '主角保持克制', kind: 'preference' })
const duplicate = store.create({ projectId: 'p1', content: '主角保持克制', kind: 'preference' })
store.create({ projectId: 'p2', content: '另一项目记忆', kind: 'fact' })

assert.equal(first.id, duplicate.id)
assert.deepEqual(store.list('p1').map((item) => item.content), ['主角保持克制'])
assert.equal(store.setImportance(first.id, 'p1', 9)?.importance, 5)
assert.equal(store.remove(first.id, 'p2'), false)
assert.equal(store.remove(first.id, 'p1'), true)
assert.equal(store.list('p1').length, 0)
db.close()
})

test('模型只有在用户明确授权时才能写入长期记忆', async () => {
const db = new DatabaseSync(':memory:')
const store = new AgentMemoryStore(db)
const denied = createControlledMemoryTools({
store,
projectId: 'p1',
turnId: 't1',
userMessage: '帮我分析人物设定'
})[0]
const allowed = createControlledMemoryTools({
store,
projectId: 'p1',
turnId: 't2',
userMessage: '请记住以后不要使用网络热梗'
})[0]

assert.equal((await denied.handler({ content: '擅自保存' }, {})).isError, true)
assert.equal(store.list('p1').length, 0)
assert.equal((await allowed.handler({ content: '不要使用网络热梗' }, {})).isError, undefined)
assert.equal(store.list('p1')[0].content, '不要使用网络热梗')
db.close()
})

test('MCP 配置不向渲染层暴露密钥,并强制使用已发现工具白名单', () => {
const db = new DatabaseSync(':memory:')
const store = new ControlledMcpStore(db)
const server = store.save({
projectId: 'p1',
name: '榜单服务',
url: 'https://example.com/mcp',
apiKey: 'secret-key'
})

assert.equal(server.hasApiKey, true)
assert.equal('apiKey' in server, false)
assert.throws(() => store.setEnabled(server.id, 'p1', true), /选择允许的工具/)

store.recordConnection(server.id, 'p1', [
{ name: 'rank_list', description: '读取榜单' },
{ name: 'book_detail', description: '读取书籍详情' }
])
const allowed = store.setAllowedTools(server.id, 'p1', ['rank_list', 'not-discovered'])
assert.deepEqual(allowed.allowedTools, ['rank_list'])
assert.equal(store.setEnabled(server.id, 'p1', true).enabled, true)
assert.equal(store.getSecret(server.id, 'p1')?.apiKey, 'secret-key')
assert.equal(store.getSecret(server.id, 'p2'), null)
db.close()
})

test('受控 MCP 只接受 HTTPS 或本机 HTTP', () => {
assert.equal(validateControlledMcpUrl('https://example.com/mcp'), 'https://example.com/mcp')
assert.equal(validateControlledMcpUrl('http://127.0.0.1:3000/mcp'), 'http://127.0.0.1:3000/mcp')
assert.equal(validateControlledMcpUrl('http://[::1]:3000/mcp'), 'http://[::1]:3000/mcp')
assert.throws(() => validateControlledMcpUrl('http://example.com/mcp'), /必须使用 HTTPS/)
assert.throws(() => validateControlledMcpUrl('file:///tmp/mcp'), /必须使用 HTTPS/)
})

test('受控 MCP 不跟随可能绕过 URL 限制的重定向', async () => {
const server = createServer((_request, response) => {
response.writeHead(302, { Location: 'http://example.com/mcp' }).end()
})
await new Promise((resolve) => server.listen(0, '127.0.0.1', resolve))
const address = server.address()
const client = new ControlledHttpMcpClient(`http://127.0.0.1:${address.port}/mcp`)
try {
await assert.rejects(() => client.listTools())
} finally {
await new Promise((resolve) => server.close(resolve))
}
})

test('受控 MCP 客户端完成握手、携带会话 ID 并调用工具', async () => {
const methods = []
const server = createServer(async (request, response) => {
let raw = ''
for await (const chunk of request) raw += chunk
const payload = JSON.parse(raw)
methods.push(payload.method)

if (payload.method === 'notifications/initialized') {
response.writeHead(202).end()
return
}
if (payload.method === 'initialize') {
response.setHeader('Mcp-Session-Id', 'session-test')
response.setHeader('Content-Type', 'application/json')
response.end(JSON.stringify({ jsonrpc: '2.0', id: payload.id, result: { protocolVersion: '2024-11-05' } }))
return
}
assert.equal(request.headers['mcp-session-id'], 'session-test')
response.setHeader('Content-Type', 'application/json')
if (payload.method === 'tools/list') {
response.end(JSON.stringify({ jsonrpc: '2.0', id: payload.id, result: { tools: [{ name: 'rank_list' }] } }))
return
}
response.end(JSON.stringify({
jsonrpc: '2.0',
id: payload.id,
result: { content: [{ type: 'text', text: '榜单结果' }] }
}))
})
await new Promise((resolve) => server.listen(0, '127.0.0.1', resolve))
const address = server.address()
const client = new ControlledHttpMcpClient(`http://127.0.0.1:${address.port}/mcp`)
try {
assert.deepEqual((await client.listTools()).map((tool) => tool.name), ['rank_list'])
assert.equal(await client.callTool('rank_list', {}), '榜单结果')
assert.deepEqual(methods, ['initialize', 'notifications/initialized', 'tools/list', 'tools/call'])
} finally {
await new Promise((resolve) => server.close(resolve))
}
})

test('子智能体工具限制任务数、并发和只读材料', () => {
const source = readFileSync(new URL('./tools/delegate-novel-tools.ts', import.meta.url), 'utf8')
assert.match(source, /const MAX_TASKS = 3/)
assert.match(source, /const MAX_CONCURRENCY = 2/)
assert.match(source, /tasks\.length < 2/)
assert.match(source, /task\.description && task\.material/)
assert.match(source, /不能调用工具、不能修改项目、不能形成长期记忆/)
})
196 changes: 196 additions & 0 deletions electron/main/ai/runtime-v2/agent-memory-store.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,196 @@
import { randomUUID } from 'node:crypto'
import type { DatabaseSync, StatementSync } from 'node:sqlite'
import type { AgentMemory, AgentMemoryKind } from '@shared/assistant-runtime'

const MAX_MEMORY_CONTENT = 1200
const MAX_MEMORIES_PER_PROJECT = 200
const VALID_KINDS = new Set<AgentMemoryKind>(['preference', 'lesson', 'fact', 'method'])

export interface AgentMemoryInput {
projectId: string
kind?: AgentMemoryKind
content: string
source?: AgentMemory['source']
importance?: number
sourceTurnId?: string
}

interface MemoryRow {
id: string
project_id: string
kind: string
content: string
source: string
importance: number
source_turn_id: string
created_at: string
updated_at: string
}

function normalizeImportance(value: unknown): number {
const number = Number(value)
return Number.isFinite(number) ? Math.min(5, Math.max(1, Math.round(number))) : 3
}

function rowToMemory(row: MemoryRow): AgentMemory {
return {
id: row.id,
projectId: row.project_id,
kind: VALID_KINDS.has(row.kind as AgentMemoryKind) ? row.kind as AgentMemoryKind : 'preference',
content: row.content,
source: row.source === 'agent' || row.source === 'system' ? row.source : 'user',
importance: row.importance,
sourceTurnId: row.source_turn_id || undefined,
createdAt: row.created_at,
updatedAt: row.updated_at
}
}

export function initAgentMemorySchema(db: DatabaseSync): void {
db.exec(`
CREATE TABLE IF NOT EXISTS assistant_memories (
id TEXT PRIMARY KEY,
project_id TEXT NOT NULL,
kind TEXT NOT NULL DEFAULT 'preference',
content TEXT NOT NULL,
source TEXT NOT NULL DEFAULT 'user',
importance INTEGER NOT NULL DEFAULT 3,
source_turn_id TEXT NOT NULL DEFAULT '',
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
) STRICT;

CREATE INDEX IF NOT EXISTS idx_assistant_memories_project
ON assistant_memories (project_id, importance DESC, updated_at DESC);
`)
}

export class AgentMemoryStore {
private readonly stmts: {
insert: StatementSync
findDuplicate: StatementSync
get: StatementSync
list: StatementSync
remove: StatementSync
updateImportance: StatementSync
count: StatementSync
pruneOne: StatementSync
}

constructor(db: DatabaseSync) {
initAgentMemorySchema(db)
this.stmts = {
insert: db.prepare(`
INSERT INTO assistant_memories
(id, project_id, kind, content, source, importance, source_turn_id, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
`),
findDuplicate: db.prepare(`
SELECT * FROM assistant_memories WHERE project_id = ? AND content = ? LIMIT 1
`),
get: db.prepare(`SELECT * FROM assistant_memories WHERE id = ? AND project_id = ?`),
list: db.prepare(`
SELECT * FROM assistant_memories
WHERE project_id = ?
ORDER BY importance DESC, updated_at DESC
LIMIT ?
`),
remove: db.prepare(`DELETE FROM assistant_memories WHERE id = ? AND project_id = ?`),
updateImportance: db.prepare(`
UPDATE assistant_memories SET importance = ?, updated_at = ?
WHERE id = ? AND project_id = ?
`),
count: db.prepare(`SELECT COUNT(*) AS count FROM assistant_memories WHERE project_id = ?`),
pruneOne: db.prepare(`
DELETE FROM assistant_memories WHERE id = (
SELECT id FROM assistant_memories WHERE project_id = ?
ORDER BY importance ASC, updated_at ASC LIMIT 1
)
`)
}
}

create(input: AgentMemoryInput): AgentMemory {
const projectId = String(input.projectId || '').trim()
const content = String(input.content || '').replace(/\s+/g, ' ').trim().slice(0, MAX_MEMORY_CONTENT)
if (!projectId) throw new Error('缺少项目 ID,无法保存创作记忆。')
if (!content) throw new Error('创作记忆内容不能为空。')

const duplicate = this.stmts.findDuplicate.get(projectId, content) as MemoryRow | undefined
if (duplicate) return rowToMemory(duplicate)

const now = new Date().toISOString()
const kind = VALID_KINDS.has(input.kind ?? 'preference') ? input.kind ?? 'preference' : 'preference'
const source = input.source === 'agent' || input.source === 'system' ? input.source : 'user'
const memory: AgentMemory = {
id: randomUUID(),
projectId,
kind,
content,
source,
importance: normalizeImportance(input.importance),
sourceTurnId: String(input.sourceTurnId || '').trim() || undefined,
createdAt: now,
updatedAt: now
}
this.stmts.insert.run(
memory.id,
memory.projectId,
memory.kind,
memory.content,
memory.source,
memory.importance,
memory.sourceTurnId ?? '',
memory.createdAt,
memory.updatedAt
)
this.prune(projectId)
return memory
}

list(projectId: string, limit = 50): AgentMemory[] {
const safeLimit = Math.min(100, Math.max(1, Math.round(Number(limit) || 50)))
const rows = this.stmts.list.all(String(projectId || '').trim(), safeLimit) as unknown as MemoryRow[]
return rows.map(rowToMemory)
}

remove(id: string, projectId: string): boolean {
return this.stmts.remove.run(String(id || ''), String(projectId || '')).changes > 0
}

setImportance(id: string, projectId: string, importance: number): AgentMemory | null {
this.stmts.updateImportance.run(
normalizeImportance(importance),
new Date().toISOString(),
String(id || ''),
String(projectId || '')
)
const row = this.stmts.get.get(String(id || ''), String(projectId || '')) as MemoryRow | undefined
return row ? rowToMemory(row) : null
}

private prune(projectId: string): void {
const row = this.stmts.count.get(projectId) as { count?: number } | undefined
let count = Number(row?.count ?? 0)
while (count > MAX_MEMORIES_PER_PROJECT) {
this.stmts.pruneOne.run(projectId)
count -= 1
}
}
}

export function formatAgentMemories(memories: AgentMemory[]): string {
if (!memories.length) return ''
const labels: Record<AgentMemoryKind, string> = {
preference: '偏好',
lesson: '教训',
fact: '事实',
method: '方法'
}
return [
'这些是用户可查看和删除的项目级长期记忆。除非用户本轮明确推翻,否则应遵守;若与当前项目事实冲突,先向用户说明。',
...memories.map((memory, index) =>
`${index + 1}. [${labels[memory.kind]}·重要度${memory.importance}] ${memory.content}`
)
].join('\n')
}
5 changes: 3 additions & 2 deletions electron/main/ai/runtime-v2/bootstrap.ts
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ import { registerAssistantIpcHandlers } from './ipc'
import { createExecutionPlanner } from './execution-plan'
import { createCommitter } from './committer'
import { registerBuiltinProviders } from './providers'
import { getSharedConversation, peekSharedConversation } from './state'
import { getAgentMemoryStore, getSharedConversation, peekSharedConversation } from './state'

export interface BootstrapAssistantRuntimeDeps {
/** 首次调用时 ensure schema,返回 DatabaseSync 实例。 */
Expand All @@ -39,7 +39,8 @@ export function bootstrapAssistantRuntime(deps: BootstrapAssistantRuntimeDeps):
registerBuiltinProviders({
contextBuilder,
snapshot: snapshotAccessor,
getConversation: () => getSharedConversation()
getConversation: () => getSharedConversation(),
getMemoryStore: () => getAgentMemoryStore()
})

// 2. 构造 execution planner
Expand Down
Loading
Loading