feat(agent): add executor gatekeeper hook
This commit is contained in:
@@ -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()
|
||||
));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user