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