Compare commits
merge into: wanghanlin:master
wanghanlin:Spring-AI-1.1.2
wanghanlin:TDesign-AI-Chat
wanghanlin:TDesign-Vue-Next-1.20.6
wanghanlin:dev
wanghanlin:feature/test
wanghanlin:master
pull from: wanghanlin:feature/test
wanghanlin:Spring-AI-1.1.2
wanghanlin:TDesign-AI-Chat
wanghanlin:TDesign-Vue-Next-1.20.6
wanghanlin:dev
wanghanlin:feature/test
wanghanlin:master
10 Commits
master
...
feature/te
| Author | SHA1 | Message | Date |
|---|---|---|---|
|
|
50fc660d08 |
docs: 更新对话接口契约与管道说明
|
3 weeks ago |
|
|
5086dd16f4 |
feat(sdk): 同步与流式对话复用同次引用来源
|
3 weeks ago |
|
|
d36f63a262 |
refactor(frontend): 管理端对话复用同次引用来源
|
3 weeks ago |
|
|
b3f6774ccb |
feat(model): 支持豆包 Seed 2.0 思考强度配置
|
3 weeks ago |
|
|
28e77223ef |
feat(chat): 同次回答携带引用来源
|
3 weeks ago |
|
|
3f93d015f8 |
refactor(chat): 移除 LLM 意图分类并默认原文检索
|
3 weeks ago |
|
|
d5a5c3ff56 |
build: 支持通过 -DskipTests=false 显式运行测试
|
3 weeks ago |
|
|
8004efc87d |
fix(frontend): 正文完成后立即解除发送状态并后台补齐引用来源
|
3 weeks ago |
|
|
e1ecc18223 |
perf(intent): doubao-seed-2.0 分类请求关闭深度思考
|
3 weeks ago |
|
|
256362e34a |
refactor(chat): 前置 FAQ 三级匹配并同步流程文档
|
3 weeks ago |
37 changed files with 2194 additions and 1060 deletions
-
44CLAUDE.md
-
58SDK-INTEGRATION.md
-
32client/README.md
-
302client/src/api.ts
-
251client/src/chat.ts
-
2client/src/config.ts
-
4client/src/dom.ts
-
22client/src/types.ts
-
123client/tests/api.test.ts
-
118client/tests/chat.test.ts
-
8client/tests/config.test.ts
-
28frontend/src/api/chat.ts
-
15frontend/src/components/MessageSources.vue
-
148frontend/src/sdk-test/SdkTestPanel.vue
-
16frontend/src/types/models.ts
-
18frontend/src/types/sse.ts
-
3frontend/src/utils/chatAdapter.ts
-
275frontend/src/utils/sse.ts
-
153frontend/src/views/ChatPanel.vue
-
33frontend/src/views/ModelConfigManager.vue
-
45frontend/src/views/PipelineFlow.vue
-
164frontend/tests/chat-protocol.test.mjs
-
6pom.xml
-
50src/main/java/com/wok/supportbot/app/AssistantApp.java
-
7src/main/java/com/wok/supportbot/app/ChatContext.java
-
121src/main/java/com/wok/supportbot/app/ChatPipeline.java
-
10src/main/java/com/wok/supportbot/app/ChatResult.java
-
59src/main/java/com/wok/supportbot/app/SourceReference.java
-
27src/main/java/com/wok/supportbot/config/ChatModelFactory.java
-
73src/main/java/com/wok/supportbot/controller/AiController.java
-
14src/main/java/com/wok/supportbot/controller/OpenApiController.java
-
18src/main/java/com/wok/supportbot/rag/RagPipeline.java
-
114src/main/java/com/wok/supportbot/service/IntentRouter.java
-
250src/test/java/com/wok/supportbot/AnswerTransportTests.java
-
177src/test/java/com/wok/supportbot/ChatModelFactoryTests.java
-
300src/test/java/com/wok/supportbot/ChatPipelineTests.java
-
166src/test/java/com/wok/supportbot/ChatResultEndpointTests.java
@ -0,0 +1,123 @@ |
|||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; |
||||
|
import { chatRequest, chatSSERequest, clearApiConfig, setApiConfig } from '../src/api'; |
||||
|
import { parseConfig } from '../src/config'; |
||||
|
import type { RagSource } from '../src/types'; |
||||
|
|
||||
|
const source: RagSource = { |
||||
|
documentId: '1234567890123456789', title: 'manual', sourceName: null, |
||||
|
chunkIndex: 0, score: null, snippet: '实际命中的文档', |
||||
|
}; |
||||
|
const encoder = new TextEncoder(); |
||||
|
const chunk = (content: string) => `data: ${JSON.stringify({ choices: [{ delta: { content } }] })}\r\n\r\n`; |
||||
|
const metadata = `data: ${JSON.stringify({ object: 'chat.completion.chunk', choices: [], sources: [source] })}\r\n\r\n`; |
||||
|
|
||||
|
function response(parts: Uint8Array[], close = true) { |
||||
|
const cancel = vi.fn(); |
||||
|
const stream = new ReadableStream<Uint8Array>({ |
||||
|
start(controller) { |
||||
|
parts.forEach(part => controller.enqueue(part)); |
||||
|
if (close) controller.close(); |
||||
|
}, |
||||
|
cancel, |
||||
|
}); |
||||
|
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(stream))); |
||||
|
return { stream, cancel }; |
||||
|
} |
||||
|
function receive(signal?: AbortSignal) { |
||||
|
const onChunk = vi.fn(); |
||||
|
const onDone = vi.fn(); |
||||
|
const onError = vi.fn(); |
||||
|
const onSources = vi.fn(); |
||||
|
const pending = chatSSERequest('问题', onChunk, onDone, onError, 7, true, undefined, signal, onSources); |
||||
|
return { pending, onChunk, onDone, onError, onSources }; |
||||
|
} |
||||
|
|
||||
|
beforeEach(() => { |
||||
|
const config = parseConfig({ integrateId: '42', requestDomain: 'https://example.test', userId: 'user' })!; |
||||
|
config.chatId = 'conversation'; |
||||
|
setApiConfig(config); |
||||
|
}); |
||||
|
afterEach(() => { clearApiConfig(); vi.unstubAllGlobals(); }); |
||||
|
|
||||
|
describe('same-request answer transport', () => { |
||||
|
it('stops on DONE without EOF, preserving split UTF-8, CRLF and sources outside text', async () => { |
||||
|
const bytes = encoder.encode(chunk('中文 ') + metadata + 'data: [DONE]\r\n\r\n' + chunk('不得出现')); |
||||
|
const { stream, cancel } = response(Array.from(bytes, byte => Uint8Array.of(byte)), false); |
||||
|
const result = receive(); |
||||
|
await result.pending; |
||||
|
expect(result.onChunk.mock.calls.flat().join('')).toBe('中文 '); |
||||
|
expect(result.onSources).toHaveBeenCalledExactlyOnceWith([source]); |
||||
|
expect(result.onDone).toHaveBeenCalledTimes(1); |
||||
|
expect(result.onError).not.toHaveBeenCalled(); |
||||
|
expect(cancel).toHaveBeenCalledTimes(1); |
||||
|
expect(stream.locked).toBe(false); |
||||
|
expect(fetch).toHaveBeenCalledTimes(1); |
||||
|
const url = new URL(vi.mocked(fetch).mock.calls[0]![0] as string); |
||||
|
expect(url.pathname).toBe('/ai/chat/stream'); |
||||
|
expect(url.searchParams.get('rewriteStrategy')).toBe('NONE'); |
||||
|
}); |
||||
|
|
||||
|
it('keeps multiline data whitespace and flushes the last unterminated event at EOF', async () => { |
||||
|
response([encoder.encode('event: status\ndata: 不显示\n\ndata: first \r\ndata: second\r\n\r\ndata: last ')]); |
||||
|
const result = receive(); |
||||
|
await result.pending; |
||||
|
expect(result.onChunk.mock.calls.flat()).toEqual([' first \nsecond', 'last ']); |
||||
|
expect(result.onDone).toHaveBeenCalledTimes(1); |
||||
|
}); |
||||
|
|
||||
|
it('recognizes EOF remainder DONE rather than rendering it', async () => { |
||||
|
response([encoder.encode(chunk('answer') + 'data: [DONE]')]); |
||||
|
const result = receive(); |
||||
|
await result.pending; |
||||
|
expect(result.onChunk).toHaveBeenCalledExactlyOnceWith('answer'); |
||||
|
expect(result.onDone).toHaveBeenCalledTimes(1); |
||||
|
}); |
||||
|
|
||||
|
it('delivers original server error once even when reader cleanup rejects', async () => { |
||||
|
const { cancel, stream } = response([encoder.encode('data: {"error":{"message":"权限不足"}}\n\n')], false); |
||||
|
cancel.mockRejectedValue(new Error('cleanup failed')); |
||||
|
const result = receive(); |
||||
|
await result.pending; |
||||
|
expect(result.onError).toHaveBeenCalledTimes(1); |
||||
|
expect(result.onError.mock.calls[0]![0].message).toBe('权限不足'); |
||||
|
expect(result.onDone).not.toHaveBeenCalled(); |
||||
|
expect(stream.locked).toBe(false); |
||||
|
}); |
||||
|
|
||||
|
it('cancels a pending read once without consuming later metadata', async () => { |
||||
|
const { stream } = response([], false); |
||||
|
const controller = new AbortController(); |
||||
|
const result = receive(controller.signal); |
||||
|
await Promise.resolve(); |
||||
|
await Promise.resolve(); |
||||
|
controller.abort(); |
||||
|
await result.pending; |
||||
|
expect(result.onDone).toHaveBeenCalledTimes(1); |
||||
|
expect(result.onError).not.toHaveBeenCalled(); |
||||
|
expect(result.onSources).not.toHaveBeenCalled(); |
||||
|
expect(stream.locked).toBe(false); |
||||
|
}); |
||||
|
|
||||
|
it('uses one synchronous JSON request with the same RAG, role and category parameters', async () => { |
||||
|
const answer = { text: 'answer', mcpEvents: [], suggestions: [], sources: [source] }; |
||||
|
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(Response.json(answer))); |
||||
|
expect(await chatRequest('问题', ['https://example.test/photo.png'], 7, true)).toEqual(answer); |
||||
|
expect(fetch).toHaveBeenCalledTimes(1); |
||||
|
const url = new URL(vi.mocked(fetch).mock.calls[0]![0] as string); |
||||
|
expect(url.pathname).toBe('/ai/chat/result'); |
||||
|
expect(Object.fromEntries(url.searchParams)).toMatchObject({ |
||||
|
roleId: '42', accountId: 'user', chatId: 'conversation', categoryId: '7', enableRag: 'true', rewriteStrategy: 'NONE', |
||||
|
}); |
||||
|
}); |
||||
|
|
||||
|
it('retains explicit rewrite configuration and disabled RAG', async () => { |
||||
|
const config = parseConfig({ integrateId: '42', requestDomain: 'https://example.test', rewriteStrategy: 'MULTI_QUERY' })!; |
||||
|
setApiConfig(config); |
||||
|
vi.stubGlobal('fetch', vi.fn().mockImplementation(() => Promise.resolve(Response.json({ text: '', mcpEvents: [], suggestions: [], sources: [] })))); |
||||
|
await chatRequest('问题', undefined, undefined, true); |
||||
|
await chatRequest('问题', undefined, undefined, false); |
||||
|
const urls = vi.mocked(fetch).mock.calls.map(([url]) => new URL(url as string)); |
||||
|
expect(urls[0]!.searchParams.get('rewriteStrategy')).toBe('MULTI_QUERY'); |
||||
|
expect(urls[1]!.searchParams.get('enableRag')).toBe('false'); |
||||
|
}); |
||||
|
}); |
||||
@ -0,0 +1,118 @@ |
|||||
|
import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; |
||||
|
import { initChat, sendQuickReply } from '../src/chat'; |
||||
|
import { setApiConfig, clearApiConfig } from '../src/api'; |
||||
|
import { parseConfig } from '../src/config'; |
||||
|
import { renderAIBubble, createEmptyAIBubble, renderSources } from '../src/dom'; |
||||
|
import { saveMessages } from '../src/storage'; |
||||
|
|
||||
|
vi.mock('../src/dom', () => ({ |
||||
|
renderUserBubble: vi.fn(), renderAIBubble: vi.fn(() => ({ id: 'sync-wrapper' })), |
||||
|
createEmptyAIBubble: vi.fn(() => ({ wrapper: { id: 'stream-wrapper' }, bubble: {} })), |
||||
|
renderSources: vi.fn(), finalizeAIBubble: vi.fn(), scrollToBottom: vi.fn(), |
||||
|
hideOfflineBanner: vi.fn(), showOfflineBanner: vi.fn(), renderErrorBubble: vi.fn(), |
||||
|
})); |
||||
|
vi.mock('../src/storage', () => ({ saveMessages: vi.fn(), clearMessages: vi.fn() })); |
||||
|
|
||||
|
class ElementStub extends EventTarget { |
||||
|
style: Record<string, string> = {}; |
||||
|
classList = { add: vi.fn(), remove: vi.fn() }; |
||||
|
setAttribute() {} |
||||
|
removeAttribute() {} |
||||
|
querySelector() { return null; } |
||||
|
querySelectorAll() { return []; } |
||||
|
} |
||||
|
const encoder = new TextEncoder(); |
||||
|
const source = { documentId: '1234567890123456789', title: 'same answer', sourceName: null, chunkIndex: null, score: null, snippet: null }; |
||||
|
const event = (value: unknown) => `data: ${JSON.stringify(value)}\n\n`; |
||||
|
const answer = event({ choices: [{ delta: { content: 'answer' } }] }); |
||||
|
const sources = event({ choices: [], sources: [source] }); |
||||
|
let clearButton: ElementStub; |
||||
|
let sender: ElementStub; |
||||
|
|
||||
|
function setup(streaming: boolean) { |
||||
|
const config = parseConfig({ integrateId: '42', requestDomain: 'https://example.test', streaming, suggestions: false })!; |
||||
|
config.chatId = 'conversation'; |
||||
|
setApiConfig(config); |
||||
|
clearButton = new ElementStub(); |
||||
|
sender = new ElementStub(); |
||||
|
initChat(config, { |
||||
|
messagesContainer: new ElementStub(), inputEl: sender, clearBtn: clearButton, |
||||
|
categorySelect: null, roleSelect: null, historyPanel: new ElementStub(), |
||||
|
welcomeEl: new ElementStub(), newMsgBtn: new ElementStub(), searchInput: null, |
||||
|
ariaLiveEl: new ElementStub(), showLoading: vi.fn(), hideLoading: vi.fn(), |
||||
|
} as unknown as Parameters<typeof initChat>[1]); |
||||
|
clearButton.dispatchEvent(new Event('click')); |
||||
|
vi.clearAllMocks(); |
||||
|
} |
||||
|
function stream(body: string) { |
||||
|
return new Response(new ReadableStream<Uint8Array>({ start(c) { c.enqueue(encoder.encode(body)); c.close(); } })); |
||||
|
} |
||||
|
|
||||
|
beforeEach(() => { vi.clearAllMocks(); }); |
||||
|
afterEach(() => { clearApiConfig(); vi.unstubAllGlobals(); }); |
||||
|
|
||||
|
describe('SDK message references belong to the producing request', () => { |
||||
|
it.each([true, false])('makes one answer request and persists/renders its sources (streaming=%s)', async streaming => { |
||||
|
setup(streaming); |
||||
|
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(streaming |
||||
|
? stream(answer + sources + 'data: [DONE]\n\n') |
||||
|
: Response.json({ text: 'answer', mcpEvents: [], suggestions: [], sources: [source] }))); |
||||
|
await sendQuickReply('question'); |
||||
|
expect(fetch).toHaveBeenCalledTimes(1); |
||||
|
expect(new URL(vi.mocked(fetch).mock.calls[0]![0] as string).pathname).toBe(streaming ? '/ai/chat/stream' : '/ai/chat/result'); |
||||
|
const saved = vi.mocked(saveMessages).mock.calls.at(-1)![1]; |
||||
|
expect(saved.at(-1)).toMatchObject({ role: 'ai', content: 'answer', sources: [source] }); |
||||
|
const wrapper = streaming |
||||
|
? vi.mocked(createEmptyAIBubble).mock.results[0]!.value.wrapper |
||||
|
: vi.mocked(renderAIBubble).mock.results[0]!.value; |
||||
|
expect(renderSources).toHaveBeenLastCalledWith(wrapper, [source]); |
||||
|
}); |
||||
|
|
||||
|
it('does not turn an empty completed stream into a second answer request', async () => { |
||||
|
setup(true); |
||||
|
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(stream('data: [DONE]\n\n'))); |
||||
|
await sendQuickReply('question'); |
||||
|
expect(fetch).toHaveBeenCalledTimes(1); |
||||
|
}); |
||||
|
|
||||
|
it('ignores a delayed synchronous answer after a new conversation begins', async () => { |
||||
|
setup(false); |
||||
|
let reply!: (response: Response) => void; |
||||
|
vi.stubGlobal('fetch', vi.fn(() => new Promise<Response>(resolve => { reply = resolve; }))); |
||||
|
const pending = sendQuickReply('old question'); |
||||
|
clearButton.dispatchEvent(new Event('click')); |
||||
|
reply(Response.json({ text: 'old answer', mcpEvents: [], suggestions: [], sources: [source] })); |
||||
|
await pending; |
||||
|
expect(saveMessages).not.toHaveBeenCalled(); |
||||
|
expect(renderSources).not.toHaveBeenCalled(); |
||||
|
expect(renderAIBubble).not.toHaveBeenCalled(); |
||||
|
}); |
||||
|
|
||||
|
it('cancels in-flight old streams without adding their sources to a new conversation', async () => { |
||||
|
setup(true); |
||||
|
const cancel = vi.fn(); |
||||
|
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(new ReadableStream<Uint8Array>({ |
||||
|
start(controller) { controller.enqueue(encoder.encode(answer)); }, cancel, |
||||
|
})))); |
||||
|
const pending = sendQuickReply('old question'); |
||||
|
await vi.waitFor(() => expect(createEmptyAIBubble).toHaveBeenCalledTimes(1)); |
||||
|
clearButton.dispatchEvent(new Event('click')); |
||||
|
await pending; |
||||
|
expect(cancel).toHaveBeenCalledTimes(1); |
||||
|
expect(saveMessages).not.toHaveBeenCalled(); |
||||
|
expect(renderSources).not.toHaveBeenCalled(); |
||||
|
}); |
||||
|
|
||||
|
it('retains partial text but removes references on user cancellation', async () => { |
||||
|
setup(true); |
||||
|
vi.stubGlobal('fetch', vi.fn().mockResolvedValue(new Response(new ReadableStream<Uint8Array>({ |
||||
|
start(controller) { controller.enqueue(encoder.encode(answer + sources)); }, |
||||
|
})))); |
||||
|
const pending = sendQuickReply('question'); |
||||
|
await vi.waitFor(() => expect(renderSources).toHaveBeenCalled()); |
||||
|
sender.dispatchEvent(new Event('stop')); |
||||
|
await pending; |
||||
|
expect(vi.mocked(saveMessages).mock.calls.at(-1)![1].at(-1)).toMatchObject({ content: 'answer', sources: [] }); |
||||
|
expect(renderSources).toHaveBeenLastCalledWith(expect.anything(), []); |
||||
|
}); |
||||
|
}); |
||||
@ -1,206 +1,139 @@ |
|||||
/** |
|
||||
* SSE 流式读取工具 —— 从 utils.js 原封不动搬移 |
|
||||
* |
|
||||
* 统一处理 Flux<String> / ServerSentEvent / SseEmitter 三种 SSE 接口。 |
|
||||
* 此文件使用原生 fetch + ReadableStream API,零框架依赖。 |
|
||||
*/ |
|
||||
import type { SSECallbacks } from '@/types/sse' |
|
||||
|
/** Shared SSE reader for plain text, OpenAI chunks and tool events. */ |
||||
|
import type { SSECallbacks, SourceReference } from '@/types/sse' |
||||
import { getToken } from '@/utils/token' |
import { getToken } from '@/utils/token' |
||||
|
|
||||
/** 构建带认证的请求头 */ |
|
||||
|
/** Preserve explicit SDK authorization instead of replacing it with the admin token. */ |
||||
function authHeaders(extra?: Record<string, string>): Record<string, string> { |
function authHeaders(extra?: Record<string, string>): Record<string, string> { |
||||
const token = getToken() |
const token = getToken() |
||||
const base: Record<string, string> = extra ? { ...extra } : {} |
|
||||
// 仅当调用方未显式提供 Authorization 时才注入管理后台 token,避免覆盖调用方传入的鉴权头(如测试面板的 SDK Token)
|
|
||||
if (token && !base['Authorization']) base['Authorization'] = `Bearer ${token}` |
|
||||
return base |
|
||||
|
const headers = { ...extra } |
||||
|
if (token && !headers['Authorization']) headers['Authorization'] = `Bearer ${token}` |
||||
|
return headers |
||||
} |
} |
||||
|
|
||||
/** 尝试从 OpenAI Chat Completions chunk 中提取 delta.content;返回 undefined 表示非 OpenAI 格式(回退纯文本) */ |
|
||||
function extractOpenAIDelta(text: string): string | undefined { |
|
||||
let obj: any |
|
||||
try { obj = JSON.parse(text) } catch { return undefined } |
|
||||
// OpenAI 错误 chunk({"error":{"message":...}}):提取错误信息作为内容,避免原始 JSON 泄漏到对话框
|
|
||||
if (obj && obj.error) { |
|
||||
const msg = obj.error.message || obj.error.type |
|
||||
return typeof msg === 'string' && msg ? msg : '服务异常' |
|
||||
|
function isSourceReference(value: unknown): value is SourceReference { |
||||
|
if (!value || typeof value !== 'object') return false |
||||
|
const nullableString = (v: unknown) => v === null || typeof v === 'string' |
||||
|
const nullableNumber = (v: unknown) => v === null || typeof v === 'number' |
||||
|
return 'documentId' in value && nullableString(value.documentId) |
||||
|
&& 'title' in value && nullableString(value.title) |
||||
|
&& 'sourceName' in value && nullableString(value.sourceName) |
||||
|
&& 'chunkIndex' in value && nullableNumber(value.chunkIndex) |
||||
|
&& 'score' in value && nullableNumber(value.score) |
||||
|
&& 'snippet' in value && nullableString(value.snippet) |
||||
|
} |
||||
|
|
||||
|
/** Returns false only for legacy plain-text payloads. Metadata never enters message text. */ |
||||
|
function dispatchOpenAI(text: string, handlers: SSECallbacks): boolean { |
||||
|
let value: unknown |
||||
|
try { value = JSON.parse(text) } catch { return false } |
||||
|
if (!value || typeof value !== 'object') return false |
||||
|
if ('error' in value && value.error && typeof value.error === 'object') { |
||||
|
const error = value.error |
||||
|
const message = 'message' in error ? error.message : 'type' in error ? error.type : undefined |
||||
|
handlers.onMessage?.(typeof message === 'string' && message ? message : '服务异常') |
||||
|
return true |
||||
} |
} |
||||
if (obj && Array.isArray(obj.choices)) { |
|
||||
const delta = obj.choices[0]?.delta |
|
||||
if (delta && typeof delta.content === 'string') return delta.content |
|
||||
return '' // OpenAI 形状但无 content(role/finish_reason 空 chunk)→ 跳过
|
|
||||
|
if (!('choices' in value) || !Array.isArray(value.choices)) return false |
||||
|
if ('sources' in value && Array.isArray(value.sources)) { |
||||
|
if (!value.sources.every(isSourceReference)) throw new Error('无效的引用来源数据') |
||||
|
handlers.onSources?.(value.sources) |
||||
} |
} |
||||
return undefined // 是 JSON 但非 OpenAI 形状 → 回退纯文本
|
|
||||
|
const content = value.choices[0]?.delta?.content |
||||
|
if (typeof content === 'string' && content) handlers.onMessage?.(content) |
||||
|
return true |
||||
} |
} |
||||
|
|
||||
/** |
|
||||
* 通用 SSE 流式读取 —— 统一处理 Flux<String> / ServerSentEvent / SseEmitter 三种 SSE 接口 |
|
||||
* |
|
||||
* @param url 请求地址 |
|
||||
* @param onChunk 每收到一段文本的回调 |
|
||||
* @param onDone 流结束的回调 |
|
||||
* @param headers 额外请求头 |
|
||||
* @param signal AbortSignal 用于取消请求(组件卸载时必须传入以释放网络资源) |
|
||||
*/ |
|
||||
export async function readSSEStream( |
|
||||
|
/** Text-only callers share the same framing, completion and cleanup semantics. */ |
||||
|
export function readSSEStream( |
||||
url: string, |
url: string, |
||||
onChunk: (text: string) => void, |
onChunk: (text: string) => void, |
||||
onDone?: () => void, |
onDone?: () => void, |
||||
headers?: Record<string, string>, |
headers?: Record<string, string>, |
||||
signal?: AbortSignal |
|
||||
|
signal?: AbortSignal, |
||||
): Promise<void> { |
): Promise<void> { |
||||
const res = await fetch(url, { headers: authHeaders(headers), signal, credentials: 'include' }) |
|
||||
if (!res.ok) throw new Error('HTTP ' + res.status) |
|
||||
const reader = res.body!.getReader() |
|
||||
const decoder = new TextDecoder() |
|
||||
let buffer = '' |
|
||||
// SSE 规范:同一事件内的多行 data: 字段用 \n 拼接,事件间用空行分隔
|
|
||||
let eventDataLines: string[] = [] |
|
||||
let currentEvent = 'message' |
|
||||
|
|
||||
const flushEvent = () => { |
|
||||
if (eventDataLines.length === 0) return |
|
||||
const text = eventDataLines.join('\n') |
|
||||
eventDataLines = [] |
|
||||
const ev = currentEvent |
|
||||
currentEvent = 'message' |
|
||||
if (text === '[DONE]') return |
|
||||
// 跳过 status / faq 等系统事件,不显示在对话框中
|
|
||||
if (ev === 'status') return |
|
||||
// OpenAI Chat Completions 格式:提取 delta.content;无 content 的空 chunk 跳过
|
|
||||
const extracted = extractOpenAIDelta(text) |
|
||||
if (extracted !== undefined) { |
|
||||
if (extracted) onChunk(extracted) |
|
||||
return |
|
||||
} |
|
||||
// 空事件视为 LLM 流式输出的换行符(Spring 将 "\n" 编码为单条空 data: 事件)
|
|
||||
onChunk(text || '\n') |
|
||||
} |
|
||||
|
|
||||
while (true) { |
|
||||
const { done, value } = await reader.read() |
|
||||
if (done) break |
|
||||
buffer += decoder.decode(value, { stream: true }) |
|
||||
const lines = buffer.split('\n') |
|
||||
buffer = lines.pop() || '' |
|
||||
for (let line of lines) { |
|
||||
// 兼容 \r\n 行结束符
|
|
||||
if (line.endsWith('\r')) line = line.slice(0, -1) |
|
||||
if (line === '') { |
|
||||
// 空行 = SSE 事件边界
|
|
||||
flushEvent() |
|
||||
} else if (line.startsWith('event:')) { |
|
||||
// 记录事件类型,用于 flushEvent 时过滤系统事件
|
|
||||
currentEvent = line.slice(6).trim() |
|
||||
} else if (line.startsWith('data:')) { |
|
||||
// 累积 data 字段;仅剥离 data: 后的一个可选空格,保留 markdown 列表缩进
|
|
||||
let data = line.slice(5) |
|
||||
if (data.startsWith(' ')) data = data.slice(1) |
|
||||
eventDataLines.push(data) |
|
||||
} else if (!line.startsWith(':')) { |
|
||||
// Flux<String> 模式(非标准 SSE),先把已累积的 SSE 事件 flush 再处理
|
|
||||
flushEvent() |
|
||||
if (line.trim()) onChunk(line) |
|
||||
} |
|
||||
} |
|
||||
} |
|
||||
// 流结束,flush 末尾未以空行收尾的事件
|
|
||||
flushEvent() |
|
||||
if (onDone) onDone() |
|
||||
|
return readSSEStreamWithEvents(url, { onMessage: onChunk, onDone }, headers, signal) |
||||
} |
} |
||||
|
|
||||
/** |
|
||||
* 增强版 SSE 流式读取(支持事件类型分发) |
|
||||
* 解析 SSE 标准的 event: 字段,将不同类型事件分发到对应回调。 |
|
||||
* |
|
||||
* @param url 请求地址 |
|
||||
* @param handlers 回调对象: |
|
||||
* - onMessage(chunk): 普通文本内容(event: message 或无 event 的 data) |
|
||||
* - onToolCallStart(data): 工具调用开始(event: tool_call_start) |
|
||||
* - onToolCallResult(data): 工具调用结果(event: tool_call_result) |
|
||||
* - onError(data): 错误事件(event: error) |
|
||||
* - onDone(): 流结束 |
|
||||
* @param headers 额外请求头 |
|
||||
* @param signal AbortSignal 用于取消请求(组件卸载时必须传入以释放网络资源) |
|
||||
*/ |
|
||||
|
/** Read complete SSE events; [DONE] terminates immediately without waiting for network EOF. */ |
||||
export async function readSSEStreamWithEvents( |
export async function readSSEStreamWithEvents( |
||||
url: string, |
url: string, |
||||
handlers: SSECallbacks, |
handlers: SSECallbacks, |
||||
headers?: Record<string, string>, |
headers?: Record<string, string>, |
||||
signal?: AbortSignal |
|
||||
|
signal?: AbortSignal, |
||||
): Promise<void> { |
): Promise<void> { |
||||
const { onMessage, onToolCallStart, onToolCallResult, onError, onDone } = handlers |
|
||||
const res = await fetch(url, { headers: authHeaders(headers), signal, credentials: 'include' }) |
const res = await fetch(url, { headers: authHeaders(headers), signal, credentials: 'include' }) |
||||
if (!res.ok) throw new Error('HTTP ' + res.status) |
|
||||
const reader = res.body!.getReader() |
|
||||
|
if (!res.ok) throw new Error((await res.text()) || 'HTTP ' + res.status) |
||||
|
if (!res.body) throw new Error('响应流为空') |
||||
|
const reader = res.body.getReader() |
||||
const decoder = new TextDecoder() |
const decoder = new TextDecoder() |
||||
let buffer = '' |
let buffer = '' |
||||
let currentEvent = 'message' |
|
||||
// SSE 规范:同一事件内的多行 data: 字段用 \n 拼接
|
|
||||
let eventDataLines: string[] = [] |
|
||||
|
let event = 'message' |
||||
|
let data: string[] = [] |
||||
|
let completed = false |
||||
|
let eof = false |
||||
|
|
||||
const flushEvent = () => { |
const flushEvent = () => { |
||||
if (eventDataLines.length === 0) return |
|
||||
const ev = currentEvent |
|
||||
currentEvent = 'message' |
|
||||
const raw = eventDataLines.join('\n') |
|
||||
eventDataLines = [] |
|
||||
if (raw === '[DONE]') return |
|
||||
// 空事件视为换行符(仅 message 事件);JSON 事件(工具调用等)空内容不应出现
|
|
||||
const text = raw || '\n' |
|
||||
|
|
||||
switch (ev) { |
|
||||
case 'tool_call_start': |
|
||||
if (onToolCallStart) onToolCallStart(JSON.parse(text)) |
|
||||
break |
|
||||
case 'tool_call_result': |
|
||||
if (onToolCallResult) onToolCallResult(JSON.parse(text)) |
|
||||
break |
|
||||
case 'error': |
|
||||
if (onError) onError(JSON.parse(text)) |
|
||||
break |
|
||||
case 'status': |
|
||||
// 系统状态事件(generating / faq_hit 等),不显示在对话框中
|
|
||||
break |
|
||||
case 'message': |
|
||||
default: { |
|
||||
// OpenAI Chat Completions 格式:提取 delta.content;无 content 的空 chunk 跳过
|
|
||||
const extracted = extractOpenAIDelta(text) |
|
||||
if (extracted !== undefined) { |
|
||||
if (extracted && onMessage) onMessage(extracted) |
|
||||
} else if (onMessage) { |
|
||||
onMessage(text) |
|
||||
} |
|
||||
break |
|
||||
} |
|
||||
|
const currentEvent = event |
||||
|
event = 'message' |
||||
|
if (!data.length) return |
||||
|
const raw = data.join('\n') |
||||
|
data = [] |
||||
|
if (raw === '[DONE]') { |
||||
|
completed = true |
||||
|
return |
||||
|
} |
||||
|
switch (currentEvent) { |
||||
|
case 'status': return |
||||
|
case 'tool_call_start': handlers.onToolCallStart?.(JSON.parse(raw)); return |
||||
|
case 'tool_call_result': handlers.onToolCallResult?.(JSON.parse(raw)); return |
||||
|
case 'error': handlers.onError?.(JSON.parse(raw)); return |
||||
|
default: |
||||
|
if (!dispatchOpenAI(raw, handlers)) handlers.onMessage?.(raw || '\n') |
||||
} |
} |
||||
} |
} |
||||
|
|
||||
while (true) { |
|
||||
const { done, value } = await reader.read() |
|
||||
if (done) break |
|
||||
buffer += decoder.decode(value, { stream: true }) |
|
||||
const lines = buffer.split('\n') |
|
||||
buffer = lines.pop() || '' |
|
||||
|
const consumeLine = (raw: string) => { |
||||
|
const line = raw.endsWith('\r') ? raw.slice(0, -1) : raw |
||||
|
if (!line) { |
||||
|
flushEvent() |
||||
|
} else if (line.startsWith('event:')) { |
||||
|
event = line.slice(6).trim() |
||||
|
} else if (line === 'data' || line.startsWith('data:')) { |
||||
|
const value = line === 'data' ? '' : line.slice(5) |
||||
|
data.push(value.startsWith(' ') ? value.slice(1) : value) |
||||
|
} else if (!line.startsWith(':') && !line.startsWith('id:') && !line.startsWith('retry:')) { |
||||
|
// Legacy Flux<String> bodies may contain unframed lines.
|
||||
|
flushEvent() |
||||
|
if (!completed && line.trim()) handlers.onMessage?.(line) |
||||
|
} |
||||
|
} |
||||
|
|
||||
for (let line of lines) { |
|
||||
if (line.endsWith('\r')) line = line.slice(0, -1) |
|
||||
if (line === '') { |
|
||||
// 空行 = SSE 事件边界
|
|
||||
flushEvent() |
|
||||
} else if (line.startsWith('event:')) { |
|
||||
currentEvent = line.slice(6).trim() |
|
||||
} else if (line.startsWith('data:')) { |
|
||||
let data = line.slice(5) |
|
||||
if (data.startsWith(' ')) data = data.slice(1) |
|
||||
eventDataLines.push(data) |
|
||||
} else if (!line.startsWith(':')) { |
|
||||
// Flux<String> 模式(非标准 SSE)
|
|
||||
flushEvent() |
|
||||
if (line.trim() && onMessage) onMessage(line) |
|
||||
|
try { |
||||
|
while (!completed) { |
||||
|
signal?.throwIfAborted() |
||||
|
const result = await reader.read() |
||||
|
signal?.throwIfAborted() |
||||
|
eof = result.done |
||||
|
buffer += eof ? decoder.decode() : decoder.decode(result.value, { stream: true }) |
||||
|
let start = 0 |
||||
|
let end: number |
||||
|
while (!completed && (end = buffer.indexOf('\n', start)) !== -1) { |
||||
|
consumeLine(buffer.slice(start, end)) |
||||
|
start = end + 1 |
||||
} |
} |
||||
|
buffer = buffer.slice(start) |
||||
|
if (eof) { |
||||
|
if (!completed && buffer) consumeLine(buffer) |
||||
|
if (!completed) flushEvent() |
||||
|
break |
||||
|
} |
||||
|
} |
||||
|
handlers.onDone?.() |
||||
|
} finally { |
||||
|
// Cancellation may itself reject or stall; it must not delay DONE or mask the original error.
|
||||
|
if (!eof) { |
||||
|
try { void reader.cancel().catch(() => {}) } catch { /* Preserve the stream outcome. */ } |
||||
} |
} |
||||
|
try { reader.releaseLock() } catch { /* Preserve the stream outcome. */ } |
||||
} |
} |
||||
flushEvent() |
|
||||
if (onDone) onDone() |
|
||||
} |
} |
||||
@ -0,0 +1,164 @@ |
|||||
|
import assert from 'node:assert/strict' |
||||
|
import { afterEach, beforeEach, test } from 'node:test' |
||||
|
import { fileURLToPath } from 'node:url' |
||||
|
import { build } from 'esbuild' |
||||
|
|
||||
|
// Exercise the production TypeScript with Vite's existing esbuild dependency, no test framework. |
||||
|
const root = fileURLToPath(new URL('../', import.meta.url)) |
||||
|
async function loadModule(entry) { |
||||
|
const result = await build({ |
||||
|
entryPoints: [root + entry], bundle: true, write: false, format: 'esm', platform: 'node', |
||||
|
alias: { '@': root + 'src' }, |
||||
|
}) |
||||
|
return import('data:text/javascript;base64,' + Buffer.from(result.outputFiles[0].text).toString('base64')) |
||||
|
} |
||||
|
const { readSSEStream, readSSEStreamWithEvents } = await loadModule('src/utils/sse.ts') |
||||
|
const { chatSync, fetchChatResult } = await loadModule('src/api/chat.ts') |
||||
|
const originalFetch = globalThis.fetch |
||||
|
const originalStorage = Object.getOwnPropertyDescriptor(globalThis, 'localStorage') |
||||
|
|
||||
|
beforeEach(() => { |
||||
|
Object.defineProperty(globalThis, 'localStorage', { |
||||
|
configurable: true, value: { getItem: () => 'admin-token' }, |
||||
|
}) |
||||
|
}) |
||||
|
afterEach(() => { |
||||
|
globalThis.fetch = originalFetch |
||||
|
if (originalStorage) Object.defineProperty(globalThis, 'localStorage', originalStorage) |
||||
|
else delete globalThis.localStorage |
||||
|
}) |
||||
|
|
||||
|
const sources = [{ |
||||
|
documentId: '9223372036854775806', title: '报销制度', sourceName: 'policy.pdf', |
||||
|
chunkIndex: 2, score: 0.125, snippet: '申请应在三十天内提交。', |
||||
|
}] |
||||
|
const metadata = JSON.stringify({ object: 'chat.completion.chunk', choices: [], sources }) |
||||
|
const answer = JSON.stringify({ choices: [{ delta: { content: '答复。' } }] }) |
||||
|
|
||||
|
function responseFor(text, { open = false, fragment = false, cancel } = {}) { |
||||
|
const bytes = new TextEncoder().encode(text) |
||||
|
const body = new ReadableStream({ |
||||
|
start(controller) { |
||||
|
if (fragment) for (const byte of bytes) controller.enqueue(Uint8Array.of(byte)) |
||||
|
else controller.enqueue(bytes) |
||||
|
if (!open) controller.close() |
||||
|
}, |
||||
|
cancel, |
||||
|
}) |
||||
|
return new Response(body) |
||||
|
} |
||||
|
|
||||
|
test('fragmented CRLF metadata preserves ID/snippet and DONE completes without EOF or cancellation settlement', { timeout: 1000 }, async () => { |
||||
|
let cancelled = 0 |
||||
|
const response = responseFor( |
||||
|
`data: ${answer}\r\n\r\ndata: ${metadata}\r\n\r\ndata: [DONE]\r\n\r\ndata: ignored\r\n\r\n`, |
||||
|
{ open: true, fragment: true, cancel() { cancelled++; return new Promise(() => {}) } }, |
||||
|
) |
||||
|
const requests = [] |
||||
|
globalThis.fetch = async (...args) => { requests.push(args); return response } |
||||
|
const events = [] |
||||
|
await readSSEStreamWithEvents('/ai/chat/stream', { |
||||
|
onMessage: text => events.push(['text', text]), |
||||
|
onSources: value => events.push(['sources', value]), |
||||
|
onDone: () => events.push(['done']), |
||||
|
}, { Authorization: 'Bearer sdk-token' }) |
||||
|
assert.deepEqual(events, [['text', '答复。'], ['sources', sources], ['done']]) |
||||
|
assert.equal(requests.length, 1) |
||||
|
assert.equal(requests[0][1].headers.Authorization, 'Bearer sdk-token') |
||||
|
assert.equal(cancelled, 1) |
||||
|
assert.equal(response.body.locked, false) |
||||
|
}) |
||||
|
|
||||
|
test('text-only facade uses identical DONE handling and never leaks empty choices or sources', { timeout: 1000 }, async () => { |
||||
|
globalThis.fetch = async () => responseFor( |
||||
|
`data: ${metadata}\n\ndata: ${answer}\n\ndata: [DONE]\n\ndata: [DONE]\n\ndata: trailing\n\n`, |
||||
|
{ open: true }, |
||||
|
) |
||||
|
const text = [] |
||||
|
let completed = 0 |
||||
|
await readSSEStream('/ai/chat/stream', value => text.push(value), () => completed++) |
||||
|
assert.deepEqual(text, ['答复。']) |
||||
|
assert.equal(completed, 1) |
||||
|
}) |
||||
|
|
||||
|
test('EOF residuals, multiline data, tool events, status and legacy text retain their semantics', async () => { |
||||
|
globalThis.fetch = async () => responseFor( |
||||
|
': heartbeat\r\nid: event-1\r\nretry: 1000\r\nevent: status\r\ndata: generating\r\n\r\n' |
||||
|
+ 'event: tool_call_start\r\ndata: {"tool":"lookup"}\r\n\r\n' |
||||
|
+ 'event: tool_call_result\r\ndata: {"tool":"lookup","result":"ok"}\r\n\r\n' |
||||
|
+ 'event: error\r\ndata: {"message":"tool unavailable"}\r\n\r\n' |
||||
|
+ 'data: first\r\ndata: indented\r\n\r\ndata:\r\n\r\nlegacy\r\ndata: 尾部', |
||||
|
{ fragment: true }, |
||||
|
) |
||||
|
const events = [] |
||||
|
await readSSEStreamWithEvents('/stream', { |
||||
|
onMessage: text => events.push(text), |
||||
|
onToolCallStart: value => events.push(value), |
||||
|
onToolCallResult: value => events.push(value), |
||||
|
onError: value => events.push(value), |
||||
|
onDone: () => events.push('done'), |
||||
|
}) |
||||
|
assert.deepEqual(events, [ |
||||
|
{ tool: 'lookup' }, { tool: 'lookup', result: 'ok' }, { message: 'tool unavailable' }, |
||||
|
'first\n indented', '\n', 'legacy', '尾部', 'done', |
||||
|
]) |
||||
|
}) |
||||
|
|
||||
|
test('EOF flushes an unterminated plain-text line and an unterminated DONE event', async () => { |
||||
|
for (const body of ['legacy tail', 'data: [DONE]']) { |
||||
|
globalThis.fetch = async () => responseFor(body) |
||||
|
const text = [] |
||||
|
let done = 0 |
||||
|
await readSSEStream('/stream', value => text.push(value), () => done++) |
||||
|
assert.deepEqual(text, body.startsWith('data:') ? [] : ['legacy tail']) |
||||
|
assert.equal(done, 1) |
||||
|
} |
||||
|
}) |
||||
|
|
||||
|
test('cleanup failure cannot replace callback errors or call onDone on an error', async () => { |
||||
|
const original = new Error('consumer failed') |
||||
|
const response = responseFor(`data: ${answer}\n\n`, { |
||||
|
open: true, cancel() { throw new Error('cleanup failed') }, |
||||
|
}) |
||||
|
globalThis.fetch = async () => response |
||||
|
let done = 0 |
||||
|
await assert.rejects(readSSEStream('/stream', () => { throw original }, () => done++), error => error === original) |
||||
|
assert.equal(done, 0) |
||||
|
assert.equal(response.body.locked, false) |
||||
|
}) |
||||
|
|
||||
|
test('empty sources for ordinary answers are delivered separately from content', async () => { |
||||
|
globalThis.fetch = async () => responseFor('data: {"choices":[],"sources":[]}\n\ndata: [DONE]\n\n') |
||||
|
const values = [] |
||||
|
await readSSEStreamWithEvents('/stream', { |
||||
|
onMessage: () => assert.fail('metadata is not text'), |
||||
|
onSources: value => values.push(value), |
||||
|
}) |
||||
|
assert.deepEqual(values, [[]]) |
||||
|
}) |
||||
|
|
||||
|
test('synchronous chat makes one result request and retains direct answer, sources, explicit strategy and signal', async () => { |
||||
|
const result = { text: '答复。', mcpEvents: [], suggestions: [], sources } |
||||
|
const calls = [] |
||||
|
globalThis.fetch = async (...args) => { calls.push(args); return Response.json(result) } |
||||
|
const controller = new AbortController() |
||||
|
assert.deepEqual(await chatSync('费用?', 'conversation-1', { |
||||
|
enableRag: true, roleId: '9223372036854775806', rewriteStrategy: 'MULTI_QUERY', categoryIds: ['123'], |
||||
|
}, controller.signal), result) |
||||
|
assert.equal(calls.length, 1) |
||||
|
const url = new URL(calls[0][0], 'https://test.invalid') |
||||
|
assert.equal(url.pathname, '/ai/chat/result') |
||||
|
assert.equal(url.searchParams.get('rewriteStrategy'), 'MULTI_QUERY') |
||||
|
assert.equal(url.searchParams.get('roleId'), '9223372036854775806') |
||||
|
assert.equal(calls[0][1].signal, controller.signal) |
||||
|
assert.equal(calls[0][1].headers.Authorization, 'Bearer admin-token') |
||||
|
}) |
||||
|
|
||||
|
test('synchronous SDK response uses the same direct contract and surfaces server errors', async () => { |
||||
|
const result = { text: '普通回答', mcpEvents: [], suggestions: [], sources: [] } |
||||
|
globalThis.fetch = async () => Response.json(result) |
||||
|
assert.deepEqual(await fetchChatResult('https://sdk.invalid/ai/chat/result', {}), result) |
||||
|
globalThis.fetch = async () => new Response('角色无访问权限', { status: 403 }) |
||||
|
await assert.rejects(fetchChatResult('/ai/chat/result', {}), /角色无访问权限/) |
||||
|
await assert.rejects(readSSEStream('/ai/chat/stream', () => {}), /角色无访问权限/) |
||||
|
}) |
||||
@ -0,0 +1,59 @@ |
|||||
|
package com.wok.supportbot.app; |
||||
|
|
||||
|
import com.fasterxml.jackson.annotation.JsonInclude; |
||||
|
import org.springframework.ai.document.Document; |
||||
|
|
||||
|
import java.util.List; |
||||
|
import java.util.Map; |
||||
|
|
||||
|
/** Public citation metadata from the documents actually used by this answer. */ |
||||
|
@JsonInclude(JsonInclude.Include.ALWAYS) |
||||
|
public record SourceReference(String documentId, String title, String sourceName, |
||||
|
Integer chunkIndex, Double score, String snippet) { |
||||
|
|
||||
|
public SourceReference { |
||||
|
if (snippet != null && snippet.length() > 160) { |
||||
|
int end = Character.isHighSurrogate(snippet.charAt(158)) ? 158 : 159; |
||||
|
snippet = snippet.substring(0, end) + "…"; |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
public static List<SourceReference> fromDocuments(List<Document> documents) { |
||||
|
if (documents == null || documents.isEmpty()) { |
||||
|
return List.of(); |
||||
|
} |
||||
|
return documents.stream().map(SourceReference::fromDocument).toList(); |
||||
|
} |
||||
|
|
||||
|
private static SourceReference fromDocument(Document document) { |
||||
|
Map<String, Object> metadata = document.getMetadata(); |
||||
|
// Preserve the published source endpoint's distance semantics; other retrievers expose score. |
||||
|
Object score = metadata.get("distance") != null ? metadata.get("distance") : metadata.get("score"); |
||||
|
return new SourceReference(stringValue(metadata.get("documentId")), |
||||
|
stringValue(metadata.get("title")), stringValue(metadata.get("sourceName")), |
||||
|
integerValue(metadata.get("chunkIndex")), doubleValue(score), document.getText()); |
||||
|
} |
||||
|
|
||||
|
private static String stringValue(Object value) { |
||||
|
return value == null ? null : value.toString(); |
||||
|
} |
||||
|
|
||||
|
private static Integer integerValue(Object value) { |
||||
|
if (value == null) return null; |
||||
|
try { |
||||
|
return Integer.valueOf(value.toString()); |
||||
|
} catch (NumberFormatException ignored) { |
||||
|
return null; |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
private static Double doubleValue(Object value) { |
||||
|
if (value == null) return null; |
||||
|
try { |
||||
|
double number = Double.parseDouble(value.toString()); |
||||
|
return Double.isFinite(number) ? number : null; |
||||
|
} catch (NumberFormatException ignored) { |
||||
|
return null; |
||||
|
} |
||||
|
} |
||||
|
} |
||||
@ -1,114 +0,0 @@ |
|||||
package com.wok.supportbot.service; |
|
||||
|
|
||||
import com.wok.supportbot.config.ChatModelFactory; |
|
||||
import lombok.AllArgsConstructor; |
|
||||
import lombok.Data; |
|
||||
import lombok.NoArgsConstructor; |
|
||||
import lombok.extern.slf4j.Slf4j; |
|
||||
import org.springframework.ai.chat.model.ChatModel; |
|
||||
import org.springframework.ai.chat.prompt.Prompt; |
|
||||
import org.springframework.ai.converter.BeanOutputConverter; |
|
||||
import org.springframework.beans.factory.annotation.Autowired; |
|
||||
import org.springframework.stereotype.Service; |
|
||||
|
|
||||
/** |
|
||||
* LLM 意图分类路由器 |
|
||||
* 使用 ChatModel 对用户问题进行意图分类,决定后续处理流程: |
|
||||
* - FAQ: 常见问题 → FaqMatchEngine 精准匹配 |
|
||||
* - RAG: 知识库检索 → 现有 RAG 流程 |
|
||||
* - CHITCHAT: 闲聊 → 简单对话 |
|
||||
* |
|
||||
* <p>结构化输出使用 Spring AI 标准组件 {@link BeanOutputConverter}: |
|
||||
* 由它把 JSON Schema 指令追加进 Prompt,并把模型返回的 JSON 反序列化为 {@link IntentResult}, |
|
||||
* 不再手写正则解析。{@code ChatClient.entity(...)} 内部即是同一套机制。 |
|
||||
*/ |
|
||||
@Service |
|
||||
@Slf4j |
|
||||
public class IntentRouter { |
|
||||
|
|
||||
@Autowired |
|
||||
private ChatModelFactory chatModelFactory; |
|
||||
|
|
||||
/** 结构化输出转换器(无状态,可安全复用):生成 Schema 指令 + 反序列化模型响应 */ |
|
||||
private static final BeanOutputConverter<IntentResult> INTENT_CONVERTER = |
|
||||
new BeanOutputConverter<>(IntentResult.class); |
|
||||
|
|
||||
/** |
|
||||
* 意图分类 Prompt 模板。 |
|
||||
* 输出格式约束(JSON Schema)由 {@code BeanOutputConverter.getFormat()} 统一追加,模板中不再硬编码。 |
|
||||
*/ |
|
||||
private static final String INTENT_PROMPT_TEMPLATE = """ |
|
||||
你是一个意图分类器。根据用户问题,判断其属于以下哪个意图: |
|
||||
- FAQ: 常见问题,如产品功能、价格、退换货政策、服务流程等标准问答 |
|
||||
- RAG: 需要查阅文档/知识库才能回答的专业问题或细节问题 |
|
||||
- CHITCHAT: 闲聊、问候、感谢、告别等非业务话题 |
|
||||
|
|
||||
用户问题: %s |
|
||||
|
|
||||
%s |
|
||||
"""; |
|
||||
|
|
||||
// ==================== 意图结果内部类 ==================== |
|
||||
|
|
||||
/** |
|
||||
* 意图分类结果 |
|
||||
*/ |
|
||||
@Data |
|
||||
@AllArgsConstructor |
|
||||
@NoArgsConstructor |
|
||||
public static class IntentResult { |
|
||||
/** 意图类型: FAQ / RAG / CHITCHAT */ |
|
||||
private String intent; |
|
||||
/** 置信度 (0.0 ~ 1.0) */ |
|
||||
private double confidence; |
|
||||
} |
|
||||
|
|
||||
// ==================== 核心路由方法 ==================== |
|
||||
|
|
||||
/** |
|
||||
* 对用户问题进行意图分类 |
|
||||
* |
|
||||
* @param userQuestion 用户问题 |
|
||||
* @return 意图分类结果 |
|
||||
*/ |
|
||||
public IntentResult route(String userQuestion) { |
|
||||
if (userQuestion == null || userQuestion.isBlank()) { |
|
||||
return new IntentResult("RAG", 0.0); |
|
||||
} |
|
||||
|
|
||||
try { |
|
||||
ChatModel chatModel = chatModelFactory.getChatModel("CHAT"); |
|
||||
String promptText = INTENT_PROMPT_TEMPLATE.formatted(userQuestion, INTENT_CONVERTER.getFormat()); |
|
||||
|
|
||||
String response = chatModel.call(new Prompt(promptText)).getResult().getOutput().getText(); |
|
||||
log.debug("意图分类原始响应: {}", response); |
|
||||
|
|
||||
IntentResult result = INTENT_CONVERTER.convert(response); |
|
||||
if (result == null || !isValidIntent(result.getIntent())) { |
|
||||
log.warn("意图分类结果无效,降级为 RAG: rawResponse={}", abbreviate(response)); |
|
||||
return new IntentResult("RAG", 0.5); |
|
||||
} |
|
||||
return result; |
|
||||
} catch (Exception e) { |
|
||||
log.warn("意图分类失败,降级为 RAG: question={}", abbreviate(userQuestion), e); |
|
||||
return new IntentResult("RAG", 0.0); |
|
||||
} |
|
||||
} |
|
||||
|
|
||||
/** |
|
||||
* 校验意图类型是否有效(模型可能返回枚举外的值) |
|
||||
*/ |
|
||||
private boolean isValidIntent(String intent) { |
|
||||
return "FAQ".equals(intent) || "RAG".equals(intent) || "CHITCHAT".equals(intent); |
|
||||
} |
|
||||
|
|
||||
/** |
|
||||
* 日志截断,避免回显整段响应 |
|
||||
*/ |
|
||||
private static String abbreviate(String text) { |
|
||||
if (text == null) { |
|
||||
return null; |
|
||||
} |
|
||||
return text.length() > 200 ? text.substring(0, 200) + "..." : text; |
|
||||
} |
|
||||
} |
|
||||
@ -0,0 +1,250 @@ |
|||||
|
package com.wok.supportbot; |
||||
|
|
||||
|
import com.fasterxml.jackson.databind.JsonNode; |
||||
|
import com.fasterxml.jackson.databind.ObjectMapper; |
||||
|
import com.wok.supportbot.app.AssistantApp; |
||||
|
import com.wok.supportbot.app.ChatContext; |
||||
|
import com.wok.supportbot.app.ChatPipeline; |
||||
|
import com.wok.supportbot.app.ChatRequest; |
||||
|
import com.wok.supportbot.app.ChatResult; |
||||
|
import com.wok.supportbot.app.SourceReference; |
||||
|
import com.wok.supportbot.chatmemory.DatabaseChatMemory; |
||||
|
import com.wok.supportbot.config.ChatModelFactory; |
||||
|
import com.wok.supportbot.config.SimpleCircuitBreaker; |
||||
|
import com.wok.supportbot.entity.LlmCallTrace; |
||||
|
import com.wok.supportbot.service.AiModelConfigService; |
||||
|
import com.wok.supportbot.service.ContentSafetyService; |
||||
|
import com.wok.supportbot.service.LlmCallTraceService; |
||||
|
import org.junit.jupiter.api.BeforeEach; |
||||
|
import org.junit.jupiter.api.Test; |
||||
|
import org.junit.jupiter.params.ParameterizedTest; |
||||
|
import org.junit.jupiter.params.provider.ValueSource; |
||||
|
import org.mockito.ArgumentCaptor; |
||||
|
import org.springframework.ai.chat.client.ChatClient; |
||||
|
import org.springframework.ai.chat.messages.AssistantMessage; |
||||
|
import org.springframework.ai.chat.model.ChatResponse; |
||||
|
import org.springframework.ai.chat.model.Generation; |
||||
|
import org.springframework.ai.document.Document; |
||||
|
import org.springframework.test.util.ReflectionTestUtils; |
||||
|
import reactor.core.publisher.Flux; |
||||
|
|
||||
|
import java.time.Duration; |
||||
|
import java.util.ArrayList; |
||||
|
import java.util.List; |
||||
|
import java.util.Map; |
||||
|
import java.util.Optional; |
||||
|
import java.util.concurrent.atomic.AtomicBoolean; |
||||
|
|
||||
|
import static org.junit.jupiter.api.Assertions.*; |
||||
|
import static org.mockito.ArgumentMatchers.*; |
||||
|
import static org.mockito.Mockito.*; |
||||
|
|
||||
|
class AnswerTransportTests { |
||||
|
private static final ObjectMapper JSON = new ObjectMapper(); |
||||
|
private AssistantApp app; |
||||
|
private ChatPipeline pipeline; |
||||
|
private ChatClient.ChatClientRequestSpec spec; |
||||
|
private ChatClient.CallResponseSpec call; |
||||
|
private ChatClient.StreamResponseSpec stream; |
||||
|
private LlmCallTraceService traces; |
||||
|
|
||||
|
@BeforeEach |
||||
|
@SuppressWarnings("unchecked") |
||||
|
void setUp() { |
||||
|
app = new AssistantApp(mock(ChatModelFactory.class), mock(DatabaseChatMemory.class)); |
||||
|
pipeline = mock(ChatPipeline.class); |
||||
|
traces = mock(LlmCallTraceService.class); |
||||
|
ContentSafetyService safety = mock(ContentSafetyService.class); |
||||
|
when(safety.mask(any())).thenAnswer(invocation -> invocation.getArgument(0)); |
||||
|
ReflectionTestUtils.setField(app, "chatPipeline", pipeline); |
||||
|
ReflectionTestUtils.setField(app, "llmCallTraceService", traces); |
||||
|
ReflectionTestUtils.setField(app, "contentSafetyService", safety); |
||||
|
ReflectionTestUtils.setField(app, "aiModelConfigService", mock(AiModelConfigService.class)); |
||||
|
ChatClient client = mock(ChatClient.class); |
||||
|
spec = mock(ChatClient.ChatClientRequestSpec.class, RETURNS_SELF); |
||||
|
call = mock(ChatClient.CallResponseSpec.class); |
||||
|
stream = mock(ChatClient.StreamResponseSpec.class); |
||||
|
when(client.prompt()).thenReturn(spec); |
||||
|
when(spec.call()).thenReturn(call); |
||||
|
when(spec.stream()).thenReturn(stream); |
||||
|
Map<String, ChatClient> cache = (Map<String, ChatClient>) ReflectionTestUtils.getField(app, "chatClientCache"); |
||||
|
cache.put("CHAT:none", client); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(booleans = {true, false}) |
||||
|
void synchronousAnswerUsesOnlyThisBuildsDocuments(boolean rag) throws Exception { |
||||
|
ChatContext ctx = context(rag); |
||||
|
List<Document> documents = rag ? List.of(document()) : List.of(); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, documents, null)); |
||||
|
when(call.chatResponse()).thenReturn(response("回答")); |
||||
|
|
||||
|
ChatResult result = app.chatWithEvents(ctx); |
||||
|
|
||||
|
assertEquals("回答", result.text()); |
||||
|
assertEquals(SourceReference.fromDocuments(documents), result.sources()); |
||||
|
assertTrue(result.suggestions().isEmpty()); |
||||
|
JsonNode json = JSON.valueToTree(result); |
||||
|
if (rag) { |
||||
|
assertEquals("9223372036854775807", json.at("/sources/0/documentId").asText()); |
||||
|
assertTrue(json.at("/sources/0/documentId").isTextual()); |
||||
|
} |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(call).chatResponse(); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(booleans = {true, false}) |
||||
|
void completedStreamCarriesExactlyOneMetadataChunkBeforeStop(boolean rag) throws Exception { |
||||
|
ChatContext ctx = context(rag); |
||||
|
List<Document> documents = rag ? List.of(document()) : List.of(); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, documents, null)); |
||||
|
when(stream.chatResponse()).thenReturn(Flux.just(response("## "), response("回答\n"))); |
||||
|
|
||||
|
Flux<String> result = app.chatStreamOpenAi(ctx); |
||||
|
verifyNoInteractions(pipeline); |
||||
|
List<String> chunks = result.collectList().block(Duration.ofSeconds(5)); |
||||
|
assertEnvelope(chunks, SourceReference.fromDocuments(documents), "## 回答\n"); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(stream).chatResponse(); |
||||
|
ArgumentCaptor<LlmCallTrace> trace = ArgumentCaptor.forClass(LlmCallTrace.class); |
||||
|
verify(traces, timeout(1000)).recordAsync(trace.capture()); |
||||
|
assertEquals("## 回答\n", trace.getValue().getAiResponse()); |
||||
|
assertEquals("COMPLETE", trace.getValue().getStatus()); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void faqHasEmptySourcesWithoutModelInvocation() throws Exception { |
||||
|
ChatContext ctx = context(true); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(), "标准答案")); |
||||
|
assertEnvelope(app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)), List.of(), "标准答案"); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verifyNoInteractions(call, stream); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void synchronousFaqHasEmptySources() { |
||||
|
ChatContext ctx = context(true); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(), "标准答案")); |
||||
|
ChatResult result = app.chatWithEvents(ctx); |
||||
|
assertEquals("标准答案", result.text()); |
||||
|
assertTrue(result.sources().isEmpty()); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verifyNoInteractions(call, stream); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void circuitFallbackDoesNotBuildOrRetrieve() throws Exception { |
||||
|
SimpleCircuitBreaker breaker = (SimpleCircuitBreaker) ReflectionTestUtils.getField(app, "aiCircuitBreaker"); |
||||
|
for (int i = 0; i < 3; i++) breaker.recordFailure(-1L); |
||||
|
ChatContext ctx = context(true); |
||||
|
assertTrue(app.chatWithEvents(ctx).sources().isEmpty()); |
||||
|
assertEnvelope(app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)), |
||||
|
List.of(), "AI 服务暂时不可用,请稍后重试。"); |
||||
|
verifyNoInteractions(pipeline, call, stream); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void streamFailureDoesNotRepeatGenerationOrExposeUnusedSources() throws Exception { |
||||
|
ChatContext ctx = context(true); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(document()), null)); |
||||
|
when(stream.chatResponse()).thenReturn(Flux.concat(Flux.just(response("部分回答")), |
||||
|
Flux.error(new IllegalStateException("模型超时")))); |
||||
|
List<String> chunks = app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)); |
||||
|
assertEnvelope(chunks, List.of(), "部分回答抱歉,AI 服务调用失败:模型超时"); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(stream).chatResponse(); |
||||
|
verifyNoInteractions(call); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void cancellationPropagatesWithoutMetadataOrSecondRequest() throws Exception { |
||||
|
ChatContext ctx = context(true); |
||||
|
AtomicBoolean cancelled = new AtomicBoolean(); |
||||
|
when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(document()), null)); |
||||
|
when(stream.chatResponse()).thenReturn(Flux.concat(Flux.just(response("首段")), Flux.<ChatResponse>never()) |
||||
|
.doOnCancel(() -> cancelled.set(true))); |
||||
|
List<String> chunks = app.chatStreamOpenAi(ctx).take(2).collectList().block(Duration.ofSeconds(5)); |
||||
|
assertEquals(2, chunks.size()); |
||||
|
assertEquals("首段", JSON.readTree(chunks.get(1)).at("/choices/0/delta/content").asText()); |
||||
|
assertFalse(chunks.stream().anyMatch(chunk -> chunk.contains("\"sources\"") || chunk.equals("[DONE]"))); |
||||
|
verify(traces, timeout(1000)).recordAsync(argThat(trace -> "CANCEL".equals(trace.getStatus()))); |
||||
|
assertTrue(cancelled.get()); |
||||
|
verify(pipeline).buildRequest(ctx); |
||||
|
verifyNoMoreInteractions(pipeline); |
||||
|
verify(stream).chatResponse(); |
||||
|
verifyNoInteractions(call); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void sourceSerializationPreservesNullableFieldsDistanceAndSnippetBoundaries() { |
||||
|
Document doc = document(); |
||||
|
SourceReference source = SourceReference.fromDocuments(List.of(doc)).get(0); |
||||
|
assertEquals("9223372036854775807", source.documentId()); |
||||
|
assertEquals(0.25, source.score()); |
||||
|
assertEquals(2, source.chunkIndex()); |
||||
|
assertEquals(160, source.snippet().length()); |
||||
|
assertTrue(source.snippet().endsWith("…")); |
||||
|
SourceReference nullable = SourceReference.fromDocuments(List.of(new Document("short"))).get(0); |
||||
|
JsonNode json = JSON.copy().setSerializationInclusion(com.fasterxml.jackson.annotation.JsonInclude.Include.NON_NULL) |
||||
|
.valueToTree(nullable); |
||||
|
assertEquals(6, json.size()); |
||||
|
assertTrue(json.get("documentId").isNull()); |
||||
|
assertTrue(json.get("score").isNull()); |
||||
|
SourceReference unicode = new SourceReference(null, null, null, null, null, "a".repeat(158) + "😀xx"); |
||||
|
assertTrue(unicode.snippet().length() <= 160); |
||||
|
assertFalse(Character.isHighSurrogate(unicode.snippet().charAt(unicode.snippet().length() - 2))); |
||||
|
ArrayList<SourceReference> mutable = new ArrayList<>(List.of(source)); |
||||
|
ChatResult result = new ChatResult("answer", null, null, mutable); |
||||
|
mutable.clear(); |
||||
|
assertEquals(List.of(source), result.sources()); |
||||
|
} |
||||
|
|
||||
|
private static void assertEnvelope(List<String> chunks, List<SourceReference> sources, String text) throws Exception { |
||||
|
assertNotNull(chunks); |
||||
|
assertEquals("[DONE]", chunks.get(chunks.size() - 1)); |
||||
|
JsonNode metadata = JSON.readTree(chunks.get(chunks.size() - 3)); |
||||
|
JsonNode stop = JSON.readTree(chunks.get(chunks.size() - 2)); |
||||
|
assertEquals("stop", stop.at("/choices/0/finish_reason").asText()); |
||||
|
assertEquals(0, metadata.get("choices").size()); |
||||
|
assertEquals(JSON.valueToTree(sources), metadata.get("sources")); |
||||
|
StringBuilder answer = new StringBuilder(); |
||||
|
int metadataCount = 0; |
||||
|
for (String raw : chunks.subList(0, chunks.size() - 1)) { |
||||
|
JsonNode chunk = JSON.readTree(raw); |
||||
|
assertEquals("chat.completion.chunk", chunk.get("object").asText()); |
||||
|
assertEquals(metadata.get("id"), chunk.get("id")); |
||||
|
assertEquals(metadata.get("model"), chunk.get("model")); |
||||
|
assertEquals(metadata.get("created"), chunk.get("created")); |
||||
|
if (chunk.has("sources")) metadataCount++; |
||||
|
answer.append(chunk.at("/choices/0/delta/content").asText("")); |
||||
|
} |
||||
|
assertEquals(1, metadataCount); |
||||
|
assertEquals(text, answer.toString()); |
||||
|
} |
||||
|
|
||||
|
private static Document document() { |
||||
|
return new Document("知识".repeat(100), Map.of("documentId", Long.MAX_VALUE, "title", "授权文档", |
||||
|
"sourceName", "manual.pdf", "chunkIndex", "2", "distance", 0.25, "score", 0.9)); |
||||
|
} |
||||
|
|
||||
|
private static ChatContext context(boolean rag) { |
||||
|
return new ChatContext("退货", "transport-chat", "CHAT", null, null, List.of(7L), |
||||
|
"NONE", rag, false, 11L, "售后", "account", null, null); |
||||
|
} |
||||
|
|
||||
|
private static ChatRequest request(ChatContext ctx, List<Document> documents, String faq) { |
||||
|
return new ChatRequest(ctx, ctx.message(), "system", Optional.ofNullable(faq), "system", |
||||
|
documents.isEmpty() ? null : "资料", documents.size(), faq != null ? "FAQ" : "RAG", |
||||
|
"VECTOR", documents, null); |
||||
|
} |
||||
|
|
||||
|
private static ChatResponse response(String text) { |
||||
|
return new ChatResponse(List.of(new Generation(new AssistantMessage(text)))); |
||||
|
} |
||||
|
} |
||||
@ -0,0 +1,177 @@ |
|||||
|
package com.wok.supportbot; |
||||
|
|
||||
|
import com.fasterxml.jackson.databind.JsonNode; |
||||
|
import com.fasterxml.jackson.databind.ObjectMapper; |
||||
|
import com.sun.net.httpserver.HttpServer; |
||||
|
import com.wok.supportbot.config.ChatModelFactory; |
||||
|
import com.wok.supportbot.entity.AiModelConfig; |
||||
|
import com.wok.supportbot.service.AiModelConfigService; |
||||
|
import org.junit.jupiter.api.AfterEach; |
||||
|
import org.junit.jupiter.api.BeforeEach; |
||||
|
import org.junit.jupiter.api.Test; |
||||
|
import org.junit.jupiter.api.extension.ExtendWith; |
||||
|
import org.junit.jupiter.params.ParameterizedTest; |
||||
|
import org.junit.jupiter.params.provider.CsvSource; |
||||
|
import org.junit.jupiter.params.provider.NullAndEmptySource; |
||||
|
import org.junit.jupiter.params.provider.ValueSource; |
||||
|
import org.mockito.InjectMocks; |
||||
|
import org.mockito.Mock; |
||||
|
import org.mockito.junit.jupiter.MockitoExtension; |
||||
|
import org.springframework.ai.chat.model.ChatModel; |
||||
|
import org.springframework.ai.chat.prompt.Prompt; |
||||
|
import org.springframework.ai.openai.OpenAiChatOptions; |
||||
|
|
||||
|
import java.io.IOException; |
||||
|
import java.net.InetSocketAddress; |
||||
|
import java.nio.charset.StandardCharsets; |
||||
|
import java.time.Duration; |
||||
|
import java.util.Map; |
||||
|
import java.util.concurrent.BlockingQueue; |
||||
|
import java.util.concurrent.LinkedBlockingQueue; |
||||
|
import java.util.concurrent.TimeUnit; |
||||
|
|
||||
|
import static org.junit.jupiter.api.Assertions.*; |
||||
|
import static org.mockito.Mockito.when; |
||||
|
|
||||
|
@ExtendWith(MockitoExtension.class) |
||||
|
class ChatModelFactoryTests { |
||||
|
|
||||
|
private static final ObjectMapper JSON = new ObjectMapper(); |
||||
|
private final BlockingQueue<JsonNode> requests = new LinkedBlockingQueue<>(); |
||||
|
private HttpServer server; |
||||
|
|
||||
|
@Mock |
||||
|
private AiModelConfigService configService; |
||||
|
@InjectMocks |
||||
|
private ChatModelFactory factory; |
||||
|
|
||||
|
@BeforeEach |
||||
|
void startLocalEndpoint() throws IOException { |
||||
|
server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0); |
||||
|
server.createContext("/", exchange -> { |
||||
|
try (exchange) { |
||||
|
JsonNode request = JSON.readTree(exchange.getRequestBody()); |
||||
|
requests.add(request); |
||||
|
boolean streaming = request.path("stream").asBoolean(); |
||||
|
String response = streaming |
||||
|
? "data: {\"id\":\"test\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"test\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"ok\"},\"finish_reason\":null}]}\n\n" |
||||
|
+ "data: {\"id\":\"test\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"test\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n" |
||||
|
+ "data: [DONE]\n\n" |
||||
|
: "{\"id\":\"test\",\"object\":\"chat.completion\",\"created\":1,\"model\":\"test\",\"choices\":[{\"index\":0,\"message\":{\"role\":\"assistant\",\"content\":\"ok\"},\"finish_reason\":\"stop\"}]}"; |
||||
|
byte[] bytes = response.getBytes(StandardCharsets.UTF_8); |
||||
|
exchange.getResponseHeaders().set("Content-Type", streaming ? "text/event-stream" : "application/json"); |
||||
|
exchange.sendResponseHeaders(200, bytes.length); |
||||
|
exchange.getResponseBody().write(bytes); |
||||
|
} |
||||
|
}); |
||||
|
server.start(); |
||||
|
} |
||||
|
|
||||
|
@AfterEach |
||||
|
void stopLocalEndpoint() { |
||||
|
if (server != null) server.stop(0); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(strings = {"doubao-seed-2-0-pro-260215", "doubao-seed-2-0-lite-260215", |
||||
|
"doubao-seed-2-0-mini-260215", "doubao-seed-2-0-code-preview-260215"}) |
||||
|
void seedChatDefaultsToMinimalOnTheWire(String modelName) throws Exception { |
||||
|
AiModelConfig config = config("CHAT", "volcengine", modelName); |
||||
|
assertRequest(model(config), "minimal"); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(strings = {"minimal", "low", "medium", "high"}) |
||||
|
void explicitSeedEffortOverridesFastDefault(String effort) throws Exception { |
||||
|
AiModelConfig config = config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215"); |
||||
|
config.setExtraConfig(Map.of("reasoningEffort", effort, "topP", 0.8)); |
||||
|
JsonNode request = assertRequest(model(config), effort); |
||||
|
assertEquals(0.8, request.path("top_p").asDouble()); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@NullAndEmptySource |
||||
|
@ValueSource(strings = {" "}) |
||||
|
void clearedSeedEffortRestoresFastDefault(String cleared) throws Exception { |
||||
|
AiModelConfig config = config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215"); |
||||
|
config.setExtraConfig(cleared == null ? Map.of() : Map.of("reasoningEffort", cleared)); |
||||
|
assertRequest(model(config), "minimal"); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@CsvSource({"CHAT,openai,gpt-4o", "CHAT,deepseek,deepseek-chat", "CHAT,moonshot,moonshot-v1-8k", |
||||
|
"CHAT,volcengine,doubao-1-5-pro-32k", "CHAT,openai,doubao-seed-2-0-pro-260215", |
||||
|
"RAG_REWRITE,volcengine,doubao-seed-2-0-pro-260215", "RAG_REWRITE,volcengine,doubao-1-5-pro-32k"}) |
||||
|
void otherModelsAndRewriteRetainProviderDefaults(String appType, String provider, String name) throws Exception { |
||||
|
assertRequest(model(config(appType, provider, name)), null); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void unsupportedRewriteModelDoesNotReceiveStaleReasoningSetting() throws Exception { |
||||
|
AiModelConfig config = config("RAG_REWRITE", "volcengine", "doubao-1-5-pro-32k"); |
||||
|
config.setExtraConfig(Map.of("reasoningEffort", "medium")); |
||||
|
assertRequest(model(config), null); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void supportedRewriteModelHonorsExplicitEffort() throws Exception { |
||||
|
AiModelConfig config = config("RAG_REWRITE", "volcengine", "doubao-seed-2-0-pro-260215"); |
||||
|
config.setExtraConfig(Map.of("reasoningEffort", "low")); |
||||
|
assertRequest(model(config), "low"); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void cacheRefreshAppliesChangedAndClearedEffort() throws Exception { |
||||
|
AiModelConfig config = config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215"); |
||||
|
ChatModel initial = model(config); |
||||
|
assertSame(initial, factory.getChatModel("CHAT")); |
||||
|
assertRequest(initial, "minimal"); |
||||
|
|
||||
|
config.setExtraConfig(Map.of("reasoningEffort", "medium")); |
||||
|
factory.clearCache(); |
||||
|
ChatModel changed = factory.getChatModel("CHAT"); |
||||
|
assertNotSame(initial, changed); |
||||
|
assertRequest(changed, "medium"); |
||||
|
|
||||
|
config.setExtraConfig(Map.of()); |
||||
|
factory.clearCache(); |
||||
|
assertRequest(factory.getChatModel("CHAT"), "minimal"); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void streamingUsesTheSameFastDefault() throws Exception { |
||||
|
ChatModel model = model(config("CHAT", "volcengine", "doubao-seed-2-0-pro-260215")); |
||||
|
var responses = model.stream(new Prompt("测试")).collectList().block(Duration.ofSeconds(5)); |
||||
|
assertNotNull(responses); |
||||
|
assertFalse(responses.isEmpty()); |
||||
|
JsonNode request = requests.poll(5, TimeUnit.SECONDS); |
||||
|
assertNotNull(request); |
||||
|
assertTrue(request.path("stream").asBoolean()); |
||||
|
assertEquals("minimal", request.path("reasoning_effort").asText()); |
||||
|
} |
||||
|
|
||||
|
private AiModelConfig config(String appType, String provider, String modelName) { |
||||
|
return AiModelConfig.builder() |
||||
|
.id(1L).appType(appType).provider(provider).modelName(modelName) |
||||
|
.apiKey("local-test-only").baseUrl("http://127.0.0.1:" + server.getAddress().getPort()) |
||||
|
.temperature(0.7).maxTokens(4096).build(); |
||||
|
} |
||||
|
|
||||
|
private ChatModel model(AiModelConfig config) { |
||||
|
when(configService.getActiveConfigWithFullKey(config.getAppType())).thenReturn(config); |
||||
|
return factory.getChatModel(config.getAppType()); |
||||
|
} |
||||
|
|
||||
|
private JsonNode assertRequest(ChatModel model, String expectedEffort) throws Exception { |
||||
|
assertEquals(expectedEffort, ((OpenAiChatOptions) model.getDefaultOptions()).getReasoningEffort()); |
||||
|
assertEquals("ok", model.call(new Prompt("测试")).getResult().getOutput().getText()); |
||||
|
JsonNode request = requests.poll(5, TimeUnit.SECONDS); |
||||
|
assertNotNull(request); |
||||
|
if (expectedEffort == null) assertFalse(request.has("reasoning_effort")); |
||||
|
else assertEquals(expectedEffort, request.path("reasoning_effort").asText()); |
||||
|
assertEquals(4096, request.path("max_tokens").asInt()); |
||||
|
assertEquals(0.7, request.path("temperature").asDouble()); |
||||
|
assertFalse(request.has("thinking")); |
||||
|
return request; |
||||
|
} |
||||
|
} |
||||
@ -0,0 +1,300 @@ |
|||||
|
package com.wok.supportbot; |
||||
|
|
||||
|
import com.wok.supportbot.app.ChatContext; |
||||
|
import com.wok.supportbot.app.ChatPipeline; |
||||
|
import com.wok.supportbot.app.ChatRequest; |
||||
|
import com.wok.supportbot.chatmemory.DatabaseChatMemory; |
||||
|
import com.wok.supportbot.config.RagPromptConfig; |
||||
|
import com.wok.supportbot.entity.KnowledgeFaq; |
||||
|
import com.wok.supportbot.rag.CategoryFilter; |
||||
|
import com.wok.supportbot.rag.RagPipeline; |
||||
|
import com.wok.supportbot.rag.preretrieval.CompressionQueryRewriter; |
||||
|
import com.wok.supportbot.rag.preretrieval.MultiQueryExpanderRewriter; |
||||
|
import com.wok.supportbot.rag.preretrieval.RewriteQueryRewriter; |
||||
|
import com.wok.supportbot.rag.preretrieval.TranslationQueryRewriter; |
||||
|
import com.wok.supportbot.service.FaqMatchEngine; |
||||
|
import com.wok.supportbot.service.FaqMatchEngine.FaqMatchResult; |
||||
|
import com.wok.supportbot.service.RagHitLogService; |
||||
|
import com.wok.supportbot.service.SystemConfigService; |
||||
|
import org.junit.jupiter.api.BeforeEach; |
||||
|
import org.junit.jupiter.api.Test; |
||||
|
import org.junit.jupiter.api.extension.ExtendWith; |
||||
|
import org.junit.jupiter.params.ParameterizedTest; |
||||
|
import org.junit.jupiter.params.provider.NullAndEmptySource; |
||||
|
import org.junit.jupiter.params.provider.ValueSource; |
||||
|
import org.mockito.ArgumentCaptor; |
||||
|
import org.mockito.Mock; |
||||
|
import org.mockito.junit.jupiter.MockitoExtension; |
||||
|
import org.springframework.ai.chat.messages.Message; |
||||
|
import org.springframework.ai.chat.messages.UserMessage; |
||||
|
import org.springframework.ai.document.Document; |
||||
|
import org.springframework.ai.vectorstore.SearchRequest; |
||||
|
import org.springframework.ai.vectorstore.VectorStore; |
||||
|
import org.springframework.test.util.ReflectionTestUtils; |
||||
|
|
||||
|
import java.util.List; |
||||
|
import java.util.Map; |
||||
|
import java.util.Optional; |
||||
|
|
||||
|
import static org.junit.jupiter.api.Assertions.*; |
||||
|
import static org.mockito.Mockito.*; |
||||
|
|
||||
|
@ExtendWith(MockitoExtension.class) |
||||
|
class ChatPipelineTests { |
||||
|
@Mock private VectorStore vectorStore; |
||||
|
@Mock private FaqMatchEngine faqMatchEngine; |
||||
|
@Mock private RagHitLogService ragHitLogService; |
||||
|
@Mock private SystemConfigService systemConfigService; |
||||
|
@Mock private DatabaseChatMemory chatMemory; |
||||
|
@Mock private RagPromptConfig ragPromptConfig; |
||||
|
@Mock private RewriteQueryRewriter rewrite; |
||||
|
@Mock private TranslationQueryRewriter translation; |
||||
|
@Mock private CompressionQueryRewriter compression; |
||||
|
@Mock private MultiQueryExpanderRewriter multiQuery; |
||||
|
|
||||
|
private final CategoryFilter categoryFilter = new CategoryFilter(); |
||||
|
private ChatPipeline pipeline; |
||||
|
|
||||
|
@BeforeEach |
||||
|
void setUp() { |
||||
|
RagPipeline rag = new RagPipeline(chatMemory); |
||||
|
ReflectionTestUtils.setField(rag, "pgVectorVectorStore", vectorStore); |
||||
|
ReflectionTestUtils.setField(rag, "faqMatchEngine", faqMatchEngine); |
||||
|
ReflectionTestUtils.setField(rag, "ragHitLogService", ragHitLogService); |
||||
|
ReflectionTestUtils.setField(rag, "ragPromptConfig", ragPromptConfig); |
||||
|
ReflectionTestUtils.setField(rag, "categoryFilter", categoryFilter); |
||||
|
ReflectionTestUtils.setField(rag, "rewriteQueryRewriter", rewrite); |
||||
|
ReflectionTestUtils.setField(rag, "translationQueryRewriter", translation); |
||||
|
ReflectionTestUtils.setField(rag, "compressionQueryRewriter", compression); |
||||
|
ReflectionTestUtils.setField(rag, "multiQueryExpanderRewriter", multiQuery); |
||||
|
pipeline = new ChatPipeline(); |
||||
|
ReflectionTestUtils.setField(pipeline, "ragPipeline", rag); |
||||
|
ReflectionTestUtils.setField(pipeline, "systemConfigService", systemConfigService); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@NullAndEmptySource |
||||
|
@ValueSource(strings = {"NONE", " "}) |
||||
|
void defaultRetrievalUsesOriginalQuestionWithoutLlmPreprocessing(String strategy) { |
||||
|
ChatContext ctx = context("退货流程是什么", true).withRewriteStrategy(strategy); |
||||
|
when(vectorStore.similaritySearch(any(SearchRequest.class))).thenReturn(List.of()); |
||||
|
|
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
|
||||
|
assertEquals("NONE", ctx.rewriteStrategy()); |
||||
|
assertEquals("RAG", request.intent()); |
||||
|
assertSame(ctx, request.ctx()); |
||||
|
assertEquals(ctx.message(), request.finalMessage()); |
||||
|
assertEquals("faq-fast-path", request.ctx().chatId()); |
||||
|
assertTrue(request.finalSystemPrompt().contains("售后客服")); |
||||
|
verifySearch(ctx, ctx.message()); |
||||
|
verify(faqMatchEngine).match(ctx.message(), ctx.categoryIds()); |
||||
|
verifyNoMoreInteractions(faqMatchEngine); |
||||
|
verifyNoPreprocessing(); |
||||
|
verify(ragHitLogService).recordMiss(ctx.chatId(), ctx.message(), "VECTOR"); |
||||
|
verifyNoMoreInteractions(ragHitLogService); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void hitDocumentsAreRetainedAndLoggedOnlyOnce() { |
||||
|
ChatContext ctx = context("退货流程是什么", true); |
||||
|
Document doc = new Document("退货说明", Map.of("documentId", "123", "title", "退货政策", "score", 0.9)); |
||||
|
when(vectorStore.similaritySearch(any(SearchRequest.class))).thenReturn(List.of(doc)); |
||||
|
when(ragPromptConfig.getAnswerRules()).thenReturn("使用知识库回答"); |
||||
|
|
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
|
||||
|
assertEquals(List.of(doc), request.hitDocuments()); |
||||
|
assertEquals("退货说明", request.ragContextText()); |
||||
|
assertEquals(1, request.hitCount()); |
||||
|
assertTrue(request.finalSystemPrompt().contains("使用知识库回答")); |
||||
|
assertTrue(request.finalSystemPrompt().endsWith("退货说明")); |
||||
|
verify(ragHitLogService).recordHit(ctx.chatId(), ctx.message(), 123L, "退货政策", "0.9", "VECTOR"); |
||||
|
verifyNoMoreInteractions(ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(strings = {"EXACT", "KEYWORD", "SEMANTIC"}) |
||||
|
void fullFaqMatchPrecedesEvenLocalGreeting(String matchType) { |
||||
|
ChatContext ctx = context("你好", true); |
||||
|
FaqMatchResult match = faqMatch("标准答案", matchType); |
||||
|
when(faqMatchEngine.match(ctx.message(), ctx.categoryIds())).thenReturn(Optional.of(match)); |
||||
|
when(systemConfigService.getValueByKey("ai_system_prompt")).thenReturn("全局提示词"); |
||||
|
|
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
|
||||
|
assertEquals("FAQ", request.intent()); |
||||
|
assertEquals(Optional.of("标准答案"), request.faqAnswer()); |
||||
|
assertSame(match, request.faqMatchResult()); |
||||
|
assertSame(ctx, request.ctx()); |
||||
|
assertEquals("全局提示词\n\n【当前角色设定】\n售后客服", request.finalSystemPrompt()); |
||||
|
verify(faqMatchEngine).match(ctx.message(), ctx.categoryIds()); |
||||
|
verifyNoMoreInteractions(faqMatchEngine); |
||||
|
verifyNoInteractions(vectorStore, ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(strings = {"你好", " HI! ", "谢谢", "再见"}) |
||||
|
void localGreetingMissSkipsRetrieval(String message) { |
||||
|
ChatContext ctx = context(message, true); |
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
assertEquals("CHITCHAT", request.intent()); |
||||
|
assertEquals(message, request.finalMessage()); |
||||
|
verify(faqMatchEngine).match(message, ctx.categoryIds()); |
||||
|
verifyNoMoreInteractions(faqMatchEngine); |
||||
|
verifyNoInteractions(vectorStore, ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void ordinaryChatDoesNotMatchFaqOrRetrieve() { |
||||
|
ChatContext ctx = context("退货流程是什么", false); |
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
assertEquals("CHAT", request.intent()); |
||||
|
assertSame(ctx, request.ctx()); |
||||
|
assertEquals(ctx.message(), request.finalMessage()); |
||||
|
assertFalse(request.faqHit()); |
||||
|
verifyNoInteractions(faqMatchEngine, vectorStore, ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void exceptionalFaqMissRetriesFullMatchAndCanReturnStandardAnswer() { |
||||
|
ChatContext ctx = context("退货流程是什么", true); |
||||
|
FaqMatchResult match = faqMatch("恢复后的标准答案", "SEMANTIC"); |
||||
|
when(faqMatchEngine.match(ctx.message(), ctx.categoryIds())) |
||||
|
.thenThrow(new IllegalStateException("暂时不可用")).thenReturn(Optional.of(match)); |
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
assertEquals("FAQ", request.intent()); |
||||
|
assertEquals(Optional.of("恢复后的标准答案"), request.faqAnswer()); |
||||
|
verify(faqMatchEngine, times(2)).match(ctx.message(), ctx.categoryIds()); |
||||
|
verifyNoInteractions(vectorStore, ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void repeatedFaqFailureStillRetrievesAndLogsOneMiss() { |
||||
|
ChatContext ctx = context("退货流程是什么", true); |
||||
|
when(faqMatchEngine.match(ctx.message(), ctx.categoryIds())).thenThrow(new IllegalStateException("不可用")); |
||||
|
assertEquals("RAG", pipeline.buildRequest(ctx).intent()); |
||||
|
verify(faqMatchEngine, times(2)).match(ctx.message(), ctx.categoryIds()); |
||||
|
verifySearch(ctx, ctx.message()); |
||||
|
verify(ragHitLogService).recordMiss(ctx.chatId(), ctx.message(), "VECTOR"); |
||||
|
verifyNoMoreInteractions(ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@ValueSource(strings = {"REWRITE", "TRANSLATION", "COMPRESSION", "MULTI_QUERY"}) |
||||
|
void explicitRewritePreservesOriginalAnswerMessageAndCategoryScope(String strategy) { |
||||
|
ChatContext ctx = context("它怎么退", true).withRewriteStrategy(strategy); |
||||
|
List<Message> history = List.of(new UserMessage("我买了一台打印机")); |
||||
|
switch (strategy) { |
||||
|
case "REWRITE" -> when(rewrite.doQueryRewrite(ctx.message())).thenReturn("打印机退货流程"); |
||||
|
case "TRANSLATION" -> when(translation.doQueryRewrite(ctx.message())).thenReturn("打印机退货流程"); |
||||
|
case "COMPRESSION" -> { |
||||
|
when(chatMemory.get(ctx.chatId(), 10)).thenReturn(history); |
||||
|
when(compression.doQueryRewrite(ctx.message(), history)).thenReturn("打印机退货流程"); |
||||
|
} |
||||
|
case "MULTI_QUERY" -> when(multiQuery.doQueryRewrite(ctx.message())).thenReturn(List.of("打印机退货流程")); |
||||
|
} |
||||
|
|
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
|
||||
|
assertEquals(strategy, ctx.rewriteStrategy()); |
||||
|
assertEquals(ctx.message(), request.finalMessage()); |
||||
|
assertSame(ctx, request.ctx()); |
||||
|
verifySearch(ctx, "打印机退货流程"); |
||||
|
verify(faqMatchEngine).match(ctx.message(), ctx.categoryIds()); |
||||
|
if ("COMPRESSION".equals(strategy)) { |
||||
|
verify(chatMemory).get(ctx.chatId(), 10); |
||||
|
verify(compression).doQueryRewrite(ctx.message(), history); |
||||
|
} else { |
||||
|
verifyNoInteractions(chatMemory); |
||||
|
} |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void explicitRewriteFailureFallsBackToOriginalQuestion() { |
||||
|
ChatContext ctx = context("退货流程是什么", true).withRewriteStrategy("REWRITE"); |
||||
|
when(rewrite.doQueryRewrite(ctx.message())).thenThrow(new IllegalStateException("重写不可用")); |
||||
|
assertEquals("RAG", pipeline.buildRequest(ctx).intent()); |
||||
|
verifySearch(ctx, ctx.message()); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void independentSourcesRetrievalDoesNotClassifyOrMatchFaq() { |
||||
|
ChatContext ctx = context("退货流程是什么", true); |
||||
|
assertTrue(pipeline.retrieveSources(ctx).isEmpty()); |
||||
|
verifySearch(ctx, ctx.message()); |
||||
|
verifyNoInteractions(faqMatchEngine); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void multipleQueriesKeepCategoryScopeAndLogMergedDocumentOnce() { |
||||
|
ChatContext ctx = context("退货流程是什么", true).withRewriteStrategy("MULTI_QUERY"); |
||||
|
Document shared = new Document("退货说明", Map.of("documentId", "123", "title", "退货政策")); |
||||
|
when(multiQuery.doQueryRewrite(ctx.message())).thenReturn(List.of("退货步骤", "退货条件")); |
||||
|
when(vectorStore.similaritySearch(any(SearchRequest.class))).thenReturn(List.of(shared)); |
||||
|
when(ragPromptConfig.getAnswerRules()).thenReturn("使用知识库回答"); |
||||
|
|
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
|
||||
|
assertEquals(List.of(shared), request.hitDocuments()); |
||||
|
assertEquals(ctx.message(), request.finalMessage()); |
||||
|
ArgumentCaptor<SearchRequest> searches = ArgumentCaptor.forClass(SearchRequest.class); |
||||
|
verify(vectorStore, times(2)).similaritySearch(searches.capture()); |
||||
|
assertEquals(List.of("退货步骤", "退货条件"), |
||||
|
searches.getAllValues().stream().map(SearchRequest::getQuery).toList()); |
||||
|
for (SearchRequest search : searches.getAllValues()) { |
||||
|
assertEquals(categoryFilter.buildExpression(ctx.categoryIds()), search.getFilterExpression()); |
||||
|
} |
||||
|
verify(ragHitLogService).recordHit(ctx.chatId(), ctx.message(), 123L, "退货政策", "", "VECTOR"); |
||||
|
verifyNoMoreInteractions(ragHitLogService); |
||||
|
verifyNoInteractions(rewrite, translation, compression, chatMemory); |
||||
|
} |
||||
|
|
||||
|
@ParameterizedTest |
||||
|
@NullAndEmptySource |
||||
|
void faqWithoutAnswerPreservesOptionalSemantics(String answer) { |
||||
|
ChatContext ctx = context("退货流程是什么", true); |
||||
|
FaqMatchResult match = faqMatch(answer, "EXACT"); |
||||
|
when(faqMatchEngine.match(ctx.message(), ctx.categoryIds())).thenReturn(Optional.of(match)); |
||||
|
|
||||
|
ChatRequest request = pipeline.buildRequest(ctx); |
||||
|
|
||||
|
assertEquals("FAQ", request.intent()); |
||||
|
assertEquals(Optional.ofNullable(answer), request.faqAnswer()); |
||||
|
assertEquals(answer != null, request.faqHit()); |
||||
|
verifyNoInteractions(vectorStore, ragHitLogService); |
||||
|
verifyNoPreprocessing(); |
||||
|
} |
||||
|
|
||||
|
private void verifySearch(ChatContext ctx, String query) { |
||||
|
ArgumentCaptor<SearchRequest> search = ArgumentCaptor.forClass(SearchRequest.class); |
||||
|
verify(vectorStore).similaritySearch(search.capture()); |
||||
|
assertEquals(query, search.getValue().getQuery()); |
||||
|
assertEquals(4, search.getValue().getTopK()); |
||||
|
assertEquals(categoryFilter.buildExpression(ctx.categoryIds()), search.getValue().getFilterExpression()); |
||||
|
} |
||||
|
|
||||
|
private void verifyNoPreprocessing() { |
||||
|
verifyNoInteractions(rewrite, translation, compression, multiQuery, chatMemory); |
||||
|
} |
||||
|
|
||||
|
private ChatContext context(String message, boolean enableRag) { |
||||
|
return ChatContext.of(message, "faq-fast-path") |
||||
|
.withSystemPrompt("售后客服") |
||||
|
.withCategoryIds(List.of(101L, 202L)) |
||||
|
.withEnableRag(enableRag); |
||||
|
} |
||||
|
|
||||
|
private FaqMatchResult faqMatch(String answer, String matchType) { |
||||
|
KnowledgeFaq faq = new KnowledgeFaq(); |
||||
|
faq.setAnswer(answer); |
||||
|
return new FaqMatchResult(faq, matchType, 0.95); |
||||
|
} |
||||
|
} |
||||
@ -0,0 +1,166 @@ |
|||||
|
package com.wok.supportbot; |
||||
|
|
||||
|
import com.wok.supportbot.app.AssistantApp; |
||||
|
import com.wok.supportbot.app.ChatContext; |
||||
|
import com.wok.supportbot.app.ChatResult; |
||||
|
import com.wok.supportbot.app.SourceReference; |
||||
|
import com.wok.supportbot.config.RoleAccessConfig; |
||||
|
import com.wok.supportbot.controller.AiController; |
||||
|
import com.wok.supportbot.controller.OpenApiController; |
||||
|
import com.wok.supportbot.entity.ApiKey; |
||||
|
import com.wok.supportbot.rag.CategoryFilter; |
||||
|
import com.wok.supportbot.security.JwtTokenProvider; |
||||
|
import com.wok.supportbot.security.SdkAuthFilter; |
||||
|
import com.wok.supportbot.security.SdkJwtTokenProvider; |
||||
|
import com.wok.supportbot.service.ConversationService; |
||||
|
import com.wok.supportbot.service.CustomerServiceRoleService; |
||||
|
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope; |
||||
|
import io.jsonwebtoken.Claims; |
||||
|
import org.junit.jupiter.api.BeforeEach; |
||||
|
import org.junit.jupiter.api.Test; |
||||
|
import org.mockito.ArgumentCaptor; |
||||
|
import org.springframework.jdbc.core.JdbcTemplate; |
||||
|
import org.springframework.mock.web.MockHttpServletRequest; |
||||
|
import org.springframework.test.util.ReflectionTestUtils; |
||||
|
import org.springframework.test.web.servlet.MockMvc; |
||||
|
import org.springframework.test.web.servlet.setup.MockMvcBuilders; |
||||
|
|
||||
|
import java.util.List; |
||||
|
import java.util.Set; |
||||
|
|
||||
|
import static org.junit.jupiter.api.Assertions.*; |
||||
|
import static org.mockito.ArgumentMatchers.any; |
||||
|
import static org.mockito.Mockito.*; |
||||
|
import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.get; |
||||
|
import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.*; |
||||
|
|
||||
|
class ChatResultEndpointTests { |
||||
|
private AssistantApp assistant; |
||||
|
private CustomerServiceRoleService roles; |
||||
|
private ConversationService conversations; |
||||
|
private RoleAccessConfig access; |
||||
|
private AiController controller; |
||||
|
private MockMvc mvc; |
||||
|
private SdkJwtTokenProvider tokens; |
||||
|
|
||||
|
@BeforeEach |
||||
|
void setUp() { |
||||
|
assistant = mock(AssistantApp.class); |
||||
|
roles = mock(CustomerServiceRoleService.class); |
||||
|
conversations = mock(ConversationService.class); |
||||
|
access = new RoleAccessConfig(); |
||||
|
controller = new AiController(); |
||||
|
ReflectionTestUtils.setField(controller, "assistantApp", assistant); |
||||
|
ReflectionTestUtils.setField(controller, "customerServiceRoleService", roles); |
||||
|
ReflectionTestUtils.setField(controller, "conversationService", conversations); |
||||
|
ReflectionTestUtils.setField(controller, "roleAccessConfig", access); |
||||
|
ReflectionTestUtils.setField(controller, "categoryFilter", mock(CategoryFilter.class)); |
||||
|
tokens = mock(SdkJwtTokenProvider.class); |
||||
|
SdkAuthFilter filter = new SdkAuthFilter(tokens, mock(JwtTokenProvider.class), mock(JdbcTemplate.class)); |
||||
|
mvc = MockMvcBuilders.standaloneSetup(controller).addFilters(filter).build(); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void resultRejectsMissingTokenAndUnauthorizedRoleBeforeGeneration() throws Exception { |
||||
|
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result").param("message", "问题")) |
||||
|
.andExpect(status().isUnauthorized()); |
||||
|
authorize(); |
||||
|
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result") |
||||
|
.header("Authorization", "Bearer valid").param("message", "问题").param("roleId", "99")) |
||||
|
.andExpect(status().isForbidden()); |
||||
|
verifyNoInteractions(assistant, roles, conversations); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void resultKeepsRoleScopeAccountBindingAndDecodedImages() throws Exception { |
||||
|
authorize(); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "授权人设", List.of(7L), List.of("tool"))); |
||||
|
SourceReference source = new SourceReference("9223372036854775807", "授权文档", null, null, null, "片段"); |
||||
|
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of(source))); |
||||
|
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result") |
||||
|
.header("Authorization", "Bearer valid").param("message", "问题").param("chatId", "chat") |
||||
|
.param("roleId", "11").param("accountId", " account ").param("systemPrompt", "越权人设") |
||||
|
.param("enableRag", "true").param("categoryIds", "99").param("imageUrls", "https%3A%2F%2Fexample.org%2Fa.png")) |
||||
|
.andExpect(status().isOk()).andExpect(content().contentTypeCompatibleWith("application/json")) |
||||
|
.andExpect(jsonPath("$.text").value("回答")) |
||||
|
.andExpect(jsonPath("$.data").doesNotExist()) |
||||
|
.andExpect(jsonPath("$.sources[0].documentId").value("9223372036854775807")) |
||||
|
.andExpect(jsonPath("$.mcpEvents").isArray()).andExpect(jsonPath("$.suggestions").isArray()); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).chatWithEvents(context.capture()); |
||||
|
ChatContext ctx = context.getValue(); |
||||
|
assertEquals(List.of(7L), ctx.categoryIds()); |
||||
|
assertEquals("授权人设", ctx.systemPrompt()); |
||||
|
assertEquals(List.of("tool"), ctx.allowedMcpTools()); |
||||
|
assertEquals(List.of("https://example.org/a.png"), ctx.imageUrls()); |
||||
|
assertEquals("NONE", ctx.rewriteStrategy()); |
||||
|
assertTrue(ctx.enableRag()); |
||||
|
verify(conversations).bindConversation("chat", "account", 11L); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void strictIsolationAndExplicitRewriteApplyToBothSyncEndpoints() { |
||||
|
access.setStrictIsolation(true); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(), List.of())); |
||||
|
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of())); |
||||
|
assertEquals("回答", controller.chatSync("问题", "chat", 11L, "account", null, true, |
||||
|
"REWRITE", null, "99", null)); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).chatWithEvents(context.capture()); |
||||
|
assertFalse(context.getValue().enableRag()); |
||||
|
assertTrue(context.getValue().categoryIds().isEmpty()); |
||||
|
assertEquals("REWRITE", context.getValue().rewriteStrategy()); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void openApiRetainsEnvelopeAndAddsSameAnswerSourcesWithNoneDefault() { |
||||
|
OpenApiController open = new OpenApiController(); |
||||
|
ReflectionTestUtils.setField(open, "assistantApp", assistant); |
||||
|
ReflectionTestUtils.setField(open, "customerServiceRoleService", roles); |
||||
|
ReflectionTestUtils.setField(open, "categoryFilter", mock(CategoryFilter.class)); |
||||
|
ReflectionTestUtils.setField(open, "roleAccessConfig", access); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(7L), List.of())); |
||||
|
SourceReference source = new SourceReference("9223372036854775807", "授权文档", null, null, null, "片段"); |
||||
|
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of(source))); |
||||
|
ApiKey key = new ApiKey(); |
||||
|
key.setId(3L); |
||||
|
key.setRoleIds("[11]"); |
||||
|
MockHttpServletRequest request = new MockHttpServletRequest(); |
||||
|
request.setAttribute("apiKey", key); |
||||
|
var result = open.chat("问题", "11", "chat", "99", null, true, request); |
||||
|
assertEquals(200, result.getStatusCode().value()); |
||||
|
assertEquals(true, result.getBody().get("success")); |
||||
|
var data = (java.util.Map<?, ?>) result.getBody().get("data"); |
||||
|
assertEquals("回答", data.get("reply")); |
||||
|
assertEquals(List.of(source), data.get("sources")); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).chatWithEvents(context.capture()); |
||||
|
assertEquals("NONE", context.getValue().rewriteStrategy()); |
||||
|
assertEquals(List.of(7L), context.getValue().categoryIds()); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
@Test |
||||
|
void explicitSourcesApiUsesTheSameCitationSerializer() { |
||||
|
ReflectionTestUtils.setField(controller, "chatPipeline", mock(com.wok.supportbot.app.ChatPipeline.class)); |
||||
|
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(7L), List.of())); |
||||
|
var document = new org.springframework.ai.document.Document("片段".repeat(100), |
||||
|
java.util.Map.of("documentId", Long.MAX_VALUE, "distance", 0.2)); |
||||
|
when(assistant.retrieveSources(any())).thenReturn(List.of(document)); |
||||
|
var result = controller.chatSources("问题", "chat", null, 11L, "account", null, "99"); |
||||
|
assertEquals(SourceReference.fromDocuments(List.of(document)), result.get("data")); |
||||
|
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class); |
||||
|
verify(assistant).retrieveSources(context.capture()); |
||||
|
assertEquals(List.of(7L), context.getValue().categoryIds()); |
||||
|
assertEquals("NONE", context.getValue().rewriteStrategy()); |
||||
|
verifyNoMoreInteractions(assistant); |
||||
|
} |
||||
|
|
||||
|
private void authorize() { |
||||
|
Claims claims = mock(Claims.class); |
||||
|
when(tokens.parseToken("valid")).thenReturn(claims); |
||||
|
when(tokens.getAllowedRoleIds(claims)).thenReturn(Set.of(11L)); |
||||
|
} |
||||
|
} |
||||
Write
Preview
Loading…
Cancel
Save
Reference in new issue