本地 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.
 
 
 
 
 
 

250 lines
12 KiB

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<String, ChatClient> cache = (Map<String, ChatClient>) 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<Document> 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<Document> 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<String> result = app.chatStreamOpenAi(ctx);
verifyNoInteractions(pipeline);
List<String> chunks = result.collectList().block(Duration.ofSeconds(5));
assertEnvelope(chunks, SourceReference.fromDocuments(documents), "## 回答\n");
verify(pipeline).buildRequest(ctx);
verifyNoMoreInteractions(pipeline);
verify(stream).chatResponse();
ArgumentCaptor<LlmCallTrace> 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<String> 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.<ChatResponse>never())
.doOnCancel(() -> cancelled.set(true)));
List<String> 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<SourceReference> 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<String> chunks, List<SourceReference> 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<Document> 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))));
}
}