Add diagnosis eval harness
This commit is contained in:
@@ -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;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user