From 8eab028fb9248e97e31cfb024129b03ee1e70918 Mon Sep 17 00:00:00 2001 From: drewOrc <36374426+drewOrc@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:02:46 +0800 Subject: [PATCH 1/4] Add Haiku 8-way runner with journal, resume and cost cap make llm calls claude-haiku-4-5-20251001 on validation 3,100 + test 5,500 with the prompt of cost-aware-hybrid-router (byte-identical, SHA-256 pinned in tests) and stores one record per query: split, index, query SHA-256, gold label, raw reply, parsed agent, parse_failed, tokens, cost, latency, attempts and request id. - Journal per identity (model, prompt, temperature, max_tokens): a rerun calls only rows without a stored reply; a changed identity reuses nothing. - Cost cap (default US$5): each call reserves an upper bound before it starts, cumulative across reruns; actual cost from usage. - Retries in code, not the SDK: 429, 5xx, 408, 409 and connection errors back off; others fail at once. Failures are journaled and exit 1. - Completion: the predictions file holds every row once and its SHA-256 matches the summary and results/llm-manifest.json before "completed 8600/8600 llm predictions" is printed. - make llm-smoke: validation rows 0-19 into results/llm-smoke/, with the extrapolated full cost. - The API key is redacted from logs, journal and exceptions. - CI installs the llm group so tests use the SDK's exception classes; a fake client stands in for the API. --- .github/workflows/ci.yml | 5 +- .gitignore | 5 + DEVLOG.md | 33 +++ Makefile | 22 +- README.md | 6 +- docs/OPERATIONS.md | 2 +- pyproject.toml | 6 +- src/tinyrouter/llm.py | 234 +++++++++++++-- src/tinyrouter/llm_run.py | 579 ++++++++++++++++++++++++++++++++++++++ tests/llm_fakes.py | 95 +++++++ tests/test_llm.py | 126 +++++++-- tests/test_llm_deps.py | 22 ++ tests/test_llm_run.py | 332 ++++++++++++++++++++++ tests/test_makefile.py | 10 + 14 files changed, 1417 insertions(+), 60 deletions(-) create mode 100644 src/tinyrouter/llm_run.py create mode 100644 tests/llm_fakes.py create mode 100644 tests/test_llm_deps.py create mode 100644 tests/test_llm_run.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 993693f..f9fa560 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -61,7 +61,10 @@ jobs: enable-cache: true - name: Install dependencies (fails if uv.lock is out of date) - run: uv sync --locked + # The llm group (anthropic SDK) is installed so the Haiku runner's + # retry and key-redaction tests use the SDK's real exception classes. + # They run against a fake client; no key is set and no API is called. + run: uv sync --locked --group llm - name: make lint (ruff check, ruff format --check, em dash check) run: make lint diff --git a/.gitignore b/.gitignore index 3df792d..325ab51 100644 --- a/.gitignore +++ b/.gitignore @@ -32,6 +32,11 @@ checkpoints/ # they are attached to a GitHub Release and checked with `make verify-logits`. results/summary.md results/logits/ +# Haiku replies: the journal and the predictions file (about 3 MB) go to the +# GitHub Release like the logits; results/llm-manifest.json and +# results/llm/haiku-8way.json are committed. The smoke run is scratch. +results/llm/*.jsonl +results/llm-smoke/ *.npz *.tmp diff --git a/DEVLOG.md b/DEVLOG.md index bb6e518..52dfb6b 100644 --- a/DEVLOG.md +++ b/DEVLOG.md @@ -4,6 +4,39 @@ --- +## 2026-09-29:步驟 4 之一,Haiku 8 類執行器(尚未實跑) + +### 本次工作 / 執行摘要 +- 新增 `src/tinyrouter/llm_run.py` 與 `make llm`、`make llm-smoke`、`make verify-llm`:validation 3,100 + test 5,500 逐筆呼叫 Haiku,逐筆寫入 journal(split、index、query 的 SHA-256、gold intent 與 agent、原始回覆、解析後標籤、`parse_failed`、tokens、花費、延遲、嘗試次數、request id)。 +- query 原文不存,只存 SHA-256:資料集公開且鎖定 revision,雜湊足以證明紀錄對應哪一列,續跑時也拿它比對資料有沒有變。 +- 續跑:journal 檔名含身分雜湊(模型、system prompt 的 SHA-256、temperature、max_tokens、user 內容格式)。身分一變就換新檔,舊回覆不會被拿來用;同身分重跑只補沒有成功紀錄的列。最後一行寫到一半(當機)會被截掉重做。 +- 成本上限:開跑前用 token counting 取 prompt 的基礎 token 數,印出上界估計;每筆開打前預留「基礎 + 每個 byte 算一個 token 的輸入、max_tokens 的輸出」的上界,累計花費(含之前幾次)加上在途預留會超過 `--max-usd` 就不開新呼叫。實際花費依回傳的 usage 計算。 +- 重試改由程式自己做(SDK 的 `max_retries=0`):429、5xx、408、409、連線錯誤指數退避並尊重 retry-after,5 次用盡記為失敗、整體 exit 1;400、401、403、404 不重試,而且停止開新呼叫。 +- 完成判定比照 `completeness.py`:預測檔從磁碟讀回,(split, index) 恰為預期集合且各一次、身分一致,SHA-256 在檔案、summary、`results/llm-manifest.json` 三處相同,才印 `completed 8600/8600 llm predictions`。預期列數寫成字面值。 +- CI 的 test job 改裝 `--group llm`,讓重試與 key 遮蔽的測試用 SDK 真正的例外類別。仍不設 key、不打 API。 +- `classify` 原本把任何例外都當可重試,改為依錯誤類型判斷。 + +### 核心發現 / 數據 +- system prompt 與 cost-aware-hybrid-router `src/routers/llm_router.py` 逐位元組相同(SHA-256 `560d22c5...5df574`,測試釘住)。模型、temperature 0、max_tokens 20、query 原樣當唯一 user 訊息,都與舊專案相同。 +- 解析規則與舊專案**不同**:舊版對回覆做子字串比對、取集合迭代到的第一個命中(順序不固定);這裡只接受完整標籤,或恰好命中一個 in-scope agent,其餘判為 oos 並標 `parse_failed`。原始回覆都有存,要用舊規則重算不必再打 API。 +- 定價:Haiku 4.5 每百萬 tokens 輸入 US$1、輸出 US$5(claude-api skill 的模型表,快取日期 2026-06-24)。粗估全量約 US$2.5,上界約 US$3.5,低於 AC6 的 US$5。 +- (無實跑數據) + +### Blockers / 遇到的問題 +- (無) + +### Next +- [ ] Drew:`.env` 放 key 後 `make llm-smoke`,看 20 筆的實際 token 與推估全量花費 +- [ ] `make llm`,把 `results/llm/haiku-8way.jsonl` 附到 Release +- [ ] 下一個 PR:不確定性、risk-coverage、fallback、oracle(RQ3、RQ4、AC6) + +### Files / Budget +- 新增:`src/tinyrouter/llm_run.py`、`tests/test_llm_run.py`、`tests/test_llm_deps.py`、`tests/llm_fakes.py` +- 修改:`src/tinyrouter/llm.py`、`tests/test_llm.py`、`tests/test_makefile.py`、`Makefile`、`.github/workflows/ci.yml`、`.gitignore`、`pyproject.toml`(只改註解)、`README.md`、`docs/OPERATIONS.md` +- API 花費:US$0 + +--- + ## 2026-09-23(夜):PR #7 審查修正 ### 本次工作 / 執行摘要 diff --git a/Makefile b/Makefile index b17fe40..03bec17 100644 --- a/Makefile +++ b/Makefile @@ -1,5 +1,5 @@ .PHONY: setup lint format test test-network smoke train evaluate ac2 pilot-lr pilot-steps baselines \ - curve oos-ablation verify-logits report clean-checkpoints + curve oos-ablation verify-logits llm-smoke llm verify-llm report clean-checkpoints CONFIG ?= configs/bert-base.yaml SEED ?= 42 @@ -112,6 +112,26 @@ oos-ablation: verify-logits: uv run python -m tinyrouter.archive +# Claude Haiku over CLINC150 in the 8-way routing space (docs/PLAN.md section 4, +# AC6). Needs ANTHROPIC_API_KEY (in .env or exported); without it they exit 2. +# Every call is journaled; a rerun calls only rows with no stored reply, and a +# change of model, prompt, temperature or max_tokens starts a fresh journal. +# A call is not started if it could take this run's identity past MAX_USD +# (default 5). Exit 1 when stopped by the cap or when a call failed after +# retries. `make llm-smoke` calls validation rows 0-19, writes under +# results/llm-smoke/ and prints the extrapolated cost of all 8,600 rows. +# Done means the whole last line `completed 8600/8600 llm predictions`. +llm-smoke: + uv run --group llm $(UV_ENV) python -m tinyrouter.llm_run --smoke $(if $(MAX_USD),--max-usd $(MAX_USD),) + +llm: + uv run --group llm $(UV_ENV) python -m tinyrouter.llm_run $(if $(MAX_USD),--max-usd $(MAX_USD),) + +# results/llm/haiku-8way.jsonl has every row once and matches its SHA-256 in +# results/llm-manifest.json and results/llm/haiku-8way.json. +verify-llm: + uv run python -m tinyrouter.llm_run --verify + report: uv run python -m tinyrouter.report diff --git a/README.md b/README.md index cbd5296..63b3237 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,9 @@ Other targets: | `make baselines` | majority-class and TF-IDF centroid baselines on every (k, seed) sample, archived like an encoder run | | `make curve MODEL=bert\|modernbert` | baselines, then 6 values of k x 3 seeds with the lr and S_min from `configs/curve.yaml`; writes `results/curves/.json`; keeps no weights | | `make oos-ablation` | ModernBERT, k=100, no OOS training rows, 3 seeds; writes `results/curves/oos-ablation.json` | +| `make llm-smoke` | Claude Haiku on validation rows 0-19 (needs `ANTHROPIC_API_KEY`); writes `results/llm-smoke/`, prints tokens, cost and the extrapolated cost of all 8,600 rows | +| `make llm` | Claude Haiku on validation 3,100 + test 5,500 in the 8-way space, one stored record per query; resumes; stops before passing `MAX_USD` (default 5); writes `results/llm/haiku-8way.{jsonl,json}` and `results/llm-manifest.json` | +| `make verify-llm` | check `results/llm/haiku-8way.jsonl`: every row once, SHA-256 equal in the file, the summary and the manifest | | `make verify-logits` | check every archive in `results/logits/` against `results/logits-manifest.json` | | `make report` | build `results/summary.md` from `results/runs/*.json` | | `make clean-checkpoints` | delete all trained weights | @@ -83,7 +86,8 @@ src/tinyrouter/ curves.py learning curves and the OOS ablation (reuses AC2 at k=100 when equivalent) baselines.py majority-class and TF-IDF centroid baselines per curve point ac2.py AC2 run over three seeds and PASS/FAIL verdict - llm.py Claude Haiku zero-shot router (baseline and fallback; not run yet) + llm.py Claude Haiku zero-shot router: prompt, retries, pricing (baseline and fallback) + llm_run.py Haiku over validation + test: journal, resume, cost cap, completion check report.py results/*.json -> results/summary.md smoke.py end-to-end wiring check tests/ pytest; `network` marker for Hub downloads diff --git a/docs/OPERATIONS.md b/docs/OPERATIONS.md index 45733d9..0512ec0 100644 --- a/docs/OPERATIONS.md +++ b/docs/OPERATIONS.md @@ -8,7 +8,7 @@ TinyRouter is a research repository with no deployment target: nothing runs as a | job | runs | required to merge | |---|---|---| -| `test` | tracked-files guard, `uv sync --locked`, `make lint`, `make test` (offline) | yes | +| `test` | tracked-files guard, `uv sync --locked --group llm` (the SDK's exception classes for the Haiku runner tests; no key, no API call), `make lint`, `make test` (offline) | yes | | `commit-hygiene` | rejects tool-attribution trailers in commit messages (patterns in `.github/disallowed-trailers.txt`) | yes | | `network` | `make test-network` and `make smoke` against the Hugging Face Hub | no | | `pr-text-hygiene` (`pr-text.yml`) | rejects the same patterns in the PR title and body, on open, edit, push and reopen | yes (added to the required checks once it is on `main`) | diff --git a/pyproject.toml b/pyproject.toml index c457f2c..5210691 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -28,8 +28,10 @@ packages = ["src/tinyrouter"] [dependency-groups] # The Claude Haiku baseline and cascade fallback. Not installed by -# `uv sync`; `uv sync --group llm` installs it. Nothing in the unit tests -# imports it (tinyrouter/llm.py defers the import to call time). +# `uv sync`; `uv sync --group llm` installs it, and `make llm` / `make +# llm-smoke` ask uv for it. tinyrouter/llm.py imports it only at call time. +# CI installs it: the retry and redaction tests raise the SDK's real +# exception classes at a fake client (they skip where it is missing). llm = [ "anthropic==1.8.0", ] diff --git a/src/tinyrouter/llm.py b/src/tinyrouter/llm.py index 114b28a..404acc1 100644 --- a/src/tinyrouter/llm.py +++ b/src/tinyrouter/llm.py @@ -1,15 +1,26 @@ """Claude Haiku zero-shot router: the LLM baseline and cascade fallback. Prompt and model carried over from cost-aware-hybrid-router -(src/routers/llm_router.py) so the two projects score the same LLM. Nothing -in this skeleton calls the API; ``classify`` needs the ``llm`` dependency -group and ANTHROPIC_API_KEY, and the unit tests exercise it with a fake -client only. +(src/routers/llm_router.py) so the two projects score the same LLM: the +system prompt is byte-identical (tests/test_llm.py pins its SHA-256), the +query is sent verbatim as the only user message, temperature 0, +max_tokens 20. The ``anthropic`` package (``llm`` dependency group) is +imported only when a client is built or an API error is classified. + +Retries are done here, not by the SDK (the client is built with +``max_retries=0``), so the attempt count and the delays are recorded and +testable. 429, 5xx (529 overloaded included), 408, 409 and connection +errors or timeouts are retried with exponential backoff, honouring +``retry-after``; anything else (400, 401, 403, 404, ...) fails at once. """ from __future__ import annotations +import hashlib +import json +import os import time +from collections.abc import Callable from dataclasses import dataclass from typing import Any @@ -17,7 +28,27 @@ MODEL = "claude-haiku-4-5-20251001" MAX_TOKENS = 20 +TEMPERATURE = 0 MAX_ATTEMPTS = 5 +MAX_BACKOFF_S = 60.0 +RETRYABLE_STATUS = frozenset({408, 409, 429}) +KEY_ENV = "ANTHROPIC_API_KEY" + +# USD per million tokens for Claude Haiku 4.5, from the claude-api skill's +# "Current Models" table (cached 2026-06-24; first-party API rates), which +# points to https://platform.claude.com/docs/en/about-claude/pricing.md. +# Cache multipliers (write 1.25x, read 0.1x of input) are from the same +# skill; this prompt is not cached, so those two should stay at zero. +PRICE_USD_PER_MTOK: dict[str, float] = { + "input": 1.00, + "output": 5.00, + "cache_write": 1.25, + "cache_read": 0.10, +} +PRICING_SOURCE = ( + "claude-api skill, Current Models table (cached 2026-06-24): Claude Haiku 4.5 " + "$1.00 input / $5.00 output per MTok; https://platform.claude.com/docs/en/about-claude/pricing.md" +) AGENT_DESCRIPTIONS: dict[str, str] = { "finance_agent": "Banking, credit cards, accounts, bills, taxes, insurance, rewards, " @@ -49,6 +80,20 @@ ) +class MissingKeyError(RuntimeError): + """ANTHROPIC_API_KEY is not set.""" + + +class LLMCallError(RuntimeError): + """A call failed for good: a non-retryable error, or retries used up.""" + + def __init__(self, detail: str, *, attempts: int, retryable: bool) -> None: + super().__init__(detail) + self.detail = detail + self.attempts = attempts + self.retryable = retryable + + @dataclass(frozen=True) class LLMPrediction: agent: str @@ -57,6 +102,61 @@ class LLMPrediction: input_tokens: int output_tokens: int latency_ms: int + cache_creation_input_tokens: int = 0 + cache_read_input_tokens: int = 0 + request_id: str | None = None + stop_reason: str | None = None + attempts: int = 1 + + @property + def cost_usd(self) -> float: + return cost_usd( + self.input_tokens, + self.output_tokens, + self.cache_creation_input_tokens, + self.cache_read_input_tokens, + ) + + +def cost_usd( + input_tokens: int, output_tokens: int, cache_write: int = 0, cache_read: int = 0 +) -> float: + """Dollar cost of one call's usage at ``PRICE_USD_PER_MTOK``.""" + p = PRICE_USD_PER_MTOK + total = ( + input_tokens * p["input"] + + output_tokens * p["output"] + + cache_write * p["cache_write"] + + cache_read * p["cache_read"] + ) + return total / 1_000_000 + + +def request_params(query: str) -> dict[str, Any]: + """The exact Messages API arguments for one query; the single place they are built.""" + return { + "model": MODEL, + "max_tokens": MAX_TOKENS, + "temperature": TEMPERATURE, + "system": SYSTEM_PROMPT, + "messages": [{"role": "user", "content": query}], + } + + +def identity() -> dict[str, object]: + """Everything that decides what the model is asked; a change invalidates stored replies.""" + return { + "model": MODEL, + "max_tokens": MAX_TOKENS, + "temperature": TEMPERATURE, + "system_prompt_sha256": hashlib.sha256(SYSTEM_PROMPT.encode("utf-8")).hexdigest(), + "user_content": "query text verbatim, one user message", + } + + +def identity_sha256() -> str: + canonical = json.dumps(identity(), sort_keys=True, separators=(",", ":")) + return hashlib.sha256(canonical.encode("utf-8")).hexdigest() def parse_agent(raw_text: str) -> tuple[str, bool]: @@ -70,38 +170,112 @@ def parse_agent(raw_text: str) -> tuple[str, bool]: return OOS, False -def classify(query: str, client: Any, sleep: Any = time.sleep) -> LLMPrediction: - """One zero-shot call with exponential backoff on API errors.""" - for attempt in range(MAX_ATTEMPTS): +def redact(text: str) -> str: + """``text`` with the API key's value, if set and present, replaced.""" + key = os.environ.get(KEY_ENV, "").strip() + if len(key) >= 8: + text = text.replace(key, "[redacted]") + return text + + +def describe_error(exc: BaseException) -> str: + """Class, status, request id and message of an API error, with the key redacted.""" + parts = [type(exc).__name__] + for attr in ("status_code", "request_id"): + value = getattr(exc, attr, None) + if value is not None: + parts.append(f"{attr}={value}") + return redact(f"{' '.join(parts)}: {exc}")[:500] + + +def is_retryable(exc: BaseException) -> bool: + """Rate limits, server errors, 408/409 and network failures; see the module docstring.""" + import anthropic + + if isinstance(exc, anthropic.APIConnectionError): + return True + if isinstance(exc, anthropic.APIStatusError): + return exc.status_code in RETRYABLE_STATUS or exc.status_code >= 500 + return False + + +def retry_delay(exc: BaseException, attempt: int) -> float: + """Seconds to wait after failed attempt ``attempt`` (1-based): 1, 2, 4, ... or retry-after.""" + delay = float(2 ** (attempt - 1)) + response = getattr(exc, "response", None) + header = getattr(response, "headers", {}).get("retry-after") if response is not None else None + try: + delay = max(delay, float(header)) if header is not None else delay + except ValueError: + pass + return min(delay, MAX_BACKOFF_S) + + +def reply_text(response: Any) -> str: + """Concatenated text blocks; an empty reply (no text block) is the empty string.""" + return "".join( + block.text for block in response.content if getattr(block, "type", "text") == "text" + ) + + +def classify( + query: str, + client: Any, + sleep: Callable[[float], None] = time.sleep, + retryable: Callable[[BaseException], bool] = is_retryable, +) -> LLMPrediction: + """One zero-shot call, retried per the module docstring; LLMCallError when it gives up.""" + for attempt in range(1, MAX_ATTEMPTS + 1): started = time.monotonic() try: - response = client.messages.create( - model=MODEL, - max_tokens=MAX_TOKENS, - temperature=0, - system=SYSTEM_PROMPT, - messages=[{"role": "user", "content": query}], - ) - except Exception as exc: # the SDK raises several error types; all are retryable here - if attempt == MAX_ATTEMPTS - 1: - raise RuntimeError(f"LLM call failed after {MAX_ATTEMPTS} attempts") from exc - sleep(2**attempt) + response = client.messages.create(**request_params(query)) + except Exception as exc: # classified below; anything unknown is not retried + can_retry = retryable(exc) + if not can_retry or attempt == MAX_ATTEMPTS: + raise LLMCallError( + describe_error(exc), attempts=attempt, retryable=can_retry + ) from None + sleep(retry_delay(exc, attempt)) continue - raw_text = response.content[0].text - agent, parsed = parse_agent(raw_text) - return LLMPrediction( - agent=agent, - raw_text=raw_text, - parsed=parsed, - input_tokens=response.usage.input_tokens, - output_tokens=response.usage.output_tokens, - latency_ms=round((time.monotonic() - started) * 1000), - ) + latency_ms = round((time.monotonic() - started) * 1000) + return prediction_from(response, latency_ms, attempt) raise AssertionError("unreachable") +def prediction_from(response: Any, latency_ms: int, attempts: int) -> LLMPrediction: + raw_text = reply_text(response) + agent, parsed = parse_agent(raw_text) + usage = response.usage + return LLMPrediction( + agent=agent, + raw_text=raw_text, + parsed=parsed, + input_tokens=int(usage.input_tokens), + output_tokens=int(usage.output_tokens), + latency_ms=latency_ms, + cache_creation_input_tokens=int(getattr(usage, "cache_creation_input_tokens", 0) or 0), + cache_read_input_tokens=int(getattr(usage, "cache_read_input_tokens", 0) or 0), + request_id=getattr(response, "_request_id", None), + stop_reason=getattr(response, "stop_reason", None), + attempts=attempts, + ) + + +def prompt_base_tokens(client: Any) -> int: + """Input tokens of the system prompt plus a one-character query (token counting endpoint).""" + counted = client.messages.count_tokens( + model=MODEL, system=SYSTEM_PROMPT, messages=[{"role": "user", "content": "x"}] + ) + return int(counted.input_tokens) + + def make_client() -> Any: - """Anthropic client from ANTHROPIC_API_KEY; requires `uv sync --group llm`.""" + """Anthropic client from ANTHROPIC_API_KEY with SDK retries off; needs the ``llm`` group.""" + if not os.environ.get(KEY_ENV, "").strip(): + raise MissingKeyError( + f"{KEY_ENV} is not set. Put it in a .env file at the repository root " + "(gitignored; the Makefile passes it to uv) or export it, then rerun." + ) import anthropic - return anthropic.Anthropic() + return anthropic.Anthropic(max_retries=0) diff --git a/src/tinyrouter/llm_run.py b/src/tinyrouter/llm_run.py new file mode 100644 index 0000000..494eda1 --- /dev/null +++ b/src/tinyrouter/llm_run.py @@ -0,0 +1,579 @@ +"""Haiku over validation 3,100 + test 5,500, one stored record per query (PLAN §4, AC6). + +The old project kept only aggregate counts, so a cascade could not be +built from it. Here every reply is kept: split, row index, SHA-256 of the +query text (the text itself is not stored; CLINC150 is public and pinned, +and the hash proves which row a record belongs to), gold intent and +agent, the raw reply, the parsed agent, ``parse_failed``, tokens, cost, +latency, attempts and the request id. + +Files, for target ``haiku-8way`` under ``results/llm/``: + +- ``haiku-8way..journal.jsonl``: append-only, one line per finished + call (and per call that failed for good, so it is on record; those are + retried next run). ```` is the start of ``llm.identity_sha256()``, + so a different model, prompt, temperature or max_tokens starts a new + journal and never reuses these replies. Rerunning after an interruption + calls the API only for rows with no successful record. +- ``haiku-8way.jsonl``: the predictions, written once every row has a + record, sorted by split and index. Not committed (like the logits + archives); its SHA-256 is in ``results/llm-manifest.json``. +- ``haiku-8way.json``: summary with identity, pricing, tokens and dollars. + +Cost cap: before each call the runner reserves an upper bound on that +call's cost (input tokens counted by the API for the prompt plus one byte +per token of query, output at ``max_tokens``) and does not start a call +that could take cumulative spend for this identity past ``--max-usd``. +Actual cost comes from each response's ``usage``. + +Completion: the predictions file, read back from disk, must hold exactly +the expected (split, index) pairs once each, all with this identity, and +its SHA-256 must agree on disk, in the manifest and in the summary. Only +then is ``completed 8600/8600 llm predictions`` printed. The expected +row counts are literals here, not read from ``data.SPLIT_FILES``. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import sys +import time +from collections import Counter +from collections.abc import Callable, Iterable +from concurrent.futures import FIRST_COMPLETED, Future, ThreadPoolExecutor, wait +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, TextIO + +from tinyrouter import llm +from tinyrouter.archive import git_state, read_manifest, utc_now, write_manifest +from tinyrouter.data import DATASET_REVISION, Split, SplitName, load_split, sha256_of +from tinyrouter.labels import AGENTS, load_label_space + +EXPECTED_ROWS: dict[str, int] = {"validation": 3100, "test": 5500} +# What the smoke run's cost is extrapolated to; equals the sum of EXPECTED_ROWS (tested). +FULL_QUERIES = 8600 +SMOKE_ROWS: dict[str, int] = {"validation": 20} +SPLIT_ORDER = ("validation", "test") +DEFAULT_MAX_USD = 5.0 +DEFAULT_WORKERS = 6 +MAX_WORKERS = 8 +MANIFEST_NAME = "llm-manifest.json" +RECORD_FIELDS: dict[str, type | tuple[type, ...]] = { + "split": str, + "index": int, + "query_sha256": str, + "gold_intent": int, + "gold_agent": str, + "raw_text": str, + "agent": str, + "parse_failed": bool, + "input_tokens": int, + "output_tokens": int, + "cache_creation_input_tokens": int, + "cache_read_input_tokens": int, + "cost_usd": float, + "latency_ms": int, + "attempts": int, + "request_id": (str, type(None)), + "stop_reason": (str, type(None)), + "identity_sha256": str, +} + +Key = tuple[str, int] +Log = Callable[[str], None] + + +class JournalError(RuntimeError): + """A stored record contradicts the data or the identity it claims.""" + + +class IncompleteError(RuntimeError): + """The predictions file is not exactly the expected rows, once each.""" + + +@dataclass(frozen=True) +class Target: + name: str + out_dir: Path + rows: dict[str, int] + manifest: Path | None + + @property + def expected_total(self) -> int: + return sum(self.rows.values()) + + def journal(self, identity_sha: str) -> Path: + return self.out_dir / f"{self.name}.{identity_sha[:12]}.journal.jsonl" + + @property + def predictions(self) -> Path: + return self.out_dir / f"{self.name}.jsonl" + + @property + def summary(self) -> Path: + return self.out_dir / f"{self.name}.json" + + +def full_target(results_root: Path) -> Target: + return Target("haiku-8way", results_root / "llm", EXPECTED_ROWS, results_root / MANIFEST_NAME) + + +def smoke_target(results_root: Path) -> Target: + return Target("haiku-8way-smoke", results_root / "llm-smoke", SMOKE_ROWS, None) + + +@dataclass(frozen=True) +class Query: + split: str + index: int + text: str + gold_intent: int + + @property + def key(self) -> Key: + return (self.split, self.index) + + @property + def text_sha256(self) -> str: + return hashlib.sha256(self.text.encode("utf-8")).hexdigest() + + +@dataclass +class Outcome: + records: dict[Key, dict[str, Any]] + spent_before: float + spent_now: float = 0.0 + calls_now: int = 0 + failures: list[dict[str, Any]] = field(default_factory=list) + stopped: str | None = None + + +def build_queries(splits: dict[str, Split], rows: dict[str, int]) -> list[Query]: + """The first ``rows[split]`` rows of each split, in split then index order.""" + queries = [] + for name in SPLIT_ORDER: + if name not in rows: + continue + split = splits[name] + if len(split) < rows[name]: + raise IncompleteError(f"{name} has {len(split)} rows, need {rows[name]}") + for i in range(rows[name]): + queries.append(Query(name, i, split.texts[i], int(split.intents[i]))) + return queries + + +def make_record(query: Query, pred: llm.LLMPrediction, identity_sha: str) -> dict[str, Any]: + space = load_label_space() + return { + "split": query.split, + "index": query.index, + "query_sha256": query.text_sha256, + "gold_intent": query.gold_intent, + "gold_agent": AGENTS[int(space.intent_to_agent_id[query.gold_intent])], + "raw_text": pred.raw_text, + "agent": pred.agent, + "parse_failed": not pred.parsed, + "input_tokens": pred.input_tokens, + "output_tokens": pred.output_tokens, + "cache_creation_input_tokens": pred.cache_creation_input_tokens, + "cache_read_input_tokens": pred.cache_read_input_tokens, + "cost_usd": pred.cost_usd, + "latency_ms": pred.latency_ms, + "attempts": pred.attempts, + "request_id": pred.request_id, + "stop_reason": pred.stop_reason, + "identity_sha256": identity_sha, + "created_at": utc_now(), + } + + +def check_record_fields(record: dict[str, Any], where: str) -> None: + for name, expected in RECORD_FIELDS.items(): + if name not in record: + raise JournalError(f"{where}: record has no {name!r}") + value = record[name] + wrong_bool = isinstance(value, bool) and expected is not bool + if wrong_bool or not isinstance(value, expected): + raise JournalError(f"{where}: {name}={value!r} is not {expected}") + + +def read_jsonl(path: Path, allow_torn_tail: bool) -> list[dict[str, Any]]: + """Parse every line; a crash can leave the last journal line half written, and only it.""" + lines = path.read_text(encoding="utf-8").splitlines() + out = [] + for n, line in enumerate(lines, 1): + try: + out.append(json.loads(line)) + except json.JSONDecodeError as exc: + if allow_torn_tail and n == len(lines): + break + raise JournalError(f"{path}:{n}: not JSON ({exc})") from None + return out + + +def trim_torn_tail(path: Path) -> None: + """Cut a half-written last line, so the next append starts on a line of its own.""" + if not path.exists(): + return + data = path.read_bytes() + if data and not data.endswith(b"\n"): + with path.open("r+b") as fh: + fh.truncate(data.rfind(b"\n") + 1) + + +def load_journal(path: Path, queries: dict[Key, Query], identity_sha: str) -> dict[Key, dict]: + """Successful records of this identity, re-parsed with the current ``parse_agent``.""" + if not path.exists(): + return {} + done: dict[Key, dict[str, Any]] = {} + for n, record in enumerate(read_jsonl(path, allow_torn_tail=True), 1): + if record.get("status") == "failed": + continue + where = f"{path}:{n}" + check_record_fields(record, where) + key = (record["split"], record["index"]) + query = queries.get(key) + if record["identity_sha256"] != identity_sha: + raise JournalError(f"{where}: identity {record['identity_sha256']} != {identity_sha}") + if query is None or record["query_sha256"] != query.text_sha256: + raise JournalError(f"{where}: {key} does not match the dataset row it names") + if record["gold_intent"] != query.gold_intent: + raise JournalError(f"{where}: {key} gold label differs from the dataset") + if key in done: + raise JournalError(f"{where}: {key} recorded twice; was the runner started twice?") + record["agent"], parsed = llm.parse_agent(record["raw_text"]) + record["parse_failed"] = not parsed + done[key] = record + return done + + +def query_cost_bound(query: Query, base_tokens: int) -> float: + """Upper bound on one call's cost: every query byte a token, output at max_tokens.""" + return llm.cost_usd(base_tokens + len(query.text.encode("utf-8")), llm.MAX_TOKENS) + + +def announce_estimate( + todo: list[Query], base_tokens: int, outcome: Outcome, max_usd: float, log: Log +) -> None: + bound = sum(query_cost_bound(q, base_tokens) for q in todo) + in_bound = sum(base_tokens + len(q.text.encode("utf-8")) for q in todo) + log( + f"estimate: {len(todo)} calls to make; prompt is {base_tokens} input tokens with a " + f"1-character query; upper bound {in_bound} input + {len(todo) * llm.MAX_TOKENS} output " + f"tokens = US${bound:.4f}. Already spent US${outcome.spent_before:.4f} on " + f"{len(outcome.records)} stored records. Cap US${max_usd:.2f}." + ) + if outcome.spent_before + bound > max_usd: + log("warning: the upper bound passes the cap; the run stops at the cap if it gets there") + + +@dataclass +class Session: + """State shared by the dispatch loop and the per-call bookkeeping of one run.""" + + client: Any + base_tokens: int + max_usd: float + workers: int + outcome: Outcome + journal: TextIO + sleep: Callable[[float], None] + log: Log + identity_sha: str + + def spent(self) -> float: + return self.outcome.spent_before + self.outcome.spent_now + + +def run( + queries: list[Query], + client: Any, + target: Target, + *, + max_usd: float, + workers: int = DEFAULT_WORKERS, + sleep: Callable[[float], None] = time.sleep, + log: Log = print, +) -> Outcome: + """Call the model for every query without a stored record, within the cost cap.""" + identity_sha = llm.identity_sha256() + by_key = {q.key: q for q in queries} + journal = target.journal(identity_sha) + records = load_journal(journal, by_key, identity_sha) + outcome = Outcome(records, spent_before=sum(r["cost_usd"] for r in records.values())) + todo = [q for q in queries if q.key not in records] + if not todo: + return outcome + base = llm.prompt_base_tokens(client) + announce_estimate(todo, base, outcome, max_usd, log) + journal.parent.mkdir(parents=True, exist_ok=True) + trim_torn_tail(journal) + with journal.open("a", encoding="utf-8") as fh, ThreadPoolExecutor(workers) as pool: + session = Session(client, base, max_usd, workers, outcome, fh, sleep, log, identity_sha) + dispatch(session, todo, pool) + return outcome + + +def dispatch(session: Session, todo: list[Query], pool: ThreadPoolExecutor) -> None: + """Keep up to ``workers`` calls in flight, never starting one the cap cannot cover.""" + outcome = session.outcome + pending = list(reversed(todo)) + inflight: dict[Future, tuple[Query, float]] = {} + reserved = 0.0 + while pending or inflight: + while pending and len(inflight) < session.workers and outcome.stopped is None: + bound = query_cost_bound(pending[-1], session.base_tokens) + if session.spent() + reserved + bound > session.max_usd: + outcome.stopped = "cost cap" + break + query = pending.pop() + call = pool.submit(llm.classify, query.text, session.client, session.sleep) + inflight[call] = (query, bound) + reserved += bound + if not inflight: + break + finished, _ = wait(inflight, return_when=FIRST_COMPLETED) + for future in finished: + query, bound = inflight.pop(future) + reserved -= bound + settle(session, future, query) + + +def settle(session: Session, future: Future, query: Query) -> None: + """Journal one finished call and update the running totals.""" + outcome = session.outcome + try: + pred = future.result() + except llm.LLMCallError as exc: + failure = { + "status": "failed", + "split": query.split, + "index": query.index, + "attempts": exc.attempts, + "retryable": exc.retryable, + "error": llm.redact(exc.detail), + "identity_sha256": session.identity_sha, + "created_at": utc_now(), + } + write_line(session.journal, failure) + outcome.failures.append(failure) + session.log(f"FAILED {query.key} after {exc.attempts} attempt(s): {failure['error']}") + if not exc.retryable and outcome.stopped is None: + outcome.stopped = "non-retryable API error" + return + record = make_record(query, pred, session.identity_sha) + write_line(session.journal, record) + outcome.records[query.key] = record + outcome.spent_now += record["cost_usd"] + outcome.calls_now += 1 + + +def write_line(fh: TextIO, record: dict[str, Any]) -> None: + fh.write(json.dumps(record, sort_keys=True, ensure_ascii=False) + "\n") + fh.flush() + + +def expected_keys(rows: dict[str, int]) -> set[Key]: + return {(split, i) for split, n in rows.items() for i in range(n)} + + +def check_predictions(path: Path, rows: dict[str, int], identity_sha: str) -> int: + """``path`` holds exactly the expected rows, once each, all of ``identity_sha``.""" + records = read_jsonl(path, allow_torn_tail=False) + for n, record in enumerate(records, 1): + check_record_fields(record, f"{path}:{n}") + if record["identity_sha256"] != identity_sha: + raise IncompleteError(f"{path}:{n}: identity {record['identity_sha256']}") + counts = Counter((r["split"], r["index"]) for r in records) + expected = expected_keys(rows) + duplicated = sorted(k for k, c in counts.items() if c > 1) + missing = sorted(expected - set(counts)) + unexpected = sorted(set(counts) - expected) + if duplicated or missing or unexpected: + raise IncompleteError( + f"{path}: expected {len(expected)} unique rows, got {len(records)}; " + f"missing {missing[:5]}{'...' if len(missing) > 5 else ''} ({len(missing)}), " + f"duplicated {duplicated[:5]} ({len(duplicated)}), unexpected {unexpected[:5]}" + ) + return len(records) + + +def write_predictions(target: Target, records: dict[Key, dict]) -> Path: + order = {name: i for i, name in enumerate(SPLIT_ORDER)} + keys = sorted(records, key=lambda k: (order[k[0]], k[1])) + tmp = target.predictions.with_name(target.predictions.name + ".tmp") + tmp.parent.mkdir(parents=True, exist_ok=True) + with tmp.open("w", encoding="utf-8") as fh: + for key in keys: + write_line(fh, records[key]) + try: + check_predictions(tmp, target.rows, llm.identity_sha256()) + except Exception: + tmp.unlink(missing_ok=True) + raise + os.replace(tmp, target.predictions) + return target.predictions + + +def totals(records: Iterable[dict[str, Any]]) -> dict[str, Any]: + records = list(records) + token_fields = ( + "input_tokens", + "output_tokens", + "cache_creation_input_tokens", + "cache_read_input_tokens", + ) + out: dict[str, Any] = {f: sum(r[f] for r in records) for f in token_fields} + out["calls"] = len(records) + out["cost_usd"] = round(sum(r["cost_usd"] for r in records), 6) + out["parse_failed"] = { + s: sum(r["parse_failed"] for r in records if r["split"] == s) for s in SPLIT_ORDER + } + return out + + +def summary_body(target: Target, outcome: Outcome, max_usd: float, sha: str) -> dict[str, Any]: + commit, dirty = git_state() + return { + "name": target.name, + "identity": llm.identity(), + "identity_sha256": llm.identity_sha256(), + "dataset_revision": DATASET_REVISION, + "rows": target.rows, + "pricing_usd_per_mtok": llm.PRICE_USD_PER_MTOK, + "pricing_source": llm.PRICING_SOURCE, + "predictions_file": target.predictions.name, + "predictions_sha256": sha, + "totals": totals(outcome.records.values()), + "last_invocation": { + "calls": outcome.calls_now, + "cost_usd": round(outcome.spent_now, 6), + "max_usd": max_usd, + "failed_calls": len(outcome.failures), + }, + "git_commit": commit, + "git_dirty": dirty, + "created_at": utc_now(), + } + + +def finalize(target: Target, outcome: Outcome, max_usd: float) -> int: + """Write predictions, summary and manifest entry, then check that all three agree.""" + path = write_predictions(target, outcome.records) + sha = sha256_of(path) + body = summary_body(target, outcome, max_usd, sha) + target.summary.write_text(json.dumps(body, indent=2) + "\n", encoding="utf-8") + if target.manifest is not None: + files = read_manifest(target.manifest) + files[path.name] = { + "sha256": sha, + "bytes": path.stat().st_size, + "identity_sha256": body["identity_sha256"], + "rows": target.rows, + "git_commit": body["git_commit"], + "created_at": body["created_at"], + } + write_manifest(target.manifest, files) + return verify(target) + + +def verify(target: Target) -> int: + """Row check plus SHA-256 agreement between file, summary and (if any) manifest.""" + identity_sha = llm.identity_sha256() + count = check_predictions(target.predictions, target.rows, identity_sha) + on_disk = sha256_of(target.predictions) + summary = json.loads(target.summary.read_text(encoding="utf-8")) + listed = on_disk + if target.manifest is not None: + listed = read_manifest(target.manifest).get(target.predictions.name, {}).get("sha256") + if not summary.get("predictions_sha256") == listed == on_disk: + raise IncompleteError( + f"{target.predictions.name} SHA-256 differs: summary " + f"{summary.get('predictions_sha256')}, manifest {listed}, file {on_disk}" + ) + return count + + +def report_spend(target: Target, outcome: Outcome, log: Log) -> None: + t = totals(outcome.records.values()) + log( + f"this run: {outcome.calls_now} calls, US${outcome.spent_now:.4f}; " + f"stored: {t['calls']}/{target.expected_total} records, {t['input_tokens']} input + " + f"{t['output_tokens']} output tokens, US${t['cost_usd']:.4f} total" + ) + if t["calls"]: + per_query = t["cost_usd"] / t["calls"] + full = FULL_QUERIES + log( + f"mean US${per_query:.6f} per query; at that rate {full} queries cost " + f"US${per_query * full:.2f}" + ) + + +def log_redacted(message: str) -> None: + print(llm.redact(message), flush=True) + + +def parse_args(argv: list[str] | None) -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Run Claude Haiku over CLINC150 (8-way routing).") + parser.add_argument("--smoke", action="store_true", help="validation rows 0-19 only") + parser.add_argument("--verify", action="store_true", help="check stored predictions only") + parser.add_argument("--max-usd", type=float, default=DEFAULT_MAX_USD) + parser.add_argument("--workers", type=int, default=DEFAULT_WORKERS) + parser.add_argument("--results-root", default="results") + args = parser.parse_args(argv) + if not 1 <= args.workers <= MAX_WORKERS: + parser.error(f"--workers must be between 1 and {MAX_WORKERS}") + if args.max_usd <= 0: + parser.error("--max-usd must be positive") + return args + + +def main( + argv: list[str] | None = None, + *, + client_factory: Callable[[], Any] = llm.make_client, + load: Callable[[SplitName], Split] = load_split, + sleep: Callable[[float], None] = time.sleep, +) -> None: + args = parse_args(argv) + root = Path(args.results_root) + target = smoke_target(root) if args.smoke else full_target(root) + what = "llm smoke predictions" if args.smoke else "llm predictions" + if args.verify: + print(f"completed {verify(target)}/{target.expected_total} {what}") + return + try: + client = client_factory() + except llm.MissingKeyError as exc: + print(f"error: {exc}", file=sys.stderr) + raise SystemExit(2) from None + splits = {name: load(name) for name in target.rows} + queries = build_queries(splits, target.rows) + outcome = run( + queries, + client, + target, + max_usd=args.max_usd, + workers=args.workers, + sleep=sleep, + log=log_redacted, + ) + report_spend(target, outcome, log_redacted) + missing = target.expected_total - len(outcome.records) + if missing: + reason = outcome.stopped or f"{len(outcome.failures)} call(s) failed after retries" + log_redacted( + f"incomplete: {missing} rows have no reply ({reason}); rerun to continue, " + "only those rows are called" + ) + raise SystemExit(1) + print(f"completed {finalize(target, outcome, args.max_usd)}/{target.expected_total} {what}") + + +if __name__ == "__main__": + main() diff --git a/tests/llm_fakes.py b/tests/llm_fakes.py new file mode 100644 index 0000000..388b0fa --- /dev/null +++ b/tests/llm_fakes.py @@ -0,0 +1,95 @@ +"""A fake Anthropic client for the Haiku runner tests; nothing here touches the network.""" + +from __future__ import annotations + +import threading +from collections.abc import Callable +from types import SimpleNamespace + +import numpy as np + +from tinyrouter.data import Split + + +def api_error(status: int, message: str = "simulated", retry_after: str | None = None): + """A real SDK status error (needs the ``llm`` group), as the client would raise it.""" + import anthropic + import httpx2 + + request = httpx2.Request("POST", "https://api.anthropic.com/v1/messages") + headers = {"retry-after": retry_after} if retry_after is not None else {} + response = httpx2.Response(status, request=request, headers=headers) + classes = { + 401: anthropic.AuthenticationError, + 429: anthropic.RateLimitError, + 500: anthropic.InternalServerError, + 529: anthropic.InternalServerError, + } + return classes.get(status, anthropic.APIStatusError)(message, response=response, body=None) + + +def connection_error(): + import anthropic + import httpx2 + + request = httpx2.Request("POST", "https://api.anthropic.com/v1/messages") + return anthropic.APIConnectionError(request=request) + + +class FakeMessages: + """``create`` replies ``reply(query)``; ``failures[query]`` is raised first, in order.""" + + def __init__( + self, + reply: Callable[[str], str] = lambda q: "auto_agent", + failures: dict[str, list[BaseException]] | None = None, + input_tokens: int = 250, + output_tokens: int = 4, + base_tokens: int = 240, + ) -> None: + self.reply = reply + self.failures = {k: list(v) for k, v in (failures or {}).items()} + self.input_tokens = input_tokens + self.output_tokens = output_tokens + self.base_tokens = base_tokens + self.calls: list[dict] = [] + self.count_calls = 0 + self._lock = threading.Lock() + + def create(self, **kwargs): + query = kwargs["messages"][0]["content"] + with self._lock: + self.calls.append(kwargs) + pending = self.failures.get(query) + error = pending.pop(0) if pending else None + if error is not None: + raise error + n = len(self.calls) + return SimpleNamespace( + content=[SimpleNamespace(type="text", text=self.reply(query))], + usage=SimpleNamespace( + input_tokens=self.input_tokens, + output_tokens=self.output_tokens, + cache_creation_input_tokens=0, + cache_read_input_tokens=None, + ), + stop_reason="end_turn", + _request_id=f"req_fake_{n}", + ) + + def count_tokens(self, **kwargs): + self.count_calls += 1 + return SimpleNamespace(input_tokens=self.base_tokens) + + def queries_called(self) -> list[str]: + return [c["messages"][0]["content"] for c in self.calls] + + +def fake_client(**kwargs) -> SimpleNamespace: + return SimpleNamespace(messages=FakeMessages(**kwargs)) + + +def fake_split(name: str, n: int) -> Split: + """``n`` rows named ``-``, alternating intents 0 and 1.""" + texts = tuple(f"{name}-{i}" for i in range(n)) + return Split(name, texts, np.array([i % 2 for i in range(n)], dtype=np.int64)) # type: ignore[arg-type] diff --git a/tests/test_llm.py b/tests/test_llm.py index 9670b66..cc811b8 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -1,56 +1,134 @@ -from types import SimpleNamespace +import hashlib import pytest +from tinyrouter import llm from tinyrouter.llm import SYSTEM_PROMPT, classify, parse_agent +anthropic = pytest.importorskip("anthropic") + +from llm_fakes import api_error, connection_error, fake_client # noqa: E402 + +# SHA-256 of SYSTEM_PROMPT in cost-aware-hybrid-router, src/routers/llm_router.py +# (github.com/drewOrc/cost-aware-hybrid-router, main, fetched 2026-09-29): the +# prompt there is built from the same AGENT_DESCRIPTIONS by the same join. +OLD_PROJECT_PROMPT_SHA256 = "560d22c59164f1125e4ec787f9002b4c9bc6bf03b4cc78472550ae850d5df574" + @pytest.mark.parametrize( ("reply", "agent", "parsed"), [ ("finance_agent", "finance_agent", True), (' "Travel_Agent". ', "travel_agent", True), + ("\n DEVICE_AGENT \n", "device_agent", True), ("oos", "oos", True), + ("OOS.", "oos", True), ("I think kitchen_agent", "kitchen_agent", True), ("finance_agent or travel_agent", "oos", False), ("no idea", "oos", False), + ("", "oos", False), ], ) def test_parse_agent(reply, agent, parsed): assert parse_agent(reply) == (agent, parsed) +def test_system_prompt_is_byte_identical_to_the_old_project(): + digest = hashlib.sha256(SYSTEM_PROMPT.encode("utf-8")).hexdigest() + assert digest == OLD_PROJECT_PROMPT_SHA256 + + def test_system_prompt_lists_all_eight_labels(): for label in ("finance_agent", "meta_agent", "oos"): assert f"- {label}:" in SYSTEM_PROMPT -class FakeMessages: - def __init__(self, failures: int): - self.failures = failures - self.calls: list[dict] = [] +def test_request_matches_the_old_project_settings(): + params = llm.request_params("book a flight") + assert params == { + "model": "claude-haiku-4-5-20251001", + "max_tokens": 20, + "temperature": 0, + "system": SYSTEM_PROMPT, + "messages": [{"role": "user", "content": "book a flight"}], + } + + +@pytest.mark.parametrize( + ("name", "value"), + [ + ("MODEL", "claude-haiku-4-5"), + ("MAX_TOKENS", 21), + ("TEMPERATURE", 1), + ("SYSTEM_PROMPT", SYSTEM_PROMPT + " "), + ], +) +def test_identity_changes_when_any_request_setting_changes(monkeypatch, name, value): + before = llm.identity_sha256() + monkeypatch.setattr(llm, name, value) + assert llm.identity_sha256() != before + + +def test_cost_uses_haiku_rates_per_million_tokens(): + assert llm.cost_usd(1_000_000, 0) == pytest.approx(1.00) + assert llm.cost_usd(0, 1_000_000) == pytest.approx(5.00) + assert llm.cost_usd(250, 4) == pytest.approx(250e-6 + 20e-6) - def create(self, **kwargs): - self.calls.append(kwargs) - if len(self.calls) <= self.failures: - raise ConnectionError("simulated") - return SimpleNamespace( - content=[SimpleNamespace(text="auto_agent")], - usage=SimpleNamespace(input_tokens=120, output_tokens=3), - ) +def test_classify_returns_prediction_with_usage_and_request_id(): + client = fake_client() + result = classify("when is my oil change due", client, sleep=lambda _: None) + assert result.agent == "auto_agent" and result.parsed + assert (result.input_tokens, result.output_tokens) == (250, 4) + assert result.request_id == "req_fake_1" + assert result.attempts == 1 + assert client.messages.calls[-1]["temperature"] == 0 -def test_classify_retries_then_returns_prediction_with_token_counts(): - fake = SimpleNamespace(messages=FakeMessages(failures=2)) + +def test_classify_retries_a_429_then_succeeds_honouring_retry_after(): + client = fake_client(failures={"hi": [api_error(429, retry_after="7"), api_error(529)]}) + sleeps: list[float] = [] + result = classify("hi", client, sleep=sleeps.append) + assert result.attempts == 3 + assert sleeps == [7.0, 2.0] + + +def test_classify_retries_connection_errors(): + client = fake_client(failures={"hi": [connection_error()]}) + assert classify("hi", client, sleep=lambda _: None).attempts == 2 + + +def test_classify_gives_up_after_max_attempts_as_a_retryable_failure(): + client = fake_client(failures={"hi": [api_error(500)] * 99}) sleeps: list[float] = [] - result = classify("when is my oil change due", fake, sleep=sleeps.append) - assert result.agent == "auto_agent" - assert (result.input_tokens, result.output_tokens) == (120, 3) - assert sleeps == [1, 2] - assert fake.messages.calls[-1]["temperature"] == 0 + with pytest.raises(llm.LLMCallError) as info: + classify("hi", client, sleep=sleeps.append) + assert info.value.attempts == llm.MAX_ATTEMPTS + assert info.value.retryable + assert sleeps == [1.0, 2.0, 4.0, 8.0] + assert len(client.messages.calls) == llm.MAX_ATTEMPTS + + +@pytest.mark.parametrize("status", [400, 401, 403, 404]) +def test_classify_does_not_retry_client_errors(status): + client = fake_client(failures={"hi": [api_error(status)]}) + with pytest.raises(llm.LLMCallError) as info: + classify("hi", client, sleep=lambda _: pytest.fail("slept on a non-retryable error")) + assert not info.value.retryable + assert len(client.messages.calls) == 1 + + +def test_retry_delay_is_capped(): + assert llm.retry_delay(api_error(429, retry_after="3600"), 1) == llm.MAX_BACKOFF_S + assert llm.retry_delay(api_error(429, retry_after="soon"), 3) == 4.0 + + +def test_make_client_refuses_without_a_key_and_names_the_variable(monkeypatch): + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + with pytest.raises(llm.MissingKeyError, match="ANTHROPIC_API_KEY is not set"): + llm.make_client() -def test_classify_gives_up_after_max_attempts(): - fake = SimpleNamespace(messages=FakeMessages(failures=99)) - with pytest.raises(RuntimeError, match="5 attempts"): - classify("hi", fake, sleep=lambda _: None) +def test_make_client_turns_sdk_retries_off(monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-test-not-a-real-key") + assert llm.make_client().max_retries == 0 diff --git a/tests/test_llm_deps.py b/tests/test_llm_deps.py new file mode 100644 index 0000000..f699f45 --- /dev/null +++ b/tests/test_llm_deps.py @@ -0,0 +1,22 @@ +"""The Haiku tests use the real SDK's exception classes, so CI must install the ``llm`` group. + +Those tests skip when ``anthropic`` is missing (a plain ``uv sync`` does not +install it). This file does not skip: in CI a missing SDK is a failure, so +the skip cannot hide every retry and leak test at once. +""" + +import importlib.util +import os +from pathlib import Path + +ROOT = Path(__file__).parent.parent + + +def test_ci_installs_the_llm_group(): + workflow = (ROOT / ".github/workflows/ci.yml").read_text(encoding="utf-8") + assert "uv sync --locked --group llm" in workflow + + +def test_anthropic_is_importable_in_ci(): + if os.environ.get("CI"): + assert importlib.util.find_spec("anthropic") is not None diff --git a/tests/test_llm_run.py b/tests/test_llm_run.py new file mode 100644 index 0000000..24b50d4 --- /dev/null +++ b/tests/test_llm_run.py @@ -0,0 +1,332 @@ +import json + +import pytest + +from tinyrouter import llm, llm_run +from tinyrouter.llm_run import IncompleteError, JournalError, Target + +pytest.importorskip("anthropic") + +from llm_fakes import api_error, fake_client, fake_split # noqa: E402 + +ROWS = {"validation": 3, "test": 2} +FAKE_KEY = "sk-ant-api03-FAKE-leak-canary-0123456789" +# Fake usage: 250 input + 4 output tokens per call. +CALL_COST = llm.cost_usd(250, 4) + + +@pytest.fixture +def small_rows(monkeypatch): + """``main`` runs the official target; shrink it to ROWS so the fake splits cover it.""" + monkeypatch.setattr(llm_run, "EXPECTED_ROWS", ROWS) + + +def small_target(tmp_path, rows=ROWS) -> Target: + return Target("haiku-test", tmp_path / "llm", rows, tmp_path / "llm-manifest.json") + + +def queries(rows=ROWS): + splits = {name: fake_split(name, n + 2) for name, n in rows.items()} + return llm_run.build_queries(splits, rows) + + +def run(tmp_path, client, rows=ROWS, **kwargs): + kwargs.setdefault("max_usd", 5.0) + kwargs.setdefault("sleep", lambda _: None) + kwargs.setdefault("log", lambda _: None) + return llm_run.run(queries(rows), client, small_target(tmp_path, rows), **kwargs) + + +def journal_lines(tmp_path, rows=ROWS): + path = small_target(tmp_path, rows).journal(llm.identity_sha256()) + return [json.loads(line) for line in path.read_text().splitlines()] + + +def test_build_queries_takes_the_first_rows_of_each_split_in_order(): + keys = [q.key for q in queries()] + assert keys == [ + ("validation", 0), + ("validation", 1), + ("validation", 2), + ("test", 0), + ("test", 1), + ] + + +def test_a_full_run_stores_every_field_and_passes_the_completion_check(tmp_path): + client = fake_client(reply=lambda q: "Kitchen_Agent" if q.endswith("0") else "hmm") + outcome = run(tmp_path, client) + target = small_target(tmp_path) + assert llm_run.finalize(target, outcome, 5.0) == 5 + records = [json.loads(line) for line in target.predictions.read_text().splitlines()] + first = records[0] + assert set(llm_run.RECORD_FIELDS) <= set(first) + assert (first["split"], first["index"], first["agent"], first["parse_failed"]) == ( + "validation", + 0, + "kitchen_agent", + False, + ) + assert first["raw_text"] == "Kitchen_Agent" + assert first["request_id"].startswith("req_fake_") + assert first["gold_intent"] == 0 and first["gold_agent"] in llm.AGENT_DESCRIPTIONS + assert records[1]["agent"] == "oos" and records[1]["parse_failed"] is True + assert "text" not in first and "validation-0" not in json.dumps(records) + + +def test_summary_and_manifest_record_spend_and_the_same_sha256(tmp_path): + outcome = run(tmp_path, fake_client()) + target = small_target(tmp_path) + llm_run.finalize(target, outcome, 5.0) + summary = json.loads(target.summary.read_text()) + manifest = json.loads(target.manifest.read_text())["files"][target.predictions.name] + assert summary["predictions_sha256"] == manifest["sha256"] + assert summary["totals"]["calls"] == 5 + assert summary["totals"]["input_tokens"] == 5 * 250 + assert summary["totals"]["cost_usd"] == pytest.approx(5 * CALL_COST) + assert summary["last_invocation"]["cost_usd"] == pytest.approx(5 * CALL_COST) + assert summary["pricing_usd_per_mtok"]["input"] == 1.00 + + +def test_a_rerun_calls_only_the_rows_without_a_stored_reply(tmp_path): + # A cap that covers two calls' upper bound (base 240 + ~12 bytes, 20 out) stops the first run. + first_cap = 2.5 * llm.cost_usd(240 + 12, llm.MAX_TOKENS) + first = fake_client() + outcome = run(tmp_path, first, max_usd=first_cap, workers=1) + assert outcome.stopped == "cost cap" and len(outcome.records) == 2 + second = fake_client() + outcome = run(tmp_path, second) + called_first = set(first.messages.queries_called()) + called_second = set(second.messages.queries_called()) + assert not called_first & called_second + assert len(called_first | called_second) == 5 + assert len(outcome.records) == 5 and outcome.spent_before == pytest.approx(2 * CALL_COST) + + +def test_a_complete_journal_makes_no_calls_at_all(tmp_path): + run(tmp_path, fake_client()) + again = fake_client() + outcome = run(tmp_path, again) + assert again.messages.calls == [] and again.messages.count_calls == 0 + assert len(outcome.records) == 5 + + +def test_changing_max_tokens_does_not_reuse_stored_replies(tmp_path, monkeypatch): + run(tmp_path, fake_client()) + monkeypatch.setattr(llm, "MAX_TOKENS", 21) + client = fake_client() + outcome = run(tmp_path, client) + assert len(client.messages.calls) == 5 + assert all(c["max_tokens"] == 21 for c in client.messages.calls) + assert outcome.spent_before == 0 + + +def test_changing_the_prompt_does_not_reuse_stored_replies(tmp_path, monkeypatch): + run(tmp_path, fake_client()) + monkeypatch.setattr(llm, "SYSTEM_PROMPT", llm.SYSTEM_PROMPT.replace("ONE", "one")) + client = fake_client() + run(tmp_path, client) + assert len(client.messages.calls) == 5 + + +def test_a_torn_last_journal_line_is_redone_and_trimmed(tmp_path): + run(tmp_path, fake_client(), workers=1) + path = small_target(tmp_path).journal(llm.identity_sha256()) + text = path.read_text() + path.write_text(text[: text.rstrip("\n").rfind("\n") + 1] + '{"split": "te') + client = fake_client() + outcome = run(tmp_path, client) + assert client.messages.queries_called() == ["test-1"] + assert len(outcome.records) == 5 + assert all( + line.startswith("{") and line.endswith("}") for line in path.read_text().splitlines() + ) + + +def test_a_journal_record_that_does_not_match_its_dataset_row_is_refused(tmp_path): + run(tmp_path, fake_client()) + path = small_target(tmp_path).journal(llm.identity_sha256()) + lines = path.read_text().splitlines() + record = json.loads(lines[0]) + record["query_sha256"] = "0" * 64 + path.write_text("\n".join([json.dumps(record), *lines[1:]]) + "\n") + with pytest.raises(JournalError, match="does not match the dataset row"): + run(tmp_path, fake_client()) + + +def test_the_cost_cap_is_never_passed_with_calls_in_flight(tmp_path): + rows = {"validation": 40} + cap = 10 * llm.cost_usd(240 + 20, llm.MAX_TOKENS) + client = fake_client(input_tokens=240 + 13, output_tokens=llm.MAX_TOKENS) + outcome = run(tmp_path, client, rows=rows, max_usd=cap, workers=8) + assert outcome.stopped == "cost cap" + assert 0 < len(client.messages.calls) < 40 + assert outcome.spent_now <= cap + assert sum(r["cost_usd"] for r in journal_lines(tmp_path, rows)) <= cap + + +def test_the_cap_counts_what_earlier_runs_spent(tmp_path): + run(tmp_path, fake_client(), max_usd=2.5 * llm.cost_usd(252, llm.MAX_TOKENS), workers=1) + client = fake_client() + outcome = run(tmp_path, client, max_usd=2 * CALL_COST) + assert client.messages.calls == [] and outcome.stopped == "cost cap" + + +def test_main_exits_1_and_reports_spend_when_the_cap_stops_the_run(small_rows, tmp_path, capsys): + with pytest.raises(SystemExit) as info: + main(tmp_path, ["--max-usd", "0.0005", "--workers", "1"], fake_client()) + assert info.value.code == 1 + out = capsys.readouterr().out + assert "cost cap" in out and "this run: 1 calls" in out + assert "completed" not in out + + +def test_a_call_that_runs_out_of_retries_is_recorded_and_the_run_exits_1( + small_rows, tmp_path, capsys +): + client = fake_client(failures={"test-1": [api_error(500)] * 99}) + with pytest.raises(SystemExit) as info: + main(tmp_path, [], client) + assert info.value.code == 1 + failed = [r for r in journal_lines_full(tmp_path) if r.get("status") == "failed"] + assert [(r["split"], r["index"], r["attempts"]) for r in failed] == [ + ("test", 1, llm.MAX_ATTEMPTS) + ] + assert "1 call(s) failed after retries" in capsys.readouterr().out + # A rerun retries only that row, and then completes. + retry = fake_client() + main(tmp_path, [], retry) + assert retry.messages.queries_called() == ["test-1"] + + +def test_a_non_retryable_error_stops_starting_new_calls(small_rows, tmp_path): + client = fake_client(failures={"validation-0": [api_error(401)]}) + with pytest.raises(SystemExit): + main(tmp_path, ["--workers", "1"], client) + assert client.messages.queries_called() == ["validation-0"] + + +def test_main_prints_completed_only_after_the_checks(small_rows, tmp_path, capsys): + main(tmp_path, [], fake_client()) + out = capsys.readouterr().out.strip().splitlines() + assert out[-1] == "completed 5/5 llm predictions" + main(tmp_path, ["--verify"], None) + assert capsys.readouterr().out.strip() == "completed 5/5 llm predictions" + + +def test_smoke_writes_its_own_files_and_leaves_the_official_ones_alone(tmp_path, capsys): + client = fake_client() + llm_run.main( + ["--smoke", "--results-root", str(tmp_path)], + client_factory=lambda: client, + load=lambda name: fake_split(name, 25), + sleep=lambda _: None, + ) + out = capsys.readouterr().out + assert out.strip().splitlines()[-1] == "completed 20/20 llm smoke predictions" + assert f"at that rate 8600 queries cost US${8600 * CALL_COST:.2f}" in out + assert len(client.messages.calls) == 20 + assert (tmp_path / "llm-smoke" / "haiku-8way-smoke.jsonl").exists() + assert not (tmp_path / "llm").exists() and not (tmp_path / "llm-manifest.json").exists() + + +def test_official_expected_rows_are_the_full_validation_and_test_splits(): + assert llm_run.EXPECTED_ROWS == {"validation": 3100, "test": 5500} + assert llm_run.FULL_QUERIES == sum(llm_run.EXPECTED_ROWS.values()) == 8600 + + +def completed_target(tmp_path) -> Target: + target = small_target(tmp_path) + llm_run.finalize(target, run(tmp_path, fake_client()), 5.0) + return target + + +def test_completion_check_fails_when_one_row_is_missing(tmp_path): + target = completed_target(tmp_path) + lines = target.predictions.read_text().splitlines() + target.predictions.write_text("\n".join(lines[:-1]) + "\n") + with pytest.raises(IncompleteError, match=r"missing \[\('test', 1\)\]"): + llm_run.check_predictions(target.predictions, ROWS, llm.identity_sha256()) + + +def test_completion_check_fails_when_one_row_is_duplicated(tmp_path): + target = completed_target(tmp_path) + lines = target.predictions.read_text().splitlines() + target.predictions.write_text("\n".join([*lines, lines[0]]) + "\n") + with pytest.raises(IncompleteError, match=r"duplicated \[\('validation', 0\)\]"): + llm_run.check_predictions(target.predictions, ROWS, llm.identity_sha256()) + + +def test_completion_check_fails_on_a_row_from_another_identity(tmp_path): + target = completed_target(tmp_path) + lines = target.predictions.read_text().splitlines() + record = json.loads(lines[0]) + record["identity_sha256"] = "f" * 64 + target.predictions.write_text("\n".join([json.dumps(record), *lines[1:]]) + "\n") + with pytest.raises(IncompleteError, match="identity"): + llm_run.check_predictions(target.predictions, ROWS, llm.identity_sha256()) + + +def test_verify_fails_when_the_file_changed_after_the_manifest_was_written(tmp_path): + target = completed_target(tmp_path) + lines = target.predictions.read_text().splitlines() + record = json.loads(lines[0]) + record["latency_ms"] += 1 + target.predictions.write_text("\n".join([json.dumps(record), *lines[1:]]) + "\n") + with pytest.raises(IncompleteError, match="SHA-256 differs"): + llm_run.verify(target) + + +def test_missing_key_exits_2_with_a_clear_message_before_loading_data( + tmp_path, capsys, monkeypatch +): + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + with pytest.raises(SystemExit) as info: + llm_run.main( + ["--results-root", str(tmp_path)], + load=lambda name: pytest.fail("loaded data without a key"), + ) + assert info.value.code == 2 + assert "ANTHROPIC_API_KEY is not set" in capsys.readouterr().err + + +def test_the_key_never_reaches_output_journal_or_exceptions( + small_rows, tmp_path, capsys, monkeypatch +): + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_KEY) + leaky = api_error(401, message=f"invalid x-api-key {FAKE_KEY}") + client = fake_client(failures={"validation-0": [leaky]}) + with pytest.raises(SystemExit): + main(tmp_path, ["--workers", "1"], client) + captured = capsys.readouterr() + assert "invalid x-api-key [redacted]" in captured.out + assert FAKE_KEY not in captured.out + captured.err + for path in tmp_path.rglob("*"): + if path.is_file(): + assert FAKE_KEY not in path.read_text() + with pytest.raises(llm.LLMCallError) as info: + llm.classify("validation-0", fake_client(failures={"validation-0": [leaky]})) + assert FAKE_KEY not in str(info.value) and info.value.__cause__ is None + assert FAKE_KEY not in repr(info.value) + + +def test_workers_and_cap_arguments_are_bounded(): + with pytest.raises(SystemExit): + llm_run.parse_args(["--workers", "9"]) + with pytest.raises(SystemExit): + llm_run.parse_args(["--max-usd", "0"]) + + +def main(tmp_path, argv, client): + llm_run.main( + [*argv, "--results-root", str(tmp_path)], + client_factory=lambda: client, + load=lambda name: fake_split(name, 4), + sleep=lambda _: None, + ) + + +def journal_lines_full(tmp_path): + target = llm_run.full_target(tmp_path) + path = target.journal(llm.identity_sha256()) + return [json.loads(line) for line in path.read_text().splitlines()] diff --git a/tests/test_makefile.py b/tests/test_makefile.py index 58e005f..98fbe8c 100644 --- a/tests/test_makefile.py +++ b/tests/test_makefile.py @@ -60,3 +60,13 @@ def test_an_explicit_setup_still_works(): result = make_dry_run("setup") assert result.returncode == 0, result.stderr assert "uv sync --locked" in result.stdout + + +def test_llm_targets_ask_for_the_llm_group_and_pass_the_cap(): + full = make_dry_run("llm", "MAX_USD=3") + assert full.returncode == 0, full.stderr + assert "uv run --group llm" in full.stdout + assert "tinyrouter.llm_run --max-usd 3" in full.stdout + smoke = make_dry_run("llm-smoke") + assert "tinyrouter.llm_run --smoke" in smoke.stdout + assert "--max-usd" not in smoke.stdout From 8e795f9d0f8d8976f2d668461ada21a0d5be9794 Mon Sep 17 00:00:00 2001 From: drewOrc <36374426+drewOrc@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:03:06 +0800 Subject: [PATCH 2/4] Check that a trimmed journal parses line by line --- tests/test_llm_run.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tests/test_llm_run.py b/tests/test_llm_run.py index 24b50d4..725aaed 100644 --- a/tests/test_llm_run.py +++ b/tests/test_llm_run.py @@ -138,8 +138,10 @@ def test_a_torn_last_journal_line_is_redone_and_trimmed(tmp_path): outcome = run(tmp_path, client) assert client.messages.queries_called() == ["test-1"] assert len(outcome.records) == 5 - assert all( - line.startswith("{") and line.endswith("}") for line in path.read_text().splitlines() + stored = [json.loads(line) for line in path.read_text().splitlines()] + assert [(r["split"], r["index"]) for r in stored][-1] == ("test", 1) + assert ( + len(llm_run.load_journal(path, {q.key: q for q in queries()}, llm.identity_sha256())) == 5 ) From 1f18c63f69e95ae649475e6fa749de24cee12256 Mon Sep 17 00:00:00 2001 From: drewOrc <36374426+drewOrc@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:20:00 +0800 Subject: [PATCH 3/4] Address PR #14 review: run lock, oos parsing, journal tail, bounds - One runner per target: flock on results/llm/.lock, exit 2 when another process holds it (two runs could each spend up to the cap). - parse_agent counts "oos" as a word among the candidates; two or more candidates are ambiguous, so "oos (not travel_agent)" is oos with parse_failed instead of travel_agent. - Repair the journal tail before loading it: a whole last record missing its newline is kept (it was paid for), only a half line is cut. - count_tokens goes through the same retries and key redaction. - settle journals unexpected exceptions as failures and stops; calls in flight at Ctrl-C are waited for and journaled; a response above its cost bound stops the run with exit 1. - Summary records the parser's hash and the cap's scope (per identity and target; the smoke run is not included). - Tests for guards that had none: UTF-8 byte bound, another identity in the journal, changed gold label, manifest-only SHA change, a row recorded twice. --- .gitignore | 1 + DEVLOG.md | 28 +++++ src/tinyrouter/llm.py | 66 ++++++++--- src/tinyrouter/llm_run.py | 238 +++++++++++++++++++++++++++----------- tests/llm_fakes.py | 14 ++- tests/test_llm.py | 29 +++++ tests/test_llm_run.py | 185 ++++++++++++++++++++++++++++- 7 files changed, 474 insertions(+), 87 deletions(-) diff --git a/.gitignore b/.gitignore index 325ab51..9e19bed 100644 --- a/.gitignore +++ b/.gitignore @@ -36,6 +36,7 @@ results/logits/ # GitHub Release like the logits; results/llm-manifest.json and # results/llm/haiku-8way.json are committed. The smoke run is scratch. results/llm/*.jsonl +results/llm/*.lock results/llm-smoke/ *.npz *.tmp diff --git a/DEVLOG.md b/DEVLOG.md index 52dfb6b..1789a61 100644 --- a/DEVLOG.md +++ b/DEVLOG.md @@ -4,6 +4,34 @@ --- +## 2026-09-29(晚):PR #14 審查修正(4 medium、6 low) + +### 本次工作 / 執行摘要 +- 單一執行鎖:`results/llm/.lock` 用 `flock(LOCK_EX|LOCK_NB)`,run 與 finalize 都要拿鎖;拿不到就 exit 2。審查實測兩個行程並行時各自花到上限,合計到 144%。 +- parser:`oos` 以單字邊界也算候選,候選數 ≥ 2 就判 oos 並標 `parse_failed`。`"oos (not travel_agent)"` 原本會被判成 travel_agent,是 RQ2 最危險的靜默誤派。 +- journal 尾行:先修尾再載入。完整但缺換行的尾行保留並補換行(那筆已付費),只有無法解析的半行才截掉。原本是先載入再截尾,會把記憶體裡算成功、磁碟上已刪掉的那筆再付一次費。 +- `count_tokens` 走同一套重試與遮罩(529 會重試、錯誤字串遮 key)。 +- `settle` 把非預期例外也記成 failed 並停跑;Ctrl-C 時先等在途呼叫回來並記帳再往外拋;回應的 input 或 output 超過上界就停跑(`bound violated`,exit 1)。 +- summary 新增 `parser_sha256`(重新解析用的是當下的 parser,而 parser 不在 identity 裡)與 `cap_scope`。 +- 補五個守門的測試:非 ASCII 的上界(bytes 不是字元數)、journal 混入別的身分、gold 被改、只改 manifest 的 SHA、同一列成功兩次。 + +### 核心發現 / 數據 +- **上限的範圍是「每個身分 × 每個 target」**:smoke 與改 prompt 之前的花費不算在內。AC6 的總花費要手動把 `results/llm/haiku-8way.json` 與 smoke 的 summary 加總。 +- 上限看不到的部分:client 端逾時但伺服器已計費的請求,重試時沿用同一份預留;每次這種逾時最多多出一個單筆上界,並行時一波約 workers 個上界。 +- (無實跑數據) + +### Blockers / 遇到的問題 +- (無) + +### Next +- [ ] 下一個分析 PR 同時報新舊兩種解析規則(舊:子字串比對取第一個命中),以及兩者判定不一致的列數 + +### Files / Budget +- `src/tinyrouter/llm.py`、`src/tinyrouter/llm_run.py`、`tests/test_llm.py`、`tests/test_llm_run.py`、`tests/llm_fakes.py`、`.gitignore` +- API 花費:US$0 + +--- + ## 2026-09-29:步驟 4 之一,Haiku 8 類執行器(尚未實跑) ### 本次工作 / 執行摘要 diff --git a/src/tinyrouter/llm.py b/src/tinyrouter/llm.py index 404acc1..c437e63 100644 --- a/src/tinyrouter/llm.py +++ b/src/tinyrouter/llm.py @@ -17,8 +17,10 @@ from __future__ import annotations import hashlib +import inspect import json import os +import re import time from collections.abc import Callable from dataclasses import dataclass @@ -33,6 +35,7 @@ MAX_BACKOFF_S = 60.0 RETRYABLE_STATUS = frozenset({408, 409, 429}) KEY_ENV = "ANTHROPIC_API_KEY" +OOS_WORD = re.compile(r"\boos\b") # USD per million tokens for Claude Haiku 4.5, from the claude-api skill's # "Current Models" table (cached 2026-06-24; first-party API rates), which @@ -160,16 +163,31 @@ def identity_sha256() -> str: def parse_agent(raw_text: str) -> tuple[str, bool]: - """Map the model's reply to an agent; unparseable replies become oos and are flagged.""" + """Map the model's reply to an agent; unparseable replies become oos and are flagged. + + A reply that names exactly one label (an agent name anywhere, or ``oos`` + as a word) gets that label. A reply naming two or more, such as + ``"oos (not travel_agent)"``, is ambiguous and becomes oos with the + flag: taking the agent there would route a query the model called out + of scope, the silent misroute RQ2 counts as the worst error. + """ text = raw_text.strip().strip("\"'`.").lower() if text in AGENTS: return text, True hits = [agent for agent in AGENTS if agent != OOS and agent in text] + if OOS_WORD.search(text): + hits.append(OOS) if len(hits) == 1: return hits[0], True return OOS, False +def parser_sha256() -> str: + """SHA-256 of ``parse_agent``'s source; stored replies are re-parsed with the current one.""" + source = inspect.getsource(parse_agent) + OOS_WORD.pattern + return hashlib.sha256(source.encode("utf-8")).hexdigest() + + def redact(text: str) -> str: """``text`` with the API key's value, if set and present, replaced.""" key = os.environ.get(KEY_ENV, "").strip() @@ -218,17 +236,20 @@ def reply_text(response: Any) -> str: ) -def classify( - query: str, - client: Any, +def with_retries( + call: Callable[[], Any], sleep: Callable[[float], None] = time.sleep, retryable: Callable[[BaseException], bool] = is_retryable, -) -> LLMPrediction: - """One zero-shot call, retried per the module docstring; LLMCallError when it gives up.""" +) -> tuple[Any, int, int]: + """``call()`` retried per the module docstring: (result, attempts, latency of the last, ms). + + Gives up with LLMCallError, whose message is ``describe_error`` (key + redacted) and which is not chained to the SDK exception. + """ for attempt in range(1, MAX_ATTEMPTS + 1): started = time.monotonic() try: - response = client.messages.create(**request_params(query)) + result = call() except Exception as exc: # classified below; anything unknown is not retried can_retry = retryable(exc) if not can_retry or attempt == MAX_ATTEMPTS: @@ -237,11 +258,24 @@ def classify( ) from None sleep(retry_delay(exc, attempt)) continue - latency_ms = round((time.monotonic() - started) * 1000) - return prediction_from(response, latency_ms, attempt) + return result, attempt, round((time.monotonic() - started) * 1000) raise AssertionError("unreachable") +def classify( + query: str, + client: Any, + sleep: Callable[[float], None] = time.sleep, + retryable: Callable[[BaseException], bool] = is_retryable, +) -> LLMPrediction: + """One zero-shot call, retried per the module docstring; LLMCallError when it gives up.""" + params = request_params(query) + response, attempts, latency_ms = with_retries( + lambda: client.messages.create(**params), sleep, retryable + ) + return prediction_from(response, latency_ms, attempts) + + def prediction_from(response: Any, latency_ms: int, attempts: int) -> LLMPrediction: raw_text = reply_text(response) agent, parsed = parse_agent(raw_text) @@ -261,10 +295,16 @@ def prediction_from(response: Any, latency_ms: int, attempts: int) -> LLMPredict ) -def prompt_base_tokens(client: Any) -> int: - """Input tokens of the system prompt plus a one-character query (token counting endpoint).""" - counted = client.messages.count_tokens( - model=MODEL, system=SYSTEM_PROMPT, messages=[{"role": "user", "content": "x"}] +def prompt_base_tokens(client: Any, sleep: Callable[[float], None] = time.sleep) -> int: + """Input tokens of the system prompt plus a one-character query (token counting endpoint). + + Retried and redacted like ``classify``; LLMCallError when it gives up. + """ + counted, _, _ = with_retries( + lambda: client.messages.count_tokens( + model=MODEL, system=SYSTEM_PROMPT, messages=[{"role": "user", "content": "x"}] + ), + sleep, ) return int(counted.input_tokens) diff --git a/src/tinyrouter/llm_run.py b/src/tinyrouter/llm_run.py index 494eda1..77de559 100644 --- a/src/tinyrouter/llm_run.py +++ b/src/tinyrouter/llm_run.py @@ -21,10 +21,23 @@ - ``haiku-8way.json``: summary with identity, pricing, tokens and dollars. Cost cap: before each call the runner reserves an upper bound on that -call's cost (input tokens counted by the API for the prompt plus one byte -per token of query, output at ``max_tokens``) and does not start a call -that could take cumulative spend for this identity past ``--max-usd``. -Actual cost comes from each response's ``usage``. +call's cost (input tokens counted by the API for the prompt plus one token +per UTF-8 byte of query, output at ``max_tokens``) and does not start a +call that could take cumulative spend past ``--max-usd``. Actual cost +comes from each response's ``usage``; a response that uses more than its +bound stops the run ("bound violated", exit 1). The cap covers one +identity of one target: the smoke run and journals of other identities +(an earlier prompt, say) are not counted, so AC6's total is this summary +plus the smoke summary, added by hand. What the cap cannot see: a request +that times out on the client after the server billed it is retried under +the same reservation, so each such timeout can add up to one bound; with +``workers`` calls in flight that is at most about ``workers`` bounds per +wave of timeouts. Calls in flight when the run is interrupted (Ctrl-C) +are waited for and journaled before the interrupt is re-raised. + +One process per target: the runner holds an exclusive ``flock`` on +``.lock`` next to the journal. A second process started while the +first runs exits 2 instead of spending from the same cap. Completion: the predictions file, read back from disk, must hold exactly the expected (split, index) pairs once each, all with this identity, and @@ -36,14 +49,16 @@ from __future__ import annotations import argparse +import fcntl import hashlib import json import os import sys import time from collections import Counter -from collections.abc import Callable, Iterable +from collections.abc import Callable, Iterable, Iterator from concurrent.futures import FIRST_COMPLETED, Future, ThreadPoolExecutor, wait +from contextlib import contextmanager from dataclasses import dataclass, field from pathlib import Path from typing import Any, TextIO @@ -62,6 +77,10 @@ DEFAULT_WORKERS = 6 MAX_WORKERS = 8 MANIFEST_NAME = "llm-manifest.json" +CAP_SCOPE = ( + "per identity and target: totals and the cap cover this target's journal for this " + "identity only; the smoke run and other identities are not included" +) RECORD_FIELDS: dict[str, type | tuple[type, ...]] = { "split": str, "index": int, @@ -95,6 +114,10 @@ class IncompleteError(RuntimeError): """The predictions file is not exactly the expected rows, once each.""" +class RunLockedError(RuntimeError): + """Another process holds this target's lock.""" + + @dataclass(frozen=True) class Target: name: str @@ -109,6 +132,10 @@ def expected_total(self) -> int: def journal(self, identity_sha: str) -> Path: return self.out_dir / f"{self.name}.{identity_sha[:12]}.journal.jsonl" + @property + def lock(self) -> Path: + return self.out_dir / f"{self.name}.lock" + @property def predictions(self) -> Path: return self.out_dir / f"{self.name}.jsonl" @@ -215,14 +242,44 @@ def read_jsonl(path: Path, allow_torn_tail: bool) -> list[dict[str, Any]]: return out -def trim_torn_tail(path: Path) -> None: - """Cut a half-written last line, so the next append starts on a line of its own.""" +@contextmanager +def exclusive(target: Target) -> Iterator[None]: + """Hold ``target.lock`` for the block; RunLockedError at once if another process has it.""" + target.lock.parent.mkdir(parents=True, exist_ok=True) + with target.lock.open("a") as fh: + try: + fcntl.flock(fh, fcntl.LOCK_EX | fcntl.LOCK_NB) + except BlockingIOError: + raise RunLockedError( + f"{target.lock} is held by another process; wait for it to finish" + ) from None + try: + yield + finally: + fcntl.flock(fh, fcntl.LOCK_UN) + + +def repair_tail(path: Path) -> None: + """Make the journal end in a newline before it is read or appended to. + + A last line without its newline is either a whole record whose newline + was not written (kept: the call was paid for; the newline is added) or + half a record (cut: that row is called again). + """ if not path.exists(): return data = path.read_bytes() - if data and not data.endswith(b"\n"): + if not data or data.endswith(b"\n"): + return + start = data.rfind(b"\n") + 1 + try: + json.loads(data[start:]) + except ValueError: with path.open("r+b") as fh: - fh.truncate(data.rfind(b"\n") + 1) + fh.truncate(start) + return + with path.open("ab") as fh: + fh.write(b"\n") def load_journal(path: Path, queries: dict[Key, Query], identity_sha: str) -> dict[Key, dict]: @@ -230,7 +287,7 @@ def load_journal(path: Path, queries: dict[Key, Query], identity_sha: str) -> di if not path.exists(): return {} done: dict[Key, dict[str, Any]] = {} - for n, record in enumerate(read_jsonl(path, allow_torn_tail=True), 1): + for n, record in enumerate(read_jsonl(path, allow_torn_tail=False), 1): if record.get("status") == "failed": continue where = f"{path}:{n}" @@ -301,20 +358,19 @@ def run( ) -> Outcome: """Call the model for every query without a stored record, within the cost cap.""" identity_sha = llm.identity_sha256() - by_key = {q.key: q for q in queries} journal = target.journal(identity_sha) - records = load_journal(journal, by_key, identity_sha) - outcome = Outcome(records, spent_before=sum(r["cost_usd"] for r in records.values())) - todo = [q for q in queries if q.key not in records] - if not todo: - return outcome - base = llm.prompt_base_tokens(client) - announce_estimate(todo, base, outcome, max_usd, log) - journal.parent.mkdir(parents=True, exist_ok=True) - trim_torn_tail(journal) - with journal.open("a", encoding="utf-8") as fh, ThreadPoolExecutor(workers) as pool: - session = Session(client, base, max_usd, workers, outcome, fh, sleep, log, identity_sha) - dispatch(session, todo, pool) + with exclusive(target): + repair_tail(journal) + records = load_journal(journal, {q.key: q for q in queries}, identity_sha) + outcome = Outcome(records, spent_before=sum(r["cost_usd"] for r in records.values())) + todo = [q for q in queries if q.key not in records] + if not todo: + return outcome + base = llm.prompt_base_tokens(client, sleep) + announce_estimate(todo, base, outcome, max_usd, log) + with journal.open("a", encoding="utf-8") as fh, ThreadPoolExecutor(workers) as pool: + session = Session(client, base, max_usd, workers, outcome, fh, sleep, log, identity_sha) + dispatch(session, todo, pool) return outcome @@ -324,52 +380,82 @@ def dispatch(session: Session, todo: list[Query], pool: ThreadPoolExecutor) -> N pending = list(reversed(todo)) inflight: dict[Future, tuple[Query, float]] = {} reserved = 0.0 - while pending or inflight: - while pending and len(inflight) < session.workers and outcome.stopped is None: - bound = query_cost_bound(pending[-1], session.base_tokens) - if session.spent() + reserved + bound > session.max_usd: - outcome.stopped = "cost cap" + try: + while pending or inflight: + while pending and len(inflight) < session.workers and outcome.stopped is None: + bound = query_cost_bound(pending[-1], session.base_tokens) + if session.spent() + reserved + bound > session.max_usd: + outcome.stopped = "cost cap" + break + query = pending.pop() + call = pool.submit(llm.classify, query.text, session.client, session.sleep) + inflight[call] = (query, bound) + reserved += bound + if not inflight: break - query = pending.pop() - call = pool.submit(llm.classify, query.text, session.client, session.sleep) - inflight[call] = (query, bound) - reserved += bound - if not inflight: - break - finished, _ = wait(inflight, return_when=FIRST_COMPLETED) - for future in finished: - query, bound = inflight.pop(future) - reserved -= bound - settle(session, future, query) + finished, _ = wait(inflight, return_when=FIRST_COMPLETED) + for future in finished: + query, bound = inflight.pop(future) + reserved -= bound + settle(session, future, query) + except BaseException: + outcome.stopped = "interrupted" + drain(session, inflight) + raise + + +def drain(session: Session, inflight: dict[Future, tuple[Query, float]]) -> None: + """Wait for calls already sent (they are paid for) and journal them.""" + for future in list(inflight): + query, _ = inflight.pop(future) + future.exception() # blocks until the call returns + settle(session, future, query) def settle(session: Session, future: Future, query: Query) -> None: """Journal one finished call and update the running totals.""" - outcome = session.outcome try: pred = future.result() except llm.LLMCallError as exc: - failure = { - "status": "failed", - "split": query.split, - "index": query.index, - "attempts": exc.attempts, - "retryable": exc.retryable, - "error": llm.redact(exc.detail), - "identity_sha256": session.identity_sha, - "created_at": utc_now(), - } - write_line(session.journal, failure) - outcome.failures.append(failure) - session.log(f"FAILED {query.key} after {exc.attempts} attempt(s): {failure['error']}") - if not exc.retryable and outcome.stopped is None: - outcome.stopped = "non-retryable API error" + stop = None if exc.retryable else "non-retryable API error" + record_failure(session, query, exc.detail, exc.attempts, exc.retryable, stop) + return + except Exception as exc: # a bug or an unexpected response shape; the call may be paid + record_failure(session, query, llm.describe_error(exc), 1, False, "unexpected error") return record = make_record(query, pred, session.identity_sha) write_line(session.journal, record) + outcome = session.outcome outcome.records[query.key] = record outcome.spent_now += record["cost_usd"] outcome.calls_now += 1 + input_bound = session.base_tokens + len(query.text.encode("utf-8")) + if pred.input_tokens > input_bound or pred.output_tokens > llm.MAX_TOKENS: + session.log( + f"BOUND VIOLATED {query.key}: {pred.input_tokens} input (bound {input_bound}), " + f"{pred.output_tokens} output (bound {llm.MAX_TOKENS}); the cap is not safe" + ) + outcome.stopped = "bound violated" + + +def record_failure( + session: Session, query: Query, detail: str, attempts: int, retryable: bool, stop: str | None +) -> None: + failure = { + "status": "failed", + "split": query.split, + "index": query.index, + "attempts": attempts, + "retryable": retryable, + "error": llm.redact(detail), + "identity_sha256": session.identity_sha, + "created_at": utc_now(), + } + write_line(session.journal, failure) + session.outcome.failures.append(failure) + session.log(f"FAILED {query.key} after {attempts} attempt(s): {failure['error']}") + if stop is not None and session.outcome.stopped is None: + session.outcome.stopped = stop def write_line(fh: TextIO, record: dict[str, Any]) -> None: @@ -446,6 +532,8 @@ def summary_body(target: Target, outcome: Outcome, max_usd: float, sha: str) -> "rows": target.rows, "pricing_usd_per_mtok": llm.PRICE_USD_PER_MTOK, "pricing_source": llm.PRICING_SOURCE, + "parser_sha256": llm.parser_sha256(), + "cap_scope": CAP_SCOPE, "predictions_file": target.predictions.name, "predictions_sha256": sha, "totals": totals(outcome.records.values()), @@ -463,6 +551,11 @@ def summary_body(target: Target, outcome: Outcome, max_usd: float, sha: str) -> def finalize(target: Target, outcome: Outcome, max_usd: float) -> int: """Write predictions, summary and manifest entry, then check that all three agree.""" + with exclusive(target): + return write_outputs(target, outcome, max_usd) + + +def write_outputs(target: Target, outcome: Outcome, max_usd: float) -> int: path = write_predictions(target, outcome.records) sha = sha256_of(path) body = summary_body(target, outcome, max_usd, sha) @@ -554,16 +647,29 @@ def main( raise SystemExit(2) from None splits = {name: load(name) for name in target.rows} queries = build_queries(splits, target.rows) - outcome = run( - queries, - client, - target, - max_usd=args.max_usd, - workers=args.workers, - sleep=sleep, - log=log_redacted, - ) + try: + outcome = run( + queries, + client, + target, + max_usd=args.max_usd, + workers=args.workers, + sleep=sleep, + log=log_redacted, + ) + except RunLockedError as exc: + print(f"error: {exc}", file=sys.stderr) + raise SystemExit(2) from None + except llm.LLMCallError as exc: + log_redacted(f"error: token count failed after {exc.attempts} attempt(s): {exc.detail}") + raise SystemExit(1) from None report_spend(target, outcome, log_redacted) + exit_unless_done(target, outcome) + print(f"completed {finalize(target, outcome, args.max_usd)}/{target.expected_total} {what}") + + +def exit_unless_done(target: Target, outcome: Outcome) -> None: + """Exit 1 when rows are missing, or when the run stopped for a reason (bound violated).""" missing = target.expected_total - len(outcome.records) if missing: reason = outcome.stopped or f"{len(outcome.failures)} call(s) failed after retries" @@ -572,7 +678,9 @@ def main( "only those rows are called" ) raise SystemExit(1) - print(f"completed {finalize(target, outcome, args.max_usd)}/{target.expected_total} {what}") + if outcome.stopped is not None: + log_redacted(f"every row has a reply, but the run stopped: {outcome.stopped}") + raise SystemExit(1) if __name__ == "__main__": diff --git a/tests/llm_fakes.py b/tests/llm_fakes.py index 388b0fa..1dc433e 100644 --- a/tests/llm_fakes.py +++ b/tests/llm_fakes.py @@ -45,7 +45,7 @@ def __init__( failures: dict[str, list[BaseException]] | None = None, input_tokens: int = 250, output_tokens: int = 4, - base_tokens: int = 240, + base_tokens: int = 250, ) -> None: self.reply = reply self.failures = {k: list(v) for k, v in (failures or {}).items()} @@ -54,6 +54,8 @@ def __init__( self.base_tokens = base_tokens self.calls: list[dict] = [] self.count_calls = 0 + self.count_failures: list[BaseException] = [] + self.response_override = None self._lock = threading.Lock() def create(self, **kwargs): @@ -65,6 +67,8 @@ def create(self, **kwargs): if error is not None: raise error n = len(self.calls) + if self.response_override is not None: + return self.response_override return SimpleNamespace( content=[SimpleNamespace(type="text", text=self.reply(query))], usage=SimpleNamespace( @@ -79,6 +83,8 @@ def create(self, **kwargs): def count_tokens(self, **kwargs): self.count_calls += 1 + if self.count_failures: + raise self.count_failures.pop(0) return SimpleNamespace(input_tokens=self.base_tokens) def queries_called(self) -> list[str]: @@ -89,7 +95,7 @@ def fake_client(**kwargs) -> SimpleNamespace: return SimpleNamespace(messages=FakeMessages(**kwargs)) -def fake_split(name: str, n: int) -> Split: - """``n`` rows named ``-``, alternating intents 0 and 1.""" - texts = tuple(f"{name}-{i}" for i in range(n)) +def fake_split(name: str, n: int, prefix: str | None = None) -> Split: + """``n`` rows named ``-``, alternating intents 0 and 1.""" + texts = tuple(f"{prefix or name}-{i}" for i in range(n)) return Split(name, texts, np.array([i % 2 for i in range(n)], dtype=np.int64)) # type: ignore[arg-type] diff --git a/tests/test_llm.py b/tests/test_llm.py index cc811b8..9ae34de 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -27,6 +27,16 @@ ("finance_agent or travel_agent", "oos", False), ("no idea", "oos", False), ("", "oos", False), + ("oos (not travel_agent)", "oos", False), + ("This is oos, not travel_agent", "oos", False), + ("travel_agent.", "travel_agent", True), + ("`travel_agent`", "travel_agent", True), + ("**travel_agent**", "travel_agent", True), + ("travel_agent\nThe query asks about flights.", "travel_agent", True), + ("**oos**", "oos", True), + ("travel agent", "oos", False), + ("out of scope", "oos", False), + ("noose", "oos", False), ], ) def test_parse_agent(reply, agent, parsed): @@ -132,3 +142,22 @@ def test_make_client_refuses_without_a_key_and_names_the_variable(monkeypatch): def test_make_client_turns_sdk_retries_off(monkeypatch): monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-test-not-a-real-key") assert llm.make_client().max_retries == 0 + + +def test_parser_hash_follows_the_parser_source(monkeypatch): + before = llm.parser_sha256() + monkeypatch.setattr(llm, "OOS_WORD", __import__("re").compile(r"oos")) + assert llm.parser_sha256() != before + + +def test_token_count_is_retried_and_its_errors_redacted(monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-api03-FAKE-count-canary-99") + client = fake_client() + client.messages.count_failures = [api_error(529), connection_error()] + sleeps: list[float] = [] + assert llm.prompt_base_tokens(client, sleeps.append) == client.messages.base_tokens + assert sleeps == [1.0, 2.0] + client.messages.count_failures = [api_error(401, "bad key sk-ant-api03-FAKE-count-canary-99")] + with pytest.raises(llm.LLMCallError) as info: + llm.prompt_base_tokens(client, sleeps.append) + assert "canary" not in str(info.value) and info.value.__cause__ is None diff --git a/tests/test_llm_run.py b/tests/test_llm_run.py index 725aaed..a11bca1 100644 --- a/tests/test_llm_run.py +++ b/tests/test_llm_run.py @@ -3,7 +3,7 @@ import pytest from tinyrouter import llm, llm_run -from tinyrouter.llm_run import IncompleteError, JournalError, Target +from tinyrouter.llm_run import IncompleteError, JournalError, RunLockedError, Target pytest.importorskip("anthropic") @@ -89,8 +89,8 @@ def test_summary_and_manifest_record_spend_and_the_same_sha256(tmp_path): def test_a_rerun_calls_only_the_rows_without_a_stored_reply(tmp_path): - # A cap that covers two calls' upper bound (base 240 + ~12 bytes, 20 out) stops the first run. - first_cap = 2.5 * llm.cost_usd(240 + 12, llm.MAX_TOKENS) + # A cap that covers two calls' upper bound (base 250 + ~12 bytes, 20 out) stops the first run. + first_cap = 2.2 * llm.cost_usd(250 + 12, llm.MAX_TOKENS) first = fake_client() outcome = run(tmp_path, first, max_usd=first_cap, workers=1) assert outcome.stopped == "cost cap" and len(outcome.records) == 2 @@ -159,7 +159,7 @@ def test_a_journal_record_that_does_not_match_its_dataset_row_is_refused(tmp_pat def test_the_cost_cap_is_never_passed_with_calls_in_flight(tmp_path): rows = {"validation": 40} cap = 10 * llm.cost_usd(240 + 20, llm.MAX_TOKENS) - client = fake_client(input_tokens=240 + 13, output_tokens=llm.MAX_TOKENS) + client = fake_client(input_tokens=240 + 12, output_tokens=llm.MAX_TOKENS, base_tokens=240) outcome = run(tmp_path, client, rows=rows, max_usd=cap, workers=8) assert outcome.stopped == "cost cap" assert 0 < len(client.messages.calls) < 40 @@ -168,7 +168,7 @@ def test_the_cost_cap_is_never_passed_with_calls_in_flight(tmp_path): def test_the_cap_counts_what_earlier_runs_spent(tmp_path): - run(tmp_path, fake_client(), max_usd=2.5 * llm.cost_usd(252, llm.MAX_TOKENS), workers=1) + run(tmp_path, fake_client(), max_usd=2.2 * llm.cost_usd(262, llm.MAX_TOKENS), workers=1) client = fake_client() outcome = run(tmp_path, client, max_usd=2 * CALL_COST) assert client.messages.calls == [] and outcome.stopped == "cost cap" @@ -332,3 +332,178 @@ def journal_lines_full(tmp_path): target = llm_run.full_target(tmp_path) path = target.journal(llm.identity_sha256()) return [json.loads(line) for line in path.read_text().splitlines()] + + +def rewrite_first_journal_record(tmp_path, **changes): + path = small_target(tmp_path).journal(llm.identity_sha256()) + lines = path.read_text().splitlines() + record = json.loads(lines[0]) + record.update(changes) + path.write_text("\n".join([json.dumps(record), *lines[1:]]) + "\n") + return path, lines + + +def test_a_journal_row_of_another_identity_is_refused(tmp_path): + run(tmp_path, fake_client()) + rewrite_first_journal_record(tmp_path, identity_sha256="e" * 64) + with pytest.raises(JournalError, match="identity"): + run(tmp_path, fake_client()) + + +def test_a_journal_row_whose_gold_label_changed_is_refused(tmp_path): + run(tmp_path, fake_client()) + rewrite_first_journal_record(tmp_path, gold_intent=7) + with pytest.raises(JournalError, match="gold label"): + run(tmp_path, fake_client()) + + +def test_a_row_recorded_twice_in_the_journal_is_refused(tmp_path): + run(tmp_path, fake_client()) + path = small_target(tmp_path).journal(llm.identity_sha256()) + first = path.read_text().splitlines()[0] + with path.open("a") as fh: + fh.write(first + "\n") + with pytest.raises(JournalError, match="recorded twice"): + run(tmp_path, fake_client()) + + +def test_verify_fails_when_only_the_manifest_sha256_changed(tmp_path): + target = completed_target(tmp_path) + body = json.loads(target.manifest.read_text()) + body["files"][target.predictions.name]["sha256"] = "0" * 64 + target.manifest.write_text(json.dumps(body)) + with pytest.raises(IncompleteError, match="SHA-256 differs"): + llm_run.verify(target) + + +def test_the_cost_bound_counts_utf8_bytes_not_characters(tmp_path): + """A non-ASCII query at its worst case (one token per byte) stays inside the cap.""" + rows = {"validation": 30} + prefix = "\u00e9" * 10 # 10 characters, 20 bytes; the query is e.g. "éééééééééé-7" + splits = {"validation": fake_split("validation", 30, prefix=prefix)} + worst = 240 + len(f"{prefix}-0".encode()) + client = fake_client(input_tokens=worst, output_tokens=llm.MAX_TOKENS, base_tokens=240) + cap = 7.5 * llm.cost_usd(worst, llm.MAX_TOKENS) + outcome = llm_run.run( + llm_run.build_queries(splits, rows), + client, + small_target(tmp_path, rows), + max_usd=cap, + workers=8, + sleep=lambda _: None, + log=lambda _: None, + ) + assert outcome.stopped == "cost cap" + assert outcome.spent_now <= cap + + +@pytest.mark.parametrize(("extra_in", "extra_out"), [(1, 0), (0, 1)]) +def test_a_response_over_its_bound_stops_the_run_and_exits_1( + small_rows, tmp_path, capsys, extra_in, extra_out +): + # validation-0 is 12 bytes; base 250. + client = fake_client(input_tokens=250 + 12 + extra_in, output_tokens=llm.MAX_TOKENS + extra_out) + with pytest.raises(SystemExit) as info: + main(tmp_path, ["--workers", "1"], client) + assert info.value.code == 1 + assert "bound violated" in capsys.readouterr().out + assert len(client.messages.calls) == 1 + stored = journal_lines_full(tmp_path) + assert len(stored) == 1 and "status" not in stored[0] + + +def test_a_second_runner_on_the_same_target_is_refused(tmp_path): + target = small_target(tmp_path) + client = fake_client() + with llm_run.exclusive(target), pytest.raises(RunLockedError): + run(tmp_path, client) + assert client.messages.calls == [] + assert len(run(tmp_path, client).records) == 5 + + +def test_main_exits_2_when_another_process_holds_the_lock(small_rows, tmp_path, capsys): + import subprocess + import sys + + target = llm_run.full_target(tmp_path) + target.out_dir.mkdir(parents=True) + holder = subprocess.Popen( + [ + sys.executable, + "-c", + "import fcntl, sys, time; f = open(sys.argv[1], 'a'); " + "fcntl.flock(f, fcntl.LOCK_EX); print('held', flush=True); time.sleep(30)", + str(target.lock), + ], + stdout=subprocess.PIPE, + text=True, + ) + try: + assert holder.stdout.readline().strip() == "held" + client = fake_client() + with pytest.raises(SystemExit) as info: + main(tmp_path, [], client) + assert info.value.code == 2 and client.messages.calls == [] + assert "held by another process" in capsys.readouterr().err + finally: + holder.kill() + holder.wait() + + +def test_a_whole_last_record_without_its_newline_is_kept(tmp_path): + run(tmp_path, fake_client(), workers=1) + path = small_target(tmp_path).journal(llm.identity_sha256()) + path.write_text(path.read_text().rstrip("\n")) + client = fake_client() + outcome = run(tmp_path, client) + assert client.messages.calls == [] and len(outcome.records) == 5 + assert path.read_text().endswith("\n") and len(path.read_text().splitlines()) == 5 + + +def test_an_unexpected_exception_is_journaled_as_a_failure_and_stops_the_run(small_rows, tmp_path): + client = fake_client() + client.messages.response_override = object() # no .content or .usage + with pytest.raises(SystemExit) as info: + main(tmp_path, ["--workers", "1"], client) + assert info.value.code == 1 + failed = journal_lines_full(tmp_path) + assert len(failed) == 1 and failed[0]["status"] == "failed" + assert "AttributeError" in failed[0]["error"] + + +def test_calls_in_flight_at_an_interrupt_are_journaled(tmp_path, monkeypatch): + real_wait = llm_run.wait + state = {"calls": 0} + + def interrupted_wait(fs, return_when): + state["calls"] += 1 + if state["calls"] == 1: + raise KeyboardInterrupt + return real_wait(fs, return_when=return_when) + + monkeypatch.setattr(llm_run, "wait", interrupted_wait) + client = fake_client() + with pytest.raises(KeyboardInterrupt): + run(tmp_path, client, workers=3) + assert len(client.messages.calls) == 3 + assert len(journal_lines(tmp_path)) == 3 + + +def test_token_count_failure_exits_1_without_the_key(small_rows, tmp_path, capsys, monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_KEY) + client = fake_client() + client.messages.count_failures = [api_error(401, f"invalid x-api-key {FAKE_KEY}")] + with pytest.raises(SystemExit) as info: + main(tmp_path, [], client) + assert info.value.code == 1 + captured = capsys.readouterr() + assert "token count failed" in captured.out + assert FAKE_KEY not in captured.out + captured.err + assert client.messages.calls == [] + + +def test_summary_records_the_parser_hash_and_the_cap_scope(tmp_path): + target = completed_target(tmp_path) + summary = json.loads(target.summary.read_text()) + assert summary["parser_sha256"] == llm.parser_sha256() + assert "smoke" in summary["cap_scope"] From cbc9b054af970090aa4e14b28c03f56d6c637f70 Mon Sep 17 00:00:00 2001 From: drewOrc <36374426+drewOrc@users.noreply.github.com> Date: Tue, 29 Sep 2026 10:21:32 +0800 Subject: [PATCH 4/4] Share the input-token bound between reservation and check; test a violation on the last row --- src/tinyrouter/llm_run.py | 13 +++++++++---- tests/llm_fakes.py | 6 +++++- tests/test_llm_run.py | 18 +++++++++++++++++- 3 files changed, 31 insertions(+), 6 deletions(-) diff --git a/src/tinyrouter/llm_run.py b/src/tinyrouter/llm_run.py index 77de559..3c0f8b3 100644 --- a/src/tinyrouter/llm_run.py +++ b/src/tinyrouter/llm_run.py @@ -308,16 +308,21 @@ def load_journal(path: Path, queries: dict[Key, Query], identity_sha: str) -> di return done +def input_token_bound(query: Query, base_tokens: int) -> int: + """Most input tokens one call can use: the prompt plus one token per UTF-8 byte of query.""" + return base_tokens + len(query.text.encode("utf-8")) + + def query_cost_bound(query: Query, base_tokens: int) -> float: - """Upper bound on one call's cost: every query byte a token, output at max_tokens.""" - return llm.cost_usd(base_tokens + len(query.text.encode("utf-8")), llm.MAX_TOKENS) + """Upper bound on one call's cost: ``input_token_bound`` in, max_tokens out.""" + return llm.cost_usd(input_token_bound(query, base_tokens), llm.MAX_TOKENS) def announce_estimate( todo: list[Query], base_tokens: int, outcome: Outcome, max_usd: float, log: Log ) -> None: bound = sum(query_cost_bound(q, base_tokens) for q in todo) - in_bound = sum(base_tokens + len(q.text.encode("utf-8")) for q in todo) + in_bound = sum(input_token_bound(q, base_tokens) for q in todo) log( f"estimate: {len(todo)} calls to make; prompt is {base_tokens} input tokens with a " f"1-character query; upper bound {in_bound} input + {len(todo) * llm.MAX_TOKENS} output " @@ -429,7 +434,7 @@ def settle(session: Session, future: Future, query: Query) -> None: outcome.records[query.key] = record outcome.spent_now += record["cost_usd"] outcome.calls_now += 1 - input_bound = session.base_tokens + len(query.text.encode("utf-8")) + input_bound = input_token_bound(query, session.base_tokens) if pred.input_tokens > input_bound or pred.output_tokens > llm.MAX_TOKENS: session.log( f"BOUND VIOLATED {query.key}: {pred.input_tokens} input (bound {input_bound}), " diff --git a/tests/llm_fakes.py b/tests/llm_fakes.py index 1dc433e..1bfbdb8 100644 --- a/tests/llm_fakes.py +++ b/tests/llm_fakes.py @@ -46,7 +46,9 @@ def __init__( input_tokens: int = 250, output_tokens: int = 4, base_tokens: int = 250, + input_tokens_for: Callable[[str], int] | None = None, ) -> None: + self.input_tokens_for = input_tokens_for self.reply = reply self.failures = {k: list(v) for k, v in (failures or {}).items()} self.input_tokens = input_tokens @@ -72,7 +74,9 @@ def create(self, **kwargs): return SimpleNamespace( content=[SimpleNamespace(type="text", text=self.reply(query))], usage=SimpleNamespace( - input_tokens=self.input_tokens, + input_tokens=( + self.input_tokens_for(query) if self.input_tokens_for else self.input_tokens + ), output_tokens=self.output_tokens, cache_creation_input_tokens=0, cache_read_input_tokens=None, diff --git a/tests/test_llm_run.py b/tests/test_llm_run.py index a11bca1..8af44eb 100644 --- a/tests/test_llm_run.py +++ b/tests/test_llm_run.py @@ -377,9 +377,11 @@ def test_verify_fails_when_only_the_manifest_sha256_changed(tmp_path): def test_the_cost_bound_counts_utf8_bytes_not_characters(tmp_path): - """A non-ASCII query at its worst case (one token per byte) stays inside the cap.""" + """A non-ASCII query at its worst case (one token per byte) is inside its bound and the cap.""" rows = {"validation": 30} prefix = "\u00e9" * 10 # 10 characters, 20 bytes; the query is e.g. "éééééééééé-7" + query = llm_run.Query("validation", 0, f"{prefix}-0", 0) + assert llm_run.input_token_bound(query, 240) == 240 + 22 splits = {"validation": fake_split("validation", 30, prefix=prefix)} worst = 240 + len(f"{prefix}-0".encode()) client = fake_client(input_tokens=worst, output_tokens=llm.MAX_TOKENS, base_tokens=240) @@ -507,3 +509,17 @@ def test_summary_records_the_parser_hash_and_the_cap_scope(tmp_path): summary = json.loads(target.summary.read_text()) assert summary["parser_sha256"] == llm.parser_sha256() assert "smoke" in summary["cap_scope"] + + +def test_a_bound_violation_on_the_last_row_still_exits_1(small_rows, tmp_path, capsys): + """Every row has a reply, but one used more tokens than its bound: not a clean finish.""" + client = fake_client( + input_tokens_for=lambda q: 250 + len(q.encode()) + (1 if q == "test-1" else 0) + ) + with pytest.raises(SystemExit) as info: + main(tmp_path, ["--workers", "1"], client) + assert info.value.code == 1 + out = capsys.readouterr().out + assert "every row has a reply, but the run stopped: bound violated" in out + assert "completed" not in out + assert len(client.messages.calls) == 5