diff --git a/src/main/java/com/wok/supportbot/service/IntentRouter.java b/src/main/java/com/wok/supportbot/service/IntentRouter.java index 5c9c88e..175b1dc 100644 --- a/src/main/java/com/wok/supportbot/service/IntentRouter.java +++ b/src/main/java/com/wok/supportbot/service/IntentRouter.java @@ -8,6 +8,7 @@ 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; @@ -33,6 +34,9 @@ public class IntentRouter { private static final BeanOutputConverter INTENT_CONVERTER = new BeanOutputConverter<>(IntentResult.class); + /** 分类只需返回 intent/confidence;不改变主回答模型的深度思考与输出长度。 */ + private static final int CLASSIFICATION_MAX_TOKENS = 128; + /** * 意图分类 Prompt 模板。 * 输出格式约束(JSON Schema)由 {@code BeanOutputConverter.getFormat()} 统一追加,模板中不再硬编码。 @@ -80,7 +84,8 @@ public class IntentRouter { ChatModel chatModel = chatModelFactory.getChatModel("CHAT"); String promptText = INTENT_PROMPT_TEMPLATE.formatted(userQuestion, INTENT_CONVERTER.getFormat()); - String response = chatModel.call(new Prompt(promptText)).getResult().getOutput().getText(); + String response = chatModel.call(classificationPrompt(chatModel, promptText)) + .getResult().getOutput().getText(); log.debug("意图分类原始响应: {}", response); IntentResult result = INTENT_CONVERTER.convert(response); @@ -95,6 +100,21 @@ public class IntentRouter { } } + 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); + } + /** * 校验意图类型是否有效(模型可能返回枚举外的值) */ diff --git a/src/test/java/com/wok/supportbot/IntentRouterTests.java b/src/test/java/com/wok/supportbot/IntentRouterTests.java new file mode 100644 index 0000000..50899fa --- /dev/null +++ b/src/test/java/com/wok/supportbot/IntentRouterTests.java @@ -0,0 +1,126 @@ +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 = 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 = 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)))); + } +}