You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
126 lines
6.0 KiB
126 lines
6.0 KiB
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))));
|
|
}
|
|
}
|