quantdiff 0.1.0rc1__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.
Files changed (45) hide show
  1. quantdiff/__init__.py +53 -0
  2. quantdiff/__main__.py +5 -0
  3. quantdiff/_http.py +151 -0
  4. quantdiff/_text.py +13 -0
  5. quantdiff/_version.py +1 -0
  6. quantdiff/api.py +340 -0
  7. quantdiff/backends/__init__.py +28 -0
  8. quantdiff/backends/_common.py +342 -0
  9. quantdiff/backends/base.py +91 -0
  10. quantdiff/backends/llamacpp.py +428 -0
  11. quantdiff/backends/ollama.py +359 -0
  12. quantdiff/backends/openai_compat.py +338 -0
  13. quantdiff/cache.py +240 -0
  14. quantdiff/card.py +1664 -0
  15. quantdiff/cli.py +377 -0
  16. quantdiff/discover.py +488 -0
  17. quantdiff/errors.py +45 -0
  18. quantdiff/metrics/__init__.py +36 -0
  19. quantdiff/metrics/codeexec.py +428 -0
  20. quantdiff/metrics/jsonschema.py +610 -0
  21. quantdiff/metrics/logit.py +214 -0
  22. quantdiff/metrics/tasks.py +114 -0
  23. quantdiff/metrics/textsim.py +66 -0
  24. quantdiff/metrics/toolcheck.py +99 -0
  25. quantdiff/png.py +360 -0
  26. quantdiff/preflight.py +365 -0
  27. quantdiff/progress.py +283 -0
  28. quantdiff/py.typed +0 -0
  29. quantdiff/report.py +780 -0
  30. quantdiff/runner.py +492 -0
  31. quantdiff/spec.py +154 -0
  32. quantdiff/stats.py +226 -0
  33. quantdiff/suites/__init__.py +462 -0
  34. quantdiff/suites/data/chat.jsonl +22 -0
  35. quantdiff/suites/data/code.jsonl +32 -0
  36. quantdiff/suites/data/json.jsonl +34 -0
  37. quantdiff/suites/data/scoring.jsonl +41 -0
  38. quantdiff/suites/data/tools.jsonl +32 -0
  39. quantdiff/types.py +322 -0
  40. quantdiff/verdict.py +1513 -0
  41. quantdiff-0.1.0rc1.dist-info/METADATA +514 -0
  42. quantdiff-0.1.0rc1.dist-info/RECORD +45 -0
  43. quantdiff-0.1.0rc1.dist-info/WHEEL +4 -0
  44. quantdiff-0.1.0rc1.dist-info/entry_points.txt +2 -0
  45. quantdiff-0.1.0rc1.dist-info/licenses/LICENSE +202 -0
