feat(harness): add information gain stop and audit

This commit is contained in:
aruo
2026-07-27 01:03:34 +08:00
parent de5a5b09d9
commit d0452184ee
92 changed files with 5019 additions and 122 deletions
@@ -0,0 +1,19 @@
package com.superbiz.agent.harness.progress;
public record CompletedToolCall(
String toolCallId,
String toolName,
String normalizedScope) {
public CompletedToolCall {
requireText(toolCallId, "toolCallId");
requireText(toolName, "toolName");
requireText(normalizedScope, "normalizedScope");
}
private static void requireText(String value, String name) {
if (value == null || value.isBlank()) {
throw new IllegalArgumentException(name + " must not be blank");
}
}
}
@@ -0,0 +1,6 @@
package com.superbiz.agent.harness.progress;
public enum DiagnosisCollectionState {
COLLECTING,
SATURATED
}
@@ -0,0 +1,13 @@
package com.superbiz.agent.harness.progress;
import com.superbiz.agent.harness.core.RunContext;
@FunctionalInterface
public interface DiagnosisProgressProjection {
DiagnosisProgressSnapshot project(RunContext context);
static DiagnosisProgressProjection empty() {
return ignored -> DiagnosisProgressSnapshot.empty();
}
}
@@ -0,0 +1,228 @@
package com.superbiz.agent.harness.progress;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.JsonNode;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.contract.SafeFallback;
import com.superbiz.agent.harness.core.RunContext;
import com.superbiz.agent.harness.tool.contract.AgentToolContracts;
import com.superbiz.agent.harness.tool.store.CanonicalInvocationStore;
import com.superbiz.agent.harness.tool.store.CanonicalToolInvocation;
import com.superbiz.agent.harness.tool.store.ToolCallKeyFactory;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map;
import java.util.Objects;
public final class DiagnosisProgressProjector implements DiagnosisProgressProjection {
private static final int MAX_FACTS = 12;
private static final int MAX_SUMMARY_CHARS = 320;
private static final int MAX_SCOPE_CHARS = 320;
private final CanonicalInvocationStore store;
private final ToolCallKeyFactory keyFactory;
private final ObjectMapper objectMapper;
public DiagnosisProgressProjector(CanonicalInvocationStore store,
ToolCallKeyFactory keyFactory,
ObjectMapper objectMapper) {
this.store = Objects.requireNonNull(store, "store must not be null");
this.keyFactory = Objects.requireNonNull(keyFactory, "keyFactory must not be null");
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
}
@Override
public DiagnosisProgressSnapshot project(RunContext context) {
Objects.requireNonNull(context, "context must not be null");
DiagnosisProgressSnapshotState state = context.progress().snapshot();
Map<String, SafeFallback.VerifiedSource> sources = new LinkedHashMap<>();
Map<String, SafeFallback.ObservedFact> facts = new LinkedHashMap<>();
List<String> limitations = new ArrayList<>();
for (CompletedToolCall completed : state.completedToolCalls()) {
CanonicalToolInvocation invocation = resolve(context, completed, limitations);
if (invocation == null) {
continue;
}
try {
projectInvocation(completed, invocation, sources, facts);
} catch (RuntimeException exception) {
addLimitation(limitations, "部分已完成的工具结果格式无法验证,未纳入已检查事实");
}
if (facts.size() >= MAX_FACTS) {
addLimitation(limitations, "已检查事实较多,展示内容已截断");
break;
}
}
return new DiagnosisProgressSnapshot(
List.copyOf(sources.values()),
List.copyOf(facts.values()),
List.copyOf(limitations),
state.stopReason());
}
private CanonicalToolInvocation resolve(RunContext context,
CompletedToolCall completed,
List<String> limitations) {
try {
String key = keyFactory.create(context.runId(), completed.toolCallId());
CanonicalToolInvocation invocation = store.find(key).orElse(null);
if (invocation == null
|| !invocation.isReferencableBy(context.runId())
|| !completed.toolCallId().equals(invocation.toolCallId())
|| !completed.toolName().equals(invocation.toolName())) {
addLimitation(limitations, "部分已完成的工具记录无法验证,未纳入已检查事实");
return null;
}
return invocation;
} catch (RuntimeException exception) {
addLimitation(limitations, "部分已完成的工具记录暂时不可读取,未纳入已检查事实");
return null;
}
}
private void projectInvocation(CompletedToolCall completed,
CanonicalToolInvocation invocation,
Map<String, SafeFallback.VerifiedSource> sources,
Map<String, SafeFallback.ObservedFact> facts) {
JsonNode root = readObject(invocation.agentResult());
String scope = publicScope(completed.toolName(), completed.normalizedScope(), root);
switch (completed.toolName()) {
case AgentToolContracts.LOOKUP_KNOWLEDGE -> projectRag(root, scope, sources, facts);
case AgentToolContracts.QUERY_LOGS -> projectLogs(root, scope, sources, facts);
case AgentToolContracts.QUERY_MYSQL -> projectMysql(root, scope, sources, facts);
default -> throw new IllegalArgumentException("Unsupported evidence Tool");
}
}
private void projectRag(JsonNode root,
String scope,
Map<String, SafeFallback.VerifiedSource> sources,
Map<String, SafeFallback.ObservedFact> facts) {
JsonNode evidence = root.path("evidence");
if (!evidence.isArray() || evidence.isEmpty()) {
addFact(sources, facts, "RAG", "knowledge_base", scope,
"该知识检索范围内未发现可用文档证据");
return;
}
for (JsonNode item : evidence) {
String source = firstNonBlank(text(item, "source"), text(item, "title"),
text(item, "document_id"), "knowledge_base");
addFact(sources, facts, "RAG", source, scope,
firstNonBlank(text(item, "excerpt"), "已找到候选知识片段"));
}
}
private void projectLogs(JsonNode root,
String scope,
Map<String, SafeFallback.VerifiedSource> sources,
Map<String, SafeFallback.ObservedFact> facts) {
String source = firstNonBlank(text(root, "source_kind"), "logs");
JsonNode events = root.path("events");
if (!events.isArray() || events.isEmpty()) {
addFact(sources, facts, "LOG", source, scope,
"该日志查询范围内未发现匹配事件");
return;
}
for (JsonNode event : events) {
addFact(sources, facts, "LOG", source, scope,
firstNonBlank(text(event, "message"), "已找到匹配日志事件"));
}
}
private void projectMysql(JsonNode root,
String scope,
Map<String, SafeFallback.VerifiedSource> sources,
Map<String, SafeFallback.ObservedFact> facts) {
String source = mysqlSource(scope);
JsonNode rows = root.path("rows");
if (!rows.isArray() || rows.isEmpty()) {
addFact(sources, facts, "MYSQL", source, scope,
"该只读数据查询范围内未发现匹配记录");
return;
}
for (JsonNode row : rows) {
addFact(sources, facts, "MYSQL", source, scope, row.toString());
}
}
private String publicScope(String toolName, String normalizedScope, JsonNode result) {
if (AgentToolContracts.LOOKUP_KNOWLEDGE.equals(toolName)) {
return bounded("query=" + text(result, "query"), MAX_SCOPE_CHARS);
}
if (AgentToolContracts.QUERY_LOGS.equals(toolName)) {
JsonNode scope = result.path("scope");
return bounded(scope.isObject() ? scope.toString() : normalizedScope, MAX_SCOPE_CHARS);
}
try {
JsonNode scope = objectMapper.readTree(normalizedScope);
return bounded("data_source=" + text(scope, "data_source"), MAX_SCOPE_CHARS);
} catch (JsonProcessingException exception) {
return "data_source=unknown";
}
}
private String mysqlSource(String scope) {
int separator = scope.indexOf('=');
return separator < 0 ? "mysql" : scope.substring(separator + 1);
}
private void addFact(Map<String, SafeFallback.VerifiedSource> sources,
Map<String, SafeFallback.ObservedFact> facts,
String sourceType,
String source,
String scope,
String summary) {
if (facts.size() >= MAX_FACTS) {
return;
}
String safeSource = bounded(source, 160);
String safeScope = bounded(scope, MAX_SCOPE_CHARS);
String safeSummary = bounded(summary, MAX_SUMMARY_CHARS);
String sourceKey = sourceType + '\u0000' + safeSource + '\u0000' + safeScope;
sources.putIfAbsent(sourceKey,
new SafeFallback.VerifiedSource(sourceType, safeSource, safeScope));
String factKey = sourceKey + '\u0000' + safeSummary;
facts.putIfAbsent(factKey,
new SafeFallback.ObservedFact(sourceType, safeSource, safeScope, safeSummary));
}
private JsonNode readObject(String value) {
try {
JsonNode root = objectMapper.readTree(value);
if (root == null || !root.isObject()) {
throw new IllegalArgumentException("Canonical Agent result must be an object");
}
return root;
} catch (JsonProcessingException exception) {
throw new IllegalArgumentException("Canonical Agent result is invalid", exception);
}
}
private static void addLimitation(List<String> limitations, String value) {
if (!limitations.contains(value)) {
limitations.add(value);
}
}
private static String firstNonBlank(String... values) {
for (String value : values) {
if (value != null && !value.isBlank()) {
return value;
}
}
return "unknown";
}
private static String text(JsonNode node, String field) {
JsonNode value = node == null ? null : node.get(field);
return value == null || value.isNull() ? "" : value.asText("");
}
private static String bounded(String value, int max) {
String safe = value == null ? "" : value;
return safe.length() <= max ? safe : safe.substring(0, max);
}
}
@@ -0,0 +1,26 @@
package com.superbiz.agent.harness.progress;
import com.superbiz.agent.harness.contract.SafeFallback;
import java.util.List;
public record DiagnosisProgressSnapshot(
List<SafeFallback.VerifiedSource> verifiedSources,
List<SafeFallback.ObservedFact> observedFacts,
List<String> limitations,
DiagnosisStopReason stopReason) {
public DiagnosisProgressSnapshot {
verifiedSources = verifiedSources == null ? List.of() : List.copyOf(verifiedSources);
observedFacts = observedFacts == null ? List.of() : List.copyOf(observedFacts);
limitations = limitations == null ? List.of() : List.copyOf(limitations);
}
public static DiagnosisProgressSnapshot empty() {
return new DiagnosisProgressSnapshot(List.of(), List.of(), List.of(), null);
}
public boolean hasObservedFacts() {
return !observedFacts.isEmpty();
}
}
@@ -0,0 +1,16 @@
package com.superbiz.agent.harness.progress;
import java.util.List;
public record DiagnosisProgressSnapshotState(
int consecutiveNoGain,
DiagnosisCollectionState collectionState,
DiagnosisStopReason stopReason,
String pendingToolCallId,
boolean stopInstructionDelivered,
List<CompletedToolCall> completedToolCalls) {
public DiagnosisProgressSnapshotState {
completedToolCalls = completedToolCalls == null ? List.of() : List.copyOf(completedToolCalls);
}
}
@@ -0,0 +1,118 @@
package com.superbiz.agent.harness.progress;
import com.superbiz.agent.harness.contract.EvidenceStatus;
import java.util.ArrayList;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Set;
public final class DiagnosisProgressTracker {
private final int stopAfterConsecutiveNoGain;
private final Set<ToolScopeIdentity> completedScopes = new LinkedHashSet<>();
private final List<CompletedToolCall> completedToolCalls = new ArrayList<>();
private int consecutiveNoGain;
private DiagnosisCollectionState collectionState = DiagnosisCollectionState.COLLECTING;
private DiagnosisStopReason stopReason;
private String pendingToolCallId;
private boolean stopInstructionDelivered;
public DiagnosisProgressTracker(int stopAfterConsecutiveNoGain) {
if (stopAfterConsecutiveNoGain <= 0) {
throw new IllegalArgumentException("stopAfterConsecutiveNoGain must be positive");
}
this.stopAfterConsecutiveNoGain = stopAfterConsecutiveNoGain;
}
public synchronized void applyPreviousObservation(PreviousObservation observation) {
if (pendingToolCallId == null) {
if (observation != null) {
throw new IllegalArgumentException("No Tool observation is pending evaluation");
}
return;
}
if (observation == null) {
throw new IllegalArgumentException("Previous Tool observation must be evaluated");
}
if (!pendingToolCallId.equals(observation.toolCallId())) {
throw new IllegalArgumentException("Previous Tool observation ID is out of order");
}
pendingToolCallId = null;
applyGain(observation.informationGain());
}
public synchronized boolean isDuplicate(String toolName, String normalizedScope) {
return completedScopes.contains(new ToolScopeIdentity(toolName, normalizedScope));
}
public synchronized void recordDuplicateScope() {
applyGain(InformationGain.NO_GAIN);
}
public synchronized void recordCompleted(CompletedToolCall call, EvidenceStatus evidenceStatus) {
if (collectionState == DiagnosisCollectionState.SATURATED) {
throw new IllegalStateException("Cannot record Tool completion after saturation");
}
if (evidenceStatus != EvidenceStatus.EVIDENCE_FOUND
&& evidenceStatus != EvidenceStatus.NO_EVIDENCE) {
throw new IllegalArgumentException("Completed Tool requires a successful evidence status");
}
ToolScopeIdentity scope = new ToolScopeIdentity(call.toolName(), call.normalizedScope());
if (!completedScopes.add(scope)) {
throw new IllegalStateException("Completed Tool scope was already recorded");
}
completedToolCalls.add(call);
if (evidenceStatus == EvidenceStatus.NO_EVIDENCE) {
applyGain(InformationGain.NO_GAIN);
} else {
pendingToolCallId = call.toolCallId();
}
}
public synchronized boolean claimStopInstruction() {
if (collectionState != DiagnosisCollectionState.SATURATED) {
return false;
}
if (stopInstructionDelivered) {
return false;
}
stopInstructionDelivered = true;
return true;
}
public synchronized void markBudgetLimitReached() {
if (stopReason == null) {
stopReason = DiagnosisStopReason.BUDGET_LIMIT_REACHED;
}
}
public synchronized DiagnosisProgressSnapshotState snapshot() {
return new DiagnosisProgressSnapshotState(
consecutiveNoGain,
collectionState,
stopReason,
pendingToolCallId,
stopInstructionDelivered,
completedToolCalls);
}
public int stopAfterConsecutiveNoGain() {
return stopAfterConsecutiveNoGain;
}
private void applyGain(InformationGain gain) {
if (collectionState == DiagnosisCollectionState.SATURATED) {
throw new IllegalStateException("Collection is already saturated");
}
if (gain == InformationGain.GAINED) {
consecutiveNoGain = 0;
return;
}
consecutiveNoGain++;
if (consecutiveNoGain >= stopAfterConsecutiveNoGain) {
collectionState = DiagnosisCollectionState.SATURATED;
stopReason = DiagnosisStopReason.INFORMATION_SATURATED;
}
}
}
@@ -0,0 +1,6 @@
package com.superbiz.agent.harness.progress;
public enum DiagnosisStopReason {
INFORMATION_SATURATED,
BUDGET_LIMIT_REACHED
}
@@ -0,0 +1,6 @@
package com.superbiz.agent.harness.progress;
public enum InformationGain {
GAINED,
NO_GAIN
}
@@ -0,0 +1,17 @@
package com.superbiz.agent.harness.progress;
import com.fasterxml.jackson.annotation.JsonProperty;
public record PreviousObservation(
@JsonProperty("tool_call_id") String toolCallId,
@JsonProperty("information_gain") InformationGain informationGain) {
public PreviousObservation {
if (toolCallId == null || toolCallId.isBlank()) {
throw new IllegalArgumentException("toolCallId must not be blank");
}
if (informationGain == null) {
throw new IllegalArgumentException("informationGain must not be null");
}
}
}
@@ -0,0 +1,13 @@
package com.superbiz.agent.harness.progress;
public record ToolScopeIdentity(String toolName, String normalizedScope) {
public ToolScopeIdentity {
if (toolName == null || toolName.isBlank()) {
throw new IllegalArgumentException("toolName must not be blank");
}
if (normalizedScope == null || normalizedScope.isBlank()) {
throw new IllegalArgumentException("normalizedScope must not be blank");
}
}
}
@@ -0,0 +1,76 @@
package com.superbiz.agent.harness.progress;
import com.fasterxml.jackson.core.JsonProcessingException;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.harness.tool.contract.AgentToolContracts;
import com.superbiz.agent.harness.tool.contract.MysqlToolRequest;
import com.superbiz.agent.harness.tool.contract.QueryLogsRequest;
import com.superbiz.agent.harness.tool.contract.RagToolRequest;
import java.util.LinkedHashMap;
import java.util.Map;
import java.util.Objects;
public final class ToolScopeNormalizer {
private static final int DEFAULT_LOG_LOOKBACK_MINUTES = 30;
private final ObjectMapper objectMapper;
public ToolScopeNormalizer(ObjectMapper objectMapper) {
this.objectMapper = Objects.requireNonNull(objectMapper, "objectMapper must not be null");
}
public String normalize(String toolName, Object input) {
Objects.requireNonNull(input, "input must not be null");
Map<String, Object> scope = switch (toolName) {
case AgentToolContracts.LOOKUP_KNOWLEDGE -> ragScope(requireType(input, RagToolRequest.class));
case AgentToolContracts.QUERY_LOGS -> logScope(requireType(input, QueryLogsRequest.class));
case AgentToolContracts.QUERY_MYSQL -> mysqlScope(requireType(input, MysqlToolRequest.class));
default -> throw new IllegalArgumentException("Unsupported evidence Tool: " + toolName);
};
try {
return objectMapper.writeValueAsString(scope);
} catch (JsonProcessingException exception) {
throw new IllegalArgumentException("Tool scope is not serializable", exception);
}
}
private static Map<String, Object> ragScope(RagToolRequest request) {
return ordered("query", normalizeText(request.query()));
}
private static Map<String, Object> logScope(QueryLogsRequest request) {
Map<String, Object> scope = new LinkedHashMap<>();
scope.put("topic", request.topic() == null ? null : request.topic().name());
scope.put("query", normalizeText(request.query()));
scope.put("lookback_minutes", request.lookbackMinutes() == null
? DEFAULT_LOG_LOOKBACK_MINUTES : request.lookbackMinutes());
return scope;
}
private static Map<String, Object> mysqlScope(MysqlToolRequest request) {
Map<String, Object> scope = new LinkedHashMap<>();
scope.put("data_source", normalizeText(request.dataSource()));
scope.put("sql", normalizeText(request.sql()));
scope.put("params", request.params());
return scope;
}
private static Map<String, Object> ordered(String name, Object value) {
Map<String, Object> result = new LinkedHashMap<>();
result.put(name, value);
return result;
}
private static String normalizeText(String value) {
return value == null ? "" : value.trim();
}
private static <T> T requireType(Object input, Class<T> type) {
if (!type.isInstance(input)) {
throw new IllegalArgumentException("Unexpected Tool input type: " + input.getClass().getName());
}
return type.cast(input);
}
}