feat(agent): add executor gatekeeper hook

This commit is contained in:
aruo
2026-07-08 02:01:49 +08:00
parent 050cbc8fee
commit c5e496e715
23 changed files with 1343 additions and 30 deletions
@@ -4,6 +4,9 @@ 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;
@@ -90,7 +93,12 @@ class VerifierInputHookTest {
when(traceSummaryService.buildVerifierTraceSummary(anyString(), anyString())).thenReturn(List.of(
Map.of("trace_ref", "trace-1", "tool_name", "query_metrics")
));
VerifierInputHook hook = new VerifierInputHook(traceSummaryService);
ToolInvocationRepository invocationRepository = mock(ToolInvocationRepository.class);
when(invocationRepository.findBySessionIdOrderByIdAsc("structured-v2-session")).thenReturn(List.of(
ToolInvocation.builder().id(101L).sessionId("structured-v2-session").toolName("query_metrics").build()
));
VerifierInputHook hook = new VerifierInputHook(traceSummaryService,
new ExecutorGatekeeperService(invocationRepository));
VerifierContextHolder.setOriginalQuery("分析 MySQL 连接池耗尽");
String executorOutput = """
@@ -131,6 +139,55 @@ class VerifierInputHookTest {
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("pass", VerifierContextHolder.getGatekeeperResult().get("status"));
}
@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(
ToolInvocation.builder().id(101L).sessionId("fabricated-invocation-session").toolName("query_metrics").build()
));
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_ids": [999],
"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("evidence.invocation_ref",
payload.path("gatekeeper_result").path("failed_rules").get(0).asText());
}
@Test
@@ -8,6 +8,7 @@ import com.superbiz.agent.agent.tool.QueryLogsTools;
import com.superbiz.agent.agent.tool.QueryMetricsTools;
import com.superbiz.agent.domain.entity.AgentStep;
import com.superbiz.agent.domain.entity.DiagnosisSession;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.AgentStepRepository;
import com.superbiz.agent.repository.DiagnosisSessionRepository;
import com.superbiz.agent.repository.ToolInvocationRepository;
@@ -21,6 +22,7 @@ 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;
@@ -33,6 +35,8 @@ 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.isNull;
import static org.mockito.Mockito.verify;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
@@ -330,6 +334,63 @@ class ChatServiceSequentialAgentTest {
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")
.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_ids": [101],
"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"));
@SuppressWarnings("unchecked")
Map<String, Object> gatekeeperResult = (Map<String, Object>) verifierEvaluation.get("gatekeeper_result");
assertEquals("pass", gatekeeperResult.get("status"));
}
@Test
void buildMethodToolsArrayIncludesLogsAndMetricsWhenAvailable() {
ChatService chatService = new ChatService();
@@ -432,6 +493,7 @@ class ChatServiceSequentialAgentTest {
when(toolTraceSummaryService.buildVerifierTraceSummary(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());
@@ -444,6 +506,7 @@ class ChatServiceSequentialAgentTest {
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");
@@ -0,0 +1,98 @@
package com.superbiz.agent.service;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.ToolInvocationRepository;
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.assertTrue;
import static org.mockito.Mockito.mock;
import static org.mockito.Mockito.when;
class ExecutorGatekeeperServiceTest {
@Test
void validatePassesForExecutorEvidenceV2WithMatchingInvocation() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findBySessionIdOrderByIdAsc("session-1")).thenReturn(List.of(
ToolInvocation.builder().id(101L).sessionId("session-1").toolName("query_metrics").build()
));
ExecutorGatekeeperService service = new ExecutorGatekeeperService(repository);
Map<String, Object> result = service.validate("session-1", validOutput(101L, "query_metrics"),
Map.of("status", "valid"));
assertEquals("pass", result.get("status"));
assertTrue(((List<?>) result.get("failed_rules")).isEmpty());
}
@Test
void validateFailsWhenRemovedFieldsArePresent() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findBySessionIdOrderByIdAsc("session-1")).thenReturn(List.of());
ExecutorGatekeeperService service = new ExecutorGatekeeperService(repository);
Map<String, Object> output = validOutput(101L, "query_metrics");
output.put("user_facing_answer", "旧版最终答案");
Map<String, Object> result = service.validate("session-1", output, Map.of("status", "valid"));
assertEquals("fail", result.get("status"));
assertTrue(((List<?>) result.get("failed_rules")).contains("schema.executor_v2"));
}
@Test
void validateFailsForFabricatedInvocationId() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findBySessionIdOrderByIdAsc("session-1")).thenReturn(List.of(
ToolInvocation.builder().id(101L).sessionId("session-1").toolName("query_metrics").build()
));
ExecutorGatekeeperService service = new ExecutorGatekeeperService(repository);
Map<String, Object> result = service.validate("session-1", validOutput(999L, "query_metrics"),
Map.of("status", "valid"));
assertEquals("fail", result.get("status"));
assertTrue(((List<?>) result.get("failed_rules")).contains("evidence.invocation_ref"));
}
@Test
void validateFailsForToolNameMismatch() {
ToolInvocationRepository repository = mock(ToolInvocationRepository.class);
when(repository.findBySessionIdOrderByIdAsc("session-1")).thenReturn(List.of(
ToolInvocation.builder().id(101L).sessionId("session-1").toolName("query_logs").build()
));
ExecutorGatekeeperService service = new ExecutorGatekeeperService(repository);
Map<String, Object> result = service.validate("session-1", validOutput(101L, "query_metrics"),
Map.of("status", "valid"));
assertEquals("fail", result.get("status"));
assertTrue(((List<?>) result.get("failed_rules")).contains("evidence.invocation_ref"));
}
@SuppressWarnings("unchecked")
private Map<String, Object> validOutput(Long invocationId, String toolName) {
return new java.util.LinkedHashMap<>(Map.of(
"answer_version", "executor_evidence_v2",
"claims", List.of(Map.of(
"claim_id", "claim-1",
"claim_type", "symptom",
"claim_text", "连接池 active 达到上限",
"support_level", "direct",
"evidence_bindings", List.of(Map.of(
"source_type", "tool_trace",
"source_id", "trace-1",
"tool_name", toolName,
"source_invocation_ids", List.of(invocationId),
"evidence_excerpt", "active=50 max=50"
))
)),
"hypotheses", List.of(),
"recommended_actions", List.of(),
"missing_info", List.of()
));
}
}