refactor(phase1): 完成包名重构 (org.example → com.superbiz.agent)
Task 4.1: 包名统一重构 - 重命名 41 个 Java 文件的包名 - 更新所有 import 语句 - 恢复枚举类(FaultCategory、DiagnosisStatus、SourceType) - 更新测试类的 import 重构范围: - domain/entity: 3 个实体类 - domain/model: 2 个数据类 - domain/enums: 3 个枚举类 - repository: 3 个接口 - service/session: 2 个类(接口 + 实现) - config: 9 个配置类 - controller: 2 个控制器 - agent/tool: 4 个工具类 - client: 1 个客户端 - Main.java: 主类 验证结果: - 编译成功,无错误 - 所有测试通过 (27/27) - ApiDocumentRepositoryTest: 7/7 ✅ - CaseLibraryRepositoryTest: 6/6 ✅ - DiagnosisRecordRepositoryTest: 6/6 ✅ - RedisSessionManagerTest: 8/8 ✅ Progress: 21/33 tasks completed (64%)
This commit is contained in:
@@ -0,0 +1,289 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import com.alibaba.cloud.ai.graph.OverAllState;
|
||||
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
|
||||
import com.alibaba.cloud.ai.graph.agent.flow.agent.SupervisorAgent;
|
||||
import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
|
||||
import com.superbiz.agent.agent.tool.DateTimeTools;
|
||||
import com.superbiz.agent.agent.tool.InternalDocsTools;
|
||||
import com.superbiz.agent.agent.tool.QueryLogsTools;
|
||||
import com.superbiz.agent.agent.tool.QueryMetricsTools;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Optional;
|
||||
|
||||
/**
|
||||
* AI Ops 智能运维服务
|
||||
* 负责多 Agent 协作的告警分析流程
|
||||
*/
|
||||
@Service
|
||||
public class AiOpsService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(AiOpsService.class);
|
||||
|
||||
@Autowired
|
||||
private DateTimeTools dateTimeTools;
|
||||
|
||||
@Autowired
|
||||
private InternalDocsTools internalDocsTools;
|
||||
|
||||
@Autowired
|
||||
private QueryMetricsTools queryMetricsTools;
|
||||
|
||||
@Autowired(required = false) // Mock 模式下才注册
|
||||
private QueryLogsTools queryLogsTools;
|
||||
|
||||
/**
|
||||
* 执行 AI Ops 告警分析流程
|
||||
*
|
||||
* @param chatModel 大模型实例
|
||||
* @param toolCallbacks 工具回调数组
|
||||
* @return 分析结果状态
|
||||
* @throws GraphRunnerException 如果 Agent 执行失败
|
||||
*/
|
||||
public Optional<OverAllState> executeAiOpsAnalysis(ChatModel chatModel, ToolCallback[] toolCallbacks) throws GraphRunnerException {
|
||||
logger.info("开始执行 AI Ops 多 Agent 协作流程");
|
||||
|
||||
// 构建 Planner 和 Executor Agent
|
||||
ReactAgent plannerAgent = buildPlannerAgent(chatModel, toolCallbacks);
|
||||
ReactAgent executorAgent = buildExecutorAgent(chatModel, toolCallbacks);
|
||||
|
||||
// 构建 Supervisor Agent
|
||||
SupervisorAgent supervisorAgent = SupervisorAgent.builder()
|
||||
.name("ai_ops_supervisor")
|
||||
.description("负责调度 Planner 与 Executor 的多 Agent 控制器")
|
||||
.model(chatModel)
|
||||
.systemPrompt(buildSupervisorSystemPrompt())
|
||||
.subAgents(List.of(plannerAgent, executorAgent))
|
||||
.build();
|
||||
|
||||
String taskPrompt = "你是企业级 SRE,接到了自动化告警排查任务。请结合工具调用,执行**规划→执行→再规划**的闭环,并最终按照固定模板输出《告警分析报告》。禁止编造虚假数据,如连续多次查询失败需诚实反馈无法完成的原因。";
|
||||
|
||||
logger.info("调用 Supervisor Agent 开始编排...");
|
||||
|
||||
Optional<OverAllState> stateOptional = supervisorAgent.invoke(taskPrompt);
|
||||
|
||||
// 添加调试代码
|
||||
if (stateOptional.isPresent()) {
|
||||
OverAllState state = stateOptional.get();
|
||||
logger.debug("Final State Keys: {}", state.data().keySet()); // 打印所有 key
|
||||
logger.debug("Planner Plan: {}", state.value("planner_plan"));
|
||||
logger.debug("Executor Feedback: {}", state.value("executor_feedback"));
|
||||
}
|
||||
|
||||
return stateOptional;
|
||||
}
|
||||
|
||||
/**
|
||||
* 从执行结果中提取最终报告文本
|
||||
*
|
||||
* @param state 执行状态
|
||||
* @return 报告文本(如果存在)
|
||||
*/
|
||||
public Optional<String> extractFinalReport(OverAllState state) {
|
||||
logger.info("开始提取最终报告...");
|
||||
|
||||
// 提取 Planner 最终输出(包含完整的告警分析报告)
|
||||
Optional<AssistantMessage> plannerFinalOutput = state.value("planner_plan")
|
||||
.filter(AssistantMessage.class::isInstance)
|
||||
.map(AssistantMessage.class::cast);
|
||||
|
||||
if (plannerFinalOutput.isPresent()) {
|
||||
String reportText = plannerFinalOutput.get().getText();
|
||||
logger.info("成功提取到 Planner 最终报告,长度: {}", reportText.length());
|
||||
return Optional.of(reportText);
|
||||
} else {
|
||||
logger.warn("未能提取到 Planner 最终报告");
|
||||
return Optional.empty();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Planner Agent
|
||||
*/
|
||||
private ReactAgent buildPlannerAgent(ChatModel chatModel, ToolCallback[] toolCallbacks) {
|
||||
return ReactAgent.builder()
|
||||
.name("planner_agent")
|
||||
.description("负责拆解告警、规划与再规划步骤")
|
||||
.model(chatModel)
|
||||
.systemPrompt(buildPlannerPrompt())
|
||||
.methodTools(buildMethodToolsArray())
|
||||
.tools(toolCallbacks)
|
||||
.outputKey("planner_plan")
|
||||
.build();
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Executor Agent
|
||||
*/
|
||||
private ReactAgent buildExecutorAgent(ChatModel chatModel, ToolCallback[] toolCallbacks) {
|
||||
return ReactAgent.builder()
|
||||
.name("executor_agent")
|
||||
.description("负责执行 Planner 的首个步骤并及时反馈")
|
||||
.model(chatModel)
|
||||
.systemPrompt(buildExecutorPrompt())
|
||||
.methodTools(buildMethodToolsArray())
|
||||
.tools(toolCallbacks)
|
||||
.outputKey("executor_feedback")
|
||||
.build();
|
||||
}
|
||||
|
||||
/**
|
||||
* 动态构建方法工具数组
|
||||
* 根据 cls.mock-enabled 决定是否包含 QueryLogsTools
|
||||
*/
|
||||
private Object[] buildMethodToolsArray() {
|
||||
if (queryLogsTools != null) {
|
||||
// Mock 模式:包含 QueryLogsTools
|
||||
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools, queryLogsTools};
|
||||
} else {
|
||||
// 真实模式:不包含 QueryLogsTools(由 MCP 提供日志查询功能)
|
||||
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools};
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Planner Agent 系统提示词
|
||||
*/
|
||||
private String buildPlannerPrompt() {
|
||||
return """
|
||||
你是 Planner Agent,同时承担 Replanner 角色,负责:
|
||||
1. 读取当前输入任务 {input} 以及 Executor 的最近反馈 {executor_feedback}。
|
||||
2. 分析 Prometheus 告警、日志、内部文档等信息,制定可执行的下一步步骤。
|
||||
3. 在执行阶段,输出 JSON,包含 decision (PLAN|EXECUTE|FINISH)、step 描述、预期要调用的工具、以及必要的上下文。
|
||||
4. 调用任何腾讯云日志/主题相关工具时,region 参数必须使用连字符格式(如 ap-guangzhou),若不确定请省略以使用默认值。
|
||||
5. 严格禁止编造数据,只能引用工具返回的真实内容;如果连续 3 次调用同一工具仍失败或返回空结果,需停止该方向并在最终报告的结论部分说明"无法完成"的原因。
|
||||
|
||||
## 最终报告输出要求(CRITICAL)
|
||||
|
||||
当 decision=FINISH 时,你必须:
|
||||
1. **不要输出 JSON 格式**
|
||||
2. **直接输出完整的 Markdown 格式报告文本**
|
||||
3. **报告必须严格遵循以下模板**:
|
||||
|
||||
```
|
||||
# 告警分析报告
|
||||
|
||||
---
|
||||
|
||||
## 📋 活跃告警清单
|
||||
|
||||
| 告警名称 | 级别 | 目标服务 | 首次触发时间 | 最新触发时间 | 状态 |
|
||||
|---------|------|----------|-------------|-------------|------|
|
||||
| [告警1名称] | [级别] | [服务名] | [时间] | [时间] | 活跃 |
|
||||
| [告警2名称] | [级别] | [服务名] | [时间] | [时间] | 活跃 |
|
||||
|
||||
---
|
||||
|
||||
## 🔍 告警根因分析1 - [告警名称]
|
||||
|
||||
### 告警详情
|
||||
- **告警级别**: [级别]
|
||||
- **受影响服务**: [服务名]
|
||||
- **持续时间**: [X分钟]
|
||||
|
||||
### 症状描述
|
||||
[根据监控指标描述症状]
|
||||
|
||||
### 日志证据
|
||||
[引用查询到的关键日志]
|
||||
|
||||
### 根因结论
|
||||
[基于证据得出的根本原因]
|
||||
|
||||
---
|
||||
|
||||
## 🛠️ 处理方案执行1 - [告警名称]
|
||||
|
||||
### 已执行的排查步骤
|
||||
1. [步骤1]
|
||||
2. [步骤2]
|
||||
|
||||
### 处理建议
|
||||
[给出具体的处理建议]
|
||||
|
||||
### 预期效果
|
||||
[说明预期的效果]
|
||||
|
||||
---
|
||||
|
||||
## 🔍 告警根因分析2 - [告警名称]
|
||||
[如果有第2个告警,重复上述格式]
|
||||
|
||||
---
|
||||
|
||||
## 📊 结论
|
||||
|
||||
### 整体评估
|
||||
[总结所有告警的整体情况]
|
||||
|
||||
### 关键发现
|
||||
- [发现1]
|
||||
- [发现2]
|
||||
|
||||
### 后续建议
|
||||
1. [建议1]
|
||||
2. [建议2]
|
||||
|
||||
### 风险评估
|
||||
[评估当前风险等级和影响范围]
|
||||
```
|
||||
|
||||
**重要提醒**:
|
||||
- 最终输出必须是纯 Markdown 文本,不要包含 JSON 结构
|
||||
- 不要使用 "finalReport": "..." 这样的格式
|
||||
- 直接从 "# 告警分析报告" 开始输出
|
||||
- 所有内容必须基于工具查询的真实数据,严禁编造
|
||||
- 如果某个步骤失败,在结论中如实说明,不要跳过
|
||||
|
||||
""";
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Executor Agent 系统提示词
|
||||
*/
|
||||
private String buildExecutorPrompt() {
|
||||
return """
|
||||
你是 Executor Agent,负责读取 Planner 最新输出 {planner_plan},只执行其中的第一步。
|
||||
- 确认步骤所需的工具与参数,尤其是 region 参数要使用连字符格式(ap-guangzhou);若 Planner 未给出则使用默认区域。
|
||||
- 调用相应的工具并收集结果,如工具返回错误或空数据,需要将失败原因、请求参数一并记录,并停止进一步调用该工具(同一工具失败达到 3 次时应直接返回 FAILED)。
|
||||
- 将日志、指标、文档等证据整理成结构化摘要,标注对应的告警名称或资源,方便 Planner 填充"告警根因分析 / 处理方案执行"章节。
|
||||
- 以 JSON 形式返回执行状态、证据以及给 Planner 的建议,写入 executor_feedback,严禁编造未实际查询到的内容。
|
||||
|
||||
|
||||
输出示例:
|
||||
{
|
||||
"status": "SUCCESS",
|
||||
"summary": "近1小时未见 error 日志,仅有 info",
|
||||
"evidence": "...",
|
||||
"nextHint": "建议转向高占用进程"
|
||||
}
|
||||
""";
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Supervisor Agent 系统提示词
|
||||
*/
|
||||
private String buildSupervisorSystemPrompt() {
|
||||
return """
|
||||
你是 AI Ops Supervisor,负责调度 planner_agent 与 executor_agent:
|
||||
1. 当需要拆解任务或重新制定策略时,调用 planner_agent。
|
||||
2. 当 planner_agent 输出 decision=EXECUTE 时,调用 executor_agent 执行第一步。
|
||||
3. 根据 executor_agent 的反馈,评估是否需要再次调用 planner_agent,直到 decision=FINISH。
|
||||
4. FINISH 后,确保向最终用户输出完整的《告警分析报告》,格式必须严格为:
|
||||
告警分析报告\n---\n# 告警处理详情\n## 活跃告警清单\n## 告警根因分析N\n## 处理方案执行N\n## 结论。
|
||||
5. 若步骤涉及腾讯云日志/主题工具,请确保使用连字符区域 ID(ap-guangzhou 等),或省略 region 以采用默认值。
|
||||
6. 如果发现 Planner/Executor 在同一方向连续 3 次调用工具仍失败或没有数据,必须终止流程,直接输出"任务无法完成"的报告,明确告知失败原因,严禁凭空编造结果。
|
||||
|
||||
只允许在 planner_agent、executor_agent 与 FINISH 之间做出选择。
|
||||
|
||||
""";
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
|
||||
import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
|
||||
import com.superbiz.agent.agent.tool.DateTimeTools;
|
||||
import com.superbiz.agent.agent.tool.InternalDocsTools;
|
||||
import com.superbiz.agent.agent.tool.QueryLogsTools;
|
||||
import com.superbiz.agent.agent.tool.QueryMetricsTools;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.tool.ToolCallback;
|
||||
import org.springframework.ai.tool.ToolCallbackProvider;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 聊天服务
|
||||
* 封装 ReactAgent 对话的公共逻辑,包括模型创建、系统提示词构建、Agent 配置等
|
||||
*/
|
||||
@Service
|
||||
public class ChatService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(ChatService.class);
|
||||
|
||||
@Autowired
|
||||
private InternalDocsTools internalDocsTools;
|
||||
|
||||
@Autowired
|
||||
private DateTimeTools dateTimeTools;
|
||||
|
||||
@Autowired
|
||||
private QueryMetricsTools queryMetricsTools;
|
||||
|
||||
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过mcp配置注入
|
||||
private QueryLogsTools queryLogsTools;
|
||||
|
||||
@Autowired(required = false)
|
||||
private ToolCallbackProvider tools;
|
||||
|
||||
@Autowired
|
||||
private ChatModel chatModel;
|
||||
|
||||
/**
|
||||
* 获取注入的 ChatModel
|
||||
*/
|
||||
public ChatModel getChatModel() {
|
||||
return chatModel;
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建系统提示词(包含历史消息)
|
||||
* @param history 历史消息列表
|
||||
* @return 完整的系统提示词
|
||||
*/
|
||||
public String buildSystemPrompt(List<Map<String, String>> history) {
|
||||
StringBuilder systemPromptBuilder = new StringBuilder();
|
||||
|
||||
// 基础系统提示
|
||||
systemPromptBuilder.append("你是一个专业的智能助手,可以获取当前时间、查询天气信息、搜索内部文档知识库,以及查询 Prometheus 告警信息。\n");
|
||||
systemPromptBuilder.append("当用户询问时间相关问题时,**必须每次都调用 getCurrentDateTime 工具**,因为时间会不断变化。即使历史消息中有时间信息,也不要直接复用,必须重新查询最新时间。\n");
|
||||
systemPromptBuilder.append("当用户需要查询公司内部文档、流程、最佳实践或技术指南时,使用 queryInternalDocs 工具。\n");
|
||||
systemPromptBuilder.append("当用户需要查询 Prometheus 告警、监控指标或系统告警状态时,使用 queryPrometheusAlerts 工具。\n");
|
||||
systemPromptBuilder.append("当用户需要查询腾讯云日志时,请调用腾讯云mcp服务查询,默认查询地域ap-guangzhou,查询时间范围为近一个月。\n\n");
|
||||
|
||||
// 添加历史消息(过滤时间查询相关内容)
|
||||
if (!history.isEmpty()) {
|
||||
systemPromptBuilder.append("--- 对话历史 ---\n");
|
||||
for (Map<String, String> msg : history) {
|
||||
String role = msg.get("role");
|
||||
String content = msg.get("content");
|
||||
|
||||
// 🔧 过滤时间查询相关的历史消息,避免 LLM 复用旧的时间信息
|
||||
if ("user".equals(role) && isTimeQuery(content)) {
|
||||
continue; // 跳过时间查询问题
|
||||
}
|
||||
if ("assistant".equals(role) && containsTimeInfo(content)) {
|
||||
continue; // 跳过包含时间信息的回答
|
||||
}
|
||||
|
||||
if ("user".equals(role)) {
|
||||
systemPromptBuilder.append("用户: ").append(content).append("\n");
|
||||
} else if ("assistant".equals(role)) {
|
||||
systemPromptBuilder.append("助手: ").append(content).append("\n");
|
||||
}
|
||||
}
|
||||
systemPromptBuilder.append("--- 对话历史结束 ---\n\n");
|
||||
}
|
||||
|
||||
systemPromptBuilder.append("请基于以上对话历史,回答用户的新问题。");
|
||||
|
||||
return systemPromptBuilder.toString();
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断是否为时间查询问题
|
||||
*/
|
||||
private boolean isTimeQuery(String content) {
|
||||
if (content == null) {
|
||||
return false;
|
||||
}
|
||||
// 匹配常见的时间查询模式
|
||||
return content.matches(".*(现在|当前|此时).*(几点|时间).*") ||
|
||||
content.matches(".*(几点|时间).*(了|呢|[??]).*") ||
|
||||
content.toLowerCase().matches(".*(what.*time|current.*time).*");
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断是否包含时间信息
|
||||
*/
|
||||
private boolean containsTimeInfo(String content) {
|
||||
if (content == null) {
|
||||
return false;
|
||||
}
|
||||
// 匹配日期时间格式:2026年5月31日、15:57、下午3点 等
|
||||
return content.matches(".*(\\d{4}年\\d{1,2}月\\d{1,2}日|\\d{1,2}:\\d{2}|[上下午]+\\d{1,2}[点时]).*");
|
||||
}
|
||||
|
||||
/**
|
||||
* 动态构建方法工具数组
|
||||
* 根据 cls.mock-enabled 决定是否包含 QueryLogsTools
|
||||
*/
|
||||
public Object[] buildMethodToolsArray() {
|
||||
if (queryLogsTools != null) {
|
||||
// Mock 模式:包含 QueryLogsTools
|
||||
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools, queryLogsTools};
|
||||
} else {
|
||||
// 真实模式:不包含 QueryLogsTools(由 MCP 提供日志查询功能)
|
||||
return new Object[]{dateTimeTools, internalDocsTools, queryMetricsTools};
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取工具回调列表,mcp服务提供的工具
|
||||
*/
|
||||
public ToolCallback[] getToolCallbacks() {
|
||||
if (tools == null) {
|
||||
return new ToolCallback[0];
|
||||
}
|
||||
return tools.getToolCallbacks();
|
||||
}
|
||||
|
||||
/**
|
||||
* 记录可用工具列表:mcp服务提供的工具
|
||||
*/
|
||||
public void logAvailableTools() {
|
||||
if (tools == null) {
|
||||
logger.info("MCP 未启用,无远程工具");
|
||||
return;
|
||||
}
|
||||
ToolCallback[] toolCallbacks = tools.getToolCallbacks();
|
||||
logger.info("可用工具列表:");
|
||||
for (ToolCallback toolCallback : toolCallbacks) {
|
||||
logger.info(">>> {}", toolCallback.getToolDefinition().name());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 创建 ReactAgent
|
||||
* @param chatModel 聊天模型
|
||||
* @param systemPrompt 系统提示词
|
||||
* @return 配置好的 ReactAgent
|
||||
*/
|
||||
public ReactAgent createReactAgent(ChatModel chatModel, String systemPrompt) {
|
||||
return ReactAgent.builder()
|
||||
.name("intelligent_assistant")
|
||||
.model(chatModel)
|
||||
.systemPrompt(systemPrompt)
|
||||
.methodTools(buildMethodToolsArray())
|
||||
.tools(getToolCallbacks())
|
||||
.build();
|
||||
}
|
||||
|
||||
/**
|
||||
* 执行 ReactAgent 对话(非流式)
|
||||
* @param agent ReactAgent 实例
|
||||
* @param question 用户问题
|
||||
* @return AI 回复
|
||||
*/
|
||||
public String executeChat(ReactAgent agent, String question) throws GraphRunnerException {
|
||||
logger.info("执行 ReactAgent.call() - 自动处理工具调用");
|
||||
var response = agent.call(question);
|
||||
String answer = response.getText();
|
||||
logger.info("ReactAgent 对话完成,答案长度: {}", answer.length());
|
||||
return answer;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,405 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import com.superbiz.agent.config.DocumentChunkConfig;
|
||||
import com.superbiz.agent.dto.DocumentChunk;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
|
||||
/**
|
||||
* 文档分片服务
|
||||
* 负责将长文档切分为多个有语义完整性的小片段
|
||||
*/
|
||||
@Service
|
||||
public class DocumentChunkService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(DocumentChunkService.class);
|
||||
|
||||
@Autowired
|
||||
private DocumentChunkConfig chunkConfig;
|
||||
|
||||
/**
|
||||
* 智能分片文档
|
||||
* 优先按照标题、段落边界进行分割,保持语义完整性
|
||||
*
|
||||
* @param content 文档内容
|
||||
* @param filePath 文件路径(用于日志)
|
||||
* @return 文档分片列表
|
||||
*/
|
||||
public List<DocumentChunk> chunkDocument(String content, String filePath) {
|
||||
List<DocumentChunk> chunks = new ArrayList<>();
|
||||
|
||||
if (content == null || content.trim().isEmpty()) {
|
||||
logger.warn("文档内容为空: {}", filePath);
|
||||
return chunks;
|
||||
}
|
||||
|
||||
// 1. 首先尝试按标题分割(Markdown格式)
|
||||
List<Section> sections = splitByHeadings(content);
|
||||
|
||||
// 2. 对每个章节进行进一步分片
|
||||
int globalChunkIndex = 0;
|
||||
for (Section section : sections) {
|
||||
List<DocumentChunk> sectionChunks = chunkSection(section, globalChunkIndex);
|
||||
chunks.addAll(sectionChunks);
|
||||
globalChunkIndex += sectionChunks.size();
|
||||
}
|
||||
|
||||
logger.info("文档分片完成: {} -> {} 个分片", filePath, chunks.size());
|
||||
return chunks;
|
||||
}
|
||||
|
||||
/**
|
||||
* 按照 Markdown 标题分割文档
|
||||
*/
|
||||
private List<Section> splitByHeadings(String content) {
|
||||
List<Section> sections = new ArrayList<>();
|
||||
|
||||
// 匹配 Markdown 标题:# 标题, ## 标题, ### 标题等
|
||||
Pattern headingPattern = Pattern.compile("^(#{1,6})\\s+(.+)$", Pattern.MULTILINE);
|
||||
Matcher matcher = headingPattern.matcher(content);
|
||||
|
||||
int lastEnd = 0;
|
||||
String currentTitle = null;
|
||||
|
||||
while (matcher.find()) {
|
||||
// 保存上一个章节
|
||||
if (lastEnd < matcher.start()) {
|
||||
String sectionContent = content.substring(lastEnd, matcher.start()).trim();
|
||||
if (!sectionContent.isEmpty()) {
|
||||
sections.add(new Section(currentTitle, sectionContent, lastEnd));
|
||||
}
|
||||
}
|
||||
|
||||
// 更新当前标题
|
||||
currentTitle = matcher.group(2).trim();
|
||||
lastEnd = matcher.start();
|
||||
}
|
||||
|
||||
// 添加最后一个章节
|
||||
if (lastEnd < content.length()) {
|
||||
String sectionContent = content.substring(lastEnd).trim();
|
||||
if (!sectionContent.isEmpty()) {
|
||||
sections.add(new Section(currentTitle, sectionContent, lastEnd));
|
||||
}
|
||||
}
|
||||
|
||||
// 如果没有找到任何标题,将整个文档作为一个章节
|
||||
if (sections.isEmpty()) {
|
||||
sections.add(new Section(null, content, 0));
|
||||
}
|
||||
|
||||
return sections;
|
||||
}
|
||||
|
||||
/**
|
||||
* 对单个章节进行分片
|
||||
* <p>
|
||||
* 核心改造(Phase 1):
|
||||
* - Token 估算替代字符计数
|
||||
* - 感知有序/无序列表结构,不在列表中间切断
|
||||
* - 软边界(maxTokens)+ 硬上限(maxTokensHard)双重控制
|
||||
* - 修复 currentStartIndex 漂移:用段落原始位置而非手工推算
|
||||
*/
|
||||
private List<DocumentChunk> chunkSection(Section section, int startChunkIndex) {
|
||||
List<DocumentChunk> chunks = new ArrayList<>();
|
||||
String content = section.content;
|
||||
String title = section.title;
|
||||
|
||||
// 短章节直接作为一个分片(用 token 估算替代字符数做短路判断)
|
||||
if (content.length() <= chunkConfig.getMaxSize()
|
||||
&& estimateTokens(content) <= chunkConfig.getMaxTokens()) {
|
||||
DocumentChunk chunk = new DocumentChunk(
|
||||
content,
|
||||
section.startIndex,
|
||||
section.startIndex + content.length(),
|
||||
startChunkIndex
|
||||
);
|
||||
chunk.setTitle(title);
|
||||
chunks.add(chunk);
|
||||
return chunks;
|
||||
}
|
||||
|
||||
// 章节内容较长,需要进一步分片
|
||||
List<String> paragraphs = splitByParagraphs(content);
|
||||
if (paragraphs.isEmpty()) {
|
||||
return chunks;
|
||||
}
|
||||
|
||||
// 定位每个段落在 section.content 中的位置(修复 index 漂移)
|
||||
List<ParagraphPos> paraPositions = locateParagraphPositions(paragraphs, content);
|
||||
|
||||
// 当前分片的段落范围
|
||||
int chunkParaStart = 0; // 当前分片第一个段落的索引(在 paragraphs 中)
|
||||
StringBuilder buffer = new StringBuilder();
|
||||
int tokenCount = 0;
|
||||
int chunkIndex = startChunkIndex;
|
||||
|
||||
for (int i = 0; i < paragraphs.size(); i++) {
|
||||
String paragraph = paragraphs.get(i);
|
||||
int paraTokens = estimateTokens(paragraph);
|
||||
|
||||
// 判断是否需要切分
|
||||
if (buffer.length() > 0 && tokenCount + paraTokens > chunkConfig.getMaxTokens()) {
|
||||
|
||||
// 检查是否处于不可中断的上下文中
|
||||
if (isInUnbreakableContext(buffer.toString(), paragraph)) {
|
||||
// 硬上限保护:即使不可中断也不能无限膨胀
|
||||
if (tokenCount + paraTokens > chunkConfig.getMaxTokensHard()) {
|
||||
logger.debug(" 触及硬上限 ({} tokens),强制切分", tokenCount + paraTokens);
|
||||
chunkParaStart = saveChunkAndGetNextStart(
|
||||
chunks, section, paraPositions,
|
||||
chunkParaStart, i, title, chunkIndex);
|
||||
chunkIndex++;
|
||||
|
||||
String prevChunkContent = chunks.get(chunks.size() - 1).getContent();
|
||||
String overlap = getOverlapText(prevChunkContent);
|
||||
buffer = new StringBuilder(overlap);
|
||||
tokenCount = estimateTokens(overlap);
|
||||
}
|
||||
// 否则:容忍超出(软边界)
|
||||
} else {
|
||||
// 安全切点:段落边界
|
||||
chunkParaStart = saveChunkAndGetNextStart(
|
||||
chunks, section, paraPositions,
|
||||
chunkParaStart, i, title, chunkIndex);
|
||||
chunkIndex++;
|
||||
|
||||
// 新分片以重叠文本开头
|
||||
String prevChunkContent = chunks.get(chunks.size() - 1).getContent();
|
||||
String overlap = getOverlapText(prevChunkContent);
|
||||
buffer = new StringBuilder(overlap);
|
||||
tokenCount = estimateTokens(overlap);
|
||||
}
|
||||
}
|
||||
|
||||
buffer.append(paragraph).append("\n\n");
|
||||
tokenCount += paraTokens;
|
||||
}
|
||||
|
||||
// 保存最后一个分片
|
||||
if (buffer.length() > 0 && chunkParaStart < paragraphs.size()) {
|
||||
String chunkContent = buffer.toString().trim();
|
||||
int actualStart = paraPositions.get(chunkParaStart).start;
|
||||
int actualEnd = paraPositions.get(paragraphs.size() - 1).end;
|
||||
DocumentChunk chunk = new DocumentChunk(
|
||||
chunkContent,
|
||||
section.startIndex + actualStart,
|
||||
section.startIndex + actualEnd,
|
||||
chunkIndex
|
||||
);
|
||||
chunk.setTitle(title);
|
||||
chunks.add(chunk);
|
||||
}
|
||||
|
||||
return chunks;
|
||||
}
|
||||
|
||||
/**
|
||||
* 保存当前分块,返回下一个分块的起始段落索引
|
||||
* <p>
|
||||
* 从 section.content 中提取原始文本(而非手工拼装),修复 index 漂移问题
|
||||
*/
|
||||
private int saveChunkAndGetNextStart(
|
||||
List<DocumentChunk> chunks,
|
||||
Section section,
|
||||
List<ParagraphPos> paraPositions,
|
||||
int fromPara,
|
||||
int toPara,
|
||||
String title,
|
||||
int chunkIndex) {
|
||||
|
||||
int actualStart = paraPositions.get(fromPara).start;
|
||||
int actualEnd = paraPositions.get(toPara - 1).end;
|
||||
String originalText = section.content.substring(actualStart, actualEnd);
|
||||
|
||||
DocumentChunk chunk = new DocumentChunk(
|
||||
originalText,
|
||||
section.startIndex + actualStart,
|
||||
section.startIndex + actualEnd,
|
||||
chunkIndex
|
||||
);
|
||||
chunk.setTitle(title);
|
||||
chunks.add(chunk);
|
||||
|
||||
return toPara; // 下一个分块的起始段落索引
|
||||
}
|
||||
|
||||
/**
|
||||
* 按段落分割文本
|
||||
*/
|
||||
private List<String> splitByParagraphs(String content) {
|
||||
List<String> paragraphs = new ArrayList<>();
|
||||
|
||||
// 按双换行符分割段落
|
||||
String[] parts = content.split("\n\n+");
|
||||
for (String part : parts) {
|
||||
String trimmed = part.trim();
|
||||
if (!trimmed.isEmpty()) {
|
||||
paragraphs.add(trimmed);
|
||||
}
|
||||
}
|
||||
|
||||
return paragraphs;
|
||||
}
|
||||
|
||||
/**
|
||||
* 定位每个段落在原始文本中的字符偏移
|
||||
*/
|
||||
private List<ParagraphPos> locateParagraphPositions(List<String> paragraphs, String sectionContent) {
|
||||
List<ParagraphPos> positions = new ArrayList<>();
|
||||
int searchFrom = 0;
|
||||
for (String p : paragraphs) {
|
||||
int idx = sectionContent.indexOf(p, searchFrom);
|
||||
if (idx >= 0) {
|
||||
positions.add(new ParagraphPos(idx, idx + p.length()));
|
||||
searchFrom = idx + p.length();
|
||||
} else {
|
||||
// fallback: 段落在原文中找不到(不应该发生)
|
||||
positions.add(new ParagraphPos(searchFrom, searchFrom + p.length()));
|
||||
searchFrom += p.length();
|
||||
}
|
||||
}
|
||||
return positions;
|
||||
}
|
||||
|
||||
/**
|
||||
* 启发式 token 估算(无需外部依赖)
|
||||
* <p>
|
||||
* 中文(BMP): ~1 字符/token
|
||||
* 英文/数字/标点: ~4 字符/token
|
||||
* 空白字符忽略
|
||||
*/
|
||||
private int estimateTokens(String text) {
|
||||
int nonCjkCount = 0;
|
||||
int cjkCount = 0;
|
||||
for (char c : text.toCharArray()) {
|
||||
if (Character.isWhitespace(c)) {
|
||||
continue;
|
||||
}
|
||||
Character.UnicodeBlock block = Character.UnicodeBlock.of(c);
|
||||
if (block == Character.UnicodeBlock.CJK_UNIFIED_IDEOGRAPHS
|
||||
|| block == Character.UnicodeBlock.CJK_UNIFIED_IDEOGRAPHS_EXTENSION_A
|
||||
|| block == Character.UnicodeBlock.CJK_UNIFIED_IDEOGRAPHS_EXTENSION_B
|
||||
|| block == Character.UnicodeBlock.CJK_COMPATIBILITY_IDEOGRAPHS) {
|
||||
cjkCount++;
|
||||
} else {
|
||||
nonCjkCount++;
|
||||
}
|
||||
}
|
||||
return cjkCount + (nonCjkCount + 3) / 4; // 非中文每 4 字符算 1 token,向上取整
|
||||
}
|
||||
|
||||
/**
|
||||
* 判断当前段落是否属于不可中断的结构
|
||||
* <p>
|
||||
* 不可中断结构包括:
|
||||
* - 有序列表项("1. ", "2. " 格式)
|
||||
* - 无序列表项("- " 或 "* " 格式)
|
||||
* - 未闭合的代码块(``` 内)
|
||||
*/
|
||||
private boolean isInUnbreakableContext(String buffer, String nextParagraph) {
|
||||
// 有序列表:判断 buffer 末尾和下一段是否都是列表项
|
||||
if (nextParagraph.matches("^\\d{1,2}\\.\\s.*")) {
|
||||
String lastLine = getLastNonEmptyLine(buffer);
|
||||
if (lastLine != null && lastLine.matches("^\\d{1,2}\\.\\s.*")) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
// 无序列表:"- " 或 "* " 格式
|
||||
if (nextParagraph.matches("^[-*]\\s.*")) {
|
||||
String lastLine = getLastNonEmptyLine(buffer);
|
||||
if (lastLine != null && lastLine.matches("^[-*]\\s.*")) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
// 代码块:``` 未闭合
|
||||
if (buffer.contains("```")) {
|
||||
int count = 0;
|
||||
for (int i = 0; i <= buffer.length() - 3; i++) {
|
||||
if (buffer.substring(i).startsWith("```")) {
|
||||
count++;
|
||||
i += 2;
|
||||
}
|
||||
}
|
||||
if (count % 2 == 1) {
|
||||
return true; // 奇数个 ``` → 在代码块内部
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取 buffer 中最后一行非空白文本
|
||||
*/
|
||||
private String getLastNonEmptyLine(String buffer) {
|
||||
String[] lines = buffer.split("\n");
|
||||
for (int i = lines.length - 1; i >= 0; i--) {
|
||||
String line = lines[i].trim();
|
||||
if (!line.isEmpty()) {
|
||||
return line;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取重叠文本
|
||||
* 从文本末尾提取指定长度的内容作为下一个分片的开头
|
||||
*/
|
||||
private String getOverlapText(String text) {
|
||||
int overlapSize = Math.min(chunkConfig.getOverlap(), text.length());
|
||||
if (overlapSize <= 0) {
|
||||
return "";
|
||||
}
|
||||
|
||||
// 从末尾提取重叠内容
|
||||
String overlap = text.substring(text.length() - overlapSize);
|
||||
|
||||
// 尝试在句子边界截断(查找最后一个句号、问号、感叹号)
|
||||
int lastSentenceEnd = Math.max(
|
||||
overlap.lastIndexOf('。'),
|
||||
Math.max(overlap.lastIndexOf('?'), overlap.lastIndexOf('!'))
|
||||
);
|
||||
|
||||
if (lastSentenceEnd > overlapSize / 2) {
|
||||
return overlap.substring(lastSentenceEnd + 1).trim();
|
||||
}
|
||||
|
||||
return overlap.trim();
|
||||
}
|
||||
|
||||
/**
|
||||
* 段落在原文中的位置
|
||||
*/
|
||||
private static class ParagraphPos {
|
||||
final int start;
|
||||
final int end;
|
||||
|
||||
ParagraphPos(int start, int end) {
|
||||
this.start = start;
|
||||
this.end = end;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 章节数据类
|
||||
*/
|
||||
private static class Section {
|
||||
String title;
|
||||
String content;
|
||||
int startIndex;
|
||||
|
||||
Section(String title, String content, int startIndex) {
|
||||
this.title = title;
|
||||
this.content = content;
|
||||
this.startIndex = startIndex;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,190 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.chat.model.ChatResponse;
|
||||
import org.springframework.ai.chat.prompt.Prompt;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.stereotype.Service;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* RAG (Retrieval-Augmented Generation) 服务
|
||||
* 结合向量检索和大语言模型生成答案
|
||||
*/
|
||||
@Service
|
||||
public class RagService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(RagService.class);
|
||||
|
||||
@Autowired
|
||||
private VectorSearchService vectorSearchService;
|
||||
|
||||
@Autowired
|
||||
private ChatModel chatModel;
|
||||
|
||||
@Value("${rag.top-k:3}")
|
||||
private int topK;
|
||||
|
||||
/**
|
||||
* 流式处理用户问题(不带历史消息)
|
||||
*
|
||||
* @param question 用户问题
|
||||
* @param callback 流式回调接口
|
||||
*/
|
||||
public void queryStream(String question, StreamCallback callback) {
|
||||
queryStream(question, new ArrayList<>(), callback);
|
||||
}
|
||||
|
||||
/**
|
||||
* 流式处理用户问题(带历史消息)
|
||||
*
|
||||
* @param question 用户问题
|
||||
* @param history 历史消息列表,格式:[{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]
|
||||
* @param callback 流式回调接口
|
||||
*/
|
||||
public void queryStream(String question, List<Map<String, String>> history, StreamCallback callback) {
|
||||
try {
|
||||
logger.info("收到 RAG 流式查询: {}", question);
|
||||
|
||||
// 1. 从向量数据库检索相关文档
|
||||
List<VectorSearchService.SearchResult> searchResults =
|
||||
vectorSearchService.searchSimilarDocuments(question, topK);
|
||||
|
||||
// 发送检索结果
|
||||
callback.onSearchResults(searchResults);
|
||||
|
||||
if (searchResults.isEmpty()) {
|
||||
logger.warn("未找到相关文档");
|
||||
callback.onComplete("抱歉,我在知识库中没有找到相关信息来回答您的问题。", "");
|
||||
return;
|
||||
}
|
||||
|
||||
// 2. 构建上下文和提示词
|
||||
String context = buildContext(searchResults);
|
||||
String prompt = buildPrompt(question, context);
|
||||
|
||||
// 3. 流式调用大语言模型(传入历史消息)
|
||||
generateAnswerStream(prompt, history, callback);
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("RAG 流式查询失败", e);
|
||||
callback.onError(e);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建上下文
|
||||
*/
|
||||
private String buildContext(List<VectorSearchService.SearchResult> searchResults) {
|
||||
StringBuilder context = new StringBuilder();
|
||||
|
||||
for (int i = 0; i < searchResults.size(); i++) {
|
||||
VectorSearchService.SearchResult result = searchResults.get(i);
|
||||
context.append("【参考资料 ").append(i + 1).append("】\n");
|
||||
context.append(result.getContent()).append("\n\n");
|
||||
}
|
||||
|
||||
return context.toString();
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建提示词
|
||||
*/
|
||||
private String buildPrompt(String question, String context) {
|
||||
return String.format(
|
||||
"你是一个专业的AI助手。请根据以下参考资料回答用户的问题。\n\n" +
|
||||
"参考资料:\n%s\n" +
|
||||
"用户问题:%s\n\n" +
|
||||
"请基于上述参考资料给出准确、详细的回答。如果参考资料中没有相关信息,请明确说明。",
|
||||
context, question
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 生成答案(流式)
|
||||
*
|
||||
* @param prompt 当前问题的提示词
|
||||
* @param history 历史消息列表
|
||||
* @param callback 流式回调接口
|
||||
*/
|
||||
private void generateAnswerStream(String prompt, List<Map<String, String>> history, StreamCallback callback) {
|
||||
// 构建消息列表:历史消息 + 当前问题
|
||||
List<Message> messages = new ArrayList<>();
|
||||
|
||||
// 添加历史消息
|
||||
for (Map<String, String> historyMsg : history) {
|
||||
String role = historyMsg.get("role");
|
||||
String content = historyMsg.get("content");
|
||||
|
||||
if ("user".equals(role)) {
|
||||
messages.add(new UserMessage(content));
|
||||
} else if ("assistant".equals(role)) {
|
||||
messages.add(new AssistantMessage(content));
|
||||
}
|
||||
}
|
||||
|
||||
// 添加当前用户问题
|
||||
messages.add(new UserMessage(prompt));
|
||||
|
||||
logger.debug("发送给AI模型的消息数量: {}(包含 {} 条历史消息)",
|
||||
messages.size(), history.size());
|
||||
|
||||
logger.info("开始调用AI模型流式接口...");
|
||||
|
||||
StringBuilder reasoningContent = new StringBuilder();
|
||||
StringBuilder finalContent = new StringBuilder();
|
||||
|
||||
Flux<ChatResponse> flux = chatModel.stream(new Prompt(messages));
|
||||
|
||||
logger.info("开始接收AI模型流式响应...");
|
||||
|
||||
flux.subscribe(
|
||||
response -> {
|
||||
if (response.getResults() != null && !response.getResults().isEmpty()) {
|
||||
String content = response.getResults().get(0).getOutput().getText();
|
||||
|
||||
if (content != null && !content.isEmpty()) {
|
||||
logger.debug("收到AI模型内容块: {}", content);
|
||||
|
||||
finalContent.append(content);
|
||||
callback.onContentChunk(content);
|
||||
|
||||
logger.debug("已调用 onContentChunk 回调");
|
||||
} else {
|
||||
logger.debug("收到空内容块,跳过");
|
||||
}
|
||||
}
|
||||
},
|
||||
error -> {
|
||||
logger.error("AI模型流式响应失败", error);
|
||||
callback.onError(new Exception("AI模型流式响应失败: " + error.getMessage(), error));
|
||||
},
|
||||
() -> {
|
||||
logger.info("AI模型流式响应完成,总内容长度: {}", finalContent.length());
|
||||
callback.onComplete(finalContent.toString(), reasoningContent.toString());
|
||||
logger.info("已调用 onComplete 回调");
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 流式回调接口
|
||||
*/
|
||||
public interface StreamCallback {
|
||||
void onSearchResults(List<VectorSearchService.SearchResult> results);
|
||||
void onReasoningChunk(String chunk);
|
||||
void onContentChunk(String chunk);
|
||||
void onComplete(String fullContent, String fullReasoning);
|
||||
void onError(Exception e);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.ai.embedding.EmbeddingModel;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* 向量嵌入服务
|
||||
* 使用阿里云 DashScope Text Embedding API
|
||||
*/
|
||||
@Service
|
||||
public class VectorEmbeddingService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(VectorEmbeddingService.class);
|
||||
|
||||
@Autowired
|
||||
private EmbeddingModel embeddingModel;
|
||||
|
||||
/**
|
||||
* 生成向量嵌入
|
||||
* 调用阿里云 DashScope Text Embedding API
|
||||
*
|
||||
* @param content 文本内容
|
||||
* @return 向量嵌入(浮点数列表)
|
||||
*/
|
||||
public List<Float> generateEmbedding(String content) {
|
||||
try {
|
||||
if (content == null || content.trim().isEmpty()) {
|
||||
logger.warn("内容为空,无法生成向量");
|
||||
throw new IllegalArgumentException("内容不能为空");
|
||||
}
|
||||
|
||||
logger.debug("开始生成向量嵌入, 内容长度: {} 字符", content.length());
|
||||
|
||||
float[] embedding = embeddingModel.embed(content);
|
||||
|
||||
List<Float> floatEmbedding = new ArrayList<>(embedding.length);
|
||||
for (float v : embedding) {
|
||||
floatEmbedding.add(v);
|
||||
}
|
||||
|
||||
logger.info("成功生成向量嵌入, 内容长度: {} 字符, 向量维度: {}",
|
||||
content.length(), floatEmbedding.size());
|
||||
|
||||
return floatEmbedding;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("生成向量嵌入失败, 内容长度: {}", content != null ? content.length() : 0, e);
|
||||
throw new RuntimeException("生成向量嵌入失败: " + e.getMessage(), e);
|
||||
}
|
||||
}
|
||||
|
||||
public List<List<Float>> generateEmbeddings(List<String> contents) {
|
||||
try {
|
||||
if (contents == null || contents.isEmpty()) {
|
||||
logger.warn("内容列表为空,无法生成向量");
|
||||
return Collections.emptyList();
|
||||
}
|
||||
|
||||
logger.info("开始批量生成向量嵌入, 数量: {}", contents.size());
|
||||
|
||||
List<float[]> embeddings = embeddingModel.embed(contents);
|
||||
|
||||
List<List<Float>> result = new ArrayList<>();
|
||||
for (float[] embedding : embeddings) {
|
||||
List<Float> floatEmbedding = new ArrayList<>(embedding.length);
|
||||
for (float v : embedding) {
|
||||
floatEmbedding.add(v);
|
||||
}
|
||||
result.add(floatEmbedding);
|
||||
}
|
||||
|
||||
logger.info("成功批量生成向量嵌入, 数量: {}, 维度: {}",
|
||||
result.size(),
|
||||
result.isEmpty() ? 0 : result.get(0).size());
|
||||
|
||||
return result;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("批量生成向量嵌入失败", e);
|
||||
throw new RuntimeException("批量生成向量嵌入失败: " + e.getMessage(), e);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 生成查询向量
|
||||
*
|
||||
* @param query 查询文本
|
||||
* @return 向量嵌入
|
||||
*/
|
||||
public List<Float> generateQueryVector(String query) {
|
||||
return generateEmbedding(query);
|
||||
}
|
||||
|
||||
/**
|
||||
* 计算两个向量的余弦相似度
|
||||
*
|
||||
* @param vector1 向量1
|
||||
* @param vector2 向量2
|
||||
* @return 余弦相似度 [-1, 1]
|
||||
*/
|
||||
public float calculateCosineSimilarity(List<Float> vector1, List<Float> vector2) {
|
||||
if (vector1.size() != vector2.size()) {
|
||||
throw new IllegalArgumentException("向量维度不匹配");
|
||||
}
|
||||
|
||||
float dotProduct = 0.0f;
|
||||
float norm1 = 0.0f;
|
||||
float norm2 = 0.0f;
|
||||
|
||||
for (int i = 0; i < vector1.size(); i++) {
|
||||
dotProduct += vector1.get(i) * vector2.get(i);
|
||||
norm1 += vector1.get(i) * vector1.get(i);
|
||||
norm2 += vector2.get(i) * vector2.get(i);
|
||||
}
|
||||
|
||||
return dotProduct / (float) (Math.sqrt(norm1) * Math.sqrt(norm2));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,351 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import io.milvus.client.MilvusServiceClient;
|
||||
import io.milvus.grpc.MutationResult;
|
||||
import io.milvus.param.R;
|
||||
import io.milvus.param.RpcStatus;
|
||||
import io.milvus.param.collection.LoadCollectionParam;
|
||||
import io.milvus.param.dml.DeleteParam;
|
||||
import io.milvus.param.dml.InsertParam;
|
||||
import lombok.Getter;
|
||||
import lombok.Setter;
|
||||
import com.superbiz.agent.constant.MilvusConstants;
|
||||
import com.superbiz.agent.dto.DocumentChunk;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.beans.factory.annotation.Value;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.io.File;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.nio.file.Paths;
|
||||
import java.time.LocalDateTime;
|
||||
import java.util.*;
|
||||
|
||||
/**
|
||||
* 向量索引服务
|
||||
* 负责读取文件、生成向量、存储到 Milvus
|
||||
*/
|
||||
@Service
|
||||
public class VectorIndexService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(VectorIndexService.class);
|
||||
|
||||
@Autowired
|
||||
private MilvusServiceClient milvusClient;
|
||||
|
||||
@Autowired
|
||||
private VectorEmbeddingService embeddingService;
|
||||
|
||||
@Autowired
|
||||
private DocumentChunkService chunkService;
|
||||
|
||||
@Value("${file.upload.path}")
|
||||
private String uploadPath;
|
||||
|
||||
/**
|
||||
* 索引指定目录下的所有文件
|
||||
*
|
||||
* @param directoryPath 目录路径(可选,默认使用配置的上传目录)
|
||||
* @return 索引结果 这里可以优化:定时重建目录下所有文件的索引
|
||||
*/
|
||||
public IndexingResult indexDirectory(String directoryPath) {
|
||||
IndexingResult result = new IndexingResult();
|
||||
result.setStartTime(LocalDateTime.now());
|
||||
|
||||
try {
|
||||
// 使用指定目录或默认上传目录
|
||||
String targetPath = (directoryPath != null && !directoryPath.trim().isEmpty())
|
||||
? directoryPath : uploadPath;
|
||||
|
||||
Path dirPath = Paths.get(targetPath).normalize();
|
||||
File directory = dirPath.toFile();
|
||||
|
||||
if (!directory.exists() || !directory.isDirectory()) {
|
||||
throw new IllegalArgumentException("目录不存在或不是有效目录: " + targetPath);
|
||||
}
|
||||
|
||||
result.setDirectoryPath(directory.getAbsolutePath());
|
||||
|
||||
// 获取所有支持的文件
|
||||
File[] files = directory.listFiles((dir, name) ->
|
||||
name.endsWith(".txt") || name.endsWith(".md")
|
||||
);
|
||||
|
||||
if (files == null || files.length == 0) {
|
||||
logger.warn("目录中没有找到支持的文件: {}", targetPath);
|
||||
result.setTotalFiles(0);
|
||||
result.setSuccess(true);
|
||||
result.setEndTime(LocalDateTime.now());
|
||||
return result;
|
||||
}
|
||||
|
||||
result.setTotalFiles(files.length);
|
||||
logger.info("开始索引目录: {}, 找到 {} 个文件", targetPath, files.length);
|
||||
|
||||
// 遍历并索引每个文件
|
||||
for (File file : files) {
|
||||
try {
|
||||
indexSingleFile(file.getAbsolutePath());
|
||||
result.incrementSuccessCount();
|
||||
logger.info("✓ 文件索引成功: {}", file.getName());
|
||||
} catch (Exception e) {
|
||||
result.incrementFailCount();
|
||||
result.addFailedFile(file.getAbsolutePath(), e.getMessage());
|
||||
logger.error("✗ 文件索引失败: {}", file.getName(), e);
|
||||
}
|
||||
}
|
||||
|
||||
result.setSuccess(result.getFailCount() == 0);
|
||||
result.setEndTime(LocalDateTime.now());
|
||||
|
||||
logger.info("目录索引完成: 总数={}, 成功={}, 失败={}",
|
||||
result.getTotalFiles(), result.getSuccessCount(), result.getFailCount());
|
||||
|
||||
return result;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("索引目录失败", e);
|
||||
result.setSuccess(false);
|
||||
result.setErrorMessage(e.getMessage());
|
||||
result.setEndTime(LocalDateTime.now());
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 索引单个文件
|
||||
*
|
||||
* @param filePath 文件路径
|
||||
* @throws Exception 索引失败时抛出异常
|
||||
*/
|
||||
public void indexSingleFile(String filePath) throws Exception {
|
||||
Path path = Paths.get(filePath).normalize();
|
||||
File file = path.toFile();
|
||||
|
||||
if (!file.exists() || !file.isFile()) {
|
||||
throw new IllegalArgumentException("文件不存在: " + filePath);
|
||||
}
|
||||
|
||||
logger.info("开始索引文件: {}", path);
|
||||
|
||||
// 1. 读取文件内容
|
||||
String content = Files.readString(path);
|
||||
logger.info("读取文件: {}, 内容长度: {} 字符", path, content.length());
|
||||
|
||||
// 2. 删除该文件的旧数据(如果存在)
|
||||
deleteExistingData(path.toString());
|
||||
|
||||
// 3. 文档分片
|
||||
List<DocumentChunk> chunks = chunkService.chunkDocument(content, path.toString());
|
||||
logger.info("文档分片完成: {} -> {} 个分片", filePath, chunks.size());
|
||||
|
||||
// 4. 为每个分片生成向量并插入 Milvus
|
||||
for (int i = 0; i < chunks.size(); i++) {
|
||||
DocumentChunk chunk = chunks.get(i);
|
||||
|
||||
try {
|
||||
// 生成向量
|
||||
List<Float> vector = embeddingService.generateEmbedding(chunk.getContent());
|
||||
|
||||
// 构建元数据(包含文件信息)
|
||||
Map<String, Object> metadata = buildMetadata(path.toString(), chunk, chunks.size());
|
||||
|
||||
// 插入到 Milvus
|
||||
insertToMilvus(chunk.getContent(), vector, metadata, chunk.getChunkIndex());
|
||||
|
||||
logger.info("✓ 分片 {}/{} 索引成功", i + 1, chunks.size());
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("✗ 分片 {}/{} 索引失败", i + 1, chunks.size(), e);
|
||||
throw new RuntimeException("分片索引失败: " + e.getMessage(), e);
|
||||
}
|
||||
}
|
||||
|
||||
logger.info("文件索引完成: {}, 共 {} 个分片", filePath, chunks.size());
|
||||
}
|
||||
|
||||
/**
|
||||
* 删除文件的旧数据(根据 metadata._source)
|
||||
*/
|
||||
private void deleteExistingData(String filePath) {
|
||||
try {
|
||||
// 使用统一的路径分隔符(正斜杠)用于Milvus存储,避免表达式解析错误
|
||||
// 将系统路径转换为统一格式
|
||||
Path path = Paths.get(filePath).normalize();
|
||||
String normalizedPath = path.toString().replace(File.separator, "/");
|
||||
|
||||
// 构建删除表达式:metadata["_source"] == "xxx"
|
||||
String expr = String.format("metadata[\"_source\"] == \"%s\"", normalizedPath);
|
||||
|
||||
logger.info("准备删除旧数据,路径: {}, 表达式: {}", normalizedPath, expr);
|
||||
|
||||
// 确保 collection 已加载(删除操作需要集合已加载)
|
||||
R<RpcStatus> loadResponse = milvusClient.loadCollection(
|
||||
LoadCollectionParam.newBuilder()
|
||||
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
|
||||
.build()
|
||||
);
|
||||
|
||||
// 状态码 65535 表示集合已经加载,这不是错误
|
||||
if (loadResponse.getStatus() != 0 && loadResponse.getStatus() != 65535) {
|
||||
logger.warn("加载 collection 失败: {}", loadResponse.getMessage());
|
||||
return;
|
||||
}
|
||||
|
||||
DeleteParam deleteParam = DeleteParam.newBuilder()
|
||||
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
|
||||
.withExpr(expr)
|
||||
.build();
|
||||
|
||||
R<MutationResult> response = milvusClient.delete(deleteParam);
|
||||
|
||||
if (response.getStatus() != 0) {
|
||||
logger.warn("删除旧数据时出现警告: {}", response.getMessage());
|
||||
} else {
|
||||
long deletedCount = response.getData().getDeleteCnt();
|
||||
logger.info("✓ 已删除文件的旧数据: {}, 删除记录数: {}", normalizedPath, deletedCount);
|
||||
}
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.warn("删除旧数据失败(可能是首次索引): {}", e.getMessage());
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建元数据(包含文件信息)
|
||||
*/
|
||||
private Map<String, Object> buildMetadata(String filePath, DocumentChunk chunk, int totalChunks) {
|
||||
Map<String, Object> metadata = new HashMap<>();
|
||||
|
||||
// 标准化路径:使用统一的路径分隔符(正斜杠)用于存储,确保跨平台一致性
|
||||
Path path = Paths.get(filePath).normalize();
|
||||
String normalizedPath = path.toString().replace(File.separator, "/");
|
||||
|
||||
// 文件信息
|
||||
Path fileName = path.getFileName();
|
||||
String fileNameStr = fileName != null ? fileName.toString() : "";
|
||||
String extension = "";
|
||||
int dotIndex = fileNameStr.lastIndexOf('.');
|
||||
if (dotIndex > 0) {
|
||||
extension = fileNameStr.substring(dotIndex);
|
||||
}
|
||||
|
||||
metadata.put("_source", normalizedPath);
|
||||
metadata.put("_extension", extension);
|
||||
metadata.put("_file_name", fileNameStr);
|
||||
|
||||
// 分片信息
|
||||
metadata.put("chunkIndex", chunk.getChunkIndex());
|
||||
metadata.put("totalChunks", totalChunks);
|
||||
|
||||
// 标题信息
|
||||
if (chunk.getTitle() != null && !chunk.getTitle().isEmpty()) {
|
||||
metadata.put("title", chunk.getTitle());
|
||||
}
|
||||
|
||||
return metadata;
|
||||
}
|
||||
|
||||
/**
|
||||
* 插入向量到 Milvus
|
||||
*/
|
||||
private void insertToMilvus(String content, List<Float> vector,
|
||||
Map<String, Object> metadata, int chunkIndex) throws Exception {
|
||||
try {
|
||||
// 确保 collection 已加载
|
||||
R<RpcStatus> loadResponse = milvusClient.loadCollection(
|
||||
LoadCollectionParam.newBuilder()
|
||||
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
|
||||
.build()
|
||||
);
|
||||
|
||||
if (loadResponse.getStatus() != 0 && loadResponse.getStatus() != 65535) {
|
||||
throw new RuntimeException("加载 collection 失败: " + loadResponse.getMessage());
|
||||
}
|
||||
|
||||
// 生成唯一 ID(使用 _source + 分片索引)
|
||||
String source = (String) metadata.get("_source");
|
||||
String id = UUID.nameUUIDFromBytes((source + "_" + chunkIndex).getBytes()).toString();
|
||||
|
||||
// 构建字段数据
|
||||
List<InsertParam.Field> fields = new ArrayList<>();
|
||||
|
||||
// ID 字段
|
||||
fields.add(new InsertParam.Field("id", Collections.singletonList(id)));
|
||||
|
||||
// content 字段
|
||||
fields.add(new InsertParam.Field("content", Collections.singletonList(content)));
|
||||
|
||||
// vector 字段
|
||||
fields.add(new InsertParam.Field("vector", Collections.singletonList(vector)));
|
||||
|
||||
// metadata 字段(JSON 对象)
|
||||
com.google.gson.Gson gson = new com.google.gson.Gson();
|
||||
com.google.gson.JsonObject metadataJson = gson.toJsonTree(metadata).getAsJsonObject();
|
||||
fields.add(new InsertParam.Field("metadata", Collections.singletonList(metadataJson)));
|
||||
|
||||
// 构建插入参数
|
||||
InsertParam insertParam = InsertParam.newBuilder()
|
||||
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
|
||||
.withFields(fields)
|
||||
.build();
|
||||
|
||||
// 执行插入
|
||||
R<MutationResult> insertResponse = milvusClient.insert(insertParam);
|
||||
|
||||
if (insertResponse.getStatus() != 0) {
|
||||
throw new RuntimeException("插入向量失败: " + insertResponse.getMessage());
|
||||
}
|
||||
|
||||
logger.debug("向量插入成功: id={}, source={}, chunk={}", id, source, chunkIndex);
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("插入向量到 Milvus 失败", e);
|
||||
throw e;
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 索引结果类
|
||||
*/
|
||||
@Getter
|
||||
public static class IndexingResult {
|
||||
@Setter
|
||||
private boolean success;
|
||||
@Setter
|
||||
private String directoryPath;
|
||||
@Setter
|
||||
private int totalFiles;
|
||||
private int successCount;
|
||||
private int failCount;
|
||||
@Setter
|
||||
private LocalDateTime startTime;
|
||||
@Setter
|
||||
private LocalDateTime endTime;
|
||||
@Setter
|
||||
private String errorMessage;
|
||||
private Map<String, String> failedFiles = new HashMap<>();
|
||||
|
||||
public void incrementSuccessCount() {
|
||||
this.successCount++;
|
||||
}
|
||||
|
||||
public void incrementFailCount() {
|
||||
this.failCount++;
|
||||
}
|
||||
|
||||
public long getDurationMs() {
|
||||
if (startTime != null && endTime != null) {
|
||||
return java.time.Duration.between(startTime, endTime).toMillis();
|
||||
}
|
||||
return 0;
|
||||
}
|
||||
|
||||
public void addFailedFile(String filePath, String error) {
|
||||
this.failedFiles.put(filePath, error);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,108 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import io.milvus.client.MilvusServiceClient;
|
||||
import io.milvus.grpc.SearchResults;
|
||||
import io.milvus.param.R;
|
||||
import io.milvus.param.dml.SearchParam;
|
||||
import io.milvus.response.SearchResultsWrapper;
|
||||
import lombok.Getter;
|
||||
import lombok.Setter;
|
||||
import com.superbiz.agent.constant.MilvusConstants;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
import org.springframework.beans.factory.annotation.Autowired;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.ArrayList;
|
||||
import java.util.Collections;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* 向量搜索服务
|
||||
* 负责从 Milvus 中搜索相似向量
|
||||
*/
|
||||
@Service
|
||||
public class VectorSearchService {
|
||||
|
||||
private static final Logger logger = LoggerFactory.getLogger(VectorSearchService.class);
|
||||
|
||||
@Autowired
|
||||
private MilvusServiceClient milvusClient;
|
||||
|
||||
@Autowired
|
||||
private VectorEmbeddingService embeddingService;
|
||||
|
||||
/**
|
||||
* 搜索相似文档
|
||||
*
|
||||
* @param query 查询文本
|
||||
* @param topK 返回最相似的K个结果
|
||||
* @return 搜索结果列表
|
||||
*/
|
||||
public List<SearchResult> searchSimilarDocuments(String query, int topK) {
|
||||
try {
|
||||
logger.info("开始搜索相似文档, 查询: {}, topK: {}", query, topK);
|
||||
|
||||
// 1. 将查询文本向量化
|
||||
List<Float> queryVector = embeddingService.generateQueryVector(query);
|
||||
logger.debug("查询向量生成成功, 维度: {}", queryVector.size());
|
||||
|
||||
// 2. 构建搜索参数
|
||||
SearchParam searchParam = SearchParam.newBuilder()
|
||||
.withCollectionName(MilvusConstants.MILVUS_COLLECTION_NAME)
|
||||
.withVectorFieldName("vector")
|
||||
.withVectors(Collections.singletonList(queryVector))
|
||||
.withTopK(topK)
|
||||
.withMetricType(io.milvus.param.MetricType.L2)
|
||||
.withOutFields(List.of("id", "content", "metadata"))
|
||||
.withParams("{\"nprobe\":10}")
|
||||
.build();
|
||||
|
||||
// 3. 执行搜索
|
||||
R<SearchResults> searchResponse = milvusClient.search(searchParam);
|
||||
|
||||
if (searchResponse.getStatus() != 0) {
|
||||
throw new RuntimeException("向量搜索失败: " + searchResponse.getMessage());
|
||||
}
|
||||
|
||||
// 4. 解析搜索结果
|
||||
SearchResultsWrapper wrapper = new SearchResultsWrapper(searchResponse.getData().getResults());
|
||||
List<SearchResult> results = new ArrayList<>();
|
||||
|
||||
for (int i = 0; i < wrapper.getRowRecords(0).size(); i++) {
|
||||
SearchResult result = new SearchResult();
|
||||
result.setId((String) wrapper.getIDScore(0).get(i).get("id"));
|
||||
result.setContent((String) wrapper.getFieldData("content", 0).get(i));
|
||||
result.setScore(wrapper.getIDScore(0).get(i).getScore());
|
||||
|
||||
// 解析 metadata
|
||||
Object metadataObj = wrapper.getFieldData("metadata", 0).get(i);
|
||||
if (metadataObj != null) {
|
||||
result.setMetadata(metadataObj.toString());
|
||||
}
|
||||
|
||||
results.add(result);
|
||||
}
|
||||
|
||||
logger.info("搜索完成, 找到 {} 个相似文档", results.size());
|
||||
return results;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("搜索相似文档失败", e);
|
||||
throw new RuntimeException("搜索失败: " + e.getMessage(), e);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 搜索结果类
|
||||
*/
|
||||
@Setter
|
||||
@Getter
|
||||
public static class SearchResult {
|
||||
private String id;
|
||||
private String content;
|
||||
private float score;
|
||||
private String metadata;
|
||||
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package com.superbiz.agent.service.session;
|
||||
|
||||
import com.superbiz.agent.domain.model.SessionContext;
|
||||
import com.superbiz.agent.domain.model.ToolCall;
|
||||
|
||||
import java.util.Optional;
|
||||
|
||||
/**
|
||||
* 会话管理器接口
|
||||
* 负责会话的创建、读取、更新和删除
|
||||
*/
|
||||
public interface SessionManager {
|
||||
|
||||
/**
|
||||
* 创建新会话
|
||||
*
|
||||
* @param sessionContext 会话上下文
|
||||
* @param ttlSeconds 会话过期时间(秒)
|
||||
* @return 会话ID
|
||||
*/
|
||||
String createSession(SessionContext sessionContext, long ttlSeconds);
|
||||
|
||||
/**
|
||||
* 获取会话
|
||||
*
|
||||
* @param sessionId 会话ID
|
||||
* @return 会话上下文(如果存在)
|
||||
*/
|
||||
Optional<SessionContext> getSession(String sessionId);
|
||||
|
||||
/**
|
||||
* 更新会话
|
||||
*
|
||||
* @param sessionContext 会话上下文
|
||||
*/
|
||||
void updateSession(SessionContext sessionContext);
|
||||
|
||||
/**
|
||||
* 删除会话
|
||||
*
|
||||
* @param sessionId 会话ID
|
||||
*/
|
||||
void deleteSession(String sessionId);
|
||||
|
||||
/**
|
||||
* 检查会话是否存在
|
||||
*
|
||||
* @param sessionId 会话ID
|
||||
* @return true 如果会话存在
|
||||
*/
|
||||
boolean exists(String sessionId);
|
||||
|
||||
/**
|
||||
* 刷新会话过期时间
|
||||
*
|
||||
* @param sessionId 会话ID
|
||||
* @param ttlSeconds 新的过期时间(秒)
|
||||
* @return true 如果刷新成功
|
||||
*/
|
||||
boolean refreshSession(String sessionId, long ttlSeconds);
|
||||
|
||||
/**
|
||||
* 添加工具调用记录到会话
|
||||
*
|
||||
* @param sessionId 会话ID
|
||||
* @param toolCall 工具调用记录
|
||||
*/
|
||||
void addToolCall(String sessionId, ToolCall toolCall);
|
||||
|
||||
/**
|
||||
* 更新会话状态
|
||||
*
|
||||
* @param sessionId 会话ID
|
||||
* @param status 新状态
|
||||
*/
|
||||
void updateStatus(String sessionId, String status);
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package com.superbiz.agent.service.session.impl;
|
||||
|
||||
import lombok.RequiredArgsConstructor;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import com.superbiz.agent.domain.model.SessionContext;
|
||||
import com.superbiz.agent.domain.model.ToolCall;
|
||||
import com.superbiz.agent.service.session.SessionManager;
|
||||
import org.springframework.data.redis.core.RedisTemplate;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.time.LocalDateTime;
|
||||
import java.util.Optional;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
|
||||
/**
|
||||
* Redis 会话管理器实现
|
||||
*/
|
||||
@Slf4j
|
||||
@Service
|
||||
@RequiredArgsConstructor
|
||||
public class RedisSessionManager implements SessionManager {
|
||||
|
||||
private static final String SESSION_KEY_PREFIX = "session:";
|
||||
|
||||
private final RedisTemplate<String, Object> redisTemplate;
|
||||
|
||||
@Override
|
||||
public String createSession(SessionContext sessionContext, long ttlSeconds) {
|
||||
String sessionId = sessionContext.getSessionId();
|
||||
if (sessionId == null || sessionId.isEmpty()) {
|
||||
throw new IllegalArgumentException("Session ID cannot be null or empty");
|
||||
}
|
||||
|
||||
sessionContext.setCreatedAt(LocalDateTime.now());
|
||||
sessionContext.setLastActiveAt(LocalDateTime.now());
|
||||
sessionContext.setTtl(ttlSeconds);
|
||||
sessionContext.setStatus("ACTIVE");
|
||||
|
||||
String key = buildKey(sessionId);
|
||||
redisTemplate.opsForValue().set(key, sessionContext, ttlSeconds, TimeUnit.SECONDS);
|
||||
|
||||
log.info("创建会话成功: sessionId={}, ttl={}秒", sessionId, ttlSeconds);
|
||||
return sessionId;
|
||||
}
|
||||
|
||||
@Override
|
||||
public Optional<SessionContext> getSession(String sessionId) {
|
||||
String key = buildKey(sessionId);
|
||||
Object value = redisTemplate.opsForValue().get(key);
|
||||
|
||||
if (value instanceof SessionContext) {
|
||||
SessionContext context = (SessionContext) value;
|
||||
log.debug("获取会话成功: sessionId={}", sessionId);
|
||||
return Optional.of(context);
|
||||
}
|
||||
|
||||
log.debug("会话不存在: sessionId={}", sessionId);
|
||||
return Optional.empty();
|
||||
}
|
||||
|
||||
@Override
|
||||
public void updateSession(SessionContext sessionContext) {
|
||||
String sessionId = sessionContext.getSessionId();
|
||||
String key = buildKey(sessionId);
|
||||
|
||||
// 获取剩余 TTL
|
||||
Long ttl = redisTemplate.getExpire(key, TimeUnit.SECONDS);
|
||||
if (ttl == null || ttl <= 0) {
|
||||
ttl = sessionContext.getTtl() != null ? sessionContext.getTtl() : 3600L;
|
||||
}
|
||||
|
||||
sessionContext.setLastActiveAt(LocalDateTime.now());
|
||||
redisTemplate.opsForValue().set(key, sessionContext, ttl, TimeUnit.SECONDS);
|
||||
|
||||
log.debug("更新会话成功: sessionId={}", sessionId);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void deleteSession(String sessionId) {
|
||||
String key = buildKey(sessionId);
|
||||
Boolean deleted = redisTemplate.delete(key);
|
||||
|
||||
if (Boolean.TRUE.equals(deleted)) {
|
||||
log.info("删除会话成功: sessionId={}", sessionId);
|
||||
} else {
|
||||
log.warn("删除会话失败,会话可能不存在: sessionId={}", sessionId);
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean exists(String sessionId) {
|
||||
String key = buildKey(sessionId);
|
||||
Boolean exists = redisTemplate.hasKey(key);
|
||||
return Boolean.TRUE.equals(exists);
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean refreshSession(String sessionId, long ttlSeconds) {
|
||||
String key = buildKey(sessionId);
|
||||
Boolean refreshed = redisTemplate.expire(key, ttlSeconds, TimeUnit.SECONDS);
|
||||
|
||||
if (Boolean.TRUE.equals(refreshed)) {
|
||||
log.debug("刷新会话过期时间成功: sessionId={}, newTtl={}秒", sessionId, ttlSeconds);
|
||||
return true;
|
||||
}
|
||||
|
||||
log.warn("刷新会话过期时间失败,会话可能不存在: sessionId={}", sessionId);
|
||||
return false;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void addToolCall(String sessionId, ToolCall toolCall) {
|
||||
Optional<SessionContext> sessionOpt = getSession(sessionId);
|
||||
|
||||
if (sessionOpt.isPresent()) {
|
||||
SessionContext context = sessionOpt.get();
|
||||
context.addToolCall(toolCall);
|
||||
updateSession(context);
|
||||
|
||||
log.debug("添加工具调用记录成功: sessionId={}, toolName={}", sessionId, toolCall.getToolName());
|
||||
} else {
|
||||
log.warn("会话不存在,无法添加工具调用记录: sessionId={}", sessionId);
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void updateStatus(String sessionId, String status) {
|
||||
Optional<SessionContext> sessionOpt = getSession(sessionId);
|
||||
|
||||
if (sessionOpt.isPresent()) {
|
||||
SessionContext context = sessionOpt.get();
|
||||
context.setStatus(status);
|
||||
updateSession(context);
|
||||
|
||||
log.debug("更新会话状态成功: sessionId={}, status={}", sessionId, status);
|
||||
} else {
|
||||
log.warn("会话不存在,无法更新状态: sessionId={}", sessionId);
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Redis key
|
||||
*/
|
||||
private String buildKey(String sessionId) {
|
||||
return SESSION_KEY_PREFIX + sessionId;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user