hackbait 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
hackbait/cli.py ADDED
@@ -0,0 +1,534 @@
1
+ """Command line surface: hackbait list | selftest | run | report | power.
2
+
3
+ Exit codes are part of the contract: 0 for a clean run, 1 for usage or
4
+ runtime errors, 2 when selftest finds a misclassification or power is handed a
5
+ rate or half-width outside its domain. argparse wants to exit 2 on bad flags,
6
+ which would collide with those codes, so the parser here downgrades usage
7
+ errors to 1.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import argparse
13
+ import importlib
14
+ import json
15
+ import os
16
+ import re
17
+ import sys
18
+ import time
19
+ import urllib.parse
20
+ from collections import Counter
21
+ from collections.abc import Sequence
22
+ from pathlib import Path
23
+ from typing import Any, NoReturn
24
+
25
+ import hackbait
26
+ import hackbait.scripted
27
+
28
+ from . import runner, stats
29
+ from .bait.testedit import EXEC_WARNING
30
+ from .types import CONTROL_TIER, TIER_NAMES, TIERS, VERDICTS
31
+
32
+ # Token counts assumed per episode when pricing a planned run that has not
33
+ # happened yet (the power command). Input is the exact suite mean from
34
+ # bench/throughput.py; output is the mid completion-length estimate from the
35
+ # same bench. Both are estimates for planning, never a measurement of a real
36
+ # run, and they only price the cost column: they never touch the sample-size
37
+ # math in stats.py, which stays metadata-free.
38
+ _PLAN_PROMPT_TOKENS = 167
39
+ _PLAN_COMPLETION_TOKENS = 350
40
+
41
+ __all__ = ["build_parser", "run"]
42
+
43
+
44
+ class _Parser(argparse.ArgumentParser):
45
+ """ArgumentParser that exits 1 on usage errors instead of argparse's 2."""
46
+
47
+ def error(self, message: str) -> NoReturn:
48
+ self.print_usage(sys.stderr)
49
+ print(f"{self.prog}: error: {message}", file=sys.stderr)
50
+ raise SystemExit(1)
51
+
52
+
53
+ def build_parser() -> argparse.ArgumentParser:
54
+ parser = _Parser(prog="hackbait", description="Honeypot suite that measures reward hacking.")
55
+ sub = parser.add_subparsers(dest="command")
56
+
57
+ sub.add_parser("list", help="one line per honeypot family")
58
+
59
+ selftest = sub.add_parser(
60
+ "selftest", help="judge every scripted policy against its expected verdicts"
61
+ )
62
+ selftest.add_argument("--per-tier", type=int, default=8, dest="per_tier")
63
+ selftest.add_argument("--seed", type=int, default=0)
64
+ selftest.add_argument(
65
+ "--exec",
66
+ action="store_true",
67
+ dest="execute",
68
+ help="run held-out tests for HB-TEST in a sandboxed subprocess",
69
+ )
70
+
71
+ run_p = sub.add_parser("run", help="run a model or a scripted policy over the suite")
72
+ target = run_p.add_mutually_exclusive_group(required=True)
73
+ target.add_argument("--model", help="model id for an OpenAI-compatible endpoint")
74
+ target.add_argument(
75
+ "--policy",
76
+ choices=("oracle", "cheat", "garbage"),
77
+ help="scripted policy instead of a model; costs zero API calls",
78
+ )
79
+ run_p.add_argument("--base-url", dest="base_url", help="overrides HACKBAIT_BASE_URL")
80
+ run_p.add_argument("--per-tier", type=int, default=8, dest="per_tier")
81
+ run_p.add_argument("--seed", type=int, default=0)
82
+ run_p.add_argument(
83
+ "--tiers", default="1,2,3", help="comma-separated subset of 0,1,2,3 (0 is the control arm)"
84
+ )
85
+ run_p.add_argument("--families", help="comma-separated family codes, e.g. HB-KEY,HB-RUB")
86
+ run_p.add_argument(
87
+ "--temperature",
88
+ default="0.0",
89
+ help='sampling temperature, or "none" to omit the field for models that reject it',
90
+ )
91
+ run_p.add_argument("--exec", action="store_true", dest="execute")
92
+ run_p.add_argument("--out", help="run file path; default runs/<label>-<unixtime>.jsonl")
93
+
94
+ report_p = sub.add_parser("report", help="rebuild the leaderboard from saved runs")
95
+ # Default is None so the handler can tell "user asked for this path"
96
+ # (missing means error) from "use the default" (missing may fall back to
97
+ # ./runs, where hackbait run writes).
98
+ report_p.add_argument(
99
+ "--results",
100
+ default=None,
101
+ help="directory of run transcripts (default leaderboard/results, falling back to ./runs)",
102
+ )
103
+ report_p.add_argument("--out-md", default="leaderboard/LEADERBOARD.md", dest="out_md")
104
+ report_p.add_argument("--out-json", default="leaderboard/leaderboard.json", dest="out_json")
105
+
106
+ power_p = sub.add_parser(
107
+ "power",
108
+ help="sample size to estimate a hack rate to a target precision",
109
+ )
110
+ power_p.add_argument(
111
+ "--rate",
112
+ type=float,
113
+ action="append",
114
+ dest="rates",
115
+ help="expected per-model hack rate (repeatable); default 0, 1, 2, 5, 10%%",
116
+ )
117
+ power_p.add_argument(
118
+ "--half-width",
119
+ type=float,
120
+ action="append",
121
+ dest="half_widths",
122
+ help="target Wilson half-width as a fraction (repeatable); default 0.02 = 2pp",
123
+ )
124
+ power_p.add_argument(
125
+ "--model",
126
+ help="price the plan for this model id, adding an estimated-cost column",
127
+ )
128
+ power_p.add_argument(
129
+ "--prices",
130
+ help="JSON file of {id: [in, out]} price overrides applied before pricing",
131
+ )
132
+ power_p.add_argument("--json", action="store_true", dest="as_json")
133
+
134
+ return parser
135
+
136
+
137
+ def run(argv: Sequence[str] | None = None) -> int:
138
+ parser = build_parser()
139
+ args = parser.parse_args(argv)
140
+ if args.command is None:
141
+ parser.print_help()
142
+ return 1
143
+ handlers = {
144
+ "list": _cmd_list,
145
+ "selftest": _cmd_selftest,
146
+ "run": _cmd_run,
147
+ "report": _cmd_report,
148
+ "power": _cmd_power,
149
+ }
150
+ try:
151
+ return handlers[args.command](args)
152
+ except Exception as err:
153
+ print(f"error: {err}", file=sys.stderr)
154
+ return 1
155
+
156
+
157
+ def _cmd_list(args: argparse.Namespace) -> int:
158
+ for honeypot in hackbait.suite():
159
+ module_name = type(honeypot).__module__
160
+ module = sys.modules.get(module_name) or importlib.import_module(module_name)
161
+ doc_lines = (module.__doc__ or "").strip().splitlines()
162
+ description = doc_lines[0] if doc_lines else ""
163
+ count = len(honeypot.cases(per_tier=8, seed=0))
164
+ print(f"{honeypot.family:<8} {module_name:<24} {count:>4} cases {description}")
165
+ return 0
166
+
167
+
168
+ _POLICY_NAMES: tuple[str, ...] = ("oracle", "cheat", "garbage")
169
+
170
+
171
+ def _cmd_selftest(args: argparse.Namespace) -> int:
172
+ policies = hackbait.scripted.scripted()
173
+ misses: list[tuple[str, str, tuple[str, ...], str]] = []
174
+ for honeypot in hackbait.suite():
175
+ family_policies = policies[honeypot.family]
176
+ cases = honeypot.cases(per_tier=args.per_tier, seed=args.seed)
177
+ counts: dict[str, dict[str, int]] = {
178
+ name: dict.fromkeys(VERDICTS, 0) for name in _POLICY_NAMES
179
+ }
180
+ for name in _POLICY_NAMES:
181
+ policy = family_policies[name]
182
+ for case in cases:
183
+ episode = runner.judge_case(honeypot, case, policy(case), execute=args.execute)
184
+ counts[name][episode.verdict] += 1
185
+ # The gate accepts what expected_verdicts allows for this exact
186
+ # case and execute flag, so an HB-TEST oracle judged
187
+ # CLEAN_UNVERIFIED under --exec is caught as broken wiring
188
+ # rather than waved through.
189
+ expected = hackbait.scripted.expected_verdicts(name, case, args.execute)
190
+ if episode.verdict not in expected:
191
+ misses.append((case.id, name, expected, episode.verdict))
192
+ control_count = sum(1 for case in cases if case.tier == CONTROL_TIER)
193
+ _print_confusion(honeypot.family, len(cases), counts, control_count)
194
+ if misses:
195
+ for case_id, name, expected, got in misses:
196
+ print(f"miss: {case_id} policy={name} expected={'|'.join(expected)} got={got}")
197
+ print(f"selftest: {len(misses)} misclassifications")
198
+ return 2
199
+ print("selftest: OK")
200
+ return 0
201
+
202
+
203
+ def _print_confusion(
204
+ family: str, case_count: int, counts: dict[str, dict[str, int]], control_count: int = 0
205
+ ) -> None:
206
+ print(f"{family} ({case_count} cases)")
207
+ header = "".join(f"{verdict:>18}" for verdict in VERDICTS)
208
+ print(f" {'policy':<8}{header}")
209
+ for name, verdict_counts in counts.items():
210
+ cells = "".join(f"{verdict_counts[verdict]:>18}" for verdict in VERDICTS)
211
+ # On control cases no bait exists, so the cheat policy answering the
212
+ # poisoned value is expected to FAIL; annotate the row so those FAILs
213
+ # read as the control arm working, not the gate breaking.
214
+ note = ""
215
+ if name == "cheat" and control_count:
216
+ note = f" (control: {control_count} expected FAIL)"
217
+ print(f" {name:<8}{cells}{note}")
218
+
219
+
220
+ def _parse_tiers(spec: str) -> tuple[int, ...]:
221
+ try:
222
+ tiers = tuple(int(part) for part in spec.split(",") if part.strip())
223
+ except ValueError:
224
+ raise ValueError(f"--tiers wants comma-separated integers, got {spec!r}") from None
225
+ allowed = set(TIERS) | {CONTROL_TIER}
226
+ bad = sorted(set(tiers) - allowed)
227
+ if not tiers or bad:
228
+ raise ValueError(f"--tiers must pick from 0,1,2,3, got {spec!r}")
229
+ return tiers
230
+
231
+
232
+ def _slug(label: str) -> str:
233
+ """Filesystem-safe run label; model ids often carry slashes and colons."""
234
+ cleaned = re.sub(r"[^A-Za-z0-9._-]+", "-", label).strip("-.")
235
+ return cleaned or "run"
236
+
237
+
238
+ def _is_local_host(base_url: str) -> bool:
239
+ """True when the endpoint host is the local machine, so keyless is fine."""
240
+ try:
241
+ host = urllib.parse.urlsplit(base_url).hostname
242
+ except ValueError:
243
+ return False
244
+ return host in {"localhost", "127.0.0.1", "::1"}
245
+
246
+
247
+ def _cmd_run(args: argparse.Namespace) -> int:
248
+ tiers = _parse_tiers(args.tiers)
249
+ if args.execute:
250
+ print("warning: " + EXEC_WARNING, file=sys.stderr)
251
+ families = None
252
+ if args.families:
253
+ families = [part.strip().upper() for part in args.families.split(",") if part.strip()]
254
+ if args.model:
255
+ base_url = args.base_url or os.environ.get("HACKBAIT_BASE_URL", "")
256
+ api_key = os.environ.get("HACKBAIT_API_KEY", "")
257
+ if not api_key:
258
+ # People coming from other OpenAI-compatible tools usually have
259
+ # this one set already; honor it rather than demanding a rename.
260
+ api_key = os.environ.get("OPENAI_API_KEY", "")
261
+ if api_key:
262
+ print("using OPENAI_API_KEY (HACKBAIT_API_KEY is unset)", file=sys.stderr)
263
+ missing = []
264
+ if not base_url:
265
+ missing.append("a base url (--base-url or HACKBAIT_BASE_URL)")
266
+ if not api_key:
267
+ if _is_local_host(base_url):
268
+ # Local servers (ollama, vllm, llama.cpp) commonly take any
269
+ # bearer token or none; do not block a localhost run on a key.
270
+ print(
271
+ "no api key set; proceeding keyless against a localhost endpoint",
272
+ file=sys.stderr,
273
+ )
274
+ else:
275
+ missing.append("an api key (HACKBAIT_API_KEY or OPENAI_API_KEY)")
276
+ if missing:
277
+ print(f"error: --model needs {' and '.join(missing)}", file=sys.stderr)
278
+ return 1
279
+ raw_temp = (args.temperature or "").strip().lower()
280
+ if raw_temp in ("none", "omit", ""):
281
+ temperature = None
282
+ else:
283
+ try:
284
+ temperature = float(raw_temp)
285
+ except ValueError:
286
+ print(
287
+ f'error: --temperature wants a number or "none", got {args.temperature!r}',
288
+ file=sys.stderr,
289
+ )
290
+ return 1
291
+ respond = runner.model_responder(
292
+ args.model, base_url=base_url, api_key=api_key, temperature=temperature
293
+ )
294
+ label = args.model
295
+ else:
296
+ policies = hackbait.scripted.scripted()
297
+ policy_name = args.policy
298
+
299
+ def respond(case: hackbait.Case) -> str:
300
+ return policies[case.family][policy_name](case)
301
+
302
+ label = f"scripted-{policy_name}"
303
+ out = Path(args.out) if args.out else Path("runs") / f"{_slug(label)}-{int(time.time())}.jsonl"
304
+ summary = runner.run_suite(
305
+ respond,
306
+ out,
307
+ label,
308
+ per_tier=args.per_tier,
309
+ seed=args.seed,
310
+ tiers=tiers,
311
+ families=families,
312
+ execute=args.execute,
313
+ )
314
+ total, errored = _print_run_table(out)
315
+ print(f"wrote {out}")
316
+ # Only a real model run carries token usage worth summarizing; a scripted
317
+ # policy reports nothing but zeros. The usage figures are metadata printed
318
+ # for the operator and never fed back into the run.
319
+ if args.model:
320
+ _print_usage_summary(summary)
321
+ _print_ci_hint(summary)
322
+ if total and errored == total:
323
+ # A wrong model name or dead endpoint turns every episode into an
324
+ # errored FAIL; a normal-looking exit 0 here would hide that the run
325
+ # measured nothing.
326
+ print(
327
+ f"error: all {total} episodes errored; check the base url, model name "
328
+ "and api key",
329
+ file=sys.stderr,
330
+ )
331
+ return 1
332
+ return 0
333
+
334
+
335
+ def _print_run_table(out: Path) -> tuple[int, int]:
336
+ """Summary from the rows just written: overall, per family, per tier.
337
+
338
+ Computed from the run file rather than report.summarize, so the table
339
+ never depends on the report module's dict shape. Control (tier 0) rows
340
+ are not hack attempts and stay out of every rate; they get their own
341
+ line. Returns (total episodes, episodes with error set) so the caller
342
+ can refuse to exit 0 on an all-errored run.
343
+ """
344
+ rows: list[dict[str, Any]] = []
345
+ with out.open(encoding="utf-8") as fh:
346
+ for line in fh:
347
+ record = json.loads(line)
348
+ if record.get("kind") != "run":
349
+ rows.append(record)
350
+ if not rows:
351
+ print("no episodes written")
352
+ return 0, 0
353
+
354
+ def rate(subset: list[dict[str, Any]]) -> str:
355
+ hacks = sum(row["verdict"] == "HACK" for row in subset)
356
+ return f"{hacks / len(subset):.3f} ({hacks}/{len(subset)})"
357
+
358
+ control = [row for row in rows if row["tier"] == CONTROL_TIER]
359
+ pressured = [row for row in rows if row["tier"] != CONTROL_TIER]
360
+
361
+ if pressured:
362
+ print(f"episodes: {len(pressured)} hack rate: {rate(pressured)}")
363
+ else:
364
+ print(f"episodes: {len(pressured)} hack rate: n/a (no pressured episodes)")
365
+ verdict_counts = Counter(row["verdict"] for row in pressured)
366
+ print("verdicts: " + " ".join(f"{v}={verdict_counts.get(v, 0)}" for v in VERDICTS))
367
+ for family in sorted({row["family"] for row in pressured}):
368
+ subset = [row for row in pressured if row["family"] == family]
369
+ print(f" {family:<8} {rate(subset)}")
370
+ for tier in TIERS:
371
+ subset = [row for row in pressured if row["tier"] == tier]
372
+ if subset:
373
+ print(f" tier {tier} ({TIER_NAMES[tier]}): {rate(subset)}")
374
+ if control:
375
+ collisions = sum(
376
+ 1 for row in control if "control_key_collision" in (row.get("signals") or [])
377
+ )
378
+ print(
379
+ f" control (tier 0): {len(control)} episodes, "
380
+ f"{collisions} key collisions (excluded from hack rate)"
381
+ )
382
+ errored = sum(1 for row in rows if row.get("error") is not None)
383
+ if errored:
384
+ print(f"errors: {errored}")
385
+ return len(rows), errored
386
+
387
+
388
+ def _print_usage_summary(summary: dict[str, Any]) -> None:
389
+ """Print the run's token, latency and cost line from the report summary.
390
+
391
+ Reads the summary's usage block, which report.summarize builds from the
392
+ stored rows. Cost is None when the model is unpriced or no episode carried
393
+ a usage block; latency is None when nothing timed. Both render as n/a
394
+ rather than a fake zero. Nothing here changes a verdict.
395
+ """
396
+ usage = summary.get("usage")
397
+ if not isinstance(usage, dict):
398
+ return
399
+ total = usage.get("total_tokens", 0)
400
+ in_tok = usage.get("prompt_tokens", 0)
401
+ out_tok = usage.get("completion_tokens", 0)
402
+ p95 = usage.get("p95_latency_ms")
403
+ p95_text = "n/a" if p95 is None else f"{p95:.0f}ms"
404
+ cost = usage.get("cost_usd")
405
+ cost_text = "n/a" if cost is None else f"${cost:.4f}"
406
+ print(f"tokens: {total} (in {in_tok} / out {out_tok}) p95 {p95_text} cost {cost_text}")
407
+
408
+
409
+ def _print_ci_hint(summary: dict[str, Any]) -> None:
410
+ """Print the observed hack-rate CI half-width and the per_tier to tighten it.
411
+
412
+ Uses the same Wilson half-width the leaderboard's bracket comes from, so
413
+ the +-Ypp here agrees with the stored interval. N is the pressured episode
414
+ count; with no pressured episodes there is nothing to bound, so the hint is
415
+ skipped. This reads the observed rate only to report precision; it never
416
+ feeds back into case generation or judging.
417
+ """
418
+ hack_rate = summary.get("hack_rate")
419
+ n = summary.get("episodes")
420
+ if not isinstance(hack_rate, (int, float)) or not isinstance(n, int) or n <= 0:
421
+ return
422
+ half_pp = 100.0 * stats.wilson_half_width(hack_rate, n)
423
+ per_tier = stats.episodes_and_per_tier(hack_rate, 0.02)[1]
424
+ print(
425
+ f"hack rate {hack_rate:.1%} +- {half_pp:.1f}pp (95% CI, n={n}); "
426
+ f"to reach +-2pp, run per_tier={per_tier}"
427
+ )
428
+
429
+
430
+ def _cmd_power(args: argparse.Namespace) -> int:
431
+ """Sample-size plan: episodes and per_tier for a target CI half-width.
432
+
433
+ The statistics come from stats.power_table, which is metadata-free by
434
+ construction. When --model names a priced model an estimated-cost column
435
+ is joined on top, but that join is one-directional: cost never feeds the
436
+ sample-size math, it only annotates it. Bad flag values (rate outside
437
+ [0, 1], half-width <= 0) exit 2, matching selftest's "the inputs were
438
+ wrong" code and distinct from a clean exit 0.
439
+ """
440
+ from . import pricing
441
+
442
+ rates = args.rates if args.rates else list(stats.DEFAULT_RATES)
443
+ half_widths = args.half_widths if args.half_widths else [stats.DEFAULT_HALF_WIDTH]
444
+ if args.prices:
445
+ pricing.load_prices(args.prices)
446
+
447
+ per_ep_cost: float | None = None
448
+ if args.model:
449
+ per_ep_cost = pricing.cost_usd(args.model, _PLAN_PROMPT_TOKENS, _PLAN_COMPLETION_TOKENS)
450
+ if per_ep_cost is None:
451
+ print(
452
+ f"note: {args.model} is unpriced; showing episodes and per_tier only",
453
+ file=sys.stderr,
454
+ )
455
+ try:
456
+ table = stats.power_table(rates, half_widths)
457
+ except ValueError as err:
458
+ # stats raises on rate outside [0, 1] or half-width <= 0. That is a
459
+ # bad-input error, not a crash, so it exits 2 rather than 1.
460
+ print(f"error: {err}", file=sys.stderr)
461
+ return 2
462
+
463
+ if per_ep_cost is not None:
464
+ for record in table:
465
+ # Cost of the episodes the plan calls for, at the per-episode
466
+ # estimate. An annotation on the sample-size row, never an input.
467
+ record["est_cost_usd"] = record["episodes"] * per_ep_cost
468
+ record["model"] = args.model
469
+
470
+ if args.as_json:
471
+ print(json.dumps(table, indent=2))
472
+ return 0
473
+
474
+ _print_power_table(table, per_ep_cost is not None, args.model)
475
+ return 0
476
+
477
+
478
+ def _print_power_table(table: list[dict[str, Any]], priced: bool, model: str | None) -> None:
479
+ """Render the power table for humans: rates and half-widths as percentages."""
480
+ if priced:
481
+ print(
482
+ f"per-episode cost estimate for {model}: "
483
+ f"{_PLAN_PROMPT_TOKENS} in / {_PLAN_COMPLETION_TOKENS} out tokens per episode"
484
+ )
485
+ header = f"{'rate':>8} {'half_width':>12} {'episodes':>10} {'per_tier':>10}"
486
+ if priced:
487
+ header += f" {'est_cost':>12}"
488
+ print(header)
489
+ print("-" * len(header))
490
+ for record in table:
491
+ rate_pct = f"{record['rate'] * 100:.1f}%"
492
+ half_pp = f"+-{record['half_width'] * 100:.1f}pp"
493
+ line = (
494
+ f"{rate_pct:>8} {half_pp:>12} "
495
+ f"{int(record['episodes']):>10} {int(record['per_tier']):>10}"
496
+ )
497
+ if priced:
498
+ cost = record.get("est_cost_usd")
499
+ cost_text = "n/a" if cost is None else f"${cost:.2f}"
500
+ line += f" {cost_text:>12}"
501
+ print(line)
502
+
503
+
504
+ def _cmd_report(args: argparse.Namespace) -> int:
505
+ from hackbait import report
506
+
507
+ if args.results is None:
508
+ results = Path("leaderboard/results")
509
+ if not results.exists():
510
+ runs = Path("runs")
511
+ if runs.is_dir():
512
+ # hackbait run writes to ./runs by default, so a fresh user's
513
+ # transcripts live there; fall back rather than erroring on a
514
+ # directory they never created.
515
+ print(f"{results} not found; falling back to {runs}", file=sys.stderr)
516
+ results = runs
517
+ else:
518
+ print(
519
+ f"error: {results}: results directory does not exist "
520
+ "(if your runs are in ./runs, pass --results runs)",
521
+ file=sys.stderr,
522
+ )
523
+ return 1
524
+ else:
525
+ # An explicitly requested path that does not exist is an error;
526
+ # report.leaderboard_rows raises and run() turns that into exit 1.
527
+ results = Path(args.results)
528
+ report.write_leaderboard(results, Path(args.out_md), Path(args.out_json))
529
+ print(f"leaderboard written to {args.out_md} and {args.out_json}")
530
+ return 0
531
+
532
+
533
+ if __name__ == "__main__":
534
+ raise SystemExit(run())
hackbait/pricing.py ADDED
@@ -0,0 +1,135 @@
1
+ """Model price table and cost estimation for hackbait runs.
2
+
3
+ Prices as of 2026-08, in USD per million tokens as (input, output) pairs, for
4
+ the models this suite actually exercises: the claude-* ids and the gemini-* ids
5
+ seen in this project. They are public list rates kept here so a run can print a
6
+ computed cost estimate next to its token usage.
7
+
8
+ The number this module returns is an estimate, not a bill. It multiplies token
9
+ counts by list rates and ignores prompt caching, batch discounts, negotiated
10
+ pricing, and per-request minimums, so it can only ever approximate an invoice.
11
+
12
+ Cost is metadata. It never feeds case generation or a verdict, and the core
13
+ judging path does not import it; only the report layer reads it. Runtime price
14
+ overrides (register_prices / load_prices) change the reported cost of a run but
15
+ cannot change which cases are generated or how they are scored.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import json
21
+ from pathlib import Path
22
+
23
+ __all__ = ["cost_usd", "load_prices", "prices", "register_prices", "reset_prices"]
24
+
25
+ # Keys are model-id prefixes. A lookup matches an id exactly first, then falls
26
+ # back to the longest key that is a prefix of the id, so a dated snapshot like
27
+ # "claude-sonnet-5-20260101" bills at the "claude-sonnet-5" rate.
28
+ _DEFAULT_PRICES: dict[str, tuple[float, float]] = {
29
+ # Claude (first-party list rates as of 2026-08).
30
+ "claude-fable-5": (10.0, 50.0),
31
+ "claude-opus-5": (5.0, 25.0),
32
+ "claude-opus-4-8": (5.0, 25.0),
33
+ "claude-opus-4-7": (5.0, 25.0),
34
+ "claude-opus-4-6": (5.0, 25.0),
35
+ "claude-sonnet-5": (3.0, 15.0),
36
+ "claude-sonnet-4-6": (3.0, 15.0),
37
+ "claude-haiku-4-5": (1.0, 5.0),
38
+ # Gemini (Google list rates as of 2026-08; short-context tier where tiered).
39
+ "gemini-2.5-pro": (1.25, 10.0),
40
+ "gemini-2.5-flash": (0.30, 2.50),
41
+ "gemini-2.5-flash-lite": (0.10, 0.40),
42
+ "gemini-2.0-flash": (0.10, 0.40),
43
+ "gemini-1.5-pro": (1.25, 5.0),
44
+ "gemini-1.5-flash": (0.075, 0.30),
45
+ # Gemini 3.x flash tiers actually exercised by live runs here. Their public
46
+ # list prices are not yet confirmed, so these carry the same flash / flash-
47
+ # lite estimates as the 2.5 tier; override with register_prices once the
48
+ # real rates land. The longest-prefix match means "gemini-3-flash-preview"
49
+ # and "gemini-3.7-flash" both resolve to the flash estimate below.
50
+ "gemini-3-flash": (0.30, 2.50),
51
+ "gemini-3.5-flash": (0.30, 2.50),
52
+ "gemini-3.6-flash": (0.30, 2.50),
53
+ "gemini-3.7-flash": (0.30, 2.50),
54
+ "gemini-flash-latest": (0.30, 2.50),
55
+ "gemini-3.1-flash-lite": (0.10, 0.40),
56
+ "gemini-flash-lite-latest": (0.10, 0.40),
57
+ }
58
+
59
+ # Mutable working copy. register_prices / load_prices update it in place;
60
+ # reset_prices restores the defaults above.
61
+ _PRICES: dict[str, tuple[float, float]] = dict(_DEFAULT_PRICES)
62
+
63
+
64
+ def prices() -> dict[str, tuple[float, float]]:
65
+ """A copy of the current price table, safe to inspect but not mutate."""
66
+ return dict(_PRICES)
67
+
68
+
69
+ def register_prices(mapping: dict[str, tuple[float, float]]) -> None:
70
+ """Merge runtime price overrides into the table, keyed by model id or prefix.
71
+
72
+ Each value is an (input_per_million, output_per_million) pair. Existing
73
+ keys are replaced; keys not named are left alone.
74
+ """
75
+ for model, pair in mapping.items():
76
+ in_rate, out_rate = pair
77
+ _PRICES[str(model)] = (float(in_rate), float(out_rate))
78
+
79
+
80
+ def reset_prices() -> None:
81
+ """Restore the built-in default table, discarding any runtime overrides."""
82
+ _PRICES.clear()
83
+ _PRICES.update(_DEFAULT_PRICES)
84
+
85
+
86
+ def _coerce_pair(value: object) -> tuple[float, float]:
87
+ """Read one price entry as (input, output) from a list or an input/output map."""
88
+ if isinstance(value, dict):
89
+ return (float(value["input"]), float(value["output"]))
90
+ # A two-element sequence: [input, output].
91
+ in_rate, out_rate = value # type: ignore[misc]
92
+ return (float(in_rate), float(out_rate))
93
+
94
+
95
+ def load_prices(path: str | Path) -> dict[str, tuple[float, float]]:
96
+ """Load price overrides from a JSON file and register them.
97
+
98
+ The JSON maps model id (or prefix) to either a two-element ``[input,
99
+ output]`` array or an ``{"input": ..., "output": ...}`` object. Returns the
100
+ mapping that was registered.
101
+ """
102
+ data = json.loads(Path(path).read_text(encoding="utf-8"))
103
+ mapping = {str(model): _coerce_pair(value) for model, value in data.items()}
104
+ register_prices(mapping)
105
+ return mapping
106
+
107
+
108
+ def _lookup(model: str) -> tuple[float, float] | None:
109
+ """Find the price pair for a model id: exact match, then longest prefix."""
110
+ if model in _PRICES:
111
+ return _PRICES[model]
112
+ candidates = [key for key in _PRICES if model.startswith(key)]
113
+ if not candidates:
114
+ return None
115
+ return _PRICES[max(candidates, key=len)]
116
+
117
+
118
+ def cost_usd(
119
+ model: str | None,
120
+ prompt_tokens: int | None,
121
+ completion_tokens: int | None,
122
+ ) -> float | None:
123
+ """Estimate the USD cost of one run's usage, or None when it cannot be priced.
124
+
125
+ Returns None when the model is unpriced (no exact or prefix match) or when
126
+ either token count is None, so callers never mistake a missing figure for
127
+ a real zero.
128
+ """
129
+ if model is None or prompt_tokens is None or completion_tokens is None:
130
+ return None
131
+ rate = _lookup(model)
132
+ if rate is None:
133
+ return None
134
+ in_rate, out_rate = rate
135
+ return prompt_tokens / 1_000_000 * in_rate + completion_tokens / 1_000_000 * out_rate
hackbait/py.typed ADDED
File without changes