feat(session): 会话存储体系实现 & Chat多Agent路由

- 新增诊断会话(diagnosis_session/agent_step/tool_invocation)三表
- AgentLoggingHook 持久化 agent_step,记录决策链和耗时
- LookupKnowledgeTool 写入 tool_invocation,记录L0/L1检索质量
- TokenTrackingChatModel 捕获真实token用量
- Chat接口支持意图路由:简单问题单Agent,复杂问题多Agent(Planner+Executor)
- Prompt外置到 src/main/resources/prompts/
- 删除旧 diagnosis_record 表及相关文件
- 新增SessionContextHolder(ThreadLocal传递sessionId)
- QuestionComplexity 复杂度判断工具
- 测试覆盖三张新表的Repository
This commit is contained in:
zhuyongxin
2026-06-26 16:22:05 +08:00
parent a74ccea5be
commit 0d9cce75f9
33 changed files with 1941 additions and 470 deletions
@@ -83,16 +83,12 @@ public class ChatController {
// 记录可用工具
chatService.logAvailableTools();
ToolCallback[] toolCallbacks = tools != null ? tools.getToolCallbacks() : new ToolCallback[0];
// 根据问题复杂度自动选择单 Agent 或多 Agent
logger.info("开始 ReactAgent 对话(支持自动工具调用)");
// 构建系统提示词(包含历史消息)
String systemPrompt = chatService.buildSystemPrompt(history);
// 创建 ReactAgent
ReactAgent agent = chatService.createReactAgent(chatModel, systemPrompt);
// 执行对话
String fullAnswer = chatService.executeChat(agent, request.getQuestion());
String fullAnswer = chatService.executeChatWithStrategy(chatModel, toolCallbacks,
request.getQuestion(), history);
// 更新会话历史
session.addMessage(request.getQuestion(), fullAnswer);
@@ -0,0 +1,64 @@
package com.superbiz.agent.domain.entity;
import jakarta.persistence.*;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.time.LocalDateTime;
/**
* Agent 决策步骤实体
* 对应表: agent_step
*/
@Entity
@Table(name = "agent_step", indexes = {
@Index(name = "idx_session_step", columnList = "session_id, step_index"),
@Index(name = "idx_agent_name", columnList = "agent_name")
})
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class AgentStep {
@Id
@GeneratedValue(strategy = GenerationType.IDENTITY)
private Long id;
@Column(name = "session_id", nullable = false, length = 64)
private String sessionId;
@Column(name = "step_index", nullable = false)
private Integer stepIndex;
@Column(name = "agent_name", nullable = false, length = 32)
private String agentName;
@Column(name = "model_input", columnDefinition = "TEXT")
private String modelInput;
@Column(name = "model_output", columnDefinition = "TEXT")
private String modelOutput;
@Column(name = "thought", columnDefinition = "TEXT")
private String thought;
@Column(name = "has_tool_call")
private Boolean hasToolCall;
@Column(name = "duration_ms")
private Integer durationMs;
@Column(name = "token_count")
private Integer tokenCount;
@Column(name = "created_at", nullable = false, updatable = false)
private LocalDateTime createdAt;
@PrePersist
protected void onCreate() {
createdAt = LocalDateTime.now();
}
}
@@ -1,128 +0,0 @@
package com.superbiz.agent.domain.entity;
import jakarta.persistence.*;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import com.superbiz.agent.domain.enums.DiagnosisStatus;
import com.superbiz.agent.domain.enums.FaultCategory;
import org.hibernate.annotations.JdbcTypeCode;
import org.hibernate.type.SqlTypes;
import java.time.LocalDateTime;
import java.util.List;
import java.util.Map;
/**
* 诊断记录实体
* 对应表: diagnosis_record
*/
@Entity
@Table(name = "diagnosis_record", indexes = {
@Index(name = "idx_business_id", columnList = "business_id"),
@Index(name = "idx_trace_id", columnList = "trace_id"),
@Index(name = "idx_session_id", columnList = "session_id"),
@Index(name = "idx_fault_category", columnList = "fault_category"),
@Index(name = "idx_error_code", columnList = "error_code"),
@Index(name = "idx_created_at", columnList = "created_at"),
@Index(name = "idx_status", columnList = "status")
})
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class DiagnosisRecord {
@Id
@GeneratedValue(strategy = GenerationType.IDENTITY)
private Long id;
@Column(name = "diagnosis_id", unique = true, nullable = false, length = 64)
private String diagnosisId;
// 关联信息
@Column(name = "session_id", length = 64)
private String sessionId;
@Column(name = "business_id", length = 128)
private String businessId;
@Column(name = "trace_id", length = 64)
private String traceId;
// 故障分类
@Enumerated(EnumType.STRING)
@Column(name = "fault_category", length = 32, columnDefinition = "VARCHAR(32)")
private FaultCategory faultCategory;
@Column(name = "fault_source", length = 128)
private String faultSource;
@Column(name = "fault_target", length = 256)
private String faultTarget;
// 错误信息
@Column(name = "error_code", length = 64)
private String errorCode;
@Column(name = "error_message", columnDefinition = "TEXT")
private String errorMessage;
@Column(name = "stack_trace", columnDefinition = "TEXT")
private String stackTrace;
// 诊断结果
@Column(name = "problem_type", length = 32)
private String problemType;
@Column(name = "root_cause", columnDefinition = "TEXT")
private String rootCause;
@Column(name = "solution", columnDefinition = "TEXT")
private String solution;
@Column(name = "report_markdown", columnDefinition = "TEXT")
private String reportMarkdown;
// 评估指标
@Enumerated(EnumType.STRING)
@Column(name = "status", length = 16, columnDefinition = "VARCHAR(16)")
private DiagnosisStatus status = DiagnosisStatus.PENDING;
@Column(name = "confidence")
private Integer confidence;
@Column(name = "duration")
private Integer duration;
// 用户反馈
@Column(name = "feedback", length = 16)
private String feedback;
// 调试字段 - JSON 类型
@JdbcTypeCode(SqlTypes.JSON)
@Column(name = "tool_calls", columnDefinition = "JSON")
private List<Map<String, Object>> toolCalls;
// 元数据
@Column(name = "created_by", length = 64)
private String createdBy;
@Column(name = "created_at", nullable = false, updatable = false)
private LocalDateTime createdAt;
@Column(name = "updated_at")
private LocalDateTime updatedAt;
@PrePersist
protected void onCreate() {
createdAt = LocalDateTime.now();
updatedAt = LocalDateTime.now();
}
@PreUpdate
protected void onUpdate() {
updatedAt = LocalDateTime.now();
}
}
@@ -0,0 +1,80 @@
package com.superbiz.agent.domain.entity;
import jakarta.persistence.*;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import org.hibernate.annotations.JdbcTypeCode;
import org.hibernate.type.SqlTypes;
import java.time.LocalDateTime;
/**
* 诊断会话实体
* 对应表: diagnosis_session
*/
@Entity
@Table(name = "diagnosis_session", indexes = {
@Index(name = "idx_created_at", columnList = "created_at"),
@Index(name = "idx_status", columnList = "status"),
@Index(name = "idx_agent_flow", columnList = "agent_flow")
})
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class DiagnosisSession {
@Id
@GeneratedValue(strategy = GenerationType.IDENTITY)
private Long id;
@Column(name = "session_id", unique = true, nullable = false, length = 64)
private String sessionId;
@Column(name = "query", nullable = false, columnDefinition = "TEXT")
private String query;
@Column(name = "status", length = 16)
private String status = "PENDING";
@Column(name = "agent_flow", length = 32)
private String agentFlow;
@Column(name = "total_duration_ms")
private Integer totalDurationMs;
@Column(name = "total_token_count")
private Integer totalTokenCount;
@Column(name = "step_count")
private Integer stepCount;
@Column(name = "tool_call_count")
private Integer toolCallCount;
@JdbcTypeCode(SqlTypes.JSON)
@Column(name = "self_evaluation", columnDefinition = "JSON")
private String selfEvaluation;
@Column(name = "feedback", length = 16)
private String feedback;
@Column(name = "created_at", nullable = false, updatable = false)
private LocalDateTime createdAt;
@Column(name = "updated_at")
private LocalDateTime updatedAt;
@PrePersist
protected void onCreate() {
createdAt = LocalDateTime.now();
updatedAt = LocalDateTime.now();
}
@PreUpdate
protected void onUpdate() {
updatedAt = LocalDateTime.now();
}
}
@@ -0,0 +1,84 @@
package com.superbiz.agent.domain.entity;
import jakarta.persistence.*;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import org.hibernate.annotations.JdbcTypeCode;
import org.hibernate.type.SqlTypes;
import java.time.LocalDateTime;
/**
* 工具调用明细实体
* 对应表: tool_invocation
*/
@Entity
@Table(name = "tool_invocation", indexes = {
@Index(name = "idx_session_id", columnList = "session_id"),
@Index(name = "idx_tool_name", columnList = "tool_name"),
@Index(name = "idx_retrieval_layer", columnList = "retrieval_layer")
})
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class ToolInvocation {
@Id
@GeneratedValue(strategy = GenerationType.IDENTITY)
private Long id;
@Column(name = "session_id", nullable = false, length = 64)
private String sessionId;
@Column(name = "step_id")
private Long stepId;
@Column(name = "tool_name", nullable = false, length = 64)
private String toolName;
@JdbcTypeCode(SqlTypes.JSON)
@Column(name = "input_params", nullable = false, columnDefinition = "JSON")
private String inputParams;
@Column(name = "output_preview", columnDefinition = "TEXT")
private String outputPreview;
@Column(name = "output_length")
private Integer outputLength;
@Column(name = "retrieval_layer", length = 8)
private String retrievalLayer;
@Column(name = "l0_match_count")
private Integer l0MatchCount;
@Column(name = "l1_match_count")
private Integer l1MatchCount;
@Column(name = "is_truncated")
private Boolean isTruncated;
@JdbcTypeCode(SqlTypes.JSON)
@Column(name = "retrieval_details", columnDefinition = "JSON")
private String retrievalDetails;
@Column(name = "duration_ms")
private Integer durationMs;
@Column(name = "success")
private Boolean success;
@Column(name = "error_message", columnDefinition = "TEXT")
private String errorMessage;
@Column(name = "created_at", nullable = false, updatable = false)
private LocalDateTime createdAt;
@PrePersist
protected void onCreate() {
createdAt = LocalDateTime.now();
}
}
@@ -1,21 +0,0 @@
package com.superbiz.agent.domain.enums;
/**
* 诊断状态枚举
*/
public enum DiagnosisStatus {
PENDING("待处理"),
RUNNING("诊断中"),
SUCCESS("成功"),
FAILED("失败");
private final String description;
DiagnosisStatus(String description) {
this.description = description;
}
public String getDescription() {
return description;
}
}
@@ -5,23 +5,39 @@ import com.alibaba.cloud.ai.graph.agent.hook.messages.AgentCommand;
import com.alibaba.cloud.ai.graph.agent.hook.HookPosition;
import com.alibaba.cloud.ai.graph.agent.hook.HookPositions;
import com.alibaba.cloud.ai.graph.RunnableConfig;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.util.SessionContextHolder;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.UserMessage;
import org.springframework.ai.chat.messages.ToolResponseMessage;
import java.util.List;
import java.util.Map;
import java.util.concurrent.ConcurrentHashMap;
/**
* Agent 日志 Hook
* 用于记录 Agent 的思考过程、消息流转
* 记录 Agent 的思考过程、消息流转 + 持久化 agent_step 到 DB
*/
@Slf4j
@HookPositions({HookPosition.BEFORE_MODEL, HookPosition.AFTER_MODEL})
public class AgentLoggingHook extends MessagesModelHook {
private int modelCallCount = 0;
private final AgentStepRepository agentStepRepository;
private final String agentName;
/** 每个 session 的步数计数器:sessionId → stepIndex */
private final ConcurrentHashMap<String, Integer> stepCounters = new ConcurrentHashMap<>();
/** beforeModel → afterModel 中间状态:sessionId_stepIndex → {stepId, startTime} */
private final ConcurrentHashMap<String, Map<String, Object>> pendingSteps = new ConcurrentHashMap<>();
public AgentLoggingHook(AgentStepRepository agentStepRepository, String agentName) {
this.agentStepRepository = agentStepRepository;
this.agentName = agentName;
}
@Override
public String getName() {
@@ -30,9 +46,12 @@ public class AgentLoggingHook extends MessagesModelHook {
@Override
public AgentCommand beforeModel(List<Message> previousMessages, RunnableConfig config) {
modelCallCount++;
String sessionId = SessionContextHolder.getSessionId();
int stepIndex = stepCounters.merge(sessionId, 0, (old, one) -> old + 1);
log.info("========================================");
log.info("*** [Agent 思考] 第 {} 轮思考开始", modelCallCount);
log.info("*** [Agent 思考] 第 {} 轮思考开始", stepIndex + 1);
log.info("*** [Agent 思考] 当前消息数量: {}", previousMessages.size());
// 打印最后几条消息
@@ -40,28 +59,51 @@ public class AgentLoggingHook extends MessagesModelHook {
if (lastN > 0) {
log.info("*** [Agent 思考] 最近 {} 条消息:", lastN);
List<Message> recentMessages = previousMessages.subList(previousMessages.size() - lastN, previousMessages.size());
for (int i = 0; i < recentMessages.size(); i++) {
Message msg = recentMessages.get(i);
String role = getMessageRole(msg);
log.info(" [{}] 角色: {}, 类型: {}", i + 1, role, msg.getClass().getSimpleName());
// Message 接口可能没有直接的 getContent() 方法,跳过内容打印
// 具体内容会在工具调用日志中体现
}
}
log.info("*** [Agent 思考] 准备调用模型...");
log.info("========================================");
// 不修改消息,直接返回
// 持久化 agent_step(beforeModel:先创建,先记 model_input 摘要)
if (sessionId != null) {
try {
String modelInputSummary = buildModelInputSummary(previousMessages);
AgentStep step = AgentStep.builder()
.sessionId(sessionId)
.stepIndex(stepIndex)
.agentName(agentName)
.modelInput(modelInputSummary)
.build();
AgentStep saved = agentStepRepository.save(step);
// 记录中间状态供 afterModel 使用
pendingSteps.put(sessionId + "_" + stepIndex, Map.of(
"stepId", saved.getId(),
"startTime", System.currentTimeMillis()
));
log.debug("agent_step 已创建: sessionId={}, stepIndex={}, id={}", sessionId, stepIndex, saved.getId());
} catch (Exception e) {
log.error("保存 agent_step 失败", e);
// 不中断 Agent 执行
}
}
return new AgentCommand(previousMessages);
}
@Override
public AgentCommand afterModel(List<Message> previousMessages, RunnableConfig config) {
String sessionId = SessionContextHolder.getSessionId();
log.info("========================================");
log.info("*** [Agent 思考] 第 {} 轮思考完成", modelCallCount);
log.info("*** [Agent 思考] 第 {} 轮思考完成", stepCounters.getOrDefault(sessionId, 0));
// 查找最后一条 AssistantMessage(模型的回复)
AssistantMessage lastAssistant = null;
@@ -72,7 +114,15 @@ public class AgentLoggingHook extends MessagesModelHook {
}
}
boolean hasToolCall = false;
if (lastAssistant != null) {
// 调试:打印 metadata
if (lastAssistant.getMetadata() != null && !lastAssistant.getMetadata().isEmpty()) {
log.info("*** [Agent 思考] 模型返回 metadata: {}", lastAssistant.getMetadata());
} else {
log.info("*** [Agent 思考] 模型返回 metadata: (空)");
}
// 打印模型返回的文本内容
String textContent = extractTextContent(lastAssistant);
if (textContent != null && !textContent.isEmpty()) {
@@ -84,6 +134,7 @@ public class AgentLoggingHook extends MessagesModelHook {
// 检查是否有工具调用
if (lastAssistant.getToolCalls() != null && !lastAssistant.getToolCalls().isEmpty()) {
hasToolCall = true;
log.info("*** [Agent 思考] 模型决定调用 {} 个工具:",
lastAssistant.getToolCalls().size());
lastAssistant.getToolCalls().forEach(toolCall -> {
@@ -100,84 +151,178 @@ public class AgentLoggingHook extends MessagesModelHook {
log.info("========================================");
// 不修改消息,直接返回
// 更新 agent_step(afterModel:补全 model_output、耗时等)
if (sessionId != null) {
int stepIndex = stepCounters.getOrDefault(sessionId, 0);
String stepKey = sessionId + "_" + stepIndex;
Map<String, Object> pending = pendingSteps.remove(stepKey);
if (pending != null) {
try {
Long stepId = (Long) pending.get("stepId");
long startTime = (long) pending.get("startTime");
int durationMs = (int) (System.currentTimeMillis() - startTime);
AgentStep step = agentStepRepository.findById(stepId).orElse(null);
if (step != null) {
String thought = extractTextContent(lastAssistant);
if (thought != null && thought.length() > 2000) {
thought = thought.substring(0, 2000);
}
step.setThought(thought);
step.setHasToolCall(hasToolCall);
step.setDurationMs(durationMs);
if (lastAssistant != null) {
String outputSummary = buildModelOutputSummary(lastAssistant);
step.setModelOutput(outputSummary);
// 读取实际 token 用量(由 TokenTrackingChatModel 写入)
Integer tokenCount = TokenUsageHolder.get();
if (tokenCount != null) {
step.setTokenCount(tokenCount);
}
}
agentStepRepository.save(step);
log.debug("agent_step 已更新: sessionId={}, stepIndex={}, duration={}ms",
sessionId, stepIndex, durationMs);
}
} catch (Exception e) {
log.error("更新 agent_step 失败", e);
}
}
}
// 清理 token 上下文
TokenUsageHolder.clear();
return new AgentCommand(previousMessages);
}
/**
* 构建模型输入摘要(前 N 条消息的 role + 截断内容)
*/
private String buildModelInputSummary(List<Message> messages) {
StringBuilder sb = new StringBuilder();
int maxMessages = Math.min(messages.size(), 5);
for (int i = messages.size() - maxMessages; i < messages.size(); i++) {
Message msg = messages.get(i);
String role = getMessageRole(msg);
String content = msg.toString();
if (content.length() > 200) {
content = content.substring(0, 200) + "...";
}
sb.append("[").append(role).append("] ").append(content).append("\n");
}
String result = sb.toString();
if (result.length() > 500) {
result = result.substring(0, 500) + "...";
}
return result;
}
/**
* 构建模型输出摘要
*/
private String buildModelOutputSummary(AssistantMessage message) {
String text = extractTextContent(message);
if (text == null) {
text = "";
}
if (text.length() > 500) {
text = text.substring(0, 500) + "...";
}
StringBuilder sb = new StringBuilder();
sb.append("{\"text\":\"").append(escapeJson(text)).append("\"");
if (message.getToolCalls() != null && !message.getToolCalls().isEmpty()) {
sb.append(",\"toolCalls\":[");
for (int i = 0; i < message.getToolCalls().size(); i++) {
if (i > 0) sb.append(",");
sb.append("{\"name\":\"").append(escapeJson(message.getToolCalls().get(i).name()))
.append("\",\"arguments\":").append(message.getToolCalls().get(i).arguments()).append("}");
}
sb.append("]");
}
sb.append("}");
return sb.toString();
}
private String escapeJson(String s) {
if (s == null) return "";
return s.replace("\\", "\\\\")
.replace("\"", "\\\"")
.replace("\n", "\\n")
.replace("\r", "\\r")
.replace("\t", "\\t");
}
/**
* 提取 AssistantMessage 的文本内容
*/
private String extractTextContent(AssistantMessage message) {
if (message == null) return null;
try {
// 方法 1: 尝试通过反射获取 text 字段
// 方法 1: 反射获取 text 字段
try {
java.lang.reflect.Field textField = message.getClass().getDeclaredField("text");
textField.setAccessible(true);
Object value = textField.get(message);
if (value != null) {
String text = value.toString();
log.debug("通过 text 字段提取成功");
return text;
return value.toString();
}
} catch (NoSuchFieldException e) {
// text 字段不存在,尝试下一种方法
// 尝试下一种方法
}
// 方法 2: 尝试 content 字段
// 方法 2: 反射获取 content 字段
try {
java.lang.reflect.Field contentField = message.getClass().getDeclaredField("content");
contentField.setAccessible(true);
Object value = contentField.get(message);
if (value != null) {
String text = value.toString();
log.debug("通过 content 字段提取成功");
return text;
return value.toString();
}
} catch (NoSuchFieldException e) {
// content 字段不存在,尝试下一种方法
// 尝试下一种方法
}
// 方法 3: 尝试调用 getText() 方法
// 方法 3: 调用 getText() 方法
try {
java.lang.reflect.Method getTextMethod = message.getClass().getMethod("getText");
Object value = getTextMethod.invoke(message);
if (value != null) {
String text = value.toString();
log.debug("通过 getText() 方法提取成功");
return text;
return value.toString();
}
} catch (NoSuchMethodException e) {
// getText() 方法不存在,尝试下一种方法
// 尝试下一种方法
}
// 方法 4: 尝试调用 getContent() 方法
// 方法 4: 调用 getContent() 方法
try {
java.lang.reflect.Method getContentMethod = message.getClass().getMethod("getContent");
Object value = getContentMethod.invoke(message);
if (value != null) {
String text = value.toString();
log.debug("通过 getContent() 方法提取成功");
return text;
return value.toString();
}
} catch (NoSuchMethodException e) {
// getContent() 方法不存在
// 方法不存在
}
// 方法 5: 打印所有字段和方法,帮助调试
// 方法 5: 打印类结构信息
log.warn("无法提取 AssistantMessage 文本内容,打印类信息:");
log.warn("类名: {}", message.getClass().getName());
log.warn("字段列表:");
for (java.lang.reflect.Field field : message.getClass().getDeclaredFields()) {
log.warn(" - {}: {}", field.getName(), field.getType().getSimpleName());
}
log.warn("方法列表:");
for (java.lang.reflect.Method method : message.getClass().getMethods()) {
if (method.getName().startsWith("get") && method.getParameterCount() == 0) {
log.warn(" - {}(): {}", method.getName(), method.getReturnType().getSimpleName());
}
}
// 方法 6: 最后尝试 toString()
// 方法 6: toString() 兜底
String toString = message.toString();
if (toString != null && !toString.startsWith("AssistantMessage@")) {
log.debug("通过 toString() 提取");
@@ -0,0 +1,45 @@
package com.superbiz.agent.hook;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.prompt.Prompt;
import reactor.core.publisher.Flux;
/**
* ChatModel 包装器 — 捕获每次模型调用的实际 token 用量
* 通过 TokenUsageHolder 传递给 AgentLoggingHook
*/
public class TokenTrackingChatModel implements ChatModel {
private final ChatModel delegate;
public TokenTrackingChatModel(ChatModel delegate) {
this.delegate = delegate;
}
@Override
public ChatResponse call(Prompt prompt) {
ChatResponse response = delegate.call(prompt);
captureTokenUsage(response);
return response;
}
@Override
public Flux<ChatResponse> stream(Prompt prompt) {
return delegate.stream(prompt);
}
private void captureTokenUsage(ChatResponse response) {
try {
if (response.getMetadata() != null && response.getMetadata().getUsage() != null) {
var usage = response.getMetadata().getUsage();
Integer total = usage.getTotalTokens();
if (total != null && total > 0) {
TokenUsageHolder.set(total);
}
}
} catch (Exception e) {
// 不中断模型调用
}
}
}
@@ -0,0 +1,22 @@
package com.superbiz.agent.hook;
/**
* Token 用量持有者(基于 ThreadLocal)
* ChatModel 调用后写入实际 token 数,AgentLoggingHook 读取
*/
public class TokenUsageHolder {
private static final ThreadLocal<Integer> TOKEN_COUNT = new ThreadLocal<>();
public static void set(Integer count) {
TOKEN_COUNT.set(count);
}
public static Integer get() {
return TOKEN_COUNT.get();
}
public static void clear() {
TOKEN_COUNT.remove();
}
}
@@ -0,0 +1,24 @@
package com.superbiz.agent.repository;
import com.superbiz.agent.domain.entity.AgentStep;
import org.springframework.data.jpa.repository.JpaRepository;
import org.springframework.stereotype.Repository;
import java.util.List;
/**
* Agent 决策步骤 Repository
*/
@Repository
public interface AgentStepRepository extends JpaRepository<AgentStep, Long> {
/**
* 根据会话ID查询所有步骤(按步骤号排序)
*/
List<AgentStep> findBySessionIdOrderByStepIndex(String sessionId);
/**
* 统计某个会话的步骤数
*/
int countBySessionId(String sessionId);
}
@@ -1,73 +0,0 @@
package com.superbiz.agent.repository;
import com.superbiz.agent.domain.enums.DiagnosisStatus;
import com.superbiz.agent.domain.enums.FaultCategory;
import com.superbiz.agent.domain.entity.DiagnosisRecord;
import org.springframework.data.domain.Page;
import org.springframework.data.domain.Pageable;
import org.springframework.data.jpa.repository.JpaRepository;
import org.springframework.stereotype.Repository;
import java.time.LocalDateTime;
import java.util.List;
import java.util.Optional;
/**
* 诊断记录 Repository
*/
@Repository
public interface DiagnosisRecordRepository extends JpaRepository<DiagnosisRecord, Long> {
/**
* 根据诊断ID查询
*/
Optional<DiagnosisRecord> findByDiagnosisId(String diagnosisId);
/**
* 根据业务ID查询
*/
Optional<DiagnosisRecord> findByBusinessId(String businessId);
/**
* 根据链路追踪ID查询
*/
Optional<DiagnosisRecord> findByTraceId(String traceId);
/**
* 根据会话ID查询所有记录
*/
List<DiagnosisRecord> findBySessionId(String sessionId);
/**
* 根据故障类别和错误码查询
*/
List<DiagnosisRecord> findByFaultCategoryAndErrorCode(FaultCategory category, String errorCode);
/**
* 根据故障类别、故障源和错误码查询
*/
List<DiagnosisRecord> findByFaultCategoryAndFaultSourceAndErrorCode(
FaultCategory category, String faultSource, String errorCode);
/**
* 根据状态查询
*/
List<DiagnosisRecord> findByStatus(DiagnosisStatus status);
/**
* 根据时间范围查询(分页)
*/
Page<DiagnosisRecord> findByCreatedAtBetween(
LocalDateTime start, LocalDateTime end, Pageable pageable);
/**
* 根据故障类别和时间范围查询(分页)
*/
Page<DiagnosisRecord> findByFaultCategoryAndCreatedAtBetween(
FaultCategory category, LocalDateTime start, LocalDateTime end, Pageable pageable);
/**
* 查询有用反馈的高置信度记录(用于生成案例)
*/
List<DiagnosisRecord> findByFeedbackAndConfidenceGreaterThanEqual(String feedback, Integer confidence);
}
@@ -0,0 +1,12 @@
package com.superbiz.agent.repository;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import org.springframework.data.jpa.repository.JpaRepository;
import org.springframework.stereotype.Repository;
import java.util.Optional;
@Repository
public interface DiagnosisSessionRepository extends JpaRepository<DiagnosisSession, Long> {
Optional<DiagnosisSession> findBySessionId(String sessionId);
}
@@ -0,0 +1,29 @@
package com.superbiz.agent.repository;
import com.superbiz.agent.domain.entity.ToolInvocation;
import org.springframework.data.jpa.repository.JpaRepository;
import org.springframework.stereotype.Repository;
import java.util.List;
/**
* 工具调用明细 Repository
*/
@Repository
public interface ToolInvocationRepository extends JpaRepository<ToolInvocation, Long> {
/**
* 根据会话ID查询所有工具调用
*/
List<ToolInvocation> findBySessionId(String sessionId);
/**
* 根据工具名查询所有调用
*/
List<ToolInvocation> findByToolName(String toolName);
/**
* 根据会话ID和工具名查询
*/
List<ToolInvocation> findBySessionIdAndToolName(String sessionId, String toolName);
}
@@ -9,6 +9,13 @@ import com.superbiz.agent.agent.tool.DateTimeTools;
import com.superbiz.agent.agent.tool.InternalDocsTools;
import com.superbiz.agent.agent.tool.QueryLogsTools;
import com.superbiz.agent.agent.tool.QueryMetricsTools;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.hook.AgentLoggingHook;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.util.SessionContextHolder;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.messages.AssistantMessage;
@@ -20,6 +27,7 @@ import com.superbiz.agent.tool.LookupKnowledgeTool;
import java.util.List;
import java.util.Optional;
import java.util.UUID;
/**
* AI Ops 智能运维服务
@@ -48,6 +56,12 @@ public class AiOpsService {
@Autowired
private AiOpsPromptProperties promptProperties;
@Autowired
private DiagnosisSessionRepository diagnosisSessionRepository;
@Autowired
private AgentStepRepository agentStepRepository;
/**
* 执行 AI Ops 告警分析流程
*
@@ -59,34 +73,65 @@ public class AiOpsService {
public Optional<OverAllState> executeAiOpsAnalysis(ChatModel chatModel, ToolCallback[] toolCallbacks) throws GraphRunnerException {
logger.info("开始执行 AI Ops 多 Agent 协作流程");
// 构建 Planner 和 Executor Agent
ReactAgent plannerAgent = buildPlannerAgent(chatModel, toolCallbacks);
ReactAgent executorAgent = buildExecutorAgent(chatModel, toolCallbacks);
String sessionId = UUID.randomUUID().toString().substring(0, 8);
long startTime = System.currentTimeMillis();
// 构建 Supervisor Agent
SupervisorAgent supervisorAgent = SupervisorAgent.builder()
.name("ai_ops_supervisor")
.description("负责调度 Planner 与 Executor 的多 Agent 控制器")
.model(chatModel)
.systemPrompt(promptProperties.getSupervisor())
.subAgents(List.of(plannerAgent, executorAgent))
// 创建诊断会话
DiagnosisSession session = DiagnosisSession.builder()
.sessionId(sessionId)
.query("AI Ops 告警分析")
.status("RUNNING")
.agentFlow("AI_OPS")
.build();
diagnosisSessionRepository.save(session);
String taskPrompt = "你是企业级 SRE,接到了自动化告警排查任务。请结合工具调用,执行**规划→执行→再规划**的闭环,并最终按照固定模板输出《告警分析报告》。禁止编造虚假数据,如连续多次查询失败需诚实反馈无法完成的原因。";
// 设置 ThreadLocal 上下文(LookupKnowledgeTool 通过此获取 sessionId)
SessionContextHolder.setSessionId(sessionId);
logger.info("调用 Supervisor Agent 开始编排...");
try {
// 构建 Planner 和 Executor Agent(每个 Agent 各自带 Hook)
ReactAgent plannerAgent = buildPlannerAgent(chatModel, toolCallbacks);
ReactAgent executorAgent = buildExecutorAgent(chatModel, toolCallbacks);
Optional<OverAllState> stateOptional = supervisorAgent.invoke(taskPrompt);
// 构建 Supervisor Agent(不加 Hook)
SupervisorAgent supervisorAgent = SupervisorAgent.builder()
.name("ai_ops_supervisor")
.description("负责调度 Planner 与 Executor 的多 Agent 控制器")
.model(chatModel)
.systemPrompt(promptProperties.getSupervisor())
.subAgents(List.of(plannerAgent, executorAgent))
.build();
// 添加调试代码
if (stateOptional.isPresent()) {
OverAllState state = stateOptional.get();
logger.debug("Final State Keys: {}", state.data().keySet()); // 打印所有 key
logger.debug("Planner Plan: {}", state.value("planner_plan"));
logger.debug("Executor Feedback: {}", state.value("executor_feedback"));
String taskPrompt = "你是企业级 SRE,接到了自动化告警排查任务。请结合工具调用,执行**规划→执行→再规划**的闭环,并最终按照固定模板输出《告警分析报告》。禁止编造虚假数据,如连续多次查询失败需诚实反馈无法完成的原因。";
logger.info("调用 Supervisor Agent 开始编排...");
Optional<OverAllState> stateOptional = supervisorAgent.invoke(taskPrompt);
long duration = System.currentTimeMillis() - startTime;
// 更新诊断会话
session.setStatus(stateOptional.isPresent() ? "SUCCESS" : "FAILED");
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
// 添加调试代码
if (stateOptional.isPresent()) {
OverAllState state = stateOptional.get();
logger.debug("Final State Keys: {}", state.data().keySet());
logger.debug("Planner Plan: {}", state.value("planner_plan"));
logger.debug("Executor Feedback: {}", state.value("executor_feedback"));
}
return stateOptional;
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
throw e;
} finally {
SessionContextHolder.clear();
}
return stateOptional;
}
/**
@@ -124,6 +169,7 @@ public class AiOpsService {
.systemPrompt(promptProperties.getPlanner())
.methodTools(buildMethodToolsArray())
.tools(toolCallbacks)
.hooks(new AgentLoggingHook(agentStepRepository, "planner"))
.outputKey("planner_plan")
.build();
}
@@ -139,6 +185,7 @@ public class AiOpsService {
.systemPrompt(promptProperties.getExecutor())
.methodTools(buildMethodToolsArray())
.tools(toolCallbacks)
.hooks(new AgentLoggingHook(agentStepRepository, "executor"))
.outputKey("executor_feedback")
.build();
}
@@ -157,4 +204,26 @@ public class AiOpsService {
return new Object[]{dateTimeTools, lookupKnowledgeTool, queryMetricsTools};
}
}
/** 从 agent_step 汇总指标回填 diagnosis_session */
private void backfillSessionMetrics(DiagnosisSession session) {
try {
List<AgentStep> steps = agentStepRepository.findBySessionIdOrderByStepIndex(session.getSessionId());
if (steps.isEmpty()) return;
int totalTokens = 0;
int stepCount = 0;
int toolCallCount = 0;
for (AgentStep s : steps) {
stepCount++;
if (s.getTokenCount() != null) totalTokens += s.getTokenCount();
if (Boolean.TRUE.equals(s.getHasToolCall())) toolCallCount++;
}
session.setTotalTokenCount(totalTokens);
session.setStepCount(stepCount);
session.setToolCallCount(toolCallCount);
} catch (Exception e) {
logger.warn("回填会话指标失败: sessionId={}", session.getSessionId(), e);
}
}
}
@@ -1,24 +1,40 @@
package com.superbiz.agent.service;
import com.alibaba.cloud.ai.graph.OverAllState;
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
import com.alibaba.cloud.ai.graph.agent.flow.agent.SupervisorAgent;
import com.alibaba.cloud.ai.graph.exception.GraphRunnerException;
import com.superbiz.agent.agent.tool.DateTimeTools;
import com.superbiz.agent.agent.tool.InternalDocsTools;
import com.superbiz.agent.agent.tool.QueryLogsTools;
import com.superbiz.agent.agent.tool.QueryMetricsTools;
import com.superbiz.agent.tool.LookupKnowledgeTool;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.hook.AgentLoggingHook;
import com.superbiz.agent.hook.TokenTrackingChatModel;
import com.superbiz.agent.hook.TokenUsageHolder;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.tool.LookupKnowledgeTool;
import com.superbiz.agent.util.QuestionComplexity;
import com.superbiz.agent.util.SessionContextHolder;
import jakarta.annotation.PostConstruct;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
import org.springframework.ai.chat.messages.AssistantMessage;
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.core.io.ClassPathResource;
import org.springframework.stereotype.Service;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.util.List;
import java.util.Map;
import java.util.Optional;
import java.util.UUID;
/**
* 聊天服务
@@ -50,6 +66,37 @@ public class ChatService {
@Autowired
private LookupKnowledgeTool lookupKnowledgeTool;
@Autowired
private DiagnosisSessionRepository diagnosisSessionRepository;
@Autowired
private AgentStepRepository agentStepRepository;
/** 多 Agent Chat 的 Prompt */
private String chatPlannerPrompt;
private String chatExecutorPrompt;
@PostConstruct
public void init() {
// 加载 Prompt
try {
chatPlannerPrompt = new String(
new ClassPathResource("prompts/chat-planner-prompt.md").getInputStream().readAllBytes(),
StandardCharsets.UTF_8);
chatExecutorPrompt = new String(
new ClassPathResource("prompts/chat-executor-prompt.md").getInputStream().readAllBytes(),
StandardCharsets.UTF_8);
logger.info("Chat 多 Agent Prompts 加载成功");
} catch (IOException e) {
logger.error("加载 Chat Prompt 文件失败", e);
throw new RuntimeException("Failed to load chat prompts", e);
}
// 包装 ChatModel 以捕获 token 用量
chatModel = new TokenTrackingChatModel(chatModel);
logger.info("ChatModel 已包装 TokenTrackingChatModel");
}
/**
* 获取注入的 ChatModel
*/
@@ -177,7 +224,7 @@ public class ChatService {
.systemPrompt(systemPrompt)
.methodTools(buildMethodToolsArray())
.tools(getToolCallbacks())
.hooks(new AgentLoggingHook()) // 添加日志 Hook
.hooks(new AgentLoggingHook(agentStepRepository, "intelligent_assistant"))
.build();
}
@@ -191,16 +238,201 @@ public class ChatService {
logger.info("========================================");
logger.info("📝 用户问题: {}", question);
String sessionId = UUID.randomUUID().toString().substring(0, 8);
long startTime = System.currentTimeMillis();
var response = agent.call(question);
long duration = System.currentTimeMillis() - startTime;
String answer = response.getText();
// 创建诊断会话
DiagnosisSession session = DiagnosisSession.builder()
.sessionId(sessionId)
.query(question)
.status("RUNNING")
.agentFlow("CHAT")
.build();
diagnosisSessionRepository.save(session);
logger.info("⏱️ 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
logger.info("========================================");
// 设置 ThreadLocal 上下文(LookupKnowledgeTool 通过此获取 sessionId)
SessionContextHolder.setSessionId(sessionId);
return answer;
try {
var response = agent.call(question);
long duration = System.currentTimeMillis() - startTime;
String answer = response.getText();
// 更新诊断会话
session.setStatus("SUCCESS");
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
logger.info("⏱️ 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
logger.info("========================================");
return answer;
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
throw e;
} finally {
SessionContextHolder.clear();
}
}
/**
* 根据问题复杂度自动选择执行策略
* @param chatModel 聊天模型
* @param toolCallbacks 工具回调
* @param question 用户问题
* @param history 历史消息
* @return AI 回复
*/
public String executeChatWithStrategy(ChatModel chatModel, ToolCallback[] toolCallbacks,
String question, List<Map<String, String>> history) throws GraphRunnerException {
if (QuestionComplexity.isComplex(question)) {
logger.info("📊 问题判定为复杂,使用多 Agent(Planner + Executor)执行");
return executeChatComplex(chatModel, toolCallbacks, question, history);
} else {
logger.info("📊 问题判定为简单,使用单 Agent 执行");
String systemPrompt = buildSystemPrompt(history);
ReactAgent agent = createReactAgent(chatModel, systemPrompt);
return executeChat(agent, question);
}
}
/**
* 多 Agent 复杂对话执行(Planner + Executor + Supervisor)
*/
public String executeChatComplex(ChatModel chatModel, ToolCallback[] toolCallbacks,
String question, List<Map<String, String>> history) throws GraphRunnerException {
String sessionId = UUID.randomUUID().toString().substring(0, 8);
long startTime = System.currentTimeMillis();
DiagnosisSession session = DiagnosisSession.builder()
.sessionId(sessionId)
.query(question)
.status("RUNNING")
.agentFlow("CHAT")
.build();
diagnosisSessionRepository.save(session);
SessionContextHolder.setSessionId(sessionId);
try {
ReactAgent planner = buildChatPlannerAgent(chatModel, toolCallbacks, history);
ReactAgent executor = buildChatExecutorAgent(chatModel, toolCallbacks, history);
SupervisorAgent supervisor = SupervisorAgent.builder()
.name("chat_supervisor")
.description("负责调度 Planner 与 Executor 的多 Agent 控制器")
.model(chatModel)
.systemPrompt("你是一个智能任务调度器。分析用户问题,调用 Planner 拆解步骤,调用 Executor 执行各步骤。")
.subAgents(List.of(planner, executor))
.build();
Optional<OverAllState> stateOptional = supervisor.invoke(question);
long duration = System.currentTimeMillis() - startTime;
String answer = null;
if (stateOptional.isPresent()) {
// 从 state 中提取 Executor 的最终输出
OverAllState state = stateOptional.get();
Optional<AssistantMessage> executorOutput = state.value("executor_feedback")
.filter(AssistantMessage.class::isInstance)
.map(AssistantMessage.class::cast);
if (executorOutput.isPresent()) {
answer = executorOutput.get().getText();
}
}
if (answer == null || answer.isBlank()) {
answer = "抱歉,多 Agent 分析未能生成有效结论。";
}
session.setStatus("SUCCESS");
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
logger.info("⏱️ 多 Agent 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
return answer;
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
logger.error("多 Agent 执行失败", e);
return "执行失败: " + e.getMessage();
} finally {
SessionContextHolder.clear();
}
}
private ReactAgent buildChatPlannerAgent(ChatModel chatModel, ToolCallback[] toolCallbacks,
List<Map<String, String>> history) {
StringBuilder prompt = new StringBuilder(chatPlannerPrompt);
if (!history.isEmpty()) {
prompt.append("\n\n--- 对话历史 ---\n");
for (Map<String, String> msg : history) {
prompt.append(msg.get("role")).append(": ").append(msg.get("content")).append("\n");
}
prompt.append("--- 对话历史结束 ---\n");
}
return ReactAgent.builder()
.name("chat_planner")
.description("负责拆解问题、规划步骤")
.model(chatModel)
.systemPrompt(prompt.toString())
// Planner 不注入工具,只能规划不能执行
.hooks(new AgentLoggingHook(agentStepRepository, "planner"))
.outputKey("planner_plan")
.build();
}
private ReactAgent buildChatExecutorAgent(ChatModel chatModel, ToolCallback[] toolCallbacks,
List<Map<String, String>> history) {
StringBuilder prompt = new StringBuilder(chatExecutorPrompt);
if (!history.isEmpty()) {
prompt.append("\n\n--- 对话历史 ---\n");
for (Map<String, String> msg : history) {
prompt.append(msg.get("role")).append(": ").append(msg.get("content")).append("\n");
}
prompt.append("--- 对话历史结束 ---\n");
}
return ReactAgent.builder()
.name("chat_executor")
.description("负责执行具体步骤并及时反馈")
.model(chatModel)
.systemPrompt(prompt.toString())
.methodTools(buildMethodToolsArray())
.tools(toolCallbacks)
.hooks(new AgentLoggingHook(agentStepRepository, "executor"))
.outputKey("executor_feedback")
.build();
}
/** 从 agent_step 汇总 token、步数等指标回填 diagnosis_session */
private void backfillSessionMetrics(DiagnosisSession session) {
try {
List<com.superbiz.agent.domain.entity.AgentStep> steps =
agentStepRepository.findBySessionIdOrderByStepIndex(session.getSessionId());
if (steps.isEmpty()) return;
int totalTokens = 0;
int stepCount = 0;
int toolCallCount = 0;
for (var s : steps) {
stepCount++;
if (s.getTokenCount() != null) totalTokens += s.getTokenCount();
if (Boolean.TRUE.equals(s.getHasToolCall())) toolCallCount++;
}
session.setTotalTokenCount(totalTokens);
session.setStepCount(stepCount);
session.setToolCallCount(toolCallCount);
} catch (Exception e) {
logger.warn("回填会话指标失败: sessionId={}", session.getSessionId(), e);
}
}
}
@@ -1,8 +1,11 @@
package com.superbiz.agent.tool;
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.VectorSearchService;
import com.superbiz.agent.util.SessionContextHolder;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.tool.annotation.Tool;
import org.springframework.beans.factory.annotation.Autowired;
@@ -25,6 +28,9 @@ public class LookupKnowledgeTool {
@Autowired
private VectorSearchService vectorSearchService;
@Autowired
private ToolInvocationRepository toolInvocationRepository;
/**
* 查询知识库文档
*
@@ -134,9 +140,120 @@ public class LookupKnowledgeTool {
log.info("========================================");
// 记录 tool_invocation(持久化检索明细)
saveToolInvocation(query, l0Matches, l1Results, highConfidence, startTime, result);
return result;
}
/**
* 保存工具调用明细到 tool_invocation 表
*/
private void saveToolInvocation(String query, List<KnowledgeEntry> l0Matches,
List<VectorSearchService.SearchResult> l1Results,
boolean highConfidence, long startTime, LookupResult result) {
try {
String sessionId = SessionContextHolder.getSessionId();
if (sessionId == null) return; // 非会话上下文不记录
boolean hasL0 = l0Matches != null && !l0Matches.isEmpty();
boolean hasL1 = l1Results != null && !l1Results.isEmpty();
long duration = System.currentTimeMillis() - startTime;
String layer;
String outputPreview = null;
int outputLength = 0;
int l0Count = 0;
int l1Count = 0;
boolean truncated = false;
if (hasL0 && !highConfidence) {
layer = "L0+L1";
l0Count = l0Matches.size();
l1Count = l1Results.size();
} else if (hasL0) {
layer = "L0";
l0Count = l0Matches.size();
} else if (hasL1) {
layer = "L1";
l1Count = l1Results.size();
} else {
layer = null;
}
// 拼接 output_preview(前500字符)
if (result != null && result.getPrimary() != null && result.getPrimary().getContent() != null) {
String content = result.getPrimary().getContent();
outputLength = content.length();
if (content.length() > 500) {
outputPreview = content.substring(0, 500) + "...";
truncated = true;
} else {
outputPreview = content;
}
} else if (l1Results != null && !l1Results.isEmpty() && l1Results.get(0).getContent() != null) {
String content = l1Results.get(0).getContent();
outputLength = content.length();
if (content.length() > 500) {
outputPreview = content.substring(0, 500) + "...";
truncated = true;
} else {
outputPreview = content;
}
}
// 构建检索明细 JSON
StringBuilder details = new StringBuilder("{");
if (hasL0) {
details.append("\"l0_titles\":[");
for (int i = 0; i < Math.min(3, l0Matches.size()); i++) {
if (i > 0) details.append(",");
details.append("\"").append(escapeJson(l0Matches.get(i).getTitle())).append("\"");
}
details.append("]");
}
if (hasL1) {
if (hasL0) details.append(",");
details.append("\"l1_scores\":[");
for (int i = 0; i < Math.min(3, l1Results.size()); i++) {
if (i > 0) details.append(",");
details.append(l1Results.get(i).getScore());
}
details.append("]");
}
details.append("}");
ToolInvocation inv = ToolInvocation.builder()
.sessionId(sessionId)
.toolName("lookup_knowledge")
.inputParams("{\"query\":\"" + escapeJson(query) + "\"}")
.outputPreview(outputPreview)
.outputLength(outputLength)
.retrievalLayer(layer)
.l0MatchCount(hasL0 ? l0Count : null)
.l1MatchCount(hasL1 ? l1Count : null)
.isTruncated(truncated)
.retrievalDetails(details.toString())
.durationMs((int) duration)
.success(true)
.build();
toolInvocationRepository.save(inv);
log.debug("tool_invocation 已保存: sessionId={}, layer={}, duration={}ms", sessionId, layer, duration);
} catch (Exception e) {
log.error("保存 tool_invocation 失败", e);
}
}
private String escapeJson(String s) {
if (s == null) return "";
return s.replace("\\", "\\\\")
.replace("\"", "\\\"")
.replace("\n", "\\n")
.replace("\r", "\\r")
.replace("\t", "\\t");
}
/**
* 组装查询结果
*
@@ -0,0 +1,44 @@
package com.superbiz.agent.util;
import java.util.List;
/**
* 问题复杂度判断
* 用于决定使用单 Agent 还是多 Agent(Planner + Executor)处理
*/
public class QuestionComplexity {
/** 复杂问题关键词 — 需要多步分析、排查、根因定位 */
private static final List<String> COMPLEX_KEYWORDS = List.of(
"排查", "分析", "为什么", "根因", "调查", "对比", "影响范围",
"原因", "故障", "告警", "诊断", "链路", "流程", "步骤",
"root cause", "troubleshoot", "investigate"
);
/** 极简问题关键词 — 快速回答,无需多 Agent */
private static final List<String> SIMPLE_KEYWORDS = List.of(
"是什么", "查一下", "什么是", "时间", "天气", "定义",
"查", "找", "what is", "define", "time"
);
/**
* 判断是否为复杂问题
*/
public static boolean isComplex(String question) {
if (question == null || question.isBlank()) return false;
String q = question.toLowerCase();
// 复杂关键词匹配 → 多 Agent
for (String kw : COMPLEX_KEYWORDS) {
if (q.contains(kw)) return true;
}
// 简单关键词匹配 → 单 Agent
for (String kw : SIMPLE_KEYWORDS) {
if (q.contains(kw)) return false;
}
// 默认:长问题(>30 字)视为复杂,短问题视为简单
return question.length() > 30;
}
}
@@ -0,0 +1,39 @@
package com.superbiz.agent.util;
/**
* 会话上下文持有者(基于 ThreadLocal)
* <p>
* 用于在执行链路中传递 sessionId 和 agentName,覆盖 AgentLoggingHook 和
* LookupKnowledgeTool 等无法直接通过 RunnableConfig 获取上下文的组件。
* <p>
* 使用规范:
* 1. 调用方(ChatService/AiOpsService)在 Agent 执行前调用 setSessionId() 和 setAgentName()
* 2. AgentLoggingHook 和工具类通过 getSessionId() / getAgentName() 读取
* 3. 必须在 finally 块中调用 clear(),防止内存泄漏和线程污染
*/
public class SessionContextHolder {
private static final ThreadLocal<String> SESSION_ID = new ThreadLocal<>();
private static final ThreadLocal<String> AGENT_NAME = new ThreadLocal<>();
public static void setSessionId(String sessionId) {
SESSION_ID.set(sessionId);
}
public static String getSessionId() {
return SESSION_ID.get();
}
public static void setAgentName(String agentName) {
AGENT_NAME.set(agentName);
}
public static String getAgentName() {
return AGENT_NAME.get();
}
public static void clear() {
SESSION_ID.remove();
AGENT_NAME.remove();
}
}