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