feat(trace): isolate chat runs
This commit is contained in:
@@ -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