Browse Source

perf(intent): doubao-seed-2.0 分类请求关闭深度思考

feature/test
wei-py 3 weeks ago
parent
commit
e1ecc18223
  1. 22
      src/main/java/com/wok/supportbot/service/IntentRouter.java
  2. 126
      src/test/java/com/wok/supportbot/IntentRouterTests.java

22
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<IntentResult> 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);
}
/**
* 校验意图类型是否有效(模型可能返回枚举外的值)
*/

126
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> 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