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

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

@ -17,17 +17,15 @@ import java.util.List;
* @param text AI 回答文本 * @param text AI 回答文本
* @param mcpEvents 本次触发的 MCP 工具调用事件,无调用时为空列表 * @param mcpEvents 本次触发的 MCP 工具调用事件,无调用时为空列表
* @param suggestions AI 推荐问题列表(0~3 条),非 LLM 路径为空 * @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 { public ChatResult {
suggestions = suggestions != null ? List.copyOf(suggestions) : List.of(); suggestions = suggestions != null ? List.copyOf(suggestions) : List.of();
mcpEvents = mcpEvents != null ? List.copyOf(mcpEvents) : 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.AssistantApp;
import com.wok.supportbot.app.ChatContext; import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.ChatPipeline; 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.app.SuggestionGenerator;
import com.wok.supportbot.cache.SuggestionCache; import com.wok.supportbot.cache.SuggestionCache;
import com.wok.supportbot.config.RoleAccessConfig; import com.wok.supportbot.config.RoleAccessConfig;
@ -12,7 +14,6 @@ import com.wok.supportbot.service.CustomerServiceRoleService;
import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope; import com.wok.supportbot.service.CustomerServiceRoleService.RoleScope;
import jakarta.annotation.Resource; import jakarta.annotation.Resource;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.document.Document;
import org.springframework.http.MediaType; import org.springframework.http.MediaType;
import org.springframework.http.ResponseEntity; import org.springframework.http.ResponseEntity;
import org.springframework.util.StringUtils; import org.springframework.util.StringUtils;
@ -28,10 +29,8 @@ import reactor.core.publisher.Flux;
import java.net.URLDecoder; import java.net.URLDecoder;
import java.nio.charset.StandardCharsets; import java.nio.charset.StandardCharsets;
import java.util.ArrayList;
import java.util.Arrays; import java.util.Arrays;
import java.util.Collections; import java.util.Collections;
import java.util.LinkedHashMap;
import java.util.List; import java.util.List;
import java.util.Map; import java.util.Map;
@ -122,21 +121,7 @@ public class AiController {
List<Long> cats = resolveCategoryIds(scope, categoryId, categoryIds); List<Long> cats = resolveCategoryIds(scope, categoryId, categoryIds);
ChatContext ctx = new ChatContext(message, chatId, "CHAT", null, null, cats, ChatContext ctx = new ChatContext(message, chatId, "CHAT", null, null, cats,
normalizeStrategy(rewriteStrategy), true, false, context.roleId(), scope.name(), context.accountId(), null, null); 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) { } catch (Exception e) {
log.error("获取 RAG 引用来源失败 [strategy={}]: {}", rewriteStrategy, e.getMessage(), e); log.error("获取 RAG 引用来源失败 [strategy={}]: {}", rewriteStrategy, e.getMessage(), e);
return Map.of("success", true, "data", List.of()); return Map.of("success", true, "data", List.of());
@ -286,9 +271,9 @@ public class AiController {
return roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty(); return roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty();
} }
/** 未指定策略时默认 MULTI_QUERY(多路扩展)。 */
/** 未指定策略时使用原始问题检索,不调用重写模型。 */
private String normalizeStrategy(String rewriteStrategy) { 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) { private AccountRoleContext resolveAccountRole(String accountId, Long fallbackRoleId) {
@ -344,15 +329,35 @@ public class AiController {
@RequestParam(required = false) Long categoryId, @RequestParam(required = false) Long categoryId,
@RequestParam(required = false) String categoryIds, @RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String imageUrls) { @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) Long categoryId,
@RequestParam(required = false) String categoryIds, @RequestParam(required = false) String categoryIds,
@RequestParam(required = false) String imageUrls) { @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); 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.AssistantApp;
import com.wok.supportbot.app.ChatContext; import com.wok.supportbot.app.ChatContext;
import com.wok.supportbot.app.ChatResult;
import com.wok.supportbot.app.SuggestionGenerator; import com.wok.supportbot.app.SuggestionGenerator;
import com.wok.supportbot.cache.SuggestionCache; import com.wok.supportbot.cache.SuggestionCache;
import com.wok.supportbot.config.RoleAccessConfig; import com.wok.supportbot.config.RoleAccessConfig;
@ -84,13 +85,16 @@ public class OpenApiController {
ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId, ChatContext ctx = buildOpenApiChatContext(message, resolvedChatId, apiKey, roleId,
categoryIds, rewriteStrategy, enableRag, false); categoryIds, rewriteStrategy, enableRag, false);
String reply = assistantApp.chat(ctx);
ChatResult reply = assistantApp.chatWithEvents(ctx);
Map<String, Object> result = new HashMap<>(); Map<String, Object> result = new HashMap<>();
result.put("success", true); result.put("success", true);
result.put("data", Map.of( 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); return ResponseEntity.ok(result);
} catch (Exception e) { } catch (Exception e) {
@ -219,9 +223,9 @@ public class OpenApiController {
boolean useRag = (enableRag == null || enableRag) // 默认启用 RAG boolean useRag = (enableRag == null || enableRag) // 默认启用 RAG
&& !(roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty()); && !(roleAccessConfig.isStrictIsolation() && scope.hasRole() && scope.categoryIds().isEmpty());
// 6. 未指定策略时默认 MULTI_QUERY
// 6. 未指定策略时直接使用原始问题检索
String strategy = (rewriteStrategy != null && !rewriteStrategy.isBlank()) String strategy = (rewriteStrategy != null && !rewriteStrategy.isBlank())
? rewriteStrategy : "MULTI_QUERY";
? rewriteStrategy : "NONE";
return new ChatContext(message, chatId, "CHAT", systemPrompt, return new ChatContext(message, chatId, "CHAT", systemPrompt,
scope.hasRole() ? scope.allowedMcpTools() : null, 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