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
@@ -8,6 +8,7 @@ import com.alibaba.cloud.ai.graph.agent.hook.messages.MessagesModelHook;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.service.ExecutorGatekeeperService;
import com.superbiz.agent.service.ToolTraceSummaryService;
import com.superbiz.agent.util.SessionContextHolder;
import com.superbiz.agent.util.VerifierContextHolder;
@@ -28,12 +29,19 @@ import java.util.Map;
public class VerifierInputHook extends MessagesModelHook {
private final ToolTraceSummaryService toolTraceSummaryService;
private final ExecutorGatekeeperService executorGatekeeperService;
private final ObjectMapper objectMapper = new ObjectMapper();
private static final TypeReference<Map<String, Object>> MAP_TYPE = new TypeReference<>() {
};
public VerifierInputHook(ToolTraceSummaryService toolTraceSummaryService) {
this(toolTraceSummaryService, null);
}
public VerifierInputHook(ToolTraceSummaryService toolTraceSummaryService,
ExecutorGatekeeperService executorGatekeeperService) {
this.toolTraceSummaryService = toolTraceSummaryService;
this.executorGatekeeperService = executorGatekeeperService;
}
@Override
@@ -60,12 +68,16 @@ public class VerifierInputHook extends MessagesModelHook {
toolTraceSummaryService.buildVerifierTraceSummary(sessionId, executorFinalAnswer);
VerifierContextHolder.setToolTraceSummary(toolTraceSummary);
Map<String, Object> gatekeeperResult = runGatekeeper(sessionId, parseResult);
VerifierContextHolder.setGatekeeperResult(gatekeeperResult);
Map<String, Object> verifierInput = new LinkedHashMap<>();
verifierInput.put("original_query", VerifierContextHolder.getOriginalQuery());
verifierInput.put("executor_final_answer", executorFinalAnswer);
verifierInput.put("executor_structured_output", parseResult.structuredOutput());
verifierInput.put("executor_output_parse_status", parseResult.status());
verifierInput.put("tool_trace_summary", toolTraceSummary);
verifierInput.put("gatekeeper_result", gatekeeperResult);
verifierInput.put("retry_context", VerifierContextHolder.getRetryContext());
String payload = objectMapper.writerWithDefaultPrettyPrinter().writeValueAsString(verifierInput);
@@ -76,6 +88,29 @@ public class VerifierInputHook extends MessagesModelHook {
}
}
private Map<String, Object> runGatekeeper(String sessionId, ExecutorOutputParseResult parseResult) {
if (executorGatekeeperService == null) {
return passGatekeeperResult();
}
try {
return executorGatekeeperService.validate(sessionId, parseResult.structuredOutput(), parseResult.status());
} catch (Exception e) {
log.error("Gatekeeper validation failed unexpectedly", e);
return executorGatekeeperService.fail("gatekeeper.internal_error",
"gatekeeper",
e.getMessage() == null ? "gatekeeper validation failed" : e.getMessage());
}
}
private Map<String, Object> passGatekeeperResult() {
Map<String, Object> result = new LinkedHashMap<>();
result.put("status", "pass");
result.put("failed_rules", List.of());
result.put("warnings", List.of());
result.put("errors", List.of());
return result;
}
private ExecutorOutputParseResult parseExecutorOutput(String executorFinalAnswer) {
if (executorFinalAnswer == null || executorFinalAnswer.isBlank()) {
return new ExecutorOutputParseResult(null, status("missing", "executor_final_answer is blank"));
@@ -112,6 +112,9 @@ public class ChatService {
@Autowired
private SelfEvaluationMergeService selfEvaluationMergeService;
@Autowired
private ExecutorGatekeeperService executorGatekeeperService;
@Value("${verifier.low-confidence-threshold:0.5}")
private double verifierLowConfidenceThreshold;
@@ -541,7 +544,7 @@ public class ChatService {
.model(chatModel)
.systemPrompt(chatVerifierPrompt)
.hooks(new AgentLoggingHook(agentStepRepository, "verifier"),
new VerifierInputHook(toolTraceSummaryService))
new VerifierInputHook(toolTraceSummaryService, executorGatekeeperService))
.outputKey("verifier_output")
.build();
}
@@ -755,6 +758,9 @@ public class ChatService {
verifierEvaluation.put("executor_structured_output", VerifierContextHolder.getExecutorStructuredOutput());
verifierEvaluation.put("tool_trace_summary",
Optional.ofNullable(VerifierContextHolder.getToolTraceSummary()).orElse(List.of()));
verifierEvaluation.put("gatekeeper_result",
Optional.ofNullable(VerifierContextHolder.getGatekeeperResult())
.orElse(Map.of("status", "pass", "failed_rules", List.of(), "warnings", List.of(), "errors", List.of())));
String merged = selfEvaluationMergeService.mergeVerifierEvaluation(session.getSelfEvaluation(), verifierEvaluation);
session.setSelfEvaluation(merged);
@@ -0,0 +1,237 @@
package com.superbiz.agent.service;
import com.superbiz.agent.domain.entity.ToolInvocation;
import com.superbiz.agent.repository.ToolInvocationRepository;
import org.springframework.stereotype.Service;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.function.Function;
import java.util.stream.Collectors;
/**
* Deterministic checks for Executor structured output before verifier reasoning.
*/
@Service
public class ExecutorGatekeeperService {
public static final String STATUS_PASS = "pass";
public static final String STATUS_WARN = "warn";
public static final String STATUS_FAIL = "fail";
public static final String RULE_SCHEMA = "schema.executor_v2";
public static final String RULE_INVOCATION_REF = "evidence.invocation_ref";
private final ToolInvocationRepository toolInvocationRepository;
public ExecutorGatekeeperService(ToolInvocationRepository toolInvocationRepository) {
this.toolInvocationRepository = toolInvocationRepository;
}
public Map<String, Object> validate(String sessionId,
Map<String, Object> structuredOutput,
Map<String, Object> parseStatus) {
GatekeeperResult result = new GatekeeperResult();
validateSchema(structuredOutput, parseStatus, result);
if (structuredOutput != null) {
validateInvocationRefs(sessionId, structuredOutput, result);
}
return result.toMap();
}
public Map<String, Object> pass() {
return new GatekeeperResult().toMap();
}
public Map<String, Object> fail(String ruleId, String target, String message) {
GatekeeperResult result = new GatekeeperResult();
result.fail(ruleId, target, message);
return result.toMap();
}
private void validateSchema(Map<String, Object> structuredOutput,
Map<String, Object> parseStatus,
GatekeeperResult result) {
String status = parseStatus == null ? "" : String.valueOf(parseStatus.getOrDefault("status", ""));
if (structuredOutput == null) {
if ("valid".equals(status)) {
result.fail(RULE_SCHEMA, "executor_structured_output", "structured output is missing after valid parse");
}
return;
}
if (!"executor_evidence_v2".equals(String.valueOf(structuredOutput.get("answer_version")))) {
result.fail(RULE_SCHEMA, "answer_version", "answer_version must be executor_evidence_v2");
}
if (structuredOutput.containsKey("diagnosis_summary")) {
result.fail(RULE_SCHEMA, "diagnosis_summary", "diagnosis_summary is removed from executor_evidence_v2");
}
if (structuredOutput.containsKey("user_facing_answer")) {
result.fail(RULE_SCHEMA, "user_facing_answer", "user_facing_answer is removed from executor_evidence_v2");
}
Object claimsValue = structuredOutput.get("claims");
if (!(claimsValue instanceof List<?> claims)) {
result.fail(RULE_SCHEMA, "claims", "claims must be an array");
return;
}
for (int i = 0; i < claims.size(); i++) {
String target = "claims[" + i + "]";
Object claimValue = claims.get(i);
if (!(claimValue instanceof Map<?, ?> claim)) {
result.fail(RULE_SCHEMA, target, "claim must be an object");
continue;
}
requireString(claim, "claim_id", target, result);
requireString(claim, "claim_type", target, result);
requireString(claim, "claim_text", target, result);
String supportLevel = stringValue(claim.get("support_level"));
if (!"direct".equals(supportLevel) && !"indirect".equals(supportLevel)) {
result.fail(RULE_SCHEMA, target + ".support_level", "support_level must be direct or indirect");
}
Object bindings = claim.get("evidence_bindings");
if (!(bindings instanceof List<?> bindingList) || bindingList.isEmpty()) {
result.fail(RULE_SCHEMA, target + ".evidence_bindings", "claims must include non-empty evidence_bindings");
}
}
requireArray(structuredOutput, "hypotheses", result);
requireArray(structuredOutput, "recommended_actions", result);
requireArray(structuredOutput, "missing_info", result);
}
private void validateInvocationRefs(String sessionId, Map<String, Object> structuredOutput, GatekeeperResult result) {
if (sessionId == null || sessionId.isBlank()) {
result.fail(RULE_INVOCATION_REF, "session_id", "session id is required to validate source_invocation_ids");
return;
}
Map<Long, ToolInvocation> validInvocations = toolInvocationRepository.findBySessionIdOrderByIdAsc(sessionId)
.stream()
.filter(invocation -> invocation.getId() != null)
.collect(Collectors.toMap(ToolInvocation::getId, Function.identity(), (left, right) -> left));
Object claimsValue = structuredOutput.get("claims");
if (!(claimsValue instanceof List<?> claims)) {
return;
}
for (int claimIndex = 0; claimIndex < claims.size(); claimIndex++) {
Object claimValue = claims.get(claimIndex);
if (!(claimValue instanceof Map<?, ?> claim)) {
continue;
}
Object bindingsValue = claim.get("evidence_bindings");
if (!(bindingsValue instanceof List<?> bindings)) {
continue;
}
for (int bindingIndex = 0; bindingIndex < bindings.size(); bindingIndex++) {
String target = "claims[" + claimIndex + "].evidence_bindings[" + bindingIndex + "]";
Object bindingValue = bindings.get(bindingIndex);
if (!(bindingValue instanceof Map<?, ?> binding)) {
result.fail(RULE_INVOCATION_REF, target, "evidence binding must be an object");
continue;
}
validateBindingInvocationIds(binding, validInvocations, target, result);
}
}
}
private void validateBindingInvocationIds(Map<?, ?> binding,
Map<Long, ToolInvocation> validInvocations,
String target,
GatekeeperResult result) {
Object idsValue = binding.get("source_invocation_ids");
if (!(idsValue instanceof List<?> ids) || ids.isEmpty()) {
result.fail(RULE_INVOCATION_REF, target + ".source_invocation_ids",
"source_invocation_ids must be a non-empty array");
return;
}
String claimedToolName = stringValue(binding.get("tool_name"));
if (claimedToolName.isBlank()) {
result.fail(RULE_INVOCATION_REF, target + ".tool_name", "tool_name is required");
}
Set<Long> checkedIds = new HashSet<>();
for (Object idValue : ids) {
Long id = asLong(idValue);
if (id == null) {
result.fail(RULE_INVOCATION_REF, target + ".source_invocation_ids",
"source_invocation_ids must contain numeric ids");
continue;
}
if (!checkedIds.add(id)) {
continue;
}
ToolInvocation invocation = validInvocations.get(id);
if (invocation == null) {
result.fail(RULE_INVOCATION_REF, target, "source_invocation_ids not found in current session: " + id);
continue;
}
if (!claimedToolName.isBlank() && !Objects.equals(claimedToolName, invocation.getToolName())) {
result.fail(RULE_INVOCATION_REF, target + ".tool_name",
"tool_name does not match invocation " + id + ": expected " + invocation.getToolName());
}
}
}
private void requireArray(Map<String, Object> output, String field, GatekeeperResult result) {
if (!(output.get(field) instanceof List<?>)) {
result.fail(RULE_SCHEMA, field, field + " must be an array");
}
}
private void requireString(Map<?, ?> object, String field, String target, GatekeeperResult result) {
if (stringValue(object.get(field)).isBlank()) {
result.fail(RULE_SCHEMA, target + "." + field, field + " is required");
}
}
private String stringValue(Object value) {
return value == null ? "" : String.valueOf(value);
}
private Long asLong(Object value) {
if (value instanceof Number number) {
return number.longValue();
}
if (value instanceof String text) {
try {
return Long.parseLong(text);
} catch (NumberFormatException ignored) {
return null;
}
}
return null;
}
private static final class GatekeeperResult {
private final List<String> failedRules = new ArrayList<>();
private final List<String> warnings = new ArrayList<>();
private final List<Map<String, Object>> errors = new ArrayList<>();
void fail(String ruleId, String target, String message) {
if (!failedRules.contains(ruleId)) {
failedRules.add(ruleId);
}
Map<String, Object> error = new LinkedHashMap<>();
error.put("rule_id", ruleId);
error.put("target", target);
error.put("message", message);
errors.add(error);
}
Map<String, Object> toMap() {
Map<String, Object> result = new LinkedHashMap<>();
result.put("status", failedRules.isEmpty() ? (warnings.isEmpty() ? STATUS_PASS : STATUS_WARN) : STATUS_FAIL);
result.put("failed_rules", failedRules);
result.put("warnings", warnings);
result.put("errors", errors);
return result;
}
}
}
@@ -14,6 +14,7 @@ public final class VerifierContextHolder {
private static final ThreadLocal<Map<String, Object>> EXECUTOR_STRUCTURED_OUTPUT = new ThreadLocal<>();
private static final ThreadLocal<Map<String, Object>> EXECUTOR_OUTPUT_PARSE_STATUS = new ThreadLocal<>();
private static final ThreadLocal<List<Map<String, Object>>> TOOL_TRACE_SUMMARY = new ThreadLocal<>();
private static final ThreadLocal<Map<String, Object>> GATEKEEPER_RESULT = new ThreadLocal<>();
private VerifierContextHolder() {
}
@@ -66,6 +67,14 @@ public final class VerifierContextHolder {
return TOOL_TRACE_SUMMARY.get();
}
public static void setGatekeeperResult(Map<String, Object> gatekeeperResult) {
GATEKEEPER_RESULT.set(gatekeeperResult);
}
public static Map<String, Object> getGatekeeperResult() {
return GATEKEEPER_RESULT.get();
}
public static void clear() {
ORIGINAL_QUERY.remove();
RETRY_CONTEXT.remove();
@@ -73,5 +82,6 @@ public final class VerifierContextHolder {
EXECUTOR_STRUCTURED_OUTPUT.remove();
EXECUTOR_OUTPUT_PARSE_STATUS.remove();
TOOL_TRACE_SUMMARY.remove();
GATEKEEPER_RESULT.remove();
}
}