steadybatch 0.2.1__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.
@@ -0,0 +1,7 @@
1
+ """steadybatch: run millions of LLM requests through batch APIs without silent failures."""
2
+
3
+ from .models import Outcome, Request, Result
4
+ from .runner import RunReport, Runner
5
+
6
+ __all__ = ["Outcome", "Request", "Result", "RunReport", "Runner"]
7
+ __version__ = "0.2.1"
steadybatch/bench.py ADDED
@@ -0,0 +1,317 @@
1
+ """Benchmark harness: run one workload on one provider and write down what happened.
2
+
3
+ steadybatch-bench run --provider openai --model gpt-4o-mini \\
4
+ --data examples/support_tickets.jsonl --schema examples/support_ticket.schema.json \\
5
+ --out runs/openai-1 --prices examples/prices.example.json
6
+
7
+ steadybatch-bench compare runs/openai-1 runs/openai-2
8
+
9
+ ``run`` writes results.jsonl (one line per record, in input order) and
10
+ summary.json. ``compare`` checks two runs of the same workload against each
11
+ other, which is how we measure reproducibility.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import argparse
17
+ import json
18
+ import logging
19
+ import statistics
20
+ import sys
21
+ from pathlib import Path
22
+ from typing import Any
23
+
24
+ from .models import Outcome, Request
25
+ from .providers import get_provider
26
+ from .runner import Runner
27
+
28
+ SYSTEM_PROMPT = (
29
+ "You read customer support conversations and extract structured facts. "
30
+ "Answer with a single JSON object and nothing else."
31
+ )
32
+
33
+
34
+ def load_records(path: Path, limit: int | None) -> list[dict[str, Any]]:
35
+ records = []
36
+ with path.open() as f:
37
+ for line in f:
38
+ if line.strip():
39
+ records.append(json.loads(line))
40
+ if limit and len(records) >= limit:
41
+ break
42
+ return records
43
+
44
+
45
+ def to_requests(records: list[dict[str, Any]], max_tokens: int, instructions: str = "") -> list[Request]:
46
+ system = SYSTEM_PROMPT + ("\n\n" + instructions.strip() if instructions.strip() else "")
47
+ return [
48
+ Request(key=str(r["key"]), system=system, max_tokens=max_tokens, temperature=0.0,
49
+ messages=[{"role": "user", "content": r["text"]}])
50
+ for r in records
51
+ ]
52
+
53
+
54
+ def accuracy(records: list[dict[str, Any]], results_by_key: dict[str, Any]) -> dict[str, float]:
55
+ """Field-level accuracy against the 'labels' in the dataset, when present."""
56
+ scores: dict[str, list[int]] = {}
57
+ for r in records:
58
+ labels = r.get("labels")
59
+ data = results_by_key.get(str(r["key"]))
60
+ if not labels or not isinstance(data, dict):
61
+ continue
62
+ for field, expected in labels.items():
63
+ scores.setdefault(field, []).append(int(data.get(field) == expected))
64
+ return {f: round(sum(v) / len(v), 4) for f, v in scores.items() if v}
65
+
66
+
67
+ def cmd_run(args: argparse.Namespace) -> int:
68
+ records = load_records(Path(args.data), args.limit)
69
+ schema = json.loads(Path(args.schema).read_text()) if args.schema else None
70
+ out = Path(args.out)
71
+ out.mkdir(parents=True, exist_ok=True)
72
+
73
+ runner = Runner(get_provider(args.provider, args.model, thinking=args.thinking), args.model, response_schema=schema,
74
+ checkpoint=out / "checkpoint.sqlite", max_attempts=args.max_attempts,
75
+ poll_every=args.poll_every, batch_size=args.batch_size,
76
+ max_open_batches=args.max_open_batches)
77
+ instructions = Path(args.instructions).read_text() if args.instructions else ""
78
+ report = runner.run(to_requests(records, args.max_tokens, instructions))
79
+
80
+ with (out / "results.jsonl").open("w") as f:
81
+ for r in report.results:
82
+ f.write(json.dumps({"key": r.key, "outcome": r.outcome.value, "data": r.data,
83
+ "error": r.error, "attempts": r.attempts, "history": r.history}) + "\n")
84
+
85
+ in_tok = sum(r.input_tokens for r in report.results)
86
+ out_tok = sum(r.output_tokens for r in report.results)
87
+ summary: dict[str, Any] = {
88
+ "provider": args.provider,
89
+ "model": args.model,
90
+ "instructions_file": args.instructions,
91
+ "thinking": args.thinking or "model default",
92
+ "max_tokens": args.max_tokens,
93
+ "records": len(records),
94
+ "ok": report.ok,
95
+ "failed_after_retries": report.failed,
96
+ "outcomes": {o.value: sum(r.outcome is o for r in report.results) for o in Outcome},
97
+ "lines_missing_seen": report.lines_missing,
98
+ "lines_errored_seen": report.lines_errored,
99
+ "lines_invalid_seen": report.lines_invalid,
100
+ "batches_submitted": report.batches_submitted,
101
+ "capacity_refusals": report.capacity_refusals,
102
+ "submit_retries": report.submit_retries,
103
+ "retried_records": sum(r.attempts > 1 for r in report.results),
104
+ "input_tokens": in_tok,
105
+ "output_tokens": out_tok,
106
+ "wall_seconds": round(report.seconds, 1),
107
+ }
108
+ if report.batch_seconds:
109
+ summary["batch_seconds_median"] = round(statistics.median(report.batch_seconds), 1)
110
+ summary["batch_seconds_max"] = round(max(report.batch_seconds), 1)
111
+ if args.prices:
112
+ p = json.loads(Path(args.prices).read_text())[args.provider]
113
+ cost = (in_tok * p["input_per_mtok"] + out_tok * p["output_per_mtok"]) / 1e6 * p.get("batch_multiplier", 1.0)
114
+ summary["cost_usd"] = round(cost, 4)
115
+ summary["cost_per_1000_records_usd"] = round(cost / max(len(records), 1) * 1000, 4)
116
+ acc = accuracy(records, {r.key: r.data for r in report.results})
117
+ if acc:
118
+ summary["field_accuracy"] = acc
119
+ (out / "summary.json").write_text(json.dumps(summary, indent=2))
120
+ print(json.dumps(summary, indent=2))
121
+ return 0
122
+
123
+
124
+ def cmd_compare(args: argparse.Namespace) -> int:
125
+ def load(d: str) -> dict[str, Any]:
126
+ path = Path(d) / "results.jsonl"
127
+ if not path.exists() and path.with_suffix(".jsonl.gz").exists():
128
+ import gzip # large published runs are stored compressed
129
+ text = gzip.decompress(path.with_suffix(".jsonl.gz").read_bytes()).decode("utf-8")
130
+ else:
131
+ text = path.read_text()
132
+ rows = [json.loads(l) for l in text.splitlines() if l]
133
+ return {r["key"]: r["data"] for r in rows if r["outcome"] == "ok"}
134
+
135
+ a, b = load(args.run_a), load(args.run_b)
136
+ both = sorted(set(a) & set(b))
137
+ exact = sum(a[k] == b[k] for k in both)
138
+ fields: dict[str, list[int]] = {}
139
+ for k in both:
140
+ if isinstance(a[k], dict) and isinstance(b[k], dict):
141
+ for f in set(a[k]) | set(b[k]):
142
+ fields.setdefault(f, []).append(int(a[k].get(f) == b[k].get(f)))
143
+ result = {
144
+ "records_in_both": len(both),
145
+ "exact_match_rate": round(exact / len(both), 4) if both else None,
146
+ "field_match_rate": {f: round(sum(v) / len(v), 4) for f, v in sorted(fields.items())},
147
+ }
148
+ print(json.dumps(result, indent=2))
149
+ return 0
150
+
151
+
152
+ REPORT_COLUMNS = [
153
+ ("run", "Run"),
154
+ ("provider", "Provider"),
155
+ ("model", "Model"),
156
+ ("records", "Records"),
157
+ ("ok", "OK"),
158
+ ("failed_after_retries", "Failed"),
159
+ ("lines_missing_seen", "Missing"),
160
+ ("lines_errored_seen", "Errored"),
161
+ ("lines_invalid_seen", "Invalid"),
162
+ ("retried_records", "Retried"),
163
+ ("batch_seconds_median", "Median batch (s)"),
164
+ ("batch_seconds_max", "Slowest batch (s)"),
165
+ ("cost_per_1000_records_usd", "Cost per 1k ($)"),
166
+ ("mean_field_accuracy", "Accuracy"),
167
+ ]
168
+
169
+
170
+ def load_summaries(run_dirs: list[str]) -> list[dict[str, Any]]:
171
+ rows = []
172
+ for d in run_dirs:
173
+ path = Path(d) / "summary.json"
174
+ if not path.exists():
175
+ print(f"skipping {d}: no summary.json", file=sys.stderr)
176
+ continue
177
+ s = json.loads(path.read_text())
178
+ s["run"] = Path(d).name
179
+ acc = s.get("field_accuracy") or {}
180
+ if acc:
181
+ s["mean_field_accuracy"] = round(sum(acc.values()) / len(acc), 4)
182
+ rows.append(s)
183
+ return rows
184
+
185
+
186
+ def write_charts(rows: list[dict[str, Any]], out: Path) -> list[str]:
187
+ try:
188
+ import matplotlib
189
+ matplotlib.use("Agg")
190
+ import matplotlib.pyplot as plt
191
+ except ImportError:
192
+ print("matplotlib not installed; skipping charts (pip install matplotlib)", file=sys.stderr)
193
+ return []
194
+ labels = [r["run"] for r in rows]
195
+ made = []
196
+
197
+ fig, ax = plt.subplots(figsize=(max(6, len(rows) * 1.2), 4))
198
+ bottom = [0] * len(rows)
199
+ for key, name in [("lines_missing_seen", "Missing"), ("lines_errored_seen", "Errored"),
200
+ ("lines_invalid_seen", "Invalid JSON")]:
201
+ vals = [100 * r.get(key, 0) / max(r.get("records", 1), 1) for r in rows]
202
+ ax.bar(labels, vals, bottom=bottom, label=name)
203
+ bottom = [b + v for b, v in zip(bottom, vals)]
204
+ ax.set_ylabel("% of records (before retries)")
205
+ ax.set_title("Problems caught per run")
206
+ ax.legend(loc="upper left", bbox_to_anchor=(1.01, 1), frameon=False)
207
+ fig.tight_layout()
208
+ fig.savefig(out / "problems.png", dpi=150)
209
+ plt.close(fig)
210
+ made.append("problems.png")
211
+
212
+ costs = [(r["run"], r["cost_per_1000_records_usd"]) for r in rows if "cost_per_1000_records_usd" in r]
213
+ if costs:
214
+ fig, ax = plt.subplots(figsize=(max(6, len(costs) * 1.2), 4))
215
+ ax.bar([c[0] for c in costs], [c[1] for c in costs])
216
+ ax.set_ylabel("USD per 1,000 records (incl. retries)")
217
+ ax.set_title("Cost per 1,000 records")
218
+ fig.tight_layout()
219
+ fig.savefig(out / "cost.png", dpi=150)
220
+ plt.close(fig)
221
+ made.append("cost.png")
222
+ return made
223
+
224
+
225
+ def add_costs(rows: list[dict[str, Any]], prices_path: str) -> None:
226
+ """Fill in cost for runs whose summary has token counts but no cost yet."""
227
+ prices = json.loads(Path(prices_path).read_text())
228
+ for r in rows:
229
+ p = prices.get(r.get("provider", ""))
230
+ if not p or "cost_usd" in r:
231
+ continue
232
+ if p.get("model") and p["model"] != r.get("model"):
233
+ continue # prices were recorded for a different model
234
+ cost = (r.get("input_tokens", 0) * p["input_per_mtok"] + r.get("output_tokens", 0) * p["output_per_mtok"]) / 1e6 * p.get("batch_multiplier", 1.0)
235
+ r["cost_usd"] = round(cost, 6)
236
+ r["cost_per_1000_records_usd"] = round(cost / max(r.get("records", 1), 1) * 1000, 4)
237
+
238
+
239
+ def cmd_report(args: argparse.Namespace) -> int:
240
+ rows = load_summaries(args.runs)
241
+ if not rows:
242
+ print("no runs with a summary.json found", file=sys.stderr)
243
+ return 1
244
+ if args.prices:
245
+ add_costs(rows, args.prices)
246
+ out = Path(args.out)
247
+ out.mkdir(parents=True, exist_ok=True)
248
+
249
+ def cell(row: dict[str, Any], key: str) -> str:
250
+ v = row.get(key)
251
+ return "" if v is None else str(v)
252
+
253
+ header = [name for _, name in REPORT_COLUMNS]
254
+ lines = ["| " + " | ".join(header) + " |", "|" + "---|" * len(header)]
255
+ for r in rows:
256
+ lines.append("| " + " | ".join(cell(r, k) for k, _ in REPORT_COLUMNS) + " |")
257
+
258
+ import csv
259
+ with (out / "summary.csv").open("w", newline="") as f:
260
+ w = csv.writer(f)
261
+ w.writerow(header)
262
+ for r in rows:
263
+ w.writerow([cell(r, k) for k, _ in REPORT_COLUMNS])
264
+
265
+ charts = [] if args.no_charts else write_charts(rows, out)
266
+ md = ["# Benchmark results", "",
267
+ "Counts of missing, errored and invalid lines are before retries; "
268
+ "OK and Failed are after retries.", "", *lines, ""]
269
+ md += [f"![{c}]({c})" for c in charts]
270
+ (out / "summary.md").write_text("\n".join(md) + "\n")
271
+ print("\n".join(lines))
272
+ print(f"\nwrote {out / 'summary.md'} and {out / 'summary.csv'}" + (f" and {len(charts)} charts" if charts else ""))
273
+ return 0
274
+
275
+
276
+ def main(argv: list[str] | None = None) -> int:
277
+ parser = argparse.ArgumentParser(prog="steadybatch-bench")
278
+ sub = parser.add_subparsers(dest="cmd", required=True)
279
+
280
+ run = sub.add_parser("run", help="run one workload on one provider")
281
+ run.add_argument("--provider", required=True, choices=["fake", "openai", "anthropic", "vllm", "gemini", "bedrock"])
282
+ run.add_argument("--model", required=True)
283
+ run.add_argument("--data", required=True)
284
+ run.add_argument("--schema")
285
+ run.add_argument("--instructions", help="text file with field definitions, added to the system prompt")
286
+ run.add_argument("--out", required=True)
287
+ run.add_argument("--limit", type=int)
288
+ run.add_argument("--max-tokens", type=int, default=512)
289
+ run.add_argument("--max-attempts", type=int, default=3)
290
+ run.add_argument("--batch-size", type=int, help="at most this many requests per batch")
291
+ run.add_argument("--max-open-batches", type=int,
292
+ help="wait for open batches to finish before submitting more than this many")
293
+ run.add_argument("--thinking", choices=["disabled", "adaptive"],
294
+ help="Anthropic only: set thinking explicitly instead of using the model default")
295
+ run.add_argument("--poll-every", type=float, default=60.0)
296
+ run.add_argument("--prices", help="JSON file with per-provider token prices")
297
+ run.set_defaults(func=cmd_run)
298
+
299
+ cmp_ = sub.add_parser("compare", help="compare two runs of the same workload")
300
+ cmp_.add_argument("run_a")
301
+ cmp_.add_argument("run_b")
302
+ cmp_.set_defaults(func=cmd_compare)
303
+
304
+ rep = sub.add_parser("report", help="turn run summaries into a table, CSV and charts")
305
+ rep.add_argument("runs", nargs="+", help="run folders that contain summary.json")
306
+ rep.add_argument("--out", default="results/latest")
307
+ rep.add_argument("--no-charts", action="store_true")
308
+ rep.add_argument("--prices", help="JSON price file; fills in cost for runs that lack it")
309
+ rep.set_defaults(func=cmd_report)
310
+
311
+ args = parser.parse_args(argv)
312
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s %(message)s")
313
+ return args.func(args)
314
+
315
+
316
+ if __name__ == "__main__":
317
+ sys.exit(main())
@@ -0,0 +1,39 @@
1
+ """Split work into batches that stay under a provider's limits.
2
+
3
+ Providers cap a batch by request count and by file size. Going over the
4
+ size cap does not always produce a clear error, so we stay under both,
5
+ with a safety margin on size.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import Callable, Iterable, Iterator, TypeVar
11
+
12
+ T = TypeVar("T")
13
+
14
+
15
+ def chunk(
16
+ items: Iterable[T],
17
+ max_count: int,
18
+ max_bytes: int,
19
+ size_of: Callable[[T], int],
20
+ margin: float = 0.9,
21
+ ) -> Iterator[list[T]]:
22
+ if max_count < 1 or max_bytes < 1:
23
+ raise ValueError("limits must be positive")
24
+ byte_budget = int(max_bytes * margin)
25
+ batch: list[T] = []
26
+ used = 0
27
+ for item in items:
28
+ size = size_of(item)
29
+ if size > byte_budget:
30
+ raise ValueError(
31
+ f"a single request is {size} bytes, over the per-batch budget of {byte_budget}"
32
+ )
33
+ if batch and (len(batch) >= max_count or used + size > byte_budget):
34
+ yield batch
35
+ batch, used = [], 0
36
+ batch.append(item)
37
+ used += size
38
+ if batch:
39
+ yield batch
steadybatch/ids.py ADDED
@@ -0,0 +1,32 @@
1
+ """Stable request IDs.
2
+
3
+ The same record with the same settings always gets the same ID. That is
4
+ what makes retries safe: if a line comes back twice, or you resubmit work
5
+ after a crash, you can tell it is the same request and not count it twice.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import hashlib
11
+ import json
12
+ from typing import Any
13
+
14
+ from .models import Request
15
+
16
+ # Anthropic allows [a-zA-Z0-9_-]{1,64}; OpenAI is looser. Stay within both.
17
+ _PREFIX = "sb-"
18
+ _HASH_CHARS = 32
19
+
20
+
21
+ def custom_id_for(request: Request, model: str, response_schema: dict[str, Any] | None = None) -> str:
22
+ payload = {
23
+ "key": request.key,
24
+ "model": model,
25
+ "messages": request.messages,
26
+ "system": request.system,
27
+ "max_tokens": request.max_tokens,
28
+ "temperature": request.temperature,
29
+ "schema": response_schema,
30
+ }
31
+ blob = json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
32
+ return _PREFIX + hashlib.sha256(blob.encode("utf-8")).hexdigest()[:_HASH_CHARS]
steadybatch/models.py ADDED
@@ -0,0 +1,85 @@
1
+ """Plain data types shared by the runner and the providers."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass, field
6
+ from enum import Enum
7
+ from typing import Any
8
+
9
+
10
+ @dataclass(frozen=True)
11
+ class Request:
12
+ """One unit of work: a record you want the model to process.
13
+
14
+ ``key`` is your own identifier for the record (a row id, a file name).
15
+ It must be unique within a job; steadybatch uses it to put results
16
+ back in order and to make retries safe.
17
+ """
18
+
19
+ key: str
20
+ messages: list[dict[str, str]]
21
+ system: str | None = None
22
+ max_tokens: int = 1024
23
+ temperature: float = 0.0
24
+
25
+
26
+ @dataclass(frozen=True)
27
+ class PreparedRequest:
28
+ """A Request plus everything a provider needs to send it."""
29
+
30
+ custom_id: str
31
+ request: Request
32
+ model: str
33
+ response_schema: dict[str, Any] | None = None
34
+
35
+
36
+ class BatchState(str, Enum):
37
+ PENDING = "pending" # accepted, not started
38
+ RUNNING = "running"
39
+ DONE = "done" # provider says it finished (results may still be incomplete)
40
+ FAILED = "failed" # the whole batch failed or expired
41
+
42
+
43
+ @dataclass
44
+ class BatchStatus:
45
+ state: BatchState
46
+ total: int = 0
47
+ succeeded: int = 0
48
+ failed: int = 0
49
+ detail: str = ""
50
+ # True when the provider refused the work because the account's queue is full,
51
+ # not because anything is wrong with the requests.
52
+ over_capacity: bool = False
53
+
54
+
55
+ @dataclass
56
+ class RawResult:
57
+ """What a provider returned for one line, before any checking."""
58
+
59
+ custom_id: str
60
+ ok: bool
61
+ text: str | None = None
62
+ error: str | None = None
63
+ input_tokens: int = 0
64
+ output_tokens: int = 0
65
+
66
+
67
+ class Outcome(str, Enum):
68
+ OK = "ok"
69
+ ERROR = "error" # the provider returned an error for this line
70
+ MISSING = "missing" # the provider never returned this line at all
71
+ INVALID = "invalid" # a response came back but failed the schema check
72
+
73
+
74
+ @dataclass
75
+ class Result:
76
+ key: str
77
+ custom_id: str
78
+ outcome: Outcome
79
+ data: Any = None
80
+ text: str | None = None
81
+ error: str | None = None
82
+ attempts: int = 0
83
+ input_tokens: int = 0
84
+ output_tokens: int = 0
85
+ history: list[str] = field(default_factory=list)
@@ -0,0 +1,35 @@
1
+ """Batch providers. Each one is imported lazily so you only need the SDKs you use."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from .base import BatchProvider
6
+ from .fake import FakeProvider, demo_answer
7
+
8
+
9
+ def get_provider(name: str, model: str | None = None, *, thinking: str | None = None) -> BatchProvider:
10
+ if name == "fake":
11
+ # Misbehaves like a real API: a few lines dropped, errored or malformed.
12
+ return FakeProvider(answer=demo_answer, drop_rate=0.01, error_rate=0.01, invalid_rate=0.01,
13
+ fail_first_attempt_only=True)
14
+ if name == "openai":
15
+ from .openai_batch import OpenAIBatch
16
+ return OpenAIBatch()
17
+ if name == "anthropic":
18
+ from .anthropic_batch import AnthropicBatch
19
+ return AnthropicBatch(thinking=thinking)
20
+ if name == "vllm":
21
+ from .vllm_offline import VLLMOffline
22
+ if not model:
23
+ raise ValueError("vllm needs a model name")
24
+ return VLLMOffline(model)
25
+ if name == "gemini":
26
+ from .gemini_batch import GeminiBatch
27
+ return GeminiBatch()
28
+ if name == "bedrock":
29
+ raise NotImplementedError(
30
+ "the Bedrock adapter is next on the list; see docs/ROADMAP.md. Contributions welcome."
31
+ )
32
+ raise ValueError(f"unknown provider: {name}")
33
+
34
+
35
+ __all__ = ["BatchProvider", "FakeProvider", "get_provider"]
@@ -0,0 +1,95 @@
1
+ """Anthropic Message Batches API.
2
+
3
+ Flow: create a batch with up to 100,000 requests in one call, poll until
4
+ processing_status is "ended", then stream results. Each result is
5
+ succeeded, errored, canceled or expired.
6
+
7
+ When a JSON Schema is given, it is sent as a structured output format
8
+ (`output_config.format`), so the reply is constrained to the schema, the
9
+ same as the OpenAI and Gemini adapters. Pass `native_schema=False` to
10
+ describe the schema in the system prompt instead. `thinking` ("disabled" or
11
+ "adaptive") is sent when given; left out, the model uses its default, and on
12
+ Claude Haiku 5.5 that default is adaptive thinking, which counts against
13
+ max_tokens. `temperature` is not sent:
14
+ current Claude models reject non-default sampling settings, and anthropic
15
+ SDK 1.x no longer takes the parameter.
16
+ """
17
+
18
+ from __future__ import annotations
19
+
20
+ import json
21
+ from typing import Iterator
22
+
23
+ from ..models import BatchState, BatchStatus, PreparedRequest, RawResult
24
+ from .base import BatchProvider
25
+
26
+
27
+ class AnthropicBatch(BatchProvider):
28
+ name = "anthropic"
29
+ max_requests_per_batch = 100_000
30
+ max_bytes_per_batch = 256 * 1024 * 1024
31
+
32
+ def __init__(self, client=None, *, native_schema: bool = True, thinking: str | None = None):
33
+ if client is None:
34
+ import anthropic # imported here so the package works without it
35
+ client = anthropic.Anthropic()
36
+ self.client = client
37
+ self.native_schema = native_schema
38
+ self.thinking = thinking
39
+
40
+ def to_line(self, req: PreparedRequest) -> dict:
41
+ params = {
42
+ "model": req.model,
43
+ "max_tokens": req.request.max_tokens,
44
+ "messages": req.request.messages,
45
+ }
46
+ if self.thinking:
47
+ params["thinking"] = {"type": self.thinking}
48
+ system = req.request.system or ""
49
+ if req.response_schema is not None and self.native_schema:
50
+ params["output_config"] = {"format": {"type": "json_schema", "schema": req.response_schema}}
51
+ elif req.response_schema is not None:
52
+ # Ask for JSON in the prompt; the runner checks every answer against the schema.
53
+ system = (system + "\n\nReply with only a JSON object that matches this JSON Schema:\n"
54
+ + json.dumps(req.response_schema)).strip()
55
+ if system:
56
+ params["system"] = system
57
+ return {"custom_id": req.custom_id, "params": params}
58
+
59
+ def submit(self, batch: list[PreparedRequest]) -> str:
60
+ created = self.client.messages.batches.create(requests=[self.to_line(r) for r in batch])
61
+ return created.id
62
+
63
+ def is_over_capacity(self, exc: Exception) -> bool:
64
+ return getattr(exc, "status_code", None) == 429
65
+
66
+ def status(self, batch_id: str) -> BatchStatus:
67
+ b = self.client.messages.batches.retrieve(batch_id)
68
+ c = b.request_counts
69
+ if b.processing_status == "ended":
70
+ state = BatchState.DONE
71
+ elif b.processing_status == "canceling":
72
+ state = BatchState.FAILED
73
+ else:
74
+ state = BatchState.RUNNING
75
+ total = c.processing + c.succeeded + c.errored + c.canceled + c.expired
76
+ return BatchStatus(state, total=total, succeeded=c.succeeded,
77
+ failed=c.errored + c.canceled + c.expired, detail=b.processing_status)
78
+
79
+ def results(self, batch_id: str) -> Iterator[RawResult]:
80
+ for entry in self.client.messages.batches.results(batch_id):
81
+ result = entry.result
82
+ if result.type == "succeeded":
83
+ msg = result.message
84
+ if getattr(msg, "stop_reason", None) == "refusal":
85
+ yield RawResult(entry.custom_id, ok=False, error="refusal")
86
+ continue
87
+ text = "".join(block.text for block in msg.content if block.type == "text")
88
+ truncated = getattr(msg, "stop_reason", None) == "max_tokens"
89
+ yield RawResult(entry.custom_id, ok=True, text=text,
90
+ error="truncated (max_tokens)" if truncated else None,
91
+ input_tokens=msg.usage.input_tokens, output_tokens=msg.usage.output_tokens)
92
+ elif result.type == "errored":
93
+ yield RawResult(entry.custom_id, ok=False, error=str(getattr(result, "error", "errored")))
94
+ else:
95
+ yield RawResult(entry.custom_id, ok=False, error=result.type)
@@ -0,0 +1,60 @@
1
+ """The interface every batch provider implements.
2
+
3
+ A provider only has to do four things: say what its limits are, submit a
4
+ list of requests, report a batch's status, and hand back whatever lines it
5
+ has. Everything else (retries, checks, ordering, checkpoints) lives in the
6
+ runner, so every provider gets the same treatment and the comparison is fair.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ from abc import ABC, abstractmethod
13
+ from typing import Iterator
14
+
15
+ from ..models import BatchStatus, PreparedRequest, RawResult
16
+
17
+
18
+ class BatchProvider(ABC):
19
+ name: str = "base"
20
+ max_requests_per_batch: int = 10_000
21
+ max_bytes_per_batch: int = 100 * 1024 * 1024
22
+
23
+ def size_of(self, req: PreparedRequest) -> int:
24
+ """Rough size of one request line in the upload file."""
25
+ return len(json.dumps(self.to_line(req)).encode("utf-8")) + 1
26
+
27
+ def to_line(self, req: PreparedRequest) -> dict:
28
+ """The provider-specific JSON for one request. Used for sizing and upload."""
29
+ return {"custom_id": req.custom_id, "messages": req.request.messages}
30
+
31
+ def is_over_capacity(self, exc: Exception) -> bool:
32
+ """True if `submit` failed because the account's queue or quota is full.
33
+
34
+ The runner then sends smaller batches instead of stopping. Providers that
35
+ can tell this apart from other errors override it.
36
+ """
37
+ return False
38
+
39
+ def is_transient(self, exc: Exception) -> bool:
40
+ """True if `submit` failed for a reason worth retrying as-is: a dropped
41
+ connection, a timeout or a server error. The SDKs retry a couple of times
42
+ on their own; a long job needs more patience than that."""
43
+ names = {c.__name__ for c in type(exc).__mro__}
44
+ if names & {"APIConnectionError", "APITimeoutError", "ConnectionError", "TimeoutError",
45
+ "ServerError", "InternalServerError", "ServiceUnavailableError"}:
46
+ return True
47
+ status = getattr(exc, "status_code", None) or getattr(exc, "code", None)
48
+ return isinstance(status, int) and status >= 500
49
+
50
+ @abstractmethod
51
+ def submit(self, batch: list[PreparedRequest]) -> str:
52
+ """Send a batch. Return the provider's batch ID."""
53
+
54
+ @abstractmethod
55
+ def status(self, batch_id: str) -> BatchStatus:
56
+ """Report where a batch is."""
57
+
58
+ @abstractmethod
59
+ def results(self, batch_id: str) -> Iterator[RawResult]:
60
+ """Yield every line the provider returned, successes and errors alike."""