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_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 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> 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 gatekeeperResult = runGatekeeper(sessionId, runId, parseResult); VerifierContextHolder.setGatekeeperResult(gatekeeperResult); Map 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 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 passGatekeeperResult() { GatekeeperRuleCatalog catalog = GatekeeperRuleCatalog.fallback(); Map 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 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 enrichExecutorStructuredOutput(Map structuredOutput, List> toolTraceSummary) { if (structuredOutput == null) { return null; } Map> invocationIdsByTool = invocationIdsByTool(toolTraceSummary); List> 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> invocationIdsByTool, List> 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 binding = (Map) rawBinding; String normalizedToolName = normalizeToolName(binding.get("tool_name")); if (!normalizedToolName.isBlank()) { binding.put("tool_name", normalizedToolName); } if (!hasInvocationId(binding)) { List 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> invocationIdsByTool(List> toolTraceSummary) { Map> idsByTool = new LinkedHashMap<>(); for (Map summary : toolTraceSummary == null ? List.>of() : toolTraceSummary) { String toolName = normalizeToolName(summary.get("tool_name")); if (toolName.isBlank()) { continue; } List ids = toLongList(summary.get("source_invocation_ids")); if (ids.isEmpty()) { continue; } idsByTool.computeIfAbsent(toolName, ignored -> new LinkedHashSet<>()).addAll(ids); } Map> result = new LinkedHashMap<>(); for (Map.Entry> entry : idsByTool.entrySet()) { result.put(entry.getKey(), new ArrayList<>(entry.getValue())); } return result; } private boolean hasInvocationId(Map binding) { if (asLong(binding.get("source_invocation_id")) != null) { return true; } return toLongList(binding.get("source_invocation_ids")).size() == 1; } private List toLongList(Object value) { if (!(value instanceof List values)) { return List.of(); } List 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 status(String status, String detail) { Map result = new LinkedHashMap<>(); result.put("status", status); result.put("detail", detail == null ? "" : detail); return result; } private String extractLastAssistantText(List 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 structuredOutput, Map status ) { } }