fix: 修复查询工具消息引用与参数校验

This commit is contained in:
Wxw-Gu
2026-09-10 11:10:47 +08:00
parent 7ae0d63a7f
commit dc6d1eb6ac
4 changed files with 202 additions and 23 deletions
+23 -6
View File
@@ -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) {
+124 -15
View File
@@ -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<string, ToolSchema>
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<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 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<string, unknown>
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<string, unknown>; 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<string, unknown> }
}
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<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)
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
@@ -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: [
+32 -2
View File
@@ -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<Awaited<ReturnType<QueryAgentProvider['chatWithTools']>>>, 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' })
})
})