fix chat session traces and document paths

This commit is contained in:
zhuyongxin
2026-07-03 13:53:16 +08:00
parent fd89d84fc0
commit 1ff7f09d25
11 changed files with 492 additions and 248 deletions
@@ -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,15 +119,30 @@ 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 告警数据
* 与 aiops-docs 文档中的告警类型对应:
@@ -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 流式对话(支持自动工具调用)");
ToolCallback[] toolCallbacks = tools != null ? tools.getToolCallbacks() : new ToolCallback[0];
// 构建系统提示词(包含历史消息)
String systemPrompt = chatService.buildSystemPrompt(history);
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());
// 创建 ReactAgent
ReactAgent agent = chatService.createReactAgent(chatModel, systemPrompt);
session.addChatMessagePair(request.getQuestion(), fullAnswer, MAX_WINDOW_SIZE);
sessionManager.updateSession(session);
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
session.getSessionId(), session.getMessagePairCount());
// 用于累积完整答案
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);
}
}
);
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 {
// 创建超长内容