fix(chat): 统一单轮检索并完善多库上下文

This commit is contained in:
evo
2026-07-21 09:35:26 +08:00
parent 42bc8fea95
commit 9c092ae3bb

View File

@@ -76,6 +76,7 @@ import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
import java.util.ArrayList; import java.util.ArrayList;
import java.util.List; import java.util.List;
import java.util.Map; import java.util.Map;
import java.util.LinkedHashMap;
import java.util.concurrent.CompletableFuture; import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentHashMap;
@@ -160,7 +161,8 @@ public class ChatServiceFacade implements IChatService {
throw new IllegalArgumentException("模型不存在: " + chatRequest.getModel()); throw new IllegalArgumentException("模型不存在: " + chatRequest.getModel());
} }
// 2. 构建上下文消息列表 // 2. 构建上下文消息列表(系统提示词 + 历史消息 + 当前用户消息)
// 注意RAG 检索增强统一在 handleAgentChat 中执行一次,此处不再重复检索
List<ChatMessage> contextMessages = buildContextMessages(chatRequest, agentVo); List<ChatMessage> contextMessages = buildContextMessages(chatRequest, agentVo);
chatRequest.setEmitter(emitter); chatRequest.setEmitter(emitter);
@@ -274,12 +276,20 @@ public class ChatServiceFacade implements IChatService {
.responseStrategy(SupervisorResponseStrategy.LAST); .responseStrategy(SupervisorResponseStrategy.LAST);
SupervisorAgent supervisor = supervisorBuilder.build(); SupervisorAgent supervisor = supervisorBuilder.build();
// 知识库增强:智能体绑定了知识库时,对 supervisor 输入做一次 RAG 增强 // 知识库增强:智能体绑定了知识库时,对 supervisor 输入做一次 RAG 增强(全程唯一一次检索)
String augmentedInput = augmentAgentInput(chatRequest, agentVo); String augmentedInput = augmentAgentInput(chatRequest, agentVo);
// 智能体自定义系统提示词supervisor builder 不支持 systemMessage前置到输入 // 组装最终 prompt系统提示词 → 多轮历史 → RAG 增强后的当前提问
String prompt = (agentVo != null && StringUtils.isNotBlank(agentVo.getSystemPrompt())) StringBuilder promptBuilder = new StringBuilder();
? agentVo.getSystemPrompt() + "\n\n" + augmentedInput if (agentVo != null && StringUtils.isNotBlank(agentVo.getSystemPrompt())) {
: augmentedInput; promptBuilder.append(agentVo.getSystemPrompt()).append("\n\n");
}
String historyText = formatHistoryMessages(chatRequest.getContextMessages(), chatRequest.getContent());
if (StringUtils.isNotBlank(historyText)) {
promptBuilder.append("以下是本次会话的历史对话,请结合上下文理解用户最新提问:\n")
.append(historyText).append("\n\n");
}
promptBuilder.append(augmentedInput);
String prompt = promptBuilder.toString();
String tokenValue = chatRequest.getTokenValue(); String tokenValue = chatRequest.getTokenValue();
@@ -308,8 +318,9 @@ public class ChatServiceFacade implements IChatService {
* 兜底 MCP 工具装配(无智能体时使用,保留原有 3 个硬编码客户端逻辑) * 兜底 MCP 工具装配(无智能体时使用,保留原有 3 个硬编码客户端逻辑)
*/ */
private ToolProvider buildDefaultMcpToolProvider(Long userId) { private ToolProvider buildDefaultMcpToolProvider(Long userId) {
String npxCommand = resolveNpxCommand();
McpTransport playwrightTransport = new StdioMcpTransport.Builder() McpTransport playwrightTransport = new StdioMcpTransport.Builder()
.command(List.of("C:\\Program Files\\nodejs\\npx.cmd", "-y", "@playwright/mcp@latest")) .command(List.of(npxCommand, "-y", "@playwright/mcp@latest"))
.logEvents(true) .logEvents(true)
.build(); .build();
McpClient playwrightMcpClient = new DefaultMcpClient.Builder() McpClient playwrightMcpClient = new DefaultMcpClient.Builder()
@@ -319,7 +330,7 @@ public class ChatServiceFacade implements IChatService {
String userDir = System.getProperty("user.dir"); String userDir = System.getProperty("user.dir");
McpTransport filesystemTransport = new StdioMcpTransport.Builder() McpTransport filesystemTransport = new StdioMcpTransport.Builder()
.command(List.of("C:\\Program Files\\nodejs\\npx.cmd", "-y", .command(List.of(npxCommand, "-y",
"@modelcontextprotocol/server-filesystem", userDir)) "@modelcontextprotocol/server-filesystem", userDir))
.logEvents(true) .logEvents(true)
.build(); .build();
@@ -333,6 +344,14 @@ public class ChatServiceFacade implements IChatService {
.build(); .build();
} }
private String resolveNpxCommand() {
String configured = System.getProperty("mcp.npx.command");
if (StringUtils.isNotBlank(configured)) return configured;
String fromEnv = System.getenv("MCP_NPX_COMMAND");
if (StringUtils.isNotBlank(fromEnv)) return fromEnv;
return System.getProperty("os.name", "").toLowerCase().contains("win") ? "npx.cmd" : "npx";
}
/** /**
* 装配磁盘 ShellSkills智能体勾选了技能名时按名过滤否则加载全部。 * 装配磁盘 ShellSkills智能体勾选了技能名时按名过滤否则加载全部。
* 无 skills 时返回 null调用方据此跳过 SkillsAgent 的 toolProvider 注入) * 无 skills 时返回 null调用方据此跳过 SkillsAgent 的 toolProvider 注入)
@@ -468,27 +487,7 @@ public class ChatServiceFacade implements IChatService {
messages.add(SystemMessage.from(agentVo.getSystemPrompt())); messages.add(SystemMessage.from(agentVo.getSystemPrompt()));
} }
// 1. 初始化当前用户消息 // 1. 从数据库查询历史对话消息(放在前面)
UserMessage userMessage = UserMessage.userMessage(chatRequest.getContent());
// 2. 知识库检索增强 (RAG):智能体的 knowledgeIds 优先,回退到请求的 knowledgeId
List<Long> knowledgeIds = collectKnowledgeIds(chatRequest, agentVo);
if (knowledgeIds != null && !knowledgeIds.isEmpty()) {
RetrievalAugmentor augmentor = buildMultiKnowledgeAugmentor(knowledgeIds);
if (augmentor != null) {
log.info("执行多知识库 RAG 流程: kids={}", knowledgeIds);
Metadata metadata = Metadata.from(userMessage, chatRequest.getSessionId(), new ArrayList<>());
AugmentationRequest augmentationRequest = new AugmentationRequest(userMessage, metadata);
AugmentationResult result = augmentor.augment(augmentationRequest);
ChatMessage augmented = result.chatMessage();
if (augmented instanceof UserMessage) {
userMessage = (UserMessage) augmented;
log.debug("RAG 增强完成UserMessage 已注入背景知识");
}
}
}
// 3. 从数据库查询历史对话消息(放在前面)
if (chatRequest.getSessionId() != null) { if (chatRequest.getSessionId() != null) {
MessageWindowChatMemory memory = createChatMemory(chatRequest.getSessionId()); MessageWindowChatMemory memory = createChatMemory(chatRequest.getSessionId());
if (memory != null) { if (memory != null) {
@@ -500,12 +499,37 @@ public class ChatServiceFacade implements IChatService {
} }
} }
// 4. 添加经过增强的用户消息(放在最后) // 2. 添加当前用户消息(放在最后RAG 增强在 handleAgentChat 中统一执行,避免重复检索
messages.add(userMessage); messages.add(UserMessage.userMessage(chatRequest.getContent()));
return messages; return messages;
} }
/**
* 将上下文消息格式化为多轮对话文本(供只接受 String 输入的 Supervisor 使用)。
* 跳过 SystemMessage系统提示词单独前置与最后一条当前用户消息单独做 RAG 增强后拼接)。
*/
private String formatHistoryMessages(List<ChatMessage> contextMessages, String currentContent) {
if (contextMessages == null || contextMessages.isEmpty()) {
return "";
}
StringBuilder sb = new StringBuilder();
int limit = contextMessages.size();
// 最后一条是当前用户消息,不纳入历史(避免与增强后的输入重复)
if (limit > 0 && contextMessages.get(limit - 1) instanceof UserMessage) {
limit--;
}
for (int i = 0; i < limit; i++) {
ChatMessage msg = contextMessages.get(i);
if (msg instanceof UserMessage userMsg) {
sb.append("用户: ").append(userMsg.singleText()).append("\n");
} else if (msg instanceof AiMessage aiMsg) {
sb.append("助手: ").append(aiMsg.text()).append("\n");
}
}
return sb.toString().trim();
}
/** /**
* 汇总本次对话要检索的知识库ID列表智能体绑定的 knowledgeIds 优先,回退到请求的 knowledgeId * 汇总本次对话要检索的知识库ID列表智能体绑定的 knowledgeIds 优先,回退到请求的 knowledgeId
*/ */
@@ -572,18 +596,35 @@ public class ChatServiceFacade implements IChatService {
@Override @Override
public List<Content> retrieve(Query query) { public List<Content> retrieve(Query query) {
List<Content> all = new ArrayList<>(); List<CompletableFuture<List<Content>>> futures = delegates.stream()
for (ContentRetriever r : delegates) { .map(r -> CompletableFuture.supplyAsync(() -> {
try { try {
List<Content> part = r.retrieve(query); List<Content> part = r.retrieve(query);
if (part != null) { return part == null ? List.<Content>of() : part;
all.addAll(part);
}
} catch (Exception e) { } catch (Exception e) {
log.warn("复合检索子检索器异常: {}", e.getMessage()); log.warn("复合检索子检索器异常: {}", e.getMessage());
return List.<Content>of();
}
})).toList();
Map<String, Content> unique = new LinkedHashMap<>();
for (CompletableFuture<List<Content>> future : futures) {
for (Content content : future.join()) {
String key = content.textSegment().metadata().getString("kid") + "|"
+ content.textSegment().metadata().getString("docId") + "|"
+ content.textSegment().metadata().getString("fid");
if (key.endsWith("null|null|null")) key = content.textSegment().text();
unique.putIfAbsent(key, content);
} }
} }
return all; List<Content> bounded = new ArrayList<>();
int chars = 0;
for (Content content : unique.values()) {
int next = content.textSegment().text().length();
if (bounded.size() >= 20 || chars + next > 24000) break;
bounded.add(content);
chars += next;
}
return bounded;
} }
} }