Browse Source

重构(AI对话): 统一对话管道,收敛接口形态,补齐Open API能力

master
wanghanlin 3 days ago
parent
commit
91c4c185b3
  1. 30
      CLAUDE.md
  2. 31
      SDK-INTEGRATION.md
  3. 654
      src/main/java/com/wok/supportbot/app/AssistantApp.java
  4. 78
      src/main/java/com/wok/supportbot/app/ChatContext.java
  5. 154
      src/main/java/com/wok/supportbot/app/ChatPipeline.java
  6. 34
      src/main/java/com/wok/supportbot/app/ChatRequest.java
  7. 17
      src/main/java/com/wok/supportbot/app/ChatResult.java
  8. 282
      src/main/java/com/wok/supportbot/controller/AiController.java
  9. 13
      src/main/java/com/wok/supportbot/controller/CustomerServiceRoleController.java
  10. 26
      src/main/java/com/wok/supportbot/controller/DocumentController.java
  11. 128
      src/main/java/com/wok/supportbot/controller/OpenApiController.java
  12. 108
      src/main/java/com/wok/supportbot/rag/CategoryFilter.java
  13. 34
      src/main/java/com/wok/supportbot/rag/RagContext.java
  14. 271
      src/main/java/com/wok/supportbot/rag/RagPipeline.java
  15. 13
      src/main/java/com/wok/supportbot/service/DocumentService.java
  16. 2
      src/main/resources/application-dev.yml
  17. 4
      src/main/resources/static/components/ChatPanel.js
  18. 10
      src/main/resources/static/js/api.js
  19. 123
      src/test/java/com/wok/supportbot/Phase1ComponentTests.java
  20. 13
      src/test/java/com/wok/supportbot/SupportBotApplicationTests.java

30
CLAUDE.md

@ -38,21 +38,37 @@ AI 智能客服系统,基于 Spring AI Alibaba + 通义千问 + PGVector,支
### Spring AI 集成模式
- **ChatClient Builder**: 所有对话通过 `ChatClient.builder(chatModelFactory.getChatModel("CHAT"))` 构建,ChatModel 由 `ChatModelFactory` 按 DB 活跃配置动态创建
- **Advisor 链**: `MessageChatMemoryAdvisor`(记忆) → `MyLoggerAdvisor`(日志) → `QuestionAnswerAdvisor`(RAG)
- **SSE 流式**: 三种实现 — Flux\<String\>、Flux\<ServerSentEvent\>、SseEmitter
- **ChatClient 构建**: 所有权在 `AssistantApp`(`getChatClient`),按 `appType` + `allowedMcpTools` 缓存不同实例
- **Advisor 链**: `ContentSafetyAdvisor`(最外层,`HIGHEST_PRECEDENCE`)→ `MessageChatMemoryAdvisor`(记忆)→ `MyLoggerAdvisor`(日志)
- **SSE 流式**: 仅保留 `Flux<String>` 形态;废弃的 `Flux<ServerSentEvent>``SseEmitter` 已移除
### ChatMemory 持久化
当前使用 `DatabaseChatMemory`(PostgreSQL 持久化),`FileBasedChatMemory`(Kryo 序列化)已注释掉。
### RAG 双模式
1. **QuestionAnswerAdvisor 模式**(生产使用): 预检索优化 + `QuestionAnswerAdvisor`
2. **RetrievalAugmentationAdvisor 模式**(实验性): `doChatWithRagEnhance()` 仅做基础 RAG 检索,无查询增强
### 统一对话管道(重构后)
对话管道由 `ChatPipeline`(编排层)+ `RagPipeline`(RAG 检索层)+ `AssistantApp`(执行层)组成:
```
用户请求
→ 鉴权/角色解析(Controller)
→ ChatPipeline.buildRequest(ChatContext)
→ IntentRouter 意图路由(CHITCHAT/FAQ/RAG)
→ RagPipeline.retrieve(FAQ 优先 → 查询重写 → 统一检索)
→ 组装 finalMessage + finalSystemPrompt + 资料块
→ AssistantApp.chat / chatStream(构建 ChatClientRequestSpec → call/stream)
```
- **ChatPipeline**: 纯编排,不持有 ChatClient;产出 `ChatRequest` 决策对象
- **RagPipeline**: 统一 RAG 检索,所有策略(含 MULTI_QUERY)均走"手动检索 + 资料块注入 system prompt"模式,不再使用 `RetrievalAugmentationAdvisor` 的 query augmenter
- **RAG 查询重写策略**: 由 `RagPipeline` 统一路由,`AssistantApp` 等旧方法已移除
- **IntentRouter**: 已在 ChatPipeline 接入,`AiController.shouldBypassKnowledgeRetrieval` 已移除
- **分类过滤**: 统一由 `CategoryFilter` 工具类处理(`parse`/`normalize`/`buildExpression`)
- **AssistantApp 入口**: `chat(ChatContext)` / `chatStream(ChatContext)` / `retrieveSources(ChatContext)`,旧方法(`doChat*`、`doChatWithRag*`)已移除
- **Open API**: `OpenApiController` 已接入 `ChatPipeline`,补齐角色/RAG/FAQ/MCP/分类隔离能力
### 文档处理管道
`DocumentService.uploadDocument()` 统一流程:文档提取 → `MyTokenTextSplitter` 分块 → `MyKeywordEnricher` AI 关键词提取 → `pgVectorVectorStore.add()` 向量化存储。每个分块的 metadata 中注入 `documentId`、`chunkIndex`、`sourceName`、`title` 以关联 `knowledge_document` 表。
### 预检索查询优化
四种策略在 `rag/preretrieval/` 下,由 `AssistantApp.doChatWithRagStrategy()` 根据 `strategy` 参数动态选择:REWRITE / TRANSLATION / COMPRESSION / MULTI_QUERY。查询重写器均为 `@Component`,仅在请求时懒加载 ChatModel,不影响启动。
## 关键配置

31
SDK-INTEGRATION.md

@ -276,9 +276,9 @@ POST /open-api/auth/token
| 接口 | 方法 | 说明 |
|---|---|---|
| `/ai/assistant_app/chat/sync` | GET | 同步对话(返回完整文本) |
| `/ai/assistant_app/chat/sse` | GET | SSE 流式对话 |
| `/ai/assistant_app/chat/rag/sse` | GET | RAG 增强流式对话 |
| `/ai/assistant_app/chat/sync` | GET | 同步对话(已统一为 ChatPipeline,支持普通/RAG 自动判断) |
| `/ai/assistant_app/chat/sse` | GET | SSE 流式对话(已统一为 ChatPipeline) |
| `/ai/assistant_app/chat/rag/sse` | GET | RAG 增强流式对话(已统一为 ChatPipeline) |
| `/ai/assistant_app/rag/sources` | GET | 获取 RAG 引用来源 |
| `/ai/sdk/conversation/list` | GET | 会话列表 |
| `/ai/sdk/conversation/{id}/messages` | GET | 会话消息 |
@ -287,18 +287,25 @@ POST /open-api/auth/token
| `/category/tree` | GET | 知识库分类树 |
| `/feedback` | POST | 消息反馈 |
**通用请求参数(对话接口):**
> **已废弃**:`/ai/assistant_app/chat/server_sent_event` 和 `/ai/assistant_app/chat/sse_emitter` 已移除,请统一使用 `/ai/assistant_app/chat/sse`
### 6.3 Open API 对话接口(第三方系统直接调用)
| 接口 | 方法 | 说明 |
|---|---|---|
| `/open-api/chat` | POST | 同步对话(已补齐角色/RAG/FAQ/MCP/分类隔离) |
| `/open-api/chat/stream` | GET | SSE 流式对话(同上) |
| `/open-api/rag/search` | POST | 知识库检索 |
**Open API 对话新增参数:**
| 参数 | 类型 | 必填 | 说明 |
|---|---|---|---|
| `message` | Query | 是 | 用户消息 |
| `chatId` | Query | 是 | 会话 ID(SDK 自动管理) |
| `roleId` | Query | 否 | 客服角色 ID |
| `accountId` | Query | 否 | 用户标识 |
| `categoryId` | Query | 否 | 知识库分类 ID |
| `rewriteStrategy` | Query | 否 | RAG 查询重写策略:REWRITE / TRANSLATION / COMPRESSION / MULTI_QUERY |
**认证方式:** 所有 `/ai/**` 请求自动在 Header 中携带 `Authorization: Bearer {token}`
| `categoryIds` | Query | 否 | 知识库分类 ID,逗号分隔(非角色场景用) |
| `rewriteStrategy` | Query | 否 | RAG 查询重写策略:REWRITE / TRANSLATION / COMPRESSION / MULTI_QUERY(默认) |
| `enableRag` | Query | 否 | 是否启用 RAG 检索(默认 true),false 时走普通对话 |
**认证方式:** 所有 `/open-api/**``/ai/**` 请求通过 `X-API-Key` Header 或 Bearer Token 鉴权。
### 6.3 管理接口(需要管理后台 JWT)

654
src/main/java/com/wok/supportbot/app/AssistantApp.java

