Files
SuperBizAgent-java/src/main/java/com/superbiz/agent/hook/VerifierInputHook.java
T

353 lines
15 KiB
Java

package com.superbiz.agent.hook;
import com.alibaba.cloud.ai.graph.RunnableConfig;
import com.alibaba.cloud.ai.graph.agent.hook.HookPosition;
import com.alibaba.cloud.ai.graph.agent.hook.HookPositions;
import com.alibaba.cloud.ai.graph.agent.hook.messages.AgentCommand;
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.GatekeeperRuleCatalog;
import com.superbiz.agent.service.ToolTraceSummaryService;
import com.superbiz.agent.util.SessionContextHolder;
import com.superbiz.agent.util.VerifierContextHolder;
import lombok.extern.slf4j.Slf4j;
import org.springframework.ai.chat.messages.AssistantMessage;
import org.springframework.ai.chat.messages.Message;
import org.springframework.ai.chat.messages.UserMessage;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Map;
import java.util.Set;
/**
* Replaces verifier history with an explicit structured payload.
*/
@Slf4j
@HookPositions(HookPosition.BEFORE_MODEL)
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
public String getName() {
return "verifier_input_hook";
}
@Override
public AgentCommand beforeModel(List<Message> previousMessages, RunnableConfig config) {
try {
String sessionId = config.metadata("sessionId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getSessionId);
String runId = config.metadata("runId")
.map(Object::toString)
.orElseGet(SessionContextHolder::getRunId);
String executorFinalAnswer = VerifierContextHolder.getExecutorFinalAnswer();
if (executorFinalAnswer == null || executorFinalAnswer.isBlank()) {
executorFinalAnswer = extractLastAssistantText(previousMessages);
}
List<Map<String, Object>> toolTraceSummary = runId == null || runId.isBlank()
? toolTraceSummaryService.buildVerifierTraceSummary(sessionId, executorFinalAnswer)
: toolTraceSummaryService.buildVerifierTraceSummaryForRun(runId, executorFinalAnswer);
VerifierContextHolder.setToolTraceSummary(toolTraceSummary);
ExecutorOutputParseResult parseResult = parseExecutorOutput(executorFinalAnswer);
parseResult = new ExecutorOutputParseResult(
enrichExecutorStructuredOutput(parseResult.structuredOutput(), toolTraceSummary),
parseResult.status()
);
VerifierContextHolder.setExecutorStructuredOutput(parseResult.structuredOutput());
VerifierContextHolder.setExecutorOutputParseStatus(parseResult.status());
Map<String, Object> gatekeeperResult = runGatekeeper(sessionId, runId, 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);
return new AgentCommand(List.of(new UserMessage(payload)));
} catch (Exception e) {
log.error("Failed to build verifier input, fallback to original messages", e);
return new AgentCommand(previousMessages);
}
}
private Map<String, Object> runGatekeeper(String sessionId, String runId, ExecutorOutputParseResult parseResult) {
if (executorGatekeeperService == null) {
return passGatekeeperResult();
}
try {
if (runId != null && !runId.isBlank()) {
return executorGatekeeperService.validateRun(runId, parseResult.structuredOutput(), parseResult.status());
}
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() {
GatekeeperRuleCatalog catalog = GatekeeperRuleCatalog.fallback();
Map<String, Object> result = new LinkedHashMap<>();
result.put("status", "pass");
result.put("severity", "none");
result.put("rule_set_version", catalog.version());
result.put("rules", catalog.auditRules());
result.put("checked_bindings", List.of());
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"));
}
String sanitized = sanitizeJsonPayload(executorFinalAnswer);
if (!looksJsonLike(sanitized)) {
return new ExecutorOutputParseResult(null, status("missing", "executor output is not JSON"));
}
try {
JsonNode root = objectMapper.readTree(sanitized);
if (!root.isObject() || !root.path("claims").isArray()) {
return new ExecutorOutputParseResult(null, status("malformed",
"executor output JSON does not match evidence-attribution contract"));
}
Map<String, Object> structuredOutput = objectMapper.convertValue(root, MAP_TYPE);
return new ExecutorOutputParseResult(structuredOutput, status("valid", "parsed executor evidence contract"));
} catch (Exception e) {
log.debug("Failed to parse executor structured output", e);
return new ExecutorOutputParseResult(null, status("malformed", e.getMessage()));
}
}
@SuppressWarnings("unchecked")
private Map<String, Object> enrichExecutorStructuredOutput(Map<String, Object> structuredOutput,
List<Map<String, Object>> toolTraceSummary) {
if (structuredOutput == null) {
return null;
}
Map<String, List<Long>> invocationIdsByTool = invocationIdsByTool(toolTraceSummary);
List<Map<String, Object>> warnings = new ArrayList<>();
enrichEvidenceBindingsInSection(structuredOutput.get("claims"), invocationIdsByTool, warnings);
enrichEvidenceBindingsInSection(structuredOutput.get("recommended_actions"), invocationIdsByTool, warnings);
if (!warnings.isEmpty()) {
structuredOutput.put("_gatekeeper_warnings", warnings);
}
return structuredOutput;
}
@SuppressWarnings("unchecked")
private void enrichEvidenceBindingsInSection(Object sectionValue,
Map<String, List<Long>> invocationIdsByTool,
List<Map<String, Object>> warnings) {
if (!(sectionValue instanceof List<?> items)) {
return;
}
for (Object itemValue : items) {
if (!(itemValue instanceof Map<?, ?> item)) {
continue;
}
Object bindingsValue = item.get("evidence_bindings");
if (!(bindingsValue instanceof List<?> bindings)) {
continue;
}
for (Object bindingValue : bindings) {
if (!(bindingValue instanceof Map<?, ?> rawBinding)) {
continue;
}
Map<String, Object> binding = (Map<String, Object>) rawBinding;
String normalizedToolName = normalizeToolName(binding.get("tool_name"));
if (!normalizedToolName.isBlank()) {
binding.put("tool_name", normalizedToolName);
}
if (!hasInvocationId(binding)) {
List<Long> ids = invocationIdsByTool.getOrDefault(normalizedToolName, List.of());
if (ids.size() == 1) {
binding.put("source_invocation_id", ids.get(0));
warnings.add(Map.of(
"rule", "evidence.invocation_auto_backfill",
"message", "source_invocation_id was auto-filled from the unique tool invocation candidate; raw_path remains missing if Executor did not provide it",
"tool_name", normalizedToolName,
"source_invocation_id", ids.get(0)
));
}
}
}
}
}
private Map<String, List<Long>> invocationIdsByTool(List<Map<String, Object>> toolTraceSummary) {
Map<String, Set<Long>> idsByTool = new LinkedHashMap<>();
for (Map<String, Object> summary : toolTraceSummary == null ? List.<Map<String, Object>>of() : toolTraceSummary) {
String toolName = normalizeToolName(summary.get("tool_name"));
if (toolName.isBlank()) {
continue;
}
List<Long> ids = toLongList(summary.get("source_invocation_ids"));
if (ids.isEmpty()) {
continue;
}
idsByTool.computeIfAbsent(toolName, ignored -> new LinkedHashSet<>()).addAll(ids);
}
Map<String, List<Long>> result = new LinkedHashMap<>();
for (Map.Entry<String, Set<Long>> entry : idsByTool.entrySet()) {
result.put(entry.getKey(), new ArrayList<>(entry.getValue()));
}
return result;
}
private boolean hasInvocationId(Map<String, Object> binding) {
if (asLong(binding.get("source_invocation_id")) != null) {
return true;
}
return toLongList(binding.get("source_invocation_ids")).size() == 1;
}
private List<Long> toLongList(Object value) {
if (!(value instanceof List<?> values)) {
return List.of();
}
List<Long> ids = new ArrayList<>();
for (Object item : values) {
Long id = asLong(item);
if (id != null) {
ids.add(id);
}
}
return ids;
}
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 String normalizeToolName(Object value) {
String toolName = value == null ? "" : String.valueOf(value);
return switch (toolName) {
case "lookupKnowledge" -> "lookup_knowledge";
case "queryLogs" -> "query_logs";
case "queryPrometheusAlerts" -> "query_metrics";
case "getAvailableLogTopics" -> "get_available_log_topics";
default -> toolName;
};
}
private String sanitizeJsonPayload(String raw) {
String trimmed = raw.trim();
int fenceStart = trimmed.indexOf("```");
if (fenceStart >= 0) {
int firstNewline = trimmed.indexOf('\n', fenceStart);
int lastFence = trimmed.indexOf("```", firstNewline + 1);
if (firstNewline >= 0 && lastFence > firstNewline) {
return trimmed.substring(firstNewline + 1, lastFence).trim();
}
}
int objectStart = trimmed.indexOf('{');
int objectEnd = trimmed.lastIndexOf('}');
if (objectStart >= 0 && objectEnd > objectStart) {
return trimmed.substring(objectStart, objectEnd + 1).trim();
}
return trimmed;
}
private boolean looksJsonLike(String text) {
return text.startsWith("{") && text.endsWith("}");
}
private Map<String, Object> status(String status, String detail) {
Map<String, Object> result = new LinkedHashMap<>();
result.put("status", status);
result.put("detail", detail == null ? "" : detail);
return result;
}
private String extractLastAssistantText(List<Message> previousMessages) {
for (int i = previousMessages.size() - 1; i >= 0; i--) {
if (previousMessages.get(i) instanceof AssistantMessage assistantMessage) {
String text = extractTextContent(assistantMessage);
if (text != null && !text.isBlank()) {
return text;
}
}
}
return "";
}
private String extractTextContent(AssistantMessage message) {
try {
try {
return message.getText();
} catch (Exception ignore) {
// Fallback for older implementations.
}
for (String methodName : List.of("getText", "getContent")) {
try {
var method = message.getClass().getMethod(methodName);
Object value = method.invoke(message);
if (value != null) {
return value.toString();
}
} catch (NoSuchMethodException ignore) {
// continue
}
}
} catch (Exception e) {
log.debug("Failed to extract verifier assistant text", e);
}
return message.toString();
}
private record ExecutorOutputParseResult(
Map<String, Object> structuredOutput,
Map<String, Object> status
) {
}
}