quantdiff/stats.py ADDED
@@ -0,0 +1,226 @@
1
+ """Paired significance tests and confidence intervals for comparing two models.
2
+
3
+ Every candidate answers the same cases and is scored on the same prompts as the reference,
4
+ so comparisons are paired: each case contributes one (reference, candidate) observation.
5
+ Pairing removes the case-to-case difficulty variation that dominates small suites, which is
6
+ why these tests separate models that an unpaired test calls "within noise".
7
+
8
+ Everything here is pure and deterministic. Bootstrap resampling uses its own seeded
9
+ generator, so the same input always yields the same interval.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import math
15
+ import random
16
+ from collections.abc import Sequence
17
+ from dataclasses import dataclass
18
+ from statistics import NormalDist
19
+ from typing import Final
20
+
21
+ __all__ = [
22
+ "DEFAULT_RESAMPLES",
23
+ "Interval",
24
+ "bootstrap_mean",
25
+ "cases_to_bound_loss",
26
+ "items_to_bound_below",
27
+ "mcnemar_exact",
28
+ "paired_proportion_diff",
29
+ ]
30
+
31
+ DEFAULT_RESAMPLES: Final = 2000
32
+ _CONFIDENCE: Final = 0.95
33
+ _Z_TWO_SIDED: Final = NormalDist().inv_cdf(1 - (1 - _CONFIDENCE) / 2)
34
+ _MAX_CASES_SEARCHED: Final = 1_000_000
35
+
36
+
37
+ @dataclass(frozen=True, slots=True)
38
+ class Interval:
39
+ """A point estimate with a 95% confidence interval, in the metric's own units."""
40
+
41
+ estimate: float
42
+ low: float
43
+ high: float
44
+
45
+ @property
46
+ def excludes_zero(self) -> bool:
47
+ return self.low > 0.0 or self.high < 0.0
48
+
49
+
50
+ def mcnemar_exact(ref: Sequence[bool], cand: Sequence[bool]) -> float:
51
+ """Two-sided exact McNemar p-value for paired pass/fail outcomes.
52
+
53
+ Only discordant pairs carry information: b cases the reference passed and the candidate
54
+ failed, c the other way round. Under equal pass rates each discordant pair is a fair
55
+ coin, so the p-value is the two-sided binomial tail of min(b, c) in b + c trials,
56
+ capped at 1. With no discordant pairs there is no evidence of a difference: 1.0.
57
+ """
58
+ _check_paired(ref, cand)
59
+ lost = sum(1 for r, c in zip(ref, cand, strict=True) if r and not c)
60
+ gained = sum(1 for r, c in zip(ref, cand, strict=True) if c and not r)
61
+ trials = lost + gained
62
+ if trials == 0:
63
+ return 1.0
64
+ tail = sum(math.comb(trials, k) for k in range(min(lost, gained) + 1))
65
+ # Integer division of exact big ints keeps full precision even for thousands of trials.
66
+ return min(1.0, 2 * tail / (1 << trials))
67
+
68
+
69
+ def paired_proportion_diff(ref: Sequence[bool], cand: Sequence[bool]) -> Interval:
70
+ """Candidate minus reference pass rate, in percentage points, with a 95% interval.
71
+
72
+ The interval is Newcombe's hybrid score method for paired proportions (method 10 in
73
+ Newcombe, Statistics in Medicine 17:2635, 1998). It combines the Wilson score interval
74
+ of each rate with the observed correlation between the paired outcomes, stays inside
75
+ [-100, 100], and keeps sensible coverage at 0% and 100% where the Wald interval
76
+ collapses to a point.
77
+ """
78
+ _check_paired(ref, cand)
79
+ if not ref:
80
+ raise ValueError("paired_proportion_diff needs at least one pair")
81
+ return _newcombe(*_table(ref, cand))
82
+
83
+
84
+ def bootstrap_mean(
85
+ values: Sequence[float],
86
+ *,
87
+ resamples: int = DEFAULT_RESAMPLES,
88
+ seed: int = 0,
89
+ ) -> Interval:
90
+ """Mean of `values` with a 95% percentile bootstrap interval. Deterministic for a seed."""
91
+ if not values:
92
+ raise ValueError("bootstrap_mean needs at least one value")
93
+ if resamples < 1:
94
+ raise ValueError("resamples must be at least 1")
95
+ n = len(values)
96
+ estimate = math.fsum(values) / n
97
+ rng = random.Random(seed) # noqa: S311 - seeded for reproducible intervals, not secrecy
98
+ means = sorted(
99
+ math.fsum(values[rng.randrange(n)] for _ in range(n)) / n for _ in range(resamples)
100
+ )
101
+ tail = (1 - _CONFIDENCE) / 2
102
+ return Interval(
103
+ estimate=estimate,
104
+ low=min(estimate, _percentile(means, tail)),
105
+ high=max(estimate, _percentile(means, 1 - tail)),
106
+ )
107
+
108
+
109
+ def items_to_bound_below(interval: Interval, items: int, bound: float) -> int | None:
110
+ """Roughly how many items would put the upper end of `interval` below `bound`.
111
+
112
+ `interval` is a 95% interval of a mean over `items` items. Its upper half-width shrinks
113
+ with the square root of the item count, so the estimate keeps the observed mean and
114
+ spread and scales the half-width by sqrt(items / n). None when the observed mean is
115
+ already at or above `bound`, since more items would only narrow the interval around it.
116
+ """
117
+ if items < 1:
118
+ raise ValueError("items must be at least 1")
119
+ room = bound - interval.estimate
120
+ if room <= 0.0:
121
+ return None
122
+ half_width = interval.high - interval.estimate
123
+ return max(items, math.floor(items * (half_width / room) ** 2) + 1)
124
+
125
+
126
+ def cases_to_bound_loss(ref: Sequence[bool], cand: Sequence[bool], margin: float) -> int | None:
127
+ """Roughly how many paired cases would put the lower end of the pass-rate difference
128
+ (candidate minus reference, in points) at or above -`margin` points.
129
+
130
+ The observed shares of the four outcomes (both pass, only the reference passes, only
131
+ the candidate passes, neither) are held fixed while the case count grows, and the
132
+ Newcombe interval of paired_proportion_diff is recomputed until its lower bound clears
133
+ the margin. None when the observed difference is already at or below -`margin`, since
134
+ more cases would only narrow the interval around it.
135
+ """
136
+ _check_paired(ref, cand)
137
+ if not ref:
138
+ raise ValueError("cases_to_bound_loss needs at least one pair")
139
+ if margin <= 0.0:
140
+ raise ValueError("margin must be positive")
141
+ n = len(ref)
142
+ shares = [count / n for count in _table(ref, cand)]
143
+
144
+ def clears(cases: int) -> bool:
145
+ both, lost, gained, neither = (share * cases for share in shares)
146
+ return _newcombe(both, lost, gained, neither).low >= -margin
147
+
148
+ if _newcombe(*shares).estimate <= -margin:
149
+ return None
150
+ if clears(n):
151
+ return n
152
+ low, high = n, 2 * n
153
+ while not clears(high):
154
+ if high >= _MAX_CASES_SEARCHED:
155
+ return None
156
+ low, high = high, 2 * high
157
+ while high - low > 1:
158
+ middle = (low + high) // 2
159
+ if clears(middle):
160
+ high = middle
161
+ else:
162
+ low = middle
163
+ return high
164
+
165
+
166
+ def _table(ref: Sequence[bool], cand: Sequence[bool]) -> tuple[int, int, int, int]:
167
+ """Counts of (both pass, only the reference passes, only the candidate passes, neither)."""
168
+ both = sum(1 for r, c in zip(ref, cand, strict=True) if r and c)
169
+ lost = sum(1 for r, c in zip(ref, cand, strict=True) if r and not c)
170
+ gained = sum(1 for r, c in zip(ref, cand, strict=True) if c and not r)
171
+ return both, lost, gained, len(ref) - both - lost - gained
172
+
173
+
174
+ def _newcombe(both: float, lost: float, gained: float, neither: float) -> Interval:
175
+ """Newcombe's hybrid score interval for a paired difference, from (possibly scaled)
176
+ cell counts of the 2x2 table, in percentage points."""
177
+ n = both + lost + gained + neither
178
+ p_ref, p_cand = (both + lost) / n, (both + gained) / n
179
+ low_ref, high_ref = _wilson(both + lost, n)
180
+ low_cand, high_cand = _wilson(both + gained, n)
181
+ margins = (both + lost) * (gained + neither) * (both + gained) * (lost + neither)
182
+ phi = 0.0 if margins == 0 else (both * neither - lost * gained) / math.sqrt(margins)
183
+ estimate = p_cand - p_ref
184
+ below = math.sqrt(
185
+ max(
186
+ 0.0,
187
+ (p_cand - low_cand) ** 2
188
+ - 2 * phi * (p_cand - low_cand) * (high_ref - p_ref)
189
+ + (high_ref - p_ref) ** 2,
190
+ )
191
+ )
192
+ above = math.sqrt(
193
+ max(
194
+ 0.0,
195
+ (high_cand - p_cand) ** 2
196
+ - 2 * phi * (high_cand - p_cand) * (p_ref - low_ref)
197
+ + (p_ref - low_ref) ** 2,
198
+ )
199
+ )
200
+ return Interval(
201
+ estimate=estimate * 100,
202
+ low=max(-1.0, estimate - below) * 100,
203
+ high=min(1.0, estimate + above) * 100,
204
+ )
205
+
206
+
207
+ def _wilson(successes: float, n: float) -> tuple[float, float]:
208
+ p = successes / n
209
+ z2 = _Z_TWO_SIDED * _Z_TWO_SIDED
210
+ centre = (p + z2 / (2 * n)) / (1 + z2 / n)
211
+ half = _Z_TWO_SIDED * math.sqrt(p * (1 - p) / n + z2 / (4 * n * n)) / (1 + z2 / n)
212
+ return max(0.0, centre - half), min(1.0, centre + half)
213
+
214
+
215
+ def _percentile(ordered: Sequence[float], q: float) -> float:
216
+ """Linear interpolation between closest ranks, as numpy's default percentile does."""
217
+ position = q * (len(ordered) - 1)
218
+ lower = math.floor(position)
219
+ upper = min(lower + 1, len(ordered) - 1)
220
+ weight = position - lower
221
+ return ordered[lower] * (1 - weight) + ordered[upper] * weight
222
+
223
+
224
+ def _check_paired(a: Sequence[object], b: Sequence[object]) -> None:
225
+ if len(a) != len(b):
226
+ raise ValueError(f"paired samples differ in length: {len(a)} and {len(b)}")
@@ -0,0 +1,462 @@
1
+ """Built-in prompt suites and strict loaders for user prompt files.
2
+
3
+ Suites are JSON Lines files: one task case (or scoring prompt) per line. Built-in
4
+ suites ship inside the package and are read with importlib.resources, so they load the
5
+ same way from a source checkout and from an installed wheel.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import hashlib
11
+ import json
12
+ import keyword
13
+ import logging
14
+ import os
15
+ import re
16
+ from collections.abc import Callable, Collection, Mapping, Sequence
17
+ from dataclasses import replace
18
+ from importlib import resources
19
+ from pathlib import Path
20
+ from typing import Final, Protocol, TypeVar
21
+
22
+ from quantdiff.errors import SuiteError
23
+ from quantdiff.metrics.jsonschema import SUPPORTED_KEYWORDS, check_schema
24
+ from quantdiff.types import JSONValue, Message, Role, ScoringPrompt, TaskCase, TaskKind, ToolSpec
25
+
26
+ __all__ = [
27
+ "BUILTIN_SUITES",
28
+ "MAX_CASES",
29
+ "MAX_FILE_BYTES",
30
+ "SCHEMA_KEYWORDS",
31
+ "case_to_dict",
32
+ "load_builtin",
33
+ "load_cases_file",
34
+ "load_scoring_prompts",
35
+ "scoring_prompts_from_cases",
36
+ "suite_digest",
37
+ ]
38
+
39
+ logger = logging.getLogger(__name__)
40
+
41
+ BUILTIN_SUITES: Final[tuple[str, ...]] = ("json", "tools", "code", "chat")
42
+ MAX_FILE_BYTES: Final = 20 * 1024 * 1024
43
+ MAX_CASES: Final = 10_000
44
+
45
+ _DATA_PACKAGE: Final = "quantdiff.suites"
46
+ _DEFAULT_MAX_TOKENS: Final = 512
47
+ _MAX_MAX_TOKENS: Final = 32_768
48
+ _ID_PATTERN: Final = re.compile(r"[A-Za-z0-9][A-Za-z0-9._-]{0,127}")
49
+ _TOOL_NAME_PATTERN: Final = re.compile(r"[A-Za-z_][A-Za-z0-9_-]{0,63}")
50
+ _ROLES: Final[Mapping[str, Role]] = {"system": "system", "user": "user", "assistant": "assistant"}
51
+ _KINDS: Final[Mapping[str, TaskKind]] = {
52
+ "json": "json",
53
+ "tools": "tools",
54
+ "code": "code",
55
+ "chat": "chat",
56
+ }
57
+ _COMMON_FIELDS: Final = frozenset({"id", "kind", "messages", "max_tokens"})
58
+ _KIND_REQUIRED: Final[Mapping[TaskKind, frozenset[str]]] = {
59
+ "json": frozenset({"json_schema"}),
60
+ "tools": frozenset({"tools", "expected_tool"}),
61
+ "code": frozenset({"entry_point", "tests"}),
62
+ "chat": frozenset(),
63
+ }
64
+ _KIND_OPTIONAL: Final[Mapping[TaskKind, frozenset[str]]] = {
65
+ "json": frozenset(),
66
+ "tools": frozenset({"expected_arguments"}),
67
+ "code": frozenset(),
68
+ "chat": frozenset(),
69
+ }
70
+ _PROMPT_FIELDS: Final = frozenset({"prompt", "id", "system", "max_tokens"})
71
+ _MESSAGES_FIELDS: Final = frozenset({"messages", "id", "max_tokens"})
72
+ _SCORING_FIELDS: Final = frozenset({"id", "text"})
73
+
74
+
75
+ class _Identified(Protocol):
76
+ @property
77
+ def id(self) -> str: ...
78
+
79
+
80
+ _T = TypeVar("_T", bound=_Identified)
81
+ Record = dict[str, JSONValue]
82
+
83
+
84
+ # Public API ------------------------------------------------------------------------------
85
+
86
+
87
+ def load_builtin(name: str) -> tuple[TaskCase, ...]:
88
+ """Load one of the suites listed in BUILTIN_SUITES."""
89
+ if name not in BUILTIN_SUITES:
90
+ raise SuiteError(f"unknown suite {name!r}; choose from: {', '.join(BUILTIN_SUITES)}")
91
+ return _parse_records(_builtin_text(f"{name}.jsonl"), f"built-in suite {name!r}", _build_case)
92
+
93
+
94
+ def load_cases_file(path: str | os.PathLike[str]) -> tuple[TaskCase, ...]:
95
+ """Load and validate a user prompts file.
96
+
97
+ Besides full task cases, each line may be the shorthand `{"prompt": "..."}` (with
98
+ optional "id", "system" and "max_tokens") or `{"messages": [...]}`; both become chat
99
+ cases.
100
+ """
101
+ file_path = Path(path)
102
+ return _parse_records(_read_text(file_path), str(file_path), _build_case)
103
+
104
+
105
+ def load_scoring_prompts(path: str | os.PathLike[str] | None = None) -> tuple[ScoringPrompt, ...]:
106
+ """Load raw-completion prompts for logit metrics; None selects the built-in set.
107
+
108
+ Each line is `{"text": "..."}` with an optional "id", or any line a prompts file
109
+ accepts, in which case the last user message is scored as raw text. So one file can
110
+ serve as both --prompts and --scoring-prompts.
111
+ """
112
+ if path is None:
113
+ return _parse_records(
114
+ _builtin_text("scoring.jsonl"), "built-in scoring prompts", _build_scoring_prompt
115
+ )
116
+ file_path = Path(path)
117
+ return _parse_records(_read_text(file_path), str(file_path), _build_scoring_prompt)
118
+
119
+
120
+ def scoring_prompts_from_cases(cases: Sequence[TaskCase]) -> tuple[ScoringPrompt, ...]:
121
+ """Score each case's last user message as raw text, keeping the case id."""
122
+ return tuple(_scoring_prompt_from_case(case) for case in cases)
123
+
124
+
125
+ def case_to_dict(case: TaskCase) -> dict[str, JSONValue]:
126
+ """Return the canonical file form of `case`; load_cases_file accepts it unchanged."""
127
+ out: dict[str, JSONValue] = {
128
+ "id": case.id,
129
+ "kind": case.kind,
130
+ "messages": [{"role": m.role, "content": m.content} for m in case.messages],
131
+ "max_tokens": case.max_tokens,
132
+ }
133
+ if case.json_schema is not None:
134
+ out["json_schema"] = case.json_schema
135
+ if case.tools:
136
+ out["tools"] = [
137
+ {"name": t.name, "description": t.description, "parameters": t.parameters}
138
+ for t in case.tools
139
+ ]
140
+ if case.kind == "tools":
141
+ out["expected_tool"] = case.expected_tool
142
+ if case.expected_arguments is not None:
143
+ out["expected_arguments"] = case.expected_arguments
144
+ if case.entry_point is not None:
145
+ out["entry_point"] = case.entry_point
146
+ if case.tests is not None:
147
+ out["tests"] = case.tests
148
+ return out
149
+
150
+
151
+ def suite_digest(cases: Sequence[TaskCase], scoring: Sequence[ScoringPrompt]) -> str:
152
+ """Return a sha256 hex digest that changes whenever any case or prompt changes."""
153
+ payload = {
154
+ "cases": [case_to_dict(case) for case in cases],
155
+ "scoring": [{"id": prompt.id, "text": prompt.text} for prompt in scoring],
156
+ }
157
+ canonical = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
158
+ return hashlib.sha256(canonical.encode("utf-8")).hexdigest()
159
+
160
+
161
+ # Reading ---------------------------------------------------------------------------------
162
+
163
+
164
+ def _builtin_text(filename: str) -> str:
165
+ return resources.files(_DATA_PACKAGE).joinpath("data").joinpath(filename).read_text("utf-8")
166
+
167
+
168
+ def _read_text(path: Path) -> str:
169
+ try:
170
+ if not path.exists():
171
+ raise SuiteError(f"prompts file not found: {path}")
172
+ if not path.is_file():
173
+ raise SuiteError(f"{path} is not a file")
174
+ if path.stat().st_size > MAX_FILE_BYTES:
175
+ raise SuiteError(f"{path} is larger than {MAX_FILE_BYTES // (1024 * 1024)} MB")
176
+ data = path.read_bytes()
177
+ except OSError as exc:
178
+ raise SuiteError(f"cannot read {path}: {exc.strerror or exc}") from None
179
+ try:
180
+ return data.decode("utf-8-sig")
181
+ except UnicodeDecodeError as exc:
182
+ raise SuiteError(f"{path} is not valid UTF-8 (byte offset {exc.start})") from None
183
+
184
+
185
+ def _parse_records(text: str, source: str, build: Callable[[Record, int], _T]) -> tuple[_T, ...]:
186
+ items: list[_T] = []
187
+ first_seen: dict[str, int] = {}
188
+ # str.splitlines would also split on U+2028 and friends, which JSON allows raw in strings.
189
+ for number, line in enumerate(text.split("\n"), start=1):
190
+ if not line.strip():
191
+ continue
192
+ if len(items) == MAX_CASES:
193
+ raise SuiteError(f"{source} has more than {MAX_CASES} entries")
194
+ try:
195
+ item = build(_decode_object(line), number)
196
+ except SuiteError as exc:
197
+ raise SuiteError(f"{source}, line {number}: {exc}") from None
198
+ if item.id in first_seen:
199
+ raise SuiteError(
200
+ f"{source}, line {number}: duplicate id {item.id!r} "
201
+ f"(first used on line {first_seen[item.id]})"
202
+ )
203
+ first_seen[item.id] = number
204
+ items.append(item)
205
+ if not items:
206
+ raise SuiteError(f"{source} contains no entries")
207
+ logger.debug("loaded %d entries from %s", len(items), source)
208
+ return tuple(items)
209
+
210
+
211
+ def _decode_object(line: str) -> Record:
212
+ try:
213
+ value = json.loads(line, object_pairs_hook=_unique_keys, parse_constant=_reject_constant)
214
+ except json.JSONDecodeError as exc:
215
+ raise SuiteError(f"invalid JSON: {exc.msg} (column {exc.colno})") from None
216
+ if not isinstance(value, dict):
217
+ raise SuiteError("each line must be a JSON object")
218
+ return value
219
+
220
+
221
+ def _unique_keys(pairs: list[tuple[str, JSONValue]]) -> Record:
222
+ record: Record = {}
223
+ for key, value in pairs:
224
+ if key in record:
225
+ raise SuiteError(f"duplicate key {key!r}")
226
+ record[key] = value
227
+ return record
228
+
229
+
230
+ def _reject_constant(name: str) -> JSONValue:
231
+ raise SuiteError(f"{name} is not valid JSON")
232
+
233
+
234
+ # Task cases ------------------------------------------------------------------------------
235
+
236
+
237
+ def _build_case(record: Record, line: int) -> TaskCase:
238
+ if "kind" in record:
239
+ return _full_case(record)
240
+ if "prompt" in record:
241
+ return _prompt_case(record, line)
242
+ if "messages" in record:
243
+ return _messages_case(record, line)
244
+ raise SuiteError('expected a "kind", "prompt" or "messages" field')
245
+
246
+
247
+ def _full_case(record: Record) -> TaskCase:
248
+ kind = _kind(record["kind"])
249
+ required = frozenset({"id", "kind", "messages"}) | _KIND_REQUIRED[kind]
250
+ _check_fields(
251
+ record, allowed=_COMMON_FIELDS | required | _KIND_OPTIONAL[kind], required=required
252
+ )
253
+ base = TaskCase(
254
+ id=_case_id(record["id"]),
255
+ kind=kind,
256
+ messages=_messages(record["messages"]),
257
+ max_tokens=_max_tokens(record.get("max_tokens", _DEFAULT_MAX_TOKENS)),
258
+ )
259
+ if kind == "json":
260
+ return _with_json_schema(base, record)
261
+ if kind == "tools":
262
+ return _with_tools(base, record)
263
+ if kind == "code":
264
+ return _with_code(base, record)
265
+ return base
266
+
267
+
268
+ def _prompt_case(record: Record, line: int) -> TaskCase:
269
+ _check_fields(record, allowed=_PROMPT_FIELDS, required={"prompt"})
270
+ messages = [Message(role="user", content=_text(record["prompt"], "prompt"))]
271
+ if "system" in record:
272
+ messages.insert(0, Message(role="system", content=_text(record["system"], "system")))
273
+ return TaskCase(
274
+ id=_case_id(record.get("id", f"prompt-{line:03d}")),
275
+ kind="chat",
276
+ messages=tuple(messages),
277
+ max_tokens=_max_tokens(record.get("max_tokens", _DEFAULT_MAX_TOKENS)),
278
+ )
279
+
280
+
281
+ def _messages_case(record: Record, line: int) -> TaskCase:
282
+ _check_fields(record, allowed=_MESSAGES_FIELDS, required={"messages"})
283
+ return TaskCase(
284
+ id=_case_id(record.get("id", f"prompt-{line:03d}")),
285
+ kind="chat",
286
+ messages=_messages(record["messages"]),
287
+ max_tokens=_max_tokens(record.get("max_tokens", _DEFAULT_MAX_TOKENS)),
288
+ )
289
+
290
+
291
+ def _with_json_schema(base: TaskCase, record: Record) -> TaskCase:
292
+ schema = _object_schema(record["json_schema"], "json_schema")
293
+ return replace(base, json_schema=schema)
294
+
295
+
296
+ def _with_tools(base: TaskCase, record: Record) -> TaskCase:
297
+ tools = _tools(record["tools"])
298
+ expected_tool = record["expected_tool"]
299
+ expected_arguments = record.get("expected_arguments")
300
+ if expected_tool is None:
301
+ if expected_arguments is not None:
302
+ raise SuiteError('"expected_arguments" requires a non-null "expected_tool"')
303
+ return replace(base, tools=tools)
304
+ by_name = {tool.name: tool for tool in tools}
305
+ if not isinstance(expected_tool, str) or expected_tool not in by_name:
306
+ raise SuiteError(f'"expected_tool" must be null or one of: {", ".join(by_name)}')
307
+ if expected_arguments is not None:
308
+ _check_expected_arguments(expected_arguments, by_name[expected_tool])
309
+ return replace(
310
+ base, tools=tools, expected_tool=expected_tool, expected_arguments=expected_arguments
311
+ )
312
+
313
+
314
+ def _with_code(base: TaskCase, record: Record) -> TaskCase:
315
+ entry_point = record["entry_point"]
316
+ if (
317
+ not isinstance(entry_point, str)
318
+ or not entry_point.isidentifier()
319
+ or keyword.iskeyword(entry_point)
320
+ ):
321
+ raise SuiteError('"entry_point" must be a valid Python function name')
322
+ return replace(base, entry_point=entry_point, tests=_text(record["tests"], "tests"))
323
+
324
+
325
+ def _kind(value: JSONValue) -> TaskKind:
326
+ kind = _KINDS.get(value) if isinstance(value, str) else None
327
+ if kind is None:
328
+ raise SuiteError(f'"kind" must be one of: {", ".join(_KINDS)}')
329
+ return kind
330
+
331
+
332
+ def _case_id(value: JSONValue) -> str:
333
+ if not isinstance(value, str) or not _ID_PATTERN.fullmatch(value):
334
+ raise SuiteError(
335
+ '"id" must be 1 to 128 characters of letters, digits, ".", "_" or "-", '
336
+ "starting with a letter or digit"
337
+ )
338
+ return value
339
+
340
+
341
+ def _max_tokens(value: JSONValue) -> int:
342
+ if isinstance(value, int) and not isinstance(value, bool) and 1 <= value <= _MAX_MAX_TOKENS:
343
+ return value
344
+ raise SuiteError(f'"max_tokens" must be an integer from 1 to {_MAX_MAX_TOKENS}')
345
+
346
+
347
+ def _text(value: JSONValue, field: str) -> str:
348
+ if not isinstance(value, str) or not value.strip():
349
+ raise SuiteError(f'"{field}" must be a non-empty string')
350
+ return value
351
+
352
+
353
+ def _messages(value: JSONValue) -> tuple[Message, ...]:
354
+ if not isinstance(value, list) or not value:
355
+ raise SuiteError('"messages" must be a non-empty array')
356
+ messages = tuple(_message(item, index) for index, item in enumerate(value))
357
+ if messages[-1].role != "user":
358
+ raise SuiteError('the last message must have role "user"')
359
+ return messages
360
+
361
+
362
+ def _message(item: JSONValue, index: int) -> Message:
363
+ where = f"messages[{index}]"
364
+ if not isinstance(item, dict):
365
+ raise SuiteError(f"{where} must be an object")
366
+ _check_fields(item, allowed={"role", "content"}, required={"role", "content"}, where=where)
367
+ role = _ROLES.get(item["role"]) if isinstance(item["role"], str) else None
368
+ if role is None:
369
+ raise SuiteError(f"{where}.role must be one of: {', '.join(_ROLES)}")
370
+ return Message(role=role, content=_text(item["content"], f"{where}.content"))
371
+
372
+
373
+ def _check_fields(
374
+ record: Mapping[str, JSONValue],
375
+ *,
376
+ allowed: Collection[str],
377
+ required: Collection[str],
378
+ where: str = "",
379
+ ) -> None:
380
+ prefix = f"{where}: " if where else ""
381
+ unknown = sorted(set(record) - set(allowed))
382
+ if unknown:
383
+ raise SuiteError(f"{prefix}unknown field {unknown[0]!r}")
384
+ missing = sorted(set(required) - set(record))
385
+ if missing:
386
+ raise SuiteError(f"{prefix}missing field {missing[0]!r}")
387
+
388
+
389
+ # Tools -----------------------------------------------------------------------------------
390
+
391
+
392
+ def _tools(value: JSONValue) -> tuple[ToolSpec, ...]:
393
+ if not isinstance(value, list) or not value:
394
+ raise SuiteError('"tools" must be a non-empty array')
395
+ tools = tuple(_tool(item, index) for index, item in enumerate(value))
396
+ names = [tool.name for tool in tools]
397
+ duplicates = sorted({name for name in names if names.count(name) > 1})
398
+ if duplicates:
399
+ raise SuiteError(f"duplicate tool name {duplicates[0]!r}")
400
+ return tools
401
+
402
+
403
+ def _tool(item: JSONValue, index: int) -> ToolSpec:
404
+ where = f"tools[{index}]"
405
+ if not isinstance(item, dict):
406
+ raise SuiteError(f"{where} must be an object")
407
+ fields = {"name", "description", "parameters"}
408
+ _check_fields(item, allowed=fields, required=fields, where=where)
409
+ name = item["name"]
410
+ if not isinstance(name, str) or not _TOOL_NAME_PATTERN.fullmatch(name):
411
+ raise SuiteError(f"{where}.name must be a function name such as get_weather")
412
+ if not isinstance(item["description"], str):
413
+ raise SuiteError(f"{where}.description must be a string")
414
+ parameters = _object_schema(item["parameters"], f"{where}.parameters")
415
+ if parameters.get("type") != "object":
416
+ raise SuiteError(f'{where}.parameters must have "type": "object"')
417
+ return ToolSpec(name=name, description=item["description"], parameters=parameters)
418
+
419
+
420
+ def _check_expected_arguments(value: JSONValue, tool: ToolSpec) -> None:
421
+ if not isinstance(value, dict):
422
+ raise SuiteError('"expected_arguments" must be an object or null')
423
+ properties = tool.parameters.get("properties", {})
424
+ unknown = sorted(set(value) - set(properties))
425
+ if unknown:
426
+ raise SuiteError(f"expected_arguments key {unknown[0]!r} is not a parameter of {tool.name}")
427
+
428
+
429
+ # JSON Schema subset ----------------------------------------------------------------------
430
+
431
+
432
+ SCHEMA_KEYWORDS: Final = SUPPORTED_KEYWORDS
433
+ """JSON Schema keywords the task metrics can check. Anything else is rejected up front."""
434
+
435
+
436
+ def _object_schema(value: JSONValue, where: str) -> dict[str, JSONValue]:
437
+ if not isinstance(value, dict):
438
+ raise SuiteError(f"{where} must be a JSON object")
439
+ check_schema(value, where)
440
+ return value
441
+
442
+
443
+ # Scoring prompts -------------------------------------------------------------------------
444
+
445
+
446
+ def _build_scoring_prompt(record: Record, line: int) -> ScoringPrompt:
447
+ if "text" in record:
448
+ _check_fields(record, allowed=_SCORING_FIELDS, required={"text"})
449
+ return ScoringPrompt(
450
+ id=_case_id(record.get("id", f"score-{line:03d}")),
451
+ text=_text(record["text"], "text"),
452
+ )
453
+ if not record.keys() & {"kind", "prompt", "messages"}:
454
+ raise SuiteError('expected a "text" or "prompt" field')
455
+ return _scoring_prompt_from_case(_build_case(record, line))
456
+
457
+
458
+ def _scoring_prompt_from_case(case: TaskCase) -> ScoringPrompt:
459
+ user_texts = [message.content for message in case.messages if message.role == "user"]
460
+ if not user_texts:
461
+ raise SuiteError(f"case {case.id!r} has no user message to score")
462
+ return ScoringPrompt(id=case.id, text=user_texts[-1])