Add diagnosis eval harness

This commit is contained in:
aruo
2026-07-04 23:51:43 +08:00
parent dc6cd32a67
commit ca5c61fabf
24 changed files with 1251 additions and 2 deletions
@@ -0,0 +1,25 @@
package com.superbiz.agent.eval;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.util.List;
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class DiagnosisEvalCase {
private String id;
private String title;
private String question;
private String traceFixture;
private List<String> expectedRootCauseKeywords;
private Integer minKeywordMatches;
private List<String> requiredEvidenceTools;
private List<String> allowedVerdicts;
private List<String> forbiddenAnswerKeywords;
}
@@ -0,0 +1,24 @@
package com.superbiz.agent.eval;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.util.List;
import java.util.Map;
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class DiagnosisEvalReport {
private int totalCases;
private int passedCases;
private double passRate;
private Map<String, Long> verdictDistribution;
private double averageToolCallCount;
private double averageDurationMs;
private List<DiagnosisEvalResult> results;
}
@@ -0,0 +1,76 @@
package com.superbiz.agent.eval;
import com.fasterxml.jackson.databind.ObjectMapper;
import java.io.IOException;
import java.nio.charset.StandardCharsets;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.Map;
public class DiagnosisEvalReportWriter {
private final ObjectMapper objectMapper;
public DiagnosisEvalReportWriter(ObjectMapper objectMapper) {
this.objectMapper = objectMapper;
}
public void writeJson(DiagnosisEvalReport report, Path outputFile) throws IOException {
Files.createDirectories(outputFile.getParent());
objectMapper.writerWithDefaultPrettyPrinter().writeValue(outputFile.toFile(), report);
}
public void writeMarkdown(DiagnosisEvalReport report, Path outputFile) throws IOException {
Files.createDirectories(outputFile.getParent());
Files.writeString(outputFile, toMarkdown(report), StandardCharsets.UTF_8);
}
public String toMarkdown(DiagnosisEvalReport report) {
StringBuilder builder = new StringBuilder();
builder.append("# Diagnosis Eval Report\n\n");
builder.append("- Total cases: ").append(report.getTotalCases()).append("\n");
builder.append("- Passed cases: ").append(report.getPassedCases()).append("\n");
builder.append("- Pass rate: ").append(String.format("%.2f%%", report.getPassRate() * 100)).append("\n");
builder.append("- Average tool calls: ").append(String.format("%.2f", report.getAverageToolCallCount())).append("\n");
builder.append("- Average duration ms: ").append(String.format("%.2f", report.getAverageDurationMs())).append("\n\n");
builder.append("## Verdict Distribution\n\n");
if (report.getVerdictDistribution() == null || report.getVerdictDistribution().isEmpty()) {
builder.append("- None\n\n");
} else {
for (Map.Entry<String, Long> entry : report.getVerdictDistribution().entrySet()) {
builder.append("- ").append(entry.getKey()).append(": ").append(entry.getValue()).append("\n");
}
builder.append("\n");
}
builder.append("## Cases\n\n");
builder.append("| Case | Result | Verdict | Keywords | Tool Calls | Duration ms | Failed Checks |\n");
builder.append("| --- | --- | --- | --- | ---: | ---: | --- |\n");
for (DiagnosisEvalResult result : report.getResults()) {
builder.append("| ")
.append(result.getCaseId())
.append(" | ")
.append(result.isPassed() ? "PASS" : "FAIL")
.append(" | ")
.append(valueOrDash(result.getVerdict()))
.append(" | ")
.append(result.getMatchedKeywordCount()).append("/").append(result.getRequiredKeywordCount())
.append(" | ")
.append(result.getToolCallCount() == null ? "-" : result.getToolCallCount())
.append(" | ")
.append(result.getDurationMs() == null ? "-" : result.getDurationMs())
.append(" | ")
.append(result.getFailedChecks() == null || result.getFailedChecks().isEmpty()
? "-"
: String.join("; ", result.getFailedChecks()))
.append(" |\n");
}
return builder.toString();
}
private String valueOrDash(String value) {
return value == null || value.isBlank() ? "-" : value;
}
}
@@ -0,0 +1,27 @@
package com.superbiz.agent.eval;
import lombok.AllArgsConstructor;
import lombok.Builder;
import lombok.Data;
import lombok.NoArgsConstructor;
import java.util.List;
import java.util.Map;
@Data
@Builder
@NoArgsConstructor
@AllArgsConstructor
public class DiagnosisEvalResult {
private String caseId;
private String title;
private boolean passed;
private List<String> failedChecks;
private String verdict;
private int matchedKeywordCount;
private int requiredKeywordCount;
private Map<String, Boolean> evidenceCoverage;
private Integer toolCallCount;
private Integer durationMs;
}
@@ -0,0 +1,217 @@
package com.superbiz.agent.eval;
import com.fasterxml.jackson.core.type.TypeReference;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.dto.DiagnosisTraceResponse;
import java.io.IOException;
import java.nio.file.Path;
import java.util.ArrayList;
import java.util.LinkedHashMap;
import java.util.LinkedHashSet;
import java.util.List;
import java.util.Locale;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.stream.Collectors;
public class DiagnosisTraceEvaluator {
private static final TypeReference<List<DiagnosisEvalCase>> CASE_LIST_TYPE = new TypeReference<>() {};
private static final String REJECT_DEGRADED_PREFIX = "当前无法基于已获取证据生成可靠结论";
private final ObjectMapper objectMapper;
public DiagnosisTraceEvaluator(ObjectMapper objectMapper) {
this.objectMapper = objectMapper;
}
public List<DiagnosisEvalCase> loadCases(Path casesFile) throws IOException {
return objectMapper.readValue(casesFile.toFile(), CASE_LIST_TYPE);
}
public DiagnosisTraceResponse loadTrace(Path traceFile) throws IOException {
return objectMapper.readValue(traceFile.toFile(), DiagnosisTraceResponse.class);
}
public DiagnosisEvalReport evaluate(List<DiagnosisEvalCase> cases, Path fixtureDir) {
List<DiagnosisEvalResult> results = new ArrayList<>();
for (DiagnosisEvalCase evalCase : cases) {
try {
DiagnosisTraceResponse trace = loadTrace(fixtureDir.resolve(evalCase.getTraceFixture()));
results.add(evaluate(evalCase, trace));
} catch (Exception e) {
results.add(DiagnosisEvalResult.builder()
.caseId(evalCase.getId())
.title(evalCase.getTitle())
.passed(false)
.failedChecks(List.of("trace fixture unavailable: " + e.getMessage()))
.verdict(null)
.matchedKeywordCount(0)
.requiredKeywordCount(size(evalCase.getExpectedRootCauseKeywords()))
.evidenceCoverage(emptyCoverage(evalCase.getRequiredEvidenceTools()))
.toolCallCount(null)
.durationMs(null)
.build());
}
}
return toReport(results);
}
public DiagnosisEvalResult evaluate(DiagnosisEvalCase evalCase, DiagnosisTraceResponse trace) {
List<String> failedChecks = new ArrayList<>();
String answer = trace.getSession() == null ? "" : nullToEmpty(trace.getSession().getAnswer());
String normalizedAnswer = answer.toLowerCase(Locale.ROOT);
int requiredKeywordCount = size(evalCase.getExpectedRootCauseKeywords());
int matchedKeywordCount = countMatches(normalizedAnswer, evalCase.getExpectedRootCauseKeywords());
int minKeywordMatches = evalCase.getMinKeywordMatches() == null
? requiredKeywordCount
: evalCase.getMinKeywordMatches();
if (matchedKeywordCount < minKeywordMatches) {
failedChecks.add("answer keyword coverage too low: " + matchedKeywordCount + "/" + minKeywordMatches);
}
for (String forbidden : safeList(evalCase.getForbiddenAnswerKeywords())) {
if (normalizedAnswer.contains(forbidden.toLowerCase(Locale.ROOT))) {
failedChecks.add("answer contains forbidden keyword: " + forbidden);
}
}
Set<String> evidenceTools = collectEvidenceTools(trace);
Map<String, Boolean> evidenceCoverage = new LinkedHashMap<>();
for (String requiredTool : safeList(evalCase.getRequiredEvidenceTools())) {
boolean present = evidenceTools.contains(requiredTool);
evidenceCoverage.put(requiredTool, present);
if (!present) {
failedChecks.add("missing required evidence tool: " + requiredTool);
}
}
String verdict = extractVerifierVerdict(trace);
if (verdict == null || verdict.isBlank()) {
failedChecks.add("missing verifier verdict");
} else if (!safeList(evalCase.getAllowedVerdicts()).isEmpty()
&& !safeList(evalCase.getAllowedVerdicts()).contains(verdict)) {
failedChecks.add("verdict not allowed: " + verdict);
}
if ("REJECT".equals(verdict) && !answer.startsWith(REJECT_DEGRADED_PREFIX)) {
failedChecks.add("reject output does not use degraded template");
}
Integer toolCallCount = trace.getToolInvocations() == null ? 0 : trace.getToolInvocations().size();
Integer durationMs = trace.getSession() == null ? null : trace.getSession().getTotalDurationMs();
return DiagnosisEvalResult.builder()
.caseId(evalCase.getId())
.title(evalCase.getTitle())
.passed(failedChecks.isEmpty())
.failedChecks(failedChecks)
.verdict(verdict)
.matchedKeywordCount(matchedKeywordCount)
.requiredKeywordCount(requiredKeywordCount)
.evidenceCoverage(evidenceCoverage)
.toolCallCount(toolCallCount)
.durationMs(durationMs)
.build();
}
private DiagnosisEvalReport toReport(List<DiagnosisEvalResult> results) {
int total = results.size();
int passed = (int) results.stream().filter(DiagnosisEvalResult::isPassed).count();
Map<String, Long> verdictDistribution = results.stream()
.map(DiagnosisEvalResult::getVerdict)
.filter(Objects::nonNull)
.collect(Collectors.groupingBy(value -> value, LinkedHashMap::new, Collectors.counting()));
double averageToolCallCount = results.stream()
.map(DiagnosisEvalResult::getToolCallCount)
.filter(Objects::nonNull)
.mapToInt(Integer::intValue)
.average()
.orElse(0.0);
double averageDurationMs = results.stream()
.map(DiagnosisEvalResult::getDurationMs)
.filter(Objects::nonNull)
.mapToInt(Integer::intValue)
.average()
.orElse(0.0);
return DiagnosisEvalReport.builder()
.totalCases(total)
.passedCases(passed)
.passRate(total == 0 ? 0.0 : (double) passed / total)
.verdictDistribution(verdictDistribution)
.averageToolCallCount(averageToolCallCount)
.averageDurationMs(averageDurationMs)
.results(results)
.build();
}
private Set<String> collectEvidenceTools(DiagnosisTraceResponse trace) {
Set<String> tools = new LinkedHashSet<>();
if (trace.getToolInvocations() != null) {
for (DiagnosisTraceResponse.ToolInvocationTrace invocation : trace.getToolInvocations()) {
if (invocation.getToolName() != null) {
tools.add(invocation.getToolName());
}
}
}
Object summaries = nestedValue(trace, "verifier_evaluation", "tool_trace_summary");
if (summaries instanceof List<?> list) {
for (Object item : list) {
if (item instanceof Map<?, ?> map && map.get("tool_name") != null) {
tools.add(String.valueOf(map.get("tool_name")));
}
}
}
return tools;
}
private String extractVerifierVerdict(DiagnosisTraceResponse trace) {
Object value = nestedValue(trace, "verifier_evaluation", "verdict");
return value == null ? null : String.valueOf(value);
}
private Object nestedValue(DiagnosisTraceResponse trace, String firstKey, String secondKey) {
if (trace.getSession() == null || trace.getSession().getSelfEvaluation() == null) {
return null;
}
Object first = trace.getSession().getSelfEvaluation().get(firstKey);
if (!(first instanceof Map<?, ?> map)) {
return null;
}
return map.get(secondKey);
}
private int countMatches(String normalizedAnswer, List<String> keywords) {
int count = 0;
for (String keyword : safeList(keywords)) {
if (normalizedAnswer.contains(keyword.toLowerCase(Locale.ROOT))) {
count++;
}
}
return count;
}
private Map<String, Boolean> emptyCoverage(List<String> tools) {
Map<String, Boolean> coverage = new LinkedHashMap<>();
for (String tool : safeList(tools)) {
coverage.put(tool, false);
}
return coverage;
}
private List<String> safeList(List<String> values) {
return values == null ? List.of() : values;
}
private int size(List<?> values) {
return values == null ? 0 : values.size();
}
private String nullToEmpty(String value) {
return value == null ? "" : value;
}
}
@@ -0,0 +1,97 @@
package com.superbiz.agent.eval;
import com.fasterxml.jackson.databind.ObjectMapper;
import com.superbiz.agent.dto.DiagnosisTraceResponse;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.io.TempDir;
import java.nio.file.Files;
import java.nio.file.Path;
import java.util.List;
import static org.junit.jupiter.api.Assertions.assertEquals;
import static org.junit.jupiter.api.Assertions.assertFalse;
import static org.junit.jupiter.api.Assertions.assertTrue;
class DiagnosisTraceEvaluatorTest {
private final ObjectMapper objectMapper = new ObjectMapper();
private final DiagnosisTraceEvaluator evaluator = new DiagnosisTraceEvaluator(objectMapper);
@Test
void evaluateFixtureReportsPassingAndMissingCases() {
List<DiagnosisEvalCase> cases = readCases();
DiagnosisEvalReport report = evaluator.evaluate(cases, Path.of("mvp/eval/fixtures"));
assertEquals(5, report.getTotalCases());
assertEquals(2, report.getPassedCases());
assertEquals(0.4, report.getPassRate(), 0.001);
assertEquals(1L, report.getVerdictDistribution().get("PASS"));
assertEquals(1L, report.getVerdictDistribution().get("LOW_CONFID"));
DiagnosisEvalResult payment = result(report, "payment-timeout");
assertTrue(payment.isPassed());
assertTrue(payment.getEvidenceCoverage().get("lookup_knowledge"));
assertTrue(payment.getEvidenceCoverage().get("query_logs"));
assertTrue(payment.getEvidenceCoverage().get("query_metrics"));
DiagnosisEvalResult missing = result(report, "redis-timeout");
assertFalse(missing.isPassed());
assertTrue(missing.getFailedChecks().get(0).contains("trace fixture unavailable"));
}
@Test
void evaluateRejectRequiresDegradedOutput() {
DiagnosisEvalCase evalCase = DiagnosisEvalCase.builder()
.id("reject-case")
.title("Reject case")
.expectedRootCauseKeywords(List.of())
.requiredEvidenceTools(List.of())
.allowedVerdicts(List.of("REJECT"))
.build();
DiagnosisTraceResponse trace = DiagnosisTraceResponse.builder()
.session(DiagnosisTraceResponse.SessionTrace.builder()
.answer("EXECUTOR_FINAL_ANSWER")
.selfEvaluation(java.util.Map.of(
"verifier_evaluation", java.util.Map.of("verdict", "REJECT")))
.build())
.toolInvocations(List.of())
.build();
DiagnosisEvalResult result = evaluator.evaluate(evalCase, trace);
assertFalse(result.isPassed());
assertTrue(result.getFailedChecks().contains("reject output does not use degraded template"));
}
@Test
void reportWriterOutputsJsonAndMarkdown(@TempDir Path tempDir) throws Exception {
DiagnosisEvalReport report = evaluator.evaluate(readCases(), Path.of("mvp/eval/fixtures"));
DiagnosisEvalReportWriter writer = new DiagnosisEvalReportWriter(objectMapper);
Path json = tempDir.resolve("eval-report.json");
Path markdown = tempDir.resolve("eval-report.md");
writer.writeJson(report, json);
writer.writeMarkdown(report, markdown);
assertTrue(Files.exists(json));
assertTrue(Files.readString(markdown).contains("# Diagnosis Eval Report"));
assertTrue(Files.readString(markdown).contains("payment-timeout"));
}
private List<DiagnosisEvalCase> readCases() {
try {
return evaluator.loadCases(Path.of("mvp/eval/cases/diagnosis-cases.json"));
} catch (Exception e) {
throw new AssertionError(e);
}
}
private DiagnosisEvalResult result(DiagnosisEvalReport report, String caseId) {
return report.getResults().stream()
.filter(item -> caseId.equals(item.getCaseId()))
.findFirst()
.orElseThrow();
}
}