From 286a07704aaa1cd7d78bce18324bb5b93c491fd0 Mon Sep 17 00:00:00 2001 From: Pratyush Sharma <56130065+pratyush618@users.noreply.github.com> Date: Fri, 13 Mar 2026 11:32:44 +0530 Subject: [PATCH 1/4] Add snapshot testing to agenteval-reporting Persist EvalResult as named JSON snapshots and compare against prior baselines to detect regressions. SnapshotStore handles save/load with path traversal validation, SnapshotReporter implements EvalReporter with baseline/compare/update modes, SnapshotRegressionException carries the RegressionReport on failure. --- agenteval-reporting/pom.xml | 4 + .../reporting/snapshot/SnapshotCaseData.java | 48 ++++++ .../snapshot/SnapshotComparisonResult.java | 23 +++ .../reporting/snapshot/SnapshotConfig.java | 85 +++++++++++ .../reporting/snapshot/SnapshotData.java | 96 ++++++++++++ .../snapshot/SnapshotRegressionException.java | 29 ++++ .../reporting/snapshot/SnapshotReporter.java | 106 +++++++++++++ .../reporting/snapshot/SnapshotScoreData.java | 18 +++ .../reporting/snapshot/SnapshotStatus.java | 15 ++ .../reporting/snapshot/SnapshotStore.java | 121 +++++++++++++++ .../snapshot/SnapshotConfigTest.java | 62 ++++++++ .../reporting/snapshot/SnapshotDataTest.java | 107 +++++++++++++ .../snapshot/SnapshotReporterTest.java | 144 ++++++++++++++++++ .../reporting/snapshot/SnapshotStoreTest.java | 128 ++++++++++++++++ spotbugs-exclude.xml | 21 +++ 15 files changed, 1007 insertions(+) create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotCaseData.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotComparisonResult.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotConfig.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotData.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotRegressionException.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotReporter.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotScoreData.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStatus.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStore.java create mode 100644 agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotConfigTest.java create mode 100644 agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotDataTest.java create mode 100644 agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotReporterTest.java create mode 100644 agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotStoreTest.java diff --git a/agenteval-reporting/pom.xml b/agenteval-reporting/pom.xml index afab760..2e24317 100644 --- a/agenteval-reporting/pom.xml +++ b/agenteval-reporting/pom.xml @@ -27,6 +27,10 @@ com.fasterxml.jackson.core jackson-databind + + com.fasterxml.jackson.datatype + jackson-datatype-jsr310 + org.mockito mockito-core diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotCaseData.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotCaseData.java new file mode 100644 index 0000000..edb6fda --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotCaseData.java @@ -0,0 +1,48 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.model.EvalScore; + +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.Objects; + +/** + * Snapshot representation of a single test case result. + * + * @param input the test case input + * @param actualOutput the agent's actual output + * @param passed whether the case passed all metrics + * @param scores per-metric score data + */ +public record SnapshotCaseData( + String input, + String actualOutput, + boolean passed, + Map scores +) { + public SnapshotCaseData { + Objects.requireNonNull(input, "input must not be null"); + scores = scores == null ? Map.of() : Map.copyOf(scores); + } + + /** + * Creates a snapshot case from an evaluation case result. + */ + public static SnapshotCaseData from(CaseResult caseResult) { + Objects.requireNonNull(caseResult, "caseResult must not be null"); + + Map scoreMap = new LinkedHashMap<>(); + for (Map.Entry entry : caseResult.scores().entrySet()) { + EvalScore s = entry.getValue(); + scoreMap.put(entry.getKey(), + new SnapshotScoreData(s.value(), s.threshold(), s.passed(), s.reason())); + } + + return new SnapshotCaseData( + caseResult.testCase().getInput(), + caseResult.testCase().getActualOutput(), + caseResult.passed(), + scoreMap); + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotComparisonResult.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotComparisonResult.java new file mode 100644 index 0000000..8a719f5 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotComparisonResult.java @@ -0,0 +1,23 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.reporting.regression.RegressionReport; + +import java.util.Objects; + +/** + * Result of comparing current evaluation results against a saved snapshot. + * + * @param snapshotName the name of the snapshot compared against + * @param status the comparison status + * @param regressionReport the detailed regression report + */ +public record SnapshotComparisonResult( + String snapshotName, + SnapshotStatus status, + RegressionReport regressionReport +) { + public SnapshotComparisonResult { + Objects.requireNonNull(snapshotName, "snapshotName must not be null"); + Objects.requireNonNull(status, "status must not be null"); + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotConfig.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotConfig.java new file mode 100644 index 0000000..d9c165a --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotConfig.java @@ -0,0 +1,85 @@ +package com.agenteval.reporting.snapshot; + +import java.nio.file.Path; + +/** + * Configuration for snapshot testing. + * + *
{@code
+ * var config = SnapshotConfig.builder()
+ *     .snapshotDirectory(Path.of("src/test/resources/agenteval-snapshots"))
+ *     .updateSnapshots(false)
+ *     .failOnRegression(true)
+ *     .regressionThreshold(0.0)
+ *     .build();
+ * }
+ */ +public final class SnapshotConfig { + + private static final Path DEFAULT_DIRECTORY = Path.of("src/test/resources/agenteval-snapshots"); + + private final Path snapshotDirectory; + private final boolean updateSnapshots; + private final boolean failOnRegression; + private final double regressionThreshold; + + private SnapshotConfig(Builder builder) { + this.snapshotDirectory = builder.snapshotDirectory; + this.updateSnapshots = builder.updateSnapshots; + this.failOnRegression = builder.failOnRegression; + this.regressionThreshold = builder.regressionThreshold; + } + + public Path snapshotDirectory() { return snapshotDirectory; } + public boolean updateSnapshots() { return updateSnapshots; } + public boolean failOnRegression() { return failOnRegression; } + public double regressionThreshold() { return regressionThreshold; } + + public static Builder builder() { + return new Builder(); + } + + public static SnapshotConfig defaults() { + return new Builder().build(); + } + + public static final class Builder { + private Path snapshotDirectory = DEFAULT_DIRECTORY; + private boolean updateSnapshots = false; + private boolean failOnRegression = true; + private double regressionThreshold = 0.0; + + private Builder() {} + + public Builder snapshotDirectory(Path directory) { + if (directory == null) { + throw new IllegalArgumentException("snapshotDirectory must not be null"); + } + this.snapshotDirectory = directory; + return this; + } + + public Builder updateSnapshots(boolean update) { + this.updateSnapshots = update; + return this; + } + + public Builder failOnRegression(boolean fail) { + this.failOnRegression = fail; + return this; + } + + public Builder regressionThreshold(double threshold) { + if (threshold < 0.0 || threshold > 1.0) { + throw new IllegalArgumentException( + "regressionThreshold must be between 0.0 and 1.0, got: " + threshold); + } + this.regressionThreshold = threshold; + return this; + } + + public SnapshotConfig build() { + return new SnapshotConfig(this); + } + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotData.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotData.java new file mode 100644 index 0000000..58dbd6f --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotData.java @@ -0,0 +1,96 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; + +import java.time.Instant; +import java.util.ArrayList; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; + +/** + * Complete snapshot of an evaluation run, persisted as JSON. + * + * @param snapshotName the logical name of this snapshot + * @param createdAt when the snapshot was created + * @param averageScore overall average score + * @param passRate overall pass rate (0.0–1.0) + * @param totalCases total number of test cases + * @param durationMs evaluation duration in milliseconds + * @param metricAverages per-metric average scores + * @param caseResults per-case snapshot data + */ +public record SnapshotData( + String snapshotName, + Instant createdAt, + double averageScore, + double passRate, + int totalCases, + long durationMs, + Map metricAverages, + List caseResults +) { + public SnapshotData { + Objects.requireNonNull(snapshotName, "snapshotName must not be null"); + Objects.requireNonNull(createdAt, "createdAt must not be null"); + metricAverages = metricAverages == null ? Map.of() : Map.copyOf(metricAverages); + caseResults = caseResults == null ? List.of() : List.copyOf(caseResults); + } + + /** + * Creates a snapshot from an evaluation result. + * + * @param name the snapshot name + * @param result the evaluation result + * @return the snapshot data + */ + public static SnapshotData from(String name, EvalResult result) { + Objects.requireNonNull(name, "name must not be null"); + Objects.requireNonNull(result, "result must not be null"); + + List cases = result.caseResults().stream() + .map(SnapshotCaseData::from) + .toList(); + + return new SnapshotData( + name, + Instant.now(), + result.averageScore(), + result.passRate(), + result.caseResults().size(), + result.durationMs(), + new LinkedHashMap<>(result.averageScoresByMetric()), + cases); + } + + /** + * Reconstructs a synthetic {@link EvalResult} for use with + * {@link com.agenteval.reporting.regression.RegressionComparison}. + */ + public EvalResult toEvalResult() { + List cases = new ArrayList<>(caseResults.size()); + + for (SnapshotCaseData snapCase : caseResults) { + AgentTestCase testCase = AgentTestCase.builder() + .input(snapCase.input()) + .actualOutput(snapCase.actualOutput()) + .build(); + + Map scores = new LinkedHashMap<>(); + for (Map.Entry entry : snapCase.scores().entrySet()) { + SnapshotScoreData sd = entry.getValue(); + scores.put(entry.getKey(), + new EvalScore(sd.value(), sd.threshold(), sd.passed(), + sd.reason() != null ? sd.reason() : "", entry.getKey())); + } + + cases.add(new CaseResult(testCase, scores, snapCase.passed())); + } + + return EvalResult.of(cases, durationMs); + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotRegressionException.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotRegressionException.java new file mode 100644 index 0000000..103ed46 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotRegressionException.java @@ -0,0 +1,29 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.reporting.ReportException; +import com.agenteval.reporting.regression.RegressionReport; + +import java.util.Objects; + +/** + * Thrown when snapshot comparison detects regressions. + */ +public class SnapshotRegressionException extends ReportException { + + private static final long serialVersionUID = 1L; + + private final transient RegressionReport regressionReport; + + public SnapshotRegressionException(String message, RegressionReport regressionReport) { + super(message); + this.regressionReport = Objects.requireNonNull(regressionReport, + "regressionReport must not be null"); + } + + /** + * Returns the regression report that triggered this exception. + */ + public RegressionReport regressionReport() { + return regressionReport; + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotReporter.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotReporter.java new file mode 100644 index 0000000..acf275a --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotReporter.java @@ -0,0 +1,106 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.core.eval.EvalResult; +import com.agenteval.reporting.EvalReporter; +import com.agenteval.reporting.regression.RegressionComparison; +import com.agenteval.reporting.regression.RegressionReport; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +import java.util.Objects; +import java.util.Optional; + +/** + * Reporter that saves evaluation results as named snapshots and compares + * against prior baselines to detect regressions. + * + *

Behavior:

+ *
    + *
  • If {@code updateSnapshots} or env {@code AGENTEVAL_UPDATE_SNAPSHOTS=true} + * → save and return (overwrite existing)
  • + *
  • If no prior snapshot exists → save as baseline
  • + *
  • Otherwise → compare and fail on regressions if configured
  • + *
+ */ +public final class SnapshotReporter implements EvalReporter { + + private static final Logger LOG = LoggerFactory.getLogger(SnapshotReporter.class); + + private final String snapshotName; + private final SnapshotStore store; + private final SnapshotConfig config; + + public SnapshotReporter(String snapshotName, SnapshotConfig config) { + this.snapshotName = Objects.requireNonNull(snapshotName, + "snapshotName must not be null"); + this.config = Objects.requireNonNull(config, "config must not be null"); + this.store = new SnapshotStore(config.snapshotDirectory()); + } + + @Override + public void report(EvalResult result) { + if (shouldUpdate()) { + LOG.info("Updating snapshot '{}'", snapshotName); + store.save(SnapshotData.from(snapshotName, result)); + return; + } + + Optional existing = store.load(snapshotName); + if (existing.isEmpty()) { + LOG.info("Creating baseline snapshot '{}'", snapshotName); + store.save(SnapshotData.from(snapshotName, result)); + return; + } + + SnapshotComparisonResult comparison = doCompare(result, existing.get()); + if (comparison.status() == SnapshotStatus.REGRESSED + && config.failOnRegression()) { + throw new SnapshotRegressionException( + "Snapshot '" + snapshotName + "' has regressions: " + + comparison.regressionReport().newFailures() + " new failure(s), " + + "overall delta: " + + String.format("%.4f", comparison.regressionReport().overallDelta()), + comparison.regressionReport()); + } + + LOG.info("Snapshot '{}' comparison: {}", snapshotName, comparison.status()); + } + + /** + * Compares an evaluation result against the stored snapshot without saving. + * + * @param result the current evaluation result + * @return the comparison result, or empty if no baseline exists + */ + public Optional compareOnly(EvalResult result) { + Optional existing = store.load(snapshotName); + if (existing.isEmpty()) { + return Optional.empty(); + } + return Optional.of(doCompare(result, existing.get())); + } + + private SnapshotComparisonResult doCompare(EvalResult current, SnapshotData baseline) { + EvalResult baselineResult = baseline.toEvalResult(); + RegressionReport report = RegressionComparison.compare(baselineResult, current); + + SnapshotStatus status; + if (report.hasRegressions()) { + status = SnapshotStatus.REGRESSED; + } else if (report.overallDelta() > config.regressionThreshold()) { + status = SnapshotStatus.IMPROVED; + } else { + status = SnapshotStatus.MATCHED; + } + + return new SnapshotComparisonResult(snapshotName, status, report); + } + + private boolean shouldUpdate() { + if (config.updateSnapshots()) { + return true; + } + String envUpdate = System.getenv("AGENTEVAL_UPDATE_SNAPSHOTS"); + return "true".equalsIgnoreCase(envUpdate); + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotScoreData.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotScoreData.java new file mode 100644 index 0000000..d131857 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotScoreData.java @@ -0,0 +1,18 @@ +package com.agenteval.reporting.snapshot; + +/** + * Jackson-serializable mirror of {@link com.agenteval.core.model.EvalScore} + * for snapshot persistence. + * + * @param value the score value (0.0–1.0) + * @param threshold the pass/fail threshold + * @param passed whether the score met the threshold + * @param reason the reason for the score + */ +public record SnapshotScoreData( + double value, + double threshold, + boolean passed, + String reason +) { +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStatus.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStatus.java new file mode 100644 index 0000000..2b100e7 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStatus.java @@ -0,0 +1,15 @@ +package com.agenteval.reporting.snapshot; + +/** + * Status of a snapshot comparison. + */ +public enum SnapshotStatus { + /** New snapshot created (no prior baseline). */ + CREATED, + /** Current results match the baseline snapshot. */ + MATCHED, + /** Current results show regressions compared to baseline. */ + REGRESSED, + /** Current results show improvement over baseline. */ + IMPROVED +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStore.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStore.java new file mode 100644 index 0000000..9a9eca5 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/snapshot/SnapshotStore.java @@ -0,0 +1,121 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.reporting.ReportException; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.SerializationFeature; +import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Objects; +import java.util.Optional; +import java.util.regex.Pattern; + +/** + * Persists and loads {@link SnapshotData} as JSON files. + * + *

Snapshot names are validated to prevent path traversal attacks. + * Files are stored as {@code .snapshot.json} in the configured directory.

+ */ +public final class SnapshotStore { + + private static final Pattern VALID_NAME = Pattern.compile("^[a-zA-Z0-9][a-zA-Z0-9_.-]*$"); + private static final String EXTENSION = ".snapshot.json"; + + private static final ObjectMapper MAPPER = new ObjectMapper() + .registerModule(new JavaTimeModule()) + .disable(SerializationFeature.WRITE_DATES_AS_TIMESTAMPS) + .enable(SerializationFeature.INDENT_OUTPUT); + + private final Path directory; + + public SnapshotStore(Path directory) { + this.directory = Objects.requireNonNull(directory, "directory must not be null"); + } + + /** + * Saves snapshot data to disk. + * + * @param snapshot the snapshot to save + * @throws ReportException if writing fails + */ + public void save(SnapshotData snapshot) { + Objects.requireNonNull(snapshot, "snapshot must not be null"); + validateName(snapshot.snapshotName()); + + try { + Files.createDirectories(directory); + Path file = resolve(snapshot.snapshotName()); + MAPPER.writeValue(file.toFile(), snapshot); + } catch (IOException e) { + throw new ReportException( + "Failed to save snapshot '" + snapshot.snapshotName() + "'", e); + } + } + + /** + * Loads a snapshot by name. + * + * @param name the snapshot name + * @return the snapshot data, or empty if not found + * @throws ReportException if reading fails + */ + public Optional load(String name) { + validateName(name); + Path file = resolve(name); + + if (!Files.exists(file)) { + return Optional.empty(); + } + + try { + return Optional.of(MAPPER.readValue(file.toFile(), SnapshotData.class)); + } catch (IOException e) { + throw new ReportException("Failed to load snapshot '" + name + "'", e); + } + } + + /** + * Checks whether a snapshot exists. + */ + public boolean exists(String name) { + validateName(name); + return Files.exists(resolve(name)); + } + + /** + * Deletes a snapshot. + * + * @param name the snapshot name + * @return true if the snapshot was deleted, false if it did not exist + * @throws ReportException if deletion fails + */ + public boolean delete(String name) { + validateName(name); + try { + return Files.deleteIfExists(resolve(name)); + } catch (IOException e) { + throw new ReportException("Failed to delete snapshot '" + name + "'", e); + } + } + + private Path resolve(String name) { + return directory.resolve(name + EXTENSION); + } + + private static void validateName(String name) { + if (name == null || name.isEmpty()) { + throw new IllegalArgumentException("Snapshot name must not be null or empty"); + } + if (!VALID_NAME.matcher(name).matches()) { + throw new IllegalArgumentException( + "Invalid snapshot name: '" + name + + "'. Must match [a-zA-Z0-9][a-zA-Z0-9_.-]*"); + } + if (name.contains("..")) { + throw new IllegalArgumentException( + "Snapshot name must not contain '..'"); + } + } +} diff --git a/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotConfigTest.java b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotConfigTest.java new file mode 100644 index 0000000..5c0c0c8 --- /dev/null +++ b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotConfigTest.java @@ -0,0 +1,62 @@ +package com.agenteval.reporting.snapshot; + +import org.junit.jupiter.api.Test; + +import java.nio.file.Path; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.assertj.core.api.Assertions.within; + +class SnapshotConfigTest { + + @Test + void defaultValues() { + SnapshotConfig config = SnapshotConfig.defaults(); + + assertThat(config.snapshotDirectory()) + .isEqualTo(Path.of("src/test/resources/agenteval-snapshots")); + assertThat(config.updateSnapshots()).isFalse(); + assertThat(config.failOnRegression()).isTrue(); + assertThat(config.regressionThreshold()).isCloseTo(0.0, within(0.001)); + } + + @Test + void customValues() { + Path customDir = Path.of("/tmp/custom-snaps"); + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(customDir) + .updateSnapshots(true) + .failOnRegression(false) + .regressionThreshold(0.05) + .build(); + + assertThat(config.snapshotDirectory()).isEqualTo(customDir); + assertThat(config.updateSnapshots()).isTrue(); + assertThat(config.failOnRegression()).isFalse(); + assertThat(config.regressionThreshold()).isCloseTo(0.05, within(0.001)); + } + + @Test + void rejectsNullDirectory() { + assertThatThrownBy(() -> SnapshotConfig.builder().snapshotDirectory(null)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void rejectsInvalidRegressionThreshold() { + assertThatThrownBy(() -> SnapshotConfig.builder().regressionThreshold(-0.1)) + .isInstanceOf(IllegalArgumentException.class); + assertThatThrownBy(() -> SnapshotConfig.builder().regressionThreshold(1.1)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void validBoundaryThresholds() { + SnapshotConfig lower = SnapshotConfig.builder().regressionThreshold(0.0).build(); + assertThat(lower.regressionThreshold()).isCloseTo(0.0, within(0.001)); + + SnapshotConfig upper = SnapshotConfig.builder().regressionThreshold(1.0).build(); + assertThat(upper.regressionThreshold()).isCloseTo(1.0, within(0.001)); + } +} diff --git a/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotDataTest.java b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotDataTest.java new file mode 100644 index 0000000..5ee8921 --- /dev/null +++ b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotDataTest.java @@ -0,0 +1,107 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import com.agenteval.reporting.regression.RegressionComparison; +import com.agenteval.reporting.regression.RegressionReport; +import org.junit.jupiter.api.Test; + +import java.util.List; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.within; + +class SnapshotDataTest { + + @Test + void fromEvalResultCapturesAllFields() { + EvalResult result = makeResult("What is Java?", "Relevancy", 0.9, 0.7, true); + SnapshotData snapshot = SnapshotData.from("test", result); + + assertThat(snapshot.snapshotName()).isEqualTo("test"); + assertThat(snapshot.averageScore()).isCloseTo(0.9, within(0.001)); + assertThat(snapshot.passRate()).isCloseTo(1.0, within(0.001)); + assertThat(snapshot.totalCases()).isEqualTo(1); + assertThat(snapshot.createdAt()).isNotNull(); + assertThat(snapshot.metricAverages()).containsEntry("Relevancy", 0.9); + } + + @Test + void fromEvalResultPreservesCaseData() { + EvalResult result = makeResult("What is Java?", "Relevancy", 0.9, 0.7, true); + SnapshotData snapshot = SnapshotData.from("test", result); + + assertThat(snapshot.caseResults()).hasSize(1); + SnapshotCaseData caseData = snapshot.caseResults().get(0); + assertThat(caseData.input()).isEqualTo("What is Java?"); + assertThat(caseData.passed()).isTrue(); + assertThat(caseData.scores()).containsKey("Relevancy"); + assertThat(caseData.scores().get("Relevancy").value()).isCloseTo(0.9, within(0.001)); + } + + @Test + void toEvalResultRoundTrip() { + EvalResult original = makeResult("What is Java?", "Relevancy", 0.85, 0.7, true); + SnapshotData snapshot = SnapshotData.from("round-trip", original); + EvalResult reconstructed = snapshot.toEvalResult(); + + assertThat(reconstructed.averageScore()) + .isCloseTo(original.averageScore(), within(0.001)); + assertThat(reconstructed.passRate()) + .isCloseTo(original.passRate(), within(0.001)); + assertThat(reconstructed.caseResults()).hasSize(1); + assertThat(reconstructed.caseResults().get(0).testCase().getInput()) + .isEqualTo("What is Java?"); + assertThat(reconstructed.caseResults().get(0).scores().get("Relevancy").value()) + .isCloseTo(0.85, within(0.001)); + } + + @Test + void toEvalResultWorksWithRegressionComparison() { + EvalResult baseline = makeResult("Q1", "M1", 0.9, 0.7, true); + SnapshotData snapshot = SnapshotData.from("baseline", baseline); + + EvalResult current = makeResult("Q1", "M1", 0.5, 0.7, false); + EvalResult reconstructedBaseline = snapshot.toEvalResult(); + + RegressionReport report = RegressionComparison.compare(reconstructedBaseline, current); + assertThat(report.hasRegressions()).isTrue(); + assertThat(report.newFailures()).isEqualTo(1); + } + + @Test + void multiCaseRoundTrip() { + AgentTestCase tc1 = AgentTestCase.builder() + .input("Q1").actualOutput("A1").build(); + AgentTestCase tc2 = AgentTestCase.builder() + .input("Q2").actualOutput("A2").build(); + + EvalScore s1 = new EvalScore(0.9, 0.7, true, "good", "M1"); + EvalScore s2 = new EvalScore(0.4, 0.7, false, "bad", "M1"); + + EvalResult result = EvalResult.of(List.of( + new CaseResult(tc1, Map.of("M1", s1), true), + new CaseResult(tc2, Map.of("M1", s2), false) + ), 150L); + + SnapshotData snapshot = SnapshotData.from("multi", result); + EvalResult reconstructed = snapshot.toEvalResult(); + + assertThat(reconstructed.caseResults()).hasSize(2); + assertThat(reconstructed.failedCases()).hasSize(1); + assertThat(reconstructed.caseResults().get(0).passed()).isTrue(); + assertThat(reconstructed.caseResults().get(1).passed()).isFalse(); + } + + private static EvalResult makeResult(String input, String metric, + double score, double threshold, boolean passed) { + AgentTestCase tc = AgentTestCase.builder() + .input(input).actualOutput("answer").build(); + EvalScore evalScore = new EvalScore(score, threshold, passed, "test reason", metric); + CaseResult cr = new CaseResult(tc, Map.of(metric, evalScore), passed); + return EvalResult.of(List.of(cr), 100L); + } +} diff --git a/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotReporterTest.java b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotReporterTest.java new file mode 100644 index 0000000..548d8a8 --- /dev/null +++ b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotReporterTest.java @@ -0,0 +1,144 @@ +package com.agenteval.reporting.snapshot; + +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import java.nio.file.Path; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class SnapshotReporterTest { + + @TempDir + Path tempDir; + + @Test + void firstRunCreatesBaseline() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .build(); + var reporter = new SnapshotReporter("baseline-test", config); + + reporter.report(makeResult(0.9)); + + var store = new SnapshotStore(tempDir); + assertThat(store.exists("baseline-test")).isTrue(); + } + + @Test + void matchingResultsDoNotThrow() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .failOnRegression(true) + .build(); + var reporter = new SnapshotReporter("stable", config); + + reporter.report(makeResult(0.9)); + // Same score should match + reporter.report(makeResult(0.9)); + } + + @Test + void regressionThrowsWhenConfigured() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .failOnRegression(true) + .build(); + var reporter = new SnapshotReporter("regress", config); + + // Create baseline with high score + reporter.report(makeResult(0.9)); + + // Lower score should regress + assertThatThrownBy(() -> reporter.report(makeResult(0.3))) + .isInstanceOf(SnapshotRegressionException.class) + .hasMessageContaining("regressions"); + } + + @Test + void regressionAllowedWhenFailOnRegressionFalse() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .failOnRegression(false) + .build(); + var reporter = new SnapshotReporter("lenient", config); + + reporter.report(makeResult(0.9)); + // Should not throw even with regression + reporter.report(makeResult(0.3)); + } + + @Test + void updateModeOverwritesExisting() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .updateSnapshots(true) + .build(); + var reporter = new SnapshotReporter("update-test", config); + var store = new SnapshotStore(tempDir); + + reporter.report(makeResult(0.9)); + reporter.report(makeResult(0.5)); + + SnapshotData loaded = store.load("update-test").orElseThrow(); + assertThat(loaded.averageScore()).isCloseTo(0.5, + org.assertj.core.api.Assertions.within(0.001)); + } + + @Test + void compareOnlyReturnsEmptyWhenNoBaseline() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .build(); + var reporter = new SnapshotReporter("no-baseline", config); + + Optional result = reporter.compareOnly(makeResult(0.9)); + assertThat(result).isEmpty(); + } + + @Test + void compareOnlyReturnsResultWhenBaselineExists() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .build(); + var reporter = new SnapshotReporter("compare-test", config); + + // Create baseline + reporter.report(makeResult(0.9)); + + // Compare without saving + Optional result = reporter.compareOnly(makeResult(0.95)); + assertThat(result).isPresent(); + assertThat(result.get().status()).isEqualTo(SnapshotStatus.IMPROVED); + } + + @Test + void improvementDetected() { + SnapshotConfig config = SnapshotConfig.builder() + .snapshotDirectory(tempDir) + .failOnRegression(true) + .build(); + var reporter = new SnapshotReporter("improve", config); + + reporter.report(makeResult(0.5)); + // Higher score should not throw + reporter.report(makeResult(0.9)); + } + + private static EvalResult makeResult(double score) { + AgentTestCase tc = AgentTestCase.builder() + .input("What is Java?").actualOutput("A programming language").build(); + boolean passed = score >= 0.7; + EvalScore evalScore = new EvalScore(score, 0.7, passed, "test", "TestMetric"); + CaseResult cr = new CaseResult(tc, Map.of("TestMetric", evalScore), passed); + return EvalResult.of(List.of(cr), 100L); + } +} diff --git a/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotStoreTest.java b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotStoreTest.java new file mode 100644 index 0000000..0244374 --- /dev/null +++ b/agenteval-reporting/src/test/java/com/agenteval/reporting/snapshot/SnapshotStoreTest.java @@ -0,0 +1,128 @@ +package com.agenteval.reporting.snapshot; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import java.nio.file.Path; +import java.time.Instant; +import java.util.List; +import java.util.Map; +import java.util.Optional; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.assertj.core.api.Assertions.within; + +class SnapshotStoreTest { + + @TempDir + Path tempDir; + + @Test + void saveAndLoadRoundTrip() { + var store = new SnapshotStore(tempDir); + var snapshot = makeSnapshot("test-snap", 0.85, 1.0, 1); + store.save(snapshot); + + Optional loaded = store.load("test-snap"); + assertThat(loaded).isPresent(); + assertThat(loaded.get().snapshotName()).isEqualTo("test-snap"); + assertThat(loaded.get().averageScore()).isCloseTo(0.85, within(0.001)); + assertThat(loaded.get().passRate()).isCloseTo(1.0, within(0.001)); + assertThat(loaded.get().totalCases()).isEqualTo(1); + } + + @Test + void loadMissingSnapshotReturnsEmpty() { + var store = new SnapshotStore(tempDir); + assertThat(store.load("nonexistent")).isEmpty(); + } + + @Test + void createsDirectoryAutomatically() { + Path nested = tempDir.resolve("sub/dir"); + var store = new SnapshotStore(nested); + store.save(makeSnapshot("auto-dir", 0.9, 1.0, 1)); + + assertThat(store.exists("auto-dir")).isTrue(); + } + + @Test + void existsReturnsTrueForSavedSnapshot() { + var store = new SnapshotStore(tempDir); + assertThat(store.exists("missing")).isFalse(); + + store.save(makeSnapshot("exists-test", 0.9, 1.0, 1)); + assertThat(store.exists("exists-test")).isTrue(); + } + + @Test + void deleteRemovesSnapshot() { + var store = new SnapshotStore(tempDir); + store.save(makeSnapshot("to-delete", 0.9, 1.0, 1)); + assertThat(store.exists("to-delete")).isTrue(); + + assertThat(store.delete("to-delete")).isTrue(); + assertThat(store.exists("to-delete")).isFalse(); + } + + @Test + void deleteNonexistentReturnsFalse() { + var store = new SnapshotStore(tempDir); + assertThat(store.delete("nope")).isFalse(); + } + + @Test + void rejectsInvalidNames() { + var store = new SnapshotStore(tempDir); + assertThatThrownBy(() -> store.save(makeSnapshot("../evil", 0.0, 0.0, 0))) + .isInstanceOf(IllegalArgumentException.class); + assertThatThrownBy(() -> store.load("")) + .isInstanceOf(IllegalArgumentException.class); + assertThatThrownBy(() -> store.load("has spaces")) + .isInstanceOf(IllegalArgumentException.class); + assertThatThrownBy(() -> store.exists(null)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void multiCaseSnapshotPreservesAllData() { + var store = new SnapshotStore(tempDir); + var cases = List.of( + new SnapshotCaseData("input1", "output1", true, + Map.of("Metric1", new SnapshotScoreData(0.9, 0.7, true, "good"))), + new SnapshotCaseData("input2", "output2", false, + Map.of("Metric1", new SnapshotScoreData(0.5, 0.7, false, "poor"))) + ); + var snapshot = new SnapshotData("multi", Instant.now(), 0.7, 0.5, 2, 200L, + Map.of("Metric1", 0.7), cases); + + store.save(snapshot); + SnapshotData loaded = store.load("multi").orElseThrow(); + + assertThat(loaded.caseResults()).hasSize(2); + assertThat(loaded.caseResults().get(0).input()).isEqualTo("input1"); + assertThat(loaded.caseResults().get(1).passed()).isFalse(); + assertThat(loaded.caseResults().get(1).scores().get("Metric1").value()) + .isCloseTo(0.5, within(0.001)); + } + + @Test + void overwritesExistingSnapshot() { + var store = new SnapshotStore(tempDir); + store.save(makeSnapshot("rewrite", 0.5, 0.5, 1)); + store.save(makeSnapshot("rewrite", 0.9, 1.0, 1)); + + SnapshotData loaded = store.load("rewrite").orElseThrow(); + assertThat(loaded.averageScore()).isCloseTo(0.9, within(0.001)); + } + + private static SnapshotData makeSnapshot(String name, double avg, double pass, int cases) { + List caseList = List.of( + new SnapshotCaseData("input", "output", pass == 1.0, + Map.of("TestMetric", new SnapshotScoreData(avg, 0.7, avg >= 0.7, "ok"))) + ); + return new SnapshotData(name, Instant.now(), avg, pass, cases, 100L, + Map.of("TestMetric", avg), caseList); + } +} diff --git a/spotbugs-exclude.xml b/spotbugs-exclude.xml index 8fe1927..8a050bf 100644 --- a/spotbugs-exclude.xml +++ b/spotbugs-exclude.xml @@ -119,4 +119,25 @@ + + + + + + + + + + + + + + + + + + From 677f685a39808b89a4d7b23a8f28bf266839b4dc Mon Sep 17 00:00:00 2001 From: Pratyush Sharma <56130065+pratyush618@users.noreply.github.com> Date: Fri, 13 Mar 2026 11:33:59 +0530 Subject: [PATCH 2/4] Add benchmark mode for multi-variant evaluation Run the same dataset against multiple config variants and compare results. AgentTestCase.toBuilder() enables isolated deep-copies per variant. Benchmark supports parallel execution via virtual threads with Semaphore-based concurrency control. BenchmarkReporter outputs console tables with [BEST]/[WORST] labels and per-metric breakdown. BenchmarkComparison bridges to RegressionComparison for variant diffs. --- .../agenteval/core/benchmark/Benchmark.java | 135 ++++++++++++++++ .../core/benchmark/BenchmarkConfig.java | 57 +++++++ .../core/benchmark/BenchmarkResult.java | 89 +++++++++++ .../core/benchmark/BenchmarkVariant.java | 74 +++++++++ .../agenteval/core/model/AgentTestCase.java | 21 +++ .../core/benchmark/BenchmarkConfigTest.java | 39 +++++ .../core/benchmark/BenchmarkResultTest.java | 109 +++++++++++++ .../core/benchmark/BenchmarkTest.java | 150 ++++++++++++++++++ .../core/benchmark/BenchmarkVariantTest.java | 102 ++++++++++++ .../model/AgentTestCaseToBuilderTest.java | 87 ++++++++++ .../benchmark/BenchmarkComparison.java | 59 +++++++ .../benchmark/BenchmarkReporter.java | 95 +++++++++++ .../benchmark/BenchmarkComparisonTest.java | 88 ++++++++++ .../benchmark/BenchmarkReporterTest.java | 97 +++++++++++ 14 files changed, 1202 insertions(+) create mode 100644 agenteval-core/src/main/java/com/agenteval/core/benchmark/Benchmark.java create mode 100644 agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkConfig.java create mode 100644 agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkResult.java create mode 100644 agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkVariant.java create mode 100644 agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkConfigTest.java create mode 100644 agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkResultTest.java create mode 100644 agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkTest.java create mode 100644 agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkVariantTest.java create mode 100644 agenteval-core/src/test/java/com/agenteval/core/model/AgentTestCaseToBuilderTest.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkComparison.java create mode 100644 agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkReporter.java create mode 100644 agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkComparisonTest.java create mode 100644 agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkReporterTest.java diff --git a/agenteval-core/src/main/java/com/agenteval/core/benchmark/Benchmark.java b/agenteval-core/src/main/java/com/agenteval/core/benchmark/Benchmark.java new file mode 100644 index 0000000..c0081fa --- /dev/null +++ b/agenteval-core/src/main/java/com/agenteval/core/benchmark/Benchmark.java @@ -0,0 +1,135 @@ +package com.agenteval.core.benchmark; + +import com.agenteval.core.eval.AgentEval; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.eval.EvaluationException; +import com.agenteval.core.model.AgentTestCase; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +import java.util.ArrayList; +import java.util.HashSet; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; +import java.util.Set; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.Future; +import java.util.concurrent.Semaphore; + +/** + * Runs the same dataset against multiple configuration variants and collects results. + * + *
{@code
+ * var result = Benchmark.run(testCases, List.of(variantA, variantB));
+ * System.out.println("Best: " + result.bestVariant());
+ * }
+ */ +public final class Benchmark { + + private static final Logger LOG = LoggerFactory.getLogger(Benchmark.class); + + private Benchmark() {} + + /** + * Runs benchmarks sequentially with default config. + */ + public static BenchmarkResult run(List testCases, + List variants) { + return run(testCases, variants, BenchmarkConfig.defaults()); + } + + /** + * Runs benchmarks with the specified config. + */ + public static BenchmarkResult run(List testCases, + List variants, + BenchmarkConfig config) { + Objects.requireNonNull(testCases, "testCases must not be null"); + Objects.requireNonNull(variants, "variants must not be null"); + Objects.requireNonNull(config, "config must not be null"); + + validateUniqueNames(variants); + + long startTime = System.currentTimeMillis(); + Map results; + + if (config.parallelVariants() && variants.size() > 1) { + results = runParallel(testCases, variants, config); + } else { + results = runSequential(testCases, variants); + } + + long totalDuration = System.currentTimeMillis() - startTime; + LOG.info("Benchmark complete: {} variants in {}ms", variants.size(), totalDuration); + + return new BenchmarkResult(results, totalDuration); + } + + private static Map runSequential( + List testCases, List variants) { + Map results = new LinkedHashMap<>(); + for (BenchmarkVariant variant : variants) { + LOG.info("Running variant '{}'", variant.name()); + results.put(variant.name(), evaluateVariant(testCases, variant)); + } + return results; + } + + private static Map runParallel( + List testCases, List variants, + BenchmarkConfig config) { + Semaphore semaphore = new Semaphore(config.maxParallelVariants()); + Map results = new LinkedHashMap<>(); + + try (ExecutorService executor = Executors.newVirtualThreadPerTaskExecutor()) { + List>> futures = new ArrayList<>(); + + for (BenchmarkVariant variant : variants) { + futures.add(executor.submit(() -> { + semaphore.acquire(); + try { + LOG.info("Running variant '{}' (parallel)", variant.name()); + EvalResult result = evaluateVariant(testCases, variant); + return Map.entry(variant.name(), result); + } finally { + semaphore.release(); + } + })); + } + + for (Future> future : futures) { + try { + Map.Entry entry = future.get(); + results.put(entry.getKey(), entry.getValue()); + } catch (Exception e) { + throw new EvaluationException("Parallel benchmark failed", e); + } + } + } + + return results; + } + + private static EvalResult evaluateVariant(List testCases, + BenchmarkVariant variant) { + // Deep-copy test cases for this variant to ensure isolation + List copied = testCases.stream() + .map(tc -> tc.toBuilder().build()) + .map(variant.casePreparer()) + .toList(); + + return AgentEval.evaluate(copied, variant.metrics(), variant.config()); + } + + private static void validateUniqueNames(List variants) { + Set seen = new HashSet<>(); + for (BenchmarkVariant v : variants) { + if (!seen.add(v.name())) { + throw new IllegalArgumentException("Duplicate variant name: " + v.name()); + } + } + } +} diff --git a/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkConfig.java b/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkConfig.java new file mode 100644 index 0000000..db25fa0 --- /dev/null +++ b/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkConfig.java @@ -0,0 +1,57 @@ +package com.agenteval.core.benchmark; + +/** + * Configuration for benchmark execution. + * + *
{@code
+ * var config = BenchmarkConfig.builder()
+ *     .parallelVariants(true)
+ *     .maxParallelVariants(4)
+ *     .build();
+ * }
+ */ +public final class BenchmarkConfig { + + private final boolean parallelVariants; + private final int maxParallelVariants; + + private BenchmarkConfig(Builder builder) { + this.parallelVariants = builder.parallelVariants; + this.maxParallelVariants = builder.maxParallelVariants; + } + + public boolean parallelVariants() { return parallelVariants; } + public int maxParallelVariants() { return maxParallelVariants; } + + public static Builder builder() { + return new Builder(); + } + + public static BenchmarkConfig defaults() { + return new Builder().build(); + } + + public static final class Builder { + private boolean parallelVariants = false; + private int maxParallelVariants = Runtime.getRuntime().availableProcessors(); + + private Builder() {} + + public Builder parallelVariants(boolean parallel) { + this.parallelVariants = parallel; + return this; + } + + public Builder maxParallelVariants(int max) { + if (max < 1) { + throw new IllegalArgumentException("maxParallelVariants must be >= 1"); + } + this.maxParallelVariants = max; + return this; + } + + public BenchmarkConfig build() { + return new BenchmarkConfig(this); + } + } +} diff --git a/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkResult.java b/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkResult.java new file mode 100644 index 0000000..400db4e --- /dev/null +++ b/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkResult.java @@ -0,0 +1,89 @@ +package com.agenteval.core.benchmark; + +import com.agenteval.core.eval.EvalResult; + +import java.util.Comparator; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Objects; + +/** + * Aggregated results from a benchmark run across multiple variants. + */ +public final class BenchmarkResult { + + private final Map variantResults; + private final long totalDurationMs; + + public BenchmarkResult(Map variantResults, long totalDurationMs) { + Objects.requireNonNull(variantResults, "variantResults must not be null"); + this.variantResults = new LinkedHashMap<>(variantResults); + this.totalDurationMs = totalDurationMs; + } + + public Map variantResults() { + return Map.copyOf(variantResults); + } + + public long totalDurationMs() { return totalDurationMs; } + + /** + * Returns the result for a specific variant. + * + * @throws IllegalArgumentException if the variant name is unknown + */ + public EvalResult resultFor(String variantName) { + EvalResult result = variantResults.get(variantName); + if (result == null) { + throw new IllegalArgumentException("Unknown variant: " + variantName); + } + return result; + } + + /** + * Returns the variant name with the highest average score. + */ + public String bestVariant() { + return variantResults.entrySet().stream() + .max(Comparator.comparingDouble(e -> e.getValue().averageScore())) + .map(Map.Entry::getKey) + .orElseThrow(() -> new IllegalStateException("No variants")); + } + + /** + * Returns the variant name with the lowest average score. + */ + public String worstVariant() { + return variantResults.entrySet().stream() + .min(Comparator.comparingDouble(e -> e.getValue().averageScore())) + .map(Map.Entry::getKey) + .orElseThrow(() -> new IllegalStateException("No variants")); + } + + /** + * Returns variant names ordered by average score (descending). + */ + public List> averageScores() { + return variantResults.entrySet().stream() + .map(e -> Map.entry(e.getKey(), e.getValue().averageScore())) + .sorted(Map.Entry.comparingByValue().reversed()) + .toList(); + } + + /** + * Returns per-metric scores grouped by metric name, then by variant. + */ + public Map> scoresByMetric() { + Map> result = new LinkedHashMap<>(); + for (Map.Entry entry : variantResults.entrySet()) { + String variant = entry.getKey(); + Map metricAvgs = entry.getValue().averageScoresByMetric(); + for (Map.Entry metricEntry : metricAvgs.entrySet()) { + result.computeIfAbsent(metricEntry.getKey(), k -> new LinkedHashMap<>()) + .put(variant, metricEntry.getValue()); + } + } + return result; + } +} diff --git a/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkVariant.java b/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkVariant.java new file mode 100644 index 0000000..3dcb812 --- /dev/null +++ b/agenteval-core/src/main/java/com/agenteval/core/benchmark/BenchmarkVariant.java @@ -0,0 +1,74 @@ +package com.agenteval.core.benchmark; + +import com.agenteval.core.config.AgentEvalConfig; +import com.agenteval.core.metric.EvalMetric; +import com.agenteval.core.model.AgentTestCase; + +import java.util.List; +import java.util.Objects; +import java.util.function.UnaryOperator; + +/** + * A named evaluation variant for benchmarking. + * + *

Each variant can specify its own config, metrics, and a case preparer + * that transforms test cases before evaluation (e.g., to call a different model).

+ * + *
{@code
+ * var variant = BenchmarkVariant.builder()
+ *     .name("gpt-4o")
+ *     .metrics(List.of(new AnswerRelevancyMetric(judge)))
+ *     .casePreparer(tc -> tc.toBuilder().actualOutput(callGpt4o(tc.getInput())).build())
+ *     .build();
+ * }
+ */ +public final class BenchmarkVariant { + + private final String name; + private final AgentEvalConfig config; + private final List metrics; + private final UnaryOperator casePreparer; + + private BenchmarkVariant(Builder builder) { + this.name = Objects.requireNonNull(builder.name, "name must not be null"); + if (builder.name.isEmpty()) { + throw new IllegalArgumentException("name must not be empty"); + } + if (builder.metrics == null || builder.metrics.isEmpty()) { + throw new IllegalArgumentException("metrics must not be null or empty"); + } + this.config = builder.config; + this.metrics = List.copyOf(builder.metrics); + this.casePreparer = builder.casePreparer; + } + + public String name() { return name; } + public AgentEvalConfig config() { return config; } + public List metrics() { return metrics; } + public UnaryOperator casePreparer() { return casePreparer; } + + public static Builder builder() { + return new Builder(); + } + + public static final class Builder { + private String name; + private AgentEvalConfig config = AgentEvalConfig.defaults(); + private List metrics; + private UnaryOperator casePreparer = UnaryOperator.identity(); + + private Builder() {} + + public Builder name(String name) { this.name = name; return this; } + public Builder config(AgentEvalConfig config) { this.config = config; return this; } + public Builder metrics(List metrics) { this.metrics = metrics; return this; } + public Builder casePreparer(UnaryOperator preparer) { + this.casePreparer = preparer; + return this; + } + + public BenchmarkVariant build() { + return new BenchmarkVariant(this); + } + } +} diff --git a/agenteval-core/src/main/java/com/agenteval/core/model/AgentTestCase.java b/agenteval-core/src/main/java/com/agenteval/core/model/AgentTestCase.java index 7c98013..4af926e 100644 --- a/agenteval-core/src/main/java/com/agenteval/core/model/AgentTestCase.java +++ b/agenteval-core/src/main/java/com/agenteval/core/model/AgentTestCase.java @@ -51,6 +51,27 @@ public static Builder builder() { return new Builder(); } + /** + * Returns a new builder pre-populated with this test case's values. + * Useful for creating modified copies (e.g., benchmark variants). + */ + public Builder toBuilder() { + Builder b = new Builder(); + b.input = this.input; + b.actualOutput = this.actualOutput; + b.expectedOutput = this.expectedOutput; + b.retrievalContext = this.retrievalContext.isEmpty() ? null : List.copyOf(this.retrievalContext); + b.context = this.context.isEmpty() ? null : List.copyOf(this.context); + b.toolCalls = this.toolCalls.isEmpty() ? null : List.copyOf(this.toolCalls); + b.expectedToolCalls = this.expectedToolCalls.isEmpty() ? null : List.copyOf(this.expectedToolCalls); + b.reasoningTrace = this.reasoningTrace.isEmpty() ? null : List.copyOf(this.reasoningTrace); + b.latencyMs = this.latencyMs; + b.tokenUsage = this.tokenUsage; + b.cost = this.cost; + b.metadata = this.metadata.isEmpty() ? null : new java.util.HashMap<>(this.metadata); + return b; + } + // --- Getters --- public String getInput() { return input; } diff --git a/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkConfigTest.java b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkConfigTest.java new file mode 100644 index 0000000..aa05832 --- /dev/null +++ b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkConfigTest.java @@ -0,0 +1,39 @@ +package com.agenteval.core.benchmark; + +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class BenchmarkConfigTest { + + @Test + void defaultValues() { + BenchmarkConfig config = BenchmarkConfig.defaults(); + assertThat(config.parallelVariants()).isFalse(); + assertThat(config.maxParallelVariants()).isPositive(); + } + + @Test + void customValues() { + BenchmarkConfig config = BenchmarkConfig.builder() + .parallelVariants(true) + .maxParallelVariants(8) + .build(); + + assertThat(config.parallelVariants()).isTrue(); + assertThat(config.maxParallelVariants()).isEqualTo(8); + } + + @Test + void rejectsZeroParallelVariants() { + assertThatThrownBy(() -> BenchmarkConfig.builder().maxParallelVariants(0)) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void rejectsNegativeParallelVariants() { + assertThatThrownBy(() -> BenchmarkConfig.builder().maxParallelVariants(-1)) + .isInstanceOf(IllegalArgumentException.class); + } +} diff --git a/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkResultTest.java b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkResultTest.java new file mode 100644 index 0000000..0edf7ab --- /dev/null +++ b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkResultTest.java @@ -0,0 +1,109 @@ +package com.agenteval.core.benchmark; + +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import org.junit.jupiter.api.Test; + +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.assertj.core.api.Assertions.within; + +class BenchmarkResultTest { + + @Test + void bestAndWorstVariant() { + Map results = new LinkedHashMap<>(); + results.put("high", makeResult(0.9)); + results.put("low", makeResult(0.3)); + + BenchmarkResult br = new BenchmarkResult(results, 500L); + + assertThat(br.bestVariant()).isEqualTo("high"); + assertThat(br.worstVariant()).isEqualTo("low"); + } + + @Test + void averageScoresOrderedDescending() { + Map results = new LinkedHashMap<>(); + results.put("low", makeResult(0.3)); + results.put("mid", makeResult(0.6)); + results.put("high", makeResult(0.9)); + + BenchmarkResult br = new BenchmarkResult(results, 500L); + + var scores = br.averageScores(); + assertThat(scores).hasSize(3); + assertThat(scores.get(0).getKey()).isEqualTo("high"); + assertThat(scores.get(1).getKey()).isEqualTo("mid"); + assertThat(scores.get(2).getKey()).isEqualTo("low"); + } + + @Test + void scoresByMetric() { + Map results = new LinkedHashMap<>(); + results.put("v1", makeResult(0.8)); + results.put("v2", makeResult(0.6)); + + BenchmarkResult br = new BenchmarkResult(results, 300L); + Map> byMetric = br.scoresByMetric(); + + assertThat(byMetric).containsKey("TestMetric"); + assertThat(byMetric.get("TestMetric").get("v1")).isCloseTo(0.8, within(0.001)); + assertThat(byMetric.get("TestMetric").get("v2")).isCloseTo(0.6, within(0.001)); + } + + @Test + void unknownVariantThrows() { + Map results = new LinkedHashMap<>(); + results.put("v1", makeResult(0.5)); + + BenchmarkResult br = new BenchmarkResult(results, 100L); + + assertThatThrownBy(() -> br.resultFor("unknown")) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Unknown variant"); + } + + @Test + void singleVariant() { + Map results = new LinkedHashMap<>(); + results.put("only", makeResult(0.7)); + + BenchmarkResult br = new BenchmarkResult(results, 100L); + + assertThat(br.bestVariant()).isEqualTo("only"); + assertThat(br.worstVariant()).isEqualTo("only"); + assertThat(br.averageScores()).hasSize(1); + } + + @Test + void resultForReturnsCorrectResult() { + Map results = new LinkedHashMap<>(); + EvalResult expected = makeResult(0.75); + results.put("target", expected); + + BenchmarkResult br = new BenchmarkResult(results, 100L); + assertThat(br.resultFor("target").averageScore()) + .isCloseTo(0.75, within(0.001)); + } + + @Test + void totalDurationTracked() { + BenchmarkResult br = new BenchmarkResult(Map.of("v", makeResult(0.5)), 12345L); + assertThat(br.totalDurationMs()).isEqualTo(12345L); + } + + private static EvalResult makeResult(double score) { + AgentTestCase tc = AgentTestCase.builder() + .input("q").actualOutput("a").build(); + EvalScore s = new EvalScore(score, 0.7, score >= 0.7, "test", "TestMetric"); + CaseResult cr = new CaseResult(tc, Map.of("TestMetric", s), score >= 0.7); + return EvalResult.of(List.of(cr), 100L); + } +} diff --git a/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkTest.java b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkTest.java new file mode 100644 index 0000000..4de5efb --- /dev/null +++ b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkTest.java @@ -0,0 +1,150 @@ +package com.agenteval.core.benchmark; + +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.metric.EvalMetric; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import org.junit.jupiter.api.Test; + +import java.util.List; +import java.util.concurrent.atomic.AtomicInteger; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.assertj.core.api.Assertions.within; + +class BenchmarkTest { + + private static final EvalMetric PASS_METRIC = new EvalMetric() { + @Override + public EvalScore evaluate(AgentTestCase testCase) { + return EvalScore.of(0.9, 0.7, "pass"); + } + + @Override + public String name() { return "PassMetric"; } + }; + + private static final EvalMetric FAIL_METRIC = new EvalMetric() { + @Override + public EvalScore evaluate(AgentTestCase testCase) { + return EvalScore.of(0.3, 0.7, "fail"); + } + + @Override + public String name() { return "FailMetric"; } + }; + + @Test + void sequentialRun() { + List cases = List.of( + AgentTestCase.builder().input("q1").actualOutput("a1").build()); + + var v1 = BenchmarkVariant.builder().name("high").metrics(List.of(PASS_METRIC)).build(); + var v2 = BenchmarkVariant.builder().name("low").metrics(List.of(FAIL_METRIC)).build(); + + BenchmarkResult result = Benchmark.run(cases, List.of(v1, v2)); + + assertThat(result.variantResults()).hasSize(2); + assertThat(result.bestVariant()).isEqualTo("high"); + assertThat(result.worstVariant()).isEqualTo("low"); + } + + @Test + void isolatedMutationsBetweenVariants() { + AgentTestCase original = AgentTestCase.builder() + .input("q1").actualOutput("original").build(); + + var mutator = BenchmarkVariant.builder() + .name("mutator") + .metrics(List.of(PASS_METRIC)) + .casePreparer(tc -> { + tc.setActualOutput("mutated"); + return tc; + }) + .build(); + + var observer = BenchmarkVariant.builder() + .name("observer") + .metrics(List.of(new EvalMetric() { + @Override + public EvalScore evaluate(AgentTestCase testCase) { + // Should see original, not mutated + return testCase.getActualOutput().equals("original") + ? EvalScore.pass("isolated") + : EvalScore.fail("leaked"); + } + + @Override + public String name() { return "IsolationCheck"; } + })) + .build(); + + BenchmarkResult result = Benchmark.run( + List.of(original), List.of(mutator, observer)); + + // Observer should see original output due to deep copy + EvalResult observerResult = result.resultFor("observer"); + assertThat(observerResult.passRate()).isCloseTo(1.0, within(0.001)); + } + + @Test + void duplicateNamesRejected() { + var v1 = BenchmarkVariant.builder().name("same").metrics(List.of(PASS_METRIC)).build(); + var v2 = BenchmarkVariant.builder().name("same").metrics(List.of(PASS_METRIC)).build(); + + List cases = List.of( + AgentTestCase.builder().input("q").build()); + + assertThatThrownBy(() -> Benchmark.run(cases, List.of(v1, v2))) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Duplicate"); + } + + @Test + void durationTracked() { + List cases = List.of( + AgentTestCase.builder().input("q").actualOutput("a").build()); + var v = BenchmarkVariant.builder().name("v").metrics(List.of(PASS_METRIC)).build(); + + BenchmarkResult result = Benchmark.run(cases, List.of(v)); + assertThat(result.totalDurationMs()).isGreaterThanOrEqualTo(0); + } + + @Test + void casePreparerApplied() { + AtomicInteger counter = new AtomicInteger(); + + var variant = BenchmarkVariant.builder() + .name("counted") + .metrics(List.of(PASS_METRIC)) + .casePreparer(tc -> { + counter.incrementAndGet(); + return tc; + }) + .build(); + + List cases = List.of( + AgentTestCase.builder().input("q1").actualOutput("a").build(), + AgentTestCase.builder().input("q2").actualOutput("a").build()); + + Benchmark.run(cases, List.of(variant)); + assertThat(counter.get()).isEqualTo(2); + } + + @Test + void parallelMode() { + List cases = List.of( + AgentTestCase.builder().input("q").actualOutput("a").build()); + var v1 = BenchmarkVariant.builder().name("p1").metrics(List.of(PASS_METRIC)).build(); + var v2 = BenchmarkVariant.builder().name("p2").metrics(List.of(PASS_METRIC)).build(); + + BenchmarkConfig config = BenchmarkConfig.builder() + .parallelVariants(true) + .maxParallelVariants(2) + .build(); + + BenchmarkResult result = Benchmark.run(cases, List.of(v1, v2), config); + assertThat(result.variantResults()).hasSize(2); + } +} diff --git a/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkVariantTest.java b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkVariantTest.java new file mode 100644 index 0000000..e804013 --- /dev/null +++ b/agenteval-core/src/test/java/com/agenteval/core/benchmark/BenchmarkVariantTest.java @@ -0,0 +1,102 @@ +package com.agenteval.core.benchmark; + +import com.agenteval.core.metric.EvalMetric; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import org.junit.jupiter.api.Test; + +import java.util.List; +import java.util.function.UnaryOperator; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class BenchmarkVariantTest { + + private static final EvalMetric DUMMY_METRIC = new EvalMetric() { + @Override + public EvalScore evaluate(AgentTestCase testCase) { + return EvalScore.pass("ok"); + } + + @Override + public String name() { return "Dummy"; } + }; + + @Test + void buildWithRequiredFields() { + BenchmarkVariant variant = BenchmarkVariant.builder() + .name("test-variant") + .metrics(List.of(DUMMY_METRIC)) + .build(); + + assertThat(variant.name()).isEqualTo("test-variant"); + assertThat(variant.metrics()).hasSize(1); + assertThat(variant.config()).isNotNull(); + assertThat(variant.casePreparer()).isNotNull(); + } + + @Test + void defaultsApplied() { + BenchmarkVariant variant = BenchmarkVariant.builder() + .name("v1") + .metrics(List.of(DUMMY_METRIC)) + .build(); + + // Default config should be AgentEvalConfig defaults + assertThat(variant.config().parallelEvaluation()).isFalse(); + // Default casePreparer should be identity + AgentTestCase tc = AgentTestCase.builder().input("q").build(); + assertThat(variant.casePreparer().apply(tc)).isSameAs(tc); + } + + @Test + void rejectsNullName() { + assertThatThrownBy(() -> BenchmarkVariant.builder() + .metrics(List.of(DUMMY_METRIC)) + .build()) + .isInstanceOf(NullPointerException.class); + } + + @Test + void rejectsEmptyName() { + assertThatThrownBy(() -> BenchmarkVariant.builder() + .name("") + .metrics(List.of(DUMMY_METRIC)) + .build()) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void rejectsEmptyMetrics() { + assertThatThrownBy(() -> BenchmarkVariant.builder() + .name("v") + .metrics(List.of()) + .build()) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void rejectsNullMetrics() { + assertThatThrownBy(() -> BenchmarkVariant.builder() + .name("v") + .build()) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void customCasePreparer() { + UnaryOperator preparer = tc -> + tc.toBuilder().actualOutput("prepared").build(); + + BenchmarkVariant variant = BenchmarkVariant.builder() + .name("prepared") + .metrics(List.of(DUMMY_METRIC)) + .casePreparer(preparer) + .build(); + + AgentTestCase tc = AgentTestCase.builder().input("q").build(); + AgentTestCase result = variant.casePreparer().apply(tc); + assertThat(result.getActualOutput()).isEqualTo("prepared"); + } +} diff --git a/agenteval-core/src/test/java/com/agenteval/core/model/AgentTestCaseToBuilderTest.java b/agenteval-core/src/test/java/com/agenteval/core/model/AgentTestCaseToBuilderTest.java new file mode 100644 index 0000000..98aac7e --- /dev/null +++ b/agenteval-core/src/test/java/com/agenteval/core/model/AgentTestCaseToBuilderTest.java @@ -0,0 +1,87 @@ +package com.agenteval.core.model; + +import org.junit.jupiter.api.Test; + +import java.math.BigDecimal; +import java.util.List; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; + +class AgentTestCaseToBuilderTest { + + @Test + void copiesAllFields() { + AgentTestCase original = AgentTestCase.builder() + .input("question") + .actualOutput("answer") + .expectedOutput("expected") + .retrievalContext(List.of("ctx1")) + .context(List.of("doc1")) + .toolCalls(List.of(new ToolCall("fn", Map.of("a", (Object) "b"), "result", 0L))) + .expectedToolCalls(List.of(ToolCall.of("fn2"))) + .latencyMs(42L) + .cost(BigDecimal.TEN) + .metadata(Map.of("k", "v")) + .build(); + + AgentTestCase copy = original.toBuilder().build(); + + assertThat(copy.getInput()).isEqualTo("question"); + assertThat(copy.getActualOutput()).isEqualTo("answer"); + assertThat(copy.getExpectedOutput()).isEqualTo("expected"); + assertThat(copy.getRetrievalContext()).containsExactly("ctx1"); + assertThat(copy.getContext()).containsExactly("doc1"); + assertThat(copy.getToolCalls()).hasSize(1); + assertThat(copy.getExpectedToolCalls()).hasSize(1); + assertThat(copy.getLatencyMs()).isEqualTo(42L); + assertThat(copy.getCost()).isEqualByComparingTo(BigDecimal.TEN); + assertThat(copy.getMetadata()).containsEntry("k", "v"); + } + + @Test + void overrideFieldsOnCopy() { + AgentTestCase original = AgentTestCase.builder() + .input("q1") + .actualOutput("a1") + .build(); + + AgentTestCase modified = original.toBuilder() + .actualOutput("new answer") + .build(); + + assertThat(modified.getInput()).isEqualTo("q1"); + assertThat(modified.getActualOutput()).isEqualTo("new answer"); + assertThat(original.getActualOutput()).isEqualTo("a1"); + } + + @Test + void copyIsIndependent() { + AgentTestCase original = AgentTestCase.builder() + .input("q1") + .actualOutput("a1") + .build(); + + AgentTestCase copy = original.toBuilder().build(); + copy.setActualOutput("mutated"); + + assertThat(original.getActualOutput()).isEqualTo("a1"); + assertThat(copy.getActualOutput()).isEqualTo("mutated"); + } + + @Test + void emptyListFieldsHandled() { + AgentTestCase original = AgentTestCase.builder() + .input("q") + .build(); + + AgentTestCase copy = original.toBuilder().build(); + + assertThat(copy.getRetrievalContext()).isEmpty(); + assertThat(copy.getContext()).isEmpty(); + assertThat(copy.getToolCalls()).isEmpty(); + assertThat(copy.getExpectedToolCalls()).isEmpty(); + assertThat(copy.getReasoningTrace()).isEmpty(); + assertThat(copy.getMetadata()).isEmpty(); + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkComparison.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkComparison.java new file mode 100644 index 0000000..9ffb801 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkComparison.java @@ -0,0 +1,59 @@ +package com.agenteval.reporting.benchmark; + +import com.agenteval.core.benchmark.BenchmarkResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.reporting.regression.RegressionComparison; +import com.agenteval.reporting.regression.RegressionReport; + +import java.util.LinkedHashMap; +import java.util.Map; +import java.util.Objects; + +/** + * Static utility for comparing benchmark variants using regression analysis. + */ +public final class BenchmarkComparison { + + private BenchmarkComparison() {} + + /** + * Compares two variants within a benchmark result. + * + * @param result the benchmark result + * @param baseline the baseline variant name + * @param comparison the comparison variant name + * @return the regression report + * @throws IllegalArgumentException if either variant name is unknown + */ + public static RegressionReport compareVariants(BenchmarkResult result, + String baseline, + String comparison) { + Objects.requireNonNull(result, "result must not be null"); + EvalResult baseResult = result.resultFor(baseline); + EvalResult compResult = result.resultFor(comparison); + return RegressionComparison.compare(baseResult, compResult); + } + + /** + * Compares all variants against a baseline variant. + * + * @param result the benchmark result + * @param baseline the baseline variant name + * @return map of variant name to regression report (excludes baseline) + * @throws IllegalArgumentException if the baseline variant name is unknown + */ + public static Map compareAllAgainst(BenchmarkResult result, + String baseline) { + Objects.requireNonNull(result, "result must not be null"); + EvalResult baseResult = result.resultFor(baseline); + Map reports = new LinkedHashMap<>(); + + for (Map.Entry entry : result.variantResults().entrySet()) { + if (!entry.getKey().equals(baseline)) { + reports.put(entry.getKey(), + RegressionComparison.compare(baseResult, entry.getValue())); + } + } + return reports; + } +} diff --git a/agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkReporter.java b/agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkReporter.java new file mode 100644 index 0000000..d3020f6 --- /dev/null +++ b/agenteval-reporting/src/main/java/com/agenteval/reporting/benchmark/BenchmarkReporter.java @@ -0,0 +1,95 @@ +package com.agenteval.reporting.benchmark; + +import com.agenteval.core.benchmark.BenchmarkResult; +import com.agenteval.core.eval.EvalResult; + +import java.io.PrintStream; +import java.util.Map; +import java.util.Objects; + +/** + * Reports benchmark results as a formatted console table with overall scores, + * [BEST]/[WORST] labels, and per-metric breakdown. + */ +public final class BenchmarkReporter { + + private static final String ANSI_RESET = "\u001B[0m"; + private static final String ANSI_GREEN = "\u001B[32m"; + private static final String ANSI_RED = "\u001B[31m"; + private static final String ANSI_BOLD = "\u001B[1m"; + + private final PrintStream out; + private final boolean ansiColors; + + public BenchmarkReporter() { + this(System.out, true); + } + + public BenchmarkReporter(PrintStream out, boolean ansiColors) { + this.out = Objects.requireNonNull(out, "out must not be null"); + this.ansiColors = ansiColors; + } + + /** + * Reports benchmark results to the configured output stream. + */ + public void report(BenchmarkResult result) { + Objects.requireNonNull(result, "result must not be null"); + + String best = result.bestVariant(); + String worst = result.worstVariant(); + boolean singleVariant = result.variantResults().size() == 1; + + printBold("=== Benchmark Results ==="); + out.printf("Variants: %d | Duration: %dms%n", + result.variantResults().size(), result.totalDurationMs()); + out.println(); + + // Overall scores table + printBold("--- Overall Scores ---"); + out.printf(" %-30s %10s %10s %s%n", "Variant", "Avg Score", "Pass Rate", ""); + for (Map.Entry entry : result.variantResults().entrySet()) { + String name = entry.getKey(); + EvalResult eval = entry.getValue(); + String label = ""; + if (!singleVariant) { + if (name.equals(best)) { + label = colorize("[BEST]", ANSI_GREEN); + } else if (name.equals(worst)) { + label = colorize("[WORST]", ANSI_RED); + } + } + out.printf(" %-30s %10.3f %9.1f%% %s%n", + name, eval.averageScore(), eval.passRate() * 100, label); + } + out.println(); + + // Per-metric breakdown + Map> byMetric = result.scoresByMetric(); + if (!byMetric.isEmpty()) { + printBold("--- Per-Metric Breakdown ---"); + for (Map.Entry> metricEntry : byMetric.entrySet()) { + out.printf(" %s:%n", metricEntry.getKey()); + for (Map.Entry variantScore : metricEntry.getValue().entrySet()) { + out.printf(" %-28s %.3f%n", + variantScore.getKey(), variantScore.getValue()); + } + } + } + } + + private void printBold(String text) { + if (ansiColors) { + out.println(ANSI_BOLD + text + ANSI_RESET); + } else { + out.println(text); + } + } + + private String colorize(String text, String color) { + if (ansiColors) { + return color + text + ANSI_RESET; + } + return text; + } +} diff --git a/agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkComparisonTest.java b/agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkComparisonTest.java new file mode 100644 index 0000000..5fd530c --- /dev/null +++ b/agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkComparisonTest.java @@ -0,0 +1,88 @@ +package com.agenteval.reporting.benchmark; + +import com.agenteval.core.benchmark.BenchmarkResult; +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import com.agenteval.reporting.regression.RegressionReport; +import org.junit.jupiter.api.Test; + +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class BenchmarkComparisonTest { + + @Test + void compareTwoVariants() { + BenchmarkResult br = makeBenchmarkResult(); + RegressionReport report = BenchmarkComparison.compareVariants(br, "baseline", "candidate"); + + assertThat(report.overallBaselineScore()).isGreaterThan(0); + assertThat(report.overallCurrentScore()).isGreaterThan(0); + } + + @Test + void compareAllAgainstBaseline() { + BenchmarkResult br = makeBenchmarkResult(); + Map reports = + BenchmarkComparison.compareAllAgainst(br, "baseline"); + + assertThat(reports).containsKey("candidate"); + assertThat(reports).doesNotContainKey("baseline"); + } + + @Test + void unknownBaselineThrows() { + BenchmarkResult br = makeBenchmarkResult(); + + assertThatThrownBy(() -> BenchmarkComparison.compareVariants(br, "nope", "candidate")) + .isInstanceOf(IllegalArgumentException.class); + } + + @Test + void regressionDetected() { + Map results = new LinkedHashMap<>(); + results.put("good", makeResult("q", 0.9)); + results.put("bad", makeResult("q", 0.3)); + + BenchmarkResult br = new BenchmarkResult(results, 200L); + RegressionReport report = BenchmarkComparison.compareVariants(br, "good", "bad"); + + assertThat(report.hasRegressions()).isTrue(); + assertThat(report.newFailures()).isEqualTo(1); + } + + @Test + void improvementDetected() { + Map results = new LinkedHashMap<>(); + results.put("old", makeResult("q", 0.5)); + results.put("new", makeResult("q", 0.95)); + + BenchmarkResult br = new BenchmarkResult(results, 200L); + RegressionReport report = BenchmarkComparison.compareVariants(br, "old", "new"); + + assertThat(report.hasRegressions()).isFalse(); + assertThat(report.newPasses()).isEqualTo(1); + } + + private static BenchmarkResult makeBenchmarkResult() { + Map results = new LinkedHashMap<>(); + results.put("baseline", makeResult("q", 0.8)); + results.put("candidate", makeResult("q", 0.85)); + return new BenchmarkResult(results, 300L); + } + + private static EvalResult makeResult(String input, double score) { + AgentTestCase tc = AgentTestCase.builder() + .input(input).actualOutput("a").build(); + boolean passed = score >= 0.7; + EvalScore s = new EvalScore(score, 0.7, passed, "test", "M1"); + CaseResult cr = new CaseResult(tc, Map.of("M1", s), passed); + return EvalResult.of(List.of(cr), 100L); + } +} diff --git a/agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkReporterTest.java b/agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkReporterTest.java new file mode 100644 index 0000000..1f48804 --- /dev/null +++ b/agenteval-reporting/src/test/java/com/agenteval/reporting/benchmark/BenchmarkReporterTest.java @@ -0,0 +1,97 @@ +package com.agenteval.reporting.benchmark; + +import com.agenteval.core.benchmark.BenchmarkResult; +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import org.junit.jupiter.api.Test; + +import java.io.ByteArrayOutputStream; +import java.io.PrintStream; +import java.nio.charset.StandardCharsets; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; + +class BenchmarkReporterTest { + + @Test + void reportsTableWithBestWorstLabels() { + BenchmarkResult result = makeBenchmarkResult(0.9, 0.4); + String output = capture(result, false); + + assertThat(output).contains("=== Benchmark Results ==="); + assertThat(output).contains("high"); + assertThat(output).contains("low"); + assertThat(output).contains("[BEST]"); + assertThat(output).contains("[WORST]"); + } + + @Test + void perMetricBreakdownShown() { + BenchmarkResult result = makeBenchmarkResult(0.9, 0.4); + String output = capture(result, false); + + assertThat(output).contains("--- Per-Metric Breakdown ---"); + assertThat(output).contains("M1"); + } + + @Test + void singleVariantNoLabels() { + Map results = new LinkedHashMap<>(); + results.put("only", makeResult(0.7)); + + BenchmarkResult br = new BenchmarkResult(results, 100L); + String output = capture(br, false); + + assertThat(output).contains("only"); + assertThat(output).doesNotContain("[BEST]"); + assertThat(output).doesNotContain("[WORST]"); + } + + @Test + void ansiColorsToggle() { + BenchmarkResult result = makeBenchmarkResult(0.9, 0.4); + + String withAnsi = capture(result, true); + String withoutAnsi = capture(result, false); + + // ANSI version should have escape codes + assertThat(withAnsi).contains("\u001B["); + assertThat(withoutAnsi).doesNotContain("\u001B["); + } + + @Test + void showsVariantCountAndDuration() { + BenchmarkResult result = makeBenchmarkResult(0.9, 0.4); + String output = capture(result, false); + + assertThat(output).contains("Variants: 2"); + assertThat(output).contains("Duration:"); + } + + private String capture(BenchmarkResult result, boolean ansi) { + var baos = new ByteArrayOutputStream(); + var ps = new PrintStream(baos, true, StandardCharsets.UTF_8); + new BenchmarkReporter(ps, ansi).report(result); + return baos.toString(StandardCharsets.UTF_8); + } + + private static BenchmarkResult makeBenchmarkResult(double highScore, double lowScore) { + Map results = new LinkedHashMap<>(); + results.put("high", makeResult(highScore)); + results.put("low", makeResult(lowScore)); + return new BenchmarkResult(results, 500L); + } + + private static EvalResult makeResult(double score) { + AgentTestCase tc = AgentTestCase.builder() + .input("q").actualOutput("a").build(); + EvalScore s = new EvalScore(score, 0.7, score >= 0.7, "test", "M1"); + CaseResult cr = new CaseResult(tc, Map.of("M1", s), score >= 0.7); + return EvalResult.of(List.of(cr), 100L); + } +} From 9fddb70cb35395b4ed22892faf67f0b22e70c0a7 Mon Sep 17 00:00:00 2001 From: Pratyush Sharma <56130065+pratyush618@users.noreply.github.com> Date: Fri, 13 Mar 2026 11:35:23 +0530 Subject: [PATCH 3/4] Add Gradle plugin module mirroring Maven plugin New agenteval-gradle-plugin with plugin ID com.agenteval.evaluate. AgentEvalPlugin registers an agenteval extension and agentEvaluate task in the verification group. Uses dev.gradleplugins:gradle-api:8.5 from Maven Central. MetricResolver and ReportFormatResolver are duplicated from the Maven plugin (small, not worth a shared module). --- agenteval-gradle-plugin/pom.xml | 62 +++++++ .../agenteval/gradle/AgentEvalExtension.java | 58 ++++++ .../com/agenteval/gradle/AgentEvalPlugin.java | 51 ++++++ .../com/agenteval/gradle/EvaluateTask.java | 170 ++++++++++++++++++ .../com/agenteval/gradle/MetricResolver.java | 97 ++++++++++ .../gradle/ReportFormatResolver.java | 53 ++++++ .../com.agenteval.evaluate.properties | 1 + .../agenteval/gradle/AgentEvalPluginTest.java | 56 ++++++ .../agenteval/gradle/EvaluateTaskTest.java | 82 +++++++++ .../agenteval/gradle/MetricResolverTest.java | 92 ++++++++++ .../gradle/ReportFormatResolverTest.java | 76 ++++++++ pom.xml | 2 + 12 files changed, 800 insertions(+) create mode 100644 agenteval-gradle-plugin/pom.xml create mode 100644 agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalExtension.java create mode 100644 agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalPlugin.java create mode 100644 agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/EvaluateTask.java create mode 100644 agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/MetricResolver.java create mode 100644 agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/ReportFormatResolver.java create mode 100644 agenteval-gradle-plugin/src/main/resources/META-INF/gradle-plugins/com.agenteval.evaluate.properties create mode 100644 agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/AgentEvalPluginTest.java create mode 100644 agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/EvaluateTaskTest.java create mode 100644 agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/MetricResolverTest.java create mode 100644 agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/ReportFormatResolverTest.java diff --git a/agenteval-gradle-plugin/pom.xml b/agenteval-gradle-plugin/pom.xml new file mode 100644 index 0000000..75da74f --- /dev/null +++ b/agenteval-gradle-plugin/pom.xml @@ -0,0 +1,62 @@ + + + 4.0.0 + + + com.agenteval + agenteval-parent + 0.1.0-SNAPSHOT + + + agenteval-gradle-plugin + AgentEval Gradle Plugin + Gradle plugin for running AgentEval evaluations + + + + dev.gradleplugins + gradle-api + 8.5 + provided + + + com.agenteval + agenteval-core + + + com.agenteval + agenteval-judge + + + com.agenteval + agenteval-metrics + + + com.agenteval + agenteval-datasets + + + com.agenteval + agenteval-reporting + + + org.slf4j + slf4j-api + + + + + + + + com.github.spotbugs + spotbugs-maven-plugin + + true + + + + + diff --git a/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalExtension.java b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalExtension.java new file mode 100644 index 0000000..1f08dfe --- /dev/null +++ b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalExtension.java @@ -0,0 +1,58 @@ +package com.agenteval.gradle; + +import org.gradle.api.provider.Property; + +/** + * Extension for the AgentEval Gradle plugin. + * + *

Conventions (defaults) are set by {@link AgentEvalPlugin#apply}.

+ * + *
{@code
+ * agenteval {
+ *     datasetPath = 'src/test/resources/golden-set.json'
+ *     configFile = 'agenteval.yaml'
+ *     reportFormats = 'console,json'
+ *     outputDirectory = 'build/agenteval'
+ *     failOnRegression = false
+ *     threshold = 0.7
+ *     metrics = 'AnswerRelevancy'
+ * }
+ * }
+ */ +public abstract class AgentEvalExtension { + + /** + * Path to the evaluation dataset file (required). + */ + public abstract Property getDatasetPath(); + + /** + * Path to the YAML configuration file. Default: "agenteval.yaml". + */ + public abstract Property getConfigFile(); + + /** + * Comma-separated report format names. Default: "console,json". + */ + public abstract Property getReportFormats(); + + /** + * Output directory for reports. Default: "build/agenteval". + */ + public abstract Property getOutputDirectory(); + + /** + * Whether to fail the build on regression. Default: false. + */ + public abstract Property getFailOnRegression(); + + /** + * Pass/fail threshold for scores. Default: 0.7. + */ + public abstract Property getThreshold(); + + /** + * Comma-separated metric names. Default: "AnswerRelevancy". + */ + public abstract Property getMetrics(); +} diff --git a/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalPlugin.java b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalPlugin.java new file mode 100644 index 0000000..ed677f0 --- /dev/null +++ b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/AgentEvalPlugin.java @@ -0,0 +1,51 @@ +package com.agenteval.gradle; + +import org.gradle.api.Plugin; +import org.gradle.api.Project; + +/** + * Gradle plugin for running AgentEval evaluations. + * + *

Registers the {@code agenteval} extension and an {@code agentEvaluate} task + * in the {@code verification} group.

+ * + *
{@code
+ * plugins {
+ *     id 'com.agenteval.evaluate'
+ * }
+ *
+ * agenteval {
+ *     datasetPath = 'src/test/resources/golden-set.json'
+ *     metrics = 'AnswerRelevancy,Faithfulness'
+ *     reportFormats = 'console,json'
+ * }
+ * }
+ */ +public class AgentEvalPlugin implements Plugin { + + @Override + public void apply(Project project) { + AgentEvalExtension extension = project.getExtensions() + .create("agenteval", AgentEvalExtension.class); + + // Set convention defaults on the extension + extension.getConfigFile().convention("agenteval.yaml"); + extension.getReportFormats().convention("console,json"); + extension.getOutputDirectory().convention("build/agenteval"); + extension.getFailOnRegression().convention(false); + extension.getThreshold().convention(0.7); + extension.getMetrics().convention("AnswerRelevancy"); + + project.getTasks().register("agentEvaluate", EvaluateTask.class, task -> { + task.setGroup("verification"); + task.setDescription("Runs AgentEval evaluations against a dataset"); + task.getDatasetPath().convention(extension.getDatasetPath()); + task.getConfigFile().convention(extension.getConfigFile()); + task.getReportFormats().convention(extension.getReportFormats()); + task.getOutputDirectory().convention(extension.getOutputDirectory()); + task.getFailOnRegression().convention(extension.getFailOnRegression()); + task.getThreshold().convention(extension.getThreshold()); + task.getMetrics().convention(extension.getMetrics()); + }); + } +} diff --git a/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/EvaluateTask.java b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/EvaluateTask.java new file mode 100644 index 0000000..8faaf95 --- /dev/null +++ b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/EvaluateTask.java @@ -0,0 +1,170 @@ +package com.agenteval.gradle; + +import com.agenteval.core.config.AgentEvalConfigLoader; +import com.agenteval.core.config.YamlConfigModel; +import com.agenteval.core.eval.CaseResult; +import com.agenteval.core.eval.EvalResult; +import com.agenteval.core.judge.JudgeModel; +import com.agenteval.core.metric.EvalMetric; +import com.agenteval.core.model.AgentTestCase; +import com.agenteval.core.model.EvalScore; +import com.agenteval.datasets.DatasetLoaders; +import com.agenteval.datasets.EvalDataset; +import com.agenteval.judge.JudgeModels; +import com.agenteval.reporting.EvalReporter; +import org.gradle.api.DefaultTask; +import org.gradle.api.GradleException; +import org.gradle.api.provider.Property; +import org.gradle.api.tasks.Input; +import org.gradle.api.tasks.Optional; +import org.gradle.api.tasks.TaskAction; + +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.HashMap; +import java.util.List; +import java.util.Locale; +import java.util.Map; + +/** + * Gradle task that runs AgentEval evaluations against a dataset. + */ +public abstract class EvaluateTask extends DefaultTask { + + @Input + public abstract Property getDatasetPath(); + + @Input + @Optional + public abstract Property getConfigFile(); + + @Input + @Optional + public abstract Property getReportFormats(); + + @Input + @Optional + public abstract Property getOutputDirectory(); + + @Input + @Optional + public abstract Property getFailOnRegression(); + + @Input + @Optional + public abstract Property getThreshold(); + + @Input + @Optional + public abstract Property getMetrics(); + + @TaskAction + public void execute() { + try { + String dsPathStr = getDatasetPath().get(); + getLogger().lifecycle("AgentEval: Loading dataset from " + dsPathStr); + + Path dsPath = Path.of(dsPathStr); + if (!Files.exists(dsPath)) { + throw new GradleException("Dataset not found: " + dsPathStr); + } + + EvalDataset dataset = DatasetLoaders.forPath(dsPath); + getLogger().lifecycle("AgentEval: Loaded " + dataset.size() + " test cases"); + + JudgeModel judge = resolveJudge(); + + List evalMetrics = resolveMetrics(judge); + getLogger().lifecycle("AgentEval: Running " + evalMetrics.size() + " metrics"); + + long start = System.currentTimeMillis(); + List caseResults = new ArrayList<>(); + for (AgentTestCase testCase : dataset.getTestCases()) { + Map scores = new HashMap<>(); + boolean allPassed = true; + for (EvalMetric metric : evalMetrics) { + EvalScore score = metric.evaluate(testCase); + score = score.withMetricName(metric.name()); + scores.put(metric.name(), score); + if (!score.passed()) { + allPassed = false; + } + } + caseResults.add(new CaseResult(testCase, scores, allPassed)); + } + long duration = System.currentTimeMillis() - start; + EvalResult result = EvalResult.of(caseResults, duration); + + Path outDir = Path.of(getOutputDirectory().get()); + Files.createDirectories(outDir); + List reporters = + ReportFormatResolver.resolve(getReportFormats().get(), outDir); + for (EvalReporter reporter : reporters) { + reporter.report(result); + } + + if (Boolean.TRUE.equals(getFailOnRegression().getOrNull()) + && result.passRate() < 1.0) { + throw new GradleException(String.format( + "AgentEval: Evaluation failed — pass rate %.1f%% (threshold: 100%%)", + result.passRate() * 100)); + } + + getLogger().lifecycle(String.format( + "AgentEval: Completed — %.1f%% pass rate, avg score %.3f", + result.passRate() * 100, result.averageScore())); + + } catch (GradleException e) { + throw e; + } catch (Exception e) { + throw new GradleException("AgentEval evaluation failed", e); + } + } + + @SuppressWarnings("IllegalCatch") + private JudgeModel resolveJudge() { + String cfgFile = getConfigFile().getOrElse("agenteval.yaml"); + Path cfgPath = Path.of(cfgFile); + if (Files.exists(cfgPath)) { + try { + YamlConfigModel model = AgentEvalConfigLoader.loadModel(cfgPath); + if (model.getJudge() != null) { + String provider = model.getJudge().getProvider(); + String modelName = model.getJudge().getModel(); + if (provider != null && modelName != null) { + return createJudge(provider, modelName); + } + } + } catch (Exception e) { + getLogger().debug("Could not load config: " + e.getMessage()); + } + } + + String provider = System.getenv("AGENTEVAL_JUDGE_PROVIDER"); + String model = System.getenv("AGENTEVAL_JUDGE_MODEL"); + if (provider != null && model != null) { + return createJudge(provider, model); + } + + return null; + } + + private JudgeModel createJudge(String provider, String model) { + return switch (provider.toLowerCase(Locale.ROOT)) { + case "openai" -> JudgeModels.openai(model); + case "anthropic" -> JudgeModels.anthropic(model); + case "ollama" -> JudgeModels.ollama(model); + default -> throw new IllegalArgumentException("Unknown judge provider: " + provider); + }; + } + + private List resolveMetrics(JudgeModel judge) { + List resolved = new ArrayList<>(); + String metricsStr = getMetrics().getOrElse("AnswerRelevancy"); + for (String name : metricsStr.split(",")) { + resolved.add(MetricResolver.resolve(name.trim(), judge)); + } + return resolved; + } +} diff --git a/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/MetricResolver.java b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/MetricResolver.java new file mode 100644 index 0000000..b440e24 --- /dev/null +++ b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/MetricResolver.java @@ -0,0 +1,97 @@ +package com.agenteval.gradle; + +import com.agenteval.core.judge.JudgeModel; +import com.agenteval.core.metric.EvalMetric; +import com.agenteval.metrics.agent.ToolSelectionAccuracyMetric; +import com.agenteval.metrics.response.AnswerRelevancyMetric; +import com.agenteval.metrics.response.BiasMetric; +import com.agenteval.metrics.response.CoherenceMetric; +import com.agenteval.metrics.response.ConcisenessMetric; +import com.agenteval.metrics.response.CorrectnessMetric; +import com.agenteval.metrics.response.FaithfulnessMetric; +import com.agenteval.metrics.response.HallucinationMetric; +import com.agenteval.metrics.response.ToxicityMetric; +import com.agenteval.metrics.rag.ContextualPrecisionMetric; +import com.agenteval.metrics.rag.ContextualRecallMetric; +import com.agenteval.metrics.rag.ContextualRelevancyMetric; +import com.agenteval.metrics.agent.TaskCompletionMetric; + +import java.util.Locale; +import java.util.Map; +import java.util.function.Function; + +/** + * Resolves metric name strings to {@link EvalMetric} instances. + * + *

Metric names are case-insensitive. All LLM-based metrics require a {@link JudgeModel}.

+ */ +public final class MetricResolver { + + private static final Map> LLM_METRICS = Map.ofEntries( + entry("answerrelevancy", AnswerRelevancyMetric::new), + entry("faithfulness", FaithfulnessMetric::new), + entry("correctness", CorrectnessMetric::new), + entry("hallucination", HallucinationMetric::new), + entry("toxicity", ToxicityMetric::new), + entry("coherence", CoherenceMetric::new), + entry("conciseness", ConcisenessMetric::new), + entry("bias", BiasMetric::new), + entry("contextualrelevancy", ContextualRelevancyMetric::new), + entry("contextualprecision", ContextualPrecisionMetric::new), + entry("contextualrecall", ContextualRecallMetric::new), + entry("taskcompletion", TaskCompletionMetric::new) + ); + + private static final Map STANDALONE_METRICS = Map.of( + "toolselectionaccuracy", new ToolSelectionAccuracyMetric() + ); + + private MetricResolver() {} + + /** + * Resolves a metric by name. LLM-based metrics use the provided judge. + * + * @param name the metric name (case-insensitive) + * @param judge the judge model for LLM-based metrics (may be null for standalone metrics) + * @return the resolved metric + * @throws IllegalArgumentException if the metric name is unknown + */ + public static EvalMetric resolve(String name, JudgeModel judge) { + String key = normalize(name); + + EvalMetric standalone = STANDALONE_METRICS.get(key); + if (standalone != null) { + return standalone; + } + + Function factory = LLM_METRICS.get(key); + if (factory != null) { + if (judge == null) { + throw new IllegalArgumentException( + "Metric '" + name + "' requires a judge model"); + } + return factory.apply(judge); + } + + throw new IllegalArgumentException("Unknown metric: " + name + + ". Available: " + availableMetrics()); + } + + /** + * Returns comma-separated list of all known metric names. + */ + public static String availableMetrics() { + var all = new java.util.TreeSet(); + all.addAll(LLM_METRICS.keySet()); + all.addAll(STANDALONE_METRICS.keySet()); + return String.join(", ", all); + } + + private static String normalize(String name) { + return name.toLowerCase(Locale.ROOT).replace("_", "").replace("-", ""); + } + + private static Map.Entry entry(String key, T value) { + return Map.entry(key, value); + } +} diff --git a/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/ReportFormatResolver.java b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/ReportFormatResolver.java new file mode 100644 index 0000000..13ed907 --- /dev/null +++ b/agenteval-gradle-plugin/src/main/java/com/agenteval/gradle/ReportFormatResolver.java @@ -0,0 +1,53 @@ +package com.agenteval.gradle; + +import com.agenteval.reporting.ConsoleReporter; +import com.agenteval.reporting.EvalReporter; +import com.agenteval.reporting.HtmlReportConfig; +import com.agenteval.reporting.HtmlReporter; +import com.agenteval.reporting.JsonReporter; +import com.agenteval.reporting.JunitXmlReporter; + +import java.nio.file.Path; +import java.util.ArrayList; +import java.util.List; +import java.util.Locale; + +/** + * Resolves report format strings to {@link EvalReporter} instances. + * + *

Supported formats: console, json, xml, html.

+ */ +public final class ReportFormatResolver { + + private ReportFormatResolver() {} + + /** + * Resolves a list of format names to reporter instances. + * + * @param formats comma-separated format names + * @param outputDirectory the directory for file-based reports + * @return the list of reporters + * @throws IllegalArgumentException if a format name is unknown + */ + public static List resolve(String formats, Path outputDirectory) { + List reporters = new ArrayList<>(); + for (String format : formats.split(",")) { + reporters.add(resolveOne(format.trim(), outputDirectory)); + } + return reporters; + } + + private static EvalReporter resolveOne(String format, Path outputDirectory) { + return switch (format.toLowerCase(Locale.ROOT)) { + case "console" -> new ConsoleReporter(); + case "json" -> new JsonReporter(outputDirectory.resolve("agenteval-report.json")); + case "xml" -> new JunitXmlReporter(outputDirectory.resolve("agenteval-report.xml")); + case "html" -> new HtmlReporter(HtmlReportConfig.builder() + .outputPath(outputDirectory.resolve("agenteval-report.html")) + .build()); + default -> throw new IllegalArgumentException( + "Unknown report format: " + format + + ". Supported: console, json, xml, html"); + }; + } +} diff --git a/agenteval-gradle-plugin/src/main/resources/META-INF/gradle-plugins/com.agenteval.evaluate.properties b/agenteval-gradle-plugin/src/main/resources/META-INF/gradle-plugins/com.agenteval.evaluate.properties new file mode 100644 index 0000000..d5fa53d --- /dev/null +++ b/agenteval-gradle-plugin/src/main/resources/META-INF/gradle-plugins/com.agenteval.evaluate.properties @@ -0,0 +1 @@ +implementation-class=com.agenteval.gradle.AgentEvalPlugin diff --git a/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/AgentEvalPluginTest.java b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/AgentEvalPluginTest.java new file mode 100644 index 0000000..7685f22 --- /dev/null +++ b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/AgentEvalPluginTest.java @@ -0,0 +1,56 @@ +package com.agenteval.gradle; + +import org.gradle.api.Project; +import org.gradle.testfixtures.ProjectBuilder; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; + +class AgentEvalPluginTest { + + @Test + void pluginAppliesSuccessfully() { + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + assertThat(project.getPlugins().hasPlugin(AgentEvalPlugin.class)).isTrue(); + } + + @Test + void extensionRegistered() { + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + AgentEvalExtension extension = project.getExtensions() + .findByType(AgentEvalExtension.class); + assertThat(extension).isNotNull(); + } + + @Test + void taskRegisteredWithCorrectTypeAndGroup() { + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + var task = project.getTasks().findByName("agentEvaluate"); + assertThat(task).isNotNull(); + assertThat(task).isInstanceOf(EvaluateTask.class); + assertThat(task.getGroup()).isEqualTo("verification"); + } + + @Test + void extensionDefaultValues() { + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + AgentEvalExtension ext = project.getExtensions() + .findByType(AgentEvalExtension.class); + assertThat(ext).isNotNull(); + + assertThat(ext.getConfigFile().get()).isEqualTo("agenteval.yaml"); + assertThat(ext.getReportFormats().get()).isEqualTo("console,json"); + assertThat(ext.getOutputDirectory().get()).isEqualTo("build/agenteval"); + assertThat(ext.getFailOnRegression().get()).isFalse(); + assertThat(ext.getThreshold().get()).isEqualTo(0.7); + assertThat(ext.getMetrics().get()).isEqualTo("AnswerRelevancy"); + } +} diff --git a/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/EvaluateTaskTest.java b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/EvaluateTaskTest.java new file mode 100644 index 0000000..ac04cca --- /dev/null +++ b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/EvaluateTaskTest.java @@ -0,0 +1,82 @@ +package com.agenteval.gradle; + +import org.gradle.api.Project; +import org.gradle.testfixtures.ProjectBuilder; +import org.junit.jupiter.api.BeforeAll; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.junit.jupiter.api.Assumptions.assumeTrue; + +class EvaluateTaskTest { + + private static boolean gradleTaskCreationSupported; + + @BeforeAll + static void checkEnvironment() { + // Gradle's ProjectBuilder may fail to inject synthetic classes on newer JDKs + // without --add-opens. Guard tests with an assumption. + try { + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + EvaluateTask task = (EvaluateTask) project.getTasks().findByName("agentEvaluate"); + gradleTaskCreationSupported = task != null; + } catch (Exception e) { + gradleTaskCreationSupported = false; + } + } + + @Test + void taskDefaultPropertyValues() { + assumeTrue(gradleTaskCreationSupported, + "Gradle task creation not supported in this JVM configuration"); + + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + EvaluateTask task = (EvaluateTask) project.getTasks().findByName("agentEvaluate"); + assertThat(task).isNotNull(); + + assertThat(task.getConfigFile().get()).isEqualTo("agenteval.yaml"); + assertThat(task.getReportFormats().get()).isEqualTo("console,json"); + assertThat(task.getOutputDirectory().get()).isEqualTo("build/agenteval"); + assertThat(task.getFailOnRegression().get()).isFalse(); + assertThat(task.getThreshold().get()).isEqualTo(0.7); + assertThat(task.getMetrics().get()).isEqualTo("AnswerRelevancy"); + } + + @Test + void datasetPathRequiredValidation() { + assumeTrue(gradleTaskCreationSupported, + "Gradle task creation not supported in this JVM configuration"); + + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + EvaluateTask task = (EvaluateTask) project.getTasks().findByName("agentEvaluate"); + assertThat(task).isNotNull(); + + assertThat(task.getDatasetPath().isPresent()).isFalse(); + } + + @Test + void extensionOverridesWireToTask() { + assumeTrue(gradleTaskCreationSupported, + "Gradle task creation not supported in this JVM configuration"); + + Project project = ProjectBuilder.builder().build(); + project.getPluginManager().apply("com.agenteval.evaluate"); + + AgentEvalExtension ext = project.getExtensions() + .findByType(AgentEvalExtension.class); + assertThat(ext).isNotNull(); + + ext.getMetrics().set("Faithfulness,Correctness"); + ext.getThreshold().set(0.8); + + EvaluateTask task = (EvaluateTask) project.getTasks().findByName("agentEvaluate"); + assertThat(task).isNotNull(); + assertThat(task.getMetrics().get()).isEqualTo("Faithfulness,Correctness"); + assertThat(task.getThreshold().get()).isEqualTo(0.8); + } +} diff --git a/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/MetricResolverTest.java b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/MetricResolverTest.java new file mode 100644 index 0000000..24cf4d1 --- /dev/null +++ b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/MetricResolverTest.java @@ -0,0 +1,92 @@ +package com.agenteval.gradle; + +import com.agenteval.core.judge.JudgeModel; +import com.agenteval.core.judge.JudgeResponse; +import com.agenteval.core.metric.EvalMetric; +import com.agenteval.core.model.TokenUsage; +import org.junit.jupiter.api.Test; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class MetricResolverTest { + + private final JudgeModel stubJudge = new JudgeModel() { + @Override + public JudgeResponse judge(String prompt) { + return new JudgeResponse(0.8, "good", TokenUsage.of(10, 5)); + } + + @Override + public String modelId() { return "stub"; } + }; + + @Test + void shouldResolveLlmMetricByName() { + EvalMetric metric = MetricResolver.resolve("AnswerRelevancy", stubJudge); + + assertThat(metric).isNotNull(); + assertThat(metric.name()).isEqualTo("AnswerRelevancy"); + } + + @Test + void shouldResolveCaseInsensitively() { + EvalMetric metric = MetricResolver.resolve("answerrelevancy", stubJudge); + + assertThat(metric).isNotNull(); + assertThat(metric.name()).isEqualTo("AnswerRelevancy"); + } + + @Test + void shouldResolveWithHyphensAndUnderscores() { + EvalMetric metric = MetricResolver.resolve("answer-relevancy", stubJudge); + assertThat(metric.name()).isEqualTo("AnswerRelevancy"); + + EvalMetric metric2 = MetricResolver.resolve("answer_relevancy", stubJudge); + assertThat(metric2.name()).isEqualTo("AnswerRelevancy"); + } + + @Test + void shouldResolveStandaloneMetricWithoutJudge() { + EvalMetric metric = MetricResolver.resolve("ToolSelectionAccuracy", null); + + assertThat(metric).isNotNull(); + assertThat(metric.name()).isEqualTo("ToolSelectionAccuracy"); + } + + @Test + void shouldThrowForUnknownMetric() { + assertThatThrownBy(() -> MetricResolver.resolve("NonexistentMetric", stubJudge)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Unknown metric"); + } + + @Test + void shouldThrowWhenLlmMetricHasNoJudge() { + assertThatThrownBy(() -> MetricResolver.resolve("Faithfulness", null)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("requires a judge model"); + } + + @Test + void shouldResolveAllKnownLlmMetrics() { + String[] metricNames = { + "Faithfulness", "Correctness", "Hallucination", + "Toxicity", "Coherence", "Conciseness", "Bias", + "ContextualRelevancy", "ContextualPrecision", "ContextualRecall", + "TaskCompletion" + }; + for (String name : metricNames) { + EvalMetric metric = MetricResolver.resolve(name, stubJudge); + assertThat(metric).as("Metric: " + name).isNotNull(); + } + } + + @Test + void availableMetricsShouldReturnNonEmptyString() { + String available = MetricResolver.availableMetrics(); + + assertThat(available).contains("answerrelevancy"); + assertThat(available).contains("faithfulness"); + } +} diff --git a/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/ReportFormatResolverTest.java b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/ReportFormatResolverTest.java new file mode 100644 index 0000000..8bfa15c --- /dev/null +++ b/agenteval-gradle-plugin/src/test/java/com/agenteval/gradle/ReportFormatResolverTest.java @@ -0,0 +1,76 @@ +package com.agenteval.gradle; + +import com.agenteval.reporting.ConsoleReporter; +import com.agenteval.reporting.EvalReporter; +import com.agenteval.reporting.HtmlReporter; +import com.agenteval.reporting.JsonReporter; +import com.agenteval.reporting.JunitXmlReporter; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import java.nio.file.Path; +import java.util.List; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +class ReportFormatResolverTest { + + @Test + void shouldResolveConsoleFormat(@TempDir Path tempDir) { + List reporters = ReportFormatResolver.resolve("console", tempDir); + + assertThat(reporters).hasSize(1); + assertThat(reporters.getFirst()).isInstanceOf(ConsoleReporter.class); + } + + @Test + void shouldResolveJsonFormat(@TempDir Path tempDir) { + List reporters = ReportFormatResolver.resolve("json", tempDir); + + assertThat(reporters).hasSize(1); + assertThat(reporters.getFirst()).isInstanceOf(JsonReporter.class); + } + + @Test + void shouldResolveXmlFormat(@TempDir Path tempDir) { + List reporters = ReportFormatResolver.resolve("xml", tempDir); + + assertThat(reporters).hasSize(1); + assertThat(reporters.getFirst()).isInstanceOf(JunitXmlReporter.class); + } + + @Test + void shouldResolveHtmlFormat(@TempDir Path tempDir) { + List reporters = ReportFormatResolver.resolve("html", tempDir); + + assertThat(reporters).hasSize(1); + assertThat(reporters.getFirst()).isInstanceOf(HtmlReporter.class); + } + + @Test + void shouldResolveMultipleFormats(@TempDir Path tempDir) { + List reporters = ReportFormatResolver.resolve("console,json,xml", tempDir); + + assertThat(reporters).hasSize(3); + assertThat(reporters.get(0)).isInstanceOf(ConsoleReporter.class); + assertThat(reporters.get(1)).isInstanceOf(JsonReporter.class); + assertThat(reporters.get(2)).isInstanceOf(JunitXmlReporter.class); + } + + @Test + void shouldBeCaseInsensitive(@TempDir Path tempDir) { + List reporters = ReportFormatResolver.resolve("CONSOLE,JSON", tempDir); + + assertThat(reporters).hasSize(2); + assertThat(reporters.get(0)).isInstanceOf(ConsoleReporter.class); + assertThat(reporters.get(1)).isInstanceOf(JsonReporter.class); + } + + @Test + void shouldThrowForUnknownFormat(@TempDir Path tempDir) { + assertThatThrownBy(() -> ReportFormatResolver.resolve("pdf", tempDir)) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Unknown report format"); + } +} diff --git a/pom.xml b/pom.xml index b8fc50b..320a5b4 100644 --- a/pom.xml +++ b/pom.xml @@ -35,6 +35,8 @@ agenteval-redteam agenteval-maven-plugin agenteval-github-actions + agenteval-gradle-plugin + agenteval-intellij From f9bbcc297891a4cf0da91274b1f1c4ddaad4b05e Mon Sep 17 00:00:00 2001 From: Pratyush Sharma <56130065+pratyush618@users.noreply.github.com> Date: Fri, 13 Mar 2026 11:36:30 +0530 Subject: [PATCH 4/4] Add IntelliJ IDEA plugin module for report viewing New agenteval-intellij with lightweight JSON report parser (no agenteval-core dependency). ReportModel/ReportParser parse agenteval-report.json using Jackson. Includes tool window, gutter icon provider for @Metric annotations, VFS file watcher, SVG icons, and plugin.xml descriptor. IntelliJ Platform-dependent UI sources are excluded from Maven compilation (require IDE SDK). --- agenteval-intellij/pom.xml | 65 ++++++ .../agenteval/intellij/AgentEvalIcons.java | 27 +++ .../intellij/AgentEvalToolWindow.java | 123 ++++++++++ .../intellij/AgentEvalToolWindowFactory.java | 23 ++ .../intellij/MetricGutterIconProvider.java | 101 +++++++++ .../agenteval/intellij/ReportFileWatcher.java | 57 +++++ .../com/agenteval/intellij/ReportModel.java | 106 +++++++++ .../com/agenteval/intellij/ReportParser.java | 69 ++++++ .../src/main/resources/META-INF/plugin.xml | 20 ++ .../main/resources/icons/agenteval-fail.svg | 4 + .../main/resources/icons/agenteval-pass.svg | 4 + .../resources/icons/agenteval-toolwindow.svg | 4 + .../agenteval/intellij/ReportModelTest.java | 113 ++++++++++ .../agenteval/intellij/ReportParserTest.java | 210 ++++++++++++++++++ 14 files changed, 926 insertions(+) create mode 100644 agenteval-intellij/pom.xml create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalIcons.java create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindow.java create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindowFactory.java create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/MetricGutterIconProvider.java create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/ReportFileWatcher.java create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/ReportModel.java create mode 100644 agenteval-intellij/src/main/java/com/agenteval/intellij/ReportParser.java create mode 100644 agenteval-intellij/src/main/resources/META-INF/plugin.xml create mode 100644 agenteval-intellij/src/main/resources/icons/agenteval-fail.svg create mode 100644 agenteval-intellij/src/main/resources/icons/agenteval-pass.svg create mode 100644 agenteval-intellij/src/main/resources/icons/agenteval-toolwindow.svg create mode 100644 agenteval-intellij/src/test/java/com/agenteval/intellij/ReportModelTest.java create mode 100644 agenteval-intellij/src/test/java/com/agenteval/intellij/ReportParserTest.java diff --git a/agenteval-intellij/pom.xml b/agenteval-intellij/pom.xml new file mode 100644 index 0000000..0fdb0c9 --- /dev/null +++ b/agenteval-intellij/pom.xml @@ -0,0 +1,65 @@ + + + 4.0.0 + + + com.agenteval + agenteval-parent + 0.1.0-SNAPSHOT + + + agenteval-intellij + AgentEval IntelliJ Plugin + IntelliJ IDEA plugin for viewing AgentEval evaluation results + + + + com.fasterxml.jackson.core + jackson-databind + provided + + + + + + + + com.github.spotbugs + spotbugs-maven-plugin + + true + + + + + org.apache.maven.plugins + maven-compiler-plugin + + + com/agenteval/intellij/AgentEvalIcons.java + com/agenteval/intellij/AgentEvalToolWindowFactory.java + com/agenteval/intellij/AgentEvalToolWindow.java + com/agenteval/intellij/ReportFileWatcher.java + com/agenteval/intellij/MetricGutterIconProvider.java + + + + + + org.apache.maven.plugins + maven-checkstyle-plugin + + + com/agenteval/intellij/AgentEvalIcons.java, + com/agenteval/intellij/AgentEvalToolWindowFactory.java, + com/agenteval/intellij/AgentEvalToolWindow.java, + com/agenteval/intellij/ReportFileWatcher.java, + com/agenteval/intellij/MetricGutterIconProvider.java + + + + + + diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalIcons.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalIcons.java new file mode 100644 index 0000000..13b63f4 --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalIcons.java @@ -0,0 +1,27 @@ +package com.agenteval.intellij; + +import com.intellij.openapi.util.IconLoader; + +import javax.swing.Icon; + +/** + * Icon constants for the AgentEval IntelliJ plugin. + * + *

Icons are loaded from the plugin's resource directory.

+ */ +public final class AgentEvalIcons { + + private AgentEvalIcons() {} + + /** Pass icon (green checkmark, 13x13). */ + public static final Icon PASS = IconLoader.getIcon( + "/icons/agenteval-pass.svg", AgentEvalIcons.class); + + /** Fail icon (red X, 13x13). */ + public static final Icon FAIL = IconLoader.getIcon( + "/icons/agenteval-fail.svg", AgentEvalIcons.class); + + /** Tool window icon (13x13). */ + public static final Icon TOOL_WINDOW = IconLoader.getIcon( + "/icons/agenteval-toolwindow.svg", AgentEvalIcons.class); +} diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindow.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindow.java new file mode 100644 index 0000000..7fee13b --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindow.java @@ -0,0 +1,123 @@ +package com.agenteval.intellij; + +import com.intellij.openapi.project.Project; +import com.intellij.ui.components.JBLabel; +import com.intellij.ui.components.JBScrollPane; +import com.intellij.ui.table.JBTable; + +import javax.swing.JComponent; +import javax.swing.JPanel; +import javax.swing.table.DefaultTableModel; +import java.awt.BorderLayout; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; + +/** + * Tool window panel showing AgentEval report results. + * + *

Displays a summary bar (pass rate, avg score, case count, duration) + * and a table of test case results with per-metric scores.

+ */ +public class AgentEvalToolWindow { + + private final JPanel content; + private final JBLabel summaryLabel; + private final JBTable resultTable; + private final Project project; + + public AgentEvalToolWindow(Project project) { + this.project = project; + this.content = new JPanel(new BorderLayout()); + this.summaryLabel = new JBLabel("No report loaded"); + this.resultTable = new JBTable(); + + content.add(summaryLabel, BorderLayout.NORTH); + content.add(new JBScrollPane(resultTable), BorderLayout.CENTER); + + refresh(); + } + + public JComponent getContent() { + return content; + } + + /** + * Refreshes the tool window by re-reading the report file. + */ + public void refresh() { + Path reportPath = findReportFile(); + if (reportPath == null || !Files.exists(reportPath)) { + summaryLabel.setText("No agenteval-report.json found"); + return; + } + + try { + ReportModel report = ReportParser.parseFile(reportPath); + updateSummary(report); + updateTable(report); + } catch (IOException e) { + summaryLabel.setText("Error reading report: " + e.getMessage()); + } + } + + private void updateSummary(ReportModel report) { + summaryLabel.setText(String.format( + "Pass Rate: %.1f%% | Avg Score: %.3f | Cases: %d | Duration: %dms", + report.getPassRate() * 100, + report.getAverageScore(), + report.getTotalCases(), + report.getDurationMs())); + } + + private void updateTable(ReportModel report) { + if (report.getCaseResults() == null || report.getCaseResults().isEmpty()) { + resultTable.setModel(new DefaultTableModel()); + return; + } + + // Collect all metric names + var firstCase = report.getCaseResults().getFirst(); + String[] metricNames = firstCase.getScores() != null + ? firstCase.getScores().keySet().toArray(String[]::new) + : new String[0]; + + // Build column names: Input, Passed, then metric columns + String[] columns = new String[2 + metricNames.length]; + columns[0] = "Input"; + columns[1] = "Passed"; + System.arraycopy(metricNames, 0, columns, 2, metricNames.length); + + // Build data + Object[][] data = new Object[report.getCaseResults().size()][columns.length]; + for (int i = 0; i < report.getCaseResults().size(); i++) { + var cr = report.getCaseResults().get(i); + data[i][0] = cr.getInput(); + data[i][1] = cr.isPassed() ? "PASS" : "FAIL"; + for (int j = 0; j < metricNames.length; j++) { + if (cr.getScores() != null && cr.getScores().containsKey(metricNames[j])) { + data[i][j + 2] = String.format("%.3f", + cr.getScores().get(metricNames[j]).getValue()); + } else { + data[i][j + 2] = "—"; + } + } + } + + resultTable.setModel(new DefaultTableModel(data, columns)); + } + + private Path findReportFile() { + if (project.getBasePath() == null) return null; + // Check common locations + Path buildDir = Path.of(project.getBasePath(), "build", "agenteval", + "agenteval-report.json"); + if (Files.exists(buildDir)) return buildDir; + + Path targetDir = Path.of(project.getBasePath(), "target", "agenteval", + "agenteval-report.json"); + if (Files.exists(targetDir)) return targetDir; + + return null; + } +} diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindowFactory.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindowFactory.java new file mode 100644 index 0000000..37518cc --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/AgentEvalToolWindowFactory.java @@ -0,0 +1,23 @@ +package com.agenteval.intellij; + +import com.intellij.openapi.project.Project; +import com.intellij.openapi.wm.ToolWindow; +import com.intellij.openapi.wm.ToolWindowFactory; +import com.intellij.ui.content.Content; +import com.intellij.ui.content.ContentFactory; +import org.jetbrains.annotations.NotNull; + +/** + * Factory for the AgentEval tool window in IntelliJ IDEA. + */ +public class AgentEvalToolWindowFactory implements ToolWindowFactory { + + @Override + public void createToolWindowContent(@NotNull Project project, + @NotNull ToolWindow toolWindow) { + AgentEvalToolWindow panel = new AgentEvalToolWindow(project); + ContentFactory contentFactory = ContentFactory.getInstance(); + Content content = contentFactory.createContent(panel.getContent(), "", false); + toolWindow.getContentManager().addContent(content); + } +} diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/MetricGutterIconProvider.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/MetricGutterIconProvider.java new file mode 100644 index 0000000..f043b20 --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/MetricGutterIconProvider.java @@ -0,0 +1,101 @@ +package com.agenteval.intellij; + +import com.intellij.codeInsight.daemon.LineMarkerInfo; +import com.intellij.codeInsight.daemon.LineMarkerProvider; +import com.intellij.openapi.editor.markup.GutterIconRenderer; +import com.intellij.openapi.project.Project; +import com.intellij.psi.PsiAnnotation; +import com.intellij.psi.PsiElement; +import com.intellij.psi.PsiIdentifier; +import com.intellij.psi.PsiModifierListOwner; +import org.jetbrains.annotations.NotNull; + +import javax.swing.Icon; +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Map; + +/** + * Shows pass/fail gutter icons on {@code @Metric} annotations by + * cross-referencing with the latest report. + */ +public class MetricGutterIconProvider implements LineMarkerProvider { + + @Override + public LineMarkerInfo getLineMarkerInfo(@NotNull PsiElement element) { + if (!(element instanceof PsiIdentifier)) { + return null; + } + + PsiElement parent = element.getParent(); + if (!(parent instanceof PsiModifierListOwner owner)) { + return null; + } + + PsiAnnotation metricAnnotation = owner.getAnnotation("com.agenteval.junit5.Metric"); + if (metricAnnotation == null) { + return null; + } + + // Extract metric name from annotation + String metricName = getAnnotationValue(metricAnnotation); + if (metricName == null) { + return null; + } + + // Check report for this metric + Map metricStatus = loadMetricStatus(element.getProject()); + Boolean passed = metricStatus.get(metricName); + if (passed == null) { + return null; + } + + Icon icon = passed ? AgentEvalIcons.PASS : AgentEvalIcons.FAIL; + String tooltip = metricName + ": " + (passed ? "PASS" : "FAIL"); + + return new LineMarkerInfo<>( + element, + element.getTextRange(), + icon, + e -> tooltip, + null, + GutterIconRenderer.Alignment.RIGHT, + () -> tooltip); + } + + private static String getAnnotationValue(PsiAnnotation annotation) { + var value = annotation.findAttributeValue("value"); + if (value == null) { + value = annotation.findAttributeValue(null); + } + if (value == null) return null; + String text = value.getText(); + // Strip quotes + if (text.startsWith("\"") && text.endsWith("\"")) { + return text.substring(1, text.length() - 1); + } + return text; + } + + private static Map loadMetricStatus(Project project) { + if (project.getBasePath() == null) return Map.of(); + + Path[] candidates = { + Path.of(project.getBasePath(), "build", "agenteval", "agenteval-report.json"), + Path.of(project.getBasePath(), "target", "agenteval", "agenteval-report.json") + }; + + for (Path path : candidates) { + if (Files.exists(path)) { + try { + ReportModel report = ReportParser.parseFile(path); + return ReportParser.extractMetricPassFail(report); + } catch (IOException e) { + return Map.of(); + } + } + } + return Map.of(); + } +} diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportFileWatcher.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportFileWatcher.java new file mode 100644 index 0000000..dd17296 --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportFileWatcher.java @@ -0,0 +1,57 @@ +package com.agenteval.intellij; + +import com.intellij.openapi.project.Project; +import com.intellij.openapi.vfs.VirtualFileEvent; +import com.intellij.openapi.vfs.VirtualFileListener; +import com.intellij.openapi.vfs.VirtualFileManager; +import com.intellij.openapi.wm.ToolWindow; +import com.intellij.openapi.wm.ToolWindowManager; +import com.intellij.ui.content.Content; +import org.jetbrains.annotations.NotNull; + +/** + * VFS listener that watches for changes to agenteval-report.json files + * and triggers a tool window refresh. + */ +public class ReportFileWatcher implements VirtualFileListener { + + private static final String REPORT_FILENAME = "agenteval-report.json"; + + private final Project project; + + public ReportFileWatcher(Project project) { + this.project = project; + } + + /** + * Registers this watcher with the VFS. + */ + public void register() { + VirtualFileManager.getInstance().addVirtualFileListener(this); + } + + @Override + public void contentsChanged(@NotNull VirtualFileEvent event) { + if (REPORT_FILENAME.equals(event.getFile().getName())) { + refreshToolWindow(); + } + } + + @Override + public void fileCreated(@NotNull VirtualFileEvent event) { + if (REPORT_FILENAME.equals(event.getFile().getName())) { + refreshToolWindow(); + } + } + + private void refreshToolWindow() { + ToolWindow toolWindow = ToolWindowManager.getInstance(project) + .getToolWindow("AgentEval"); + if (toolWindow != null) { + Content content = toolWindow.getContentManager().getContent(0); + if (content != null && content.getComponent() instanceof AgentEvalToolWindow panel) { + panel.refresh(); + } + } + } +} diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportModel.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportModel.java new file mode 100644 index 0000000..fb37274 --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportModel.java @@ -0,0 +1,106 @@ +package com.agenteval.intellij; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.List; +import java.util.Map; + +/** + * Lightweight model for parsing AgentEval JSON reports in the IntelliJ plugin. + * + *

Matches the output format of {@code JsonReporter}. No dependency on agenteval-core.

+ */ +@JsonIgnoreProperties(ignoreUnknown = true) +public final class ReportModel { + + @JsonProperty("averageScore") + private double averageScore; + + @JsonProperty("passRate") + private double passRate; + + @JsonProperty("totalCases") + private int totalCases; + + @JsonProperty("failedCases") + private int failedCases; + + @JsonProperty("durationMs") + private long durationMs; + + @JsonProperty("metricAverages") + private Map metricAverages; + + @JsonProperty("caseResults") + private List caseResults; + + public ReportModel() {} + + public double getAverageScore() { return averageScore; } + public double getPassRate() { return passRate; } + public int getTotalCases() { return totalCases; } + public int getFailedCases() { return failedCases; } + public long getDurationMs() { return durationMs; } + public Map getMetricAverages() { return metricAverages; } + public List getCaseResults() { return caseResults; } + + /** + * Returns true if the overall evaluation passed (all cases passed). + */ + public boolean isOverallPass() { + return failedCases == 0; + } + + /** + * A single test case result within the report. + */ + @JsonIgnoreProperties(ignoreUnknown = true) + public static final class CaseResultModel { + + @JsonProperty("input") + private String input; + + @JsonProperty("passed") + private boolean passed; + + @JsonProperty("averageScore") + private double averageScore; + + @JsonProperty("scores") + private Map scores; + + public CaseResultModel() {} + + public String getInput() { return input; } + public boolean isPassed() { return passed; } + public double getAverageScore() { return averageScore; } + public Map getScores() { return scores; } + } + + /** + * A metric score within a case result. + */ + @JsonIgnoreProperties(ignoreUnknown = true) + public static final class ScoreModel { + + @JsonProperty("value") + private double value; + + @JsonProperty("threshold") + private double threshold; + + @JsonProperty("passed") + private boolean passed; + + @JsonProperty("reason") + private String reason; + + public ScoreModel() {} + + public double getValue() { return value; } + public double getThreshold() { return threshold; } + public boolean isPassed() { return passed; } + public String getReason() { return reason; } + } +} diff --git a/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportParser.java b/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportParser.java new file mode 100644 index 0000000..5476258 --- /dev/null +++ b/agenteval-intellij/src/main/java/com/agenteval/intellij/ReportParser.java @@ -0,0 +1,69 @@ +package com.agenteval.intellij; + +import com.fasterxml.jackson.databind.ObjectMapper; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.LinkedHashMap; +import java.util.Map; + +/** + * Parses AgentEval JSON reports into {@link ReportModel} instances. + */ +public final class ReportParser { + + private static final ObjectMapper MAPPER = new ObjectMapper(); + + private ReportParser() {} + + /** + * Parses a JSON string into a report model. + * + * @param json the JSON string + * @return the parsed report model + * @throws IOException if parsing fails + */ + public static ReportModel parse(String json) throws IOException { + return MAPPER.readValue(json, ReportModel.class); + } + + /** + * Parses a JSON report file. + * + * @param path the path to the JSON file + * @return the parsed report model + * @throws IOException if reading or parsing fails + */ + public static ReportModel parseFile(Path path) throws IOException { + String json = Files.readString(path); + return parse(json); + } + + /** + * Extracts per-metric pass/fail status from a report. + * + *

A metric is considered passing if its average score is above the + * average threshold across all cases for that metric.

+ * + * @param report the report model + * @return map of metric name to pass/fail status + */ + public static Map extractMetricPassFail(ReportModel report) { + Map result = new LinkedHashMap<>(); + if (report.getCaseResults() == null || report.getCaseResults().isEmpty()) { + return result; + } + + // Aggregate pass/fail per metric: metric passes if all cases pass that metric + Map allPassed = new LinkedHashMap<>(); + for (ReportModel.CaseResultModel cr : report.getCaseResults()) { + if (cr.getScores() == null) continue; + for (Map.Entry entry : cr.getScores().entrySet()) { + allPassed.merge(entry.getKey(), entry.getValue().isPassed(), + (existing, current) -> existing && current); + } + } + return allPassed; + } +} diff --git a/agenteval-intellij/src/main/resources/META-INF/plugin.xml b/agenteval-intellij/src/main/resources/META-INF/plugin.xml new file mode 100644 index 0000000..2c76bbf --- /dev/null +++ b/agenteval-intellij/src/main/resources/META-INF/plugin.xml @@ -0,0 +1,20 @@ + + com.agenteval.intellij + AgentEval + AgentEval + View AgentEval evaluation results in IntelliJ IDEA + + com.intellij.modules.platform + com.intellij.modules.java + + + + + + + diff --git a/agenteval-intellij/src/main/resources/icons/agenteval-fail.svg b/agenteval-intellij/src/main/resources/icons/agenteval-fail.svg new file mode 100644 index 0000000..68010c9 --- /dev/null +++ b/agenteval-intellij/src/main/resources/icons/agenteval-fail.svg @@ -0,0 +1,4 @@ + + + + diff --git a/agenteval-intellij/src/main/resources/icons/agenteval-pass.svg b/agenteval-intellij/src/main/resources/icons/agenteval-pass.svg new file mode 100644 index 0000000..63aa9b5 --- /dev/null +++ b/agenteval-intellij/src/main/resources/icons/agenteval-pass.svg @@ -0,0 +1,4 @@ + + + + diff --git a/agenteval-intellij/src/main/resources/icons/agenteval-toolwindow.svg b/agenteval-intellij/src/main/resources/icons/agenteval-toolwindow.svg new file mode 100644 index 0000000..4149382 --- /dev/null +++ b/agenteval-intellij/src/main/resources/icons/agenteval-toolwindow.svg @@ -0,0 +1,4 @@ + + + + diff --git a/agenteval-intellij/src/test/java/com/agenteval/intellij/ReportModelTest.java b/agenteval-intellij/src/test/java/com/agenteval/intellij/ReportModelTest.java new file mode 100644 index 0000000..9d0e26c --- /dev/null +++ b/agenteval-intellij/src/test/java/com/agenteval/intellij/ReportModelTest.java @@ -0,0 +1,113 @@ +package com.agenteval.intellij; + +import org.junit.jupiter.api.Test; + +import java.io.IOException; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.within; + +class ReportModelTest { + + @Test + void metricAveragesAccessible() throws IOException { + String json = """ + { + "averageScore": 0.85, + "passRate": 1.0, + "totalCases": 1, + "failedCases": 0, + "durationMs": 100, + "metricAverages": {"M1": 0.9, "M2": 0.8}, + "caseResults": [] + } + """; + + ReportModel report = ReportParser.parse(json); + + assertThat(report.getMetricAverages()).hasSize(2); + assertThat(report.getMetricAverages().get("M1")).isCloseTo(0.9, within(0.001)); + assertThat(report.getMetricAverages().get("M2")).isCloseTo(0.8, within(0.001)); + } + + @Test + void overallPassWhenNoFailures() throws IOException { + String json = """ + { + "averageScore": 0.9, + "passRate": 1.0, + "totalCases": 2, + "failedCases": 0, + "durationMs": 200, + "metricAverages": {}, + "caseResults": [] + } + """; + + ReportModel report = ReportParser.parse(json); + assertThat(report.isOverallPass()).isTrue(); + } + + @Test + void overallFailWhenHasFailures() throws IOException { + String json = """ + { + "averageScore": 0.5, + "passRate": 0.5, + "totalCases": 2, + "failedCases": 1, + "durationMs": 200, + "metricAverages": {}, + "caseResults": [] + } + """; + + ReportModel report = ReportParser.parse(json); + assertThat(report.isOverallPass()).isFalse(); + } + + @Test + void scoreModelFieldsAccessible() throws IOException { + String json = """ + { + "averageScore": 0.9, + "passRate": 1.0, + "totalCases": 1, + "failedCases": 0, + "durationMs": 50, + "metricAverages": {"M1": 0.9}, + "caseResults": [{ + "input": "Q", "passed": true, "averageScore": 0.9, + "scores": {"M1": {"value": 0.9, "threshold": 0.7, "passed": true, "reason": "good"}} + }] + } + """; + + ReportModel report = ReportParser.parse(json); + var score = report.getCaseResults().getFirst().getScores().get("M1"); + + assertThat(score.getValue()).isCloseTo(0.9, within(0.001)); + assertThat(score.getThreshold()).isCloseTo(0.7, within(0.001)); + assertThat(score.isPassed()).isTrue(); + assertThat(score.getReason()).isEqualTo("good"); + } + + @Test + void unknownFieldsIgnored() throws IOException { + String json = """ + { + "averageScore": 0.9, + "passRate": 1.0, + "totalCases": 1, + "failedCases": 0, + "durationMs": 50, + "metricAverages": {}, + "caseResults": [], + "extraField": "ignored" + } + """; + + ReportModel report = ReportParser.parse(json); + assertThat(report.getAverageScore()).isCloseTo(0.9, within(0.001)); + } +} diff --git a/agenteval-intellij/src/test/java/com/agenteval/intellij/ReportParserTest.java b/agenteval-intellij/src/test/java/com/agenteval/intellij/ReportParserTest.java new file mode 100644 index 0000000..4524ca5 --- /dev/null +++ b/agenteval-intellij/src/test/java/com/agenteval/intellij/ReportParserTest.java @@ -0,0 +1,210 @@ +package com.agenteval.intellij; + +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.io.TempDir; + +import java.io.IOException; +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Map; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.assertj.core.api.Assertions.within; + +class ReportParserTest { + + @Test + void parseSingleCaseReport() throws IOException { + String json = """ + { + "averageScore": 0.9, + "passRate": 1.0, + "totalCases": 1, + "failedCases": 0, + "durationMs": 150, + "metricAverages": {"Relevancy": 0.9}, + "caseResults": [{ + "input": "What is Java?", + "passed": true, + "averageScore": 0.9, + "scores": { + "Relevancy": { + "value": 0.9, + "threshold": 0.7, + "passed": true, + "reason": "Good answer" + } + } + }] + } + """; + + ReportModel report = ReportParser.parse(json); + + assertThat(report.getAverageScore()).isCloseTo(0.9, within(0.001)); + assertThat(report.getPassRate()).isCloseTo(1.0, within(0.001)); + assertThat(report.getTotalCases()).isEqualTo(1); + assertThat(report.getFailedCases()).isEqualTo(0); + assertThat(report.getDurationMs()).isEqualTo(150L); + assertThat(report.getMetricAverages()).containsEntry("Relevancy", 0.9); + assertThat(report.getCaseResults()).hasSize(1); + + var cr = report.getCaseResults().getFirst(); + assertThat(cr.getInput()).isEqualTo("What is Java?"); + assertThat(cr.isPassed()).isTrue(); + assertThat(cr.getScores().get("Relevancy").getValue()).isCloseTo(0.9, within(0.001)); + assertThat(cr.getScores().get("Relevancy").getReason()).isEqualTo("Good answer"); + } + + @Test + void parseMultiCaseReport() throws IOException { + String json = """ + { + "averageScore": 0.7, + "passRate": 0.5, + "totalCases": 2, + "failedCases": 1, + "durationMs": 300, + "metricAverages": {"M1": 0.7}, + "caseResults": [ + {"input": "Q1", "passed": true, "averageScore": 0.9, + "scores": {"M1": {"value": 0.9, "threshold": 0.7, "passed": true, "reason": "ok"}}}, + {"input": "Q2", "passed": false, "averageScore": 0.5, + "scores": {"M1": {"value": 0.5, "threshold": 0.7, "passed": false, "reason": "bad"}}} + ] + } + """; + + ReportModel report = ReportParser.parse(json); + + assertThat(report.getTotalCases()).isEqualTo(2); + assertThat(report.getFailedCases()).isEqualTo(1); + assertThat(report.getCaseResults()).hasSize(2); + assertThat(report.getCaseResults().get(0).isPassed()).isTrue(); + assertThat(report.getCaseResults().get(1).isPassed()).isFalse(); + } + + @Test + void parseFailedCases() throws IOException { + String json = """ + { + "averageScore": 0.3, + "passRate": 0.0, + "totalCases": 1, + "failedCases": 1, + "durationMs": 100, + "metricAverages": {"M1": 0.3}, + "caseResults": [{ + "input": "Q1", "passed": false, "averageScore": 0.3, + "scores": {"M1": {"value": 0.3, "threshold": 0.7, "passed": false, "reason": "poor"}} + }] + } + """; + + ReportModel report = ReportParser.parse(json); + assertThat(report.isOverallPass()).isFalse(); + assertThat(report.getFailedCases()).isEqualTo(1); + } + + @Test + void parseMissingFieldsHandledGracefully() throws IOException { + String json = """ + { + "averageScore": 0.0, + "passRate": 0.0, + "totalCases": 0, + "failedCases": 0, + "durationMs": 0, + "metricAverages": {}, + "caseResults": [] + } + """; + + ReportModel report = ReportParser.parse(json); + assertThat(report.getCaseResults()).isEmpty(); + assertThat(report.getMetricAverages()).isEmpty(); + } + + @Test + void malformedJsonThrowsException() { + assertThatThrownBy(() -> ReportParser.parse("not valid json")) + .isInstanceOf(IOException.class); + } + + @Test + void parseFile(@TempDir Path tempDir) throws IOException { + String json = """ + { + "averageScore": 0.85, + "passRate": 1.0, + "totalCases": 1, + "failedCases": 0, + "durationMs": 50, + "metricAverages": {"M1": 0.85}, + "caseResults": [{ + "input": "Q", "passed": true, "averageScore": 0.85, + "scores": {"M1": {"value": 0.85, "threshold": 0.7, "passed": true, "reason": "ok"}} + }] + } + """; + Path file = tempDir.resolve("report.json"); + Files.writeString(file, json); + + ReportModel report = ReportParser.parseFile(file); + assertThat(report.getAverageScore()).isCloseTo(0.85, within(0.001)); + } + + @Test + void extractMetricPassFailAllPass() throws IOException { + String json = """ + { + "averageScore": 0.9, + "passRate": 1.0, + "totalCases": 1, + "failedCases": 0, + "durationMs": 50, + "metricAverages": {"M1": 0.9, "M2": 0.8}, + "caseResults": [{ + "input": "Q", "passed": true, "averageScore": 0.85, + "scores": { + "M1": {"value": 0.9, "threshold": 0.7, "passed": true, "reason": "ok"}, + "M2": {"value": 0.8, "threshold": 0.7, "passed": true, "reason": "ok"} + } + }] + } + """; + + ReportModel report = ReportParser.parse(json); + Map status = ReportParser.extractMetricPassFail(report); + + assertThat(status).containsEntry("M1", true); + assertThat(status).containsEntry("M2", true); + } + + @Test + void extractMetricPassFailMixed() throws IOException { + String json = """ + { + "averageScore": 0.65, + "passRate": 0.5, + "totalCases": 2, + "failedCases": 1, + "durationMs": 100, + "metricAverages": {"M1": 0.65}, + "caseResults": [ + {"input": "Q1", "passed": true, "averageScore": 0.9, + "scores": {"M1": {"value": 0.9, "threshold": 0.7, "passed": true, "reason": "ok"}}}, + {"input": "Q2", "passed": false, "averageScore": 0.4, + "scores": {"M1": {"value": 0.4, "threshold": 0.7, "passed": false, "reason": "bad"}}} + ] + } + """; + + ReportModel report = ReportParser.parse(json); + Map status = ReportParser.extractMetricPassFail(report); + + // M1 failed in at least one case, so overall it's false + assertThat(status).containsEntry("M1", false); + } +}