From 28e77223efa6b9fcc7e8e6fbeccbad71e1d59dd1 Mon Sep 17 00:00:00 2001 From: wei-py Date: Mon, 14 Sep 2026 14:13:29 +0800 Subject: [PATCH] =?UTF-8?q?feat(chat):=20=E5=90=8C=E6=AC=A1=E5=9B=9E?= =?UTF-8?q?=E7=AD=94=E6=90=BA=E5=B8=A6=E5=BC=95=E7=94=A8=E6=9D=A5=E6=BA=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../com/wok/supportbot/app/AssistantApp.java | 50 +++- .../com/wok/supportbot/app/ChatResult.java | 10 +- .../wok/supportbot/app/SourceReference.java | 59 +++++ .../supportbot/controller/AiController.java | 73 +++-- .../controller/OpenApiController.java | 14 +- .../wok/supportbot/AnswerTransportTests.java | 250 ++++++++++++++++++ .../supportbot/ChatResultEndpointTests.java | 166 ++++++++++++ 7 files changed, 560 insertions(+), 62 deletions(-) create mode 100644 src/main/java/com/wok/supportbot/app/SourceReference.java create mode 100644 src/test/java/com/wok/supportbot/AnswerTransportTests.java create mode 100644 src/test/java/com/wok/supportbot/ChatResultEndpointTests.java diff --git a/src/main/java/com/wok/supportbot/app/AssistantApp.java b/src/main/java/com/wok/supportbot/app/AssistantApp.java index a3a4010..0660352 100644 --- a/src/main/java/com/wok/supportbot/app/AssistantApp.java +++ b/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 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 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 { *

* 复用 {@link #chatStream(ChatContext)} 的完整编排逻辑(熔断早退 / FAQ 命中早退 / * 正常流式调用 / 空白缓冲 / 埋点),差异在于把每个文本片段包装为 OpenAI 标准 JSON chunk: - * 首片 delta 携带 role=assistant,流结束时追加 finish_reason=stop 的 chunk 与 [DONE]。 + * 首片 delta 携带 role=assistant,正文结束后追加 sources 元数据、finish_reason=stop 与 [DONE]。 *

* 每个 Flux 元素即一个完整 JSON 字符串,Spring WebFlux 自动加 data: 前缀。 * @@ -460,6 +461,10 @@ public class AssistantApp { * @return OpenAI 标准格式的流式回答 */ public Flux chatStreamOpenAi(ChatContext ctx) { + return Flux.defer(() -> buildOpenAiStream(ctx)).subscribeOn(Schedulers.boundedElastic()); + } + + private Flux 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 sources) { + Map 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 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]"); } diff --git a/src/main/java/com/wok/supportbot/app/ChatResult.java b/src/main/java/com/wok/supportbot/app/ChatResult.java index dc5e5ab..a13b860 100644 --- a/src/main/java/com/wok/supportbot/app/ChatResult.java +++ b/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 mcpEvents, List suggestions) { - - /** 向后兼容构造器(无 suggestions) */ - public ChatResult(String text, List mcpEvents) { - this(text, mcpEvents, List.of()); - } +public record ChatResult(String text, List mcpEvents, + List suggestions, List 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(); } } diff --git a/src/main/java/com/wok/supportbot/app/SourceReference.java b/src/main/java/com/wok/supportbot/app/SourceReference.java new file mode 100644 index 0000000..24ec971 --- /dev/null +++ b/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 fromDocuments(List documents) { + if (documents == null || documents.isEmpty()) { + return List.of(); + } + return documents.stream().map(SourceReference::fromDocument).toList(); + } + + private static SourceReference fromDocument(Document document) { + Map 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; + } + } +} diff --git a/src/main/java/com/wok/supportbot/controller/AiController.java b/src/main/java/com/wok/supportbot/controller/AiController.java index bcdc1c9..804f011 100644 --- a/src/main/java/com/wok/supportbot/controller/AiController.java +++ b/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 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 docs = assistantApp.retrieveSources(ctx); - List> out = new ArrayList<>(); - for (Document doc : docs) { - Map meta = doc.getMetadata(); - Map 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); } diff --git a/src/main/java/com/wok/supportbot/controller/OpenApiController.java b/src/main/java/com/wok/supportbot/controller/OpenApiController.java index cdc1a07..41f0ed2 100644 --- a/src/main/java/com/wok/supportbot/controller/OpenApiController.java +++ b/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 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, diff --git a/src/test/java/com/wok/supportbot/AnswerTransportTests.java b/src/test/java/com/wok/supportbot/AnswerTransportTests.java new file mode 100644 index 0000000..5334007 --- /dev/null +++ b/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 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)))); + } +} diff --git a/src/test/java/com/wok/supportbot/ChatResultEndpointTests.java b/src/test/java/com/wok/supportbot/ChatResultEndpointTests.java new file mode 100644 index 0000000..7859158 --- /dev/null +++ b/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 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 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 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 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)); + } +}