feat(trace): isolate chat runs
This commit is contained in:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user