Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -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.
*
* <pre>{@code
* var result = Benchmark.run(testCases, List.of(variantA, variantB));
* System.out.println("Best: " + result.bestVariant());
* }</pre>
*/
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<AgentTestCase> testCases,
List<BenchmarkVariant> variants) {
return run(testCases, variants, BenchmarkConfig.defaults());
}

/**
* Runs benchmarks with the specified config.
*/
public static BenchmarkResult run(List<AgentTestCase> testCases,
List<BenchmarkVariant> 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<String, EvalResult> 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<String, EvalResult> runSequential(
List<AgentTestCase> testCases, List<BenchmarkVariant> variants) {
Map<String, EvalResult> 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<String, EvalResult> runParallel(
List<AgentTestCase> testCases, List<BenchmarkVariant> variants,
BenchmarkConfig config) {
Semaphore semaphore = new Semaphore(config.maxParallelVariants());
Map<String, EvalResult> results = new LinkedHashMap<>();

try (ExecutorService executor = Executors.newVirtualThreadPerTaskExecutor()) {
List<Future<Map.Entry<String, EvalResult>>> 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<Map.Entry<String, EvalResult>> future : futures) {
try {
Map.Entry<String, EvalResult> 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<AgentTestCase> testCases,
BenchmarkVariant variant) {
// Deep-copy test cases for this variant to ensure isolation
List<AgentTestCase> 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<BenchmarkVariant> variants) {
Set<String> seen = new HashSet<>();
for (BenchmarkVariant v : variants) {
if (!seen.add(v.name())) {
throw new IllegalArgumentException("Duplicate variant name: " + v.name());
}
}
}
}
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
package com.agenteval.core.benchmark;

/**
* Configuration for benchmark execution.
*
* <pre>{@code
* var config = BenchmarkConfig.builder()
* .parallelVariants(true)
* .maxParallelVariants(4)
* .build();
* }</pre>
*/
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);
}
}
}
Original file line number Diff line number Diff line change
@@ -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<String, EvalResult> variantResults;
private final long totalDurationMs;

public BenchmarkResult(Map<String, EvalResult> variantResults, long totalDurationMs) {
Objects.requireNonNull(variantResults, "variantResults must not be null");
this.variantResults = new LinkedHashMap<>(variantResults);
this.totalDurationMs = totalDurationMs;
}

public Map<String, EvalResult> 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<Map.Entry<String, Double>> averageScores() {
return variantResults.entrySet().stream()
.map(e -> Map.entry(e.getKey(), e.getValue().averageScore()))
.sorted(Map.Entry.<String, Double>comparingByValue().reversed())
.toList();
}

/**
* Returns per-metric scores grouped by metric name, then by variant.
*/
public Map<String, Map<String, Double>> scoresByMetric() {
Map<String, Map<String, Double>> result = new LinkedHashMap<>();
for (Map.Entry<String, EvalResult> entry : variantResults.entrySet()) {
String variant = entry.getKey();
Map<String, Double> metricAvgs = entry.getValue().averageScoresByMetric();
for (Map.Entry<String, Double> metricEntry : metricAvgs.entrySet()) {
result.computeIfAbsent(metricEntry.getKey(), k -> new LinkedHashMap<>())
.put(variant, metricEntry.getValue());
}
}
return result;
}
}
Original file line number Diff line number Diff line change
@@ -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.
*
* <p>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).</p>
*
* <pre>{@code
* var variant = BenchmarkVariant.builder()
* .name("gpt-4o")
* .metrics(List.of(new AnswerRelevancyMetric(judge)))
* .casePreparer(tc -> tc.toBuilder().actualOutput(callGpt4o(tc.getInput())).build())
* .build();
* }</pre>
*/
public final class BenchmarkVariant {

private final String name;
private final AgentEvalConfig config;
private final List<EvalMetric> metrics;
private final UnaryOperator<AgentTestCase> 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<EvalMetric> metrics() { return metrics; }
public UnaryOperator<AgentTestCase> 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<EvalMetric> metrics;
private UnaryOperator<AgentTestCase> 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<EvalMetric> metrics) { this.metrics = metrics; return this; }
public Builder casePreparer(UnaryOperator<AgentTestCase> preparer) {
this.casePreparer = preparer;
return this;
}

public BenchmarkVariant build() {
return new BenchmarkVariant(this);
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -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; }
Expand Down
Loading