fix chat session traces and document paths
This commit is contained in:
@@ -2,6 +2,7 @@ package com.superbiz.agent.agent.tool;
|
||||
|
||||
import com.fasterxml.jackson.annotation.JsonProperty;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import com.superbiz.agent.service.ToolInvocationRecorder;
|
||||
import lombok.Data;
|
||||
import org.slf4j.Logger;
|
||||
import org.slf4j.LoggerFactory;
|
||||
@@ -34,6 +35,11 @@ public class QueryLogsTools {
|
||||
public static final String TOOL_GET_AVAILABLE_LOG_TOPICS = "getAvailableLogTopics";
|
||||
|
||||
private final ObjectMapper objectMapper = new ObjectMapper();
|
||||
private final ToolInvocationRecorder toolInvocationRecorder;
|
||||
|
||||
public QueryLogsTools(ToolInvocationRecorder toolInvocationRecorder) {
|
||||
this.toolInvocationRecorder = toolInvocationRecorder;
|
||||
}
|
||||
|
||||
@Value("${cls.mock-enabled:false}")
|
||||
private boolean mockEnabled;
|
||||
@@ -55,6 +61,7 @@ public class QueryLogsTools {
|
||||
"Call this tool first before querying logs to understand what log topics are available. " +
|
||||
"Returns a list of log topics with their names, descriptions, and example queries.")
|
||||
public String getAvailableLogTopics() {
|
||||
long startTime = System.currentTimeMillis();
|
||||
logger.info("获取可用的日志主题列表");
|
||||
|
||||
try {
|
||||
@@ -123,11 +130,15 @@ public class QueryLogsTools {
|
||||
|
||||
output.setMessage(String.format("共有 %d 个可用的日志主题。建议使用默认地域 'ap-guangzhou' 或省略 region 参数", topics.size()));
|
||||
|
||||
return objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
|
||||
String response = objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
|
||||
recordInvocation(startTime, "get_available_log_topics", null, null, null, response, true, null, "logs");
|
||||
return response;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("获取日志主题列表失败", e);
|
||||
return "{\"success\":false,\"message\":\"获取日志主题列表失败: " + e.getMessage() + "\"}";
|
||||
String response = "{\"success\":false,\"message\":\"获取日志主题列表失败: " + e.getMessage() + "\"}";
|
||||
recordInvocation(startTime, "get_available_log_topics", null, null, null, response, false, e.getMessage(), "logs");
|
||||
return response;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -164,6 +175,7 @@ public class QueryLogsTools {
|
||||
@ToolParam(description = "查询条件,支持 Lucene 语法,如 level:ERROR OR cpu_usage:>80;为空时返回该主题近 5 条核心日志") String query,
|
||||
@ToolParam(description = "返回日志条数,默认20,最大100") Integer limit) {
|
||||
|
||||
long startTime = System.currentTimeMillis();
|
||||
int actualLimit = (limit == null || limit <= 0) ? 20 : Math.min(limit, 100);
|
||||
|
||||
String safeQuery = query == null ? "" : query;
|
||||
@@ -178,7 +190,10 @@ public class QueryLogsTools {
|
||||
logger.info("使用 Mock 数据,返回 {} 条日志", logEntries.size());
|
||||
} else {
|
||||
// 真实模式:调用 CLS API(这里预留接口,后续实现)
|
||||
return buildErrorResponse("CLS 真实查询尚未实现,请启用 mock 模式进行测试");
|
||||
String response = buildErrorResponse("CLS 真实查询尚未实现,请启用 mock 模式进行测试");
|
||||
recordInvocation(startTime, safeQuery, region, logTopic, actualLimit, response, false,
|
||||
"CLS 真实查询尚未实现,请启用 mock 模式进行测试", normalizeTopicDomain(logTopic));
|
||||
return response;
|
||||
}
|
||||
|
||||
// 构建成功响应
|
||||
@@ -193,15 +208,51 @@ public class QueryLogsTools {
|
||||
|
||||
String jsonResult = objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
|
||||
logger.info("日志查询完成: 找到 {} 条日志", logEntries.size());
|
||||
recordInvocation(startTime, safeQuery, region, logTopic, actualLimit, jsonResult,
|
||||
!logEntries.isEmpty(), logEntries.isEmpty() ? "未找到匹配的日志" : null,
|
||||
normalizeTopicDomain(logTopic));
|
||||
|
||||
return jsonResult;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("查询日志失败", e);
|
||||
return buildErrorResponse("查询失败: " + e.getMessage());
|
||||
String response = buildErrorResponse("查询失败: " + e.getMessage());
|
||||
recordInvocation(startTime, safeQuery, region, logTopic, actualLimit, response, false,
|
||||
e.getMessage(), normalizeTopicDomain(logTopic));
|
||||
return response;
|
||||
}
|
||||
}
|
||||
|
||||
private void recordInvocation(long startTime, String query, String region, String logTopic, Integer limit,
|
||||
String output, boolean success, String errorMessage, String topicDomain) {
|
||||
Map<String, Object> input = new HashMap<>();
|
||||
input.put("query", query == null || query.isBlank() ? "DEFAULT_QUERY" : query);
|
||||
if (region != null) {
|
||||
input.put("region", region);
|
||||
}
|
||||
if (logTopic != null) {
|
||||
input.put("log_topic", logTopic);
|
||||
}
|
||||
if (limit != null) {
|
||||
input.put("limit", limit);
|
||||
}
|
||||
input.put("mock_enabled", mockEnabled);
|
||||
|
||||
toolInvocationRecorder.recordEvidenceTool(
|
||||
"query_logs",
|
||||
input,
|
||||
output,
|
||||
success,
|
||||
startTime,
|
||||
errorMessage,
|
||||
topicDomain
|
||||
);
|
||||
}
|
||||
|
||||
private String normalizeTopicDomain(String logTopic) {
|
||||
return logTopic == null || logTopic.isBlank() ? "logs" : logTopic;
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Mock 日志数据
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package com.superbiz.agent.agent.tool;
|
||||
|
||||
import com.fasterxml.jackson.annotation.JsonProperty;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import com.superbiz.agent.service.ToolInvocationRecorder;
|
||||
import lombok.Data;
|
||||
import okhttp3.OkHttpClient;
|
||||
import okhttp3.Request;
|
||||
@@ -30,6 +31,11 @@ public class QueryMetricsTools {
|
||||
public static final String TOOL_QUERY_PROMETHEUS_ALERTS = "queryPrometheusAlerts";
|
||||
|
||||
private final ObjectMapper objectMapper = new ObjectMapper();
|
||||
private final ToolInvocationRecorder toolInvocationRecorder;
|
||||
|
||||
public QueryMetricsTools(ToolInvocationRecorder toolInvocationRecorder) {
|
||||
this.toolInvocationRecorder = toolInvocationRecorder;
|
||||
}
|
||||
|
||||
@Value("${prometheus.base-url}")
|
||||
private String prometheusBaseUrl;
|
||||
@@ -59,6 +65,7 @@ public class QueryMetricsTools {
|
||||
"This tool retrieves all currently active/firing alerts including their labels, annotations, state, and values. " +
|
||||
"Use this tool when you need to check what alerts are currently firing, investigate alert conditions, or monitor alert status.")
|
||||
public String queryPrometheusAlerts() {
|
||||
long startTime = System.currentTimeMillis();
|
||||
logger.info("开始查询 Prometheus 活动告警, Mock模式: {}", mockEnabled);
|
||||
|
||||
try {
|
||||
@@ -73,7 +80,9 @@ public class QueryMetricsTools {
|
||||
PrometheusAlertsResult result = fetchPrometheusAlerts();
|
||||
|
||||
if (!"success".equals(result.getStatus())) {
|
||||
return buildErrorResponse("Prometheus API 返回非成功状态: " + result.getStatus(), result.getError());
|
||||
String response = buildErrorResponse("Prometheus API 返回非成功状态: " + result.getStatus(), result.getError());
|
||||
recordInvocation(startTime, response, false, result.getError());
|
||||
return response;
|
||||
}
|
||||
|
||||
// 转换为简化格式,对于相同的 alertname,只保留第一个
|
||||
@@ -110,14 +119,29 @@ public class QueryMetricsTools {
|
||||
|
||||
String jsonResult = objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(output);
|
||||
logger.info("Prometheus 告警查询完成: 找到 {} 个告警", simplifiedAlerts.size());
|
||||
recordInvocation(startTime, jsonResult, true, null);
|
||||
|
||||
return jsonResult;
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("查询 Prometheus 告警失败", e);
|
||||
return buildErrorResponse("查询失败", e.getMessage());
|
||||
String response = buildErrorResponse("查询失败", e.getMessage());
|
||||
recordInvocation(startTime, response, false, e.getMessage());
|
||||
return response;
|
||||
}
|
||||
}
|
||||
|
||||
private void recordInvocation(long startTime, String output, boolean success, String errorMessage) {
|
||||
toolInvocationRecorder.recordEvidenceTool(
|
||||
"query_metrics",
|
||||
Map.of("query", "active_prometheus_alerts", "mock_enabled", mockEnabled),
|
||||
output,
|
||||
success,
|
||||
startTime,
|
||||
errorMessage,
|
||||
"prometheus_alerts"
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* 构建 Mock 告警数据
|
||||
|
||||
@@ -1,32 +1,30 @@
|
||||
package com.superbiz.agent.controller;
|
||||
|
||||
import com.alibaba.cloud.ai.graph.NodeOutput;
|
||||
import com.alibaba.cloud.ai.graph.OverAllState;
|
||||
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
|
||||
import com.alibaba.cloud.ai.graph.streaming.OutputType;
|
||||
import com.alibaba.cloud.ai.graph.streaming.StreamingOutput;
|
||||
import lombok.Getter;
|
||||
import lombok.Setter;
|
||||
import com.superbiz.agent.domain.model.SessionContext;
|
||||
import com.superbiz.agent.service.AiOpsService;
|
||||
import com.superbiz.agent.service.ChatService;
|
||||
import com.superbiz.agent.service.session.SessionManager;
|
||||
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.beans.factory.annotation.Value;
|
||||
import org.springframework.http.MediaType;
|
||||
import org.springframework.http.ResponseEntity;
|
||||
import org.springframework.web.bind.annotation.*;
|
||||
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
|
||||
import reactor.core.publisher.Flux;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.time.LocalDateTime;
|
||||
import java.time.ZoneId;
|
||||
import java.util.*;
|
||||
import java.util.concurrent.ConcurrentHashMap;
|
||||
import java.util.concurrent.ExecutorService;
|
||||
import java.util.concurrent.Executors;
|
||||
import java.util.concurrent.locks.ReentrantLock;
|
||||
|
||||
/**
|
||||
* 统一 API 控制器
|
||||
@@ -44,17 +42,20 @@ public class ChatController {
|
||||
@Autowired
|
||||
private ChatService chatService;
|
||||
|
||||
@Autowired
|
||||
private SessionManager sessionManager;
|
||||
|
||||
@Autowired(required = false)
|
||||
private ToolCallbackProvider tools;
|
||||
|
||||
private final ExecutorService executor = Executors.newCachedThreadPool();
|
||||
|
||||
// 存储会话信息
|
||||
private final Map<String, SessionInfo> sessions = new ConcurrentHashMap<>();
|
||||
|
||||
// 最大历史消息窗口大小(成对计算:用户消息+AI回复=1对)
|
||||
private static final int MAX_WINDOW_SIZE = 6;
|
||||
|
||||
@Value("${session.ttl-seconds:3600}")
|
||||
private long sessionTtlSeconds;
|
||||
|
||||
/**
|
||||
* 普通对话接口(支持工具调用)
|
||||
* 与 /chat_react 逻辑一致,但直接返回完整结果而非流式输出
|
||||
@@ -71,10 +72,10 @@ public class ChatController {
|
||||
}
|
||||
|
||||
// 获取或创建会话
|
||||
SessionInfo session = getOrCreateSession(request.getId());
|
||||
SessionContext session = getOrCreateSession(request.getId());
|
||||
|
||||
// 获取历史消息
|
||||
List<Map<String, String>> history = session.getHistory();
|
||||
List<Map<String, String>> history = session.getMessageHistorySnapshot();
|
||||
logger.info("会话历史消息对数: {}", history.size() / 2);
|
||||
|
||||
// 获取注入的 ChatModel
|
||||
@@ -88,13 +89,14 @@ public class ChatController {
|
||||
// 根据问题复杂度自动选择单 Agent 或多 Agent
|
||||
logger.info("开始 ReactAgent 对话(支持自动工具调用)");
|
||||
ChatService.ChatResult result = chatService.executeChatWithStrategy(chatModel, toolCallbacks,
|
||||
request.getQuestion(), history);
|
||||
request.getQuestion(), history, session.getSessionId());
|
||||
String fullAnswer = result.answer();
|
||||
|
||||
// 更新会话历史
|
||||
session.addMessage(request.getQuestion(), fullAnswer);
|
||||
session.addChatMessagePair(request.getQuestion(), fullAnswer, MAX_WINDOW_SIZE);
|
||||
sessionManager.updateSession(session);
|
||||
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
|
||||
request.getId(), session.getMessagePairCount());
|
||||
session.getSessionId(), session.getMessagePairCount());
|
||||
|
||||
return ResponseEntity.ok(ApiResponse.success(ChatResponse.success(fullAnswer, result.sessionId())));
|
||||
|
||||
@@ -116,9 +118,11 @@ public class ChatController {
|
||||
return ResponseEntity.ok(ApiResponse.error("会话ID不能为空"));
|
||||
}
|
||||
|
||||
SessionInfo session = sessions.get(request.getId());
|
||||
if (session != null) {
|
||||
session.clearHistory();
|
||||
Optional<SessionContext> session = sessionManager.getSession(request.getId());
|
||||
if (session.isPresent()) {
|
||||
SessionContext context = session.get();
|
||||
context.clearMessageHistory();
|
||||
sessionManager.updateSession(context);
|
||||
return ResponseEntity.ok(ApiResponse.success("会话历史已清空"));
|
||||
} else {
|
||||
return ResponseEntity.ok(ApiResponse.error("会话不存在"));
|
||||
@@ -131,8 +135,8 @@ public class ChatController {
|
||||
}
|
||||
|
||||
/**
|
||||
* ReactAgent 对话接口(SSE 流式模式,支持多轮对话,支持自动工具调用,例如获取当前时间,查询日志,告警等)
|
||||
* 支持 session 管理,保留对话历史
|
||||
* 对话接口(SSE 流式模式)
|
||||
* 与 /chat 使用同一条 ChatService 策略链路,区别仅在于通过 SSE 分块返回最终答案。
|
||||
*/
|
||||
@PostMapping(value = "/chat_stream", produces = "text/event-stream;charset=UTF-8")
|
||||
public SseEmitter chatStream(@RequestBody ChatRequest request) {
|
||||
@@ -155,10 +159,10 @@ public class ChatController {
|
||||
logger.info("收到 ReactAgent 对话请求 - SessionId: {}, Question: {}", request.getId(), request.getQuestion());
|
||||
|
||||
// 获取或创建会话
|
||||
SessionInfo session = getOrCreateSession(request.getId());
|
||||
SessionContext session = getOrCreateSession(request.getId());
|
||||
|
||||
// 获取历史消息
|
||||
List<Map<String, String>> history = session.getHistory();
|
||||
List<Map<String, String>> history = session.getMessageHistorySnapshot();
|
||||
logger.info("ReactAgent 会话历史消息对数: {}", history.size() / 2);
|
||||
|
||||
// 获取注入的 ChatModel
|
||||
@@ -167,92 +171,25 @@ public class ChatController {
|
||||
// 记录可用工具
|
||||
chatService.logAvailableTools();
|
||||
|
||||
logger.info("开始 ReactAgent 流式对话(支持自动工具调用)");
|
||||
|
||||
// 构建系统提示词(包含历史消息)
|
||||
String systemPrompt = chatService.buildSystemPrompt(history);
|
||||
|
||||
// 创建 ReactAgent
|
||||
ReactAgent agent = chatService.createReactAgent(chatModel, systemPrompt);
|
||||
|
||||
// 用于累积完整答案
|
||||
StringBuilder fullAnswerBuilder = new StringBuilder();
|
||||
|
||||
// 使用 agent.stream() 进行流式对话
|
||||
Flux<NodeOutput> stream = agent.stream(request.getQuestion());
|
||||
|
||||
stream.subscribe(
|
||||
output -> {
|
||||
try {
|
||||
// 检查是否为 StreamingOutput 类型
|
||||
if (output instanceof StreamingOutput streamingOutput) {
|
||||
OutputType type = streamingOutput.getOutputType();
|
||||
|
||||
// 处理模型推理的流式输出
|
||||
if (type == OutputType.AGENT_MODEL_STREAMING) {
|
||||
// 流式增量内容,逐步显示
|
||||
String chunk = streamingOutput.message().getText();
|
||||
if (chunk != null && !chunk.isEmpty()) {
|
||||
fullAnswerBuilder.append(chunk);
|
||||
|
||||
// 实时发送到前端
|
||||
emitter.send(SseEmitter.event()
|
||||
.name("message")
|
||||
.data(SseMessage.content(chunk), MediaType.APPLICATION_JSON));
|
||||
|
||||
logger.info("发送流式内容: {}", chunk);
|
||||
}
|
||||
} else if (type == OutputType.AGENT_MODEL_FINISHED) {
|
||||
// 模型推理完成
|
||||
logger.info("模型输出完成");
|
||||
} else if (type == OutputType.AGENT_TOOL_FINISHED) {
|
||||
// 工具调用完成
|
||||
logger.info("工具调用完成: {}", output.node());
|
||||
} else if (type == OutputType.AGENT_HOOK_FINISHED) {
|
||||
// Hook 执行完成
|
||||
logger.debug("Hook 执行完成: {}", output.node());
|
||||
}
|
||||
}
|
||||
} catch (IOException e) {
|
||||
logger.error("发送流式消息失败", e);
|
||||
throw new RuntimeException(e);
|
||||
}
|
||||
},
|
||||
error -> {
|
||||
// 错误处理
|
||||
logger.error("ReactAgent 流式对话失败", error);
|
||||
try {
|
||||
emitter.send(SseEmitter.event()
|
||||
.name("message")
|
||||
.data(SseMessage.error(error.getMessage()), MediaType.APPLICATION_JSON));
|
||||
} catch (IOException ex) {
|
||||
logger.error("发送错误消息失败", ex);
|
||||
}
|
||||
emitter.completeWithError(error);
|
||||
},
|
||||
() -> {
|
||||
// 完成处理
|
||||
try {
|
||||
String fullAnswer = fullAnswerBuilder.toString();
|
||||
logger.info("ReactAgent 流式对话完成 - SessionId: {}, 答案长度: {}",
|
||||
request.getId(), fullAnswer.length());
|
||||
|
||||
// 更新会话历史
|
||||
session.addMessage(request.getQuestion(), fullAnswer);
|
||||
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
|
||||
request.getId(), session.getMessagePairCount());
|
||||
|
||||
// 发送完成标记
|
||||
emitter.send(SseEmitter.event()
|
||||
.name("message")
|
||||
.data(SseMessage.done(), MediaType.APPLICATION_JSON));
|
||||
emitter.complete();
|
||||
} catch (IOException e) {
|
||||
logger.error("发送完成消息失败", e);
|
||||
emitter.completeWithError(e);
|
||||
}
|
||||
}
|
||||
);
|
||||
ToolCallback[] toolCallbacks = tools != null ? tools.getToolCallbacks() : new ToolCallback[0];
|
||||
|
||||
logger.info("开始统一 ChatService 对话(SSE 分块返回)");
|
||||
ChatService.ChatResult result = chatService.executeChatWithStrategy(chatModel, toolCallbacks,
|
||||
request.getQuestion(), history, session.getSessionId());
|
||||
String fullAnswer = result.answer() == null ? "" : result.answer();
|
||||
logger.info("统一 ChatService 对话完成 - SessionId: {}, 答案长度: {}",
|
||||
result.sessionId(), fullAnswer.length());
|
||||
|
||||
session.addChatMessagePair(request.getQuestion(), fullAnswer, MAX_WINDOW_SIZE);
|
||||
sessionManager.updateSession(session);
|
||||
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
|
||||
session.getSessionId(), session.getMessagePairCount());
|
||||
|
||||
sendContentChunks(emitter, fullAnswer);
|
||||
emitter.send(SseEmitter.event()
|
||||
.name("message")
|
||||
.data(SseMessage.done(), MediaType.APPLICATION_JSON));
|
||||
emitter.complete();
|
||||
|
||||
} catch (Exception e) {
|
||||
logger.error("ReactAgent 对话初始化失败", e);
|
||||
@@ -365,12 +302,13 @@ public class ChatController {
|
||||
try {
|
||||
logger.info("收到获取会话信息请求 - SessionId: {}", sessionId);
|
||||
|
||||
SessionInfo session = sessions.get(sessionId);
|
||||
if (session != null) {
|
||||
Optional<SessionContext> session = sessionManager.getSession(sessionId);
|
||||
if (session.isPresent()) {
|
||||
SessionContext context = session.get();
|
||||
SessionInfoResponse response = new SessionInfoResponse();
|
||||
response.setSessionId(sessionId);
|
||||
response.setMessagePairCount(session.getMessagePairCount());
|
||||
response.setCreateTime(session.createTime);
|
||||
response.setMessagePairCount(context.getMessagePairCount());
|
||||
response.setCreateTime(toEpochMillis(context.getCreatedAt()));
|
||||
return ResponseEntity.ok(ApiResponse.success(response));
|
||||
} else {
|
||||
return ResponseEntity.ok(ApiResponse.error("会话不存在"));
|
||||
@@ -384,107 +322,39 @@ public class ChatController {
|
||||
|
||||
// ==================== 辅助方法 ====================
|
||||
|
||||
private SessionInfo getOrCreateSession(String sessionId) {
|
||||
if (sessionId == null || sessionId.isEmpty()) {
|
||||
sessionId = UUID.randomUUID().toString();
|
||||
}
|
||||
return sessions.computeIfAbsent(sessionId, SessionInfo::new);
|
||||
private SessionContext getOrCreateSession(String sessionId) {
|
||||
String resolvedSessionId = (sessionId == null || sessionId.isEmpty())
|
||||
? UUID.randomUUID().toString()
|
||||
: sessionId;
|
||||
return sessionManager.getSession(resolvedSessionId)
|
||||
.orElseGet(() -> {
|
||||
SessionContext context = SessionContext.builder()
|
||||
.sessionId(resolvedSessionId)
|
||||
.status("ACTIVE")
|
||||
.ttl(sessionTtlSeconds)
|
||||
.build();
|
||||
sessionManager.createSession(context, sessionTtlSeconds);
|
||||
return context;
|
||||
});
|
||||
}
|
||||
|
||||
// ==================== 内部类 ====================
|
||||
|
||||
/**
|
||||
* 会话信息
|
||||
* 管理单个会话的历史消息,支持自动清理和线程安全
|
||||
*/
|
||||
private static class SessionInfo {
|
||||
private final String sessionId;
|
||||
// 存储历史消息对:[{"role": "user", "content": "..."}, {"role": "assistant", "content": "..."}]
|
||||
private final List<Map<String, String>> messageHistory;
|
||||
private final long createTime;
|
||||
private final ReentrantLock lock;
|
||||
|
||||
public SessionInfo(String sessionId) {
|
||||
this.sessionId = sessionId;
|
||||
this.messageHistory = new ArrayList<>();
|
||||
this.createTime = System.currentTimeMillis();
|
||||
this.lock = new ReentrantLock();
|
||||
private long toEpochMillis(LocalDateTime time) {
|
||||
if (time == null) {
|
||||
return 0L;
|
||||
}
|
||||
return time.atZone(ZoneId.systemDefault()).toInstant().toEpochMilli();
|
||||
}
|
||||
|
||||
/**
|
||||
* 添加一对消息(用户问题 + AI回复)
|
||||
* 自动管理历史消息窗口大小
|
||||
*/
|
||||
public void addMessage(String userQuestion, String aiAnswer) {
|
||||
lock.lock();
|
||||
try {
|
||||
// 添加用户消息
|
||||
Map<String, String> userMsg = new HashMap<>();
|
||||
userMsg.put("role", "user");
|
||||
userMsg.put("content", userQuestion);
|
||||
messageHistory.add(userMsg);
|
||||
|
||||
// 添加AI回复
|
||||
Map<String, String> assistantMsg = new HashMap<>();
|
||||
assistantMsg.put("role", "assistant");
|
||||
assistantMsg.put("content", aiAnswer);
|
||||
messageHistory.add(assistantMsg);
|
||||
|
||||
// 自动清理:保持最多 MAX_WINDOW_SIZE 对消息
|
||||
// 每对消息包含2条记录(user + assistant)
|
||||
int maxMessages = MAX_WINDOW_SIZE * 2;
|
||||
while (messageHistory.size() > maxMessages) {
|
||||
// 成对删除最旧的消息(删除前2条)
|
||||
messageHistory.remove(0); // 删除最旧的用户消息
|
||||
if (!messageHistory.isEmpty()) {
|
||||
messageHistory.remove(0); // 删除对应的AI回复
|
||||
}
|
||||
}
|
||||
|
||||
logger.debug("会话 {} 更新历史消息,当前消息对数: {}",
|
||||
sessionId, messageHistory.size() / 2);
|
||||
|
||||
} finally {
|
||||
lock.unlock();
|
||||
}
|
||||
private void sendContentChunks(SseEmitter emitter, String content) throws IOException {
|
||||
if (content == null || content.isEmpty()) {
|
||||
return;
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取历史消息(线程安全)
|
||||
* 返回副本以避免并发修改
|
||||
*/
|
||||
public List<Map<String, String>> getHistory() {
|
||||
lock.lock();
|
||||
try {
|
||||
return new ArrayList<>(messageHistory);
|
||||
} finally {
|
||||
lock.unlock();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 清空历史消息
|
||||
*/
|
||||
public void clearHistory() {
|
||||
lock.lock();
|
||||
try {
|
||||
messageHistory.clear();
|
||||
logger.info("会话 {} 历史消息已清空", sessionId);
|
||||
} finally {
|
||||
lock.unlock();
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取当前消息对数
|
||||
*/
|
||||
public int getMessagePairCount() {
|
||||
lock.lock();
|
||||
try {
|
||||
return messageHistory.size() / 2;
|
||||
} finally {
|
||||
lock.unlock();
|
||||
}
|
||||
int chunkSize = 80;
|
||||
for (int i = 0; i < content.length(); i += chunkSize) {
|
||||
int end = Math.min(i + chunkSize, content.length());
|
||||
emitter.send(SseEmitter.event()
|
||||
.name("message")
|
||||
.data(SseMessage.content(content.substring(i, end)), MediaType.APPLICATION_JSON));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -8,7 +8,9 @@ import lombok.NoArgsConstructor;
|
||||
import java.io.Serializable;
|
||||
import java.time.LocalDateTime;
|
||||
import java.util.ArrayList;
|
||||
import java.util.HashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
/**
|
||||
* 会话上下文数据类
|
||||
@@ -53,6 +55,12 @@ public class SessionContext implements Serializable {
|
||||
@Builder.Default
|
||||
private List<ToolCall> toolCalls = new ArrayList<>();
|
||||
|
||||
/**
|
||||
* 聊天消息历史:[{"role":"user","content":"..."}, {"role":"assistant","content":"..."}]
|
||||
*/
|
||||
@Builder.Default
|
||||
private List<Map<String, String>> messageHistory = new ArrayList<>();
|
||||
|
||||
/**
|
||||
* 会话创建时间
|
||||
*/
|
||||
@@ -79,6 +87,64 @@ public class SessionContext implements Serializable {
|
||||
this.lastActiveAt = LocalDateTime.now();
|
||||
}
|
||||
|
||||
/**
|
||||
* 添加一对聊天消息,并按消息对数裁剪窗口。
|
||||
*/
|
||||
public void addChatMessagePair(String userQuestion, String assistantAnswer, int maxPairCount) {
|
||||
if (this.messageHistory == null) {
|
||||
this.messageHistory = new ArrayList<>();
|
||||
}
|
||||
|
||||
Map<String, String> userMessage = new HashMap<>();
|
||||
userMessage.put("role", "user");
|
||||
userMessage.put("content", userQuestion);
|
||||
this.messageHistory.add(userMessage);
|
||||
|
||||
Map<String, String> assistantMessage = new HashMap<>();
|
||||
assistantMessage.put("role", "assistant");
|
||||
assistantMessage.put("content", assistantAnswer);
|
||||
this.messageHistory.add(assistantMessage);
|
||||
|
||||
int maxMessages = Math.max(maxPairCount, 0) * 2;
|
||||
while (maxMessages > 0 && this.messageHistory.size() > maxMessages) {
|
||||
this.messageHistory.remove(0);
|
||||
if (!this.messageHistory.isEmpty()) {
|
||||
this.messageHistory.remove(0);
|
||||
}
|
||||
}
|
||||
this.lastActiveAt = LocalDateTime.now();
|
||||
}
|
||||
|
||||
/**
|
||||
* 获取聊天历史副本,避免调用方直接修改内部列表。
|
||||
*/
|
||||
public List<Map<String, String>> getMessageHistorySnapshot() {
|
||||
if (this.messageHistory == null || this.messageHistory.isEmpty()) {
|
||||
return new ArrayList<>();
|
||||
}
|
||||
List<Map<String, String>> snapshot = new ArrayList<>();
|
||||
for (Map<String, String> message : this.messageHistory) {
|
||||
snapshot.add(new HashMap<>(message));
|
||||
}
|
||||
return snapshot;
|
||||
}
|
||||
|
||||
/**
|
||||
* 清空聊天历史。
|
||||
*/
|
||||
public void clearMessageHistory() {
|
||||
if (this.messageHistory == null) {
|
||||
this.messageHistory = new ArrayList<>();
|
||||
} else {
|
||||
this.messageHistory.clear();
|
||||
}
|
||||
this.lastActiveAt = LocalDateTime.now();
|
||||
}
|
||||
|
||||
public int getMessagePairCount() {
|
||||
return this.messageHistory == null ? 0 : this.messageHistory.size() / 2;
|
||||
}
|
||||
|
||||
/**
|
||||
* 更新最后活跃时间
|
||||
*/
|
||||
|
||||
@@ -272,19 +272,18 @@ public class ChatService {
|
||||
* @return ChatResult(answer + sessionId)
|
||||
*/
|
||||
public ChatResult executeChat(ReactAgent agent, String question) throws GraphRunnerException {
|
||||
return executeChat(agent, question, null);
|
||||
}
|
||||
|
||||
public ChatResult executeChat(ReactAgent agent, String question, String requestedSessionId) throws GraphRunnerException {
|
||||
logger.info("========================================");
|
||||
logger.info("📝 用户问题: {}", question);
|
||||
|
||||
String sessionId = UUID.randomUUID().toString().substring(0, 8);
|
||||
String sessionId = resolveSessionId(requestedSessionId);
|
||||
long startTime = System.currentTimeMillis();
|
||||
|
||||
// 创建诊断会话
|
||||
DiagnosisSession session = DiagnosisSession.builder()
|
||||
.sessionId(sessionId)
|
||||
.query(question)
|
||||
.status("RUNNING")
|
||||
.agentFlow("CHAT")
|
||||
.build();
|
||||
// 创建或更新诊断会话
|
||||
DiagnosisSession session = startDiagnosisSession(sessionId, question);
|
||||
diagnosisSessionRepository.save(session);
|
||||
|
||||
// 设置 ThreadLocal 上下文(LookupKnowledgeTool 通过此获取 sessionId)
|
||||
@@ -335,14 +334,20 @@ public class ChatService {
|
||||
*/
|
||||
public ChatResult executeChatWithStrategy(ChatModel chatModel, ToolCallback[] toolCallbacks,
|
||||
String question, List<Map<String, String>> history) throws GraphRunnerException {
|
||||
return executeChatWithStrategy(chatModel, toolCallbacks, question, history, null);
|
||||
}
|
||||
|
||||
public ChatResult executeChatWithStrategy(ChatModel chatModel, ToolCallback[] toolCallbacks,
|
||||
String question, List<Map<String, String>> history,
|
||||
String requestedSessionId) throws GraphRunnerException {
|
||||
if (QuestionComplexity.isComplex(question)) {
|
||||
logger.info("📊 问题判定为复杂,使用多 Agent(Planner + Executor)执行");
|
||||
return executeChatComplex(chatModel, toolCallbacks, question, history);
|
||||
return executeChatComplex(chatModel, toolCallbacks, question, history, requestedSessionId);
|
||||
} else {
|
||||
logger.info("📊 问题判定为简单,使用单 Agent 执行");
|
||||
String systemPrompt = buildSystemPrompt(history);
|
||||
ReactAgent agent = createReactAgent(chatModel, systemPrompt);
|
||||
return executeChat(agent, question);
|
||||
return executeChat(agent, question, requestedSessionId);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -351,15 +356,16 @@ public class ChatService {
|
||||
*/
|
||||
public ChatResult executeChatComplex(ChatModel chatModel, ToolCallback[] toolCallbacks,
|
||||
String question, List<Map<String, String>> history) throws GraphRunnerException {
|
||||
String sessionId = UUID.randomUUID().toString().substring(0, 8);
|
||||
return executeChatComplex(chatModel, toolCallbacks, question, history, null);
|
||||
}
|
||||
|
||||
public ChatResult executeChatComplex(ChatModel chatModel, ToolCallback[] toolCallbacks,
|
||||
String question, List<Map<String, String>> history,
|
||||
String requestedSessionId) throws GraphRunnerException {
|
||||
String sessionId = resolveSessionId(requestedSessionId);
|
||||
long startTime = System.currentTimeMillis();
|
||||
|
||||
DiagnosisSession session = DiagnosisSession.builder()
|
||||
.sessionId(sessionId)
|
||||
.query(question)
|
||||
.status("RUNNING")
|
||||
.agentFlow("CHAT")
|
||||
.build();
|
||||
DiagnosisSession session = startDiagnosisSession(sessionId, question);
|
||||
diagnosisSessionRepository.save(session);
|
||||
|
||||
SessionContextHolder.setSessionId(sessionId);
|
||||
@@ -528,6 +534,30 @@ public class ChatService {
|
||||
return agent.call(input, config).getText();
|
||||
}
|
||||
|
||||
private String resolveSessionId(String requestedSessionId) {
|
||||
if (requestedSessionId != null && !requestedSessionId.isBlank()) {
|
||||
return requestedSessionId;
|
||||
}
|
||||
return UUID.randomUUID().toString().substring(0, 8);
|
||||
}
|
||||
|
||||
private DiagnosisSession startDiagnosisSession(String sessionId, String question) {
|
||||
DiagnosisSession session = diagnosisSessionRepository.findBySessionId(sessionId)
|
||||
.orElseGet(() -> DiagnosisSession.builder()
|
||||
.sessionId(sessionId)
|
||||
.agentFlow("CHAT")
|
||||
.build());
|
||||
session.setQuery(question);
|
||||
session.setStatus("RUNNING");
|
||||
session.setAgentFlow("CHAT");
|
||||
session.setAnswer(null);
|
||||
session.setTotalDurationMs(null);
|
||||
session.setTotalTokenCount(null);
|
||||
session.setStepCount(null);
|
||||
session.setToolCallCount(null);
|
||||
return session;
|
||||
}
|
||||
|
||||
private String buildPlannerInput(String question, String retryContext) {
|
||||
if (retryContext == null || retryContext.isBlank()) {
|
||||
return question;
|
||||
|
||||
@@ -260,7 +260,8 @@ public class DocumentManagementService {
|
||||
private String saveToLocal(MultipartFile file, String fileName, String category) {
|
||||
try {
|
||||
// 1. 构建目标路径
|
||||
Path categoryDir = Paths.get(knowledgeBasePath, category);
|
||||
Path baseDir = Paths.get(knowledgeBasePath).normalize();
|
||||
Path categoryDir = baseDir.resolve(category).normalize();
|
||||
Files.createDirectories(categoryDir);
|
||||
|
||||
Path targetPath = categoryDir.resolve(fileName);
|
||||
@@ -268,8 +269,9 @@ public class DocumentManagementService {
|
||||
// 2. 保存文件
|
||||
file.transferTo(targetPath.toFile());
|
||||
|
||||
log.info("文件已保存到本地: {}", targetPath);
|
||||
return targetPath.toString();
|
||||
String relativePath = baseDir.relativize(targetPath.normalize()).toString().replace("\\", "/");
|
||||
log.info("文件已保存到本地: {}, storedPath={}", targetPath, relativePath);
|
||||
return relativePath;
|
||||
|
||||
} catch (IOException e) {
|
||||
throw new DocumentProcessException(
|
||||
@@ -287,7 +289,7 @@ public class DocumentManagementService {
|
||||
private void cleanupLocalFile(String localPath) {
|
||||
if (localPath != null) {
|
||||
try {
|
||||
Files.deleteIfExists(Paths.get(localPath));
|
||||
Files.deleteIfExists(resolveLocalPath(localPath));
|
||||
log.info("已清理本地文件: {}", localPath);
|
||||
} catch (IOException e) {
|
||||
log.warn("清理本地文件失败: {}", localPath, e);
|
||||
@@ -357,7 +359,7 @@ public class DocumentManagementService {
|
||||
// 删除本地文件
|
||||
if (doc.getFilePath() != null) {
|
||||
try {
|
||||
Files.deleteIfExists(Paths.get(doc.getFilePath()));
|
||||
Files.deleteIfExists(resolveLocalPath(doc.getFilePath()));
|
||||
log.info("本地文件已删除: {}", doc.getFilePath());
|
||||
} catch (IOException e) {
|
||||
log.warn("删除本地文件失败: {}", doc.getFilePath(), e);
|
||||
@@ -407,6 +409,26 @@ public class DocumentManagementService {
|
||||
return null;
|
||||
}
|
||||
|
||||
private Path resolveLocalPath(String filePath) {
|
||||
Path path = Paths.get(filePath).normalize();
|
||||
if (path.isAbsolute()) {
|
||||
return path;
|
||||
}
|
||||
|
||||
Path basePath = Paths.get(knowledgeBasePath).toAbsolutePath().normalize();
|
||||
Path baseName = basePath.getFileName();
|
||||
if (baseName != null && path.startsWith(baseName) && basePath.getParent() != null) {
|
||||
return basePath.getParent().resolve(path).normalize();
|
||||
}
|
||||
|
||||
Path pathFromWorkingDir = path.toAbsolutePath().normalize();
|
||||
if (pathFromWorkingDir.startsWith(basePath)) {
|
||||
return pathFromWorkingDir;
|
||||
}
|
||||
|
||||
return basePath.resolve(path).normalize();
|
||||
}
|
||||
|
||||
/**
|
||||
* 转换为响应 DTO
|
||||
*/
|
||||
|
||||
@@ -160,7 +160,12 @@ public class KnowledgeIndexService {
|
||||
|
||||
public String readDocument(String filePath, int maxChars) {
|
||||
try {
|
||||
Path fullPath = Paths.get(knowledgeBasePath, filePath);
|
||||
Path fullPath = resolveDocumentPath(filePath);
|
||||
if (!Files.exists(fullPath)) {
|
||||
log.warn("读取文档失败,文件不存在: basePath={}, filePath={}, resolvedPath={}",
|
||||
knowledgeBasePath, filePath, fullPath);
|
||||
return null;
|
||||
}
|
||||
String content = Files.readString(fullPath);
|
||||
|
||||
if (content.length() > maxChars) {
|
||||
@@ -170,11 +175,35 @@ public class KnowledgeIndexService {
|
||||
return content;
|
||||
|
||||
} catch (IOException e) {
|
||||
log.error("读取文档失败: {}/{}", knowledgeBasePath, filePath, e);
|
||||
log.error("读取文档失败: basePath={}, filePath={}", knowledgeBasePath, filePath, e);
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
Path resolveDocumentPath(String filePath) {
|
||||
if (filePath == null || filePath.isBlank()) {
|
||||
throw new IllegalArgumentException("filePath cannot be blank");
|
||||
}
|
||||
|
||||
Path path = Paths.get(filePath).normalize();
|
||||
if (path.isAbsolute()) {
|
||||
return path;
|
||||
}
|
||||
|
||||
Path basePath = Paths.get(knowledgeBasePath).toAbsolutePath().normalize();
|
||||
Path baseName = basePath.getFileName();
|
||||
if (baseName != null && path.startsWith(baseName) && basePath.getParent() != null) {
|
||||
return basePath.getParent().resolve(path).normalize();
|
||||
}
|
||||
|
||||
Path pathFromWorkingDir = path.toAbsolutePath().normalize();
|
||||
if (pathFromWorkingDir.startsWith(basePath)) {
|
||||
return pathFromWorkingDir;
|
||||
}
|
||||
|
||||
return basePath.resolve(path).normalize();
|
||||
}
|
||||
|
||||
public void addToIndex(KnowledgeEntry entry) {
|
||||
knowledgeIndex.add(entry);
|
||||
log.debug("文档已添加到 L0 索引: title={}", entry.getTitle());
|
||||
|
||||
@@ -0,0 +1,93 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import com.fasterxml.jackson.core.JsonProcessingException;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import com.superbiz.agent.domain.entity.ToolInvocation;
|
||||
import com.superbiz.agent.repository.ToolInvocationRepository;
|
||||
import com.superbiz.agent.util.SessionContextHolder;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
import org.springframework.stereotype.Service;
|
||||
|
||||
import java.util.LinkedHashMap;
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
import java.util.UUID;
|
||||
|
||||
/**
|
||||
* Central persistence point for agent evidence tool invocations.
|
||||
*/
|
||||
@Slf4j
|
||||
@Service
|
||||
public class ToolInvocationRecorder {
|
||||
|
||||
private static final int OUTPUT_PREVIEW_LIMIT = 500;
|
||||
|
||||
private final ToolInvocationRepository toolInvocationRepository;
|
||||
private final ObjectMapper objectMapper;
|
||||
|
||||
public ToolInvocationRecorder(ToolInvocationRepository toolInvocationRepository, ObjectMapper objectMapper) {
|
||||
this.toolInvocationRepository = toolInvocationRepository;
|
||||
this.objectMapper = objectMapper;
|
||||
}
|
||||
|
||||
public void save(ToolInvocation invocation) {
|
||||
try {
|
||||
if (invocation.getSessionId() == null || invocation.getSessionId().isBlank()) {
|
||||
invocation.setSessionId(SessionContextHolder.getSessionId());
|
||||
}
|
||||
if (invocation.getSessionId() == null || invocation.getSessionId().isBlank()) {
|
||||
log.debug("Skip tool_invocation without sessionId: tool={}", invocation.getToolName());
|
||||
return;
|
||||
}
|
||||
toolInvocationRepository.save(invocation);
|
||||
} catch (Exception e) {
|
||||
log.error("保存 tool_invocation 失败: tool={}", invocation.getToolName(), e);
|
||||
}
|
||||
}
|
||||
|
||||
public void recordEvidenceTool(String toolName,
|
||||
Map<String, Object> inputParams,
|
||||
String output,
|
||||
boolean success,
|
||||
long startTimeMillis,
|
||||
String errorMessage,
|
||||
String topicDomain) {
|
||||
String outputPreview = preview(output);
|
||||
Map<String, Object> details = new LinkedHashMap<>();
|
||||
details.put("trace_id", UUID.randomUUID().toString());
|
||||
if (topicDomain != null && !topicDomain.isBlank()) {
|
||||
details.put("retrieved_domains", List.of(topicDomain));
|
||||
}
|
||||
|
||||
ToolInvocation invocation = ToolInvocation.builder()
|
||||
.toolName(toolName)
|
||||
.inputParams(toJson(inputParams == null ? Map.of() : inputParams))
|
||||
.outputPreview(outputPreview)
|
||||
.outputLength(output == null ? 0 : output.length())
|
||||
.isTruncated(output != null && output.length() > OUTPUT_PREVIEW_LIMIT)
|
||||
.retrievalDetails(toJson(details))
|
||||
.durationMs((int) Math.max(0, System.currentTimeMillis() - startTimeMillis))
|
||||
.success(success)
|
||||
.errorMessage(errorMessage)
|
||||
.build();
|
||||
save(invocation);
|
||||
}
|
||||
|
||||
private String preview(String output) {
|
||||
if (output == null) {
|
||||
return null;
|
||||
}
|
||||
return output.length() <= OUTPUT_PREVIEW_LIMIT
|
||||
? output
|
||||
: output.substring(0, OUTPUT_PREVIEW_LIMIT) + "...";
|
||||
}
|
||||
|
||||
private String toJson(Map<String, Object> value) {
|
||||
try {
|
||||
return objectMapper.writeValueAsString(value);
|
||||
} catch (JsonProcessingException e) {
|
||||
log.debug("tool_invocation JSON 序列化失败", e);
|
||||
return "{}";
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -3,8 +3,8 @@ package com.superbiz.agent.tool;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import com.superbiz.agent.domain.entity.ToolInvocation;
|
||||
import com.superbiz.agent.dto.*;
|
||||
import com.superbiz.agent.repository.ToolInvocationRepository;
|
||||
import com.superbiz.agent.service.KnowledgeIndexService;
|
||||
import com.superbiz.agent.service.ToolInvocationRecorder;
|
||||
import com.superbiz.agent.service.VectorSearchService;
|
||||
import com.superbiz.agent.util.SessionContextHolder;
|
||||
import lombok.extern.slf4j.Slf4j;
|
||||
@@ -49,7 +49,7 @@ public class LookupKnowledgeTool {
|
||||
private VectorSearchService vectorSearchService;
|
||||
|
||||
@Autowired
|
||||
private ToolInvocationRepository toolInvocationRepository;
|
||||
private ToolInvocationRecorder toolInvocationRecorder;
|
||||
|
||||
@Autowired
|
||||
private RetrievedDocTracker retrievedDocTracker;
|
||||
@@ -417,7 +417,7 @@ public class LookupKnowledgeTool {
|
||||
.success(true)
|
||||
.build();
|
||||
|
||||
toolInvocationRepository.save(inv);
|
||||
toolInvocationRecorder.save(inv);
|
||||
log.debug("tool_invocation 已保存: sessionId={}, layer={}, relevanceLevel={}, duration={}ms",
|
||||
sessionId, layer, result != null ? result.getRelevanceLevel() : null, duration);
|
||||
} catch (Exception e) {
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
package com.superbiz.agent.service;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.junit.jupiter.api.io.TempDir;
|
||||
import org.springframework.mock.web.MockMultipartFile;
|
||||
import org.springframework.test.util.ReflectionTestUtils;
|
||||
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
class DocumentManagementServiceTest {
|
||||
|
||||
@TempDir
|
||||
Path tempDir;
|
||||
|
||||
@Test
|
||||
void saveToLocalStoresRelativePathUnderKnowledgeBase() throws Exception {
|
||||
DocumentManagementService service = new DocumentManagementService();
|
||||
ReflectionTestUtils.setField(service, "knowledgeBasePath", tempDir.toString());
|
||||
MockMultipartFile file = new MockMultipartFile(
|
||||
"file",
|
||||
"runbook.md",
|
||||
"text/markdown",
|
||||
"runbook content".getBytes()
|
||||
);
|
||||
|
||||
String storedPath = ReflectionTestUtils.invokeMethod(
|
||||
service,
|
||||
"saveToLocal",
|
||||
file,
|
||||
"runbook.md",
|
||||
"payment"
|
||||
);
|
||||
|
||||
assertEquals("payment/runbook.md", storedPath);
|
||||
assertTrue(Files.exists(tempDir.resolve("payment").resolve("runbook.md")));
|
||||
}
|
||||
}
|
||||
@@ -4,8 +4,6 @@ import com.superbiz.agent.dto.KnowledgeEntry;
|
||||
import org.junit.jupiter.api.BeforeEach;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.junit.jupiter.api.io.TempDir;
|
||||
import org.mockito.Mock;
|
||||
import org.mockito.MockitoAnnotations;
|
||||
import org.springframework.test.util.ReflectionTestUtils;
|
||||
|
||||
import java.nio.file.Files;
|
||||
@@ -21,17 +19,13 @@ class KnowledgeIndexServiceTest {
|
||||
|
||||
private KnowledgeIndexService service;
|
||||
|
||||
@Mock
|
||||
private FrontmatterParser frontmatterParser;
|
||||
|
||||
@TempDir
|
||||
Path tempDir;
|
||||
|
||||
@BeforeEach
|
||||
void setUp() {
|
||||
MockitoAnnotations.openMocks(this);
|
||||
service = new KnowledgeIndexService();
|
||||
ReflectionTestUtils.setField(service, "frontmatterParser", frontmatterParser);
|
||||
ReflectionTestUtils.setField(service, "knowledgeBasePath", tempDir.toString());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -144,6 +138,30 @@ class KnowledgeIndexServiceTest {
|
||||
assertTrue(result.contains("Test content"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void testReadDocument_relativePathUnderBasePath() throws Exception {
|
||||
Path categoryDir = tempDir.resolve("payment");
|
||||
Files.createDirectories(categoryDir);
|
||||
Path testFile = categoryDir.resolve("relative.md");
|
||||
Files.writeString(testFile, "Relative content");
|
||||
|
||||
String result = service.readDocument("payment/relative.md", 100);
|
||||
|
||||
assertEquals("Relative content", result);
|
||||
}
|
||||
|
||||
@Test
|
||||
void testReadDocument_legacyPathAlreadyContainsBasePath() throws Exception {
|
||||
Path categoryDir = tempDir.resolve("payment");
|
||||
Files.createDirectories(categoryDir);
|
||||
Path testFile = categoryDir.resolve("legacy.md");
|
||||
Files.writeString(testFile, "Legacy content");
|
||||
|
||||
String result = service.readDocument(tempDir.getFileName() + "/payment/legacy.md", 100);
|
||||
|
||||
assertEquals("Legacy content", result);
|
||||
}
|
||||
|
||||
@Test
|
||||
void testReadDocument_exceedsMaxChars() throws Exception {
|
||||
// 创建超长内容
|
||||
|
||||
Reference in New Issue
Block a user