Browse Source

refactor(chat): 移除 LLM 意图分类并默认原文检索

feature/test
wei-py 3 weeks ago
parent
commit
3f93d015f8
  1. 7
      src/main/java/com/wok/supportbot/app/ChatContext.java
  2. 87
      src/main/java/com/wok/supportbot/app/ChatPipeline.java
  3. 18
      src/main/java/com/wok/supportbot/rag/RagPipeline.java
  4. 134
      src/main/java/com/wok/supportbot/service/IntentRouter.java
  5. 317
      src/test/java/com/wok/supportbot/ChatPipelineTests.java
  6. 126
      src/test/java/com/wok/supportbot/IntentRouterTests.java

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

@ -17,7 +17,7 @@ import java.util.List;
* @param systemPrompt 角色人设/系统提示词,可为 null
* @param allowedMcpTools 允许的 MCP 工具名列表;null=无角色允许全部;空=有角色但无授权(禁止);["*"]=全部
* @param categoryIds 知识库分类隔离范围,可为空(不限制)
* @param rewriteStrategy RAG 查询重写策略(REWRITE / TRANSLATION / COMPRESSION / MULTI_QUERY),可为 null
* @param rewriteStrategy RAG 查询重写策略,默认 NONE 原文检索;可显式选择 REWRITE / TRANSLATION / COMPRESSION / MULTI_QUERY
* @param enableRag 是否启用 RAG 检索;false=普通对话
* @param streaming 是否流式输出
* @param roleId 客服角色 ID(可空,供调用追踪等使用)
@ -47,12 +47,15 @@ public record ChatContext(
private static final String DEFAULT_APP_TYPE = "CHAT";
/**
* 紧凑构造器:规范化默认值,保证 appType 非空。
* 紧凑构造器:规范化应用类型与检索默认值。
*/
public ChatContext {
if (appType == null || appType.isBlank()) {
appType = DEFAULT_APP_TYPE;
}
if (rewriteStrategy == null || rewriteStrategy.isBlank()) {
rewriteStrategy = "NONE";
}
}
/**

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

@ -3,9 +3,7 @@ package com.wok.supportbot.app;
import com.wok.supportbot.rag.RagContext;
import com.wok.supportbot.rag.RagPipeline;
import com.wok.supportbot.service.FaqMatchEngine.FaqMatchResult;
import com.wok.supportbot.service.IntentRouter;
import com.wok.supportbot.service.SystemConfigService;
import com.wok.supportbot.service.RagHitLogService;
import jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.document.Document;
@ -20,19 +18,17 @@ import java.util.Optional;
/**
* 统一对话管道(编排层)。
* <p>
* 编排一次完整对话的决策流程:FAQ 优先 → 意图路由 → RAG 检索 → 组装系统提示词与用户消息,
* 编排一次完整对话的决策流程:FAQ 优先 → 本地寒暄判断 → RAG 检索 → 组装系统提示词与用户消息,
* 产出 {@link ChatRequest} 交由 {@code AssistantApp} 执行实际的 {@code call()} / {@code stream()}。
* <p>
* 设计说明:本类为纯编排层,不持有 ChatClient(ChatClient 构建与 Advisor 链装配仍在
* {@code AssistantApp}),因此 {@code call} / {@code stream} 由 {@code AssistantApp} 承担,
* 避免 {@code ChatPipeline} ↔ {@code AssistantApp} 循环依赖。
* <p>
* 接入 {@link IntentRouter} 替代原 {@code AiController.shouldBypassKnowledgeRetrieval} 的硬编码寒暄词判断:
* 寒暄词列表保留为快速路径与兜底,IntentRouter 负责细粒度意图分类,二者命中其一即跳过 KB 检索。
* 完整 FAQ 匹配后仅使用本地寒暄判断,其他请求直接检索,不调用 LLM 做意图分类。
* <p>
* {@code @pipeline} orchestration-layer order=0<br>
* {@code @pipeline-step} buildRequest: FAQ优先 → 意图路由 → RAG检索 → 提示词组装<br>
* {@code @pipeline-step} routeIntent: 寒暄词快速路径 → IntentRouter LLM分类 → 降级RAG<br>
* {@code @pipeline-step} buildRequest: FAQ优先 → 本地寒暄判断 → 默认原文检索 → 提示词组装<br>
* {@code @pipeline-step} effectiveSystem: DB全局提示词 + 角色人设 动态组合<br>
* 同步至: frontend/src/views/PipelineFlow.vue, CLAUDE.md ASCII管道图
*/
@ -40,29 +36,20 @@ import java.util.Optional;
@Slf4j
public class ChatPipeline {
/** IntentRouter 判为 CHITCHAT 的置信度阈值,低于此值视为不确定,继续走 RAG */
private static final double CHITCHAT_CONFIDENCE_THRESHOLD = 0.6;
@Resource
private IntentRouter intentRouter;
@Resource
private RagPipeline ragPipeline;
@Resource
private SystemConfigService systemConfigService;
@Resource
private RagHitLogService ragHitLogService;
/**
* 编排一次对话请求,产出执行决策。
* <p>
* 决策分支:
* <ul>
* <li>未启用 RAG(普通对话 / 严格隔离下 KB 拒绝)→ 用原始 message、基础 system</li>
* <li>FAQ 命中 → 直接返回标准答案,不调用 IntentRouter 或 ChatClient</li>
* <li>FAQ 未命中的寒暄/闲聊(IntentRouter 或寒暄词命中)→ 跳过 KB 检索</li>
* <li>FAQ 命中 → 直接返回标准答案,不调用 ChatClient</li>
* <li>FAQ 未命中的本地寒暄词 → 跳过 KB 检索</li>
* <li>RAG 生成 → 资料块注入 system,原始 message 作为 user 消息(重写查询仅用于检索)</li>
* </ul>
*
@ -79,22 +66,18 @@ public class ChatPipeline {
globalPrompt, null, null, "CHAT", null, null, null);
}
// 完整 FAQ 三级匹配前置:标准答案命中时省去 LLM 意图分类,并保留角色分类隔离。
// 完整 FAQ 三级匹配前置,保留角色分类隔离。
RagPipeline.FaqMatchOutcome faqOutcome = ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds());
if (faqOutcome.result().isPresent()) {
FaqMatchResult faqMatch = faqOutcome.result().get();
log.info("FAQ 命中标准答案,跳过意图分类: chatId={}, matchType={}", ctx.chatId(), faqMatch.getMatchType());
log.info("FAQ 命中标准答案: chatId={}, matchType={}", ctx.chatId(), faqMatch.getMatchType());
return new ChatRequest(ctx, ctx.message(), baseSystem,
Optional.ofNullable(faqMatch.getFaq().getAnswer()),
globalPrompt, null, null, "FAQ", null, null, faqMatch);
}
// FAQ 未命中后再路由:寒暄词快速路径 → IntentRouter LLM 分类。
IntentRouter.IntentResult intent = routeIntent(ctx.message());
// 寒暄/闲聊:IntentRouter 判定 CHITCHAT 高置信,跳过 KB 检索
if (intent != null && "CHITCHAT".equals(intent.getIntent())
&& intent.getConfidence() >= CHITCHAT_CONFIDENCE_THRESHOLD) {
// 仅本地寒暄词绕过检索;业务问题默认不做任何 LLM 预处理。
if (isChitchat(ctx.message())) {
return new ChatRequest(ctx, ctx.message(), baseSystem, Optional.empty(),
globalPrompt, null, null, "CHITCHAT", null, null, null);
}
@ -102,21 +85,6 @@ public class ChatPipeline {
// 仅干净完成的 FAQ 匹配可跳过;异常降级的 miss 仍由 RAG 重试,避免误跳。
RagContext rag = ragPipeline.retrieve(ctx, faqOutcome.completedCleanly());
// 记录 RAG 检索日志到 rag_hit_log 表(供知识库分析看板使用)
if (!rag.faqHit() && rag.documents() != null && !rag.documents().isEmpty()) {
String searchMode = ctx.rewriteStrategy() != null ? ctx.rewriteStrategy() : "VECTOR";
for (Document doc : rag.documents()) {
String docIdStr = String.valueOf(doc.getMetadata().getOrDefault("documentId", ""));
Long documentId = null;
try { if (!docIdStr.isEmpty()) documentId = Long.parseLong(docIdStr); } catch (NumberFormatException ignored) { }
String title = String.valueOf(doc.getMetadata().getOrDefault("title", ""));
String score = String.valueOf(doc.getMetadata().getOrDefault("score", ""));
ragHitLogService.recordHit(ctx.chatId(), ctx.message(), documentId, title, score, searchMode);
}
} else if (!rag.faqHit()) {
String searchMode = ctx.rewriteStrategy() != null ? ctx.rewriteStrategy() : "VECTOR";
ragHitLogService.recordMiss(ctx.chatId(), ctx.message(), searchMode);
}
if (rag.faqHit()) {
return new ChatRequest(ctx, ctx.message(), baseSystem, rag.faqAnswer(),
globalPrompt, null, null, "FAQ", null, null, rag.faqMatchResult());
@ -129,44 +97,9 @@ public class ChatPipeline {
rag.searchMode(), rag.documents(), null);
}
/**
* 意图路由:先用寒暄词列表做快速路径,未命中再调 IntentRouter 做 LLM 分类。
* 异常时返回 null(调用方默认走 RAG 检索)。
*/
private IntentRouter.IntentResult routeIntent(String message) {
// 寒暄词快速路径(零 LLM 开销)
if (isChitchat(message)) {
return new IntentRouter.IntentResult("CHITCHAT", 1.0);
}
try {
return intentRouter.route(message);
} catch (Exception e) {
log.debug("意图路由异常,沿用 RAG 检索: {}", e.getMessage());
return null;
}
}
/**
* 判断是否跳过 KB 检索(保留旧方法签名供 retrieveSources 等使用)。
*/
private boolean shouldBypassRag(String message) {
if (isChitchat(message)) {
return true;
}
try {
IntentRouter.IntentResult intent = intentRouter.route(message);
return "CHITCHAT".equals(intent.getIntent()) && intent.getConfidence() >= CHITCHAT_CONFIDENCE_THRESHOLD;
} catch (Exception e) {
log.debug("意图路由异常,沿用 RAG: {}", e.getMessage());
}
return false;
}
/**
* 寒暄词快速判断:问候/感谢/告别等短消息无知识库检索意图。
* 与原 {@code AiController.shouldBypassKnowledgeRetrieval} 逻辑一致,作为 IntentRouter 的快速路径与兜底。
* <p>
* 供 {@code buildRequest} 与"引用来源"等不需 LLM 意图分类的场景共用。
* 供 {@code buildRequest} 与独立引用检索接口共用,不调用 LLM。
*/
public boolean isChitchat(String message) {
if (!StringUtils.hasText(message)) {

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

@ -37,7 +37,7 @@ import java.util.stream.Collectors;
* 收敛原本分散在 {@code AssistantApp} 中按策略分支的检索逻辑:
* <ul>
* <li>FAQ 优先匹配(复用 {@link FaqMatchEngine} 三级匹配)</li>
* <li>查询重写(按 {@code rewriteStrategy} 复用 {@code rag/preretrieval/*} 四种 rewriter)</li>
* <li>默认 NONE 原文直检索;仅显式选择时调用 {@code rag/preretrieval/*} 四种 rewriter</li>
* <li>统一检索:{@code MULTI_QUERY} 扩展多查询后按文档 ID 去重合并,其余策略单查询检索</li>
* <li>统一资料块模板 {@link #buildRagContextBlock},替代原 {@code buildRetrievalAdvisor.qaTemplate}
* 与 {@code buildRagSystemPrompt} 两份回答模板</li>
@ -48,10 +48,10 @@ import java.util.stream.Collectors;
* 不再使用 {@code RetrievalAugmentationAdvisor} 的 query augmenter 自动注入,
* 消除上下文注入位置随策略不同而不同的不一致。
* <p>
* 阶段一作为旁路组件存在,旧 {@code AssistantApp} RAG 路径未改动;阶段二由 {@code ChatPipeline} 接入。
* 由 {@code ChatPipeline} 编排调用;RAG 命中/未命中日志仅在本管道记录一次。
* <p>
* {@code @pipeline} rag-layer order=1<br>
* {@code @pipeline-step} retrieve: FAQ优先匹配 → 查询重写/扩展 → similaritySearch(PGVector) → 资料拼接<br>
* {@code @pipeline-step} retrieve: FAQ优先匹配 → 默认原文/显式重写 → similaritySearch(PGVector) → 资料拼接<br>
* {@code @pipeline-step} similaritySearch: 纯向量检索 topK=4 + CategoryFilter 分类过滤<br>
* 注意: HybridSearchService/RrfFusion/RerankerService 尚未接入本管道,当前仅单路向量检索。<br>
* 同步至: frontend/src/views/PipelineFlow.vue RAG 子图
@ -113,7 +113,7 @@ public class RagPipeline {
/**
* 执行一次统一的 RAG 检索。
* <p>
* 流程:FAQ 优先 → 查询重写/扩展 → 统一检索 → 拼接资料文本。
* 流程:FAQ 优先 → 默认原文(显式选择才重写/扩展)→ 统一检索 → 拼接资料文本。
*
* @param ctx 对话上下文(使用 {@code message / chatId / rewriteStrategy / categoryIds})
* @return 检索结果;FAQ 命中时 documents 与 contextText 为空,rewrittenQuery 为原始 message
@ -125,10 +125,10 @@ public class RagPipeline {
/**
* 执行一次统一的 RAG 检索(可跳过前序已做过的 FAQ 匹配)。
* <p>
* 流程:FAQ 优先(未匹配过时)→ 查询重写/扩展 → 统一检索 → 拼接资料文本。
* 流程:FAQ 优先(未匹配过时)→ 默认原文(显式选择才重写/扩展)→ 统一检索 → 拼接资料文本。
*
* @param ctx 对话上下文(使用 {@code message / chatId / rewriteStrategy / categoryIds})
* @param faqAlreadyMatched 编排层是否已在前序阶段(意图路由之前)干净跑过完整 FAQ 三级匹配;
* @param faqAlreadyMatched 编排层是否已在前序阶段干净跑过完整 FAQ 三级匹配;
* true 时跳过 retrieve 内重复的 FAQ 匹配,避免同一请求重复做 FAQ 语义 embedding
* @return 检索结果;FAQ 命中时 documents 与 contextText 为空,rewrittenQuery 为原始 message
*/
@ -149,10 +149,10 @@ public class RagPipeline {
}
/**
* 仅检索知识库片段(跳过 FAQ 匹配),用于"引用来源"展示。
* 仅检索知识库片段(跳过 FAQ 匹配),用于已发布的独立来源检索接口。
* <p>
* 与 {@link #retrieve} 共用同一套查询重写与检索逻辑,确保来源即答案所依据的片段,
* 但不触发 FAQ 优先匹配——来源接口的语义是展示 KB 片段,FAQ 命中时本就无 KB 来源。
* 与 {@link #retrieve} 共用查询重写与检索逻辑;对话引用直接复用当次生成使用的文档,
* 不调用此方法二次检索。
*
* @param ctx 对话上下文
* @return 命中的知识库片段(含 metadata),无命中返回空列表

134
src/main/java/com/wok/supportbot/service/IntentRouter.java

@ -1,134 +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.ai.openai.OpenAiChatOptions;
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);
/** 分类只需返回 intent/confidence;不改变主回答模型的深度思考与输出长度。 */
private static final int CLASSIFICATION_MAX_TOKENS = 128;
/**
* 意图分类 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(classificationPrompt(chatModel, 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 Prompt classificationPrompt(ChatModel chatModel, String promptText) {
// Seed 2.0 默认开启深度思考;官方将 minimal 映射为关闭思考。
// 仅为已支持的模型系列覆盖本次分类请求,其他提供商不发送此参数。
// https://www.volcengine.com/docs/82379/1449737
if (chatModel.getDefaultOptions() instanceof OpenAiChatOptions defaults
&& defaults.getModel() != null
&& defaults.getModel().startsWith("doubao-seed-2-0-")) {
return new Prompt(promptText, OpenAiChatOptions.builder()
.reasoningEffort("minimal")
.maxTokens(CLASSIFICATION_MAX_TOKENS)
.build());
}
return new Prompt(promptText);
}
/**
* 校验意图类型是否有效(模型可能返回枚举外的值)
*/
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;
}
}

317
src/test/java/com/wok/supportbot/ChatPipelineTests.java

@ -3,26 +3,37 @@ 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.RagContext;
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.IntentRouter;
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.CsvSource;
import org.junit.jupiter.params.provider.NullAndEmptySource;
import org.junit.jupiter.params.provider.ValueSource;
import org.mockito.InjectMocks;
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.*;
@ -30,180 +41,254 @@ import static org.mockito.Mockito.*;
@ExtendWith(MockitoExtension.class)
class ChatPipelineTests {
@Mock
private IntentRouter intentRouter;
@Mock
private RagPipeline ragPipeline;
@Mock
private SystemConfigService systemConfigService;
@Mock
private RagHitLogService ragHitLogService;
@InjectMocks
@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 configureGlobalPrompt() {
when(systemConfigService.getValueByKey("ai_system_prompt")).thenReturn("全局提示词");
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
@ValueSource(strings = {"EXACT", "KEYWORD", "SEMANTIC"})
void faqHitSkipsIntentRoutingAndRetrieval(String matchType) {
ChatContext ctx = context("退货流程是什么", true);
FaqMatchResult match = faqMatch("标准答案", matchType);
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.of(match), true));
@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("FAQ", request.intent());
assertEquals(Optional.of("标准答案"), request.faqAnswer());
assertSame(match, request.faqMatchResult());
assertEquals("NONE", ctx.rewriteStrategy());
assertEquals("RAG", request.intent());
assertSame(ctx, request.ctx());
assertEquals(ctx.message(), request.finalMessage());
assertEquals("全局提示词\n\n【当前角色设定】\n售后客服", request.finalSystemPrompt());
verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
verifyNoMoreInteractions(ragPipeline);
verifyNoInteractions(intentRouter, ragHitLogService);
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 ordinaryChatDoesNotMatchFaqOrRouteIntent() {
ChatContext ctx = context("退货流程是什么", false);
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("CHAT", request.intent());
assertFalse(request.faqHit());
assertSame(ctx, request.ctx());
assertEquals(ctx.message(), request.finalMessage());
assertNull(request.ragContextText());
verifyNoInteractions(ragPipeline, intentRouter, ragHitLogService);
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
@CsvSource({"RAG, 0.9", "FAQ, 0.95", "FAQ, 0.5", "CHITCHAT, 0.59"})
void cleanFaqMissRoutesThenRetrievesWithoutRepeatingFaq(String intent, double confidence) {
ChatContext ctx = context("退货流程是什么", true);
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.empty(), true));
when(intentRouter.route(ctx.message())).thenReturn(new IntentRouter.IntentResult(intent, confidence));
stubRagRetrieval(ctx, true);
@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("RAG", request.intent());
assertEquals("FAQ", request.intent());
assertEquals(Optional.of("标准答案"), request.faqAnswer());
assertSame(match, request.faqMatchResult());
assertSame(ctx, request.ctx());
assertEquals(ctx.message(), request.finalMessage());
assertEquals("退货说明", request.ragContextText());
assertTrue(request.finalSystemPrompt().endsWith("\n资料:退货说明"));
assertEquals(1, request.hitCount());
var order = inOrder(ragPipeline, intentRouter);
order.verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
order.verify(intentRouter).route(ctx.message());
order.verify(ragPipeline).retrieve(ctx, true);
order.verify(ragPipeline).buildRagContextBlock("退货说明");
verifyNoMoreInteractions(ragPipeline, intentRouter);
assertEquals("全局提示词\n\n【当前角色设定】\n售后客服", request.finalSystemPrompt());
verify(faqMatchEngine).match(ctx.message(), ctx.categoryIds());
verifyNoMoreInteractions(faqMatchEngine);
verifyNoInteractions(vectorStore, ragHitLogService);
verifyNoPreprocessing();
}
@ParameterizedTest
@ValueSource(booleans = {true, false})
void faqMissKeepsHighConfidenceChitchatPath(boolean completedCleanly) {
ChatContext ctx = context("和我聊聊天吧", true);
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.empty(), completedCleanly));
when(intentRouter.route(ctx.message())).thenReturn(new IntentRouter.IntentResult("CHITCHAT", 0.6));
@ValueSource(strings = {"你好", " HI! ", "谢谢", "再见"})
void localGreetingMissSkipsRetrieval(String message) {
ChatContext ctx = context(message, true);
ChatRequest request = pipeline.buildRequest(ctx);
assertEquals("CHITCHAT", request.intent());
assertFalse(request.faqHit());
assertNull(request.ragContextText());
var order = inOrder(ragPipeline, intentRouter);
order.verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
order.verify(intentRouter).route(ctx.message());
verifyNoMoreInteractions(ragPipeline, intentRouter);
verifyNoInteractions(ragHitLogService);
assertEquals(message, request.finalMessage());
verify(faqMatchEngine).match(message, ctx.categoryIds());
verifyNoMoreInteractions(faqMatchEngine);
verifyNoInteractions(vectorStore, ragHitLogService);
verifyNoPreprocessing();
}
@Test
void localGreetingStillMatchesFaqOnceWithoutCallingIntentLlm() {
ChatContext ctx = context("你好", true);
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.empty(), true));
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();
}
assertEquals("CHITCHAT", request.intent());
verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
verifyNoMoreInteractions(ragPipeline);
verifyNoInteractions(intentRouter, ragHitLogService);
@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 exceptionalFaqMissAllowsRagRetryToReturnFaq() {
void repeatedFaqFailureStillRetrievesAndLogsOneMiss() {
ChatContext ctx = context("退货流程是什么", true);
FaqMatchResult retryMatch = faqMatch("重试命中的标准答案", "SEMANTIC");
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.empty(), false));
when(intentRouter.route(ctx.message())).thenReturn(new IntentRouter.IntentResult("FAQ", 0.95));
when(ragPipeline.retrieve(ctx, false)).thenReturn(new RagContext(
Optional.of("重试命中的标准答案"), List.of(), "", ctx.message(), "VECTOR", retryMatch));
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("FAQ", request.intent());
assertEquals(Optional.of("重试命中的标准答案"), request.faqAnswer());
assertSame(retryMatch, request.faqMatchResult());
verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
verify(ragPipeline).retrieve(ctx, false);
verifyNoMoreInteractions(ragPipeline);
verifyNoInteractions(ragHitLogService);
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);
}
}
@ParameterizedTest
@ValueSource(booleans = {true, false})
void intentFailureFallsBackToRagWithFaqCompletionFlag(boolean completedCleanly) {
@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);
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.empty(), completedCleanly));
when(intentRouter.route(ctx.message())).thenThrow(new IllegalStateException("意图服务不可用"));
stubRagRetrieval(ctx, completedCleanly);
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("RAG", request.intent());
verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
verify(ragPipeline).retrieve(ctx, completedCleanly);
verify(ragPipeline).buildRagContextBlock("退货说明");
verifyNoMoreInteractions(ragPipeline);
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 faqWithoutAnswerPreservesExistingOptionalSemantics(String answer) {
void faqWithoutAnswerPreservesOptionalSemantics(String answer) {
ChatContext ctx = context("退货流程是什么", true);
FaqMatchResult match = faqMatch(answer, "EXACT");
when(ragPipeline.tryFaqMatchClean(ctx.message(), ctx.categoryIds()))
.thenReturn(new RagPipeline.FaqMatchOutcome(Optional.of(match), true));
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());
assertSame(match, request.faqMatchResult());
verify(ragPipeline).tryFaqMatchClean(ctx.message(), ctx.categoryIds());
verifyNoMoreInteractions(ragPipeline);
verifyNoInteractions(intentRouter, ragHitLogService);
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))
.withRewriteStrategy("MULTI_QUERY")
.withEnableRag(enableRag);
}
@ -212,10 +297,4 @@ class ChatPipelineTests {
faq.setAnswer(answer);
return new FaqMatchResult(faq, matchType, 0.95);
}
private void stubRagRetrieval(ChatContext ctx, boolean faqAlreadyMatched) {
when(ragPipeline.retrieve(ctx, faqAlreadyMatched)).thenReturn(new RagContext(
Optional.empty(), List.of(new Document("退货说明")), "退货说明", "改写后的检索问题", "VECTOR", null));
when(ragPipeline.buildRagContextBlock("退货说明")).thenReturn("\n资料:退货说明");
}
}

