231 lines
10 KiB
Java
231 lines
10 KiB
Java
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.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")
|
|
));
|
|
VerifierInputHook hook = new VerifierInputHook(traceSummaryService);
|
|
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_ids": [101],
|
|
"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());
|
|
}
|
|
|
|
@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());
|
|
}
|
|
}
|