Browse Source

feat(chat): 同次回答携带引用来源

feature/test
wei-py 3 weeks ago
parent
commit
28e77223ef
  1. 50
      src/main/java/com/wok/supportbot/app/AssistantApp.java
  2. 10
      src/main/java/com/wok/supportbot/app/ChatResult.java
  3. 59
      src/main/java/com/wok/supportbot/app/SourceReference.java
  4. 73
      src/main/java/com/wok/supportbot/controller/AiController.java
  5. 14
      src/main/java/com/wok/supportbot/controller/OpenApiController.java
  6. 250
      src/test/java/com/wok/supportbot/AnswerTransportTests.java
  7. 166
      src/test/java/com/wok/supportbot/ChatResultEndpointTests.java

50
src/main/java/com/wok/supportbot/app/AssistantApp.java

@ -36,6 +36,7 @@ import reactor.core.Disposable;
import reactor.core.publisher.Flux;
import reactor.core.publisher.FluxSink;
import reactor.core.publisher.SignalType;
import reactor.core.scheduler.Schedulers;
import java.util.ArrayList;
import java.util.Collections;
@ -314,19 +315,19 @@ public class AssistantApp {
*/
public ChatResult chatWithEvents(ChatContext ctx) {
long startNanos = System.nanoTime();
// 熔断:全局 AI 调用处于熔断状态,直接返回降级提示(不做 buildRequest,避免熔断期间仍走意图路由/检索)
// 熔断时直接返回降级提示,不再编排或检索。
if (aiCircuitBreaker.isOpen(AI_CIRCUIT_KEY)) {
log.warn("AI 调用熔断中,返回降级提示");
recordTrace(ctx, null, CIRCUIT_OPEN_MESSAGE, 0, "BYPASS",
new TraceMeta("CIRCUIT_BREAK", "AI 服务熔断降级", null, null, null, null));
return new ChatResult(CIRCUIT_OPEN_MESSAGE, List.of());
return new ChatResult(CIRCUIT_OPEN_MESSAGE, List.of(), List.of(), List.of());
}
ChatRequest req = chatPipeline.buildRequest(ctx);
if (req.faqHit()) {
String faqAnswer = req.faqAnswer().get();
recordTrace(ctx, req, faqAnswer, 0, "FAQ",
new TraceMeta(null, null, null, null, null, null));
return new ChatResult(faqAnswer, List.of());
return new ChatResult(faqAnswer, List.of(), List.of(), List.of());
}
// 显式事件收集器 + 轮次计数器,通过 toolContext 传给 McpToolCallback,规避 Reactor 跨线程丢 ThreadLocal 的问题
List<ToolCallEvent> events = new CopyOnWriteArrayList<>();
@ -354,14 +355,14 @@ public class AssistantApp {
events));
// 推荐问题已不再由主回复同步生成,改由 SuggestionGenerator 异步按需生成
return new ChatResult(text, events, List.of());
return new ChatResult(text, events, List.of(), SourceReference.fromDocuments(req.hitDocuments()));
} catch (Exception e) {
aiCircuitBreaker.recordFailure(AI_CIRCUIT_KEY);
log.error("AI 同步调用失败: chatId={}, error={}", ctx.chatId(), e.getMessage());
String fallback = "抱歉,AI 服务调用失败:" + e.getMessage();
recordTrace(ctx, req, fallback, elapsedMillis(startNanos), "ERROR",
new TraceMeta(classifyError(e), maskError(e.getMessage()), null, null, null, events));
return new ChatResult(fallback, List.of());
return new ChatResult(fallback, List.of(), List.of(), List.of());
}
}
@ -376,7 +377,7 @@ public class AssistantApp {
*/
public Flux<String> chatStream(ChatContext ctx) {
long startNanos = System.nanoTime();
// 熔断:全局 AI 调用处于熔断状态(不做 buildRequest,避免熔断期间仍走意图路由/检索)
// 熔断时直接返回降级提示,不再编排或检索。
if (aiCircuitBreaker.isOpen(AI_CIRCUIT_KEY)) {
log.warn("AI 调用熔断中(流式),返回降级提示");
recordTrace(ctx, null, CIRCUIT_OPEN_MESSAGE, 0, "BYPASS",
@ -452,7 +453,7 @@ public class AssistantApp {
* <p>
* 复用 {@link #chatStream(ChatContext)} 的完整编排逻辑(熔断早退 / FAQ 命中早退 /
* 正常流式调用 / 空白缓冲 / 埋点),差异在于把每个文本片段包装为 OpenAI 标准 JSON chunk:
* 首片 delta 携带 role=assistant,流结束时追加 finish_reason=stop 的 chunk 与 [DONE]。
* 首片 delta 携带 role=assistant,正文结束后追加 sources 元数据、finish_reason=stop 与 [DONE]。
* <p>
* 每个 Flux 元素即一个完整 JSON 字符串,Spring WebFlux 自动加 data: 前缀。
*
@ -460,6 +461,10 @@ public class AssistantApp {
* @return OpenAI 标准格式的流式回答
*/
public Flux<String> chatStreamOpenAi(ChatContext ctx) {
return Flux.defer(() -> buildOpenAiStream(ctx)).subscribeOn(Schedulers.boundedElastic());
}
private Flux<String> buildOpenAiStream(ChatContext ctx) {
long startNanos = System.nanoTime();
// OpenAI 标准 chunk 的公共元信息:同一次流式回答共享 id / created / model
String completionId = "chatcmpl-" + UUID.randomUUID().toString().replace("-", "");
@ -472,7 +477,7 @@ public class AssistantApp {
log.warn("获取活跃模型配置失败,model 回退 unknown: chatId={}, error={}", ctx.chatId(), e.getMessage());
}
String model = (cfg != null && cfg.getModelName() != null) ? cfg.getModelName() : "unknown";
// 熔断:全局 AI 调用处于熔断状态(不做 buildRequest,避免熔断期间仍走意图路由/检索)
// 熔断时直接返回降级提示,不再编排或检索。
if (aiCircuitBreaker.isOpen(AI_CIRCUIT_KEY)) {
log.warn("AI 调用熔断中(OpenAI 流式),返回降级提示");
recordTrace(ctx, null, CIRCUIT_OPEN_MESSAGE, 0, "BYPASS",
@ -518,7 +523,8 @@ public class AssistantApp {
});
// 聚合所有分片用于埋点(在 doFinally 时取完整回复文本)
StringBuilder aggregated = new StringBuilder();
return preserveTrailingWhitespace(rawStream)
// JSON 编码会保留正文空白,无需为 SSE 行尾 trim 缓冲 token。
return rawStream.filter(chunk -> !chunk.isEmpty())
.doOnNext(aggregated::append)
.map(chunk -> buildOpenAiChunk(completionId, model, created, chunk, false, null))
.doOnComplete(() -> aiCircuitBreaker.recordSuccess(AI_CIRCUIT_KEY))
@ -540,12 +546,12 @@ public class AssistantApp {
usage != null ? usage.getTotalTokens() : null,
events));
})
// 首片(仅 role=assistant、无 content)在流订阅时立即发出,确保 SSE 响应头/首字节及时 flush。
// 推理模型(如 doubao-seed)思考阶段 delta.content 为空、被 preserveTrailingWhitespace 吞掉,
// 若不提前发首片,思考阶段将无任何字节输出,前端等待首字节会触发 60s 超时。
// 编排结束后先发送 role 协议帧,避免模型思考期间连接完全静默。
// 此帧没有正文,不代表用户已收到首个回答 token。
.startWith(buildOpenAiChunk(completionId, model, created, "", true, null))
// 流正常结束时追加 finish_reason=stop 的 chunk 与 [DONE]
// 来源只取本次编排命中;元数据不经过正文聚合与 trace。
.concatWith(Flux.just(
buildSourcesChunk(completionId, model, created, SourceReference.fromDocuments(req.hitDocuments())),
buildOpenAiChunk(completionId, model, created, "", false, "stop"),
"[DONE]"))
// 错误兜底:脱敏错误信息,避免泄露内部细节(首片 role 已提前发出,此处不再带 role)
@ -594,8 +600,23 @@ public class AssistantApp {
}
}
private String buildSourcesChunk(String id, String model, long created, List<SourceReference> sources) {
Map<String, Object> chunk = new LinkedHashMap<>();
chunk.put("id", id);
chunk.put("object", "chat.completion.chunk");
chunk.put("created", created);
chunk.put("model", model);
chunk.put("choices", List.of());
chunk.put("sources", sources);
try {
return OBJECT_MAPPER.writeValueAsString(chunk);
} catch (Exception e) {
throw new IllegalStateException("序列化引用来源失败", e);
}
}
/**
* 组装 OpenAI 格式的早退/兜底流:内容 chunk + finish_reason=stop + [DONE]。
* 组装 OpenAI 格式的早退/兜底流:内容 + 空 sources + finish_reason=stop + [DONE]。
* 用于熔断降级、FAQ 命中与错误兜底三种场景。
*
* @param id chunk 唯一 ID
@ -608,6 +629,7 @@ public class AssistantApp {
private Flux<String> openAiFallbackStream(String id, String model, long created, String content, boolean withRole) {
return Flux.just(
buildOpenAiChunk(id, model, created, content, withRole, null),
buildSourcesChunk(id, model, created, List.of()),
buildOpenAiChunk(id, model, created, "", false, "stop"),
"[DONE]");
}

10
src/main/java/com/wok/supportbot/app/ChatResult.java

@ -17,17 +17,15 @@ import java.util.List;
* @param text AI 回答文本
* @param mcpEvents 本次触发的 MCP 工具调用事件,无调用时为空列表
* @param suggestions AI 推荐问题列表(0~3 条),非 LLM 路径为空
* @param sources 本次答案实际使用的知识库片段,不进行附加检索
*/
public record ChatResult(String text, List<McpToolCallback.ToolCallEvent> mcpEvents, List<String> suggestions) {
/** 向后兼容构造器(无 suggestions) */
public ChatResult(String text, List<McpToolCallback.ToolCallEvent> mcpEvents) {
this(text, mcpEvents, List.of());
}
public record ChatResult(String text, List<McpToolCallback.ToolCallEvent> mcpEvents,
List<String> suggestions, List<SourceReference> sources) {
/** 紧凑构造器:保证不可变性 */
public ChatResult {
suggestions = suggestions != null ? List.copyOf(suggestions) : List.of();
mcpEvents = mcpEvents != null ? List.copyOf(mcpEvents) : List.of();
sources = sources != null ? List.copyOf(sources) : List.of();
}
}

59
src/main/java/com/wok/supportbot/app/SourceReference.java

@ -0,0 +1,59 @@
package com.wok.supportbot.app;
import com.fasterxml.jackson.annotation.JsonInclude;
import org.springframework.ai.document.Document;
import java.util.List;
import java.util.Map;
/** Public citation metadata from the documents actually used by this answer. */
@JsonInclude(JsonInclude.Include.ALWAYS)
public record SourceReference(String documentId, String title, String sourceName,
Integer chunkIndex, Double score, String snippet) {
public SourceReference {
if (snippet != null && snippet.length() > 160) {
int end = Character.isHighSurrogate(snippet.charAt(158)) ? 158 : 159;
snippet = snippet.substring(0, end) + "…";
}
}
public static List<SourceReference> fromDocuments(List<Document> documents) {
if (documents == null || documents.isEmpty()) {
return List.of();
}
return documents.stream().map(SourceReference::fromDocument).toList();
}
private static SourceReference fromDocument(Document document) {
Map<String, Object> metadata = document.getMetadata();
// Preserve the published source endpoint's distance semantics; other retrievers expose score.
Object score = metadata.get("distance") != null ? metadata.get("distance") : metadata.get("score");
return new SourceReference(stringValue(metadata.get("documentId")),
stringValue(metadata.get("title")), stringValue(metadata.get("sourceName")),
integerValue(metadata.get("chunkIndex")), doubleValue(score), document.getText());
}
private static String stringValue(Object value) {
return value == null ? null : value.toString();
}
private static Integer integerValue(Object value) {
if (value == null) return null;
try {
return Integer.valueOf(value.toString());
} catch (NumberFormatException ignored) {
return null;
}
}
private static Double doubleValue(Object value) {
if (value == null) return null;
try {
double number = Double.parseDouble(value.toString());
return Double.isFinite(number) ? number : null;
} catch (NumberFormatException ignored) {
return null;
}
}
}

73
src/main/java/com/wok/supportbot/controller/AiController.java

@ -3,6 +3,8 @@ package com.wok.supportbot.controller;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.ChatPipeline;
import com.wok.supportbot.app.ChatResult;
import com.wok.supportbot.app.SourceReference;
import com.wok.supportbot.app.SuggestionGenerator;
import com.wok.supportbot.cache.SuggestionCache;
import com.wok.supportbot.config.RoleAccessConfig;
@ -12,7 +14,6 @@ import com.wok.supportbot.service.CustomerServiceRoleService;
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope;
import jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.document.Document;
import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity;
import org.springframework.util.StringUtils;
@ -28,10 +29,8 @@ import reactor.core.publisher.Flux;
import java.net.URLDecoder;
import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.Arrays;
import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
@ -122,21 +121,7 @@ public class AiController {
List<Long> cats = resolveCategoryIds(scope, categoryId, categoryIds);
ChatContext ctx = new ChatContext(message, chatId, "CHAT", null, null, cats,
normalizeStrategy(rewriteStrategy), true, false, context.roleId(), scope.name(), context.accountId(), null, null);
List<Document> docs = assistantApp.retrieveSources(ctx);
List<Map<String, Object>> out = new ArrayList<>();
for (Document doc : docs) {
Map<String, Object> meta = doc.getMetadata();
Map<String, Object> item = new LinkedHashMap<>();
item.put("documentId", meta.get("documentId"));
item.put("title", meta.get("title"));
item.put("sourceName", meta.get("sourceName"));
item.put("chunkIndex", meta.get("chunkIndex"));
item.put("score", meta.get("distance"));
String text = doc.getText();
item.put("snippet", text != null && text.length() > 160 ? text.substring(0, 160) + "…" : text);
out.add(item);
}
return Map.of("success", true, "data", out);
return Map.of("success", true, "data", SourceReference.fromDocuments(assistantApp.retrieveSources(ctx)));
} catch (Exception e) {
log.error("获取 RAG 引用来源失败 [strategy={}]: {}", rewriteStrategy, e.getMessage(), e);
return Map.of("success", true, "data", List.of());
@ -286,9 +271,9 @@ public class AiController {
return roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty();
}
/** 未指定策略时默认 MULTI_QUERY(多路扩展)。 */
/** 未指定策略时使用原始问题检索,不调用重写模型。 */
private String normalizeStrategy(String rewriteStrategy) {
return (rewriteStrategy != null && !rewriteStrategy.isEmpty()) ? rewriteStrategy : "MULTI_QUERY";
return StringUtils.hasText(rewriteStrategy) ? rewriteStrategy : "NONE";
}
private AccountRoleContext resolveAccountRole(String accountId, Long fallbackRoleId) {
@ -344,15 +329,35 @@ public class AiController {
@RequestParam(required = false) Long categoryId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String imageUrls) {
ChatContext ctx;
if (Boolean.TRUE.equals(enableRag)) {
ctx = buildRagChatContext(message, chatId, rewriteStrategy, roleId, accountId,
categoryId, categoryIds, systemPrompt);
} else {
ctx = buildChatContext(message, chatId, roleId, accountId, systemPrompt);
}
ctx = ctx.withImageUrls(parseImageUrls(imageUrls));
return assistantApp.chat(ctx);
return chatResult(message, chatId, roleId, accountId, systemPrompt, enableRag,
rewriteStrategy, categoryId, categoryIds, imageUrls).text();
}
/** 同步完整结果;与文本接口共享角色隔离、会话绑定和图片处理。 */
@GetMapping(value = "/chat/result", produces = MediaType.APPLICATION_JSON_VALUE)
public ChatResult chatResult(
@RequestParam String message,
@RequestParam(required = false) String chatId,
@RequestParam(required = false) Long roleId,
@RequestParam(required = false) String accountId,
@RequestParam(required = false) String systemPrompt,
@RequestParam(required = false) Boolean enableRag,
@RequestParam(required = false) String rewriteStrategy,
@RequestParam(required = false) Long categoryId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String imageUrls) {
return assistantApp.chatWithEvents(buildStandardChatContext(message, chatId, roleId, accountId,
systemPrompt, enableRag, rewriteStrategy, categoryId, categoryIds, imageUrls));
}
private ChatContext buildStandardChatContext(String message, String chatId, Long roleId, String accountId,
String systemPrompt, Boolean enableRag, String rewriteStrategy, Long categoryId,
String categoryIds, String imageUrls) {
ChatContext ctx = Boolean.TRUE.equals(enableRag)
? buildRagChatContext(message, chatId, rewriteStrategy, roleId, accountId,
categoryId, categoryIds, systemPrompt)
: buildChatContext(message, chatId, roleId, accountId, systemPrompt);
return ctx.withImageUrls(parseImageUrls(imageUrls));
}
/**
@ -371,14 +376,8 @@ public class AiController {
@RequestParam(required = false) Long categoryId,
@RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String imageUrls) {
ChatContext ctx;
if (Boolean.TRUE.equals(enableRag)) {
ctx = buildRagChatContext(message, chatId, rewriteStrategy, roleId, accountId,
categoryId, categoryIds, systemPrompt);
} else {
ctx = buildChatContext(message, chatId, roleId, accountId, systemPrompt);
}
ctx = ctx.withImageUrls(parseImageUrls(imageUrls)).withStreaming(true);
ChatContext ctx = buildStandardChatContext(message, chatId, roleId, accountId, systemPrompt,
enableRag, rewriteStrategy, categoryId, categoryIds, imageUrls).withStreaming(true);
return assistantApp.chatStreamOpenAi(ctx);
}

14
src/main/java/com/wok/supportbot/controller/OpenApiController.java

@ -2,6 +2,7 @@ package com.wok.supportbot.controller;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.ChatResult;
import com.wok.supportbot.app.SuggestionGenerator;
import com.wok.supportbot.cache.SuggestionCache;
import com.wok.supportbot.config.RoleAccessConfig;
@ -84,13 +85,16 @@ public class OpenApiController {
ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId,
categoryIds, rewriteStrategy, enableRag, false);
String reply = assistantApp.chat(ctx);
ChatResult reply = assistantApp.chatWithEvents(ctx);
Map<String, Object> result = new HashMap<>();
result.put("success", true);
result.put("data", Map.of(
"reply", reply,
"chatId", resolvedChatId
"reply", reply.text(),
"chatId", resolvedChatId,
"mcpEvents", reply.mcpEvents(),
"suggestions", reply.suggestions(),
"sources", reply.sources()
));
return ResponseEntity.ok(result);
} catch (Exception e) {
@ -219,9 +223,9 @@ public class OpenApiController {
boolean useRag = (enableRag == null || enableRag) // 默认启用 RAG
&& !(roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty());
// 6. 未指定策略时默认 MULTI_QUERY
// 6. 未指定策略时直接使用原始问题检索
String strategy = (rewriteStrategy != null && !rewriteStrategy.isBlank())
? rewriteStrategy : "MULTI_QUERY";
? rewriteStrategy : "NONE";
return new ChatContext(message, chatId, "CHAT", systemPrompt,
scope.hasRole() ? scope.allowedMcpTools() : null,

250
src/test/java/com/wok/supportbot/AnswerTransportTests.java

@ -0,0 +1,250 @@
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))));
}
}

166
src/test/java/com/wok/supportbot/ChatResultEndpointTests.java

@ -0,0 +1,166 @@
package com.wok.supportbot;
import com.wok.supportbot.app.AssistantApp;
import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.ChatResult;
import com.wok.supportbot.app.SourceReference;
import com.wok.supportbot.config.RoleAccessConfig;
import com.wok.supportbot.controller.AiController;
import com.wok.supportbot.controller.OpenApiController;
import com.wok.supportbot.entity.ApiKey;
import com.wok.supportbot.rag.CategoryFilter;
import com.wok.supportbot.security.JwtTokenProvider;
import com.wok.supportbot.security.SdkAuthFilter;
import com.wok.supportbot.security.SdkJwtTokenProvider;
import com.wok.supportbot.service.ConversationService;
import com.wok.supportbot.service.CustomerServiceRoleService;
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope;
import io.jsonwebtoken.Claims;
import org.junit.jupiter.api.BeforeEach;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.jdbc.core.JdbcTemplate;
import org.springframework.mock.web.MockHttpServletRequest;
import org.springframework.test.util.ReflectionTestUtils;
import org.springframework.test.web.servlet.MockMvc;
import org.springframework.test.web.servlet.setup.MockMvcBuilders;
import java.util.List;
import java.util.Set;
import static org.junit.jupiter.api.Assertions.*;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.Mockito.*;
import static org.springframework.test.web.servlet.request.MockMvcRequestBuilders.get;
import static org.springframework.test.web.servlet.result.MockMvcResultMatchers.*;
class ChatResultEndpointTests {
private AssistantApp assistant;
private CustomerServiceRoleService roles;
private ConversationService conversations;
private RoleAccessConfig access;
private AiController controller;
private MockMvc mvc;
private SdkJwtTokenProvider tokens;
@BeforeEach
void setUp() {
assistant = mock(AssistantApp.class);
roles = mock(CustomerServiceRoleService.class);
conversations = mock(ConversationService.class);
access = new RoleAccessConfig();
controller = new AiController();
ReflectionTestUtils.setField(controller, "assistantApp", assistant);
ReflectionTestUtils.setField(controller, "customerServiceRoleService", roles);
ReflectionTestUtils.setField(controller, "conversationService", conversations);
ReflectionTestUtils.setField(controller, "roleAccessConfig", access);
ReflectionTestUtils.setField(controller, "categoryFilter", mock(CategoryFilter.class));
tokens = mock(SdkJwtTokenProvider.class);
SdkAuthFilter filter = new SdkAuthFilter(tokens, mock(JwtTokenProvider.class), mock(JdbcTemplate.class));
mvc = MockMvcBuilders.standaloneSetup(controller).addFilters(filter).build();
}
@Test
void resultRejectsMissingTokenAndUnauthorizedRoleBeforeGeneration() throws Exception {
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result").param("message", "问题"))
.andExpect(status().isUnauthorized());
authorize();
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result")
.header("Authorization", "Bearer valid").param("message", "问题").param("roleId", "99"))
.andExpect(status().isForbidden());
verifyNoInteractions(assistant, roles, conversations);
}
@Test
void resultKeepsRoleScopeAccountBindingAndDecodedImages() throws Exception {
authorize();
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "授权人设", List.of(7L), List.of("tool")));
SourceReference source = new SourceReference("9223372036854775807", "授权文档", null, null, null, "片段");
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of(source)));
mvc.perform(get("/ai/chat/result").servletPath("/ai/chat/result")
.header("Authorization", "Bearer valid").param("message", "问题").param("chatId", "chat")
.param("roleId", "11").param("accountId", " account ").param("systemPrompt", "越权人设")
.param("enableRag", "true").param("categoryIds", "99").param("imageUrls", "https%3A%2F%2Fexample.org%2Fa.png"))
.andExpect(status().isOk()).andExpect(content().contentTypeCompatibleWith("application/json"))
.andExpect(jsonPath("$.text").value("回答"))
.andExpect(jsonPath("$.data").doesNotExist())
.andExpect(jsonPath("$.sources[0].documentId").value("9223372036854775807"))
.andExpect(jsonPath("$.mcpEvents").isArray()).andExpect(jsonPath("$.suggestions").isArray());
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class);
verify(assistant).chatWithEvents(context.capture());
ChatContext ctx = context.getValue();
assertEquals(List.of(7L), ctx.categoryIds());
assertEquals("授权人设", ctx.systemPrompt());
assertEquals(List.of("tool"), ctx.allowedMcpTools());
assertEquals(List.of("https://example.org/a.png"), ctx.imageUrls());
assertEquals("NONE", ctx.rewriteStrategy());
assertTrue(ctx.enableRag());
verify(conversations).bindConversation("chat", "account", 11L);
verifyNoMoreInteractions(assistant);
}
@Test
void strictIsolationAndExplicitRewriteApplyToBothSyncEndpoints() {
access.setStrictIsolation(true);
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(), List.of()));
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of()));
assertEquals("回答", controller.chatSync("问题", "chat", 11L, "account", null, true,
"REWRITE", null, "99", null));
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class);
verify(assistant).chatWithEvents(context.capture());
assertFalse(context.getValue().enableRag());
assertTrue(context.getValue().categoryIds().isEmpty());
assertEquals("REWRITE", context.getValue().rewriteStrategy());
verifyNoMoreInteractions(assistant);
}
@Test
void openApiRetainsEnvelopeAndAddsSameAnswerSourcesWithNoneDefault() {
OpenApiController open = new OpenApiController();
ReflectionTestUtils.setField(open, "assistantApp", assistant);
ReflectionTestUtils.setField(open, "customerServiceRoleService", roles);
ReflectionTestUtils.setField(open, "categoryFilter", mock(CategoryFilter.class));
ReflectionTestUtils.setField(open, "roleAccessConfig", access);
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(7L), List.of()));
SourceReference source = new SourceReference("9223372036854775807", "授权文档", null, null, null, "片段");
when(assistant.chatWithEvents(any())).thenReturn(new ChatResult("回答", List.of(), List.of(), List.of(source)));
ApiKey key = new ApiKey();
key.setId(3L);
key.setRoleIds("[11]");
MockHttpServletRequest request = new MockHttpServletRequest();
request.setAttribute("apiKey", key);
var result = open.chat("问题", "11", "chat", "99", null, true, request);
assertEquals(200, result.getStatusCode().value());
assertEquals(true, result.getBody().get("success"));
var data = (java.util.Map<?, ?>) result.getBody().get("data");
assertEquals("回答", data.get("reply"));
assertEquals(List.of(source), data.get("sources"));
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class);
verify(assistant).chatWithEvents(context.capture());
assertEquals("NONE", context.getValue().rewriteStrategy());
assertEquals(List.of(7L), context.getValue().categoryIds());
verifyNoMoreInteractions(assistant);
}
@Test
void explicitSourcesApiUsesTheSameCitationSerializer() {
ReflectionTestUtils.setField(controller, "chatPipeline", mock(com.wok.supportbot.app.ChatPipeline.class));
when(roles.getRoleScope(11L)).thenReturn(new RoleScope(true, "售后", "人设", List.of(7L), List.of()));
var document = new org.springframework.ai.document.Document("片段".repeat(100),
java.util.Map.of("documentId", Long.MAX_VALUE, "distance", 0.2));
when(assistant.retrieveSources(any())).thenReturn(List.of(document));
var result = controller.chatSources("问题", "chat", null, 11L, "account", null, "99");
assertEquals(SourceReference.fromDocuments(List.of(document)), result.get("data"));
ArgumentCaptor<ChatContext> context = ArgumentCaptor.forClass(ChatContext.class);
verify(assistant).retrieveSources(context.capture());
assertEquals(List.of(7L), context.getValue().categoryIds());
assertEquals("NONE", context.getValue().rewriteStrategy());
verifyNoMoreInteractions(assistant);
}
private void authorize() {
Claims claims = mock(Claims.class);
when(tokens.parseToken("valid")).thenReturn(claims);
when(tokens.getAllowedRoleIds(claims)).thenReturn(Set.of(11L));
}
}
Loading…
Cancel
Save