diff --git a/src/main/java/com/superbiz/agent/agent/tool/QueryLogsTools.java b/src/main/java/com/superbiz/agent/agent/tool/QueryLogsTools.java index 0953c31..af4ac39 100644 --- a/src/main/java/com/superbiz/agent/agent/tool/QueryLogsTools.java +++ b/src/main/java/com/superbiz/agent/agent/tool/QueryLogsTools.java @@ -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 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 日志数据 diff --git a/src/main/java/com/superbiz/agent/agent/tool/QueryMetricsTools.java b/src/main/java/com/superbiz/agent/agent/tool/QueryMetricsTools.java index 83585ae..e9f044d 100644 --- a/src/main/java/com/superbiz/agent/agent/tool/QueryMetricsTools.java +++ b/src/main/java/com/superbiz/agent/agent/tool/QueryMetricsTools.java @@ -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 告警数据 diff --git a/src/main/java/com/superbiz/agent/controller/ChatController.java b/src/main/java/com/superbiz/agent/controller/ChatController.java index 24eeddf..4f058f2 100644 --- a/src/main/java/com/superbiz/agent/controller/ChatController.java +++ b/src/main/java/com/superbiz/agent/controller/ChatController.java @@ -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 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> history = session.getHistory(); + List> 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 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> history = session.getHistory(); + List> 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 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 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> 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 userMsg = new HashMap<>(); - userMsg.put("role", "user"); - userMsg.put("content", userQuestion); - messageHistory.add(userMsg); - - // 添加AI回复 - Map 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> 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)); } } diff --git a/src/main/java/com/superbiz/agent/domain/model/SessionContext.java b/src/main/java/com/superbiz/agent/domain/model/SessionContext.java index 6425491..e93fdb5 100644 --- a/src/main/java/com/superbiz/agent/domain/model/SessionContext.java +++ b/src/main/java/com/superbiz/agent/domain/model/SessionContext.java @@ -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 toolCalls = new ArrayList<>(); + /** + * 聊天消息历史:[{"role":"user","content":"..."}, {"role":"assistant","content":"..."}] + */ + @Builder.Default + private List> 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 userMessage = new HashMap<>(); + userMessage.put("role", "user"); + userMessage.put("content", userQuestion); + this.messageHistory.add(userMessage); + + Map 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> getMessageHistorySnapshot() { + if (this.messageHistory == null || this.messageHistory.isEmpty()) { + return new ArrayList<>(); + } + List> snapshot = new ArrayList<>(); + for (Map 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; + } + /** * 更新最后活跃时间 */ diff --git a/src/main/java/com/superbiz/agent/service/ChatService.java b/src/main/java/com/superbiz/agent/service/ChatService.java index c5ea785..e1731fe 100644 --- a/src/main/java/com/superbiz/agent/service/ChatService.java +++ b/src/main/java/com/superbiz/agent/service/ChatService.java @@ -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> history) throws GraphRunnerException { + return executeChatWithStrategy(chatModel, toolCallbacks, question, history, null); + } + + public ChatResult executeChatWithStrategy(ChatModel chatModel, ToolCallback[] toolCallbacks, + String question, List> 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> 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> 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; diff --git a/src/main/java/com/superbiz/agent/service/DocumentManagementService.java b/src/main/java/com/superbiz/agent/service/DocumentManagementService.java index 1c44f39..da5dbda 100644 --- a/src/main/java/com/superbiz/agent/service/DocumentManagementService.java +++ b/src/main/java/com/superbiz/agent/service/DocumentManagementService.java @@ -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 */ diff --git a/src/main/java/com/superbiz/agent/service/KnowledgeIndexService.java b/src/main/java/com/superbiz/agent/service/KnowledgeIndexService.java index 8c69555..088d0f5 100644 --- a/src/main/java/com/superbiz/agent/service/KnowledgeIndexService.java +++ b/src/main/java/com/superbiz/agent/service/KnowledgeIndexService.java @@ -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()); diff --git a/src/main/java/com/superbiz/agent/service/ToolInvocationRecorder.java b/src/main/java/com/superbiz/agent/service/ToolInvocationRecorder.java new file mode 100644 index 0000000..893b160 --- /dev/null +++ b/src/main/java/com/superbiz/agent/service/ToolInvocationRecorder.java @@ -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 inputParams, + String output, + boolean success, + long startTimeMillis, + String errorMessage, + String topicDomain) { + String outputPreview = preview(output); + Map 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 value) { + try { + return objectMapper.writeValueAsString(value); + } catch (JsonProcessingException e) { + log.debug("tool_invocation JSON 序列化失败", e); + return "{}"; + } + } +} diff --git a/src/main/java/com/superbiz/agent/tool/LookupKnowledgeTool.java b/src/main/java/com/superbiz/agent/tool/LookupKnowledgeTool.java index 1730b74..1bf36f7 100644 --- a/src/main/java/com/superbiz/agent/tool/LookupKnowledgeTool.java +++ b/src/main/java/com/superbiz/agent/tool/LookupKnowledgeTool.java @@ -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) { diff --git a/src/test/java/com/superbiz/agent/service/DocumentManagementServiceTest.java b/src/test/java/com/superbiz/agent/service/DocumentManagementServiceTest.java new file mode 100644 index 0000000..584332d --- /dev/null +++ b/src/test/java/com/superbiz/agent/service/DocumentManagementServiceTest.java @@ -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"))); + } +} diff --git a/src/test/java/com/superbiz/agent/service/KnowledgeIndexServiceTest.java b/src/test/java/com/superbiz/agent/service/KnowledgeIndexServiceTest.java index 72a69f3..1a85961 100644 --- a/src/test/java/com/superbiz/agent/service/KnowledgeIndexServiceTest.java +++ b/src/test/java/com/superbiz/agent/service/KnowledgeIndexServiceTest.java @@ -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 { // 创建超长内容