from __future__ import annotations
import argparse
import json
import sys
from dataclasses import dataclass
from pathlib import Path
DEFAULT_GATE_GROUPS = ("zone_eval_scan", "zone_eval_evaluate")
NS_PER_MS = 1_000_000.0
NS_PER_US = 1_000.0
@dataclass(frozen=True)
class Sample:
group: str
bench_id: str
times_ns: list[float]
@property
def name(self) -> str:
return f"{self.group}/{self.bench_id}"
def percentile(self, pct: float) -> float:
ordered = sorted(self.times_ns)
rank = max(1, min(len(ordered), -(-int(pct * len(ordered)) // 100)))
return ordered[rank - 1]
def read_samples(report: Path, baseline: str) -> list[Sample]:
found: list[Sample] = []
for path in sorted(report.glob(f"*/*/{baseline}/sample.json")):
group = path.parent.parent.parent.name
bench_id = path.parent.parent.name
try:
doc = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
print(f"error: {path} is unreadable: {exc}", file=sys.stderr)
raise SystemExit(2) from exc
iters = doc.get("iters")
times = doc.get("times")
if not isinstance(iters, list) or not isinstance(times, list) or not iters:
print(
f"error: {path} carries no iters[]/times[] arrays — "
"this gate reads sample.json, not estimates.json",
file=sys.stderr,
)
raise SystemExit(2)
if len(iters) != len(times):
print(
f"error: {path} has {len(iters)} iters and {len(times)} times",
file=sys.stderr,
)
raise SystemExit(2)
per_iter = [t / i for t, i in zip(times, iters) if i]
if not per_iter:
print(f"error: {path} holds no usable samples", file=sys.stderr)
raise SystemExit(2)
found.append(Sample(group=group, bench_id=bench_id, times_ns=per_iter))
return found
def main() -> int:
for stream in (sys.stdout, sys.stderr):
reconfigure = getattr(stream, "reconfigure", None)
if reconfigure is not None:
reconfigure(encoding="utf-8", errors="replace")
parser = argparse.ArgumentParser(
description="Fail the build when a gated benchmark's p99 leaves its budget.",
)
parser.add_argument(
"--report",
type=Path,
required=True,
help="criterion's report directory (target/criterion)",
)
parser.add_argument(
"--p99-ms",
type=float,
required=True,
help="the p99 budget, in milliseconds; a benchmark at or above it fails",
)
parser.add_argument(
"--baseline",
default="ci",
help="the baseline name passed to `cargo bench -- --save-baseline` (default: ci)",
)
parser.add_argument(
"--gate-group",
action="append",
default=None,
metavar="GROUP",
help=(
"a criterion group to hold to the budget; repeatable. "
f"Default: {', '.join(DEFAULT_GATE_GROUPS)}"
),
)
args = parser.parse_args()
if args.p99_ms <= 0:
print("error: --p99-ms must be positive", file=sys.stderr)
return 2
if not args.report.is_dir():
print(f"error: --report {args.report} is not a directory", file=sys.stderr)
return 2
gated_groups = tuple(args.gate_group) if args.gate_group else DEFAULT_GATE_GROUPS
budget_ns = args.p99_ms * NS_PER_MS
samples = read_samples(args.report, args.baseline)
gated = [s for s in samples if s.group in gated_groups]
ungated = [s for s in samples if s.group not in gated_groups]
print(f"Bench gate — p99 budget {args.p99_ms} ms, baseline '{args.baseline}'")
print(f"Report: {args.report}")
print()
failures: list[tuple[Sample, float]] = []
if gated:
print("Gated (per-event):")
for sample in sorted(gated, key=lambda s: s.name):
p99 = sample.percentile(99)
p50 = sample.percentile(50)
over = p99 >= budget_ns
mark = "FAIL" if over else "ok"
print(
f" [{mark:>4}] {sample.name}: "
f"p50 {p50 / NS_PER_US:.2f} µs, p99 {p99 / NS_PER_US:.2f} µs "
f"({len(sample.times_ns)} samples)"
)
if over:
failures.append((sample, p99))
print()
if ungated:
print("Reported only (not gated — see the module docstring):")
for sample in sorted(ungated, key=lambda s: s.name):
p99 = sample.percentile(99)
p50 = sample.percentile(50)
print(
f" [info] {sample.name}: "
f"p50 {p50 / NS_PER_MS:.3f} ms, p99 {p99 / NS_PER_MS:.3f} ms"
)
print()
if not gated:
print(
"error: no gated benchmark found under "
f"{args.report}/*/*/{args.baseline}/sample.json",
file=sys.stderr,
)
print(
" Looked for group(s): " + ", ".join(gated_groups),
file=sys.stderr,
)
if samples:
seen = ", ".join(sorted({s.group for s in samples}))
print(f" Groups present: {seen}", file=sys.stderr)
else:
print(" No sample.json at all — did `cargo bench` run?", file=sys.stderr)
print(
" A gate that measures nothing is a failure, never a pass.",
file=sys.stderr,
)
return 1
if failures:
print(
f"error: {len(failures)} benchmark(s) at or above the "
f"{args.p99_ms} ms p99 budget:",
file=sys.stderr,
)
for sample, p99 in failures:
print(
f" {sample.name}: p99 {p99 / NS_PER_US:.2f} µs "
f">= {budget_ns / NS_PER_US:.2f} µs",
file=sys.stderr,
)
return 1
print(f"All {len(gated)} gated benchmark(s) inside the {args.p99_ms} ms p99 budget.")
return 0
if __name__ == "__main__":
sys.exit(main())