|
|
@@ -0,0 +1,378 @@
|
|
|
+package com.zsjz.ai.module.agent.llm;
|
|
|
+
|
|
|
+import com.zsjz.ai.common.exception.ServerException;
|
|
|
+import com.zsjz.ai.module.agent.entity.AgentModel;
|
|
|
+import com.zsjz.ai.module.agent.mapper.AgentModelMapper;
|
|
|
+import com.zsjz.ai.module.agent.service.AgentModelFactory;
|
|
|
+import com.zsjz.ai.module.agent.service.AgentModelService;
|
|
|
+import io.agentscope.core.message.ContentBlock;
|
|
|
+import io.agentscope.core.message.Msg;
|
|
|
+import io.agentscope.core.message.MsgRole;
|
|
|
+import io.agentscope.core.message.TextBlock;
|
|
|
+import io.agentscope.core.model.ChatResponse;
|
|
|
+import io.agentscope.core.model.ChatUsage;
|
|
|
+import io.agentscope.core.model.GenerateOptions;
|
|
|
+import io.agentscope.core.model.Model;
|
|
|
+import io.agentscope.core.model.ToolSchema;
|
|
|
+import org.junit.jupiter.api.BeforeEach;
|
|
|
+import org.junit.jupiter.api.DisplayName;
|
|
|
+import org.junit.jupiter.api.Test;
|
|
|
+import org.mockito.ArgumentCaptor;
|
|
|
+import reactor.core.publisher.Flux;
|
|
|
+
|
|
|
+import java.time.Duration;
|
|
|
+import java.util.List;
|
|
|
+
|
|
|
+import static org.junit.jupiter.api.Assertions.assertEquals;
|
|
|
+import static org.junit.jupiter.api.Assertions.assertFalse;
|
|
|
+import static org.junit.jupiter.api.Assertions.assertNull;
|
|
|
+import static org.junit.jupiter.api.Assertions.assertThrows;
|
|
|
+import static org.junit.jupiter.api.Assertions.assertTrue;
|
|
|
+import static org.mockito.ArgumentMatchers.any;
|
|
|
+import static org.mockito.ArgumentMatchers.anyList;
|
|
|
+import static org.mockito.ArgumentMatchers.eq;
|
|
|
+import static org.mockito.Mockito.mock;
|
|
|
+import static org.mockito.Mockito.verify;
|
|
|
+import static org.mockito.Mockito.when;
|
|
|
+
|
|
|
+/**
|
|
|
+ * {@link LlmService} 的契约测试。
|
|
|
+ *
|
|
|
+ * <p>守住四件事:
|
|
|
+ * <ol>
|
|
|
+ * <li><b>模型解析</b>:不传 modelId 走默认对话模型,传了就走指定模型;两者都拿不到时
|
|
|
+ * 必须抛可读异常(而不是 NPE / 返回 null);</li>
|
|
|
+ * <li><b>流式分片拼接</b>:{@code Model#stream} 吐的是增量,服务要拼成完整文本,
|
|
|
+ * 且 usage 要带出来;</li>
|
|
|
+ * <li><b>结构化输出</b>:容忍 {@code ```json} 围栏与夹带的解释文字;解析不了要抛异常;</li>
|
|
|
+ * <li><b>失败语义</b>:一律 {@link ServerException},绝不静默返回 null
|
|
|
+ * (与 IntentService/FollowupService 的 fail-open 不同,这是调用方主动要结果的场景)。</li>
|
|
|
+ * </ol>
|
|
|
+ */
|
|
|
+class LlmServiceTest {
|
|
|
+
|
|
|
+ private AgentModelService agentModelService;
|
|
|
+ private AgentModelMapper agentModelMapper;
|
|
|
+ private AgentModelFactory agentModelFactory;
|
|
|
+ private LlmService llmService;
|
|
|
+
|
|
|
+ @BeforeEach
|
|
|
+ void setUp() {
|
|
|
+ agentModelService = mock(AgentModelService.class);
|
|
|
+ agentModelMapper = mock(AgentModelMapper.class);
|
|
|
+ agentModelFactory = mock(AgentModelFactory.class);
|
|
|
+ llmService = new LlmService(agentModelService, agentModelMapper, agentModelFactory);
|
|
|
+ }
|
|
|
+
|
|
|
+ // ==================== 辅助 ====================
|
|
|
+
|
|
|
+ private AgentModel config(Long id, String provider, String modelId) {
|
|
|
+ AgentModel c = new AgentModel();
|
|
|
+ c.setId(id);
|
|
|
+ c.setProvider(provider);
|
|
|
+ c.setModelId(modelId);
|
|
|
+ c.setName("测试模型");
|
|
|
+ return c;
|
|
|
+ }
|
|
|
+
|
|
|
+ /**
|
|
|
+ * 让 resolveModel 走通:默认模型 ID = 1,配置存在,Model 实例已 mock
|
|
|
+ */
|
|
|
+ private Model stubDefaultModel() {
|
|
|
+ AgentModel cfg = config(1L, "openai", "gpt-4o-mini");
|
|
|
+ when(agentModelService.getDefaultModelId()).thenReturn(1L);
|
|
|
+ when(agentModelMapper.selectById(1L)).thenReturn(cfg);
|
|
|
+ Model model = mock(Model.class);
|
|
|
+ when(agentModelFactory.create(cfg)).thenReturn(model);
|
|
|
+ return model;
|
|
|
+ }
|
|
|
+
|
|
|
+ private static ChatResponse text(String content) {
|
|
|
+ return ChatResponse.builder()
|
|
|
+ .content(List.of(TextBlock.builder().text(content).build()))
|
|
|
+ .build();
|
|
|
+ }
|
|
|
+
|
|
|
+ // ==================== 1. 基本调用 ====================
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chat(prompt):用默认模型,把流式分片拼成完整文本")
|
|
|
+ void chatWithDefaultModel() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(
|
|
|
+ text("本案共"), text(" 38263 条"), text("通话记录")));
|
|
|
+
|
|
|
+ String result = llmService.chat("统计一下通话");
|
|
|
+
|
|
|
+ assertEquals("本案共 38263 条通话记录", result);
|
|
|
+
|
|
|
+ @SuppressWarnings("unchecked")
|
|
|
+ ArgumentCaptor<List<Msg>> captor = ArgumentCaptor.forClass(List.class);
|
|
|
+ verify(model).stream(captor.capture(), anyList(), any());
|
|
|
+ List<Msg> messages = captor.getValue();
|
|
|
+ assertEquals(1, messages.size(), "没给 system 时只应有一条 user 消息");
|
|
|
+ assertEquals(MsgRole.USER, messages.get(0).getRole());
|
|
|
+ assertEquals("统计一下通话", messages.get(0).getTextContent());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chat(modelId, system, user):system 在前,且用的是指定模型")
|
|
|
+ void chatWithExplicitModelAndSystem() {
|
|
|
+ AgentModel cfg = config(7L, "dashscope", "qwen-plus");
|
|
|
+ when(agentModelMapper.selectById(7L)).thenReturn(cfg);
|
|
|
+ Model model = mock(Model.class);
|
|
|
+ when(agentModelFactory.create(cfg)).thenReturn(model);
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(text("已分类")));
|
|
|
+
|
|
|
+ String result = llmService.chat(7L, "你是金融分类助手", "给这笔交易打标签");
|
|
|
+
|
|
|
+ assertEquals("已分类", result);
|
|
|
+ // 显式传了 modelId 就不该再去查默认模型
|
|
|
+ verify(agentModelService, org.mockito.Mockito.never()).getDefaultModelId();
|
|
|
+
|
|
|
+ @SuppressWarnings("unchecked")
|
|
|
+ ArgumentCaptor<List<Msg>> captor = ArgumentCaptor.forClass(List.class);
|
|
|
+ verify(model).stream(captor.capture(), anyList(), any());
|
|
|
+ List<Msg> messages = captor.getValue();
|
|
|
+ assertEquals(2, messages.size());
|
|
|
+ assertEquals(MsgRole.SYSTEM, messages.get(0).getRole());
|
|
|
+ assertEquals("你是金融分类助手", messages.get(0).getTextContent());
|
|
|
+ assertEquals(MsgRole.USER, messages.get(1).getRole());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("messages 优先于 system/user")
|
|
|
+ void messagesWinOverSystemUser() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(text("ok")));
|
|
|
+
|
|
|
+ List<Msg> history = List.of(
|
|
|
+ Msg.builder().role(MsgRole.USER).textContent("第一轮").build(),
|
|
|
+ Msg.builder().role(MsgRole.ASSISTANT).textContent("回答").build(),
|
|
|
+ Msg.builder().role(MsgRole.USER).textContent("第二轮").build());
|
|
|
+ llmService.chat(LlmRequest.builder()
|
|
|
+ .system("不该出现")
|
|
|
+ .user("也不该出现")
|
|
|
+ .messages(history)
|
|
|
+ .build());
|
|
|
+
|
|
|
+ @SuppressWarnings("unchecked")
|
|
|
+ ArgumentCaptor<List<Msg>> captor = ArgumentCaptor.forClass(List.class);
|
|
|
+ verify(model).stream(captor.capture(), anyList(), any());
|
|
|
+ assertEquals(3, captor.getValue().size());
|
|
|
+ assertEquals("第一轮", captor.getValue().get(0).getTextContent());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chatDetail 带出 usage 与耗时")
|
|
|
+ void chatDetailCarriesUsage() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(
|
|
|
+ ChatResponse.builder()
|
|
|
+ .content(List.of(TextBlock.builder().text("结果").build()))
|
|
|
+ .usage(ChatUsage.builder().inputTokens(120).outputTokens(8).build())
|
|
|
+ .build()));
|
|
|
+
|
|
|
+ LlmResult result = llmService.chatDetail(LlmRequest.builder().user("hi").build());
|
|
|
+
|
|
|
+ assertEquals("结果", result.text());
|
|
|
+ assertEquals(1L, result.modelId());
|
|
|
+ assertEquals("openai:gpt-4o-mini", result.modelName());
|
|
|
+ assertEquals(120, result.inputTokens());
|
|
|
+ assertEquals(8, result.outputTokens());
|
|
|
+ assertEquals(128, result.totalTokens());
|
|
|
+ assertTrue(result.elapsedMs() >= 0);
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("采样参数透传给 GenerateOptions")
|
|
|
+ void samplingOptionsAreForwarded() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(text("x")));
|
|
|
+
|
|
|
+ llmService.chat(LlmRequest.builder()
|
|
|
+ .user("hi").temperature(0.3).maxTokens(256).build());
|
|
|
+
|
|
|
+ ArgumentCaptor<GenerateOptions> captor = ArgumentCaptor.forClass(GenerateOptions.class);
|
|
|
+ verify(model).stream(anyList(), anyList(), captor.capture());
|
|
|
+ assertEquals(0.3, captor.getValue().getTemperature());
|
|
|
+ assertEquals(256, captor.getValue().getMaxTokens());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("stream() 原样下发分片")
|
|
|
+ void streamEmitsChunks() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(
|
|
|
+ text("逐"), text("字"), text("上屏")));
|
|
|
+
|
|
|
+ List<String> chunks = llmService.stream(LlmRequest.builder().user("hi").build())
|
|
|
+ .collectList().block();
|
|
|
+
|
|
|
+ assertEquals(List.of("逐", "字", "上屏"), chunks);
|
|
|
+ }
|
|
|
+
|
|
|
+ // ==================== 2. 结构化输出 ====================
|
|
|
+
|
|
|
+ /** 结构化输出的目标类型 */
|
|
|
+ public static class TagResult {
|
|
|
+ public String label;
|
|
|
+ public double confidence;
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chatAs:剥掉 ```json 围栏后反序列化")
|
|
|
+ void chatAsStripsFence() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(
|
|
|
+ text("```json\n{\"label\":\"大额交易\",\"confidence\":0.92}\n```")));
|
|
|
+
|
|
|
+ TagResult result = llmService.chatAs(LlmRequest.builder().user("打标签").build(), TagResult.class);
|
|
|
+
|
|
|
+ assertEquals("大额交易", result.label);
|
|
|
+ assertEquals(0.92, result.confidence);
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chatAs:容忍前后夹带的解释文字,且把 schema 注入 system")
|
|
|
+ void chatAsToleratesProseAndInjectsSchema() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(
|
|
|
+ text("好的,分析结果如下:{\"label\":\"夜间通话\",\"confidence\":0.8} 以上。")));
|
|
|
+
|
|
|
+ TagResult result = llmService.chatAs(LlmRequest.builder()
|
|
|
+ .system("你是分类助手").user("打标签").build(), TagResult.class);
|
|
|
+
|
|
|
+ assertEquals("夜间通话", result.label);
|
|
|
+
|
|
|
+ @SuppressWarnings("unchecked")
|
|
|
+ ArgumentCaptor<List<Msg>> captor = ArgumentCaptor.forClass(List.class);
|
|
|
+ verify(model).stream(captor.capture(), anyList(), any());
|
|
|
+ String system = captor.getValue().get(0).getTextContent();
|
|
|
+ assertTrue(system.startsWith("你是分类助手"), "原 system 必须保留: " + system);
|
|
|
+ assertTrue(system.contains("只输出一个 JSON 对象"), "缺少 JSON 约束: " + system);
|
|
|
+ assertTrue(system.contains("schema"), "缺少 schema 说明: " + system);
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chatAs:模型没吐 JSON 时抛 500 而不是返回 null")
|
|
|
+ void chatAsFailsOnNonJson() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(text("抱歉,我无法完成")));
|
|
|
+
|
|
|
+ ServerException e = assertThrows(ServerException.class,
|
|
|
+ () -> llmService.chatAs(LlmRequest.builder().user("x").build(), TagResult.class));
|
|
|
+ assertEquals(500, e.getCode());
|
|
|
+ assertTrue(e.getMessage().contains("未返回合法 JSON"), e.getMessage());
|
|
|
+ }
|
|
|
+
|
|
|
+ // ==================== 3. 失败语义 ====================
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("没有默认模型时抛 400 并给出可操作的提示")
|
|
|
+ void noDefaultModel() {
|
|
|
+ when(agentModelService.getDefaultModelId()).thenReturn(null);
|
|
|
+
|
|
|
+ ServerException e = assertThrows(ServerException.class, () -> llmService.chat("hi"));
|
|
|
+ assertEquals(400, e.getCode());
|
|
|
+ assertTrue(e.getMessage().contains("默认"), e.getMessage());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("模型配置不存在时抛 404")
|
|
|
+ void modelConfigMissing() {
|
|
|
+ when(agentModelMapper.selectById(99L)).thenReturn(null);
|
|
|
+
|
|
|
+ ServerException e = assertThrows(ServerException.class, () -> llmService.chat(99L, "hi"));
|
|
|
+ assertEquals(404, e.getCode());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("缺提示词时抛 400")
|
|
|
+ void missingPrompt() {
|
|
|
+ stubDefaultModel();
|
|
|
+
|
|
|
+ ServerException e = assertThrows(ServerException.class,
|
|
|
+ () -> llmService.chat(LlmRequest.builder().build()));
|
|
|
+ assertEquals(400, e.getCode());
|
|
|
+ assertTrue(e.getMessage().contains("提示词"), e.getMessage());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("模型迟迟不返回时按 timeout 抛 504")
|
|
|
+ void timeoutBecomes504() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.never());
|
|
|
+
|
|
|
+ ServerException e = assertThrows(ServerException.class, () -> llmService.chat(
|
|
|
+ LlmRequest.builder().user("hi").timeout(Duration.ofMillis(150)).build()));
|
|
|
+ assertEquals(504, e.getCode());
|
|
|
+ assertTrue(e.getMessage().contains("超时"), e.getMessage());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("模型内部异常包成 500,不把 Reactor 的栈甩给调用方")
|
|
|
+ void modelErrorBecomes500() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(
|
|
|
+ Flux.error(new IllegalStateException("401 invalid api key")));
|
|
|
+
|
|
|
+ ServerException e = assertThrows(ServerException.class, () -> llmService.chat("hi"));
|
|
|
+ assertEquals(500, e.getCode());
|
|
|
+ assertTrue(e.getMessage().contains("401 invalid api key"), e.getMessage());
|
|
|
+ }
|
|
|
+
|
|
|
+ // ==================== 4. 边界 ====================
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("非文本块(thinking / tool_use)不进正文")
|
|
|
+ void nonTextBlocksAreSkipped() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ ContentBlock thinking = mock(ContentBlock.class);
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(
|
|
|
+ ChatResponse.builder().content(List.of(thinking)).build(),
|
|
|
+ text("正文")));
|
|
|
+
|
|
|
+ assertEquals("正文", llmService.chat("hi"));
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("timeout 默认 60s,显式设置时以显式值为准")
|
|
|
+ void timeoutDefaults() {
|
|
|
+ assertEquals(Duration.ofSeconds(60), LlmRequest.builder().build().timeoutOrDefault());
|
|
|
+ assertEquals(Duration.ofSeconds(5),
|
|
|
+ LlmRequest.builder().timeout(Duration.ofSeconds(5)).build().timeoutOrDefault());
|
|
|
+ assertNull(LlmRequest.builder().build().getTimeout());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("工具 schema 传空列表:直连不带工具")
|
|
|
+ void noToolsAreSent() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(text("x")));
|
|
|
+
|
|
|
+ llmService.chat("hi");
|
|
|
+
|
|
|
+ @SuppressWarnings("unchecked")
|
|
|
+ ArgumentCaptor<List<ToolSchema>> captor = ArgumentCaptor.forClass(List.class);
|
|
|
+ verify(model).stream(anyList(), captor.capture(), any());
|
|
|
+ assertTrue(captor.getValue().isEmpty(), "直连调用不该带任何工具");
|
|
|
+ assertFalse(captor.getValue() == null);
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("chat(null) 抛 400 而不是 NPE")
|
|
|
+ void nullRequest() {
|
|
|
+ ServerException e = assertThrows(ServerException.class, () -> llmService.chat((LlmRequest) null));
|
|
|
+ assertEquals(400, e.getCode());
|
|
|
+ }
|
|
|
+
|
|
|
+ @Test
|
|
|
+ @DisplayName("stream() 的 Flux 可直接订阅,不需要 block")
|
|
|
+ void streamIsLazy() {
|
|
|
+ Model model = stubDefaultModel();
|
|
|
+ when(model.stream(anyList(), anyList(), any())).thenReturn(Flux.just(text("a")));
|
|
|
+
|
|
|
+ assertEquals(1, llmService.stream(LlmRequest.builder().user("hi").build()).count().block());
|
|
|
+ }
|
|
|
+}
|