feat: 新增本地查询模型调用

This commit is contained in:
Wxw-Gu
2026-09-10 09:35:42 +08:00
parent c759bcd30e
commit 7ae0d63a7f
9 changed files with 401 additions and 0 deletions
+27
View File
@@ -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 会直接返回错误,不会回退到另一套模型配置。
+1
View File
@@ -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')
},
+1
View File
@@ -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",
+38
View File
@@ -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<string, unknown>): Promise<QueryAgentToolResult> {
const paths: Record<string, string> = {
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<string, unknown>
const status = typeof payload.status === 'string' ? payload.status : 'execution_failed'
return { ...payload, status }
}
async function main(): Promise<void> {
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()
})
+95
View File
@@ -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<string, unknown> }
}
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<Record<string, unknown>>,
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<Record<string, unknown>>,
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<OpenAIResponsePayload>(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,
@@ -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<Record<string, unknown>>,
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<string, unknown>
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<string, unknown>
) => Promise<QueryAgentToolResult>
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<string, unknown>): Record<string, unknown> {
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<string, unknown>)
return value
}
const output: Record<string, unknown> = {}
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<string, unknown>).some(([key, child]) => FORBIDDEN_INPUT_KEYS.has(key) || containsForbiddenKey(child))
}
function validateToolInput(name: string, value: unknown): Record<string, unknown> {
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<string, unknown>
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<QueryAgentPocResult> {
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<Record<string, unknown>> = [
{ 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<string, unknown> = {}
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})` }
}
}
+13
View File
@@ -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<string, unknown>
}
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 }
@@ -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'))
@@ -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<Awaited<ReturnType<QueryAgentProvider['chatWithTools']>>>, 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()
})
})