package com.wok.supportbot; import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; import com.wok.supportbot.app.AssistantApp; import com.wok.supportbot.app.ChatContext; import com.wok.supportbot.app.ChatPipeline; import com.wok.supportbot.app.ChatRequest; import com.wok.supportbot.app.ChatResult; import com.wok.supportbot.app.SourceReference; import com.wok.supportbot.chatmemory.DatabaseChatMemory; import com.wok.supportbot.config.ChatModelFactory; import com.wok.supportbot.config.SimpleCircuitBreaker; import com.wok.supportbot.entity.LlmCallTrace; import com.wok.supportbot.service.AiModelConfigService; import com.wok.supportbot.service.ContentSafetyService; import com.wok.supportbot.service.LlmCallTraceService; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import org.junit.jupiter.params.provider.ValueSource; import org.mockito.ArgumentCaptor; import org.springframework.ai.chat.client.ChatClient; import org.springframework.ai.chat.messages.AssistantMessage; import org.springframework.ai.chat.model.ChatResponse; import org.springframework.ai.chat.model.Generation; import org.springframework.ai.document.Document; import org.springframework.test.util.ReflectionTestUtils; import reactor.core.publisher.Flux; import java.time.Duration; import java.util.ArrayList; import java.util.List; import java.util.Map; import java.util.Optional; import java.util.concurrent.atomic.AtomicBoolean; import static org.junit.jupiter.api.Assertions.*; import static org.mockito.ArgumentMatchers.*; import static org.mockito.Mockito.*; class AnswerTransportTests { private static final ObjectMapper JSON = new ObjectMapper(); private AssistantApp app; private ChatPipeline pipeline; private ChatClient.ChatClientRequestSpec spec; private ChatClient.CallResponseSpec call; private ChatClient.StreamResponseSpec stream; private LlmCallTraceService traces; @BeforeEach @SuppressWarnings("unchecked") void setUp() { app = new AssistantApp(mock(ChatModelFactory.class), mock(DatabaseChatMemory.class)); pipeline = mock(ChatPipeline.class); traces = mock(LlmCallTraceService.class); ContentSafetyService safety = mock(ContentSafetyService.class); when(safety.mask(any())).thenAnswer(invocation -> invocation.getArgument(0)); ReflectionTestUtils.setField(app, "chatPipeline", pipeline); ReflectionTestUtils.setField(app, "llmCallTraceService", traces); ReflectionTestUtils.setField(app, "contentSafetyService", safety); ReflectionTestUtils.setField(app, "aiModelConfigService", mock(AiModelConfigService.class)); ChatClient client = mock(ChatClient.class); spec = mock(ChatClient.ChatClientRequestSpec.class, RETURNS_SELF); call = mock(ChatClient.CallResponseSpec.class); stream = mock(ChatClient.StreamResponseSpec.class); when(client.prompt()).thenReturn(spec); when(spec.call()).thenReturn(call); when(spec.stream()).thenReturn(stream); Map cache = (Map) ReflectionTestUtils.getField(app, "chatClientCache"); cache.put("CHAT:none", client); } @ParameterizedTest @ValueSource(booleans = {true, false}) void synchronousAnswerUsesOnlyThisBuildsDocuments(boolean rag) throws Exception { ChatContext ctx = context(rag); List documents = rag ? List.of(document()) : List.of(); when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, documents, null)); when(call.chatResponse()).thenReturn(response("回答")); ChatResult result = app.chatWithEvents(ctx); assertEquals("回答", result.text()); assertEquals(SourceReference.fromDocuments(documents), result.sources()); assertTrue(result.suggestions().isEmpty()); JsonNode json = JSON.valueToTree(result); if (rag) { assertEquals("9223372036854775807", json.at("/sources/0/documentId").asText()); assertTrue(json.at("/sources/0/documentId").isTextual()); } verify(pipeline).buildRequest(ctx); verifyNoMoreInteractions(pipeline); verify(call).chatResponse(); } @ParameterizedTest @ValueSource(booleans = {true, false}) void completedStreamCarriesExactlyOneMetadataChunkBeforeStop(boolean rag) throws Exception { ChatContext ctx = context(rag); List documents = rag ? List.of(document()) : List.of(); when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, documents, null)); when(stream.chatResponse()).thenReturn(Flux.just(response("## "), response("回答\n"))); Flux result = app.chatStreamOpenAi(ctx); verifyNoInteractions(pipeline); List chunks = result.collectList().block(Duration.ofSeconds(5)); assertEnvelope(chunks, SourceReference.fromDocuments(documents), "## 回答\n"); verify(pipeline).buildRequest(ctx); verifyNoMoreInteractions(pipeline); verify(stream).chatResponse(); ArgumentCaptor trace = ArgumentCaptor.forClass(LlmCallTrace.class); verify(traces, timeout(1000)).recordAsync(trace.capture()); assertEquals("## 回答\n", trace.getValue().getAiResponse()); assertEquals("COMPLETE", trace.getValue().getStatus()); } @Test void faqHasEmptySourcesWithoutModelInvocation() throws Exception { ChatContext ctx = context(true); when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(), "标准答案")); assertEnvelope(app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)), List.of(), "标准答案"); verify(pipeline).buildRequest(ctx); verifyNoMoreInteractions(pipeline); verifyNoInteractions(call, stream); } @Test void synchronousFaqHasEmptySources() { ChatContext ctx = context(true); when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(), "标准答案")); ChatResult result = app.chatWithEvents(ctx); assertEquals("标准答案", result.text()); assertTrue(result.sources().isEmpty()); verify(pipeline).buildRequest(ctx); verifyNoMoreInteractions(pipeline); verifyNoInteractions(call, stream); } @Test void circuitFallbackDoesNotBuildOrRetrieve() throws Exception { SimpleCircuitBreaker breaker = (SimpleCircuitBreaker) ReflectionTestUtils.getField(app, "aiCircuitBreaker"); for (int i = 0; i < 3; i++) breaker.recordFailure(-1L); ChatContext ctx = context(true); assertTrue(app.chatWithEvents(ctx).sources().isEmpty()); assertEnvelope(app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)), List.of(), "AI 服务暂时不可用,请稍后重试。"); verifyNoInteractions(pipeline, call, stream); } @Test void streamFailureDoesNotRepeatGenerationOrExposeUnusedSources() throws Exception { ChatContext ctx = context(true); when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(document()), null)); when(stream.chatResponse()).thenReturn(Flux.concat(Flux.just(response("部分回答")), Flux.error(new IllegalStateException("模型超时")))); List chunks = app.chatStreamOpenAi(ctx).collectList().block(Duration.ofSeconds(5)); assertEnvelope(chunks, List.of(), "部分回答抱歉,AI 服务调用失败:模型超时"); verify(pipeline).buildRequest(ctx); verifyNoMoreInteractions(pipeline); verify(stream).chatResponse(); verifyNoInteractions(call); } @Test void cancellationPropagatesWithoutMetadataOrSecondRequest() throws Exception { ChatContext ctx = context(true); AtomicBoolean cancelled = new AtomicBoolean(); when(pipeline.buildRequest(ctx)).thenReturn(request(ctx, List.of(document()), null)); when(stream.chatResponse()).thenReturn(Flux.concat(Flux.just(response("首段")), Flux.never()) .doOnCancel(() -> cancelled.set(true))); List chunks = app.chatStreamOpenAi(ctx).take(2).collectList().block(Duration.ofSeconds(5)); assertEquals(2, chunks.size()); assertEquals("首段", JSON.readTree(chunks.get(1)).at("/choices/0/delta/content").asText()); assertFalse(chunks.stream().anyMatch(chunk -> chunk.contains("\"sources\"") || chunk.equals("[DONE]"))); verify(traces, timeout(1000)).recordAsync(argThat(trace -> "CANCEL".equals(trace.getStatus()))); assertTrue(cancelled.get()); verify(pipeline).buildRequest(ctx); verifyNoMoreInteractions(pipeline); verify(stream).chatResponse(); verifyNoInteractions(call); } @Test void sourceSerializationPreservesNullableFieldsDistanceAndSnippetBoundaries() { Document doc = document(); SourceReference source = SourceReference.fromDocuments(List.of(doc)).get(0); assertEquals("9223372036854775807", source.documentId()); assertEquals(0.25, source.score()); assertEquals(2, source.chunkIndex()); assertEquals(160, source.snippet().length()); assertTrue(source.snippet().endsWith("…")); SourceReference nullable = SourceReference.fromDocuments(List.of(new Document("short"))).get(0); JsonNode json = JSON.copy().setSerializationInclusion(com.fasterxml.jackson.annotation.JsonInclude.Include.NON_NULL) .valueToTree(nullable); assertEquals(6, json.size()); assertTrue(json.get("documentId").isNull()); assertTrue(json.get("score").isNull()); SourceReference unicode = new SourceReference(null, null, null, null, null, "a".repeat(158) + "😀xx"); assertTrue(unicode.snippet().length() <= 160); assertFalse(Character.isHighSurrogate(unicode.snippet().charAt(unicode.snippet().length() - 2))); ArrayList mutable = new ArrayList<>(List.of(source)); ChatResult result = new ChatResult("answer", null, null, mutable); mutable.clear(); assertEquals(List.of(source), result.sources()); } private static void assertEnvelope(List chunks, List sources, String text) throws Exception { assertNotNull(chunks); assertEquals("[DONE]", chunks.get(chunks.size() - 1)); JsonNode metadata = JSON.readTree(chunks.get(chunks.size() - 3)); JsonNode stop = JSON.readTree(chunks.get(chunks.size() - 2)); assertEquals("stop", stop.at("/choices/0/finish_reason").asText()); assertEquals(0, metadata.get("choices").size()); assertEquals(JSON.valueToTree(sources), metadata.get("sources")); StringBuilder answer = new StringBuilder(); int metadataCount = 0; for (String raw : chunks.subList(0, chunks.size() - 1)) { JsonNode chunk = JSON.readTree(raw); assertEquals("chat.completion.chunk", chunk.get("object").asText()); assertEquals(metadata.get("id"), chunk.get("id")); assertEquals(metadata.get("model"), chunk.get("model")); assertEquals(metadata.get("created"), chunk.get("created")); if (chunk.has("sources")) metadataCount++; answer.append(chunk.at("/choices/0/delta/content").asText("")); } assertEquals(1, metadataCount); assertEquals(text, answer.toString()); } private static Document document() { return new Document("知识".repeat(100), Map.of("documentId", Long.MAX_VALUE, "title", "授权文档", "sourceName", "manual.pdf", "chunkIndex", "2", "distance", 0.25, "score", 0.9)); } private static ChatContext context(boolean rag) { return new ChatContext("退货", "transport-chat", "CHAT", null, null, List.of(7L), "NONE", rag, false, 11L, "售后", "account", null, null); } private static ChatRequest request(ChatContext ctx, List documents, String faq) { return new ChatRequest(ctx, ctx.message(), "system", Optional.ofNullable(faq), "system", documents.isEmpty() ? null : "资料", documents.size(), faq != null ? "FAQ" : "RAG", "VECTOR", documents, null); } private static ChatResponse response(String text) { return new ChatResponse(List.of(new Generation(new AssistantMessage(text)))); } }