test(graph): establish diagnosis stategraph suite
This commit is contained in:
+80
-2
@@ -14,6 +14,7 @@ import java.util.Deque;
|
||||
|
||||
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.ArgumentMatchers.anyMap;
|
||||
import static org.mockito.ArgumentMatchers.eq;
|
||||
import static org.mockito.Mockito.mock;
|
||||
@@ -21,7 +22,44 @@ import static org.mockito.Mockito.times;
|
||||
import static org.mockito.Mockito.verify;
|
||||
import static org.mockito.Mockito.when;
|
||||
|
||||
class DiagnosisRealGraphIntegrationTest {
|
||||
class DiagnosisGraphNodeContractTest {
|
||||
|
||||
@Test
|
||||
void legalSnapshotAfterToolFailureRemainsCompletedAndReachesGatekeeper()
|
||||
throws Exception {
|
||||
ExecutorGatekeeperService gatekeeper = mock(ExecutorGatekeeperService.class);
|
||||
when(gatekeeper.validateRun(eq("run-tool-limited"), anyMap(), anyMap()))
|
||||
.thenReturn(Map.of(
|
||||
"status", "pass",
|
||||
"severity", "none",
|
||||
"checked_bindings", List.of(),
|
||||
"failed_rules", List.of()));
|
||||
String limitedOutput = """
|
||||
{"answer_version":"executor_evidence_v2","claims":[],
|
||||
"hypotheses":[],"recommended_actions":[],
|
||||
"missing_info":["query_logs failed: timeout; no evidence available"]}
|
||||
""";
|
||||
String verifierOutput = """
|
||||
{"verdict":"LOW_CONFID","groundedness_score":0.0,
|
||||
"critical_fact_count":0,"claim_checks":[],"facts_checked":[],
|
||||
"rationale":"工具失败已被合法表达为证据限制"}
|
||||
""";
|
||||
DiagnosisGraphActions actions = new DiagnosisRealGraphActionsFactory().create(
|
||||
constant(plannerOutput()), constant(limitedOutput),
|
||||
constant(verifierOutput), constant(composerOutput()), gatekeeper);
|
||||
|
||||
OverAllState state = new DiagnosisGraphFactory().compile(actions)
|
||||
.invoke(initialState(), config("run-tool-limited"))
|
||||
.orElseThrow();
|
||||
|
||||
assertEquals("COMPLETED", DiagnosisGraphState.stringValue(
|
||||
state, DiagnosisGraphState.EXECUTOR_STATUS));
|
||||
assertEquals(List.of("planner", "executor", "gatekeeper", "verified_input",
|
||||
"verifier", "composer"),
|
||||
eventNodes(state));
|
||||
verify(gatekeeper, times(1)).validateRun(
|
||||
eq("run-tool-limited"), anyMap(), anyMap());
|
||||
}
|
||||
|
||||
@Test
|
||||
void legalNoEvidenceWithVerifiedBindingContinuesUnderLowConfidenceCeiling()
|
||||
@@ -128,7 +166,20 @@ class DiagnosisRealGraphIntegrationTest {
|
||||
.thenReturn(Map.of(
|
||||
"status", "fail",
|
||||
"severity", "reject",
|
||||
"checked_bindings", List.of(),
|
||||
"checked_bindings", List.of(
|
||||
Map.of(
|
||||
"claim_id", "c1",
|
||||
"status", "pass",
|
||||
"tool_name", "query_metrics",
|
||||
"source_invocation_id", 17L,
|
||||
"raw_path", "$.alerts[0]",
|
||||
"matched_text", "active=50 max=50"),
|
||||
Map.of(
|
||||
"claim_id", "c2",
|
||||
"status", "fail",
|
||||
"tool_name", "query_logs",
|
||||
"source_invocation_id", 18L,
|
||||
"raw_path", "$.results[0]")),
|
||||
"failed_rules", List.of("evidence.raw_path")));
|
||||
DiagnosisGraphActions actions = new DiagnosisRealGraphActionsFactory().create(
|
||||
constant(plannerOutput()),
|
||||
@@ -148,6 +199,33 @@ class DiagnosisRealGraphIntegrationTest {
|
||||
eventNodes(state));
|
||||
}
|
||||
|
||||
@Test
|
||||
void verifierInvalidOutputSetsExecutionStatusWithoutFabricatedVerdict()
|
||||
throws Exception {
|
||||
ExecutorGatekeeperService gatekeeper = mock(ExecutorGatekeeperService.class);
|
||||
when(gatekeeper.validateRun(eq("run-verifier-invalid"), anyMap(), anyMap()))
|
||||
.thenReturn(passResult());
|
||||
QueueInvoker verifier = new QueueInvoker("not-json", "still-not-json");
|
||||
DiagnosisGraphActions actions = new DiagnosisRealGraphActionsFactory().create(
|
||||
constant(plannerOutput()), constant(executorOutput()), verifier,
|
||||
constant(composerOutput()), gatekeeper);
|
||||
|
||||
OverAllState state = new DiagnosisGraphFactory().compile(actions)
|
||||
.invoke(initialState(), config("run-verifier-invalid"))
|
||||
.orElseThrow();
|
||||
|
||||
assertEquals("INVALID_OUTPUT", DiagnosisGraphState.stringValue(
|
||||
state, DiagnosisGraphState.VERIFIER_STATUS));
|
||||
assertTrue(state.value(DiagnosisGraphState.VERIFIER_MODEL_VERDICT).isEmpty());
|
||||
assertTrue(state.value(DiagnosisGraphState.EFFECTIVE_VERDICT).isEmpty());
|
||||
assertFalse(DiagnosisGraphState.stringValue(
|
||||
state, DiagnosisGraphState.FINAL_ANSWER)
|
||||
.contains("连接数达到上限"));
|
||||
assertEquals(List.of("planner", "executor", "gatekeeper", "verified_input",
|
||||
"verifier", "verifier", "fallback"),
|
||||
eventNodes(state));
|
||||
}
|
||||
|
||||
@Test
|
||||
void exhaustedComposerRetryUsesVerifiedMaterialWithoutRerunningPredecessors()
|
||||
throws Exception {
|
||||
+44
@@ -0,0 +1,44 @@
|
||||
package com.superbiz.agent.graph.diagnosis;
|
||||
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.util.List;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertFalse;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
|
||||
class DiagnosisGraphTestSuiteStructureTest {
|
||||
|
||||
private static final Path TEST_ROOT = Path.of(
|
||||
"src", "test", "java", "com", "superbiz", "agent");
|
||||
|
||||
@Test
|
||||
void authoritativeGraphLayersExistWithoutSequentialOrHookImplementationTests()
|
||||
throws IOException {
|
||||
List<Path> authoritativeTests = List.of(
|
||||
graphTest("DiagnosisGraphWorkflowTest.java"),
|
||||
graphTest("DiagnosisGraphNodeContractTest.java"),
|
||||
TEST_ROOT.resolve(Path.of(
|
||||
"service", "ChatServiceGraphIntegrationTest.java")));
|
||||
|
||||
authoritativeTests.forEach(path -> assertTrue(
|
||||
Files.isRegularFile(path), "missing authoritative test: " + path));
|
||||
assertFalse(Files.exists(TEST_ROOT.resolve(Path.of(
|
||||
"service", "ChatServiceSequentialAgentTest.java"))));
|
||||
assertFalse(Files.exists(TEST_ROOT.resolve(Path.of(
|
||||
"hook", "VerifierInputHookTest.java"))));
|
||||
|
||||
for (Path path : authoritativeTests) {
|
||||
String source = Files.readString(path);
|
||||
assertFalse(source.contains("SequentialAgent"), path.toString());
|
||||
assertFalse(source.contains("VerifierInputHook"), path.toString());
|
||||
}
|
||||
}
|
||||
|
||||
private Path graphTest(String fileName) {
|
||||
return TEST_ROOT.resolve(Path.of("graph", "diagnosis", fileName));
|
||||
}
|
||||
}
|
||||
+36
-8
@@ -22,7 +22,7 @@ import java.util.stream.Stream;
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertFalse;
|
||||
|
||||
class DiagnosisGraphRoutingTest {
|
||||
class DiagnosisGraphWorkflowTest {
|
||||
|
||||
private static final String THREAD_ID = "run-routing-test";
|
||||
|
||||
@@ -193,15 +193,15 @@ class DiagnosisGraphRoutingTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
void gatekeeperLowConfidenceWithVerifiedBindingRunsVerifiedInput()
|
||||
void gatekeeperCeilingLowConfidenceWithVerifiedBindingDoesNotRetryEvidence()
|
||||
throws Exception {
|
||||
ScriptedDiagnosisGraphActions script = new ScriptedDiagnosisGraphActions();
|
||||
planner(script, PlannerStatus.COMPLETED);
|
||||
executor(script, ExecutorStatus.COMPLETED);
|
||||
gatekeeper(script, GatekeeperStatus.LOW_CONFID, 1);
|
||||
verifiedInput(script);
|
||||
verifier(script, VerifierStatus.COMPLETED, Verdict.LOW_CONFID,
|
||||
Verdict.LOW_CONFID, Map.of());
|
||||
verifier(script, VerifierStatus.COMPLETED, Verdict.PASS,
|
||||
Verdict.LOW_CONFID, Verdict.LOW_CONFID, Map.of());
|
||||
composerCompleted(script);
|
||||
|
||||
OverAllState state = run(script);
|
||||
@@ -209,6 +209,12 @@ class DiagnosisGraphRoutingTest {
|
||||
assertEquals(1, script.calls(DiagnosisGraphTopology.Node.VERIFIED_INPUT));
|
||||
assertEquals(1, script.calls(DiagnosisGraphTopology.Node.VERIFIER));
|
||||
assertEquals(0, script.calls(DiagnosisGraphTopology.Node.EVIDENCE_RETRY));
|
||||
assertEquals(Verdict.PASS,
|
||||
DiagnosisGraphState.enumValue(
|
||||
state,
|
||||
DiagnosisGraphState.VERIFIER_MODEL_VERDICT,
|
||||
Verdict.class)
|
||||
.orElseThrow());
|
||||
assertEquals(Verdict.LOW_CONFID,
|
||||
DiagnosisGraphState.enumValue(
|
||||
state,
|
||||
@@ -272,7 +278,7 @@ class DiagnosisGraphRoutingTest {
|
||||
}
|
||||
|
||||
@Test
|
||||
void evidenceRetryRunsOnceAndResetsPlannerStageCounter()
|
||||
void secondLowConfidenceDoesNotRetryEvidenceAgainAndResetsPlannerCounter()
|
||||
throws Exception {
|
||||
ScriptedDiagnosisGraphActions script = new ScriptedDiagnosisGraphActions();
|
||||
planner(script, PlannerStatus.INVALID_OUTPUT);
|
||||
@@ -294,6 +300,7 @@ class DiagnosisGraphRoutingTest {
|
||||
assertEquals(4, script.calls(DiagnosisGraphTopology.Node.PLANNER));
|
||||
assertEquals(2, script.calls(DiagnosisGraphTopology.Node.EXECUTOR));
|
||||
assertEquals(1, script.calls(DiagnosisGraphTopology.Node.EVIDENCE_RETRY));
|
||||
assertEquals(1, script.calls(DiagnosisGraphTopology.Node.COMPOSER));
|
||||
assertEquals(1, DiagnosisGraphState.intValue(
|
||||
state, DiagnosisGraphState.EVIDENCE_RETRY_COUNT));
|
||||
assertEquals(1, DiagnosisGraphState.intValue(
|
||||
@@ -436,10 +443,16 @@ class DiagnosisGraphRoutingTest {
|
||||
CompiledGraph graph = new DiagnosisGraphFactory().compile(script.actions());
|
||||
assertEquals(DiagnosisGraphTopology.RECURSION_LIMIT,
|
||||
graph.getMaxIterations());
|
||||
return graph.invoke(
|
||||
OverAllState state = graph.invoke(
|
||||
initialState(),
|
||||
RunnableConfig.builder().threadId(THREAD_ID).build())
|
||||
.orElseThrow();
|
||||
assertEquals(script.sequence(), DiagnosisGraphState.listValue(
|
||||
state, DiagnosisGraphState.ORCHESTRATION_EVENTS)
|
||||
.stream()
|
||||
.map(event -> ((OrchestrationEvent) event).node())
|
||||
.toList());
|
||||
return state;
|
||||
}
|
||||
|
||||
private Map<String, Object> initialState() {
|
||||
@@ -548,10 +561,25 @@ class DiagnosisGraphRoutingTest {
|
||||
Verdict verdict,
|
||||
Verdict ceiling,
|
||||
Map<String, Object> verifierOutput) {
|
||||
verifier(script, status, verdict, verdict, ceiling, verifierOutput);
|
||||
}
|
||||
|
||||
private void verifier(
|
||||
ScriptedDiagnosisGraphActions script,
|
||||
VerifierStatus status,
|
||||
Verdict modelVerdict,
|
||||
Verdict effectiveVerdict,
|
||||
Verdict ceiling,
|
||||
Map<String, Object> verifierOutput) {
|
||||
Map<String, Object> update = new java.util.LinkedHashMap<>();
|
||||
update.put(DiagnosisGraphState.VERIFIER_STATUS, status.name());
|
||||
if (verdict != null) {
|
||||
update.put(DiagnosisGraphState.EFFECTIVE_VERDICT, verdict.name());
|
||||
if (modelVerdict != null) {
|
||||
update.put(DiagnosisGraphState.VERIFIER_MODEL_VERDICT,
|
||||
modelVerdict.name());
|
||||
}
|
||||
if (effectiveVerdict != null) {
|
||||
update.put(DiagnosisGraphState.EFFECTIVE_VERDICT,
|
||||
effectiveVerdict.name());
|
||||
}
|
||||
if (ceiling != null) {
|
||||
update.put(DiagnosisGraphState.VERIFIER_VERDICT_CEILING,
|
||||
@@ -1,380 +0,0 @@
|
||||
package com.superbiz.agent.hook;
|
||||
|
||||
import com.alibaba.cloud.ai.graph.RunnableConfig;
|
||||
import com.alibaba.cloud.ai.graph.agent.hook.messages.AgentCommand;
|
||||
import com.fasterxml.jackson.databind.JsonNode;
|
||||
import com.fasterxml.jackson.databind.ObjectMapper;
|
||||
import com.superbiz.agent.domain.entity.ToolInvocation;
|
||||
import com.superbiz.agent.repository.ToolInvocationRepository;
|
||||
import com.superbiz.agent.service.ExecutorGatekeeperService;
|
||||
import com.superbiz.agent.service.ToolTraceSummaryService;
|
||||
import com.superbiz.agent.util.VerifierContextHolder;
|
||||
import org.junit.jupiter.api.AfterEach;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.messages.AssistantMessage;
|
||||
import org.springframework.ai.chat.messages.Message;
|
||||
import org.springframework.ai.chat.messages.UserMessage;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Map;
|
||||
|
||||
import static org.junit.jupiter.api.Assertions.assertEquals;
|
||||
import static org.junit.jupiter.api.Assertions.assertFalse;
|
||||
import static org.junit.jupiter.api.Assertions.assertNotNull;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
import static org.mockito.ArgumentMatchers.anyString;
|
||||
import static org.mockito.Mockito.mock;
|
||||
import static org.mockito.Mockito.when;
|
||||
|
||||
class VerifierInputHookTest {
|
||||
|
||||
private final ObjectMapper objectMapper = new ObjectMapper();
|
||||
|
||||
@AfterEach
|
||||
void tearDown() {
|
||||
VerifierContextHolder.clear();
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelAddsStructuredExecutorOutputWhenJsonContractIsValid() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of(
|
||||
Map.of("trace_ref", "trace-1", "tool_name", "query_metrics")
|
||||
));
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService);
|
||||
VerifierContextHolder.setOriginalQuery("分析 MySQL 连接池耗尽");
|
||||
|
||||
String executorOutput = """
|
||||
{
|
||||
"answer_version": "executor_evidence_v1",
|
||||
"diagnosis_summary": "连接池已满,但缺少泄漏证据。",
|
||||
"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_ids": [101],
|
||||
"evidence_excerpt": "active=50 max=50"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"hypotheses": [],
|
||||
"recommended_actions": [],
|
||||
"missing_info": ["缺少泄漏检测日志"],
|
||||
"user_facing_answer": "已确认连接池 active 达到上限。"
|
||||
}
|
||||
""";
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage(executorOutput)),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "structured-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
assertEquals("valid", payload.path("executor_output_parse_status").path("status").asText());
|
||||
assertEquals("executor_evidence_v1",
|
||||
payload.path("executor_structured_output").path("answer_version").asText());
|
||||
assertEquals("连接池 active 达到上限",
|
||||
payload.path("executor_structured_output").path("claims").get(0).path("claim_text").asText());
|
||||
assertNotNull(VerifierContextHolder.getExecutorStructuredOutput());
|
||||
assertEquals("valid", VerifierContextHolder.getExecutorOutputParseStatus().get("status"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelAddsStructuredExecutorOutputWhenV2ContractHasNoUserFacingAnswer() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of(
|
||||
Map.of("trace_ref", "trace-1", "tool_name", "query_metrics")
|
||||
));
|
||||
ToolInvocationRepository invocationRepository = mock(ToolInvocationRepository.class);
|
||||
when(invocationRepository.findBySessionIdOrderByIdAsc("structured-v2-session")).thenReturn(List.of(
|
||||
invocation(101L, "structured-v2-session", "query_metrics", "$.alerts[0]", "active=50 max=50")
|
||||
));
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService,
|
||||
new ExecutorGatekeeperService(invocationRepository));
|
||||
VerifierContextHolder.setOriginalQuery("分析 MySQL 连接池耗尽");
|
||||
|
||||
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": []
|
||||
}
|
||||
""";
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage(executorOutput)),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "structured-v2-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
assertEquals("valid", payload.path("executor_output_parse_status").path("status").asText());
|
||||
assertEquals("executor_evidence_v2",
|
||||
payload.path("executor_structured_output").path("answer_version").asText());
|
||||
assertFalse(payload.path("executor_structured_output").has("user_facing_answer"));
|
||||
assertEquals("连接池 active 达到上限",
|
||||
payload.path("executor_structured_output").path("claims").get(0).path("claim_text").asText());
|
||||
assertEquals("pass", payload.path("gatekeeper_result").path("status").asText());
|
||||
assertEquals("none", payload.path("gatekeeper_result").path("severity").asText());
|
||||
assertEquals("pass", VerifierContextHolder.getGatekeeperResult().get("status"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelBackfillsOnlyUniqueInvocationIdAndDoesNotPassWithoutRawPath() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of(
|
||||
Map.of(
|
||||
"trace_ref", "metrics-1",
|
||||
"tool_name", "query_metrics",
|
||||
"source_invocation_ids", List.of(101L)
|
||||
)
|
||||
));
|
||||
ToolInvocationRepository invocationRepository = mock(ToolInvocationRepository.class);
|
||||
when(invocationRepository.findBySessionIdOrderByIdAsc("backfill-session")).thenReturn(List.of(
|
||||
invocation(101L, "backfill-session", "query_metrics", "$.alerts[0]",
|
||||
"CPU 使用率持续超过 80%,当前值为 92%")
|
||||
));
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService,
|
||||
new ExecutorGatekeeperService(invocationRepository));
|
||||
|
||||
String executorOutput = """
|
||||
{
|
||||
"answer_version": "executor_evidence_v2",
|
||||
"claims": [
|
||||
{
|
||||
"claim_id": "claim-1",
|
||||
"claim_type": "symptom",
|
||||
"claim_text": "payment-service CPU 使用率超过 92%",
|
||||
"support_level": "direct",
|
||||
"evidence_bindings": [
|
||||
{
|
||||
"source_type": "tool_trace",
|
||||
"source_id": "prometheus-alert-HighCPUUsage",
|
||||
"tool_name": "queryPrometheusAlerts",
|
||||
"evidence_excerpt": "CPU 使用率持续超过 80%,当前值为 92%"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"hypotheses": [],
|
||||
"recommended_actions": [
|
||||
{
|
||||
"action_text": "restart payment-service",
|
||||
"reason": "cpu alert is firing",
|
||||
"evidence_bindings": [
|
||||
{
|
||||
"source_type": "tool_trace",
|
||||
"source_id": "prometheus-alert-HighCPUUsage",
|
||||
"tool_name": "queryPrometheusAlerts",
|
||||
"evidence_excerpt": "CPU usage is 92%"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"missing_info": []
|
||||
}
|
||||
""";
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage(executorOutput)),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "backfill-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
JsonNode binding = payload.path("executor_structured_output")
|
||||
.path("claims").get(0)
|
||||
.path("evidence_bindings").get(0);
|
||||
assertEquals("query_metrics", binding.path("tool_name").asText());
|
||||
assertEquals(101L, binding.path("source_invocation_id").asLong());
|
||||
JsonNode actionBinding = payload.path("executor_structured_output")
|
||||
.path("recommended_actions").get(0)
|
||||
.path("evidence_bindings").get(0);
|
||||
assertEquals("query_metrics", actionBinding.path("tool_name").asText());
|
||||
assertEquals(101L, actionBinding.path("source_invocation_id").asLong());
|
||||
assertFalse(binding.has("raw_path"));
|
||||
assertEquals("fail", payload.path("gatekeeper_result").path("status").asText());
|
||||
assertEquals("low_confid", payload.path("gatekeeper_result").path("severity").asText());
|
||||
assertEquals("evidence.invocation_auto_backfill",
|
||||
payload.path("gatekeeper_result").path("warnings").get(0).path("rule").asText());
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelAddsFailingGatekeeperResultForFabricatedInvocationId() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of());
|
||||
ToolInvocationRepository invocationRepository = mock(ToolInvocationRepository.class);
|
||||
when(invocationRepository.findBySessionIdOrderByIdAsc("fabricated-invocation-session")).thenReturn(List.of(
|
||||
invocation(101L, "fabricated-invocation-session", "query_metrics", "$.alerts[0]", "active=50 max=50")
|
||||
));
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService,
|
||||
new ExecutorGatekeeperService(invocationRepository));
|
||||
|
||||
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",
|
||||
"tool_name": "query_metrics",
|
||||
"source_invocation_id": 999,
|
||||
"raw_path": "$.alerts[0]",
|
||||
"evidence_excerpt": "active=50 max=50"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"hypotheses": [],
|
||||
"recommended_actions": [],
|
||||
"missing_info": []
|
||||
}
|
||||
""";
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage(executorOutput)),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "fabricated-invocation-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
assertEquals("fail", payload.path("gatekeeper_result").path("status").asText());
|
||||
assertEquals("reject", payload.path("gatekeeper_result").path("severity").asText());
|
||||
assertEquals("evidence.invocation_ref",
|
||||
payload.path("gatekeeper_result").path("failed_rules").get(0).asText());
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelExtractsStructuredOutputFromPrefixedJsonFence() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of());
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService);
|
||||
|
||||
String executorOutput = """
|
||||
现在我已经收集了足够的数据,最终输出如下。
|
||||
|
||||
```json
|
||||
{
|
||||
"answer_version": "executor_evidence_v1",
|
||||
"diagnosis_summary": "已确认连接池 active 达到上限。",
|
||||
"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_ids": [101],
|
||||
"evidence_excerpt": "active=50 max=50"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"hypotheses": [],
|
||||
"recommended_actions": [],
|
||||
"missing_info": [],
|
||||
"user_facing_answer": "已确认连接池 active 达到上限。"
|
||||
}
|
||||
```
|
||||
""";
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage(executorOutput)),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "fenced-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
assertEquals("valid", payload.path("executor_output_parse_status").path("status").asText());
|
||||
assertEquals("claim-1",
|
||||
payload.path("executor_structured_output").path("claims").get(0).path("claim_id").asText());
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelMarksMalformedJsonAndKeepsRawAnswerFallback() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of());
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService);
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage("{\"diagnosis_summary\":\"缺少 claims\"}")),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "malformed-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
assertEquals("malformed", payload.path("executor_output_parse_status").path("status").asText());
|
||||
assertTrue(payload.path("executor_structured_output").isNull());
|
||||
assertEquals("{\"diagnosis_summary\":\"缺少 claims\"}", payload.path("executor_final_answer").asText());
|
||||
assertEquals("malformed", VerifierContextHolder.getExecutorOutputParseStatus().get("status"));
|
||||
}
|
||||
|
||||
@Test
|
||||
void beforeModelMarksPlainTextAsMissingStructuredOutput() throws Exception {
|
||||
ToolTraceSummaryService traceSummaryService = mock(ToolTraceSummaryService.class);
|
||||
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of());
|
||||
VerifierInputHook hook = new VerifierInputHook(traceSummaryService);
|
||||
|
||||
AgentCommand command = hook.beforeModel(
|
||||
List.of(new AssistantMessage("普通自然语言答案")),
|
||||
RunnableConfig.builder().addMetadata("sessionId", "plain-session").build()
|
||||
);
|
||||
|
||||
JsonNode payload = readPayload(command);
|
||||
assertEquals("missing", payload.path("executor_output_parse_status").path("status").asText());
|
||||
assertTrue(payload.path("executor_structured_output").isNull());
|
||||
assertFalse(payload.path("executor_final_answer").asText().isBlank());
|
||||
}
|
||||
|
||||
private JsonNode readPayload(AgentCommand command) throws Exception {
|
||||
var field = AgentCommand.class.getDeclaredField("messages");
|
||||
field.setAccessible(true);
|
||||
@SuppressWarnings("unchecked")
|
||||
List<Message> messages = (List<Message>) field.get(command);
|
||||
assertEquals(1, messages.size());
|
||||
Message message = messages.get(0);
|
||||
assertTrue(message instanceof UserMessage);
|
||||
return objectMapper.readTree(((UserMessage) message).getText());
|
||||
}
|
||||
|
||||
private ToolInvocation invocation(Long id, String sessionId, String toolName, String rawPath, String text) {
|
||||
return ToolInvocation.builder()
|
||||
.id(id)
|
||||
.sessionId(sessionId)
|
||||
.toolName(toolName)
|
||||
.retrievalDetails("{\"evidence_refs\":[{\"raw_path\":\"" + rawPath
|
||||
+ "\",\"text\":\"" + text + "\"}]}")
|
||||
.build();
|
||||
}
|
||||
}
|
||||
@@ -19,6 +19,7 @@ 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 com.superbiz.agent.util.SessionContextHolder;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.mockito.ArgumentCaptor;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
@@ -33,6 +34,7 @@ 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.assertNull;
|
||||
import static org.junit.jupiter.api.Assertions.assertTrue;
|
||||
import static org.mockito.ArgumentMatchers.any;
|
||||
import static org.mockito.ArgumentMatchers.anyString;
|
||||
@@ -76,6 +78,8 @@ class ChatServiceGraphIntegrationTest {
|
||||
assertFalse(saved.getSelfEvaluation().contains("tool_trace_summary"));
|
||||
verify(fixture.evaluationService).evaluateRun(result.runId(), result.answer());
|
||||
verify(fixture.retrievedDocTracker).clearSession(sessionId);
|
||||
assertNull(SessionContextHolder.getSessionId());
|
||||
assertNull(SessionContextHolder.getRunId());
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -113,6 +117,8 @@ class ChatServiceGraphIntegrationTest {
|
||||
assertEquals(result.answer(), saved.getAnswer());
|
||||
verify(fixture.evaluationService, never()).evaluateRun(anyString(), anyString());
|
||||
verify(fixture.retrievedDocTracker).clearSession(sessionId);
|
||||
assertNull(SessionContextHolder.getSessionId());
|
||||
assertNull(SessionContextHolder.getRunId());
|
||||
}
|
||||
|
||||
@Test
|
||||
|
||||
Reference in New Issue
Block a user