353 lines
15 KiB
Java
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
|
|
) {
|
|
}
|
|
}
|