mirror of
https://gitcode.com/ageerle/ruoyi-ai.git
synced 2026-09-13 00:14:59 +00:00
fix(chat): 统一单轮检索并完善多库上下文
This commit is contained in:
@@ -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;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user