@ -2,49 +2,26 @@ package com.wok.supportbot.app;
import com.wok.supportbot.advisor.ContentSafetyAdvisor;
import com.wok.supportbot.advisor.MyLoggerAdvisor;
import com.wok.supportbot.advisor.ReReadingAdvisor;
import com.wok.supportbot.chatmemory.DatabaseChatMemory;
import com.wok.supportbot.config.ChatModelFactory;
import com.wok.supportbot.config.RagPromptConfig;
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.mcp.McpToolCallback;
import com.wok.supportbot.mcp.McpToolCallbackAdapter;
import jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.chat.client.ChatClient;
import org.springframework.ai.chat.client.advisor.MessageChatMemoryAdvisor;
import org.springframework.ai.chat.client.advisor.api.Advisor;
import org.springframework.ai.chat.memory.ChatMemory;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.PromptTemplate;
import org.springframework.ai.document.Document;
import org.springframework.ai.rag.advisor.RetrievalAugmentationAdvisor;
import org.springframework.ai.rag.generation.augmentation.ContextualQueryAugmenter;
import org.springframework.ai.rag.retrieval.search.VectorStoreDocumentRetriever;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.ai.vectorstore.filter.FilterExpressionBuilder;
import org.springframework.beans.factory.annotation.Value;
import org.springframework.stereotype.Component;
import org.springframework.util.StringUtils;
import reactor.core.publisher.Flux;
import java.util.ArrayList;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.concurrent.ConcurrentHashMap;
import java.util.stream.Collectors;
import static org.springframework.ai.chat.memory.ChatMemory.CONVERSATION_ID;
@ -59,21 +36,14 @@ import static org.springframework.ai.chat.memory.ChatMemory.CONVERSATION_ID;
@Slf4j
public class AssistantApp {
@Resource
private VectorStore pgVectorVectorStore;
@Resource
private ContentSafetyAdvisor contentSafetyAdvisor;
@Resource
private FaqMatchEngine faqMatchEngine;
@Resource
private McpToolCallbackAdapter mcpToolCallbackAdapter;
/** RAG 回答规则(保真护栏),通过 application.yml 的 knowledge.rag.answer-rules 配置 */
@Resource
private RagPromptConfig ragPromptConfig;
private ChatPipeline chatPipeline;
/** MCP 工具开关,默认启用,可通过 application.yml 的 chat.mcp.enabled 关闭 */
@Value("${chat.mcp.enabled:true}")
@ -158,119 +128,82 @@ public class AssistantApp {
log.info("AssistantApp ChatClient cache cleared");
}
/**
* AI 基础对话支持多轮对话记忆
*
* @param message
* @param chatId
* @return
*/
public String doChat(String message, String chatId) {
return doChat(message, chatId, null);
}
// ==================== 统一入口 ====================
/**
* AI 基础对话支持多轮对话记忆 + 角色系统提示词
* 同步对话新入口委托 {@link ChatPipeline} 编排
*
* @param message 用户消息保持原样不做包装避免污染会话记忆
* @param chatId 会话ID
* @param systemPrompt 角色人设/风格作为系统提示词叠加在基础提示词之上为空则仅用基础提示词
* @return AI 回答
* @param ctx 对话上下文
* @return AI 回答文本
*/
public String doChat(String message, String chatId, String systemPrompt) {
return doChat(message, chatId, systemPrompt, null);
public String chat(ChatContext ctx) {
return chatWithEvents(ctx).text();
}
/**
* AI 基础对话支持多轮对话记忆 + 角色系统提示词 + MCP 工具权限
* 同步对话 + MCP 工具调用事件新入口
* {@link #chat(ChatContext)} 多返回本次触发的工具调用事件供需要展示调用过程的场景
*
* @param message 用户消息
* @param chatId 会话ID
* @param systemPrompt 角色人设/风格
* @param allowedMcpTools 允许的 MCP 工具列表null=不注册工具
* @return AI 回答
* @param ctx 对话上下文
* @return 回答文本 + MCP 事件
*/
public String doChat(String message, String chatId, String systemPrompt, List<String> allowedMcpTools) {
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT", allowedMcpTools)
.prompt()
.user(message)
.advisors(s -> s.param(CONVERSATION_ID, chatId));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
public ChatResult chatWithEvents(ChatContext ctx) {
ChatRequest req = chatPipeline.buildRequest(ctx);
if (req.faqHit()) {
return new ChatResult(req.faqAnswer().get(), List.of());
}
ChatResponse chatResponse = spec.call().chatResponse();
return chatResponse.getResult().getOutput().getText();
}
/**
* 组合系统提示词基础客服提示词 + 角色人设
* 角色人设作为附加段落叠加既保留客服基线约束又让角色风格生效
*/
private String effectiveSystem(String rolePrompt) {
if (!StringUtils.hasText(rolePrompt)) {
return SYSTEM_PROMPT;
McpToolCallback.resetEvents();
McpToolCallback.resetCallRounds();
ChatClient.ChatClientRequestSpec spec = getChatClient(ctx.appType(), ctx.allowedMcpTools())
.prompt()
.user(req.finalMessage())
.advisors(s -> s.param(CONVERSATION_ID, ctx.chatId()));
if (StringUtils.hasText(req.finalSystemPrompt())) {
spec = spec.system(req.finalSystemPrompt());
}
return SYSTEM_PROMPT + "\n\n【当前角色设定】\n" + rolePrompt;
}
/**
* AI 基础对话支持多轮对话记忆SSE 流式传输
*
* @param message
* @param chatId
* @return
*/
public Flux<String> doChatByStream(String message, String chatId) {
return doChatByStream(message, chatId, null);
String text = spec.call().chatResponse().getResult().getOutput().getText();
return new ChatResult(text, McpToolCallback.drainEvents());
}
/**
* AI 基础对话多轮记忆 + 角色系统提示词SSE 流式传输
* 流式对话新入口委托 {@link ChatPipeline} 编排
* 自动重置 MCP 事件并在流末尾追加工具调用事件行
*
* @param message 用户消息保持原样
* @param chatId 会话ID
* @param systemPrompt 角色人设/风格作为系统提示词叠加为空则仅用基础提示词
* @param ctx 对话上下文
* @return 流式回答
*/
public Flux<String> doChatByStream(String message, String chatId, String systemPrompt) {
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT")
public Flux<String> chatStream(ChatContext ctx) {
ChatRequest req = chatPipeline.buildRequest(ctx);
if (req.faqHit()) {
return Flux.just(req.faqAnswer().get());
}
McpToolCallback.resetEvents();
McpToolCallback.resetCallRounds();
ChatClient.ChatClientRequestSpec spec = getChatClient(ctx.appType(), ctx.allowedMcpTools())
.prompt()
.user(message)
.advisors(s -> s.param(CONVERSATION_ID, chatId));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
.user(req.finalMessage())
.advisors(s -> s.param(CONVERSATION_ID, ctx.chatId()));
if (StringUtils.hasText(req.finalSystemPrompt())) {
spec = spec.system(req.finalSystemPrompt());
}
return spec.stream().content();
return appendMcpToolEvents(spec.stream().content());
}
/**
* AI 基础对话多轮记忆 + 角色系统提示词 + MCP 工具权限SSE 流式传输
* 统一检索引用来源新入口委托 {@link ChatPipeline#retrieveSources}
*
* @param message 用户消息
* @param chatId 会话ID
* @param systemPrompt 角色人设/风格
* @param allowedMcpTools 允许的 MCP 工具列表null=不注册工具
* @return 流式回答
* @param ctx 对话上下文使用 message / rewriteStrategy / categoryIds
* @return 命中的知识库片段 metadatadocumentId/title/sourceName/chunkIndex/distance
*/
public Flux<String> doChatByStream(String message, String chatId, String systemPrompt, List<String> allowedMcpTools) {
// 重置 MCP 工具调用状态事件收集器 + 轮次计数器
McpToolCallback.resetEvents();
McpToolCallback.resetCallRounds();
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT", allowedMcpTools)
.prompt()
.user(message)
.advisors(s -> s.param(CONVERSATION_ID, chatId));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
}
return appendMcpToolEvents(spec.stream().content());
public List<Document> retrieveSources(ChatContext ctx) {
return chatPipeline.retrieveSources(ctx);
}
// ==================== 内部工具方法 ====================
/**
* SSE 文本流末尾追加 MCP 工具调用事件
* 前端通过解析 "event: tool_call_start" / "event: tool_call_result" 行来展示工具调用过程
* 事件格式遵循 SSE 标准 readSSEStreamWithEvents 的解析逻辑匹配
* SSE 文本流末尾追加 MCP 工具调用事件行
* 事件格式遵循 SSE 标准前端由 readSSEStreamWithEvents 解析
*/
private Flux<String> appendMcpToolEvents(Flux<String> contentFlux) {
return contentFlux.concatMap(chunk -> Flux.just(chunk))
@ -300,491 +233,4 @@ public class AssistantApp {
return s.replace("\\", "\\\\").replace("\"", "\\\"")
.replace("\n", "\\n").replace("\r", "\\r");
}
// ==================== FAQ 优先匹配 ====================
/**
* 尝试 FAQ 三级匹配精确关键词语义命中则返回标准答案
* RAG 对话入口处优先调用避免不必要的知识库检索开销
*
* @param message 用户消息
* @return FAQ 标准答案未命中时返回 empty
*/
private Optional<String> tryFaqMatch(String message) {
try {
return faqMatchEngine.match(message)
.map(result -> result.getFaq().getAnswer());
} catch (Exception e) {
log.warn("FAQ 匹配异常,降级到 RAG: {}", e.getMessage());
return Optional.empty();
}
}
// AI 恋爱知识库问答功能
@Resource
RewriteQueryRewriter rewriteQueryRewriter;
@Resource
CompressionQueryRewriter compressionQueryRewriter;
@Resource
MultiQueryExpanderRewriter multiQueryExpanderRewriter;
@Resource
TranslationQueryRewriter translationQueryRewriter;
/**
* RAG 知识库进行对话
*
* @param message
* @param chatId
* @return
*/
public String doChatWithRag(String message, String chatId) {
// 在预检索阶段系统接收用户的原始查询通过查询转换和查询扩展等方法对其进行优化输出增强的用户查询
// String rewrittenMessage = translationQueryRewriter.doQueryRewrite(message);
String rewrittenMessage = rewriteQueryRewriter.doQueryRewrite(message);
ChatResponse chatResponse = getChatClient("CHAT")
.prompt()
.user(rewrittenMessage)
.advisors(spec -> spec.param(CONVERSATION_ID, chatId))
// 应用 RAG 知识库问答
.advisors(buildRetrievalAdvisor(4, Collections.emptyList()))
.call()
.chatResponse();
return chatResponse.getResult().getOutput().getText();
}
/**
* RAG 知识库进行对话支持动态选择查询重写策略
*
* @param message 用户消息
* @param chatId 会话ID
* @param strategy 查询重写策略NONE/REWRITE/TRANSLATION/COMPRESSION/MULTI_QUERY
* @return AI 回答
*/
public String doChatWithRagStrategy(String message, String chatId, String strategy) {
return doChatWithRagStrategy(message, chatId, strategy, Collections.emptyList());
}
public String doChatWithRagStrategy(String message, String chatId, String strategy, Long categoryId) {
List<Long> categoryIds = categoryId != null ? List.of(categoryId) : Collections.emptyList();
return doChatWithRagStrategy(message, chatId, strategy, categoryIds);
}
public String doChatWithRagStrategy(String message, String chatId, String strategy, List<Long> categoryIds) {
return doChatWithRagStrategy(message, chatId, strategy, categoryIds, null);
}
public String doChatWithRagStrategy(String message, String chatId, String strategy, List<Long> categoryIds, String systemPrompt) {
// FAQ 优先匹配三级匹配精确关键词语义命中则直接返回标准答案
Optional<String> faqAnswer = tryFaqMatch(message);
if (faqAnswer.isPresent()) {
log.info("FAQ 命中,直接返回标准答案: chatId={}", chatId);
return faqAnswer.get();
}
// 对于 MULTI_QUERY 策略需要使用特殊的处理方式
if ("MULTI_QUERY".equalsIgnoreCase(strategy)) {
return doChatWithMultiQueryRag(message, chatId, categoryIds, systemPrompt);
}
// 其他策略单查询处理
String rewrittenMessage = rewriteQuery(message, chatId, strategy);
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT")
.prompt()
.user(rewrittenMessage)
.advisors(s -> s.param(CONVERSATION_ID, chatId))
// 应用 RAG 知识库问答
.advisors(buildRetrievalAdvisor(4, categoryIds));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
}
ChatResponse chatResponse = spec.call().chatResponse();
return chatResponse.getResult().getOutput().getText();
}
/**
* RAG 知识库对话支持查询重写策略 + MCP 工具权限
*
* @param message 用户消息
* @param chatId 会话ID
* @param strategy 查询重写策略
* @param categoryIds 知识库分类过滤
* @param systemPrompt 角色人设/风格
* @param allowedMcpTools 允许的 MCP 工具列表null=不注册工具
* @return AI 回答
*/
public String doChatWithRagStrategy(String message, String chatId, String strategy, List<Long> categoryIds,
String systemPrompt, List<String> allowedMcpTools) {
// FAQ 优先匹配
Optional<String> faqAnswer = tryFaqMatch(message);
if (faqAnswer.isPresent()) {
log.info("FAQ 命中,直接返回标准答案: chatId={}", chatId);
return faqAnswer.get();
}
if ("MULTI_QUERY".equalsIgnoreCase(strategy)) {
return doChatWithMultiQueryRag(message, chatId, categoryIds, systemPrompt, allowedMcpTools);
}
String rewrittenMessage = rewriteQuery(message, chatId, strategy);
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT", allowedMcpTools)
.prompt()
.user(rewrittenMessage)
.advisors(s -> s.param(CONVERSATION_ID, chatId))
.advisors(buildRetrievalAdvisor(4, categoryIds));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
}
ChatResponse chatResponse = spec.call().chatResponse();
return chatResponse.getResult().getOutput().getText();
}
/**
* 根据策略对查询做预检索改写MULTI_QUERY 不走此方法
*/
private String rewriteQuery(String message, String chatId, String strategy) {
if (strategy == null || strategy.isEmpty()) {
return message;
}
try {
switch (strategy.toUpperCase()) {
case "REWRITE":
return rewriteQueryRewriter.doQueryRewrite(message);
case "TRANSLATION":
return translationQueryRewriter.doQueryRewrite(message);
case "COMPRESSION":
// 查询压缩需要对话历史从会话记忆中取最近若干条传入
// 利用多轮上下文把指代不清的追问补全为独立查询
List<Message> history = chatMemory.get(chatId, 10);
return compressionQueryRewriter.doQueryRewrite(message, history);
case "NONE":
default:
return message;
}
} catch (Exception e) {
// 查询重写失败时降级为原始查询避免 RAG_REWRITE 配置异常 InvalidApiKey导致整次对话不可用
log.warn("查询重写失败 [strategy={}, chatId={}],降级使用原始查询: {}", strategy, chatId, e.getMessage());
return message;
}
}
/**
* RAG 知识库进行对话支持查询重写策略SSE 流式传输
*
* @param message 用户消息
* @param chatId 会话ID
* @param strategy 查询重写策略NONE/REWRITE/TRANSLATION/COMPRESSION/MULTI_QUERY
* @param categoryIds 知识库分类过滤
* @param systemPrompt 角色人设/风格
* @return 流式回答
*/
public Flux<String> doChatWithRagStrategyByStream(String message, String chatId, String strategy,
List<Long> categoryIds, String systemPrompt) {
// FAQ 优先匹配命中则直接以流式形式返回标准答案
Optional<String> faqAnswer = tryFaqMatch(message);
if (faqAnswer.isPresent()) {
log.info("FAQ 命中(流式),直接返回标准答案: chatId={}", chatId);
return Flux.just(faqAnswer.get());
}
// 对于 MULTI_QUERY 策略需要先手动检索合并再流式生成
if ("MULTI_QUERY".equalsIgnoreCase(strategy)) {
return doChatWithMultiQueryRagByStream(message, chatId, categoryIds, systemPrompt);
}
String rewrittenMessage = rewriteQuery(message, chatId, strategy);
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT")
.prompt()
.user(rewrittenMessage)
.advisors(s -> s.param(CONVERSATION_ID, chatId))
.advisors(buildRetrievalAdvisor(4, categoryIds));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
}
return spec.stream().content();
}
/**
* RAG 知识库进行对话支持查询重写策略 + MCP 工具权限SSE 流式传输
*
* @param message 用户消息
* @param chatId 会话ID
* @param strategy 查询重写策略
* @param categoryIds 知识库分类过滤
* @param systemPrompt 角色人设/风格
* @param allowedMcpTools 允许的 MCP 工具列表null=不注册工具
* @return 流式回答
*/
public Flux<String> doChatWithRagStrategyByStream(String message, String chatId, String strategy,
List<Long> categoryIds, String systemPrompt,
List<String> allowedMcpTools) {
// 重置 MCP 工具调用状态
McpToolCallback.resetEvents();
McpToolCallback.resetCallRounds();
// FAQ 优先匹配命中则直接以流式形式返回标准答案
Optional<String> faqAnswer = tryFaqMatch(message);
if (faqAnswer.isPresent()) {
log.info("FAQ 命中(流式),直接返回标准答案: chatId={}", chatId);
return Flux.just(faqAnswer.get());
}
// 对于 MULTI_QUERY 策略需要先手动检索合并再流式生成
if ("MULTI_QUERY".equalsIgnoreCase(strategy)) {
return doChatWithMultiQueryRagByStream(message, chatId, categoryIds, systemPrompt, allowedMcpTools);
}
String rewrittenMessage = rewriteQuery(message, chatId, strategy);
ChatClient.ChatClientRequestSpec spec = getChatClient("CHAT", allowedMcpTools)
.prompt()
.user(rewrittenMessage)
.advisors(s -> s.param(CONVERSATION_ID, chatId))
.advisors(buildRetrievalAdvisor(4, categoryIds));
if (StringUtils.hasText(systemPrompt)) {
spec = spec.system(effectiveSystem(systemPrompt));
}
return appendMcpToolEvents(spec.stream().content());
}
/**
* 使用多路查询扩展的 RAG 对话
* 将原始查询扩展为多个语义不同的查询分别检索后按文档ID去重合并
* 再把合并后的资料注入 system 提示词不污染用户消息与会话记忆
*
* @param message 用户消息保持原样
* @param chatId 会话ID
* @param categoryIds 知识库分类过滤
* @param systemPrompt 角色人设/风格
* @return AI 回答
*/
private String doChatWithMultiQueryRag(String message, String chatId, List<Long> categoryIds, String systemPrompt) {
return doChatWithMultiQueryRag(message, chatId, categoryIds, systemPrompt, null);
}
private String doChatWithMultiQueryRag(String message, String chatId, List<Long> categoryIds,
String systemPrompt, List<String> allowedMcpTools) {
String ragSystem = buildMultiQueryRagSystem(message, categoryIds, systemPrompt);
ChatResponse chatResponse = getChatClient("CHAT", allowedMcpTools)
.prompt()
.system(ragSystem)
.user(message)
.advisors(s -> s.param(CONVERSATION_ID, chatId))
.call()
.chatResponse();
return chatResponse.getResult().getOutput().getText();
}
private Flux<String> doChatWithMultiQueryRagByStream(String message, String chatId, List<Long> categoryIds, String systemPrompt) {
return doChatWithMultiQueryRagByStream(message, chatId, categoryIds, systemPrompt, null);
}
private Flux<String> doChatWithMultiQueryRagByStream(String message, String chatId, List<Long> categoryIds,
String systemPrompt, List<String> allowedMcpTools) {
String ragSystem = buildMultiQueryRagSystem(message, categoryIds, systemPrompt);
return appendMcpToolEvents(getChatClient("CHAT", allowedMcpTools)
.prompt()
.system(ragSystem)
.user(message)
.advisors(s -> s.param(CONVERSATION_ID, chatId))
.stream()
.content());
}
/**
* 多路查询扩展 + 检索合并组合出注入了知识库资料的系统提示词
* 将原始查询扩展为多个语义不同的查询分别检索后按文档ID去重合并
* 资料注入 system不污染用户消息与会话记忆供同步与流式两条路径复用
*/
private String buildMultiQueryRagSystem(String message, List<Long> categoryIds, String systemPrompt) {
List<Document> mergedDocs = retrieveMultiQueryDocs(message, categoryIds);
log.info("多路检索合并后文档数: {}", mergedDocs.size());
// 拼接资料为上下文组合系统提示词
String context = mergedDocs.stream()
.map(Document::getText)
.filter(StringUtils::hasText)
.collect(Collectors.joining("\n\n---\n\n"));
return buildRagSystemPrompt(systemPrompt, context);
}
/**
* 多路查询扩展 + 检索去重合并返回命中的文档 RAG 回答与"引用来源"复用
*/
private List<Document> retrieveMultiQueryDocs(String message, List<Long> categoryIds) {
// 多路扩展依赖 RAG_REWRITE 模型扩展失败时退回原问题检索避免整次 RAG 不可用
List<String> expandedQueries;
try {
expandedQueries = multiQueryExpanderRewriter.doQueryRewrite(message);
} catch (Exception e) {
log.warn("多路查询扩展失败,降级为原始问题检索: {}", e.getMessage());
expandedQueries = List.of(message);
}
if (expandedQueries == null || expandedQueries.isEmpty()) {
expandedQueries = List.of(message);
}
log.info("多路查询扩展结果: {}", expandedQueries);
// 对每个查询分别检索按文档 ID 去重合并封顶 maxDocs避免上下文膨胀
final int maxDocs = 8;
Map<String, Document> merged = new LinkedHashMap<>();
for (String query : expandedQueries) {
if (!StringUtils.hasText(query) || merged.size() >= maxDocs) {
continue;
}
List<Document> docs = pgVectorVectorStore.similaritySearch(buildRagSearchRequest(4, categoryIds, query));
if (docs == null) {
continue;
}
for (Document doc : docs) {
if (merged.size() >= maxDocs) {
break;
}
merged.putIfAbsent(doc.getId(), doc);
}
}
return new ArrayList<>(merged.values());
}
/**
* 检索本次问题命中的知识库片段不生成回答用于在回答下方展示"引用来源"
* 复用与 RAG 回答完全相同的查询改写策略与分类范围确保来源即答案所依据的片段
*
* @param strategy 查询重写策略与回答保持一致
* @return 命中的文档片段 metadatadocumentId/title/sourceName/chunkIndex/distance
*/
public List<Document> retrieveRagSources(String message, String chatId, String strategy, List<Long> categoryIds) {
if (!StringUtils.hasText(message)) {
return Collections.emptyList();
}
if ("MULTI_QUERY".equalsIgnoreCase(strategy)) {
return retrieveMultiQueryDocs(message, categoryIds);
}
String rewritten = rewriteQuery(message, chatId, strategy);
List<Document> docs = pgVectorVectorStore.similaritySearch(buildRagSearchRequest(4, categoryIds, rewritten));
return docs != null ? docs : Collections.emptyList();
}
/**
* 组合 RAG 场景的系统提示词基础提示词 + 角色人设 + 检索到的资料
* 资料为空时不声称依据资料避免诱导模型编造
*/
private String buildRagSystemPrompt(String rolePrompt, String context) {
String base = effectiveSystem(rolePrompt);
if (!StringUtils.hasText(context)) {
return base;
}
return base + "\n\n【RAG回答硬性规则】\n"
+ ragPromptConfig.getAnswerRules() + "\n\n"
+ "【知识库资料】\n" + context;
}
private SearchRequest buildRagSearchRequest(int topK, Long categoryId) {
List<Long> categoryIds = categoryId != null ? List.of(categoryId) : Collections.emptyList();
return buildRagSearchRequest(topK, categoryIds);
}
private SearchRequest buildRagSearchRequest(int topK, List<Long> categoryIds) {
return ragSearchRequestBuilder(topK, categoryIds).build();
}
/**
* 带查询文本的检索请求供手动向量检索多路查询使用
* QuestionAnswerAdvisor 会自行设置查询文本故那条路径不需要此重载
*/
private SearchRequest buildRagSearchRequest(int topK, List<Long> categoryIds, String query) {
SearchRequest.Builder builder = ragSearchRequestBuilder(topK, categoryIds);
if (StringUtils.hasText(query)) {
builder.query(query);
}
return builder.build();
}
private SearchRequest.Builder ragSearchRequestBuilder(int topK, List<Long> categoryIds) {
SearchRequest.Builder builder = SearchRequest.builder()
.similarityThreshold(0.0)
.topK(topK);
Filter.Expression filterExpression = buildCategoryFilterExpression(categoryIds);
if (filterExpression != null) {
builder.filterExpression(filterExpression);
}
return builder;
}
private Advisor buildRetrievalAdvisor(int topK, List<Long> categoryIds) {
Filter.Expression filterExpression = buildCategoryFilterExpression(categoryIds);
// 自定义模板替换 Spring AI 默认模板注入与 MULTI_QUERY 路径一致的 RAG 回答规则
PromptTemplate qaTemplate = new PromptTemplate(
"【RAG回答硬性规则】\n" + ragPromptConfig.getAnswerRules() + "\n\n"
+ "【知识库资料】\n---------------------\n{context}\n---------------------\n\n"
+ "用户问题:{query}\n\n请按上述规则回答:");
return RetrievalAugmentationAdvisor.builder()
.documentRetriever(new VectorStoreDocumentRetriever(
pgVectorVectorStore,
0.0,
topK,
() -> filterExpression))
.queryAugmenter(ContextualQueryAugmenter.builder()
.allowEmptyContext(false)
.promptTemplate(qaTemplate)
.build())
.build();
}
private Filter.Expression buildCategoryFilterExpression(List<Long> categoryIds) {
List<Object> values = normalizeCategoryIds(categoryIds).stream()
.map(value -> (Object) value)
.toList();
if (values.isEmpty()) {
return null;
}
FilterExpressionBuilder filterBuilder = new FilterExpressionBuilder();
return filterBuilder.in("categoryId", values).build();
}
private List<String> normalizeCategoryIds(List<Long> categoryIds) {
if (categoryIds == null || categoryIds.isEmpty()) {
return Collections.emptyList();
}
return categoryIds.stream()
.filter(java.util.Objects::nonNull)
.filter(id -> id > 0)
.map(String::valueOf)
.distinct()
.toList();
}
/**
* RAG 知识库进行对话(另外一种使用方式)
*
* @param message
* @param chatId
* @return
*/
public String doChatWithRagEnhance(String message, String chatId) {
Advisor retrievalAugmentationAdvisor = RetrievalAugmentationAdvisor.builder()
.documentRetriever(VectorStoreDocumentRetriever.builder()
.vectorStore(pgVectorVectorStore)
.similarityThreshold(0.5)
.topK(4)
.build())
.queryAugmenter(ContextualQueryAugmenter.builder()
.allowEmptyContext(false) // 不允许模型在没有找到相关文档的情况下也生成回答
.build())
.build();
ChatResponse chatResponse = getChatClient("CHAT")
.prompt()
.user(message)
.advisors(spec -> spec.param(CONVERSATION_ID, chatId))
// 应用 RAG 知识库问答
.advisors(retrievalAugmentationAdvisor)
.call()
.chatResponse();
return chatResponse.getResult().getOutput().getText();
}
}

78
src/main/java/com/wok/supportbot/app/ChatContext.java

@ -0,0 +1,78 @@
package com.wok.supportbot.app;
import java.util.List;
/**
* 一次对话请求的完整上下文值对象
* <p>
* 统一封装所有对话入口普通对话 / RAG 对话 / Open API / 来源检索所需的参数
* {@code ChatPipeline} / {@code RagPipeline} 编排使用替代原本散落在
* {@code AssistantApp} doChat* 方法中按参数个数重载的组合爆炸
* <p>
* 不可变需要调整某一字段时使用 {@code withXxx} 派生新实例
*
* @param message 用户原始消息
* @param chatId 会话 ID多轮记忆键
* @param appType ChatModel 应用类型CHAT / RAG_REWRITE 默认 CHAT
* @param systemPrompt 角色人设/系统提示词可为 null
* @param allowedMcpTools 允许的 MCP 工具名列表null/=允许所有["*"]=全部
* @param categoryIds 知识库分类隔离范围可为空不限制
* @param rewriteStrategy RAG 查询重写策略REWRITE / TRANSLATION / COMPRESSION / MULTI_QUERY可为 null
* @param enableRag 是否启用 RAG 检索false=普通对话
* @param streaming 是否流式输出
*/
public record ChatContext(
String message,
String chatId,
String appType,
String systemPrompt,
List<String> allowedMcpTools,
List<Long> categoryIds,
String rewriteStrategy,
boolean enableRag,
boolean streaming
) {
/** 默认应用类型 */
private static final String DEFAULT_APP_TYPE = "CHAT";
/**
* 紧凑构造器规范化默认值保证 appType 非空
*/
public ChatContext {
if (appType == null || appType.isBlank()) {
appType = DEFAULT_APP_TYPE;
}
}
/**
* 便捷构造仅指定核心字段其余取默认值 RAG非流式
*/
public static ChatContext of(String message, String chatId) {
return new ChatContext(message, chatId, DEFAULT_APP_TYPE, null, null, null, null, false, false);
}
public ChatContext withSystemPrompt(String systemPrompt) {
return new ChatContext(message, chatId, appType, systemPrompt, allowedMcpTools, categoryIds, rewriteStrategy, enableRag, streaming);
}
public ChatContext withAllowedMcpTools(List<String> allowedMcpTools) {
return new ChatContext(message, chatId, appType, systemPrompt, allowedMcpTools, categoryIds, rewriteStrategy, enableRag, streaming);
}
public ChatContext withCategoryIds(List<Long> categoryIds) {
return new ChatContext(message, chatId, appType, systemPrompt, allowedMcpTools, categoryIds, rewriteStrategy, enableRag, streaming);
}
public ChatContext withRewriteStrategy(String rewriteStrategy) {
return new ChatContext(message, chatId, appType, systemPrompt, allowedMcpTools, categoryIds, rewriteStrategy, enableRag, streaming);
}
public ChatContext withEnableRag(boolean enableRag) {
return new ChatContext(message, chatId, appType, systemPrompt, allowedMcpTools, categoryIds, rewriteStrategy, enableRag, streaming);
}
public ChatContext withStreaming(boolean streaming) {
return new ChatContext(message, chatId, appType, systemPrompt, allowedMcpTools, categoryIds, rewriteStrategy, enableRag, streaming);
}
}

154
src/main/java/com/wok/supportbot/app/ChatPipeline.java

@ -0,0 +1,154 @@
package com.wok.supportbot.app;
import com.wok.supportbot.rag.RagContext;
import com.wok.supportbot.rag.RagPipeline;
import com.wok.supportbot.service.IntentRouter;
import jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.document.Document;
import org.springframework.stereotype.Component;
import org.springframework.util.StringUtils;
import java.util.Collections;
import java.util.List;
import java.util.Locale;
import java.util.Optional;
/**
* 统一对话管道编排层
* <p>
* 编排一次完整对话的决策流程意图路由 FAQ 优先 RAG 检索 组装系统提示词与用户消息
* 产出 {@link ChatRequest} 交由 {@code AssistantApp} 执行实际的 {@code call()} / {@code stream()}
* <p>
* 设计说明本类为纯编排层不持有 ChatClientChatClient 构建与 Advisor 链装配仍在
* {@code AssistantApp}因此 {@code call} / {@code stream} {@code AssistantApp} 承担
* 避免 {@code ChatPipeline} {@code AssistantApp} 循环依赖
* <p>
* 接入 {@link IntentRouter} 替代原 {@code AiController.shouldBypassKnowledgeRetrieval} 的硬编码寒暄词判断
* 寒暄词列表保留为快速路径与兜底IntentRouter 负责细粒度意图分类二者命中其一即跳过 KB 检索
*/
@Component
@Slf4j
public class ChatPipeline {
/** IntentRouter 判为 CHITCHAT 的置信度阈值,低于此值视为不确定,继续走 RAG */
private static final double CHITCHAT_CONFIDENCE_THRESHOLD = 0.6;
@Resource
private IntentRouter intentRouter;
@Resource
private RagPipeline ragPipeline;
/**
* 编排一次对话请求产出执行决策
* <p>
* 决策分支
* <ul>
* <li>未启用 RAG普通对话 / 严格隔离下 KB 拒绝 用原始 message基础 system</li>
* <li>寒暄/闲聊IntentRouter 或寒暄词命中 同上跳过 KB 检索</li>
* <li>FAQ 命中 直接返回标准答案不调用 ChatClient</li>
* <li>RAG 生成 资料块注入 system重写后查询作为 user 消息</li>
* </ul>
*
* @param ctx 对话上下文
* @return 执行决策
*/
public ChatRequest buildRequest(ChatContext ctx) {
String baseSystem = effectiveSystem(ctx.systemPrompt());
// 普通对话enableRag=false Controller isKbDenied 强制置 false 的情况
if (!ctx.enableRag()) {
return new ChatRequest(ctx, ctx.message(), baseSystem, Optional.empty());
}
// 意图路由寒暄/闲聊跳过 KB 检索避免召回"最相似但实际无关"的片段
if (shouldBypassRag(ctx.message())) {
return new ChatRequest(ctx, ctx.message(), baseSystem, Optional.empty());
}
// RAG 检索 FAQ 优先匹配
RagContext rag = ragPipeline.retrieve(ctx);
if (rag.faqHit()) {
return new ChatRequest(ctx, ctx.message(), baseSystem, rag.faqAnswer());
}
// RAG 生成资料块注入 system重写后查询作为 user 消息
String finalSystem = baseSystem + ragPipeline.buildRagContextBlock(rag.contextText());
return new ChatRequest(ctx, rag.rewrittenQuery(), finalSystem, Optional.empty());
}
/**
* 判断是否跳过 KB 检索
* <p>
* 先用寒暄词列表做快速路径 LLM 开销同时作为 IntentRouter 失败的兜底
* 未命中再调 {@link IntentRouter} 做细粒度意图分类
*/
private boolean shouldBypassRag(String message) {
if (isChitchat(message)) {
return true;
}
try {
IntentRouter.IntentResult intent = intentRouter.route(message);
if ("CHITCHAT".equals(intent.getIntent()) && intent.getConfidence() >= CHITCHAT_CONFIDENCE_THRESHOLD) {
return true;
}
} catch (Exception e) {
log.debug("意图路由异常,沿用 RAG: {}", e.getMessage());
}
return false;
}
/**
* 寒暄词快速判断问候/感谢/告别等短消息无知识库检索意图
* 与原 {@code AiController.shouldBypassKnowledgeRetrieval} 逻辑一致作为 IntentRouter 的快速路径与兜底
* <p>
* {@code buildRequest} "引用来源"等不需 LLM 意图分类的场景共用
*/
public boolean isChitchat(String message) {
if (!StringUtils.hasText(message)) {
return true;
}
String normalized = message.trim()
.toLowerCase(Locale.ROOT)
.replaceAll("[\\s,。!?!?,.;;::、~~…]+", "");
if (normalized.isEmpty()) {
return true;
}
if (normalized.length() > 12) {
return false;
}
return List.of(
"你好", "您好", "hello", "hi", "哈喽", "嗨", "在吗",
"早上好", "上午好", "中午好", "下午好", "晚上好",
"谢谢", "感谢", "多谢", "好的", "好", "嗯", "嗯嗯",
"再见", "拜拜", "辛苦了"
).contains(normalized);
}
/**
* 组合系统提示词基础客服提示词 + 角色人设
* {@code AssistantApp.effectiveSystem} 保持一致阶段四统一为单一实现
* 当前基础提示词为空串仅有角色人设时返回 {@code "\n\n【当前角色设定】\n" + rolePrompt}
*/
private String effectiveSystem(String rolePrompt) {
if (!StringUtils.hasText(rolePrompt)) {
return "";
}
return "\n\n【当前角色设定】\n" + rolePrompt;
}
/**
* 统一检索引用来源不含 FAQ 匹配与意图路由
* 寒暄词无 KB 检索意图时直接返回空列表避免向量库召回"最相似但实际无关"的片段
*
* @param ctx 对话上下文使用 message / rewriteStrategy / categoryIds
* @return 命中的知识库片段寒暄词或无命中时返回空列表
*/
public List<Document> retrieveSources(ChatContext ctx) {
if (!StringUtils.hasText(ctx.message()) || isChitchat(ctx.message())) {
return Collections.emptyList();
}
return ragPipeline.retrieveDocuments(ctx);
}
}

34
src/main/java/com/wok/supportbot/app/ChatRequest.java

@ -0,0 +1,34 @@
package com.wok.supportbot.app;
import java.util.Optional;
/**
* 一次对话经 {@link ChatPipeline} 编排后的执行决策值对象
* <p>
* {@link ChatPipeline#buildRequest} 产出本对象{@code AssistantApp} 据此构造
* {@code ChatClientRequestSpec} 并执行 {@code call()} / {@code stream()}
* <p>
* 字段含义
* <ul>
* <li>{@link #faqAnswer()}FAQ 命中时直接返回标准答案跳过 ChatClient 调用</li>
* <li>{@link #finalMessage()}传给模型的用户消息RAG 场景为重写后的查询否则为原始 message</li>
* <li>{@link #finalSystemPrompt()}传给模型的系统提示词角色人设 + RAG 资料块可为空</li>
* </ul>
*
* @param ctx 原始上下文
* @param finalMessage 传给模型的用户消息
* @param finalSystemPrompt 传给模型的系统提示词可为 null/
* @param faqAnswer FAQ 命中答案未命中为 {@link Optional#empty()}
*/
public record ChatRequest(
ChatContext ctx,
String finalMessage,
String finalSystemPrompt,
Optional<String> faqAnswer
) {
/** FAQ 是否命中 */
public boolean faqHit() {
return faqAnswer != null && faqAnswer.isPresent();
}
}

17
src/main/java/com/wok/supportbot/app/ChatResult.java

@ -0,0 +1,17 @@
package com.wok.supportbot.app;
import com.wok.supportbot.mcp.McpToolCallback;
import java.util.List;
/**
* 同步对话的完整结果值对象
* <p>
* 除了回答文本还携带本次触发的 MCP 工具调用事件让同步对话也能像流式对话一样
* 展示工具调用过程原仅流式 {@code appendMcpToolEvents} 追加事件同步路径无事件
*
* @param text AI 回答文本
* @param mcpEvents 本次触发的 MCP 工具调用事件无调用时为空列表
*/
public record ChatResult(String text, List<McpToolCallback.ToolCallEvent> mcpEvents) {
}

282
src/main/java/com/wok/supportbot/controller/AiController.java

@ -1,7 +1,10 @@
package com.wok.supportbot.controller;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.ChatPipeline;
import com.wok.supportbot.config.RoleAccessConfig;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.service.ConversationService;
import com.wok.supportbot.service.CustomerServiceRoleService;
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope;
@ -9,7 +12,6 @@ import jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.document.Document;
import org.springframework.http.MediaType;
import org.springframework.http.codec.ServerSentEvent;
import org.springframework.util.StringUtils;
import org.springframework.web.bind.annotation.DeleteMapping;
import org.springframework.web.bind.annotation.GetMapping;
@ -19,16 +21,12 @@ import org.springframework.web.bind.annotation.RequestBody;
import org.springframework.web.bind.annotation.RequestMapping;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
import reactor.core.publisher.Flux;
import java.io.IOException;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Locale;
import java.util.Map;
@ -45,163 +43,74 @@ public class AiController {
private ConversationService conversationService;
@Resource
private RoleAccessConfig roleAccessConfig;
@Resource
private CategoryFilter categoryFilter;
@Resource
private ChatPipeline chatPipeline;
/**
* 同步调用 AI 智能客服应用
*
* @param message 用户消息
* @param chatId 会话ID
* @param roleId 客服角色ID可选命中角色时由服务端强制套用该角色人设
* @param systemPrompt 角色人设兜底仅在未命中角色时生效
* @return AI 回答
* 同步调用 AI 智能客服应用已委托 ChatPipeline 新管道后续版本移除
* @deprecated 请使用 {@link AssistantApp#chat(ChatContext)}
*/
@GetMapping("/assistant_app/chat/sync")
@Deprecated
public String doChatWithAssistantAppSync(String message, String chatId, Long roleId, String accountId, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
String result = assistantApp.doChat(message, chatId, resolveSystemPrompt(scope, systemPrompt), scope.allowedMcpTools());
return result;
ChatContext ctx = buildChatContext(message, chatId, roleId, accountId, systemPrompt);
return assistantApp.chat(ctx);
}
/**
* SSE 流式调用 AI 智能客服应用
* 返回Flux 响应式؜对象并且添加 SSE 对应的 MediaType
*
* @param message
* @param chatId
* @return
* SSE 流式调用 AI 智能客服应用已委托 ChatPipeline 新管道后续版本移除
* @deprecated 请使用 {@link AssistantApp#chatStream(ChatContext)}
*/
@GetMapping(value = "/assistant_app/chat/sse", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
@Deprecated
public Flux<String> doChatWithLoveAppSSE(String message, String chatId, Long roleId, String accountId, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
return assistantApp.doChatByStream(message, chatId, resolveSystemPrompt(scope, systemPrompt), scope.allowedMcpTools());
ChatContext ctx = buildChatContext(message, chatId, roleId, accountId, systemPrompt);
return assistantApp.chatStream(ctx);
}
/**
* SSE 流式调用 AI 智能客服应用
* 返回 Flux 对象并且؜设置泛型为 ServerSentEvent使用这种方式可以省略 MediaType
*
* @param message
* @param chatId
* @return
*/
@GetMapping(value = "/assistant_app/chat/server_sent_event")
public Flux<ServerSentEvent<String>> doChatWithAssistantAppServerSentEvent(String message, String chatId, Long roleId, String accountId, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
return assistantApp.doChatByStream(message, chatId, resolveSystemPrompt(scope, systemPrompt), scope.allowedMcpTools())
.map(chunk -> ServerSentEvent.<String>builder()
.data(chunk)
.build());
}
/**
* SSE 流式调用 AI 智能客服应用
* 使用 SSEEmiter؜通过 send 方法持续向 SseEmitter 发送消息
*
* @param message
* @param chatId
* @return
*/
@GetMapping(value = "/assistant_app/chat/sse_emitter")
public SseEmitter doChatWithAssistantAppServerSseEmitter(String message, String chatId, Long roleId, String accountId, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
// 创建一个超时时间较长的 SseEmitter
SseEmitter sseEmitter = new SseEmitter(180000L); // 3 分钟超时
// 获取 Flux 响应式数据流并且直接通过订阅推送给 SseEmitter
assistantApp.doChatByStream(message, chatId, resolveSystemPrompt(scope, systemPrompt), scope.allowedMcpTools())
.subscribe(chunk -> {
try {
sseEmitter.send(chunk);
} catch (IOException e) {
sseEmitter.completeWithError(e);
}
}, sseEmitter::completeWithError, sseEmitter::complete);
// 返回
return sseEmitter;
}
/**
* RAG 知识库同步对话支持查询重写策略
*
* @param message 用户消息
* @param chatId 会话ID
* @param rewriteStrategy 查询重写策略可选NONE/REWRITE/TRANSLATION/COMPRESSION/MULTI_QUERY默认为 MULTI_QUERY
* @param roleId 客服角色ID可选命中角色时由服务端强制限定检索范围与人设客户端无法跨域
* @param categoryId 单个分类过滤仅未命中角色时生效
* @param categoryIds 多个分类过滤逗号分隔仅未命中角色时生效
* @param systemPrompt 角色人设兜底仅未命中角色时生效
* @return AI 回答
* RAG 知识库同步对话已委托 ChatPipeline 新管道后续版本移除
* @deprecated 请使用 {@link AssistantApp#chat(ChatContext)}
*/
@GetMapping("/assistant_app/chat/rag/sync")
@Deprecated
public String doChatWithRagSync(String message, String chatId, String rewriteStrategy, Long roleId, String accountId, Long categoryId, String categoryIds, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
String sys = resolveSystemPrompt(scope, systemPrompt);
// 严格隔离未授权任何知识库的角色退回普通对话绝不检索 KB
if (isKbDenied(scope) || shouldBypassKnowledgeRetrieval(message)) {
return assistantApp.doChat(message, chatId, sys, scope.allowedMcpTools());
}
try {
return assistantApp.doChatWithRagStrategy(
message, chatId, normalizeStrategy(rewriteStrategy),
resolveCategoryIds(scope, categoryId, categoryIds), sys, scope.allowedMcpTools());
} catch (Exception e) {
log.error("RAG 对话失败 [strategy={}, chatId={}]: {}", rewriteStrategy, chatId, e.getMessage(), e);
return "抱歉,知识库检索出现异常,请稍后重试。";
}
ChatContext ctx = buildRagChatContext(message, chatId, rewriteStrategy, roleId, accountId, categoryId, categoryIds, systemPrompt);
return assistantApp.chat(ctx);
}
/**
* RAG 知识库流式对话SSE支持查询重写策略
*
* @param message 用户消息
* @param chatId 会话ID
* @param rewriteStrategy 查询重写策略可选NONE/REWRITE/TRANSLATION/COMPRESSION/MULTI_QUERY默认为 MULTI_QUERY
* @param roleId 客服角色ID可选命中角色时由服务端强制限定检索范围与人设客户端无法跨域
* @param categoryId 单个分类过滤仅未命中角色时生效
* @param categoryIds 多个分类过滤逗号分隔仅未命中角色时生效
* @param systemPrompt 角色人设兜底仅未命中角色时生效
* @return 流式 AI 回答
* RAG 知识库流式对话已委托 ChatPipeline 新管道后续版本移除
* @deprecated 请使用 {@link AssistantApp#chatStream(ChatContext)}
*/
@GetMapping(value = "/assistant_app/chat/rag/sse", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
@Deprecated
public Flux<String> doChatWithRagSSE(String message, String chatId, String rewriteStrategy, Long roleId, String accountId, Long categoryId, String categoryIds, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
String sys = resolveSystemPrompt(scope, systemPrompt);
// 严格隔离未授权任何知识库的角色退回普通流式对话绝不检索 KB
if (isKbDenied(scope) || shouldBypassKnowledgeRetrieval(message)) {
return assistantApp.doChatByStream(message, chatId, sys, scope.allowedMcpTools());
}
return assistantApp.doChatWithRagStrategyByStream(
message, chatId, normalizeStrategy(rewriteStrategy),
resolveCategoryIds(scope, categoryId, categoryIds), sys, scope.allowedMcpTools());
ChatContext ctx = buildRagChatContext(message, chatId, rewriteStrategy, roleId, accountId, categoryId, categoryIds, systemPrompt);
return assistantApp.chatStream(ctx);
}
/**
* RAG 引用来源返回本次问题命中的知识库片段不生成回答供回答下方展示来源
* 复用与 RAG 回答相同的角色范围分类过滤与查询改写策略确保来源即答案所依据的片段
* RAG 引用来源已委托 ChatPipeline 新管道后续版本移除
* @deprecated 请使用 {@link AssistantApp#retrieveSources(ChatContext)}
*/
@GetMapping("/assistant_app/rag/sources")
@Deprecated
public Map<String, Object> getRagSources(String message, String chatId, String rewriteStrategy,
Long roleId, String accountId, Long categoryId, String categoryIds) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
if (message == null || message.isBlank() || isKbDenied(scope) || shouldBypassKnowledgeRetrieval(message)) {
if (message == null || message.isBlank() || isKbDenied(scope) || chatPipeline.isChitchat(message)) {
return Map.of("success", true, "data", List.of());
}
try {
List<Long> cats = resolveCategoryIds(scope, categoryId, categoryIds);
List<Document> docs = assistantApp.retrieveRagSources(message, chatId, normalizeStrategy(rewriteStrategy), cats);
ChatContext ctx = new ChatContext(message, chatId, "CHAT", null, null, cats,
normalizeStrategy(rewriteStrategy), true, false);
List<Document> docs = assistantApp.retrieveSources(ctx);
List<Map<String, Object>> out = new ArrayList<>();
for (Document doc : docs) {
Map<String, Object> meta = doc.getMetadata();
@ -222,18 +131,32 @@ public class AiController {
}
}
// ==================== ChatContext 构建辅助 ====================
/** 构造普通对话的 ChatContext(enableRag=false)。 */
private ChatContext buildChatContext(String message, String chatId, Long roleId, String accountId, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
return new ChatContext(message, chatId, "CHAT", resolveSystemPrompt(scope, systemPrompt),
scope.allowedMcpTools(), null, null, false, false);
}
/** 构造 RAG 对话的 ChatContext(含严格隔离判断:KbDenied 则 enableRag=false)。 */
private ChatContext buildRagChatContext(String message, String chatId, String rewriteStrategy,
Long roleId, String accountId, Long categoryId, String categoryIds, String systemPrompt) {
AccountRoleContext context = resolveAccountRole(accountId, roleId);
bindConversation(chatId, context);
RoleScope scope = customerServiceRoleService.getRoleScope(context.roleId());
String sys = resolveSystemPrompt(scope, systemPrompt);
boolean enableRag = !isKbDenied(scope);
List<Long> cats = resolveCategoryIds(scope, categoryId, categoryIds);
return new ChatContext(message, chatId, "CHAT", sys, scope.allowedMcpTools(), cats,
normalizeStrategy(rewriteStrategy), enableRag, false);
}
// ==================== SDK 会话管理接口带账户归属校验 ====================
/**
* SDK 获取会话列表分页强制按账户 + 角色过滤
* 仅返回属于指定 accountId + roleId 的会话防止跨账户数据泄露
*
* @param page 页码默认1
* @param size 每页大小默认20
* @param accountId 外部账户 IDSDK userId必传
* @param roleId 角色 IDSDK integrateId必传
* @return 分页会话列表
*/
@GetMapping("/sdk/conversation/list")
public Map<String, Object> sdkListConversations(
@RequestParam(defaultValue = "1") int page,
@ -256,15 +179,6 @@ public class AiController {
}
}
/**
* SDK 获取会话消息含所有权校验
* 先验证会话是否属于指定 accountId + roleId通过后才返回消息
*
* @param conversationId 会话 ID
* @param accountId 外部账户 ID必传
* @param roleId 角色 ID必传
* @return 消息列表
*/
@GetMapping("/sdk/conversation/{id}/messages")
public Map<String, Object> sdkGetConversationMessages(
@PathVariable("id") String conversationId,
@ -282,15 +196,6 @@ public class AiController {
}
}
/**
* SDK 删除会话含所有权校验
* 仅允许删除属于自己的会话
*
* @param conversationId 会话 ID
* @param accountId 外部账户 ID必传
* @param roleId 角色 ID必传
* @return 删除结果
*/
@DeleteMapping("/sdk/conversation/{id}")
public Map<String, Object> sdkDeleteConversation(
@PathVariable("id") String conversationId,
@ -307,15 +212,6 @@ public class AiController {
}
}
/**
* SDK 导出会话含所有权校验返回文本内容
* 通过 accountId + roleId 验证会话归属后导出
*
* @param conversationId 会话 ID
* @param accountId 外部账户 ID必传
* @param roleId 角色 ID必传
* @return 导出的会话文本
*/
@GetMapping(value = "/sdk/conversation/{id}/export", produces = MediaType.TEXT_PLAIN_VALUE + ";charset=UTF-8")
public String sdkExportConversation(
@PathVariable("id") String conversationId,
@ -332,14 +228,6 @@ public class AiController {
}
}
/**
* SDK 截断会话含所有权校验
* 用于"编辑历史消息并重发"场景仅允许截断属于自己的会话
*
* @param conversationId 会话 ID
* @param body { "userTurn": 第几条用户消息, "accountId": "xxx", "roleId": 1 }
* @return 逻辑删除的消息条数
*/
@PostMapping("/sdk/conversation/{id}/truncate")
public Map<String, Object> sdkTruncateConversation(
@PathVariable("id") String conversationId,
@ -360,52 +248,23 @@ public class AiController {
}
}
/**
* 严格隔离模式下命中角色但未绑定任何知识库分类 拒绝检索任何 KB
* 非严格模式下永远返回 false未绑定 = 可检索全部沿用通用角色语义
*/
// ==================== 私有辅助方法 ====================
/** 严格隔离模式下,命中角色但未绑定任何知识库分类,拒绝检索 KB。 */
private boolean isKbDenied(RoleScope scope) {
return roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty();
}
/**
* 寒暄感谢告别这类短消息没有知识库检索意图直接走普通对话
* 否则向量库会被迫召回一个最相似但实际无关的片段导致回答下方出现误导性来源
*/
private boolean shouldBypassKnowledgeRetrieval(String message) {
if (!StringUtils.hasText(message)) {
return true;
}
String normalized = message.trim()
.toLowerCase(Locale.ROOT)
.replaceAll("[\\s,。!?!?,.;;::、~~…]+", "");
if (normalized.isEmpty()) {
return true;
}
if (normalized.length() > 12) {
return false;
}
return List.of(
"你好", "您好", "hello", "hi", "哈喽", "嗨", "在吗",
"早上好", "上午好", "中午好", "下午好", "晚上好",
"谢谢", "感谢", "多谢", "好的", "好", "嗯", "嗯嗯",
"再见", "拜拜", "辛苦了"
).contains(normalized);
}
/** 未指定策略时默认 MULTI_QUERY(多路扩展)。 */
private String normalizeStrategy(String rewriteStrategy) {
return (rewriteStrategy != null && !rewriteStrategy.isEmpty()) ? rewriteStrategy : "MULTI_QUERY";
}
/**
* 命中角色时用 roleId 套用角色人设和知识库范围
* accountId SDK 宿主侧的外部用户标识后端不再创建或查询本地账号也不允许它覆盖 roleId
*/
private AccountRoleContext resolveAccountRole(String accountId, Long fallbackRoleId) {
String externalAccountId = StringUtils.hasText(accountId) ? accountId.trim() : null;
return new AccountRoleContext(externalAccountId, fallbackRoleId);
}
private void bindConversation(String chatId, AccountRoleContext context) {
conversationService.bindConversation(chatId, context.accountId(), context.roleId());
}
@ -417,11 +276,7 @@ public class AiController {
return fallbackSystemPrompt;
}
/**
* 解析最终的检索分类范围
* 命中角色时强制使用角色绑定的分类忽略客户端传入防止越权跨域
* 未命中角色时沿用客户端传入的分类便于直接调试/无角色调用
*/
/** 命中角色时强制使用角色绑定的分类(忽略客户端传入),防止越权;未命中角色沿用客户端传入。 */
private List<Long> resolveCategoryIds(RoleScope scope, Long categoryId, String categoryIds) {
if (scope.hasRole()) {
return scope.categoryIds();
@ -434,16 +289,7 @@ public class AiController {
}
private List<Long> parseCategoryIds(String categoryIds) {
if (categoryIds == null || categoryIds.isBlank()) {
return Collections.emptyList();
}
return Arrays.stream(categoryIds.split(","))
.map(String::trim)
.filter(item -> !item.isEmpty())
.map(Long::valueOf)
.filter(id -> id > 0)
.distinct()
.toList();
return categoryFilter.parse(categoryIds);
}
private record AccountRoleContext(String accountId, Long roleId) {

13
src/main/java/com/wok/supportbot/controller/CustomerServiceRoleController.java

@ -1,5 +1,6 @@
package com.wok.supportbot.controller;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.service.CustomerServiceRoleService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.ResponseEntity;
@ -22,6 +23,8 @@ public class CustomerServiceRoleController {
@Autowired
private CustomerServiceRoleService customerServiceRoleService;
@Autowired
private CategoryFilter categoryFilter;
@GetMapping("/role/list")
@PreAuthorize("hasRole('admin')")
@ -124,15 +127,7 @@ public class CustomerServiceRoleController {
}
private List<Long> parseCategoryIds(Object rawCategoryIds) {
if (!(rawCategoryIds instanceof List<?> list)) {
return List.of();
}
return list.stream()
.filter(Objects::nonNull)
.map(item -> item instanceof Number number ? number.longValue() : Long.parseLong(item.toString()))
.filter(id -> id > 0)
.distinct()
.toList();
return categoryFilter.parse(rawCategoryIds);
}
private static String str(Object v) {

26
src/main/java/com/wok/supportbot/controller/DocumentController.java

@ -4,6 +4,7 @@ import com.wok.supportbot.entity.CategoryNode;
import com.wok.supportbot.entity.KnowledgeCategory;
import com.wok.supportbot.entity.KnowledgeDocument;
import com.wok.supportbot.entity.SearchResult;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.service.DocumentService;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.ResponseEntity;
@ -27,6 +28,8 @@ public class DocumentController {
@Autowired
private DocumentService documentService;
@Autowired
private CategoryFilter categoryFilter;
// ==================== 上传校验常量 ====================
@ -676,28 +679,7 @@ public class DocumentController {
}
private List<Long> parseCategoryIds(Object rawCategoryIds) {
if (rawCategoryIds == null) {
return List.of();
}
if (rawCategoryIds instanceof List<?> list) {
return list.stream()
.filter(Objects::nonNull)
.map(item -> item instanceof Number number ? number.longValue() : Long.parseLong(item.toString()))
.filter(id -> id > 0)
.distinct()
.toList();
}
String text = rawCategoryIds.toString();
if (text.isBlank()) {
return List.of();
}
return Arrays.stream(text.split(","))
.map(String::trim)
.filter(item -> !item.isEmpty())
.map(Long::parseLong)
.filter(id -> id > 0)
.distinct()
.toList();
return categoryFilter.parse(rawCategoryIds);
}
// ==================== 标签管理 ====================

128
src/main/java/com/wok/supportbot/controller/OpenApiController.java

@ -1,16 +1,24 @@
package com.wok.supportbot.controller;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.config.RoleAccessConfig;
import com.wok.supportbot.entity.ApiKey;
import com.wok.supportbot.entity.SearchResult;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.rag.HybridSearchService;
import com.wok.supportbot.rag.SearchMode;
import com.wok.supportbot.service.ApiKeyService;
import com.wok.supportbot.service.CustomerServiceRoleService;
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import jakarta.servlet.http.HttpServletRequest;
import lombok.extern.slf4j.Slf4j;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity;
import org.springframework.util.StringUtils;
import org.springframework.web.bind.annotation.*;
import reactor.core.publisher.Flux;
@ -37,20 +45,28 @@ public class OpenApiController {
@Autowired
private ApiKeyService apiKeyService;
@Autowired
private CustomerServiceRoleService customerServiceRoleService;
@Autowired
private CategoryFilter categoryFilter;
@Autowired
private RoleAccessConfig roleAccessConfig;
private static final ObjectMapper objectMapper = new ObjectMapper();
/**
* 同步对话接口
*
* @param message 用户消息
* @param roleId 客服角色ID可选
* @param chatId 会话ID可选不传则自动生成
* @param request HTTP 请求含已鉴权的 API Key 信息
* @return AI 回答
* 同步对话接口已接入 ChatPipeline补齐角色/RAG/FAQ/MCP/分类隔离能力
*/
@PostMapping("/chat")
public ResponseEntity<Map<String, Object>> chat(
@RequestParam String message,
@RequestParam(required = false) String roleId,
@RequestParam(required = false) String chatId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String rewriteStrategy,
@RequestParam(required = false) Boolean enableRag,
HttpServletRequest request) {
try {
ApiKey apiKey = getApiKeyFromRequest(request);
@ -58,7 +74,9 @@ public class OpenApiController {
? chatId
: "openapi-" + apiKey.getId() + "-" + System.currentTimeMillis();
String reply = assistantApp.doChat(message, resolvedChatId);
ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId,
categoryIds, rewriteStrategy, enableRag, false);
String reply = assistantApp.chat(ctx);
Map<String, Object> result = new HashMap<>();
result.put("success", true);
@ -77,25 +95,25 @@ public class OpenApiController {
}
/**
* SSE 流式对话接口
*
* @param message 用户消息
* @param chatId 会话ID可选
* @param request HTTP 请求
* @return SSE 流式回答
* SSE 流式对话接口已接入 ChatPipeline补齐角色/RAG/FAQ/MCP/分类隔离能力
*/
@GetMapping(value = "/chat/stream", produces = MediaType.TEXT_EVENT_STREAM_VALUE)
public Flux<String> chatStream(
@RequestParam String message,
@RequestParam(required = false) String roleId,
@RequestParam(required = false) String chatId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String rewriteStrategy,
@RequestParam(required = false) Boolean enableRag,
HttpServletRequest request) {
ApiKey apiKey = getApiKeyFromRequest(request);
String resolvedChatId = (chatId != null && !chatId.isBlank())
? chatId
: "openapi-stream-" + apiKey.getId() + "-" + System.currentTimeMillis();
return assistantApp.doChatByStream(message, resolvedChatId);
ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId,
categoryIds, rewriteStrategy, enableRag, true);
return assistantApp.chatStream(ctx);
}
/**
@ -137,7 +155,7 @@ public class OpenApiController {
}
/**
* request attribute 获取已鉴权的 API Key 信息
* request attribute 获取已鉴权的 API Key 信息
*/
private ApiKey getApiKeyFromRequest(HttpServletRequest request) {
Object attr = request.getAttribute("apiKey");
@ -146,4 +164,82 @@ public class OpenApiController {
}
throw new IllegalStateException("API Key 鉴权信息缺失");
}
/**
* 构造 Open API 对话的 ChatContext含角色解析严格隔离判断分类过滤
*/
private ChatContext buildOpenApiChatContext(String message, String chatId, ApiKey apiKey,
String roleIdStr, String categoryIds, String rewriteStrategy, Boolean enableRag, boolean streaming) {
// 1. 解析 API Key 允许的角色列表
List<Long> allowedRoleIds = parseAllowedRoleIds(apiKey.getRoleIds());
Long roleId = resolveRoleId(roleIdStr, allowedRoleIds);
// 2. 获取角色范围
RoleScope scope = customerServiceRoleService.getRoleScope(roleId);
// 3. 解析系统提示词角色人设
String systemPrompt = scope.hasRole() && StringUtils.hasText(scope.systemPrompt())
? scope.systemPrompt() : null;
// 4. 解析分类隔离范围角色有绑定则强制使用否则使用客户端传入
List<Long> catIds = resolveCategoryIds(scope, categoryIds);
// 5. 严格隔离角色无分类则拒绝检索 KB
boolean useRag = (enableRag == null || enableRag) // 默认启用 RAG
&& !(roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty());
// 6. 未指定策略时默认 MULTI_QUERY
String strategy = (rewriteStrategy != null && !rewriteStrategy.isBlank())
? rewriteStrategy : "MULTI_QUERY";
return new ChatContext(message, chatId, "CHAT", systemPrompt, scope.allowedMcpTools(),
catIds, strategy, useRag, streaming);
}
/**
* 解析 API Key 绑定的角色 ID 列表JSON 字符串如 "[1,2,3]" List<Long>
* null / 空数组表示不限制返回空列表
*/
private List<Long> parseAllowedRoleIds(String roleIdsJson) {
if (roleIdsJson == null || roleIdsJson.isBlank() || "[]".equals(roleIdsJson.trim())) {
return Collections.emptyList();
}
try {
return objectMapper.readValue(roleIdsJson, new TypeReference<List<Long>>() {});
} catch (Exception e) {
log.warn("解析 API Key role_ids 失败: {}", roleIdsJson, e);
return Collections.emptyList();
}
}
/**
* 从客户端传入的 roleId 参数解析有效角色 ID
* API Key 绑定了角色列表则仅允许使用绑定的角色未绑定时允许任意角色
* 传入 roleId 不在允许列表内时降级为不指定角色无角色人设/分类/MCP 权限
*/
private Long resolveRoleId(String roleIdStr, List<Long> allowedRoleIds) {
if (roleIdStr == null || roleIdStr.isBlank()) {
return null;
}
Long roleId;
try {
roleId = Long.valueOf(roleIdStr.trim());
} catch (NumberFormatException e) {
return null;
}
if (allowedRoleIds.isEmpty()) {
return roleId;
}
return allowedRoleIds.contains(roleId) ? roleId : null;
}
/**
* 解析分类隔离范围角色有绑定时强制使用角色的分类否则使用客户端传入的 categoryIds
*/
private List<Long> resolveCategoryIds(RoleScope scope, String categoryIds) {
if (scope.hasRole()) {
return scope.categoryIds();
}
return categoryFilter.parse(categoryIds);
}
}

108
src/main/java/com/wok/supportbot/rag/CategoryFilter.java

@ -0,0 +1,108 @@
package com.wok.supportbot.rag;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.ai.vectorstore.filter.FilterExpressionBuilder;
import org.springframework.stereotype.Component;
import java.util.Arrays;
import java.util.Collections;
import java.util.List;
import java.util.Objects;
/**
* 分类过滤统一工具
* <p>
* 收敛原本分散在 {@code AiController.parseCategoryIds(String)}
* {@code DocumentController.parseCategoryIds(Object)}
* {@code CustomerServiceRoleController.parseCategoryIds(Object)}
* {@code AssistantApp.normalizeCategoryIds / buildCategoryFilterExpression}
* {@code DocumentService.normalizeCategoryIds} 等多处的分类 ID 解析归一化与过滤表达式构建逻辑
* <p>
* 统一规则
* <ul>
* <li>过滤 null &lt;= 0 的非法值</li>
* <li>去重</li>
* <li>空集合统一返回 {@link Collections#emptyList()}过滤表达式返回 {@code null}表示不加分类约束</li>
* </ul>
*/
@Component
public class CategoryFilter {
/** 向量库 metadata 中存储分类 ID 的字段名 */
private static final String CATEGORY_ID_KEY = "categoryId";
/**
* 解析逗号分隔的分类 ID 字符串 "1,2,3"
*
* @param categoryIds 逗号分隔字符串可为 null/空白
* @return 去重且 &gt; 0 Long 列表永不返回 null
*/
public List<Long> parse(String categoryIds) {
if (categoryIds == null || categoryIds.isBlank()) {
return Collections.emptyList();
}
return Arrays.stream(categoryIds.split(","))
.map(String::trim)
.filter(item -> !item.isEmpty())
.map(Long::valueOf)
.filter(id -> id > 0)
.distinct()
.toList();
}
/**
* 解析前端传入的原始分类 ID 对象兼容 List数字或字符串元素与逗号分隔字符串两种形式
*
* @param rawCategoryIds List 或字符串可为 null
* @return 去重且 &gt; 0 Long 列表永不返回 null
*/
public List<Long> parse(Object rawCategoryIds) {
if (rawCategoryIds == null) {
return Collections.emptyList();
}
if (rawCategoryIds instanceof List<?> list) {
return list.stream()
.filter(Objects::nonNull)
.map(item -> item instanceof Number number ? number.longValue() : Long.parseLong(item.toString()))
.filter(id -> id > 0)
.distinct()
.toList();
}
return parse(rawCategoryIds.toString());
}
/**
* Long 分类 ID 列表归一化为字符串列表 SQL 片段或日志使用
*
* @param categoryIds Long 列表可为 null
* @return 去重且 &gt; 0 的字符串列表永不返回 null
*/
public List<String> normalize(List<Long> categoryIds) {
if (categoryIds == null || categoryIds.isEmpty()) {
return Collections.emptyList();
}
return categoryIds.stream()
.filter(Objects::nonNull)
.filter(id -> id > 0)
.map(String::valueOf)
.distinct()
.toList();
}
/**
* 构建向量库分类过滤表达式{@code categoryId IN (...)}
*
* @param categoryIds Long 列表可为 null
* @return 过滤表达式分类为空时返回 {@code null}表示不加分类约束
*/
public Filter.Expression buildExpression(List<Long> categoryIds) {
List<Object> values = normalize(categoryIds).stream()
.map(value -> (Object) value)
.toList();
if (values.isEmpty()) {
return null;
}
FilterExpressionBuilder filterBuilder = new FilterExpressionBuilder();
return filterBuilder.in(CATEGORY_ID_KEY, values).build();
}
}

34
src/main/java/com/wok/supportbot/rag/RagContext.java

@ -0,0 +1,34 @@
package com.wok.supportbot.rag;
import org.springframework.ai.document.Document;
import java.util.List;
import java.util.Optional;
/**
* RAG 检索结果值对象
* <p>
* {@link RagPipeline#retrieve} 产出 {@code ChatPipeline} 决定后续动作
* <ul>
* <li>{@link #faqAnswer()} 命中 直接返回标准答案跳过生成</li>
* <li>未命中 {@link #documents()} 展示引用来源 {@link #contextText()} 注入系统提示词
* {@link #rewrittenQuery()} 作为传给模型的用户消息</li>
* </ul>
*
* @param faqAnswer FAQ 命中的标准答案未命中为 {@link Optional#empty()}
* @param documents 检索命中的知识库片段 metadata可为空
* @param contextText 拼接后的资料文本 {@code "\n\n---\n\n"} 分隔无资料时为空串
* @param rewrittenQuery 传给模型的用户消息MULTI_QUERY 为原始 message其余策略为重写后的查询
*/
public record RagContext(
Optional<String> faqAnswer,
List<Document> documents,
String contextText,
String rewrittenQuery
) {
/** FAQ 是否命中 */
public boolean faqHit() {
return faqAnswer != null && faqAnswer.isPresent();
}
}

271
src/main/java/com/wok/supportbot/rag/RagPipeline.java

@ -0,0 +1,271 @@
package com.wok.supportbot.rag;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.chatmemory.DatabaseChatMemory;
import com.wok.supportbot.config.RagPromptConfig;
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 jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.document.Document;
import org.springframework.ai.vectorstore.SearchRequest;
import org.springframework.ai.vectorstore.VectorStore;
import org.springframework.ai.vectorstore.filter.Filter;
import org.springframework.stereotype.Component;
import org.springframework.util.StringUtils;
import java.util.ArrayList;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.stream.Collectors;
/**
* 统一 RAG 检索管道
* <p>
* 收敛原本分散在 {@code AssistantApp} 中按策略分支的检索逻辑
* <ul>
* <li>FAQ 优先匹配复用 {@link FaqMatchEngine} 三级匹配</li>
* <li>查询重写 {@code rewriteStrategy} 复用 {@code rag/preretrieval/*} 四种 rewriter</li>
* <li>统一检索{@code MULTI_QUERY} 扩展多查询后按文档 ID 去重合并其余策略单查询检索</li>
* <li>统一资料块模板 {@link #buildRagContextBlock}替代原 {@code buildRetrievalAdvisor.qaTemplate}
* {@code buildRagSystemPrompt} 两份回答模板</li>
* </ul>
* <p>
* 统一为"手动检索 + 资料块注入系统提示词"模式所有策略都先由本管道检索出文档
* 再由调用方把 {@link RagContext#contextText()} 组装进系统提示词
* 不再使用 {@code RetrievalAugmentationAdvisor} query augmenter 自动注入
* 消除上下文注入位置随策略不同而不同的不一致
* <p>
* 阶段一作为旁路组件存在 {@code AssistantApp} RAG 路径未改动阶段二由 {@code ChatPipeline} 接入
*/
@Component
@Slf4j
public class RagPipeline {
/** 单路检索 topK,与原 AssistantApp 保持一致 */
private static final int TOP_K = 4;
/** 多路检索合并后的文档封顶数,避免上下文膨胀 */
private static final int MAX_DOCS = 8;
/** COMPRESSION 策略读取的对话历史条数 */
private static final int COMPRESSION_HISTORY_SIZE = 10;
@Resource
private VectorStore pgVectorVectorStore;
@Resource
private FaqMatchEngine faqMatchEngine;
@Resource
private RagPromptConfig ragPromptConfig;
@Resource
private CategoryFilter categoryFilter;
@Resource
private RewriteQueryRewriter rewriteQueryRewriter;
@Resource
private TranslationQueryRewriter translationQueryRewriter;
@Resource
private CompressionQueryRewriter compressionQueryRewriter;
@Resource
private MultiQueryExpanderRewriter multiQueryExpanderRewriter;
private final DatabaseChatMemory chatMemory;
public RagPipeline(DatabaseChatMemory chatMemory) {
this.chatMemory = chatMemory;
}
/**
* 执行一次统一的 RAG 检索
* <p>
* 流程FAQ 优先 查询重写/扩展 统一检索 拼接资料文本
*
* @param ctx 对话上下文使用 {@code message / chatId / rewriteStrategy / categoryIds}
* @return 检索结果FAQ 命中时 documents contextText 为空rewrittenQuery 为原始 message
*/
public RagContext retrieve(ChatContext ctx) {
// 1. FAQ 优先匹配命中则直接返回标准答案跳过检索与生成
Optional<String> faqAnswer = tryFaqMatch(ctx.message());
if (faqAnswer.isPresent()) {
log.info("FAQ 命中,跳过知识库检索: chatId={}", ctx.chatId());
return new RagContext(faqAnswer, Collections.emptyList(), "", ctx.message());
}
// 2. 统一检索 + 组装结果
return retrieveDocumentsAndContext(ctx);
}
/**
* 仅检索知识库片段跳过 FAQ 匹配用于"引用来源"展示
* <p>
* {@link #retrieve} 共用同一套查询重写与检索逻辑确保来源即答案所依据的片段
* 但不触发 FAQ 优先匹配来源接口的语义是展示 KB 片段FAQ 命中时本就无 KB 来源
*
* @param ctx 对话上下文
* @return 命中的知识库片段 metadata无命中返回空列表
*/
public List<Document> retrieveDocuments(ChatContext ctx) {
return retrieveDocumentsAndContext(ctx).documents();
}
/**
* 统一检索 + 拼接资料文本不含 FAQ 匹配 {@link #retrieve} {@link #retrieveDocuments} 复用
*/
private RagContext retrieveDocumentsAndContext(ChatContext ctx) {
List<Document> docs;
String rewrittenQuery;
if ("MULTI_QUERY".equalsIgnoreCase(ctx.rewriteStrategy())) {
// MULTI_QUERY资料注入 systemuser 消息保持原始 message
docs = retrieveMultiQueryDocs(ctx.message(), ctx.categoryIds());
rewrittenQuery = ctx.message();
} else {
rewrittenQuery = rewriteQuery(ctx.message(), ctx.chatId(), ctx.rewriteStrategy());
docs = similaritySearch(rewrittenQuery, ctx.categoryIds());
}
String contextText = joinContext(docs);
return new RagContext(Optional.empty(), docs, contextText, rewrittenQuery);
}
/**
* 构造统一的 RAG 资料块文本供调用方追加到系统提示词
* <p>
* 替代原 {@code buildRetrievalAdvisor} qaTemplate {@code buildRagSystemPrompt} 两份模板
* 所有策略共用同一份"回答硬性规则 + 知识库资料"格式
*
* @param contextText 检索拼接的资料文本为空时返回空串调用方据此决定是否追加
*/
public String buildRagContextBlock(String contextText) {
if (!StringUtils.hasText(contextText)) {
return "";
}
return "\n\n【RAG回答硬性规则】\n" + ragPromptConfig.getAnswerRules()
+ "\n\n【知识库资料】\n" + contextText;
}
// ==================== FAQ 匹配 ====================
/**
* 尝试 FAQ 三级匹配精确关键词语义命中返回标准答案
* 异常时降级为未命中继续走 RAG 检索
*/
private Optional<String> tryFaqMatch(String message) {
try {
return faqMatchEngine.match(message)
.map(result -> result.getFaq().getAnswer());
} catch (Exception e) {
log.warn("FAQ 匹配异常,降级到 RAG: {}", e.getMessage());
return Optional.empty();
}
}
// ==================== 查询重写 ====================
/**
* 按策略改写查询MULTI_QUERY 不走此方法 {@link #retrieveMultiQueryDocs} 自行扩展
* 重写失败时降级为原始查询避免 RAG_REWRITE 配置异常导致整次对话不可用
*/
private String rewriteQuery(String message, String chatId, String strategy) {
if (strategy == null || strategy.isEmpty()) {
return message;
}
try {
switch (strategy.toUpperCase()) {
case "REWRITE":
return rewriteQueryRewriter.doQueryRewrite(message);
case "TRANSLATION":
return translationQueryRewriter.doQueryRewrite(message);
case "COMPRESSION":
// 查询压缩需要对话历史从会话记忆取最近若干条把指代不清的追问补全为独立查询
List<Message> history = chatMemory.get(chatId, COMPRESSION_HISTORY_SIZE);
return compressionQueryRewriter.doQueryRewrite(message, history);
case "NONE":
default:
return message;
}
} catch (Exception e) {
log.warn("查询重写失败 [strategy={}, chatId={}],降级使用原始查询: {}", strategy, chatId, e.getMessage());
return message;
}
}
// ==================== 多路检索 ====================
/**
* 多路查询扩展 + 分别检索 + 按文档 ID 去重合并
* 扩展失败时退回原问题检索避免整次 RAG 不可用
*/
private List<Document> retrieveMultiQueryDocs(String message, List<Long> categoryIds) {
List<String> expandedQueries;
try {
expandedQueries = multiQueryExpanderRewriter.doQueryRewrite(message);
} catch (Exception e) {
log.warn("多路查询扩展失败,降级为原始问题检索: {}", e.getMessage());
expandedQueries = List.of(message);
}
if (expandedQueries == null || expandedQueries.isEmpty()) {
expandedQueries = List.of(message);
}
log.info("多路查询扩展结果: {}", expandedQueries);
Map<String, Document> merged = new LinkedHashMap<>();
for (String query : expandedQueries) {
if (!StringUtils.hasText(query) || merged.size() >= MAX_DOCS) {
continue;
}
List<Document> docs = similaritySearch(query, categoryIds);
for (Document doc : docs) {
if (merged.size() >= MAX_DOCS) {
break;
}
merged.putIfAbsent(doc.getId(), doc);
}
}
return new ArrayList<>(merged.values());
}
// ==================== 底层检索 ====================
/**
* 单查询向量检索附带分类过滤
*/
private List<Document> similaritySearch(String query, List<Long> categoryIds) {
if (!StringUtils.hasText(query)) {
return Collections.emptyList();
}
SearchRequest.Builder builder = SearchRequest.builder()
.similarityThreshold(0.0)
.topK(TOP_K);
Filter.Expression filterExpression = categoryFilter.buildExpression(categoryIds);
if (filterExpression != null) {
builder.filterExpression(filterExpression);
}
List<Document> docs = pgVectorVectorStore.similaritySearch(builder.query(query).build());
return docs != null ? docs : Collections.emptyList();
}
/**
* 把文档列表拼接为资料文本过滤空内容用分隔符区分片段
*/
private String joinContext(List<Document> docs) {
if (docs == null || docs.isEmpty()) {
return "";
}
return docs.stream()
.map(Document::getText)
.filter(StringUtils::hasText)
.collect(Collectors.joining("\n\n---\n\n"));
}
}

13
src/main/java/com/wok/supportbot/service/DocumentService.java

@ -76,6 +76,9 @@ public class DocumentService {
@Autowired
private com.wok.supportbot.config.FileStorageConfig fileStorageConfig;
@Autowired
private com.wok.supportbot.rag.CategoryFilter categoryFilter;
// ==================== 文档上传 ====================
/**
@ -888,15 +891,7 @@ public class DocumentService {
private com.wok.supportbot.rag.HybridSearchService hybridSearchService;
private List<String> normalizeCategoryIds(List<Long> categoryIds) {
if (categoryIds == null || categoryIds.isEmpty()) {
return Collections.emptyList();
}
return categoryIds.stream()
.filter(Objects::nonNull)
.filter(id -> id > 0)
.map(String::valueOf)
.distinct()
.collect(Collectors.toList());
return categoryFilter.normalize(categoryIds);
}
// ==================== 统计 ====================

2
src/main/resources/application-dev.yml

@ -10,7 +10,7 @@ server:
spring:
datasource:
url: jdbc:postgresql://192.168.1.18:5432/support_bot
url: jdbc:postgresql://192.168.1.49:5432/support_bot
username: postgres
password: supportbot123
sql:

4
src/main/resources/static/components/ChatPanel.js

@ -85,8 +85,6 @@ export default {
<select class="select chat-mode" v-model="mode">
<option value="sync">&#x540C;&#x6B65;&#x8C03;&#x7528;</option>
<option value="sse">SSE &#x6D41;&#x5F0F;</option>
<option value="sse2">ServerSentEvent</option>
<option value="sse3">SseEmitter</option>
</select>
</div>
</header>
@ -361,7 +359,7 @@ export default {
} else if (mode.value === 'sync') {
assistantMsg.content = await chatSync(text, cid, roleId)
} else {
const url = chatSSEUrl(text, cid, mode.value, roleId)
const url = chatSSEUrl(text, cid, roleId)
await readSSEStream(url, async (chunk) => {
assistantMsg.content += chunk
await scrollToBottom()

10
src/main/resources/static/js/api.js

@ -147,17 +147,11 @@ export function chatRagSync(message, chatId, strategy, roleId, accountId) {
* 获取普通 SSE 流式对话 URL
* @param {string} message
* @param {string} chatId
* @param {'sse'|'sse2'|'sse3'} mode
* @param {string} [roleId]
* @returns {string}
*/
export function chatSSEUrl(message, chatId, mode, roleId, accountId) {
const pathMap = {
sse: '/ai/assistant_app/chat/sse',
sse2: '/ai/assistant_app/chat/server_sent_event',
sse3: '/ai/assistant_app/chat/sse_emitter'
}
let url = API_BASE + `${pathMap[mode]}?message=${encodeURIComponent(message)}&chatId=${encodeURIComponent(chatId)}`
export function chatSSEUrl(message, chatId, roleId, accountId) {
let url = API_BASE + `/ai/assistant_app/chat/sse?message=${encodeURIComponent(message)}&chatId=${encodeURIComponent(chatId)}`
if (roleId) url += `&roleId=${encodeURIComponent(roleId)}`
if (accountId) url += `&accountId=${encodeURIComponent(accountId)}`
return url

123
src/test/java/com/wok/supportbot/Phase1ComponentTests.java

@ -0,0 +1,123 @@
package com.wok.supportbot;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.rag.RagContext;
import com.wok.supportbot.rag.RagPipeline;
import jakarta.annotation.Resource;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.DisplayName;
import org.junit.jupiter.api.Test;
import org.springframework.boot.test.context.SpringBootTest;
import java.util.Collections;
import java.util.List;
/**
* 阶段一公共组件验证测试
* 需要运行中的 PostgreSQLRagPipeline 检索依赖向量库
*/
@SpringBootTest
class Phase1ComponentTests {
@Resource
private CategoryFilter categoryFilter;
@Resource
private RagPipeline ragPipeline;
// ==================== 1.1 CategoryFilter ====================
@Test @DisplayName("CF-01 parse(String) 正常逗号分隔")
void parseStringNormal() {
Assertions.assertEquals(List.of(1L, 2L, 3L), categoryFilter.parse("1,2,3"));
}
@Test @DisplayName("CF-02 parse(String) 含空格")
void parseStringWithSpaces() {
Assertions.assertEquals(List.of(1L, 2L, 3L), categoryFilter.parse(" 1 , 2 , 3 "));
}
@Test @DisplayName("CF-03 parse(String) null/空白")
void parseStringNullOrBlank() {
Assertions.assertEquals(Collections.emptyList(), categoryFilter.parse((String) null));
Assertions.assertEquals(Collections.emptyList(), categoryFilter.parse(""));
Assertions.assertEquals(Collections.emptyList(), categoryFilter.parse(" "));
}
@Test @DisplayName("CF-04/05 parse(Object) List")
void parseObjectList() {
// 数字元素
Assertions.assertEquals(List.of(1L, 2L, 3L), categoryFilter.parse((Object) List.of(1, 2, 3)));
// 字符串元素
Assertions.assertEquals(List.of(1L, 2L, 3L), categoryFilter.parse((Object) List.of("1", "2", "3")));
}
@Test @DisplayName("CF-06/07 parse(Object) 字符串 / null")
void parseObjectStringOrNull() {
Assertions.assertEquals(List.of(4L, 5L, 6L), categoryFilter.parse((Object) "4,5,6"));
Assertions.assertEquals(Collections.emptyList(), categoryFilter.parse((Object) null));
}
@Test @DisplayName("CF-08/09 normalize")
void normalize() {
Assertions.assertEquals(List.of("1", "2", "3"), categoryFilter.normalize(List.of(1L, 2L, 3L)));
Assertions.assertEquals(Collections.emptyList(), categoryFilter.normalize(null));
Assertions.assertEquals(Collections.emptyList(), categoryFilter.normalize(Collections.emptyList()));
}
@Test @DisplayName("CF-10/11 buildExpression")
void buildExpression() {
Assertions.assertNotNull(categoryFilter.buildExpression(List.of(1L, 2L)));
Assertions.assertNull(categoryFilter.buildExpression(Collections.emptyList()));
}
// ==================== 1.2 RagPipeline ====================
@Test @DisplayName("RP-08 未指定策略 → 原样查询")
void ragDefaultStrategy() {
ChatContext ctx = new ChatContext("退换货流程是什么", "test-rp-01", "CHAT",
null, null, null, null, true, false);
RagContext result = ragPipeline.retrieve(ctx);
Assertions.assertFalse(result.faqHit(), "不应触发 FAQ(除非该问题恰好命中 FAQ 库)");
Assertions.assertNotNull(result.documents(), "documents 不应为 null");
Assertions.assertNotNull(result.rewrittenQuery(), "rewrittenQuery 不应为 null");
}
@Test @DisplayName("RP-11/12 分类过滤")
void ragCategoryFilter() {
// 有分类过滤
ChatContext ctx1 = new ChatContext("平台使用说明", "test-rp-11", "CHAT",
null, null, List.of(1L), null, true, false);
RagContext r1 = ragPipeline.retrieve(ctx1);
Assertions.assertNotNull(r1.documents());
System.out.println(">>> 分类过滤(1) 命中数: " + r1.documents().size());
// 无分类过滤
ChatContext ctx2 = new ChatContext("平台使用说明", "test-rp-12", "CHAT",
null, null, null, null, true, false);
RagContext r2 = ragPipeline.retrieve(ctx2);
Assertions.assertNotNull(r2.documents());
System.out.println(">>> 无分类过滤 命中数: " + r2.documents().size());
}
@Test @DisplayName("RP-14 retrieveDocuments 跳过 FAQ")
void ragRetrieveDocuments() {
ChatContext ctx = new ChatContext("平台使用说明", "test-rp-14", "CHAT",
null, null, null, null, true, false);
var docs = ragPipeline.retrieveDocuments(ctx);
Assertions.assertNotNull(docs);
System.out.println(">>> retrieveDocuments 命中数: " + docs.size());
}
@Test @DisplayName("RP-04 REWRITE 策略")
void ragRewriteStrategy() {
ChatContext ctx = new ChatContext("订单退款流程", "test-rp-04", "CHAT",
null, null, null, "REWRITE", true, false);
RagContext result = ragPipeline.retrieve(ctx);
Assertions.assertNotNull(result.rewrittenQuery());
System.out.println(">>> REWRITE 重写前: 订单退款流程");
System.out.println(">>> REWRITE 重写后: " + result.rewrittenQuery());
Assertions.assertNotNull(result.documents());
}
}

13
src/test/java/com/wok/supportbot/SupportBotApplicationTests.java

@ -1,6 +1,7 @@
package com.wok.supportbot;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import jakarta.annotation.Resource;
import org.junit.jupiter.api.Assertions;
import org.junit.jupiter.api.Test;
@ -18,17 +19,17 @@ class SupportBotApplicationTests {
String chatId = UUID.randomUUID().toString();
// 第一轮商品咨询
String message = "你好,我想买一台适合学生用的笔记本电脑,有推荐吗?";
String answer = assistantApp.doChat(message, chatId);
String answer = assistantApp.chat(ChatContext.of(message, chatId));
Assertions.assertNotNull(answer);
// 第二轮物流问题
message = "我上周买的那台电脑现在还没到,能查一下物流吗?";
answer = assistantApp.doChat(message, chatId);
answer = assistantApp.chat(ChatContext.of(message, chatId));
Assertions.assertNotNull(answer);
// 第三轮售后问题
message = "电脑到了,但有点问题。你刚刚说的售后流程能再说一遍吗?";
answer = assistantApp.doChat(message, chatId);
answer = assistantApp.chat(ChatContext.of(message, chatId));
Assertions.assertNotNull(answer);
}
@ -36,7 +37,8 @@ class SupportBotApplicationTests {
void doChatWithRag() {
String chatId = "1069b88d-eb85-47ac-bd2e-c393d118a5aa";
String message = "我之前询问了你什么问题?";
String answer = assistantApp.doChatWithRag(message, chatId);
String answer = assistantApp.chat(new ChatContext(message, chatId, "CHAT", null, null,
null, null, true, false));
Assertions.assertNotNull(answer);
}
@ -44,7 +46,8 @@ class SupportBotApplicationTests {
void doChatWithRagEnhance() {
String chatId = "1069b88d-eb85-47ac-bd2e-c393d118a5aa";
String message = "我之前询问了你什么问题?";
String answer = assistantApp.doChatWithRagEnhance(message, chatId);
String answer = assistantApp.chat(new ChatContext(message, chatId, "CHAT", null, null,
null, null, true, false));
Assertions.assertNotNull(answer);
}

Loading…
Cancel
Save