From dc6d1eb6acd392e675700bdebf909f595c740b40 Mon Sep 17 00:00:00 2001 From: Wxw-Gu Date: Thu, 10 Sep 2026 11:10:47 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E6=9F=A5=E8=AF=A2?= =?UTF-8?q?=E5=B7=A5=E5=85=B7=E6=B6=88=E6=81=AF=E5=BC=95=E7=94=A8=E4=B8=8E?= =?UTF-8?q?=E5=8F=82=E6=95=B0=E6=A0=A1=E9=AA=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/main/services/local-query-api-service.ts | 29 +++- src/main/services/query-agent-poc-service.ts | 139 +++++++++++++++++-- tests/unit/local-query-api-service.test.ts | 23 +++ tests/unit/query-agent-poc-service.test.ts | 34 ++++- 4 files changed, 202 insertions(+), 23 deletions(-) diff --git a/src/main/services/local-query-api-service.ts b/src/main/services/local-query-api-service.ts index e5f1523..0aed8c0 100644 --- a/src/main/services/local-query-api-service.ts +++ b/src/main/services/local-query-api-service.ts @@ -21,13 +21,30 @@ function kindOf(message: FormattedMessage): QueryMessageType { if (message.content?.trim()) return 'text' return 'other' } -function toRef(conversationId: string, messageId: string): string { - return Buffer.from(JSON.stringify({ c: conversationId, m: messageId }), 'utf8').toString('base64url') +interface CanonicalMessageIdentity { + conversationId: string + messageId: string } -function fromRef(value: string): { c: string; m: string } | null { + +function normalizeMessageIdentity(conversationId: string, messageId: string): CanonicalMessageIdentity | null { + const normalizedConversationId = conversationId.trim() + const normalizedMessageId = messageId.trim().replace(/^local:/, '') + if (!normalizedConversationId || !normalizedMessageId) return null + return { conversationId: normalizedConversationId, messageId: normalizedMessageId } +} + +function toRef(conversationId: string, messageId: string): string { + const identity = normalizeMessageIdentity(conversationId, messageId) + if (!identity) throw new Error('消息引用无效') + return Buffer.from(JSON.stringify({ c: identity.conversationId, m: identity.messageId }), 'utf8').toString('base64url') +} + +function fromRef(value: string): CanonicalMessageIdentity | null { try { const parsed = JSON.parse(Buffer.from(value, 'base64url').toString('utf8')) - return typeof parsed?.c === 'string' && typeof parsed?.m === 'string' ? parsed : null + return typeof parsed?.c === 'string' && typeof parsed?.m === 'string' + ? normalizeMessageIdentity(parsed.c, parsed.m) + : null } catch { return null } } function contactView(contact: FormattedContact) { return { displayName: contact.m_nsNickName || contact.m_nsUsrName, type: contact.type } as const } @@ -93,8 +110,8 @@ export class LocalQueryApiService { } async context(request: MessageContextRequest) { const ref = fromRef(request.messageRef); if (!ref) return { status: 'invalid_request' as const } - const before = Math.min(CONTEXT_MAX, Math.max(0, request.before ?? 10)); const after = Math.min(CONTEXT_MAX, Math.max(0, request.after ?? 10)); const messages = await listMessagesAsync(ref.c); const index = messages.findIndex((message) => message.id === ref.m); if (index < 0) return { status: 'contact_not_found' as const } - const contact = (await listContactsAsync()).find((item) => item.md5 === ref.c); if (!contact) return { status: 'contact_not_found' as const }; const map = (message: FormattedMessage) => toQueryMessage(ref.c, message, contact) + const before = Math.min(CONTEXT_MAX, Math.max(0, request.before ?? 10)); const after = Math.min(CONTEXT_MAX, Math.max(0, request.after ?? 10)); const messages = await listMessagesAsync(ref.conversationId); const index = messages.findIndex((message) => normalizeMessageIdentity(ref.conversationId, message.id)?.messageId === ref.messageId); if (index < 0) return { status: 'contact_not_found' as const } + const contact = (await listContactsAsync()).find((item) => item.md5 === ref.conversationId); if (!contact) return { status: 'contact_not_found' as const }; const map = (message: FormattedMessage) => toQueryMessage(ref.conversationId, message, contact) return { status: 'completed' as const, anchor: map(messages[index]), before: messages.slice(Math.max(0, index - before), index).map(map), after: messages.slice(index + 1, index + 1 + after).map(map) } } async overview(request: ConversationOverviewRequest) { diff --git a/src/main/services/query-agent-poc-service.ts b/src/main/services/query-agent-poc-service.ts index 1dff7a5..c2c465d 100644 --- a/src/main/services/query-agent-poc-service.ts +++ b/src/main/services/query-agent-poc-service.ts @@ -4,6 +4,29 @@ 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']) +interface ToolSchema { + type?: string + required?: string[] + additionalProperties?: boolean + properties?: Record + items?: ToolSchema + enum?: unknown[] + minLength?: number + minimum?: number + maximum?: number + minItems?: number + maxItems?: number +} + +export interface ToolArgumentValidationError { + status: 'invalid_tool_arguments' + field: string + constraint: string + expected?: unknown + actual?: unknown + [key: string]: unknown +} + export interface QueryAgentProvider { getRuntimeConfig(): { configured: boolean; providerName: string; model: string; modelName: string } chatWithTools( @@ -79,12 +102,89 @@ function containsForbiddenKey(value: unknown): boolean { 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 actualType(value: unknown): string { + if (value === null) return 'null' + if (Array.isArray(value)) return 'array' + return typeof value +} + +function schemaTypeMatches(value: unknown, type: string): boolean { + if (type === 'object') return Boolean(value && typeof value === 'object' && !Array.isArray(value)) + if (type === 'integer') return typeof value === 'number' && Number.isInteger(value) + if (type === 'number') return typeof value === 'number' && Number.isFinite(value) + return actualType(value) === type +} + +function validateSchema(value: unknown, schema: ToolSchema, field = '$'): ToolArgumentValidationError | undefined { + if (schema.type && !schemaTypeMatches(value, schema.type)) { + return { status: 'invalid_tool_arguments', field, constraint: 'type', expected: schema.type, actual: actualType(value) } + } + if (schema.enum && !schema.enum.some((allowed) => Object.is(allowed, value))) { + return { status: 'invalid_tool_arguments', field, constraint: 'enum', expected: schema.enum, actual: value } + } + if (typeof value === 'string') { + if (schema.minLength !== undefined && value.length < schema.minLength) { + return { status: 'invalid_tool_arguments', field, constraint: 'minLength', expected: schema.minLength, actual: value.length } + } + } + if (typeof value === 'number') { + if (schema.minimum !== undefined && value < schema.minimum) { + return { status: 'invalid_tool_arguments', field, constraint: 'minimum', expected: schema.minimum, actual: value } + } + if (schema.maximum !== undefined && value > schema.maximum) { + return { status: 'invalid_tool_arguments', field, constraint: 'maximum', expected: schema.maximum, actual: value } + } + } + if (Array.isArray(value)) { + if (schema.minItems !== undefined && value.length < schema.minItems) { + return { status: 'invalid_tool_arguments', field, constraint: 'minItems', expected: schema.minItems, actual: value.length } + } + if (schema.maxItems !== undefined && value.length > schema.maxItems) { + return { status: 'invalid_tool_arguments', field, constraint: 'maxItems', expected: schema.maxItems, actual: value.length } + } + if (schema.items) { + for (let index = 0; index < value.length; index += 1) { + const error = validateSchema(value[index], schema.items, `${field}[${index}]`) + if (error) return error + } + } + } + if (value && typeof value === 'object' && !Array.isArray(value)) { + const objectValue = value as Record + for (const required of schema.required || []) { + if (!(required in objectValue)) { + return { status: 'invalid_tool_arguments', field: field === '$' ? required : `${field}.${required}`, constraint: 'required', expected: true, actual: false } + } + } + const properties = schema.properties || {} + if (schema.additionalProperties === false) { + for (const key of Object.keys(objectValue)) { + if (!(key in properties)) { + return { status: 'invalid_tool_arguments', field: field === '$' ? key : `${field}.${key}`, constraint: 'additionalProperties', expected: false, actual: true } + } + } + } + for (const [key, childSchema] of Object.entries(properties)) { + if (key in objectValue) { + const error = validateSchema(objectValue[key], childSchema, field === '$' ? key : `${field}.${key}`) + if (error) return error + } + } + } + return undefined +} + +export function validateToolArguments(name: string, value: unknown): { input?: Record; error?: ToolArgumentValidationError } { + const definition = LOCAL_QUERY_TOOL_DEFINITIONS.find((tool) => tool.name === name) + if (!definition) return { error: { status: 'invalid_tool_arguments', field: '$', constraint: 'tool', expected: 'supported tool', actual: name } } + if (!value || typeof value !== 'object' || Array.isArray(value)) { + return { error: { status: 'invalid_tool_arguments', field: '$', constraint: 'type', expected: 'object', actual: actualType(value) } } + } + if (containsForbiddenKey(value)) { + return { error: { status: 'invalid_tool_arguments', field: '$', constraint: 'forbidden_field' } } + } + const error = validateSchema(value, definition.parameters as ToolSchema) + return error ? { error } : { input: value as Record } } function resultCount(result: QueryAgentToolResult): { resultCount?: number; evidenceCount?: number } { @@ -134,23 +234,32 @@ export class QueryAgentPocService { 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 toolResult: QueryAgentToolResult | undefined 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) + 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) { - toolResult = { status: 'invalid_request', error: error instanceof Error ? error.message : '工具调用失败' } + if (!toolResult) toolResult = { status: 'invalid_request', error: error instanceof Error ? error.message : '工具调用失败' } } + const completedToolResult = toolResult || { status: 'invalid_request', error: '工具调用失败' } 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) }) + 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) }) } } result.firstModelMs = firstModelAt ? firstModelAt - startedAt : undefined diff --git a/tests/unit/local-query-api-service.test.ts b/tests/unit/local-query-api-service.test.ts index f5629a0..d544c37 100644 --- a/tests/unit/local-query-api-service.test.ts +++ b/tests/unit/local-query-api-service.test.ts @@ -47,6 +47,29 @@ describe('LocalQueryApiService', () => { expect(result.messages?.[0]).not.toHaveProperty('text') }) + it('round-trips opaque refs from messages, search, and overview through context', async () => { + fixture.contacts.splice(1) + const queried = await service.messages({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, order: 'asc', limit: 1 }) + const queriedRef = queried.messages?.[0]?.messageRef + expect(queriedRef).toEqual(expect.any(String)) + await expect(service.context({ messageRef: queriedRef!, before: 0, after: 0 })).resolves.toMatchObject({ + status: 'completed', + anchor: { messageRef: queriedRef } + }) + + knowledge.search.mockResolvedValueOnce({ state: 'ready', evidence: [{ conversationId: 'md5-bobo', messageId: 'local:m1', timestamp: 1, sender: 'BOBO', sourceKind: 'text', text: '你好' }], conversationRetrieval: { totalMessages: 2, chunkCount: 1, complete: true }, voiceCoverage: undefined }) + const searched = await service.search({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, query: '你好' }) + const searchedRef = searched.evidence?.[0]?.messageRef + expect(searchedRef).toBe(queriedRef) + await expect(service.context({ messageRef: searchedRef!, before: 0, after: 0 })).resolves.toMatchObject({ status: 'completed' }) + + knowledge.search.mockResolvedValueOnce({ state: 'ready', evidence: [{ conversationId: 'md5-bobo', messageId: 'local:m1', timestamp: 1, sender: 'BOBO', sourceKind: 'text', text: '你好' }], conversationRetrieval: { totalMessages: 2, chunkCount: 1, complete: true }, voiceCoverage: undefined }) + const overview = await service.overview({ target: { query: 'BOBO' }, timeRange: { kind: 'all' } }) + const overviewRef = overview.evidence?.[0]?.messageRef + expect(overviewRef).toBe(queriedRef) + await expect(service.context({ messageRef: overviewRef!, before: 0, after: 0 })).resolves.toMatchObject({ status: 'completed' }) + }) + it('sorts overview evidence while preserving the selected set and exposes sampling', async () => { fixture.contacts.splice(1) knowledge.search.mockResolvedValueOnce({ state: 'ready', evidence: [ diff --git a/tests/unit/query-agent-poc-service.test.ts b/tests/unit/query-agent-poc-service.test.ts index 70196ee..ca2aae2 100644 --- a/tests/unit/query-agent-poc-service.test.ts +++ b/tests/unit/query-agent-poc-service.test.ts @@ -1,5 +1,5 @@ import { describe, expect, it, vi } from 'vitest' -import { QueryAgentPocService, type QueryAgentProvider } from '../../src/main/services/query-agent-poc-service' +import { QueryAgentPocService, type QueryAgentProvider, validateToolArguments } from '../../src/main/services/query-agent-poc-service' function provider(responses: Array>>, configured = true): QueryAgentProvider { return { @@ -28,7 +28,7 @@ describe('QueryAgentPocService', () => { 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.traces.every((trace) => trace.status === 'invalid_tool_arguments')).toBe(true) expect(result.error).toContain('最大工具调用次数') expect(execute).not.toHaveBeenCalled() }) @@ -39,4 +39,34 @@ describe('QueryAgentPocService', () => { expect(result.error).toContain('尚未配置') expect(configuredProvider.chatWithTools).not.toHaveBeenCalled() }) + + it('accepts schema-valid variants and rejects invalid arguments before executing a tool', async () => { + const execute = vi.fn(async () => ({ status: 'completed', returnedCount: 0 })) + const valid = await new QueryAgentPocService(provider([ + { success: true, toolCalls: [{ id: 'valid', name: 'search_messages', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, query: '答应', variants: ['承诺', '保证', '说好', '一定'] }) }] }, + { success: true, data: 'done' } + ]), execute).run('test') + expect(valid.traces[0]).toMatchObject({ status: 'completed', toolName: 'search_messages' }) + expect(execute).toHaveBeenCalledTimes(1) + + execute.mockClear() + const invalid = await new QueryAgentPocService(provider([ + { success: true, toolCalls: [{ id: 'invalid', name: 'search_messages', arguments: JSON.stringify({ target: { query: 'BOBO' }, timeRange: { kind: 'all' }, query: '答应', variants: ['1', '2', '3', '4', '5'] }) }] }, + { success: true, data: '修正后完成' } + ]), execute).run('test') + expect(invalid.traces[0]).toMatchObject({ status: 'invalid_tool_arguments' }) + expect(invalid.traces[0].resultCount).toBeUndefined() + expect(execute).not.toHaveBeenCalled() + }) + + it('enforces generic nested tool schema constraints', () => { + const base = { target: { query: 'BOBO' }, timeRange: { kind: 'all' }, query: 'x' } + expect(validateToolArguments('search_messages', { ...base, variants: ['1', '2', '3', '4'] }).error).toBeUndefined() + expect(validateToolArguments('search_messages', { ...base, variants: ['1', '2', '3', '4', '5'] }).error).toMatchObject({ field: 'variants', constraint: 'maxItems', expected: 4, actual: 5 }) + expect(validateToolArguments('search_messages', { ...base, extra: true }).error).toMatchObject({ field: 'extra', constraint: 'additionalProperties' }) + expect(validateToolArguments('query_messages', { target: { query: 'BOBO' }, timeRange: { kind: 'all' }, direction: 'sideways' }).error).toMatchObject({ field: 'direction', constraint: 'enum' }) + expect(validateToolArguments('query_messages', { target: { query: '' }, timeRange: { kind: 'all' } }).error).toMatchObject({ field: 'target.query', constraint: 'minLength' }) + expect(validateToolArguments('query_messages', { target: { query: 'BOBO' }, timeRange: { kind: 'all' }, limit: 0 }).error).toMatchObject({ field: 'limit', constraint: 'minimum' }) + expect(validateToolArguments('message_context', { messageRef: 'opaque', before: 51 }).error).toMatchObject({ field: 'before', constraint: 'maximum' }) + }) })