126
src/test/java/com/wok/supportbot/IntentRouterTests.java

@ -1,126 +0,0 @@
package com.wok.supportbot;
import com.wok.supportbot.config.ChatModelFactory;
import com.wok.supportbot.service.IntentRouter;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.openai.OpenAiChatModel;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.openai.api.OpenAiApi;
import org.springframework.http.MediaType;
import org.springframework.test.web.client.MockRestServiceServer;
import org.springframework.web.client.RestClient;
import org.springframework.ai.openai.OpenAiChatOptions;
import org.springframework.test.util.ReflectionTestUtils;
import java.util.List;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.ArgumentMatchers.any;
import static org.springframework.test.web.client.match.MockRestRequestMatchers.jsonPath;
import static org.springframework.test.web.client.match.MockRestRequestMatchers.requestTo;
import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess;
import static org.mockito.Mockito.*;
class IntentRouterTests {
private final ChatModelFactory factory = mock(ChatModelFactory.class);
private final ChatModel model = mock(ChatModel.class);
private final IntentRouter router = new IntentRouter();
@BeforeEach
void setUp() {
ReflectionTestUtils.setField(router, "chatModelFactory", factory);
when(factory.getChatModel("CHAT")).thenReturn(model);
}
@Test
void seedClassificationDisablesThinkingWithoutChangingAnswerDefaults() {
OpenAiChatOptions defaults = OpenAiChatOptions.builder()
.model("doubao-seed-2-0-mini-260428")
.reasoningEffort("medium").maxTokens(2000).temperature(0.5).build();
when(model.getDefaultOptions()).thenReturn(defaults);
when(model.call(any(Prompt.class))).thenReturn(response("{\"intent\":\"RAG\",\"confidence\":0.95}"));
IntentRouter.IntentResult result = router.route("打印小票的排版规则是什么");
assertEquals("RAG", result.getIntent());
assertEquals(0.95, result.getConfidence());
ArgumentCaptor<Prompt> prompt = ArgumentCaptor.forClass(Prompt.class);
verify(model).call(prompt.capture());
OpenAiChatOptions options = assertInstanceOf(OpenAiChatOptions.class, prompt.getValue().getOptions());
assertEquals("minimal", options.getReasoningEffort());
assertEquals(128, options.getMaxTokens());
assertTrue(prompt.getValue().getContents().contains("打印小票的排版规则是什么"));
assertEquals("medium", defaults.getReasoningEffort());
assertEquals(2000, defaults.getMaxTokens());
assertEquals(0.5, defaults.getTemperature());
}
@Test
void otherModelsDoNotReceiveSeedSpecificOptions() {
when(model.getDefaultOptions()).thenReturn(OpenAiChatOptions.builder()
.model("deepseek-v4-flash").maxTokens(2000).build());
when(model.call(any(Prompt.class))).thenReturn(response("{\"intent\":\"CHITCHAT\",\"confidence\":0.9}"));
assertEquals("CHITCHAT", router.route("谢谢你的帮助").getIntent());
ArgumentCaptor<Prompt> prompt = ArgumentCaptor.forClass(Prompt.class);
verify(model).call(prompt.capture());
assertNull(prompt.getValue().getOptions());
}
@Test
void invalidOrTruncatedClassificationStillFallsBackToRetrieval() {
when(model.getDefaultOptions()).thenReturn(OpenAiChatOptions.builder()
.model("doubao-seed-2-0-mini-260428").build());
when(model.call(any(Prompt.class))).thenReturn(response("{\"intent\":"));
assertEquals("RAG", router.route("如何办理退款").getIntent());
}
@Test
void actualOpenAiRequestSendsFastClassificationOptions() {
RestClient.Builder restClient = RestClient.builder();
MockRestServiceServer server = MockRestServiceServer.bindTo(restClient).build();
OpenAiChatModel realModel = OpenAiChatModel.builder()
.openAiApi(OpenAiApi.builder().baseUrl("http://localhost").apiKey("test-key")
.restClientBuilder(restClient).build())
.defaultOptions(OpenAiChatOptions.builder().model("doubao-seed-2-0-mini-260428")
.maxTokens(2000).temperature(0.5).build())
.build();
when(factory.getChatModel("CHAT")).thenReturn(realModel);
server.expect(requestTo("http://localhost/v1/chat/completions"))
.andExpect(jsonPath("$.model").value("doubao-seed-2-0-mini-260428"))
.andExpect(jsonPath("$.reasoning_effort").value("minimal"))
.andExpect(jsonPath("$.max_tokens").value(128))
.andExpect(jsonPath("$.temperature").value(0.5))
.andRespond(withSuccess("""
{"id":"classification","object":"chat.completion","created":1,
"model":"doubao-seed-2-0-mini-260428",
"choices":[{"index":0,"message":{"role":"assistant",
"content":"{\\"intent\\":\\"RAG\\",\\"confidence\\":0.9}"},"finish_reason":"stop"}]}
""", MediaType.APPLICATION_JSON));
IntentRouter.IntentResult result = router.route("如何办理退款");
assertEquals("RAG", result.getIntent());
assertEquals(0.9, result.getConfidence());
server.verify();
assertEquals(2000, realModel.getDefaultOptions().getMaxTokens());
}
@Test
void blankQuestionDoesNotCallModel() {
assertEquals("RAG", router.route(" ").getIntent());
verifyNoInteractions(model);
verify(factory, never()).getChatModel(any());
}
private ChatResponse response(String text) {
return new ChatResponse(List.of(new Generation(new AssistantMessage(text))));
}
}
Loading…
Cancel
Save