本地 RAG 知识库
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

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))));
}
}