From 376ea79ff6b55f6a9d26015daf9c3b2fd6f2c13b Mon Sep 17 00:00:00 2001 From: Wxw-Gu Date: Thu, 10 Sep 2026 15:35:12 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E4=BC=98=E5=8C=96=E9=97=AE=E9=97=AE?= =?UTF-8?q?=E5=BE=AE=E4=BF=A1=E5=B7=A5=E5=85=B7=E8=A7=84=E5=88=92?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/main/services/ai-provider-service.ts | 3 +- src/main/services/query-agent-poc-service.ts | 85 +++++++++++++++---- src/shared/local-query-api.ts | 8 +- tests/unit/ai-provider-search-consent.test.ts | 21 +++++ tests/unit/query-agent-poc-service.test.ts | 48 +++++++++++ 5 files changed, 144 insertions(+), 21 deletions(-) diff --git a/src/main/services/ai-provider-service.ts b/src/main/services/ai-provider-service.ts index 89368d6..b0ebecb 100644 --- a/src/main/services/ai-provider-service.ts +++ b/src/main/services/ai-provider-service.ts @@ -924,8 +924,7 @@ async function requestOpenAICompatibleWithTools( body: JSON.stringify({ model, messages, - tools, - tool_choice: 'auto', + ...(tools.length ? { tools, tool_choice: 'auto' } : { tool_choice: 'none' }), temperature: provider.advanced.temperature, max_tokens: modelMaxTokens(provider, model), ...(provider.advanced.thinking === 'disabled' ? { thinking: { type: 'disabled' } } : {}) diff --git a/src/main/services/query-agent-poc-service.ts b/src/main/services/query-agent-poc-service.ts index c2c465d..f469fe5 100644 --- a/src/main/services/query-agent-poc-service.ts +++ b/src/main/services/query-agent-poc-service.ts @@ -6,6 +6,7 @@ const FORBIDDEN_INPUT_KEYS = new Set(['apiKey', 'authorization', 'token', 'datab interface ToolSchema { type?: string + description?: string required?: string[] additionalProperties?: boolean properties?: Record @@ -69,10 +70,18 @@ export type QueryAgentToolExecutor = ( input: Record ) => Promise -const SYSTEM_PROMPT = `你是 TraceMemo 的本地聊天查询助手。你只能使用提供的四个 Query Tool 获取事实。 -Tool 返回的是事实来源。不得编造未返回的消息、猜测联系人、修改 resolvedTimeRange,或把 partial/unknown 当作 complete。不得把 sampled Evidence 当作完整聊天。 -当 coverage complete 且结果为 0 时,可以说明当前可读取的完整范围没有找到;当 coverage partial/unknown 且结果为 0 时,必须说明无法确认绝对不存在。 -缺少必要信息时直接用自然语言澄清;超出工具能力时说明不能可靠完成,并给出当前工具可以执行的替代方向。最终回答只基于工具结果。` +const SYSTEM_PROMPT = `你是 TraceMemo 的本地聊天查询助手,只能使用提供的四个 Query Tool 获取事实,最终回答只基于 Tool Result。 + +规划原则: +- 先判断问题需要哪种证据,再调用最少的 Tool。每次收到 Tool Result 后都判断“当前 Evidence 是否已经足以给出有边界的回答”;足够就立即回答,不为追求绝对完整继续调查。 +- query_messages 是精确事实查询,适用于能用联系人、时间、方向、消息类型、顺序等结构条件表达的问题。earliest/latest 等时间边界也是结构条件,必须使用 order 与 limit 精确查询,不能使用抽样 overview。结果已经回答问题时,不要追加 conversation_overview。 +- search_messages 是语义 Evidence 检索,适用于结构条件无法确定答案的问题。第一次尽量在一个调用中给出高质量 query 和最多 4 个 variants;调用前确认 variants 数组长度不超过 4。一次有效 search 后不得再次 search,已有相关 Evidence 就直接判断。 +- conversation_overview 只用于真正需要理解一个时间范围内整体聊了什么、主要话题或整体互动的 broad summary。它返回 temporal coverage sample,不代表完整聊天,也不是语义检索不足时的默认 fallback。 +- message_context 只用于已经找到一条有价值 Evidence、但单条内容缺少前后语境而无法判断真实含义的情况。不要把它当作默认确认步骤;上下文足够后立即回答。 +- 普通聊天查询不是 exhaustive investigation。经过合理的 search 或可选 context 仍不足以形成强结论时,直接说明证据范围和不确定性,不要循环调用 search、overview、context。 + +事实边界:不得编造未返回的消息、猜测联系人、修改 resolvedTimeRange,或把 partial/unknown 当作 complete。coverage complete 且结果为 0 时,可以说明当前可读取的完整范围没有找到;coverage partial/unknown 且结果为 0 时,必须说明无法确认绝对不存在。不要把 sampled Evidence 当作完整聊天,也不要把 source message count 和 selected evidence count 混为一谈。 +缺少必要信息时用自然语言澄清;超出工具能力时说明不能可靠完成,并给出当前工具可以执行的替代方向。` function toolDefinitions(): AIChatToolDefinition[] { return LOCAL_QUERY_TOOL_DEFINITIONS.map((tool) => ({ @@ -81,6 +90,10 @@ function toolDefinitions(): AIChatToolDefinition[] { })) } +function toolDefinition(name: string): AIChatToolDefinition[] { + return toolDefinitions().filter((tool) => tool.function.name === name) +} + function sanitizeInput(input: Record): Record { const sanitizeValue = (value: unknown): unknown => { if (typeof value === 'string') return value.length > 240 ? `${value.slice(0, 240)}...` : value @@ -194,6 +207,39 @@ function resultCount(result: QueryAgentToolResult): { resultCount?: number; evid } } +function messageRecordForModel(value: unknown): unknown { + if (!value || typeof value !== 'object' || Array.isArray(value)) return value + const record = value as Record + return record.messageType || !record.sourceKind ? record : { ...record, messageType: record.sourceKind } +} + +function toolResultForModel(name: string, result: QueryAgentToolResult, callsUsed: number, nextTools: AIChatToolDefinition[]): QueryAgentToolResult { + const visible: QueryAgentToolResult = { ...result } + if (Array.isArray(result.messages)) visible.messages = result.messages.map(messageRecordForModel) + if (Array.isArray(result.evidence)) visible.evidence = result.evidence.map(messageRecordForModel) + if (result.anchor) visible.anchor = messageRecordForModel(result.anchor) + if (Array.isArray(result.before)) visible.before = result.before.map(messageRecordForModel) + if (Array.isArray(result.after)) visible.after = result.after.map(messageRecordForModel) + visible._agent = { + toolName: name, + toolCallsUsed: callsUsed, + toolCallsRemaining: Math.max(0, MAX_TOOL_CALLS - callsUsed), + availableNextTools: nextTools.map((tool) => tool.function.name), + instruction: nextTools.length + ? '先判断当前 Evidence 是否足以回答;足够就立即回答,只在含义仍有明确歧义时使用当前可用 Tool。' + : '工具阶段已经结束。必须直接给出有边界的最终回答,不得再调用 Tool。' + } + return visible +} + +function nextToolDefinitions(name: string, result: QueryAgentToolResult): AIChatToolDefinition[] { + if (result.status === 'invalid_tool_arguments') return toolDefinition(name) + if (name === 'search_messages' && result.status === 'completed' && resultCount(result).evidenceCount) { + return toolDefinition('message_context') + } + return [] +} + export class QueryAgentPocService { constructor(private readonly provider: QueryAgentProvider, private readonly executeTool: QueryAgentToolExecutor) {} @@ -209,7 +255,7 @@ export class QueryAgentPocService { { role: 'system', content: SYSTEM_PROMPT }, { role: 'user', content: trimmed } ] - const tools = toolDefinitions() + let tools = toolDefinitions() let firstModelAt: number | undefined let finalModelDuration: number | undefined while (result.toolCallCount < MAX_TOOL_CALLS) { @@ -237,17 +283,22 @@ export class QueryAgentPocService { let toolResult: QueryAgentToolResult | undefined let traceInput: Record = {} try { - let parsed: unknown - try { parsed = JSON.parse(call.arguments || '{}') } catch { - toolResult = { status: 'invalid_tool_arguments', field: '$', constraint: 'json' } + if (!tools.some((tool) => tool.function.name === call.name)) { + toolResult = { status: 'invalid_tool_arguments', field: '$', constraint: 'tool_availability', expected: tools.map((tool) => tool.function.name), actual: call.name } } if (!toolResult) { - const validated = validateToolArguments(call.name, parsed) - if (validated.error) { - toolResult = validated.error - } else { - traceInput = validated.input || {} - toolResult = await this.executeTool(call.name, traceInput) + let parsed: unknown + try { parsed = JSON.parse(call.arguments || '{}') } catch { + toolResult = { status: 'invalid_tool_arguments', field: '$', constraint: 'json' } + } + if (!toolResult) { + const validated = validateToolArguments(call.name, parsed) + if (validated.error) { + toolResult = validated.error + } else { + traceInput = validated.input || {} + toolResult = await this.executeTool(call.name, traceInput) + } } } } catch (error) { @@ -259,7 +310,11 @@ export class QueryAgentPocService { result.toolTotalMs += durationMs const counts = resultCount(completedToolResult) result.traces.push({ toolName: call.name, input: sanitizeInput(traceInput), durationMs, status: completedToolResult.status, ...counts }) - messages.push({ role: 'tool', tool_call_id: call.id, name: call.name, content: JSON.stringify(completedToolResult) }) + const nextTools = completedToolResult.constraint === 'tool_availability' + ? tools + : nextToolDefinitions(call.name, completedToolResult) + messages.push({ role: 'tool', tool_call_id: call.id, name: call.name, content: JSON.stringify(toolResultForModel(call.name, completedToolResult, result.toolCallCount, nextTools)) }) + tools = nextTools } } result.firstModelMs = firstModelAt ? firstModelAt - startedAt : undefined diff --git a/src/shared/local-query-api.ts b/src/shared/local-query-api.ts index 82ca9f9..c09abbd 100644 --- a/src/shared/local-query-api.ts +++ b/src/shared/local-query-api.ts @@ -20,10 +20,10 @@ export interface LocalQueryToolDefinition { 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 } } } + { name: 'query_messages', description: '精确读取符合联系人、时间、方向、消息类型、顺序等结构条件的消息;适合具体事实和 earliest/latest 等时间边界查询,边界查询使用 order 与 limit。', 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;首次调用应集中提供高质量 query 和 variants。已有相关 Evidence 时直接判断,不要重复搜索。', parameters: { type: 'object', required: ['target', 'timeRange', 'query'], additionalProperties: false, properties: { target: targetSchema, timeRange: timeRangeSchema, query: { type: 'string', minLength: 1 }, variants: { type: 'array', description: '可选的语义变体,数组元素总数不得超过 4;调用前必须自行计数。', maxItems: 4, items: { type: 'string' } }, limit: { type: 'integer', minimum: 1, maximum: 200 } } } }, + { name: 'message_context', description: '补充已找到的单条有价值 Evidence 的前后消息;仅在该 Evidence 缺少语境、无法判断含义时使用,不是默认确认步骤。', 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: '提取指定联系人和时间范围的整体聊天覆盖样本;只用于 broad summary,不是语义搜索 fallback,也不能确定 earliest/latest 等精确时间边界。', 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' } diff --git a/tests/unit/ai-provider-search-consent.test.ts b/tests/unit/ai-provider-search-consent.test.ts index 02a5b60..d7b26d8 100644 --- a/tests/unit/ai-provider-search-consent.test.ts +++ b/tests/unit/ai-provider-search-consent.test.ts @@ -122,6 +122,27 @@ describe('AI Search provider identity', () => { } }) + it('explicitly disables tool calls when orchestration provides no tools', 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: 'done' } }], usage: {} }), { + status: 200, + headers: { 'content-type': 'application/json' } + }) + ) + vi.stubGlobal('fetch', fetchMock) + try { + const result = await service.chatWithTools([{ role: 'user', content: 'answer now' }], []) + expect(result).toMatchObject({ success: true, data: 'done', toolCalls: [] }) + const request = JSON.parse(String(fetchMock.mock.calls[0]?.[1]?.body)) as { tools?: unknown[]; tool_choice: string } + expect(request.tool_choice).toBe('none') + expect(request).not.toHaveProperty('tools') + } 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 index ca2aae2..649c55a 100644 --- a/tests/unit/query-agent-poc-service.test.ts +++ b/tests/unit/query-agent-poc-service.test.ts @@ -23,6 +23,54 @@ describe('QueryAgentPocService', () => { expect(JSON.stringify(result)).not.toContain('secret-ref') }) + it('presents planning boundaries and remaining budget with tool results', async () => { + const execute = vi.fn(async () => ({ + status: 'completed', + evidenceCount: 1, + sourceCoverage: { state: 'complete', sourceMessageCount: 10 }, + selection: { mode: 'temporal_coverage', selectedEvidenceCount: 1, sampled: true }, + evidence: [{ messageRef: 'opaque', timestamp: 1, sender: 'BOBO', sourceKind: 'image', text: 'caption' }] + })) + const configuredProvider = provider([ + { success: true, toolCalls: [{ id: 'call-1', name: 'search_messages', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, query: 'topic' }) }] }, + { success: true, data: '根据这条证据可以回答。' } + ]) + await new QueryAgentPocService(configuredProvider, execute).run('查找相关记录') + const calls = vi.mocked(configuredProvider.chatWithTools).mock.calls + const firstMessages = calls[0]?.[0] || [] + expect(String(firstMessages[0]?.content)).toContain('Evidence 是否已经足以') + const toolMessage = calls[1]?.[0].find((message) => message.role === 'tool') + const presented = JSON.parse(String(toolMessage?.content)) as Record + expect(presented._agent).toMatchObject({ toolName: 'search_messages', toolCallsUsed: 1, toolCallsRemaining: 4, availableNextTools: ['message_context'] }) + expect(presented.evidence[0]).toMatchObject({ sourceKind: 'image', messageType: 'image', text: 'caption' }) + expect(presented.sourceCoverage).toEqual({ state: 'complete', sourceMessageCount: 10 }) + expect(calls[1]?.[1].map((tool) => tool.function.name)).toEqual(['message_context']) + }) + + it('removes tools after a sufficient exact result', async () => { + const configuredProvider = provider([ + { success: true, toolCalls: [{ id: 'call-1', name: 'query_messages', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, limit: 1 }) }] }, + { success: true, data: '完成' } + ]) + await new QueryAgentPocService(configuredProvider, vi.fn(async () => ({ status: 'completed', returnedCount: 1 }))).run('第一条消息') + expect(vi.mocked(configuredProvider.chatWithTools).mock.calls[1]?.[1]).toEqual([]) + }) + + it('does not execute a tool that is unavailable after the stopping boundary', async () => { + const execute = vi.fn(async () => ({ status: 'completed', returnedCount: 1 })) + const configuredProvider = provider([ + { success: true, toolCalls: [{ id: 'call-1', name: 'query_messages', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, limit: 1 }) }] }, + { success: true, toolCalls: [{ id: 'call-2', name: 'conversation_overview', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' } }) }] }, + { success: true, data: '完成' } + ]) + const result = await new QueryAgentPocService(configuredProvider, execute).run('第一条消息') + expect(execute).toHaveBeenCalledTimes(1) + expect(result.traces[1]).toMatchObject({ toolName: 'conversation_overview', status: 'invalid_tool_arguments' }) + const thirdCallMessages = vi.mocked(configuredProvider.chatWithTools).mock.calls[2]?.[0] || [] + const unavailableResult = thirdCallMessages.findLast((message) => message.role === 'tool') + expect(String(unavailableResult?.content)).toContain('tool_availability') + }) + 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: '{}' }] }))