mirror of
https://wget.la/https://github.com/Wxw-Gu/WechatExplorer
synced 2026-10-03 18:33:14 +08:00
fix: 修复查询工具消息引用与参数校验
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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: [
|
||||
|
||||
@@ -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' })
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user