feat(trace): isolate chat runs

This commit is contained in:
zhuyongxin
2026-07-10 19:02:04 +08:00
parent 6fdbd34bab
commit 26d5529280
21 changed files with 610 additions and 134 deletions
@@ -96,10 +96,11 @@ public class ChatController {
// 更新会话历史
session.addChatMessagePair(request.getQuestion(), fullAnswer, MAX_WINDOW_SIZE);
sessionManager.updateSession(session);
chatService.syncChatSessionMetadata(session.getSessionId(), session.getMessagePairCount());
logger.info("已更新会话历史 - SessionId: {}, 当前消息对数: {}",
session.getSessionId(), session.getMessagePairCount());
return ResponseEntity.ok(ApiResponse.success(ChatResponse.success(fullAnswer, result.sessionId())));
return ResponseEntity.ok(ApiResponse.success(ChatResponse.success(fullAnswer, result.sessionId(), result.runId())));
} catch (Exception e) {
logger.error("对话失败", e);
@@ -413,12 +414,14 @@ public class ChatController {
private String answer;
private String errorMessage;
private String sessionId;
private String runId;
public static ChatResponse success(String answer, String sessionId) {
public static ChatResponse success(String answer, String sessionId, String runId) {
ChatResponse response = new ChatResponse();
response.setSuccess(true);
response.setAnswer(answer);
response.setSessionId(sessionId);
response.setRunId(runId);
return response;
}
@@ -47,11 +47,13 @@ public class AgentLoggingHook extends MessagesModelHook {
@Override
public AgentCommand beforeModel(List<Message> previousMessages, RunnableConfig config) {
String sessionId = resolveSessionId(config);
String runId = resolveRunId(config);
String traceScopeId = traceScopeId(sessionId, runId);
boolean hasSession = sessionId != null;
int stepIndex = 0;
if (hasSession) {
stepIndex = stepCounters.merge(sessionId, 0, (oldValue, ignored) -> oldValue + 1);
if (traceScopeId != null) {
stepIndex = stepCounters.merge(traceScopeId, 0, (oldValue, ignored) -> oldValue + 1);
}
log.info("========================================");
@@ -73,12 +75,13 @@ public class AgentLoggingHook extends MessagesModelHook {
try {
AgentStep step = AgentStep.builder()
.sessionId(sessionId)
.runId(runId)
.stepIndex(stepIndex)
.agentName(agentName)
.modelInput(buildModelInputSummary(previousMessages))
.build();
AgentStep saved = agentStepRepository.save(step);
pendingSteps.put(sessionId + "_" + stepIndex, Map.of(
pendingSteps.put(stepKey(traceScopeId, stepIndex), Map.of(
"stepId", saved.getId(),
"startTime", System.currentTimeMillis()
));
@@ -93,7 +96,9 @@ public class AgentLoggingHook extends MessagesModelHook {
@Override
public AgentCommand afterModel(List<Message> previousMessages, RunnableConfig config) {
String sessionId = resolveSessionId(config);
int stepIndex = sessionId == null ? 0 : stepCounters.getOrDefault(sessionId, 0);
String runId = resolveRunId(config);
String traceScopeId = traceScopeId(sessionId, runId);
int stepIndex = traceScopeId == null ? 0 : stepCounters.getOrDefault(traceScopeId, 0);
log.info("========================================");
log.info("*** [AgentTrace] agent={}, phase=after_model, stepIndex={}", agentName, stepIndex);
@@ -122,7 +127,7 @@ public class AgentLoggingHook extends MessagesModelHook {
log.info("========================================");
if (sessionId != null) {
String stepKey = sessionId + "_" + stepIndex;
String stepKey = stepKey(traceScopeId, stepIndex);
Map<String, Object> pending = pendingSteps.remove(stepKey);
if (pending != null) {
try {
@@ -162,6 +167,23 @@ public class AgentLoggingHook extends MessagesModelHook {
.orElseGet(SessionContextHolder::getSessionId);
}
private String resolveRunId(RunnableConfig config) {
return config.metadata("runId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getRunId);
}
private String traceScopeId(String sessionId, String runId) {
if (runId != null && !runId.isBlank()) {
return runId;
}
return sessionId;
}
private String stepKey(String traceScopeId, int stepIndex) {
return traceScopeId + "_" + stepIndex;
}
private AssistantMessage findLastAssistant(List<Message> previousMessages) {
for (int i = previousMessages.size() - 1; i >= 0; i--) {
if (previousMessages.get(i) instanceof AssistantMessage assistantMessage) {
@@ -59,13 +59,17 @@ public class VerifierInputHook extends MessagesModelHook {
String sessionId = config.metadata("sessionId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getSessionId);
String runId = config.metadata("runId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getRunId);
String executorFinalAnswer = VerifierContextHolder.getExecutorFinalAnswer();
if (executorFinalAnswer == null || executorFinalAnswer.isBlank()) {
executorFinalAnswer = extractLastAssistantText(previousMessages);
}
List<Map<String, Object>> toolTraceSummary =
toolTraceSummaryService.buildVerifierTraceSummary(sessionId, executorFinalAnswer);
List<Map<String, Object>> toolTraceSummary = runId == null || runId.isBlank()
? toolTraceSummaryService.buildVerifierTraceSummary(sessionId, executorFinalAnswer)
: toolTraceSummaryService.buildVerifierTraceSummaryForRun(runId, executorFinalAnswer);
VerifierContextHolder.setToolTraceSummary(toolTraceSummary);
ExecutorOutputParseResult parseResult = parseExecutorOutput(executorFinalAnswer);
@@ -76,7 +80,7 @@ public class VerifierInputHook extends MessagesModelHook {
VerifierContextHolder.setExecutorStructuredOutput(parseResult.structuredOutput());
VerifierContextHolder.setExecutorOutputParseStatus(parseResult.status());
Map<String, Object> gatekeeperResult = runGatekeeper(sessionId, parseResult);
Map<String, Object> gatekeeperResult = runGatekeeper(sessionId, runId, parseResult);
VerifierContextHolder.setGatekeeperResult(gatekeeperResult);
Map<String, Object> verifierInput = new LinkedHashMap<>();
@@ -96,11 +100,14 @@ public class VerifierInputHook extends MessagesModelHook {
}
}
private Map<String, Object> runGatekeeper(String sessionId, ExecutorOutputParseResult parseResult) {
private Map<String, Object> runGatekeeper(String sessionId, String runId, ExecutorOutputParseResult parseResult) {
if (executorGatekeeperService == null) {
return passGatekeeperResult();
}
try {
if (runId != null && !runId.isBlank()) {
return executorGatekeeperService.validateRun(runId, parseResult.structuredOutput(), parseResult.status());
}
return executorGatekeeperService.validate(sessionId, parseResult.structuredOutput(), parseResult.status());
} catch (Exception e) {
log.error("Gatekeeper validation failed unexpectedly", e);
@@ -14,14 +14,16 @@ 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.DiagnosisSession;
import com.superbiz.agent.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.hook.AgentLoggingHook;
import com.superbiz.agent.hook.PlannerSkillMetadataHook;
import com.superbiz.agent.hook.TokenTrackingChatModel;
import com.superbiz.agent.hook.TokenUsageHolder;
import com.superbiz.agent.hook.VerifierInputHook;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ChatSessionRepository;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
import com.superbiz.agent.tool.LookupKnowledgeTool;
import com.superbiz.agent.tool.RetrievedDocTracker;
@@ -43,6 +45,7 @@ import org.springframework.stereotype.Service;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.time.LocalDateTime;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
@@ -62,8 +65,8 @@ public class ChatService {
private static final String DEGRADED_PREFIX = "当前无法基于已获取证据生成可靠结论,建议人工介入。";
private static final String CHAT_PROMPT_AUDIT_VERSION = "chat-prompts-v1";
/** 封装 answer + 后端生成的 sessionId,用于 feedback 关联 */
public record ChatResult(String answer, String sessionId) {}
/** 封装 answer + 后端生成的 sessionId/runId,用于 feedback 关联 */
public record ChatResult(String answer, String sessionId, String runId) {}
@Autowired
private InternalDocsTools internalDocsTools;
@@ -74,7 +77,7 @@ public class ChatService {
@Autowired
private QueryMetricsTools queryMetricsTools;
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过mcp配置注入
@Autowired(required = false) // Mock 模式下才注册,所以设置为 optional,真实环境通过 mcp 配置注入
private QueryLogsTools queryLogsTools;
@Autowired(required = false)
@@ -87,7 +90,10 @@ public class ChatService {
private LookupKnowledgeTool lookupKnowledgeTool;
@Autowired
private DiagnosisSessionRepository diagnosisSessionRepository;
private ChatSessionRepository chatSessionRepository;
@Autowired
private DiagnosisRunRepository diagnosisRunRepository;
@Autowired
private AgentStepRepository agentStepRepository;
@@ -131,7 +137,7 @@ public class ChatService {
@PostConstruct
public void init() {
// 加载 Prompt
// 鍔犺浇 Prompt
try {
chatPlannerPrompt = new String(
new ClassPathResource("prompts/chat-planner-prompt.md").getInputStream().readAllBytes(),
@@ -186,7 +192,7 @@ public class ChatService {
String role = msg.get("role");
String content = msg.get("content");
// 🔧 过滤时间查询相关的历史消息,避免 LLM 复用旧的时间信息
// 过滤时间查询相关的历史消息,避免 LLM 复用旧的时间信息
if ("user".equals(role) && isTimeQuery(content)) {
continue; // 跳过时间查询问题
}
@@ -229,7 +235,7 @@ public class ChatService {
return false;
}
// 匹配日期时间格式:2026年5月31日、15:57、下午3点 等
return content.matches(".*(\\d{4}年\\d{1,2}月\\d{1,2}日|\\d{1,2}:\\d{2}|[上下午]+\\d{1,2}[点时]).*");
return content.matches(".*(\\d{4}.*\\d{1,2}.*\\d{1,2}.*|\\d{1,2}:\\d{2}|[上下]午\\s*\\d{1,2}[点时]).*");
}
/**
@@ -250,7 +256,7 @@ public class ChatService {
}
/**
* 获取工具回调列表,mcp服务提供的工具
* 获取工具回调列表,mcp 服务提供的工具
*/
public ToolCallback[] getToolCallbacks() {
if (tools == null) {
@@ -260,7 +266,7 @@ public class ChatService {
}
/**
* 记录可用工具列表:mcp服务提供的工具
* 记录可用工具列表:mcp 服务提供的工具
*/
public void logAvailableTools() {
if (tools == null) {
@@ -295,7 +301,7 @@ public class ChatService {
* 执行 ReactAgent 对话(非流式)
* @param agent ReactAgent 实例
* @param question 用户问题
* @return ChatResult(answer + sessionId)
* @return ChatResult(answer + sessionId + runId)
*/
public ChatResult executeChat(ReactAgent agent, String question) throws GraphRunnerException {
return executeChat(agent, question, null);
@@ -303,22 +309,24 @@ public class ChatService {
public ChatResult executeChat(ReactAgent agent, String question, String requestedSessionId) throws GraphRunnerException {
logger.info("========================================");
logger.info("📝 用户问题: {}", question);
logger.info("用户问题: {}", question);
String sessionId = resolveSessionId(requestedSessionId);
String runId = newRunId();
long startTime = System.currentTimeMillis();
// 创建或更新诊断会话
DiagnosisSession session = startDiagnosisSession(sessionId, question);
diagnosisSessionRepository.save(session);
// 创建或更新 Chat Session 元数据,并创建本次诊断 run
ensureChatSession(sessionId, null);
DiagnosisRun run = startDiagnosisRun(sessionId, runId, question);
// 设置 ThreadLocal 上下文(LookupKnowledgeTool 通过此获取 sessionId)
SessionContextHolder.setSessionId(sessionId);
// 设置 ThreadLocal 上下文,供工具和 Hook 读取 sessionId/runId
SessionContextHolder.setContext(sessionId, runId);
try {
// 通过 RunnableConfig 将 sessionId 传入 Hook(线程安全,异步也兼容)
// 通过 RunnableConfig 将 sessionId/runId 传入 Hook
var config = RunnableConfig.builder()
.addMetadata("sessionId", sessionId)
.addMetadata("runId", runId)
.build();
var response = agent.call(question, config);
@@ -326,23 +334,27 @@ public class ChatService {
String answer = response.getText();
// 更新诊断会话
session.setStatus("SUCCESS");
session.setAnswer(answer);
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
// 更新诊断 run
run.setStatus("SUCCESS");
run.setAnswer(answer);
run.setTotalDurationMs((int) duration);
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
evaluationService.evaluate(sessionId, answer);
evaluationService.evaluateRun(runId, answer);
logger.info("⏱️ 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
logger.info("总耗时: {} ms", duration);
logger.info("输出长度: {} 字符", answer.length());
logger.info("========================================");
return new ChatResult(answer, sessionId);
return new ChatResult(answer, sessionId, runId);
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
String errorAnswer = "Execution failed: " + e.getMessage();
run.setStatus("FAILED");
run.setAnswer(errorAnswer);
run.setTotalDurationMs((int) (System.currentTimeMillis() - startTime));
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
throw e;
} finally {
retrievedDocTracker.clearSession(sessionId);
@@ -356,7 +368,7 @@ public class ChatService {
* @param toolCallbacks 工具回调
* @param question 用户问题
* @param history 历史消息
* @return ChatResult(answer + sessionId)
* @return ChatResult(answer + sessionId + runId)
*/
public ChatResult executeChatWithStrategy(ChatModel chatModel, ToolCallback[] toolCallbacks,
String question, List<Map<String, String>> history) throws GraphRunnerException {
@@ -367,10 +379,10 @@ public class ChatService {
String question, List<Map<String, String>> history,
String requestedSessionId) throws GraphRunnerException {
if (QuestionComplexity.isComplex(question)) {
logger.info("📊 问题判定为复杂,使用多 Agent(Planner + Executor)执行");
logger.info("问题判定为复杂,使用多 Agent(Planner + Executor)执行");
return executeChatComplex(chatModel, toolCallbacks, question, history, requestedSessionId);
} else {
logger.info("📊 问题判定为简单,使用单 Agent 执行");
logger.info("问题判定为简单,使用单 Agent 执行");
String systemPrompt = buildSystemPrompt(history);
ReactAgent agent = createReactAgent(chatModel, systemPrompt);
return executeChat(agent, question, requestedSessionId);
@@ -389,12 +401,13 @@ public class ChatService {
String question, List<Map<String, String>> history,
String requestedSessionId) throws GraphRunnerException {
String sessionId = resolveSessionId(requestedSessionId);
String runId = newRunId();
long startTime = System.currentTimeMillis();
DiagnosisSession session = startDiagnosisSession(sessionId, question);
diagnosisSessionRepository.save(session);
ensureChatSession(sessionId, history == null ? null : history.size() / 2);
DiagnosisRun run = startDiagnosisRun(sessionId, runId, question);
SessionContextHolder.setSessionId(sessionId);
SessionContextHolder.setContext(sessionId, runId);
VerifierContextHolder.setOriginalQuery(question);
VerifierContextHolder.setRetryContext(null);
VerifierContextHolder.setExecutorFinalAnswer(null);
@@ -405,6 +418,7 @@ public class ChatService {
String answer = null;
RunnableConfig config = RunnableConfig.builder()
.addMetadata("sessionId", sessionId)
.addMetadata("runId", runId)
.build();
for (int round = 1; round <= 2; round++) {
@@ -427,7 +441,7 @@ public class ChatService {
finalDecision = buildVerifierFallbackDecision(round, "workflow 未返回有效状态");
ComposerRenderResult renderResult = buildFixedFallbackAnswer(question, finalDecision);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
@@ -449,21 +463,21 @@ public class ChatService {
finalDecision = buildVerifierFallbackDecision(round, "verifier_output 缺失或无法解析");
ComposerRenderResult renderResult = buildFixedFallbackAnswer(question, finalDecision);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
if ("PASS".equals(finalDecision.verdict())) {
ComposerRenderResult renderResult = composeFinalAnswer(chatModel, question, finalDecision, config);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
if ("REJECT".equals(finalDecision.verdict())) {
ComposerRenderResult renderResult = composeFinalAnswer(chatModel, question, finalDecision, config);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
@@ -473,12 +487,12 @@ public class ChatService {
if (!shouldRetry) {
ComposerRenderResult renderResult = composeFinalAnswer(chatModel, question, finalDecision, config);
answer = renderResult.answer();
persistVerifierEvaluation(session, finalDecision, round, renderResult.audit());
persistVerifierEvaluation(run, finalDecision, round, renderResult.audit());
break;
}
retryContext = buildRetryContext(finalDecision);
persistVerifierEvaluation(session, finalDecision, round);
persistVerifierEvaluation(run, finalDecision, round);
}
long duration = System.currentTimeMillis() - startTime;
@@ -487,24 +501,28 @@ public class ChatService {
answer = "抱歉,多 Agent 分析未能生成有效结论。";
}
session.setStatus("SUCCESS");
session.setAnswer(answer);
session.setTotalDurationMs((int) duration);
backfillSessionMetrics(session);
diagnosisSessionRepository.save(session);
run.setStatus("SUCCESS");
run.setAnswer(answer);
run.setTotalDurationMs((int) duration);
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
evaluationService.evaluate(sessionId, answer);
evaluationService.evaluateRun(runId, answer);
logger.info("⏱️ 多 Agent 总耗时: {} ms", duration);
logger.info("📏 输出长度: {} 字符", answer.length());
logger.info("多 Agent 总耗时: {} ms", duration);
logger.info("输出长度: {} 字符", answer.length());
return new ChatResult(answer, sessionId);
return new ChatResult(answer, sessionId, runId);
} catch (Exception e) {
session.setStatus("FAILED");
diagnosisSessionRepository.save(session);
String errorAnswer = "Execution failed: " + e.getMessage();
run.setStatus("FAILED");
run.setAnswer(errorAnswer);
run.setTotalDurationMs((int) (System.currentTimeMillis() - startTime));
backfillRunMetrics(run);
diagnosisRunRepository.save(run);
logger.error("多 Agent 执行失败", e);
return new ChatResult("执行失败: " + e.getMessage(), sessionId);
return new ChatResult(errorAnswer, sessionId, runId);
} finally {
retrievedDocTracker.clearSession(sessionId);
SessionContextHolder.clear();
@@ -516,7 +534,7 @@ public class ChatService {
String retryContext) {
StringBuilder prompt = new StringBuilder(chatPlannerPrompt);
// 注入 knowledge map
// 娉ㄥ叆 knowledge map
String knowledgeMap = knowledgeDomainService.buildKnowledgeMap();
if (!knowledgeMap.isBlank()) {
prompt.append("\n\n## 可用知识库\n\n").append(knowledgeMap);
@@ -530,7 +548,7 @@ public class ChatService {
prompt.append("--- 对话历史结束 ---\n");
}
if (retryContext != null && !retryContext.isBlank()) {
prompt.append("\n\n--- 本轮补证据约束 ---\n").append(retryContext).append("\n");
prompt.append("\n\n--- 鏈疆琛ヨ瘉鎹害鏉?---\n").append(retryContext).append("\n");
}
return ReactAgent.builder()
.name("chat_planner")
@@ -580,7 +598,7 @@ public class ChatService {
prompt.append("--- 对话历史结束 ---\n");
}
if (retryContext != null && !retryContext.isBlank()) {
prompt.append("\n\n--- 本轮补证据约束 ---\n").append(retryContext).append("\n");
prompt.append("\n\n--- 鏈疆琛ヨ瘉鎹害鏉?---\n").append(retryContext).append("\n");
}
return ReactAgent.builder()
.name("chat_executor")
@@ -620,21 +638,40 @@ public class ChatService {
return UUID.randomUUID().toString().substring(0, 8);
}
private DiagnosisSession startDiagnosisSession(String sessionId, String question) {
DiagnosisSession session = diagnosisSessionRepository.findBySessionId(sessionId)
.orElseGet(() -> DiagnosisSession.builder()
private String newRunId() {
return "run-" + UUID.randomUUID();
}
public void syncChatSessionMetadata(String sessionId, Integer messagePairCount) {
if (sessionId == null || sessionId.isBlank()) {
return;
}
ensureChatSession(sessionId, messagePairCount);
}
private ChatSession ensureChatSession(String sessionId, Integer messagePairCount) {
ChatSession chatSession = chatSessionRepository.findBySessionId(sessionId)
.orElseGet(() -> ChatSession.builder()
.sessionId(sessionId)
.agentFlow("CHAT")
.status("ACTIVE")
.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;
chatSession.setStatus("ACTIVE");
chatSession.setLastActiveAt(LocalDateTime.now());
if (messagePairCount != null) {
chatSession.setMessagePairCount(messagePairCount);
}
return chatSessionRepository.save(chatSession);
}
private DiagnosisRun startDiagnosisRun(String sessionId, String runId, String question) {
DiagnosisRun run = DiagnosisRun.builder()
.runId(runId)
.sessionId(sessionId)
.query(question)
.status("RUNNING")
.agentFlow("CHAT")
.build();
return diagnosisRunRepository.save(run);
}
private String buildWorkflowInput(String question, String retryContext) {
@@ -867,11 +904,11 @@ public class ChatService {
.orElse(null);
}
private void persistVerifierEvaluation(DiagnosisSession session, VerifierDecision decision, int round) {
persistVerifierEvaluation(session, decision, round, null);
private void persistVerifierEvaluation(DiagnosisRun run, VerifierDecision decision, int round) {
persistVerifierEvaluation(run, decision, round, null);
}
private void persistVerifierEvaluation(DiagnosisSession session, VerifierDecision decision, int round,
private void persistVerifierEvaluation(DiagnosisRun run, VerifierDecision decision, int round,
Map<String, Object> composerOutput) {
if (decision == null) {
return;
@@ -899,9 +936,9 @@ public class ChatService {
verifierEvaluation.put("composer_output", composerOutput);
}
String merged = selfEvaluationMergeService.mergeVerifierEvaluation(session.getSelfEvaluation(), verifierEvaluation);
session.setSelfEvaluation(merged);
diagnosisSessionRepository.save(session);
String merged = selfEvaluationMergeService.mergeVerifierEvaluation(run.getSelfEvaluation(), verifierEvaluation);
run.setSelfEvaluation(merged);
diagnosisRunRepository.save(run);
}
private Map<String, Object> promptAuditSnapshot() {
@@ -1244,7 +1281,10 @@ public class ChatService {
private List<String> buildNextStepSuggestionsFromTrace() {
List<String> suggestions = new ArrayList<>();
List<Map<String, Object>> toolSummary = toolTraceSummaryService.buildVerifierTraceSummary(SessionContextHolder.getSessionId(), null);
String runId = SessionContextHolder.getRunId();
List<Map<String, Object>> toolSummary = runId == null || runId.isBlank()
? toolTraceSummaryService.buildVerifierTraceSummary(SessionContextHolder.getSessionId(), null)
: toolTraceSummaryService.buildVerifierTraceSummaryForRun(runId, null);
boolean hasKnowledgeTool = toolSummary.stream().anyMatch(item -> "lookup_knowledge".equals(item.get("tool_name")));
boolean hasFailedEvidence = toolSummary.stream().anyMatch(item -> !Boolean.TRUE.equals(item.get("success")));
@@ -1309,11 +1349,11 @@ public class ChatService {
private record ComposerRenderResult(String answer, Map<String, Object> audit) {
}
/** 从 agent_step 和 tool_invocation 汇总指标回填 diagnosis_session */
private void backfillSessionMetrics(DiagnosisSession session) {
/** 从 agent_step 和 tool_invocation 汇总指标回填 diagnosis_run */
private void backfillRunMetrics(DiagnosisRun run) {
try {
List<com.superbiz.agent.domain.entity.AgentStep> steps =
agentStepRepository.findBySessionIdOrderByStepIndex(session.getSessionId());
agentStepRepository.findByRunIdOrderByStepIndex(run.getRunId());
int totalTokens = 0;
int stepCount = 0;
@@ -1321,12 +1361,12 @@ public class ChatService {
stepCount++;
if (s.getTokenCount() != null) totalTokens += s.getTokenCount();
}
long toolCallCount = toolInvocationRepository.countBySessionId(session.getSessionId());
session.setTotalTokenCount(totalTokens);
session.setStepCount(stepCount);
session.setToolCallCount(Math.toIntExact(toolCallCount));
long toolCallCount = toolInvocationRepository.countByRunId(run.getRunId());
run.setTotalTokenCount(totalTokens);
run.setStepCount(stepCount);
run.setToolCallCount(Math.toIntExact(toolCallCount));
} catch (Exception e) {
logger.warn("回填会话指标失败: sessionId={}", session.getSessionId(), e);
logger.warn("Failed to backfill run metrics: runId={}", run.getRunId(), e);
}
}
}
@@ -1,7 +1,9 @@
package com.superbiz.agent.service;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
import org.slf4j.Logger;
@@ -29,6 +31,9 @@ public class EvaluationService {
@Autowired
private DiagnosisSessionRepository diagnosisSessionRepository;
@Autowired
private DiagnosisRunRepository diagnosisRunRepository;
@Autowired
private ToolInvocationRepository toolInvocationRepository;
@@ -40,7 +45,7 @@ public class EvaluationService {
diagnosisSessionRepository.findBySessionId(sessionId).ifPresent(session -> {
try {
List<ToolInvocation> toolInvocations = toolInvocationRepository.findBySessionId(sessionId);
Map<String, Object> ruleEvaluation = evaluateWithRules(session, toolInvocations);
Map<String, Object> ruleEvaluation = evaluateWithRules(session.getStatus(), toolInvocations);
String merged = selfEvaluationMergeService.mergeRuleEvaluation(session.getSelfEvaluation(), ruleEvaluation);
session.setSelfEvaluation(merged);
diagnosisSessionRepository.save(session);
@@ -51,14 +56,30 @@ public class EvaluationService {
});
}
@Async
public void evaluateRun(String runId, String answer) {
diagnosisRunRepository.findByRunId(runId).ifPresent(run -> {
try {
List<ToolInvocation> toolInvocations = toolInvocationRepository.findByRunIdOrderByIdAsc(runId);
Map<String, Object> ruleEvaluation = evaluateWithRules(run.getStatus(), toolInvocations);
String merged = selfEvaluationMergeService.mergeRuleEvaluation(run.getSelfEvaluation(), ruleEvaluation);
run.setSelfEvaluation(merged);
diagnosisRunRepository.save(run);
logger.info("证据评分已写入: runId={}, result={}", runId, merged);
} catch (Exception e) {
logger.error("评分失败: runId={}", runId, e);
}
});
}
// -------------------------------------------------------------------------
// 规则引擎(事实层)
// -------------------------------------------------------------------------
private Map<String, Object> evaluateWithRules(DiagnosisSession session, List<ToolInvocation> invocations) {
private Map<String, Object> evaluateWithRules(String status, List<ToolInvocation> invocations) {
List<Map<String, Object>> factors = new ArrayList<>();
if ("FAILED".equals(session.getStatus())) {
if ("FAILED".equals(status)) {
factors.add(factor("execution_failed", -100, "执行失败"));
return buildResult(0, factors);
}
@@ -61,7 +61,25 @@ public class ExecutorGatekeeperService {
GatekeeperResult result = new GatekeeperResult(ruleCatalog);
validateSchema(structuredOutput, parseStatus, result);
if (structuredOutput != null) {
validateInvocationRefs(sessionId, structuredOutput, result);
List<ToolInvocation> invocations = sessionId == null || sessionId.isBlank()
? List.of()
: toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId);
validateInvocationRefs("session_id", sessionId, invocations, structuredOutput, result);
importWarnings(structuredOutput, result);
}
return result.toMap();
}
public Map<String, Object> validateRun(String runId,
Map<String, Object> structuredOutput,
Map<String, Object> parseStatus) {
GatekeeperResult result = new GatekeeperResult(ruleCatalog);
validateSchema(structuredOutput, parseStatus, result);
if (structuredOutput != null) {
List<ToolInvocation> invocations = runId == null || runId.isBlank()
? List.of()
: toolInvocationRepository.findByRunIdOrderByIdAsc(runId);
validateInvocationRefs("run_id", runId, invocations, structuredOutput, result);
importWarnings(structuredOutput, result);
}
return result.toMap();
@@ -134,15 +152,18 @@ public class ExecutorGatekeeperService {
requireArray(structuredOutput, "missing_info", result);
}
private void validateInvocationRefs(String sessionId, Map<String, Object> structuredOutput, GatekeeperResult result) {
if (sessionId == null || sessionId.isBlank()) {
result.fail(RULE_INVOCATION_REF, "session_id", "session id is required to validate source_invocation_id",
private void validateInvocationRefs(String scopeName,
String scopeId,
List<ToolInvocation> invocations,
Map<String, Object> structuredOutput,
GatekeeperResult result) {
if (scopeId == null || scopeId.isBlank()) {
result.fail(RULE_INVOCATION_REF, scopeName, scopeName + " is required to validate source_invocation_id",
SEVERITY_LOW_CONFID);
return;
}
Map<Long, ToolInvocation> validInvocations = toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId)
.stream()
Map<Long, ToolInvocation> validInvocations = invocations.stream()
.filter(invocation -> invocation.getId() != null)
.collect(Collectors.toMap(ToolInvocation::getId, Function.identity(), (left, right) -> left));
Object claimsValue = structuredOutput.get("claims");
@@ -48,6 +48,9 @@ public class ToolInvocationRecorder {
if (invocation.getSessionId() == null || invocation.getSessionId().isBlank()) {
invocation.setSessionId(SessionContextHolder.getSessionId());
}
if (invocation.getRunId() == null || invocation.getRunId().isBlank()) {
invocation.setRunId(SessionContextHolder.getRunId());
}
if (invocation.getSessionId() == null || invocation.getSessionId().isBlank()) {
log.debug("Skip tool_invocation without sessionId: tool={}", invocation.getToolName());
return;
@@ -42,6 +42,19 @@ public class ToolTraceSummaryService {
}
List<ToolInvocation> invocations = toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId);
return buildVerifierTraceSummary(invocations, executorFinalAnswer);
}
public List<Map<String, Object>> buildVerifierTraceSummaryForRun(String runId, String executorFinalAnswer) {
if (runId == null || runId.isBlank()) {
return List.of();
}
List<ToolInvocation> invocations = toolInvocationRepository.findByRunIdOrderByIdAsc(runId);
return buildVerifierTraceSummary(invocations, executorFinalAnswer);
}
private List<Map<String, Object>> buildVerifierTraceSummary(List<ToolInvocation> invocations, String executorFinalAnswer) {
if (invocations.isEmpty()) {
return List.of();
}
@@ -14,16 +14,31 @@ package com.superbiz.agent.util;
public class SessionContextHolder {
private static final ThreadLocal<String> SESSION_ID = new ThreadLocal<>();
private static final ThreadLocal<String> RUN_ID = new ThreadLocal<>();
public static void setContext(String sessionId, String runId) {
setSessionId(sessionId);
setRunId(runId);
}
public static void setSessionId(String sessionId) {
SESSION_ID.set(sessionId);
}
public static void setRunId(String runId) {
RUN_ID.set(runId);
}
public static String getSessionId() {
return SESSION_ID.get();
}
public static String getRunId() {
return RUN_ID.get();
}
public static void clear() {
SESSION_ID.remove();
RUN_ID.remove();
}
}
@@ -0,0 +1,32 @@
package com.superbiz.agent.controller;
import com.superbiz.agent.service.ChatService;
import org.junit.jupiter.api.Test;
import org.springframework.http.ResponseEntity;
import org.springframework.test.util.ReflectionTestUtils;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verifyNoInteractions;
class ChatControllerTest {
@Test
void blankChatRequestReturnsErrorBeforeCreatingRun() {
ChatController controller = new ChatController();
ChatService chatService = mock(ChatService.class);
ReflectionTestUtils.setField(controller, "chatService", chatService);
ChatController.ChatRequest request = new ChatController.ChatRequest();
request.setId("invalid-chat-session");
request.setQuestion(" ");
ResponseEntity<ChatController.ApiResponse<ChatController.ChatResponse>> response = controller.chat(request);
ChatController.ChatResponse body = response.getBody().getData();
assertFalse(body.isSuccess());
assertEquals("问题内容不能为空", body.getErrorMessage());
verifyNoInteractions(chatService);
}
}
@@ -7,10 +7,12 @@ import com.superbiz.agent.agent.tool.DateTimeTools;
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.DiagnosisSession;
import com.superbiz.agent.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ChatSessionRepository;
import com.superbiz.agent.repository.DiagnosisRunRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
import com.superbiz.agent.tool.LookupKnowledgeTool;
import com.superbiz.agent.tool.RetrievedDocTracker;
@@ -31,11 +33,15 @@ import java.util.concurrent.atomic.AtomicInteger;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNotEquals;
import static org.junit.jupiter.api.Assertions.assertSame;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.ArgumentMatchers.isNull;
import static org.mockito.Mockito.atLeast;
import static org.mockito.Mockito.atLeastOnce;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
@@ -58,8 +64,73 @@ class ChatServiceSequentialAgentTest {
assertTrue(result.answer().contains("连接池 active 达到上限"));
assertFalse(result.answer().contains("\"answer_version\""));
assertEquals("sequential-test-session", result.sessionId());
assertTrue(result.runId().startsWith("run-"));
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "chat_composer"), chatModel.agentCalls);
assertTrue(chatModel.sawVerifierPrompt);
ChatSessionRepository chatSessionRepository =
(ChatSessionRepository) ReflectionTestUtils.getField(chatService, "chatSessionRepository");
DiagnosisRunRepository diagnosisRunRepository =
(DiagnosisRunRepository) ReflectionTestUtils.getField(chatService, "diagnosisRunRepository");
EvaluationService evaluationService =
(EvaluationService) ReflectionTestUtils.getField(chatService, "evaluationService");
ArgumentCaptor<ChatSession> chatSessionCaptor = ArgumentCaptor.forClass(ChatSession.class);
verify(chatSessionRepository, atLeastOnce()).save(chatSessionCaptor.capture());
assertEquals("sequential-test-session", chatSessionCaptor.getValue().getSessionId());
ArgumentCaptor<DiagnosisRun> runCaptor = ArgumentCaptor.forClass(DiagnosisRun.class);
verify(diagnosisRunRepository, atLeastOnce()).save(runCaptor.capture());
DiagnosisRun savedRun = runCaptor.getValue();
assertEquals(result.runId(), savedRun.getRunId());
assertEquals("sequential-test-session", savedRun.getSessionId());
assertEquals("SUCCESS", savedRun.getStatus());
assertEquals(result.answer(), savedRun.getAnswer());
verify(evaluationService).evaluateRun(eq(result.runId()), eq(result.answer()));
}
@Test
void executeChatComplexCreatesDistinctRunsForSameSessionAcrossTurns() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel firstRoundModel = new ScriptedChatModel();
ScriptedChatModel secondRoundModel = new ScriptedChatModel();
String sessionId = "sequential-same-session";
ChatService.ChatResult first = chatService.executeChatComplex(
firstRoundModel,
new ToolCallback[0],
"第一轮:请分析支付超时",
List.of(),
sessionId
);
ChatService.ChatResult second = chatService.executeChatComplex(
secondRoundModel,
new ToolCallback[0],
"第二轮:基于上一轮结论列出缺失证据",
List.of(
Map.of("role", "user", "content", "第一轮:请分析支付超时"),
Map.of("role", "assistant", "content", first.answer())
),
sessionId
);
assertEquals(sessionId, first.sessionId());
assertEquals(sessionId, second.sessionId());
assertNotEquals(first.runId(), second.runId());
DiagnosisRunRepository diagnosisRunRepository =
(DiagnosisRunRepository) ReflectionTestUtils.getField(chatService, "diagnosisRunRepository");
ArgumentCaptor<DiagnosisRun> runCaptor = ArgumentCaptor.forClass(DiagnosisRun.class);
verify(diagnosisRunRepository, atLeast(2)).save(runCaptor.capture());
List<String> savedRunIds = runCaptor.getAllValues().stream()
.filter(run -> sessionId.equals(run.getSessionId()))
.map(DiagnosisRun::getRunId)
.distinct()
.toList();
assertEquals(2, savedRunIds.size());
assertTrue(savedRunIds.contains(first.runId()));
assertTrue(savedRunIds.contains(second.runId()));
}
@Test
@@ -645,9 +716,12 @@ class ChatServiceSequentialAgentTest {
private ChatService createChatService() {
ChatService chatService = new ChatService();
DiagnosisSessionRepository diagnosisSessionRepository = mock(DiagnosisSessionRepository.class);
when(diagnosisSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty());
when(diagnosisSessionRepository.save(any(DiagnosisSession.class))).thenAnswer(invocation -> invocation.getArgument(0));
ChatSessionRepository chatSessionRepository = mock(ChatSessionRepository.class);
when(chatSessionRepository.findBySessionId(anyString())).thenReturn(Optional.empty());
when(chatSessionRepository.save(any(ChatSession.class))).thenAnswer(invocation -> invocation.getArgument(0));
DiagnosisRunRepository diagnosisRunRepository = mock(DiagnosisRunRepository.class);
when(diagnosisRunRepository.save(any(DiagnosisRun.class))).thenAnswer(invocation -> invocation.getArgument(0));
AtomicInteger stepId = new AtomicInteger(1);
AgentStepRepository agentStepRepository = mock(AgentStepRepository.class);
@@ -660,13 +734,20 @@ class ChatServiceSequentialAgentTest {
});
when(agentStepRepository.findById(any())).thenReturn(Optional.of(new AgentStep()));
when(agentStepRepository.findBySessionIdOrderByStepIndex(anyString())).thenReturn(List.of());
when(agentStepRepository.findByRunIdOrderByStepIndex(anyString())).thenReturn(List.of());
ToolInvocationRepository toolInvocationRepository = mock(ToolInvocationRepository.class);
when(toolInvocationRepository.countBySessionId(anyString())).thenReturn(0L);
when(toolInvocationRepository.countByRunId(anyString())).thenReturn(0L);
when(toolInvocationRepository.findBySessionIdOrderByIdAsc(anyString())).thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
when(toolInvocationRepository.findByRunIdOrderByIdAsc(anyString())).thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
EvaluationService evaluationService = mock(EvaluationService.class);
RetrievedDocTracker retrievedDocTracker = mock(RetrievedDocTracker.class);
@@ -674,6 +755,7 @@ class ChatServiceSequentialAgentTest {
when(knowledgeDomainService.buildKnowledgeMap()).thenReturn("");
ToolTraceSummaryService toolTraceSummaryService = mock(ToolTraceSummaryService.class);
when(toolTraceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of());
when(toolTraceSummaryService.buildVerifierTraceSummaryForRun(anyString(), anyString())).thenReturn(List.of());
SelfEvaluationMergeService selfEvaluationMergeService = mock(SelfEvaluationMergeService.class);
when(selfEvaluationMergeService.mergeVerifierEvaluation(any(), any())).thenReturn("{}");
ExecutorGatekeeperService executorGatekeeperService = new ExecutorGatekeeperService(toolInvocationRepository);
@@ -681,7 +763,8 @@ class ChatServiceSequentialAgentTest {
ReflectionTestUtils.setField(chatService, "dateTimeTools", new DateTimeTools());
ReflectionTestUtils.setField(chatService, "lookupKnowledgeTool", new LookupKnowledgeTool());
ReflectionTestUtils.setField(chatService, "queryLogsTools", new QueryLogsTools(mock(ToolInvocationRecorder.class)));
ReflectionTestUtils.setField(chatService, "diagnosisSessionRepository", diagnosisSessionRepository);
ReflectionTestUtils.setField(chatService, "chatSessionRepository", chatSessionRepository);
ReflectionTestUtils.setField(chatService, "diagnosisRunRepository", diagnosisRunRepository);
ReflectionTestUtils.setField(chatService, "agentStepRepository", agentStepRepository);
ReflectionTestUtils.setField(chatService, "toolInvocationRepository", toolInvocationRepository);
ReflectionTestUtils.setField(chatService, "evaluationService", evaluationService);
@@ -11,10 +11,29 @@ import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
class ExecutorGatekeeperServiceTest {
@Test
void validateRunUsesRunScopedToolRows() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findByRunIdOrderByIdAsc("run-gatekeeper-1")).thenReturn(List.of(
invocation(101L, "query_metrics", "$.alerts[0]",
"HighCPUUsage firing, service=payment-service, current=92%")
));
ExecutorGatekeeperService service = new ExecutorGatekeeperService(repository);
Map<String, Object> result = service.validateRun("run-gatekeeper-1",
validOutput(101L, "query_metrics", "$.alerts[0]",
"HighCPUUsage firing, service=payment-service, current=92%"),
Map.of("status", "valid"));
assertEquals("pass", result.get("status"));
verify(repository).findByRunIdOrderByIdAsc("run-gatekeeper-1");
}
@Test
void ruleCatalogLoadsDefaultMetadata() {
GatekeeperRuleCatalog catalog = GatekeeperRuleCatalog.loadDefault(new com.fasterxml.jackson.databind.ObjectMapper());
@@ -26,6 +26,36 @@ class ToolInvocationRecorderTest {
private final ObjectMapper objectMapper = new ObjectMapper();
@Test
void recordEvidenceToolWritesRunIdFromExecutionContext() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.save(any(ToolInvocation.class))).thenAnswer(invocation -> invocation.getArgument(0));
ToolInvocationRecorder recorder = new ToolInvocationRecorder(repository, new ObjectMapper());
SessionContextHolder.setContext("recorder-run-session", "run-recorder-1");
try {
recorder.recordEvidenceTool(
"query_metrics",
Map.of("query", "active_prometheus_alerts"),
"{\"success\":true,\"alerts\":[]}",
true,
12,
null,
"prometheus_alerts",
ToolInvocationRecorder.EVIDENCE_STATUS_NO_EVIDENCE,
Map.of("metric_family", "prometheus_alerts")
);
} finally {
SessionContextHolder.clear();
}
ArgumentCaptor<ToolInvocation> captor = ArgumentCaptor.forClass(ToolInvocation.class);
verify(repository).save(captor.capture());
ToolInvocation saved = captor.getValue();
assertEquals("recorder-run-session", saved.getSessionId());
assertEquals("run-recorder-1", saved.getRunId());
}
@Test
void recordEvidenceToolPreservesNoEvidenceSemantics() throws Exception {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
@@ -15,6 +15,32 @@ import static org.mockito.Mockito.when;
class ToolTraceSummaryServiceTest {
@Test
void buildVerifierTraceSummaryForRunUsesRunScopedToolRows() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findByRunIdOrderByIdAsc("run-summary-1")).thenReturn(List.of(
ToolInvocation.builder()
.id(101L)
.sessionId("session-1")
.runId("run-summary-1")
.toolName("query_metrics")
.inputParams("{\"query\":\"active_prometheus_alerts\"}")
.outputPreview("active=50 max=50")
.retrievalDetails("{\"retrieved_domains\":[\"prometheus_alerts\"],\"evidence_status\":\"supported\"}")
.success(true)
.build()
));
ToolTraceSummaryService service = new ToolTraceSummaryService(repository);
List<Map<String, Object>> summaries = service.buildVerifierTraceSummaryForRun(
"run-summary-1", "active=50 max=50");
assertEquals(1, summaries.size());
assertEquals("query_metrics", summaries.get(0).get("tool_name"));
assertEquals(List.of(101L), summaries.get(0).get("source_invocation_ids"));
}
@Test
void buildVerifierTraceSummaryTreatsNoEvidenceAsGapWithoutLosingSuccessfulEvidence() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);