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
@@ -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);