feat(graph): cut over chat diagnosis stategraph

This commit is contained in:
zhuyongxin
2026-07-17 18:30:08 +08:00
parent 1460dd1e99
commit 99e490f227
36 changed files with 2640 additions and 1623 deletions
@@ -0,0 +1,39 @@
package com.superbiz.agent.config;
import org.junit.jupiter.api.Test;
import java.io.IOException;
import java.io.InputStream;
import java.nio.charset.StandardCharsets;
import java.util.Arrays;
import java.util.stream.Collectors;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNotNull;
class DiagnosisRunSchemaContractTest {
private static final String MIGRATION =
"db/migration/V012__add_diagnosis_run_orchestration_trace.sql";
@Test
void v012AddsOnlyOneNullableJsonColumnToDiagnosisRun() throws IOException {
try (InputStream stream = Thread.currentThread()
.getContextClassLoader()
.getResourceAsStream(MIGRATION)) {
assertNotNull(stream, "missing migration " + MIGRATION);
String sql = new String(stream.readAllBytes(), StandardCharsets.UTF_8);
String executableSql = Arrays.stream(sql.split("\\R"))
.map(String::trim)
.filter(line -> !line.isEmpty() && !line.startsWith("--"))
.collect(Collectors.joining(" "))
.replaceAll("\\s+", " ");
assertEquals(
"ALTER TABLE diagnosis_run ADD COLUMN orchestration_trace JSON NULL "
+ "COMMENT 'Compact run-scoped StateGraph orchestration summary' "
+ "AFTER self_evaluation;",
executableSql);
}
}
}
@@ -0,0 +1,259 @@
package com.superbiz.agent.graph.diagnosis;
import com.alibaba.cloud.ai.graph.CompiledGraph;
import com.alibaba.cloud.ai.graph.RunnableConfig;
import org.junit.jupiter.api.Test;
import reactor.core.publisher.Flux;
import java.util.List;
import java.util.Map;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyMap;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
class ChatDiagnosisGraphRuntimeTest {
@Test
void successfulRunReturnsFinalAnswerTraceAndCurrentThreadId() throws Exception {
ScriptedDiagnosisGraphActions script = passPath();
ChatDiagnosisGraphRuntime.Result result =
new ChatDiagnosisGraphRuntime().execute(
script.actions(),
"分析连接池",
"session-runtime",
"run-runtime");
assertEquals("已确认连接数达到上限。", result.answer());
assertEquals(DiagnosisGraphTopology.Node.COMPOSER,
result.trace().finalNode());
assertEquals("composer_completed", result.trace().terminationReason());
assertEquals(false, result.trace().degraded());
assertEquals(List.of("run-runtime"),
script.threadIds(DiagnosisGraphTopology.Node.PLANNER));
assertEquals(List.of("run-runtime"),
script.threadIds(DiagnosisGraphTopology.Node.EXECUTOR));
assertEquals(List.of("run-runtime"),
script.threadIds(DiagnosisGraphTopology.Node.GATEKEEPER));
assertEquals(List.of("run-runtime"),
script.threadIds(DiagnosisGraphTopology.Node.VERIFIER));
assertEquals(List.of("run-runtime"),
script.threadIds(DiagnosisGraphTopology.Node.COMPOSER));
for (String node : List.of(
DiagnosisGraphTopology.Node.PLANNER,
DiagnosisGraphTopology.Node.EXECUTOR,
DiagnosisGraphTopology.Node.GATEKEEPER,
DiagnosisGraphTopology.Node.VERIFIED_INPUT,
DiagnosisGraphTopology.Node.VERIFIER,
DiagnosisGraphTopology.Node.COMPOSER)) {
assertEquals(List.of("session-runtime"),
script.metadata(node, "sessionId"));
assertEquals(List.of("run-runtime"),
script.metadata(node, "runId"));
}
Map<String, Object> initial = script.states(
DiagnosisGraphTopology.Node.PLANNER).get(0);
assertEquals(Map.of(
"query", "分析连接池",
"original_query", "分析连接池"),
initial.get(DiagnosisGraphState.DIAGNOSIS_CONTEXT));
assertEquals("NORMAL", initial.get(DiagnosisGraphState.PLANNER_MODE));
assertEquals(0, initial.get(DiagnosisGraphState.PLANNER_RETRY_COUNT));
assertEquals(0, initial.get(DiagnosisGraphState.VERIFIER_RETRY_COUNT));
assertEquals(0, initial.get(DiagnosisGraphState.COMPOSER_RETRY_COUNT));
assertEquals(0, initial.get(DiagnosisGraphState.EVIDENCE_RETRY_COUNT));
assertEquals(List.of(), initial.get(DiagnosisGraphState.ORCHESTRATION_EVENTS));
}
@Test
void executionFailureExposesOnlyRealPartialTrace() {
ScriptedDiagnosisGraphActions afterPlannerFailure =
new ScriptedDiagnosisGraphActions()
.step(DiagnosisGraphTopology.Node.PLANNER, "COMPLETED",
"planner_completed", Map.of(
DiagnosisGraphState.PLANNER_STATUS, "COMPLETED",
DiagnosisGraphState.PLANNER_PLAN,
Map.of("plan", List.of("查询指标"))));
ChatDiagnosisGraphRuntime.ExecutionFailure partial = assertThrows(
ChatDiagnosisGraphRuntime.ExecutionFailure.class,
() -> new ChatDiagnosisGraphRuntime().execute(
afterPlannerFailure.actions(), "问题", "session-partial", "run-partial"));
assertEquals(DiagnosisGraphTopology.Node.PLANNER,
partial.partialTrace().finalNode());
assertEquals("planner_completed",
partial.partialTrace().terminationReason());
ChatDiagnosisGraphRuntime.ExecutionFailure beforeFirstEvent = assertThrows(
ChatDiagnosisGraphRuntime.ExecutionFailure.class,
() -> new ChatDiagnosisGraphRuntime().execute(
new ScriptedDiagnosisGraphActions().actions(),
"问题", "session-empty", "run-empty"));
assertNull(beforeFirstEvent.partialTrace());
}
@Test
void handledPreVerificationFallbackReturnsSafeDegradedResult() throws Exception {
ScriptedDiagnosisGraphActions script = new ScriptedDiagnosisGraphActions()
.step(DiagnosisGraphTopology.Node.PLANNER, "NON_RETRYABLE_FAILED",
"planner_failed", Map.of(
DiagnosisGraphState.PLANNER_STATUS,
"NON_RETRYABLE_FAILED"))
.step(DiagnosisGraphTopology.Node.FALLBACK, "COMPLETED",
"fallback_completed", Map.of(
DiagnosisGraphState.FINAL_ANSWER,
"当前无法形成可信诊断。"));
ChatDiagnosisGraphRuntime.Result result =
new ChatDiagnosisGraphRuntime().execute(
script.actions(), "问题", "session-pre", "run-pre");
assertEquals("当前无法形成可信诊断。", result.answer());
assertEquals(DiagnosisGraphTopology.Node.FALLBACK,
result.trace().finalNode());
assertEquals(true, result.trace().degraded());
assertEquals(0, script.calls(DiagnosisGraphTopology.Node.VERIFIER));
}
@Test
void handledPostVerificationFallbackReturnsSafeDegradedResult() throws Exception {
ScriptedDiagnosisGraphActions script = new ScriptedDiagnosisGraphActions()
.step(DiagnosisGraphTopology.Node.PLANNER, "COMPLETED",
"planner_completed", Map.of(
DiagnosisGraphState.PLANNER_STATUS, "COMPLETED"))
.step(DiagnosisGraphTopology.Node.EXECUTOR, "COMPLETED",
"executor_completed", Map.of(
DiagnosisGraphState.EXECUTOR_STATUS, "COMPLETED"))
.step(DiagnosisGraphTopology.Node.GATEKEEPER, "PASS",
"gatekeeper_pass", Map.of(
DiagnosisGraphState.GATEKEEPER_STATUS, "PASS",
DiagnosisGraphState.VERIFIED_BINDING_COUNT, 1))
.step(DiagnosisGraphTopology.Node.VERIFIED_INPUT, "COMPLETED",
"verified_input_built", Map.of())
.step(DiagnosisGraphTopology.Node.VERIFIER, "INVALID_OUTPUT",
"verifier_invalid_output", Map.of(
DiagnosisGraphState.VERIFIER_STATUS,
"INVALID_OUTPUT"))
.step(DiagnosisGraphTopology.Node.VERIFIER, "INVALID_OUTPUT",
"verifier_invalid_output", Map.of(
DiagnosisGraphState.VERIFIER_STATUS,
"INVALID_OUTPUT"))
.step(DiagnosisGraphTopology.Node.FALLBACK, "COMPLETED",
"fallback_completed", Map.of(
DiagnosisGraphState.FINAL_ANSWER,
"已降级为安全答复。"));
ChatDiagnosisGraphRuntime.Result result =
new ChatDiagnosisGraphRuntime().execute(
script.actions(), "问题", "session-post", "run-post");
assertEquals("已降级为安全答复。", result.answer());
assertEquals(DiagnosisGraphTopology.Node.FALLBACK,
result.trace().finalNode());
assertEquals(true, result.trace().degraded());
assertEquals(2, script.calls(DiagnosisGraphTopology.Node.VERIFIER));
}
@Test
void blankAnswerFailsAndKeepsOnlyTheRealFallbackTrace() {
ScriptedDiagnosisGraphActions script = new ScriptedDiagnosisGraphActions()
.step(DiagnosisGraphTopology.Node.PLANNER, "NON_RETRYABLE_FAILED",
"planner_failed", Map.of(
DiagnosisGraphState.PLANNER_STATUS,
"NON_RETRYABLE_FAILED"))
.step(DiagnosisGraphTopology.Node.FALLBACK, "COMPLETED",
"fallback_completed", Map.of(
DiagnosisGraphState.FINAL_ANSWER, " "));
ChatDiagnosisGraphRuntime.ExecutionFailure failure = assertThrows(
ChatDiagnosisGraphRuntime.ExecutionFailure.class,
() -> new ChatDiagnosisGraphRuntime().execute(
script.actions(), "问题", "session-blank", "run-blank"));
assertEquals("diagnosis graph returned no final answer",
failure.getCause().getMessage());
assertEquals(DiagnosisGraphTopology.Node.FALLBACK,
failure.partialTrace().finalNode());
assertEquals(2, failure.partialState().data()
.get(DiagnosisGraphState.ORCHESTRATION_EVENTS) instanceof List<?> events
? events.size() : 0);
}
@Test
void emptyGraphStreamFailsWithoutFabricatingStateOrTrace() throws Exception {
DiagnosisGraphFactory graphFactory = mock(DiagnosisGraphFactory.class);
CompiledGraph graph = mock(CompiledGraph.class);
when(graphFactory.compile(any(DiagnosisGraphActions.class)))
.thenReturn(graph);
when(graph.stream(anyMap(), any(RunnableConfig.class)))
.thenReturn(Flux.empty());
ChatDiagnosisGraphRuntime runtime = new ChatDiagnosisGraphRuntime(
graphFactory, new DiagnosisOrchestrationTraceBuilder());
ChatDiagnosisGraphRuntime.ExecutionFailure failure = assertThrows(
ChatDiagnosisGraphRuntime.ExecutionFailure.class,
() -> runtime.execute(
new ScriptedDiagnosisGraphActions().actions(),
"问题", "session-empty-stream", "run-empty-stream"));
assertEquals("diagnosis graph returned no state",
failure.getCause().getMessage());
assertNull(failure.partialState());
assertNull(failure.partialTrace());
}
private ScriptedDiagnosisGraphActions passPath() {
return new ScriptedDiagnosisGraphActions()
.step(DiagnosisGraphTopology.Node.PLANNER, "COMPLETED",
"planner_completed", Map.of(
DiagnosisGraphState.PLANNER_STATUS, "COMPLETED",
DiagnosisGraphState.PLANNER_PLAN,
Map.of("plan", List.of("查询指标"))))
.step(DiagnosisGraphTopology.Node.EXECUTOR, "COMPLETED",
"executor_completed", Map.of(
DiagnosisGraphState.EXECUTOR_STATUS, "COMPLETED",
DiagnosisGraphState.EXECUTOR_OUTPUT,
Map.of("answer_version", "executor_evidence_v2",
"claims", List.of())))
.step(DiagnosisGraphTopology.Node.GATEKEEPER, "PASS",
"gatekeeper_pass", Map.of(
DiagnosisGraphState.GATEKEEPER_STATUS, "PASS",
DiagnosisGraphState.GATEKEEPER_RESULT,
Map.of("status", "pass", "severity", "none"),
DiagnosisGraphState.VERIFIED_BINDING_COUNT, 1,
DiagnosisGraphState.VERIFIER_VERDICT_CEILING, "PASS"))
.step(DiagnosisGraphTopology.Node.VERIFIED_INPUT, "COMPLETED",
"verified_input_built", Map.of(
DiagnosisGraphState.VERIFIED_EXECUTOR_OUTPUT,
Map.of("answer_version", "executor_evidence_v2",
"claims", List.of()),
DiagnosisGraphState.VERIFIED_EVIDENCE, List.of()))
.step(DiagnosisGraphTopology.Node.VERIFIER, "COMPLETED",
"verifier_completed", Map.of(
DiagnosisGraphState.VERIFIER_STATUS, "COMPLETED",
DiagnosisGraphState.VERIFIER_MODEL_VERDICT, "PASS",
DiagnosisGraphState.EFFECTIVE_VERDICT, "PASS",
DiagnosisGraphState.VERIFIER_OUTPUT, Map.of(
"verdict", "PASS",
"groundedness_score", 1.0,
"critical_fact_count", 0,
"claim_checks", List.of(),
"facts_checked", List.of(),
"rationale", "证据充分")))
.step(DiagnosisGraphTopology.Node.COMPOSER, "COMPLETED",
"composer_completed", Map.of(
DiagnosisGraphState.COMPOSER_STATUS, "COMPLETED",
DiagnosisGraphState.COMPOSER_OUTPUT,
Map.of("status", "valid"),
DiagnosisGraphState.FINAL_ANSWER,
"已确认连接数达到上限。"));
}
}
@@ -0,0 +1,39 @@
package com.superbiz.agent.graph.diagnosis;
import org.junit.jupiter.api.Test;
import org.springframework.core.io.ClassPathResource;
import java.nio.charset.StandardCharsets;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
class ChatVerifierPromptContractTest {
@Test
void promptDeclaresOnlyVerifiedGraphInputContract() throws Exception {
String prompt = new String(
new ClassPathResource("prompts/chat-verifier-prompt.md")
.getInputStream().readAllBytes(),
StandardCharsets.UTF_8);
for (String allowed : List.of(
"diagnosis_context",
"verified_executor_output",
"verified_evidence",
"gatekeeper_audit",
"verdict_ceiling",
"retry_context")) {
assertTrue(prompt.contains("`" + allowed + "`"), allowed);
}
for (String forbidden : List.of(
"executor_final_answer",
"executor_structured_output",
"executor_output_parse_status",
"tool_trace_summary",
"VerifierInputHook")) {
assertFalse(prompt.contains(forbidden), forbidden);
}
}
}
@@ -0,0 +1,91 @@
package com.superbiz.agent.graph.diagnosis;
import com.alibaba.cloud.ai.graph.OverAllState;
import org.junit.jupiter.api.Test;
import java.util.List;
import java.util.Map;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
class DiagnosisGraphResultMapperTest {
private final DiagnosisGraphResultMapper mapper =
new DiagnosisGraphResultMapper();
@Test
void completedVerifierMapsCompatibilityFieldsFromVerifiedStateOnly() {
Map<String, Object> verifiedOutput = Map.of(
"answer_version", "executor_evidence_v2",
"claims", List.of(Map.of("claim_id", "c1")));
List<Map<String, Object>> verifiedEvidence = List.of(Map.of(
"claim_id", "c1",
"source_invocation_id", 17,
"tool_name", "query_metrics",
"raw_path", "$.alerts[0]",
"matched_text", "active=50 max=50"));
Map<String, Object> gatekeeper = Map.of(
"status", "fail",
"severity", "low_confid",
"checked_bindings", verifiedEvidence);
Map<String, Object> verifierOutput = Map.of(
"verdict", "PASS",
"groundedness_score", 1.0,
"critical_fact_count", 1,
"claim_checks", List.of(Map.of(
"claim_id", "c1",
"verification", "direct_observation")),
"facts_checked", List.of(),
"rationale", "模型认为证据充分");
OverAllState state = new OverAllState(Map.ofEntries(
Map.entry(DiagnosisGraphState.PLANNER_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.EXECUTOR_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.GATEKEEPER_STATUS, "LOW_CONFID"),
Map.entry(DiagnosisGraphState.GATEKEEPER_RESULT, gatekeeper),
Map.entry(DiagnosisGraphState.VERIFIED_EXECUTOR_OUTPUT, verifiedOutput),
Map.entry(DiagnosisGraphState.VERIFIED_EVIDENCE, verifiedEvidence),
Map.entry(DiagnosisGraphState.VERIFIER_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.VERIFIER_MODEL_VERDICT, "PASS"),
Map.entry(DiagnosisGraphState.EFFECTIVE_VERDICT, "LOW_CONFID"),
Map.entry(DiagnosisGraphState.VERIFIER_OUTPUT, verifierOutput),
Map.entry(DiagnosisGraphState.COMPOSER_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.COMPOSER_OUTPUT, Map.of("status", "valid")),
Map.entry(DiagnosisGraphState.EVIDENCE_RETRY_COUNT, 0),
Map.entry(DiagnosisGraphState.EXECUTOR_OUTPUT, Map.of(
"raw_secret", "must-not-persist"))));
Map<String, Object> evaluation = mapper.verifierEvaluation(
state, Map.of("version", "chat-prompts-v2"));
assertEquals("COMPLETED", evaluation.get("verifier_status"));
assertEquals("PASS", evaluation.get("model_verdict"));
assertEquals("LOW_CONFID", evaluation.get("effective_verdict"));
assertEquals("LOW_CONFID", evaluation.get("verdict"));
assertEquals(verifiedOutput, evaluation.get("executor_structured_output"));
assertEquals(verifiedEvidence, evaluation.get("verified_evidence"));
assertEquals(gatekeeper, evaluation.get("gatekeeper_result"));
assertEquals(Map.of("status", "valid"), evaluation.get("composer_output"));
assertEquals(Map.of("version", "chat-prompts-v2"), evaluation.get("prompt_audit"));
assertEquals(Map.of("status", "valid", "detail", "graph executor completed"),
evaluation.get("executor_output_parse_status"));
assertFalse(evaluation.containsKey("tool_trace_summary"));
assertFalse(String.valueOf(evaluation).contains("raw_secret"));
}
@Test
void preVerificationFallbackDoesNotFabricateVerdict() {
OverAllState state = new OverAllState(Map.of(
DiagnosisGraphState.PLANNER_STATUS, "INVALID_OUTPUT",
DiagnosisGraphState.FAILURE_REASON, "planner_invalid_output"));
Map<String, Object> evaluation = mapper.verifierEvaluation(
state, Map.of("version", "chat-prompts-v2"));
assertEquals("INVALID_OUTPUT", evaluation.get("planner_status"));
assertEquals("planner_invalid_output", evaluation.get("failure_reason"));
assertFalse(evaluation.containsKey("verdict"));
assertFalse(evaluation.containsKey("model_verdict"));
assertFalse(evaluation.containsKey("effective_verdict"));
}
}
@@ -70,6 +70,16 @@ final class ScriptedDiagnosisGraphActions {
return List.copyOf(node(node).threadIds);
}
List<Object> metadata(String node, String key) {
return node(node).metadata.stream()
.map(values -> values.get(key))
.toList();
}
List<Map<String, Object>> states(String node) {
return List.copyOf(node(node).states);
}
private ScriptedNode node(String name) {
ScriptedNode node = nodes.get(name);
if (node == null) {
@@ -83,6 +93,8 @@ final class ScriptedDiagnosisGraphActions {
private final String name;
private final Deque<Step> steps = new ArrayDeque<>();
private final List<String> threadIds = new ArrayList<>();
private final List<Map<String, Object>> metadata = new ArrayList<>();
private final List<Map<String, Object>> states = new ArrayList<>();
private int calls;
private ScriptedNode(String name) {
@@ -104,6 +116,8 @@ final class ScriptedDiagnosisGraphActions {
calls++;
sequence.add(name);
threadIds.add(config.threadId().orElse(null));
metadata.add(new LinkedHashMap<>(config.metadata().orElse(Map.of())));
states.add(new LinkedHashMap<>(state.data()));
Map<String, Object> update = new LinkedHashMap<>(step.update());
update.put(DiagnosisGraphState.ORCHESTRATION_EVENTS,
@@ -0,0 +1,301 @@
package com.superbiz.agent.service;
import com.alibaba.cloud.ai.graph.OverAllState;
import com.alibaba.cloud.ai.graph.action.AsyncNodeActionWithConfig;
import com.superbiz.agent.agent.tool.DateTimeTools;
import com.superbiz.agent.agent.tool.QueryLogsTools;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
import com.superbiz.agent.graph.diagnosis.ChatDiagnosisGraphRuntime;
import com.superbiz.agent.graph.diagnosis.DiagnosisGraphState;
import com.superbiz.agent.graph.diagnosis.DiagnosisGraphActions;
import com.superbiz.agent.graph.diagnosis.DiagnosisGraphTopology;
import com.superbiz.agent.graph.diagnosis.DiagnosisOrchestrationTrace;
import com.superbiz.agent.graph.diagnosis.OrchestrationEvent;
import com.superbiz.agent.repository.AgentStepRepository;
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;
import org.junit.jupiter.api.Test;
import org.mockito.ArgumentCaptor;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.test.util.ReflectionTestUtils;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CompletableFuture;
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.assertNotNull;
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.Mockito.atLeast;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.when;
class ChatServiceGraphIntegrationTest {
@Test
void graphSuccessPreservesChatResultRunTraceMetricsAndEvaluation() throws Exception {
Fixture fixture = fixture();
String sessionId = "graph-success-session";
String question = "请分析订单支付超时";
when(fixture.runtime.execute(any(), eq(question), eq(sessionId), anyString()))
.thenAnswer(invocation -> successResult(invocation.getArgument(3), false));
ChatService.ChatResult result = fixture.service.executeChatComplex(
fixture.chatModel, new ToolCallback[0], question, List.of(), sessionId);
assertEquals(sessionId, result.sessionId());
assertTrue(result.runId().startsWith("run-"));
assertEquals("诊断完成。", result.answer());
DiagnosisRun saved = lastSavedRun(fixture.diagnosisRunRepository);
assertEquals(result.runId(), saved.getRunId());
assertEquals(sessionId, saved.getSessionId());
assertEquals("CHAT", saved.getAgentFlow());
assertEquals("SUCCESS", saved.getStatus());
assertEquals("诊断完成。", saved.getAnswer());
assertEquals(12, saved.getTotalTokenCount());
assertEquals(1, saved.getStepCount());
assertEquals(2, saved.getToolCallCount());
assertNotNull(saved.getOrchestrationTrace());
assertTrue(saved.getOrchestrationTrace().contains("\"final_node\":\"composer\""));
assertTrue(saved.getSelfEvaluation().contains("\"verdict\":\"PASS\""));
assertTrue(saved.getSelfEvaluation().contains("\"prompt_audit\""));
assertFalse(saved.getSelfEvaluation().contains("tool_trace_summary"));
verify(fixture.evaluationService).evaluateRun(result.runId(), result.answer());
verify(fixture.retrievedDocTracker).clearSession(sessionId);
}
@Test
void handledFallbackIsSuccessfulAndUsesDegradedTraceWithoutFabricatedVerdict() throws Exception {
Fixture fixture = fixture();
String sessionId = "graph-fallback-session";
when(fixture.runtime.execute(any(), anyString(), eq(sessionId), anyString()))
.thenAnswer(invocation -> successResult(invocation.getArgument(3), true));
ChatService.ChatResult result = fixture.service.executeChatComplex(
fixture.chatModel, new ToolCallback[0], "无法建立可信材料", List.of(), sessionId);
DiagnosisRun saved = lastSavedRun(fixture.diagnosisRunRepository);
assertEquals("SUCCESS", saved.getStatus());
assertEquals("当前无法给出可信结论。", result.answer());
assertTrue(saved.getOrchestrationTrace().contains("\"degraded\":true"));
assertTrue(saved.getSelfEvaluation().contains("planner_status"));
assertFalse(saved.getSelfEvaluation().contains("\"verdict\""));
verify(fixture.evaluationService).evaluateRun(result.runId(), result.answer());
}
@Test
void unhandledGraphFailureMarksRunFailedAndDoesNotEvaluate() throws Exception {
Fixture fixture = fixture();
String sessionId = "graph-failed-session";
when(fixture.runtime.execute(any(), anyString(), eq(sessionId), anyString()))
.thenThrow(new IllegalStateException("graph exploded"));
ChatService.ChatResult result = fixture.service.executeChatComplex(
fixture.chatModel, new ToolCallback[0], "触发失败", List.of(), sessionId);
DiagnosisRun saved = lastSavedRun(fixture.diagnosisRunRepository);
assertEquals("FAILED", saved.getStatus());
assertTrue(result.answer().contains("graph exploded"));
assertEquals(result.answer(), saved.getAnswer());
verify(fixture.evaluationService, never()).evaluateRun(anyString(), anyString());
verify(fixture.retrievedDocTracker).clearSession(sessionId);
}
@Test
void blankGraphAnswerMarksRunFailedAndPersistsOnlyRealPartialTrace()
throws Exception {
Fixture fixture = fixture();
String sessionId = "graph-blank-answer-session";
when(fixture.runtime.execute(any(), anyString(), eq(sessionId), anyString()))
.thenAnswer(invocation -> new ChatDiagnosisGraphRuntime().execute(
blankAnswerActions(),
invocation.getArgument(1),
invocation.getArgument(2),
invocation.getArgument(3)));
ChatService.ChatResult result = fixture.service.executeChatComplex(
fixture.chatModel, new ToolCallback[0], "空答案", List.of(), sessionId);
DiagnosisRun saved = lastSavedRun(fixture.diagnosisRunRepository);
assertEquals("FAILED", saved.getStatus());
assertTrue(result.answer().contains("diagnosis graph execution failed"));
assertNotNull(saved.getOrchestrationTrace());
assertTrue(saved.getOrchestrationTrace().contains(
"\"final_node\":\"fallback\""));
assertFalse(saved.getOrchestrationTrace().contains("verifier"));
verify(fixture.evaluationService, never())
.evaluateRun(anyString(), anyString());
}
@Test
void sameSessionCreatesDistinctRunIdsAndKeepsTracePerRun() throws Exception {
Fixture fixture = fixture();
String sessionId = "graph-multi-run-session";
when(fixture.runtime.execute(any(), anyString(), eq(sessionId), anyString()))
.thenAnswer(invocation -> successResult(invocation.getArgument(3), false));
ChatService.ChatResult first = fixture.service.executeChatComplex(
fixture.chatModel, new ToolCallback[0], "第一轮", List.of(), sessionId);
ChatService.ChatResult second = fixture.service.executeChatComplex(
fixture.chatModel, new ToolCallback[0], "第二轮", List.of(), sessionId);
assertNotEquals(first.runId(), second.runId());
ArgumentCaptor<DiagnosisRun> captor = ArgumentCaptor.forClass(DiagnosisRun.class);
verify(fixture.diagnosisRunRepository, atLeast(4)).save(captor.capture());
List<DiagnosisRun> completed = captor.getAllValues().stream()
.filter(run -> "SUCCESS".equals(run.getStatus()))
.toList();
assertTrue(completed.stream().anyMatch(run -> first.runId().equals(run.getRunId())));
assertTrue(completed.stream().anyMatch(run -> second.runId().equals(run.getRunId())));
assertTrue(completed.stream().allMatch(run ->
run.getOrchestrationTrace().contains(run.getRunId())));
}
private ChatDiagnosisGraphRuntime.Result successResult(String runId, boolean degraded) {
OverAllState state = degraded
? new OverAllState(Map.of(
DiagnosisGraphState.PLANNER_STATUS, "INVALID_OUTPUT",
DiagnosisGraphState.FAILURE_REASON, "planner_invalid_output"))
: new OverAllState(Map.ofEntries(
Map.entry(DiagnosisGraphState.PLANNER_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.EXECUTOR_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.GATEKEEPER_STATUS, "PASS"),
Map.entry(DiagnosisGraphState.VERIFIER_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.VERIFIER_MODEL_VERDICT, "PASS"),
Map.entry(DiagnosisGraphState.EFFECTIVE_VERDICT, "PASS"),
Map.entry(DiagnosisGraphState.VERIFIER_OUTPUT, Map.of(
"verdict", "PASS",
"groundedness_score", 1.0,
"critical_fact_count", 0,
"claim_checks", List.of(),
"facts_checked", List.of(),
"rationale", "证据充分")),
Map.entry(DiagnosisGraphState.COMPOSER_STATUS, "COMPLETED"),
Map.entry(DiagnosisGraphState.COMPOSER_OUTPUT, Map.of("status", "valid"))));
DiagnosisOrchestrationTrace trace = new DiagnosisOrchestrationTrace(
"stategraph-v1",
List.of(),
degraded ? DiagnosisGraphTopology.Node.FALLBACK
: DiagnosisGraphTopology.Node.COMPOSER,
degraded ? "fallback_completed" : "composer_completed-" + runId,
degraded,
0);
return new ChatDiagnosisGraphRuntime.Result(
state,
degraded ? "当前无法给出可信结论。" : "诊断完成。",
trace);
}
private DiagnosisGraphActions blankAnswerActions() {
AsyncNodeActionWithConfig unexpected = (state, config) ->
CompletableFuture.failedFuture(new AssertionError(
"unexpected graph node"));
AsyncNodeActionWithConfig planner = (state, config) ->
CompletableFuture.completedFuture(Map.of(
DiagnosisGraphState.PLANNER_STATUS,
"NON_RETRYABLE_FAILED",
DiagnosisGraphState.ORCHESTRATION_EVENTS,
List.of(new OrchestrationEvent(
DiagnosisGraphTopology.Node.PLANNER,
"NON_RETRYABLE_FAILED",
"planner_failed",
1))));
AsyncNodeActionWithConfig fallback = (state, config) ->
CompletableFuture.completedFuture(Map.of(
DiagnosisGraphState.FINAL_ANSWER, " ",
DiagnosisGraphState.ORCHESTRATION_EVENTS,
List.of(new OrchestrationEvent(
DiagnosisGraphTopology.Node.FALLBACK,
"COMPLETED",
"fallback_completed",
1))));
return new DiagnosisGraphActions(
planner,
unexpected,
unexpected,
unexpected,
unexpected,
unexpected,
unexpected,
fallback);
}
private DiagnosisRun lastSavedRun(DiagnosisRunRepository repository) {
ArgumentCaptor<DiagnosisRun> captor = ArgumentCaptor.forClass(DiagnosisRun.class);
verify(repository, atLeast(2)).save(captor.capture());
return captor.getAllValues().get(captor.getAllValues().size() - 1);
}
private Fixture fixture() throws Exception {
ChatService service = new ChatService();
ChatDiagnosisGraphRuntime runtime = mock(ChatDiagnosisGraphRuntime.class);
ChatSessionRepository chatSessionRepository = mock(ChatSessionRepository.class);
when(chatSessionRepository.findBySessionId(anyString())).thenReturn(java.util.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));
AgentStepRepository agentStepRepository = mock(AgentStepRepository.class);
when(agentStepRepository.findByRunIdOrderByStepIndex(anyString()))
.thenReturn(List.of(AgentStep.builder().tokenCount(12).build()));
ToolInvocationRepository toolInvocationRepository = mock(ToolInvocationRepository.class);
when(toolInvocationRepository.countByRunId(anyString())).thenReturn(2L);
EvaluationService evaluationService = mock(EvaluationService.class);
RetrievedDocTracker retrievedDocTracker = mock(RetrievedDocTracker.class);
KnowledgeDomainService knowledgeDomainService = mock(KnowledgeDomainService.class);
when(knowledgeDomainService.buildKnowledgeMap()).thenReturn("");
ReflectionTestUtils.setField(service, "dateTimeTools", new DateTimeTools());
ReflectionTestUtils.setField(service, "lookupKnowledgeTool", new LookupKnowledgeTool());
ReflectionTestUtils.setField(service, "queryLogsTools",
new QueryLogsTools(mock(ToolInvocationRecorder.class)));
ReflectionTestUtils.setField(service, "chatSessionRepository", chatSessionRepository);
ReflectionTestUtils.setField(service, "diagnosisRunRepository", diagnosisRunRepository);
ReflectionTestUtils.setField(service, "agentStepRepository", agentStepRepository);
ReflectionTestUtils.setField(service, "toolInvocationRepository", toolInvocationRepository);
ReflectionTestUtils.setField(service, "evaluationService", evaluationService);
ReflectionTestUtils.setField(service, "retrievedDocTracker", retrievedDocTracker);
ReflectionTestUtils.setField(service, "knowledgeDomainService", knowledgeDomainService);
ReflectionTestUtils.setField(service, "selfEvaluationMergeService", new SelfEvaluationMergeService());
ReflectionTestUtils.setField(service, "executorGatekeeperService",
mock(ExecutorGatekeeperService.class));
ReflectionTestUtils.setField(service, "diagnosisGraphRuntime", runtime);
ReflectionTestUtils.setField(service, "chatPlannerPrompt", "PLANNER_TEST_PROMPT");
ReflectionTestUtils.setField(service, "chatExecutorPrompt", "EXECUTOR_TEST_PROMPT");
ReflectionTestUtils.setField(service, "chatVerifierPrompt", "VERIFIER_TEST_PROMPT");
ReflectionTestUtils.setField(service, "chatComposerPrompt", "COMPOSER_TEST_PROMPT");
return new Fixture(
service,
runtime,
mock(ChatModel.class),
diagnosisRunRepository,
evaluationService,
retrievedDocTracker);
}
private record Fixture(
ChatService service,
ChatDiagnosisGraphRuntime runtime,
ChatModel chatModel,
DiagnosisRunRepository diagnosisRunRepository,
EvaluationService evaluationService,
RetrievedDocTracker retrievedDocTracker) {
}
}
@@ -1,918 +0,0 @@
package com.superbiz.agent.service;
import com.alibaba.cloud.ai.graph.agent.ReactAgent;
import com.alibaba.cloud.ai.graph.skills.registry.SkillRegistry;
import com.alibaba.cloud.ai.graph.skills.registry.classpath.ClasspathSkillRegistry;
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.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.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;
import org.junit.jupiter.api.Test;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.model.ChatModel;
import org.springframework.ai.chat.model.ChatResponse;
import org.springframework.ai.chat.model.Generation;
import org.springframework.ai.chat.prompt.Prompt;
import org.springframework.ai.tool.ToolCallback;
import org.springframework.test.util.ReflectionTestUtils;
import org.mockito.ArgumentCaptor;
import java.util.List;
import java.util.Map;
import java.util.Optional;
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;
class ChatServiceSequentialAgentTest {
@Test
void executeChatComplexInvokesSequentialWorkflow() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel();
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析订单支付超时的原因,并给出修复建议",
List.of(),
"sequential-test-session"
);
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
void executeChatComplexDoesNotRetryLowConfidenceByDefault() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel("""
{
"verdict": "LOW_CONFID",
"groundedness_score": 0.1,
"critical_fact_count": 1,
"facts_checked": [
{
"fact": "missing direct evidence",
"is_critical": true,
"verification": "no_evidence",
"detail": "scripted evidence gap",
"evidence_refs": []
}
],
"rationale": "scripted low confidence"
}
""");
chatModel.composerOutput = "not-json";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析订单支付超时的原因,并给出修复建议",
List.of(),
"sequential-low-confidence-session"
);
assertTrue(result.answer().startsWith("以下结论基于当前已获取证据"));
assertFalse(result.answer().contains("EXECUTOR_FINAL_ANSWER"));
assertTrue(result.answer().contains("当前缺口"));
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "chat_composer"), chatModel.agentCalls);
}
@Test
void executeChatComplexLowConfidenceConfirmedFactsOnlyUseDirectEvidence() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel("""
{
"verdict": "LOW_CONFID",
"groundedness_score": 0.37,
"critical_fact_count": 3,
"facts_checked": [
{
"fact": "连接池耗尽 active=50/50",
"is_critical": true,
"verification": "direct_evidence",
"detail": "log evidence",
"evidence_refs": []
},
{
"fact": "临时扩容连接池到 80",
"is_critical": true,
"verification": "indirect_support",
"detail": "suggestion inferred from evidence",
"evidence_refs": []
},
{
"fact": "OOM 导致连接泄漏",
"is_critical": true,
"verification": "no_evidence",
"detail": "missing OOM log",
"evidence_refs": []
}
],
"rationale": "scripted low confidence"
}
""");
chatModel.composerOutput = "not-json";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-low-confid-direct-only-session"
);
assertTrue(result.answer().contains("已确认信息:\n- 连接池耗尽 active=50/50"));
assertTrue(result.answer().contains("80"));
assertTrue(result.answer().contains("suggestion inferred from evidence"));
assertTrue(result.answer().contains("missing OOM log"));
assertFalse(result.answer().contains("EXECUTOR_FINAL_ANSWER"));
}
@Test
void executeChatComplexFallsBackToLowConfidenceWhenVerifierOutputMissing() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel("", "");
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析订单支付超时的原因,并给出修复建议",
List.of(),
"sequential-missing-verifier-session"
);
assertTrue(result.answer().startsWith("以下结论基于当前已获取证据"));
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "chat_verifier"), chatModel.agentCalls);
}
@Test
void executeChatComplexFallsBackToLowConfidenceWhenVerifierJsonInvalid() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel("not-json");
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析订单支付超时的原因,并给出修复建议",
List.of(),
"sequential-invalid-verifier-session"
);
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier"), chatModel.agentCalls);
}
@Test
void executeChatComplexRejectOutputDoesNotLeakExecutorAnswer() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel("""
{
"verdict": "REJECT",
"groundedness_score": 0.0,
"critical_fact_count": 1,
"facts_checked": [
{
"fact": "payment timeout root cause",
"is_critical": true,
"verification": "contradicted",
"detail": "scripted contradiction",
"evidence_refs": []
}
],
"rationale": "scripted reject"
}
""");
chatModel.composerOutput = "not-json";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析订单支付超时的原因,并给出修复建议",
List.of(),
"sequential-reject-session"
);
assertTrue(result.answer().startsWith("当前无法基于已获取证据生成可靠结论"));
assertFalse(result.answer().contains("EXECUTOR_FINAL_ANSWER"));
}
@Test
void executeChatComplexRunsPlannerExecutorVerifierInFixedOrder() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel();
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析订单支付超时的原因,并给出修复建议",
List.of(),
"sequential-workflow-session"
);
assertTrue(result.answer().contains("连接池 active 达到上限"));
assertFalse(result.answer().contains("\"answer_version\""));
assertEquals(List.of("chat_planner", "chat_executor", "chat_verifier", "chat_composer"), chatModel.agentCalls);
assertTrue(chatModel.sawVerifierPrompt);
}
@Test
void verifierReceivesStructuredExecutorPayloadFields() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel();
chatModel.executorOutput = """
{
"answer_version": "executor_evidence_v2",
"claims": [
{
"claim_id": "claim-1",
"claim_type": "symptom",
"claim_text": "连接池 active 达到上限",
"support_level": "direct",
"evidence_bindings": [
{
"source_type": "tool_trace",
"source_id": "trace-1",
"tool_name": "query_metrics",
"source_invocation_id": 101,
"raw_path": "$.alerts[0]",
"evidence_excerpt": "active=50 max=50"
}
]
}
],
"hypotheses": [],
"recommended_actions": [],
"missing_info": []
}
""";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-structured-executor-session"
);
assertTrue(result.answer().contains("连接池 active 达到上限"));
assertFalse(result.answer().contains("\"answer_version\""));
assertTrue(chatModel.verifierPromptText.contains("\"executor_structured_output\""));
assertTrue(chatModel.verifierPromptText.contains("\"executor_output_parse_status\""));
assertTrue(chatModel.verifierPromptText.contains("\"status\" : \"valid\""));
assertTrue(chatModel.verifierPromptText.contains("连接池 active 达到上限"));
}
@Test
void executeChatComplexRendersExecutorEvidenceV2InsteadOfRawJsonOnPass() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel();
chatModel.composerOutput = "not-json";
chatModel.executorOutput = """
{
"answer_version": "executor_evidence_v2",
"claims": [
{
"claim_id": "claim-1",
"claim_type": "symptom",
"claim_text": "连接池 active 达到上限",
"support_level": "direct",
"evidence_bindings": [
{
"source_type": "tool_trace",
"source_id": "trace-1",
"tool_name": "query_metrics",
"source_invocation_id": 101,
"raw_path": "$.alerts[0]",
"evidence_excerpt": "active=50 max=50"
}
]
}
],
"hypotheses": [
{
"hypothesis_text": "连接泄漏可能参与了连接池耗尽",
"basis": "已有连接池满载证据,但缺少泄漏检测日志",
"needed_evidence": ["连接泄漏检测日志"]
}
],
"recommended_actions": [
{
"action_text": "补充查询连接池泄漏检测日志",
"reason": "用于确认是否存在连接未释放"
}
],
"missing_info": ["缺少连接泄漏检测日志"]
}
""";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-v2-render-session"
);
assertTrue(result.answer().contains("已确认信息"));
assertTrue(result.answer().contains("连接池 active 达到上限"));
assertTrue(result.answer().contains("建议下一步"));
assertFalse(result.answer().contains("\"answer_version\""));
assertFalse(result.answer().contains("executor_evidence_v2"));
}
@Test
void executeChatComplexPersistsGatekeeperResultInVerifierEvaluation() throws Exception {
ChatService chatService = createChatService();
SelfEvaluationMergeService mergeService =
(SelfEvaluationMergeService) ReflectionTestUtils.getField(chatService, "selfEvaluationMergeService");
ToolInvocationRepository invocationRepository =
(ToolInvocationRepository) ReflectionTestUtils.getField(chatService, "toolInvocationRepository");
when(invocationRepository.findBySessionIdOrderByIdAsc("sequential-gatekeeper-persist-session"))
.thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.sessionId("sequential-gatekeeper-persist-session")
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
ScriptedChatModel chatModel = new ScriptedChatModel();
chatModel.executorOutput = """
{
"answer_version": "executor_evidence_v2",
"claims": [
{
"claim_id": "claim-1",
"claim_type": "symptom",
"claim_text": "连接池 active 达到上限",
"support_level": "direct",
"evidence_bindings": [
{
"source_type": "tool_trace",
"source_id": "trace-1",
"tool_name": "query_metrics",
"source_invocation_id": 101,
"raw_path": "$.alerts[0]",
"evidence_excerpt": "active=50 max=50"
}
]
}
],
"hypotheses": [],
"recommended_actions": [],
"missing_info": []
}
""";
chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-gatekeeper-persist-session"
);
ArgumentCaptor<Map<String, Object>> captor = ArgumentCaptor.forClass(Map.class);
verify(mergeService).mergeVerifierEvaluation(isNull(), captor.capture());
Map<String, Object> verifierEvaluation = captor.getValue();
assertTrue(verifierEvaluation.containsKey("gatekeeper_result"));
assertTrue(verifierEvaluation.containsKey("prompt_audit"));
@SuppressWarnings("unchecked")
Map<String, Object> gatekeeperResult = (Map<String, Object>) verifierEvaluation.get("gatekeeper_result");
assertEquals("pass", gatekeeperResult.get("status"));
assertEquals("none", gatekeeperResult.get("severity"));
@SuppressWarnings("unchecked")
Map<String, Object> promptAudit = (Map<String, Object>) verifierEvaluation.get("prompt_audit");
assertEquals("chat-prompts-v1", promptAudit.get("version"));
@SuppressWarnings("unchecked")
List<Map<String, Object>> prompts = (List<Map<String, Object>>) promptAudit.get("prompts");
assertEquals(4, prompts.size());
assertTrue(prompts.stream().anyMatch(prompt ->
"chat_executor".equals(prompt.get("name"))
&& "chat-executor-v2".equals(prompt.get("version"))));
}
@Test
void executeChatComplexMapsClaimChecksToFactsCheckedAndPersistsBoth() throws Exception {
ChatService chatService = createChatService();
SelfEvaluationMergeService mergeService =
(SelfEvaluationMergeService) ReflectionTestUtils.getField(chatService, "selfEvaluationMergeService");
ToolInvocationRepository invocationRepository =
(ToolInvocationRepository) ReflectionTestUtils.getField(chatService, "toolInvocationRepository");
when(invocationRepository.findBySessionIdOrderByIdAsc("sequential-claim-check-session"))
.thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.sessionId("sequential-claim-check-session")
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
ScriptedChatModel chatModel = new ScriptedChatModel("""
{
"verdict": "LOW_CONFID",
"groundedness_score": 0.32,
"critical_fact_count": 6,
"claim_checks": [
{"claim_id":"claim-1","claim_text":"CPU 使用率 92%","claim_type":"symptom","verification":"direct_observation","detail":"direct","evidence_refs":[{"trace_ref":"trace-1","tool_name":"query_metrics","source_invocation_ids":[101],"note":"cpu"}]},
{"claim_id":"claim-2","claim_text":"CPU 过高可能导致超时","claim_type":"risk","verification":"reasonable_inference","detail":"inference","evidence_refs":[]},
{"claim_id":"claim-3","claim_text":"CPU 是唯一根因","claim_type":"root_cause","verification":"overstated","detail":"too strong","evidence_refs":[]},
{"claim_id":"claim-4","claim_text":"缺少线程池证据","claim_type":"symptom","verification":"unsupported","detail":"missing","evidence_refs":[]},
{"claim_id":"claim-5","claim_text":"出现证据外错误码 ERR_FAKE","claim_type":"symptom","verification":"external_unknown","detail":"external","evidence_refs":[]},
{"claim_id":"claim-6","claim_text":"证据显示 CPU 很低","claim_type":"symptom","verification":"contradicted","detail":"conflict","evidence_refs":[]}
],
"facts_checked": [],
"rationale": "claim checks drive compatibility"
}
""");
chatModel.executorOutput = validExecutorV2Output();
chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-claim-check-session"
);
ArgumentCaptor<Map<String, Object>> captor = ArgumentCaptor.forClass(Map.class);
verify(mergeService).mergeVerifierEvaluation(isNull(), captor.capture());
Map<String, Object> verifierEvaluation = captor.getValue();
@SuppressWarnings("unchecked")
List<Map<String, Object>> claimChecks = (List<Map<String, Object>>) verifierEvaluation.get("claim_checks");
@SuppressWarnings("unchecked")
List<Map<String, Object>> factsChecked = (List<Map<String, Object>>) verifierEvaluation.get("facts_checked");
assertEquals(6, claimChecks.size());
assertEquals(6, factsChecked.size());
assertEquals("direct_evidence", factsChecked.get(0).get("verification"));
assertEquals("indirect_support", factsChecked.get(1).get("verification"));
assertEquals("indirect_support", factsChecked.get(2).get("verification"));
assertEquals("no_evidence", factsChecked.get(3).get("verification"));
assertEquals("no_evidence", factsChecked.get(4).get("verification"));
assertEquals("contradicted", factsChecked.get(5).get("verification"));
assertTrue(String.valueOf(factsChecked.get(0).get("fact")).startsWith("claim-1:"));
}
@Test
void executeChatComplexDowngradesPassToRejectWhenGatekeeperInvocationRefFails() throws Exception {
ChatService chatService = createChatService();
SelfEvaluationMergeService mergeService =
(SelfEvaluationMergeService) ReflectionTestUtils.getField(chatService, "selfEvaluationMergeService");
ToolInvocationRepository invocationRepository =
(ToolInvocationRepository) ReflectionTestUtils.getField(chatService, "toolInvocationRepository");
when(invocationRepository.findBySessionIdOrderByIdAsc("sequential-gatekeeper-fail-session"))
.thenReturn(List.of(ToolInvocation.builder()
.id(101L)
.sessionId("sequential-gatekeeper-fail-session")
.toolName("query_metrics")
.retrievalDetails(evidenceRefs("$.alerts[0]", "active=50 max=50"))
.build()));
ScriptedChatModel chatModel = new ScriptedChatModel("""
{
"verdict": "PASS",
"groundedness_score": 1.0,
"critical_fact_count": 1,
"claim_checks": [
{"claim_id":"claim-1","claim_text":"连接池 active 达到上限","claim_type":"symptom","verification":"direct_observation","detail":"direct","evidence_refs":[]}
],
"facts_checked": [],
"rationale": "model tried pass"
}
""");
chatModel.composerOutput = "not-json";
chatModel.executorOutput = """
{
"answer_version": "executor_evidence_v2",
"claims": [
{
"claim_id": "claim-1",
"claim_type": "symptom",
"claim_text": "连接池 active 达到上限",
"support_level": "direct",
"evidence_bindings": [
{
"source_type": "tool_trace",
"tool_name": "query_metrics",
"source_invocation_id": 999,
"raw_path": "$.alerts[0]",
"evidence_excerpt": "active=50 max=50"
}
]
}
],
"hypotheses": [],
"recommended_actions": [],
"missing_info": []
}
""";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-gatekeeper-fail-session"
);
assertTrue(result.answer().startsWith("当前无法基于已获取证据生成可靠结论"));
ArgumentCaptor<Map<String, Object>> captor = ArgumentCaptor.forClass(Map.class);
verify(mergeService).mergeVerifierEvaluation(isNull(), captor.capture());
assertEquals("REJECT", captor.getValue().get("verdict"));
}
@Test
void executeChatComplexDowngradesPassToLowConfidenceWhenExecutorOutputMalformed() throws Exception {
ChatService chatService = createChatService();
SelfEvaluationMergeService mergeService =
(SelfEvaluationMergeService) ReflectionTestUtils.getField(chatService, "selfEvaluationMergeService");
ScriptedChatModel chatModel = new ScriptedChatModel("""
{
"verdict": "PASS",
"groundedness_score": 1.0,
"critical_fact_count": 0,
"claim_checks": [],
"facts_checked": [],
"rationale": "model tried pass"
}
""");
chatModel.composerOutput = "not-json";
chatModel.executorOutput = "{ not-json";
ChatService.ChatResult result = chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"请分析 MySQL 连接池耗尽",
List.of(),
"sequential-malformed-pass-session"
);
assertTrue(result.answer().startsWith("以下结论基于当前已获取证据"));
ArgumentCaptor<Map<String, Object>> captor = ArgumentCaptor.forClass(Map.class);
verify(mergeService).mergeVerifierEvaluation(isNull(), captor.capture());
assertEquals("LOW_CONFID", captor.getValue().get("verdict"));
}
@Test
void buildMethodToolsArrayIncludesLogsAndMetricsWhenAvailable() {
ChatService chatService = new ChatService();
DateTimeTools dateTimeTools = new DateTimeTools();
LookupKnowledgeTool lookupKnowledgeTool = new LookupKnowledgeTool();
QueryLogsTools queryLogsTools = new QueryLogsTools(mock(ToolInvocationRecorder.class));
QueryMetricsTools queryMetricsTools = new QueryMetricsTools(mock(ToolInvocationRecorder.class));
ReflectionTestUtils.setField(chatService, "dateTimeTools", dateTimeTools);
ReflectionTestUtils.setField(chatService, "lookupKnowledgeTool", lookupKnowledgeTool);
ReflectionTestUtils.setField(chatService, "queryLogsTools", queryLogsTools);
ReflectionTestUtils.setField(chatService, "queryMetricsTools", queryMetricsTools);
Object[] methodTools = chatService.buildMethodToolsArray();
assertEquals(4, methodTools.length);
assertSame(dateTimeTools, methodTools[0]);
assertSame(lookupKnowledgeTool, methodTools[1]);
assertSame(queryLogsTools, methodTools[2]);
assertSame(queryMetricsTools, methodTools[3]);
}
@Test
void createReactAgentInjectsSkillCatalogThroughAlibabaHook() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel();
SkillRegistry skillRegistry = ClasspathSkillRegistry.builder()
.classpathPath("skills")
.basePath("target/test-skills-cache")
.build();
ReflectionTestUtils.setField(chatService, "skillRegistry", skillRegistry);
ReactAgent agent = chatService.createReactAgent(chatModel, "BASE_TEST_PROMPT");
agent.call("diagnose mysql connection pool exhaustion");
assertTrue(chatModel.promptText.contains("BASE_TEST_PROMPT"));
assertTrue(chatModel.promptText.contains("## Skills System"));
assertTrue(chatModel.promptText.contains("diagnose-mysql-connection-pool"));
assertTrue(chatModel.promptText.contains("read_skill"));
}
@Test
void plannerGetsSkillMetadataAndExecutorGetsReadSkillTool() throws Exception {
ChatService chatService = createChatService();
ScriptedChatModel chatModel = new ScriptedChatModel();
SkillRegistry skillRegistry = ClasspathSkillRegistry.builder()
.classpathPath("skills")
.basePath("target/test-skills-cache")
.build();
ReflectionTestUtils.setField(chatService, "skillRegistry", skillRegistry);
chatService.executeChatComplex(
chatModel,
new ToolCallback[0],
"diagnose mysql connection pool exhaustion",
List.of(),
"planner-skill-metadata-session"
);
assertTrue(chatModel.plannerPromptText.contains("\"skill_catalog\""));
assertTrue(chatModel.plannerPromptText.contains("diagnose-mysql-connection-pool"));
assertTrue(chatModel.plannerPromptText.contains("\"selected_skill\""));
assertFalse(chatModel.plannerPromptText.contains("## Skills System"));
assertFalse(chatModel.plannerPromptText.contains("read_skill"));
assertTrue(chatModel.executorPromptText.contains("## Skills System"));
assertTrue(chatModel.executorPromptText.contains("diagnose-mysql-connection-pool"));
assertTrue(chatModel.executorPromptText.contains("read_skill"));
assertTrue(chatModel.executorPromptText.contains("只允许对该 skill 调用一次 read_skill"));
assertFalse(chatModel.verifierPromptText.contains("diagnose-mysql-connection-pool"));
assertFalse(chatModel.verifierPromptText.contains("read_skill"));
}
private ChatService createChatService() {
ChatService chatService = new ChatService();
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);
when(agentStepRepository.save(any(AgentStep.class))).thenAnswer(invocation -> {
AgentStep step = invocation.getArgument(0);
if (step.getId() == null) {
step.setId((long) stepId.getAndIncrement());
}
return step;
});
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);
KnowledgeDomainService knowledgeDomainService = mock(KnowledgeDomainService.class);
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);
ReflectionTestUtils.setField(chatService, "dateTimeTools", new DateTimeTools());
ReflectionTestUtils.setField(chatService, "lookupKnowledgeTool", new LookupKnowledgeTool());
ReflectionTestUtils.setField(chatService, "queryLogsTools", new QueryLogsTools(mock(ToolInvocationRecorder.class)));
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);
ReflectionTestUtils.setField(chatService, "retrievedDocTracker", retrievedDocTracker);
ReflectionTestUtils.setField(chatService, "knowledgeDomainService", knowledgeDomainService);
ReflectionTestUtils.setField(chatService, "toolTraceSummaryService", toolTraceSummaryService);
ReflectionTestUtils.setField(chatService, "selfEvaluationMergeService", selfEvaluationMergeService);
ReflectionTestUtils.setField(chatService, "executorGatekeeperService", executorGatekeeperService);
ReflectionTestUtils.setField(chatService, "verifierLowConfidenceThreshold", 0.5d);
ReflectionTestUtils.setField(chatService, "chatPlannerPrompt", "PLANNER_TEST_PROMPT");
ReflectionTestUtils.setField(chatService, "chatExecutorPrompt", "EXECUTOR_TEST_PROMPT");
ReflectionTestUtils.setField(chatService, "chatVerifierPrompt", "VERIFIER_TEST_PROMPT");
ReflectionTestUtils.setField(chatService, "chatComposerPrompt", "COMPOSER_TEST_PROMPT");
return chatService;
}
private String validExecutorV2Output() {
return """
{
"answer_version": "executor_evidence_v2",
"claims": [
{
"claim_id": "claim-1",
"claim_type": "symptom",
"claim_text": "连接池 active 达到上限",
"support_level": "direct",
"evidence_bindings": [
{
"source_type": "tool_trace",
"source_id": "trace-1",
"tool_name": "query_metrics",
"source_invocation_id": 101,
"raw_path": "$.alerts[0]",
"evidence_excerpt": "active=50 max=50"
}
]
}
],
"hypotheses": [],
"recommended_actions": [],
"missing_info": []
}
""";
}
private String evidenceRefs(String rawPath, String text) {
return "{\"evidence_refs\":[{\"raw_path\":\"" + rawPath + "\",\"text\":\"" + text + "\"}]}";
}
private static final class ScriptedChatModel implements ChatModel {
private final java.util.ArrayList<String> agentCalls = new java.util.ArrayList<>();
private String promptText = "";
private String plannerPromptText = "";
private String executorPromptText = "";
private String verifierPromptText = "";
private String composerPromptText = "";
private String executorOutput = """
{
"answer_version": "executor_evidence_v2",
"claims": [
{
"claim_id": "claim-1",
"claim_type": "symptom",
"claim_text": "连接池 active 达到上限",
"support_level": "direct",
"evidence_bindings": [
{
"source_type": "tool_trace",
"source_id": "trace-1",
"tool_name": "query_metrics",
"source_invocation_id": 101,
"raw_path": "$.alerts[0]",
"evidence_excerpt": "active=50 max=50"
}
]
}
],
"hypotheses": [],
"recommended_actions": [],
"missing_info": []
}
""";
private String composerOutput = """
{
"answer_summary": "已确认连接池 active 达到上限。",
"recommended_actions": [
{
"action_text": "补充查询连接池泄漏检测日志",
"reason": "用于确认是否存在连接未释放"
}
],
"user_facing_answer": "已确认连接池 active 达到上限。建议补充查询连接池泄漏检测日志。"
}
""";
private boolean sawVerifierPrompt;
private final java.util.List<String> verifierOutputs;
private int verifierOutputIndex;
private ScriptedChatModel() {
this("""
{
"verdict": "PASS",
"groundedness_score": 1.0,
"critical_fact_count": 1,
"claim_checks": [
{"claim_id":"claim-1","claim_text":"连接池 active 达到上限","claim_type":"symptom","verification":"direct_observation","detail":"covered by scripted verifier","evidence_refs":[]}
],
"facts_checked": [],
"rationale": "scripted pass"
}
""");
}
private ScriptedChatModel(String verifierOutput) {
this.verifierOutputs = java.util.List.of(verifierOutput);
}
private ScriptedChatModel(String... verifierOutputs) {
this.verifierOutputs = java.util.List.of(verifierOutputs);
}
@Override
public ChatResponse call(Prompt prompt) {
promptText = prompt.getContents();
String text;
if (promptText.contains("PLANNER_TEST_PROMPT")) {
agentCalls.add("chat_planner");
plannerPromptText = promptText;
text = "PLANNER_PLAN";
} else if (promptText.contains("EXECUTOR_TEST_PROMPT")) {
agentCalls.add("chat_executor");
executorPromptText = promptText;
text = executorOutput;
} else if (promptText.contains("VERIFIER_TEST_PROMPT")) {
agentCalls.add("chat_verifier");
verifierPromptText = promptText;
sawVerifierPrompt = true;
int index = Math.min(verifierOutputIndex, verifierOutputs.size() - 1);
text = verifierOutputs.get(index);
verifierOutputIndex++;
} else if (promptText.contains("COMPOSER_TEST_PROMPT")) {
agentCalls.add("chat_composer");
composerPromptText = promptText;
text = composerOutput;
} else {
text = "UNEXPECTED_PROMPT";
}
return new ChatResponse(List.of(new Generation(new AssistantMessage(text))));
}
}
}
@@ -1,6 +1,7 @@
package com.superbiz.agent.service;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.fasterxml.jackson.databind.JsonNode;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.ChatSession;
import com.superbiz.agent.domain.entity.DiagnosisRun;
@@ -20,6 +21,7 @@ import java.util.List;
import java.util.Optional;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertNull;
import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.junit.jupiter.api.Assertions.assertTrue;
@@ -89,9 +91,54 @@ class DiagnosisTraceServiceTest {
assertEquals("first question", response.getRun().getQuery());
assertEquals(List.of("planner"),
response.getSteps().stream().map(DiagnosisTraceResponse.AgentStepTrace::getAgentName).toList());
assertEquals("COMPOSER",
response.getRun().getOrchestrationTrace().get("final_node"));
verify(diagnosisRunRepository, never()).findFirstBySessionIdOrderByCreatedAtDescIdDesc(any());
}
@Test
void orchestrationTraceIsParsedOnlyOnRunProjection() {
String sessionId = "trace-session-orchestration";
DiagnosisRun run = run(4L, sessionId, "run-orchestration", "question",
LocalDateTime.of(2026, 7, 17, 12, 0));
when(diagnosisRunRepository.findBySessionIdAndRunId(sessionId, run.getRunId()))
.thenReturn(Optional.of(run));
when(agentStepRepository.findByRunIdOrderByStepIndex(run.getRunId())).thenReturn(List.of());
when(toolInvocationRepository.findByRunIdOrderByIdAsc(run.getRunId())).thenReturn(List.of());
DiagnosisTraceResponse response = service.getTrace(sessionId, run.getRunId());
JsonNode json = new ObjectMapper().findAndRegisterModules()
.valueToTree(response);
assertEquals("stategraph-v1",
response.getRun().getOrchestrationTrace().get("version"));
assertFalse(json.has("orchestrationTrace"));
assertFalse(json.path("session").has("orchestrationTrace"));
assertTrue(json.path("run").path("orchestrationTrace").isObject());
assertFalse(json.path("run").has("orchestrationTraceRaw"));
}
@Test
void nullOrInvalidOrchestrationTraceFailsClosedOnRunProjection() {
String sessionId = "trace-session-invalid-orchestration";
DiagnosisRun run = run(5L, sessionId, "run-invalid-orchestration", "question",
LocalDateTime.of(2026, 7, 17, 12, 1));
run.setOrchestrationTrace("not-json");
when(diagnosisRunRepository.findBySessionIdAndRunId(sessionId, run.getRunId()))
.thenReturn(Optional.of(run));
when(agentStepRepository.findByRunIdOrderByStepIndex(run.getRunId())).thenReturn(List.of());
when(toolInvocationRepository.findByRunIdOrderByIdAsc(run.getRunId())).thenReturn(List.of());
DiagnosisTraceResponse invalid = service.getTrace(sessionId, run.getRunId());
assertNull(invalid.getRun().getOrchestrationTrace());
run.setOrchestrationTrace(null);
DiagnosisTraceResponse historical = service.getTrace(sessionId, run.getRunId());
assertNull(historical.getRun().getOrchestrationTrace());
}
@Test
void getTraceWithRunIdReturnsExactSecondRun() {
String sessionId = "trace-session-exact";
@@ -218,6 +265,7 @@ class DiagnosisTraceServiceTest {
.agentFlow("CHAT")
.answer("answer for " + runId)
.selfEvaluation("{\"verifier_evaluation\":{\"verdict\":\"PASS\"}}")
.orchestrationTrace("{\"version\":\"stategraph-v1\",\"transitions\":[],\"final_node\":\"COMPOSER\",\"termination_reason\":\"composer_completed\",\"degraded\":false,\"evidence_retry_count\":0}")
.feedback("useful")
.stepCount(1)
.toolCallCount(1)