diff --git a/docs/agent/query-agent-poc.md b/docs/agent/query-agent-poc.md new file mode 100644 index 0000000..accbda5 --- /dev/null +++ b/docs/agent/query-agent-poc.md @@ -0,0 +1,27 @@ +# Query Agent POC + +这是独立的开发测试入口,不会修改生产“问问微信”执行链。 + +先启动 TraceMemo,并在 API Center 开启 Local HTTP API。然后在仓库根目录运行 例: + +```bash +pnpm poc:query-agent -- "我和BOBO第一次聊了什么" +``` + +也可以直接传入其他自然语言问题: + +```bash +pnpm poc:query-agent -- "BOBO上个月有没有给我发过文件" +``` + +POC 使用设置页当前默认 AI Provider、模型、Base URL 和安全存储中的 API Key。Local Query API 仍使用现有 Bearer Token;POC 输出不会打印 Token、API Key、数据库路径或内部消息 ID。 + +输出为 JSON,包含: + +- `question`、`provider`、`model` +- `modelCallCount`、`toolCallCount` +- `firstModelMs`、`toolTotalMs`、`finalModelMs`、`totalMs` +- 每次工具调用的名称、脱敏参数、耗时、状态和结果数量 +- 最终 `answer` 或错误信息 + +工具调用最多 5 次,只允许 `query_messages`、`search_messages`、`message_context`、`conversation_overview`。未配置 AI Provider、Local Query API 未启动或当前 Provider 协议不支持 tools 时,POC 会直接返回错误,不会回退到另一套模型配置。 diff --git a/electron.vite.config.ts b/electron.vite.config.ts index aa0755c..2b7139e 100644 --- a/electron.vite.config.ts +++ b/electron.vite.config.ts @@ -9,6 +9,7 @@ export default defineConfig({ input: { index: resolve('src/main/index.ts'), reportTemplateTest: resolve('src/main/report-template-test-entry.ts'), + queryAgentPoc: resolve('src/main/query-agent-poc-entry.ts'), voiceRecognitionWorker: resolve('src/main/voice-pipeline/voice-recognition-worker.ts'), knowledgeWorker: resolve('src/main/knowledge/knowledge-worker.ts') }, diff --git a/package.json b/package.json index 2194bb0..83958c5 100644 --- a/package.json +++ b/package.json @@ -41,6 +41,7 @@ "predev": "node -e \"require('electron')\"", "dev": "node scripts/ensure-env.cjs && node scripts/build-wechat-connector.cjs && electron-vite dev", "dev:update": "cross-env TRACEMEMO_UPDATE_SIMULATION=true pnpm dev", + "poc:query-agent": "electron-vite build && electron out/main/queryAgentPoc.js", "test:wechat-connector": "go -C services/wechat-connector test ./... && go -C services/wechat-connector vet ./...", "test:unit": "vitest run --config vitest.unit.config.ts", "test:component": "vitest run --config vitest.component.config.ts", diff --git a/src/main/query-agent-poc-entry.ts b/src/main/query-agent-poc-entry.ts new file mode 100644 index 0000000..3bd2463 --- /dev/null +++ b/src/main/query-agent-poc-entry.ts @@ -0,0 +1,38 @@ +import './app-data-bootstrap' +import { app } from 'electron' +import { apiTokenStore } from './api-token-store' +import { AIProviderService } from './services/ai-provider-service' +import { QueryAgentPocService, type QueryAgentToolResult } from './services/query-agent-poc-service' + +const baseUrl = (process.env.TRACEMEMO_QUERY_API_BASE || 'http://127.0.0.1:6131/api/v1').replace(/\/+$/, '') +const question = process.argv.slice(2).join(' ').trim() + +async function callQueryApi(name: string, input: Record): Promise { + const paths: Record = { + query_messages: '/query/messages', + search_messages: '/query/search', + message_context: '/query/message-context', + conversation_overview: '/query/conversation-overview' + } + const path = paths[name] + if (!path) throw new Error(`不允许的工具: ${name}`) + const token = apiTokenStore.getTokenForAuthentication() + if (!token) throw new Error('Local Query API Token 不可用') + const response = await fetch(`${baseUrl}${path}`, { method: 'POST', headers: { authorization: `Bearer ${token}`, 'content-type': 'application/json' }, body: JSON.stringify(input) }) + const payload = await response.json().catch(() => ({})) as Record + const status = typeof payload.status === 'string' ? payload.status : 'execution_failed' + return { ...payload, status } +} + +async function main(): Promise { + await app.whenReady() + const service = new QueryAgentPocService(new AIProviderService(), callQueryApi) + const result = await service.run(question) + process.stdout.write(`${JSON.stringify(result, null, 2)}\n`) + app.quit() +} + +void main().catch((error) => { + process.stderr.write(`${JSON.stringify({ error: error instanceof Error ? error.message : String(error) }, null, 2)}\n`) + app.quit() +}) diff --git a/src/main/services/ai-provider-service.ts b/src/main/services/ai-provider-service.ts index 48bcf8b..89368d6 100644 --- a/src/main/services/ai-provider-service.ts +++ b/src/main/services/ai-provider-service.ts @@ -25,6 +25,15 @@ interface AIProviderMetadataFile { type AIMessagePart = { type: 'text'; text: string } | { type: 'image'; dataUrl: string } type AIMessage = { role: string; content: string | AIMessagePart[] } type AIChatDeltaHandler = (delta: string) => void +export interface AIChatToolCall { + id: string + name: string + arguments: string +} +export interface AIChatToolDefinition { + type: 'function' + function: { name: string; description: string; parameters: Record } +} type AIRequestResult = { data: string finishReason?: string @@ -36,6 +45,7 @@ interface OpenAIResponsePayload { message?: { content?: string | Array<{ type?: string; text?: string }> | null reasoning_content?: string + tool_calls?: Array<{ id?: string; function?: { name?: string; arguments?: string } }> } finish_reason?: string }> @@ -303,6 +313,39 @@ export class AIProviderService { } } + async chatWithTools( + messages: Array>, + tools: AIChatToolDefinition[], + options?: AIChatRequestOptions, + signal?: AbortSignal + ): Promise<{ + success: boolean + data?: string + toolCalls?: AIChatToolCall[] + usage?: { input?: number; output?: number; total?: number; estimated?: boolean } + error?: string + }> { + try { + if (options?.apiKey) throw new Error('Tool Calling 不支持 legacy provider 配置') + const resolved = this.resolveProvider(options) + if (resolved.provider.type !== 'openai-compatible' || resolved.provider.advanced.apiProtocol === 'responses') { + throw new Error('当前 Provider 协议暂不支持 Tool Calling POC') + } + const result = await requestOpenAICompatibleWithTools( + resolved.provider, + resolved.key, + resolved.model, + messages, + tools, + signal + ) + return { success: true, ...result } + } catch (error) { + if (signal?.aborted) throw error + return { success: false, error: safeAIError(error) } + } + } + /** * 多模态图片理解。 * 输入:text + image parts 的 messages,返回 AI 文本响应。 @@ -862,6 +905,58 @@ async function requestOpenAIResponses( ) } +async function requestOpenAICompatibleWithTools( + provider: AIProviderSummary, + apiKey: string, + model: string, + messages: Array>, + tools: AIChatToolDefinition[], + signal?: AbortSignal +): Promise<{ data: string; toolCalls: AIChatToolCall[]; usage?: AIRequestResult['usage'] }> { + const endpoint = provider.baseUrl.endsWith('/chat/completions') + ? provider.baseUrl + : `${provider.baseUrl.replace(/\/+$/, '')}/chat/completions` + return fetchWithTimeout( + endpoint, + { + method: 'POST', + headers: buildHeaders(provider, apiKey), + body: JSON.stringify({ + model, + messages, + tools, + tool_choice: 'auto', + temperature: provider.advanced.temperature, + max_tokens: modelMaxTokens(provider, model), + ...(provider.advanced.thinking === 'disabled' ? { thinking: { type: 'disabled' } } : {}) + }) + }, + provider.advanced.timeoutMs, + signal, + async (response) => { + const payload = await parseJsonResponse(response) + if (!response.ok) { + throw new AIProviderRequestError(payload.error?.message || `AI 请求失败 (${response.status})`, { + status: response.status, + code: payload.error?.code, + type: payload.error?.type, + responseBody: payload.error + }) + } + const message = payload.choices?.[0]?.message + return { + data: openAIMessageText(message?.content), + toolCalls: (message?.tool_calls || []).flatMap((call, index) => { + const name = call.function?.name + if (!name) return [] + return [{ id: call.id || `tool-call-${index + 1}`, name, arguments: call.function?.arguments || '{}' }] + }), + usage: toOpenAIUsage(payload.usage) + } + } + ) +} + async function requestAnthropic( provider: AIProviderSummary, apiKey: string, diff --git a/src/main/services/query-agent-poc-service.ts b/src/main/services/query-agent-poc-service.ts new file mode 100644 index 0000000..1dff7a5 --- /dev/null +++ b/src/main/services/query-agent-poc-service.ts @@ -0,0 +1,160 @@ +import type { AIChatToolCall, AIChatToolDefinition } from './ai-provider-service' +import { LOCAL_QUERY_TOOL_DEFINITIONS } from '../../shared/local-query-api' + +const MAX_TOOL_CALLS = 5 +const FORBIDDEN_INPUT_KEYS = new Set(['apiKey', 'authorization', 'token', 'databasePath', 'sql', 'wxid', 'md5']) + +export interface QueryAgentProvider { + getRuntimeConfig(): { configured: boolean; providerName: string; model: string; modelName: string } + chatWithTools( + messages: Array>, + tools: AIChatToolDefinition[] + ): Promise<{ success: boolean; data?: string; toolCalls?: AIChatToolCall[]; usage?: { input?: number; output?: number; total?: number; estimated?: boolean }; error?: string }> +} + +export interface QueryAgentToolResult { + status: string + [key: string]: unknown +} + +export interface QueryAgentTraceItem { + toolName: string + input: Record + durationMs: number + status: string + resultCount?: number + evidenceCount?: number +} + +export interface QueryAgentPocResult { + question: string + provider: string + model: string + modelCallCount: number + toolCallCount: number + firstModelMs?: number + toolTotalMs: number + finalModelMs?: number + totalMs: number + traces: QueryAgentTraceItem[] + answer?: string + error?: string +} + +export type QueryAgentToolExecutor = ( + name: string, + input: Record +) => Promise + +const SYSTEM_PROMPT = `你是 TraceMemo 的本地聊天查询助手。你只能使用提供的四个 Query Tool 获取事实。 +Tool 返回的是事实来源。不得编造未返回的消息、猜测联系人、修改 resolvedTimeRange,或把 partial/unknown 当作 complete。不得把 sampled Evidence 当作完整聊天。 +当 coverage complete 且结果为 0 时,可以说明当前可读取的完整范围没有找到;当 coverage partial/unknown 且结果为 0 时,必须说明无法确认绝对不存在。 +缺少必要信息时直接用自然语言澄清;超出工具能力时说明不能可靠完成,并给出当前工具可以执行的替代方向。最终回答只基于工具结果。` + +function toolDefinitions(): AIChatToolDefinition[] { + return LOCAL_QUERY_TOOL_DEFINITIONS.map((tool) => ({ + type: 'function' as const, + function: { name: tool.name, description: tool.description, parameters: tool.parameters } + })) +} + +function sanitizeInput(input: Record): Record { + const sanitizeValue = (value: unknown): unknown => { + if (typeof value === 'string') return value.length > 240 ? `${value.slice(0, 240)}...` : value + if (Array.isArray(value)) return value.slice(0, 8).map(sanitizeValue) + if (value && typeof value === 'object') return sanitizeInput(value as Record) + return value + } + const output: Record = {} + for (const [key, value] of Object.entries(input)) { + if (key === 'messageRef') output[key] = '[opaque-message-ref]' + else if (!FORBIDDEN_INPUT_KEYS.has(key)) output[key] = sanitizeValue(value) + } + return output +} + +function containsForbiddenKey(value: unknown): boolean { + if (Array.isArray(value)) return value.some(containsForbiddenKey) + if (!value || typeof value !== 'object') return false + return Object.entries(value as Record).some(([key, child]) => FORBIDDEN_INPUT_KEYS.has(key) || containsForbiddenKey(child)) +} + +function validateToolInput(name: string, value: unknown): Record { + if (!LOCAL_QUERY_TOOL_DEFINITIONS.some((tool) => tool.name === name)) throw new Error(`不允许的工具: ${name}`) + if (!value || typeof value !== 'object' || Array.isArray(value)) throw new Error('工具参数必须是 JSON 对象') + const input = value as Record + if (containsForbiddenKey(input)) throw new Error('工具参数包含受限字段') + return input +} + +function resultCount(result: QueryAgentToolResult): { resultCount?: number; evidenceCount?: number } { + return { + resultCount: typeof result.returnedCount === 'number' ? result.returnedCount : undefined, + evidenceCount: typeof result.evidenceCount === 'number' ? result.evidenceCount : Array.isArray(result.evidence) ? result.evidence.length : undefined + } +} + +export class QueryAgentPocService { + constructor(private readonly provider: QueryAgentProvider, private readonly executeTool: QueryAgentToolExecutor) {} + + async run(question: string): Promise { + const trimmed = question.trim() + const startedAt = Date.now() + const runtime = this.provider.getRuntimeConfig() + const result: QueryAgentPocResult = { question: trimmed, provider: runtime.providerName, model: runtime.modelName || runtime.model, modelCallCount: 0, toolCallCount: 0, toolTotalMs: 0, totalMs: 0, traces: [] } + if (!trimmed) return { ...result, error: '请输入查询问题', totalMs: Date.now() - startedAt } + if (!runtime.configured) return { ...result, error: '当前 AI Provider 尚未配置', totalMs: Date.now() - startedAt } + + const messages: Array> = [ + { role: 'system', content: SYSTEM_PROMPT }, + { role: 'user', content: trimmed } + ] + const tools = toolDefinitions() + let firstModelAt: number | undefined + let finalModelDuration: number | undefined + while (result.toolCallCount < MAX_TOOL_CALLS) { + const modelStartedAt = Date.now() + const model = await this.provider.chatWithTools(messages, tools) + result.modelCallCount += 1 + const modelDuration = Date.now() - modelStartedAt + if (firstModelAt === undefined) firstModelAt = Date.now() + if (!model.success) return { ...result, error: model.error || '模型调用失败', firstModelMs: firstModelAt - startedAt, totalMs: Date.now() - startedAt } + const calls = model.toolCalls || [] + if (calls.length === 0) { + finalModelDuration = modelDuration + result.answer = model.data?.trim() || '模型未返回答案' + result.firstModelMs = firstModelAt - startedAt + result.finalModelMs = finalModelDuration + result.totalMs = Date.now() - startedAt + return result + } + if (result.toolCallCount + calls.length > MAX_TOOL_CALLS) { + return { ...result, error: `超过最大工具调用次数(${MAX_TOOL_CALLS})`, firstModelMs: firstModelAt - startedAt, totalMs: Date.now() - startedAt } + } + messages.push({ role: 'assistant', content: model.data || '', tool_calls: calls.map((call) => ({ id: call.id, type: 'function', function: { name: call.name, arguments: call.arguments } })) }) + for (const call of calls) { + const inputStartedAt = Date.now() + let toolResult: QueryAgentToolResult + let traceInput: Record = {} + try { + let parsed: unknown + try { parsed = JSON.parse(call.arguments || '{}') } catch { throw new Error('工具参数 JSON 无效') } + const input = validateToolInput(call.name, parsed) + traceInput = input + toolResult = await this.executeTool(call.name, input) + } catch (error) { + toolResult = { status: 'invalid_request', error: error instanceof Error ? error.message : '工具调用失败' } + } + const durationMs = Date.now() - inputStartedAt + result.toolCallCount += 1 + result.toolTotalMs += durationMs + const counts = resultCount(toolResult) + result.traces.push({ toolName: call.name, input: sanitizeInput(traceInput), durationMs, status: toolResult.status, ...counts }) + messages.push({ role: 'tool', tool_call_id: call.id, name: call.name, content: JSON.stringify(toolResult) }) + } + } + result.firstModelMs = firstModelAt ? firstModelAt - startedAt : undefined + result.totalMs = Date.now() - startedAt + return { ...result, error: `超过最大工具调用次数(${MAX_TOOL_CALLS})` } + } +} diff --git a/src/shared/local-query-api.ts b/src/shared/local-query-api.ts index a6fb7f8..82ca9f9 100644 --- a/src/shared/local-query-api.ts +++ b/src/shared/local-query-api.ts @@ -12,6 +12,19 @@ export type QueryMessageType = | 'sticker' | 'system' | 'other' +export interface LocalQueryToolDefinition { + name: 'query_messages' | 'search_messages' | 'message_context' | 'conversation_overview' + description: string + parameters: Record +} +const targetSchema = { type: 'object', properties: { query: { type: 'string', minLength: 1 } }, required: ['query'], additionalProperties: false } +const timeRangeSchema = { type: 'object', properties: { kind: { enum: ['all', 'today', 'yesterday', 'this_week', 'last_7_days', 'this_month', 'previous_month', 'this_year', 'previous_year', 'absolute'] }, startTime: { type: 'number' }, endTime: { type: 'number' } }, required: ['kind'], additionalProperties: false } +export const LOCAL_QUERY_TOOL_DEFINITIONS: LocalQueryToolDefinition[] = [ + { name: 'query_messages', description: '按联系人、时间、方向、消息类型和顺序精确读取消息。', parameters: { type: 'object', required: ['target', 'timeRange'], additionalProperties: false, properties: { target: targetSchema, timeRange: timeRangeSchema, direction: { enum: ['any', 'from_target', 'to_target'] }, messageTypes: { type: 'array', items: { enum: ['text', 'image', 'voice', 'video', 'file', 'link', 'sticker', 'system', 'other'] } }, order: { enum: ['asc', 'desc'] }, limit: { type: 'integer', minimum: 1, maximum: 200 }, excludeSystem: { type: 'boolean' } } } }, + { name: 'search_messages', description: '搜索指定联系人和时间范围内与主题相关的聊天 Evidence。', parameters: { type: 'object', required: ['target', 'timeRange', 'query'], additionalProperties: false, properties: { target: targetSchema, timeRange: timeRangeSchema, query: { type: 'string', minLength: 1 }, variants: { type: 'array', maxItems: 4, items: { type: 'string' } }, limit: { type: 'integer', minimum: 1, maximum: 200 } } } }, + { name: 'message_context', description: '获取已找到消息的前后上下文。', parameters: { type: 'object', required: ['messageRef'], additionalProperties: false, properties: { messageRef: { type: 'string', minLength: 1 }, before: { type: 'integer', minimum: 0, maximum: 50 }, after: { type: 'integer', minimum: 0, maximum: 50 } } } }, + { name: 'conversation_overview', description: '获取指定联系人和时间范围的聊天覆盖样本。', parameters: { type: 'object', required: ['target', 'timeRange'], additionalProperties: false, properties: { target: targetSchema, timeRange: timeRangeSchema } } } +] export type QueryTimeRange = | { kind: 'all' | 'today' | 'yesterday' | 'this_week' | 'last_7_days' | 'this_month' | 'previous_month' | 'this_year' | 'previous_year' } | { kind: 'absolute'; startTime?: number; endTime?: number } diff --git a/tests/unit/ai-provider-search-consent.test.ts b/tests/unit/ai-provider-search-consent.test.ts index fd7a270..02a5b60 100644 --- a/tests/unit/ai-provider-search-consent.test.ts +++ b/tests/unit/ai-provider-search-consent.test.ts @@ -98,6 +98,30 @@ describe('AI Search provider identity', () => { vi.unstubAllGlobals() }) + it('serializes bounded tool definitions and parses tool calls through the selected provider', async () => { + const service = new AIProviderService() + service.save({ ...provider('https://tools.example.test/v1'), type: 'openai-compatible' }) + const fetchMock = vi.fn().mockResolvedValue( + new Response(JSON.stringify({ choices: [{ message: { content: '', tool_calls: [{ id: 'call-1', function: { name: 'query_messages', arguments: '{"limit":1}' } }] } }], usage: {} }), { + status: 200, + headers: { 'content-type': 'application/json' } + }) + ) + vi.stubGlobal('fetch', fetchMock) + try { + const result = await service.chatWithTools( + [{ role: 'user', content: 'find the first message' }], + [{ type: 'function', function: { name: 'query_messages', description: 'read messages', parameters: { type: 'object' } } }] + ) + expect(result).toMatchObject({ success: true, toolCalls: [{ id: 'call-1', name: 'query_messages', arguments: '{"limit":1}' }] }) + const request = JSON.parse(String(fetchMock.mock.calls[0]?.[1]?.body)) as { model: string; tools: unknown[]; tool_choice: string } + expect(request).toMatchObject({ model: 'fixture-model', tool_choice: 'auto' }) + expect(request.tools).toHaveLength(1) + } finally { + vi.unstubAllGlobals() + } + }) + it('aborts the provider fetch when the caller cancels an AI request', async () => { const service = new AIProviderService() service.save(provider('http://127.0.0.1:11434')) diff --git a/tests/unit/query-agent-poc-service.test.ts b/tests/unit/query-agent-poc-service.test.ts new file mode 100644 index 0000000..70196ee --- /dev/null +++ b/tests/unit/query-agent-poc-service.test.ts @@ -0,0 +1,42 @@ +import { describe, expect, it, vi } from 'vitest' +import { QueryAgentPocService, type QueryAgentProvider } from '../../src/main/services/query-agent-poc-service' + +function provider(responses: Array>>, configured = true): QueryAgentProvider { + return { + getRuntimeConfig: () => ({ configured, providerName: 'Fixture Provider', model: 'fixture-model', modelName: 'Fixture Model' }), + chatWithTools: vi.fn(async () => responses.shift() || { success: true, data: 'done' }) + } +} + +describe('QueryAgentPocService', () => { + it('runs a bounded model -> tool -> model loop and records sanitized trace', async () => { + const execute = vi.fn(async () => ({ status: 'completed', returnedCount: 1, messages: [{ messageRef: 'secret-ref' }] })) + const service = new QueryAgentPocService(provider([ + { success: true, toolCalls: [{ id: 'call-1', name: 'query_messages', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, limit: 1 }) }] }, + { success: true, data: '第一条消息是图片。' } + ]), execute) + const result = await service.run('我和 BOBO 最开始聊了什么') + expect(result.answer).toBe('第一条消息是图片。') + expect(result.modelCallCount).toBe(2) + expect(result.toolCallCount).toBe(1) + expect(result.traces[0]).toMatchObject({ toolName: 'query_messages', status: 'completed', resultCount: 1 }) + expect(JSON.stringify(result)).not.toContain('secret-ref') + }) + + it('rejects unknown tools and stops after five calls', async () => { + const execute = vi.fn(async () => ({ status: 'completed' })) + const responses = Array.from({ length: 6 }, () => ({ success: true, toolCalls: [{ id: 'x', name: 'unknown', arguments: '{}' }] })) + const result = await new QueryAgentPocService(provider(responses), execute).run('test') + expect(result.toolCallCount).toBe(5) + expect(result.traces.every((trace) => trace.status === 'invalid_request')).toBe(true) + expect(result.error).toContain('最大工具调用次数') + expect(execute).not.toHaveBeenCalled() + }) + + it('clarifies unavailable configuration without making a model call', async () => { + const configuredProvider = provider([], false) + const result = await new QueryAgentPocService(configuredProvider, vi.fn()).run('test') + expect(result.error).toContain('尚未配置') + expect(configuredProvider.chatWithTools).not.toHaveBeenCalled() + }) +})