253 lines
11 KiB
Java
253 lines
11 KiB
Java
package com.superbiz.agent.eval;
|
|
|
|
import java.util.ArrayList;
|
|
import java.util.Comparator;
|
|
import java.util.LinkedHashMap;
|
|
import java.util.LinkedHashSet;
|
|
import java.util.List;
|
|
import java.util.Map;
|
|
import java.util.Objects;
|
|
import java.util.Set;
|
|
import java.util.function.Function;
|
|
import java.util.stream.Collectors;
|
|
|
|
public class DiagnosisEvalBaselineDiffer {
|
|
|
|
private static final String REGRESSION = "REGRESSION";
|
|
private static final String IMPROVEMENT = "IMPROVEMENT";
|
|
private static final String CHANGED = "CHANGED";
|
|
|
|
public DiagnosisEvalDiffReport compare(DiagnosisEvalReport baseline, DiagnosisEvalReport current) {
|
|
List<DiagnosisEvalDiffItem> items = new ArrayList<>();
|
|
|
|
compareDouble(items, "aggregate", null, "passRate",
|
|
baseline.getPassRate(), current.getPassRate(), true);
|
|
compareDouble(items, "aggregate", null, "averageToolCallCount",
|
|
baseline.getAverageToolCallCount(), current.getAverageToolCallCount(), false);
|
|
compareDouble(items, "aggregate", null, "averageDurationMs",
|
|
baseline.getAverageDurationMs(), current.getAverageDurationMs(), false);
|
|
compareVerdictDistribution(items, baseline.getVerdictDistribution(), current.getVerdictDistribution());
|
|
compareCases(items, safeResults(baseline), safeResults(current));
|
|
|
|
int regressionCount = countType(items, REGRESSION);
|
|
int improvementCount = countType(items, IMPROVEMENT);
|
|
int changedCount = countType(items, CHANGED);
|
|
|
|
return DiagnosisEvalDiffReport.builder()
|
|
.baselineTotalCases(baseline.getTotalCases())
|
|
.currentTotalCases(current.getTotalCases())
|
|
.baselinePassedCases(baseline.getPassedCases())
|
|
.currentPassedCases(current.getPassedCases())
|
|
.baselinePassRate(baseline.getPassRate())
|
|
.currentPassRate(current.getPassRate())
|
|
.regressionCount(regressionCount)
|
|
.improvementCount(improvementCount)
|
|
.changedCount(changedCount)
|
|
.hasRegression(regressionCount > 0)
|
|
.items(items)
|
|
.build();
|
|
}
|
|
|
|
private void compareVerdictDistribution(List<DiagnosisEvalDiffItem> items,
|
|
Map<String, Long> baseline,
|
|
Map<String, Long> current) {
|
|
Set<String> verdicts = new LinkedHashSet<>();
|
|
verdicts.addAll(safeMap(baseline).keySet());
|
|
verdicts.addAll(safeMap(current).keySet());
|
|
for (String verdict : verdicts) {
|
|
long baselineCount = safeMap(baseline).getOrDefault(verdict, 0L);
|
|
long currentCount = safeMap(current).getOrDefault(verdict, 0L);
|
|
if (baselineCount != currentCount) {
|
|
items.add(item(CHANGED, "aggregate", null, "verdictDistribution." + verdict,
|
|
String.valueOf(baselineCount), String.valueOf(currentCount),
|
|
(double) currentCount - baselineCount,
|
|
"verdict count changed for " + verdict));
|
|
}
|
|
}
|
|
}
|
|
|
|
private void compareCases(List<DiagnosisEvalDiffItem> items,
|
|
List<DiagnosisEvalResult> baselineResults,
|
|
List<DiagnosisEvalResult> currentResults) {
|
|
Map<String, DiagnosisEvalResult> baselineById = byCaseId(baselineResults);
|
|
Map<String, DiagnosisEvalResult> currentById = byCaseId(currentResults);
|
|
Set<String> caseIds = new LinkedHashSet<>();
|
|
caseIds.addAll(baselineById.keySet());
|
|
caseIds.addAll(currentById.keySet());
|
|
|
|
for (String caseId : caseIds) {
|
|
DiagnosisEvalResult baseline = baselineById.get(caseId);
|
|
DiagnosisEvalResult current = currentById.get(caseId);
|
|
if (baseline == null) {
|
|
items.add(item(CHANGED, "case", caseId, "casePresence",
|
|
"missing", "present", null, "new case appears in current report"));
|
|
continue;
|
|
}
|
|
if (current == null) {
|
|
items.add(item(REGRESSION, "case", caseId, "casePresence",
|
|
"present", "missing", null, "baseline case is missing from current report"));
|
|
continue;
|
|
}
|
|
|
|
comparePassState(items, baseline, current);
|
|
compareVerdict(items, baseline, current);
|
|
compareInteger(items, caseId, "matchedKeywordCount",
|
|
baseline.getMatchedKeywordCount(), current.getMatchedKeywordCount(), true);
|
|
compareInteger(items, caseId, "toolCallCount",
|
|
baseline.getToolCallCount(), current.getToolCallCount(), false);
|
|
compareInteger(items, caseId, "durationMs",
|
|
baseline.getDurationMs(), current.getDurationMs(), false);
|
|
compareEvidenceCoverage(items, baseline, current);
|
|
}
|
|
}
|
|
|
|
private void comparePassState(List<DiagnosisEvalDiffItem> items,
|
|
DiagnosisEvalResult baseline,
|
|
DiagnosisEvalResult current) {
|
|
if (baseline.isPassed() == current.isPassed()) {
|
|
return;
|
|
}
|
|
String type = baseline.isPassed() ? REGRESSION : IMPROVEMENT;
|
|
items.add(item(type, "case", baseline.getCaseId(), "passed",
|
|
String.valueOf(baseline.isPassed()), String.valueOf(current.isPassed()), null,
|
|
baseline.getCaseId() + " pass state changed"));
|
|
}
|
|
|
|
private void compareVerdict(List<DiagnosisEvalDiffItem> items,
|
|
DiagnosisEvalResult baseline,
|
|
DiagnosisEvalResult current) {
|
|
if (Objects.equals(baseline.getVerdict(), current.getVerdict())) {
|
|
return;
|
|
}
|
|
int baselineRank = verdictRank(baseline.getVerdict());
|
|
int currentRank = verdictRank(current.getVerdict());
|
|
String type = currentRank < baselineRank ? REGRESSION : currentRank > baselineRank ? IMPROVEMENT : CHANGED;
|
|
items.add(item(type, "case", baseline.getCaseId(), "verdict",
|
|
value(baseline.getVerdict()), value(current.getVerdict()), (double) currentRank - baselineRank,
|
|
baseline.getCaseId() + " verdict changed"));
|
|
}
|
|
|
|
private void compareEvidenceCoverage(List<DiagnosisEvalDiffItem> items,
|
|
DiagnosisEvalResult baseline,
|
|
DiagnosisEvalResult current) {
|
|
Set<String> tools = new LinkedHashSet<>();
|
|
tools.addAll(safeMap(baseline.getEvidenceCoverage()).keySet());
|
|
tools.addAll(safeMap(current.getEvidenceCoverage()).keySet());
|
|
for (String tool : tools) {
|
|
boolean baselineCovered = Boolean.TRUE.equals(safeMap(baseline.getEvidenceCoverage()).get(tool));
|
|
boolean currentCovered = Boolean.TRUE.equals(safeMap(current.getEvidenceCoverage()).get(tool));
|
|
if (baselineCovered == currentCovered) {
|
|
continue;
|
|
}
|
|
String type = baselineCovered ? REGRESSION : IMPROVEMENT;
|
|
items.add(item(type, "case", baseline.getCaseId(), "evidenceCoverage." + tool,
|
|
String.valueOf(baselineCovered), String.valueOf(currentCovered), null,
|
|
baseline.getCaseId() + " evidence coverage changed for " + tool));
|
|
}
|
|
}
|
|
|
|
private void compareDouble(List<DiagnosisEvalDiffItem> items,
|
|
String scope,
|
|
String caseId,
|
|
String metric,
|
|
double baseline,
|
|
double current,
|
|
boolean higherIsBetter) {
|
|
if (Double.compare(baseline, current) == 0) {
|
|
return;
|
|
}
|
|
double delta = current - baseline;
|
|
String type = classifyDelta(delta, higherIsBetter);
|
|
items.add(item(type, scope, caseId, metric,
|
|
String.valueOf(baseline), String.valueOf(current), delta,
|
|
metric + " changed"));
|
|
}
|
|
|
|
private void compareInteger(List<DiagnosisEvalDiffItem> items,
|
|
String caseId,
|
|
String metric,
|
|
Integer baseline,
|
|
Integer current,
|
|
boolean higherIsBetter) {
|
|
if (Objects.equals(baseline, current)) {
|
|
return;
|
|
}
|
|
if (baseline == null || current == null) {
|
|
items.add(item(CHANGED, "case", caseId, metric,
|
|
value(baseline), value(current), null, caseId + " " + metric + " changed"));
|
|
return;
|
|
}
|
|
int delta = current - baseline;
|
|
items.add(item(classifyDelta(delta, higherIsBetter), "case", caseId, metric,
|
|
String.valueOf(baseline), String.valueOf(current), (double) delta,
|
|
caseId + " " + metric + " changed"));
|
|
}
|
|
|
|
private String classifyDelta(double delta, boolean higherIsBetter) {
|
|
if (delta == 0.0) {
|
|
return CHANGED;
|
|
}
|
|
boolean improved = higherIsBetter ? delta > 0 : delta < 0;
|
|
return improved ? IMPROVEMENT : REGRESSION;
|
|
}
|
|
|
|
private DiagnosisEvalDiffItem item(String type,
|
|
String scope,
|
|
String caseId,
|
|
String metric,
|
|
String baselineValue,
|
|
String currentValue,
|
|
Double delta,
|
|
String message) {
|
|
return DiagnosisEvalDiffItem.builder()
|
|
.type(type)
|
|
.scope(scope)
|
|
.caseId(caseId)
|
|
.metric(metric)
|
|
.baselineValue(baselineValue)
|
|
.currentValue(currentValue)
|
|
.delta(delta)
|
|
.message(message)
|
|
.build();
|
|
}
|
|
|
|
private Map<String, DiagnosisEvalResult> byCaseId(List<DiagnosisEvalResult> results) {
|
|
return results.stream()
|
|
.sorted(Comparator.comparing(DiagnosisEvalResult::getCaseId))
|
|
.collect(Collectors.toMap(
|
|
DiagnosisEvalResult::getCaseId,
|
|
Function.identity(),
|
|
(left, right) -> right,
|
|
LinkedHashMap::new));
|
|
}
|
|
|
|
private List<DiagnosisEvalResult> safeResults(DiagnosisEvalReport report) {
|
|
return report.getResults() == null ? List.of() : report.getResults();
|
|
}
|
|
|
|
private <T> Map<String, T> safeMap(Map<String, T> value) {
|
|
return value == null ? Map.of() : value;
|
|
}
|
|
|
|
private int countType(List<DiagnosisEvalDiffItem> items, String type) {
|
|
return (int) items.stream().filter(item -> type.equals(item.getType())).count();
|
|
}
|
|
|
|
private int verdictRank(String verdict) {
|
|
if ("PASS".equals(verdict)) {
|
|
return 3;
|
|
}
|
|
if ("LOW_CONFID".equals(verdict)) {
|
|
return 2;
|
|
}
|
|
if ("REJECT".equals(verdict)) {
|
|
return 1;
|
|
}
|
|
return 0;
|
|
}
|
|
|
|
private String value(Object value) {
|
|
return value == null ? "-" : String.valueOf(value);
|
|
}
|
|
}
|