6 changed files with 222 additions and 467 deletions
-
7src/main/java/com/wok/supportbot/app/ChatContext.java
-
87src/main/java/com/wok/supportbot/app/ChatPipeline.java
-
18src/main/java/com/wok/supportbot/rag/RagPipeline.java
-
134src/main/java/com/wok/supportbot/service/IntentRouter.java
-
317src/test/java/com/wok/supportbot/ChatPipelineTests.java
-
126src/test/java/com/wok/supportbot/IntentRouterTests.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; |
|||
} |
|||
} |
|||
@ -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)))); |
|||
} |
|||
} |
|||
Write
Preview
Loading…
Cancel
Save
Reference in new issue