Reduce run-llm to a pure batch executor
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -4,19 +4,25 @@
|
||||
# Authors:
|
||||
# imacat@mail.imacat.idv.tw (imacat), 2026/7/30
|
||||
# AI assistance: Claude Code (Anthropic)
|
||||
"""The generic batch runner for one LLM analysis step.
|
||||
"""The generic batch executor for one LLM analysis step.
|
||||
|
||||
Sends every input item to the Anthropic Messages Batch API twice with
|
||||
the same system prompt, reconciles the disagreeing items with a third
|
||||
arbitration batch, and archives every artifact self-contained under
|
||||
``<runs_dir>/<phase>/<YYYYMMDD-HHMM>-<prompt-stem>/``, where the
|
||||
base directory of the run archives is given as the positional
|
||||
command-line argument.
|
||||
One invocation is one definition file plus one input, sent to the
|
||||
Anthropic Messages Batch API exactly once, archived self-contained
|
||||
under the destination directory given by the three positional
|
||||
command-line arguments: prompt, input, archive_dir. The tool
|
||||
knows nothing about run counts or protocols: run identity --
|
||||
run1/run2, arbitration -- lives entirely in the caller's command
|
||||
list, per the research plan. A rerun of an already existing
|
||||
destination requires ``--replace``; any other directory is never
|
||||
touched.
|
||||
|
||||
Comparing runs and reconciling disagreements are the responsibility
|
||||
of separate subcommands, not this one.
|
||||
"""
|
||||
import argparse
|
||||
import enum
|
||||
import hashlib
|
||||
import json
|
||||
import shutil
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import asdict, dataclass
|
||||
@@ -31,21 +37,8 @@ from ..config import get_settings
|
||||
MODEL: str = "claude-sonnet-4-6"
|
||||
TEMPERATURE: float = 0.0
|
||||
THINKING: dict[str, str] = {"type": "disabled"}
|
||||
SCRIPT_VERSION: str = "run_llm.py 1.0.0"
|
||||
SCRIPT_VERSION: str = "run_llm.py 3.0.0"
|
||||
POLL_INTERVAL_SECONDS: float = 60.0
|
||||
ARBITRATION_TEMPLATE: str = (
|
||||
"<item>\n{content}\n</item>\n"
|
||||
"<run1>\n{run1}\n</run1>\n"
|
||||
"<run2>\n{run2}\n</run2>")
|
||||
|
||||
|
||||
class Source(enum.StrEnum):
|
||||
"""The source of a final record."""
|
||||
|
||||
AGREED = "agreed"
|
||||
"""The record takes the run text the two runs agreed on."""
|
||||
ARBITRATION = "arbitration"
|
||||
"""The record takes the arbitration output."""
|
||||
|
||||
|
||||
class InputFormatError(Exception):
|
||||
@@ -174,28 +167,26 @@ def parse_args(argv: list[str] | None) -> argparse.Namespace:
|
||||
:return: The parsed arguments.
|
||||
"""
|
||||
parser: argparse.ArgumentParser = argparse.ArgumentParser(
|
||||
description="Run one LLM step: 2 runs + 1 arbitration.")
|
||||
description="Run one LLM definition file against one input"
|
||||
" and archive the result.")
|
||||
parser.add_argument(
|
||||
"runs_dir", type=Path,
|
||||
help="the base directory of the run archives")
|
||||
parser.add_argument(
|
||||
"--prompt", required=True, type=Path,
|
||||
"prompt", type=Path,
|
||||
help="the prompt definition file, used as the system prompt")
|
||||
parser.add_argument(
|
||||
"--arbitration-prompt", required=True, type=Path,
|
||||
help="the arbitration prompt definition file")
|
||||
parser.add_argument(
|
||||
"--input", required=True, type=Path,
|
||||
"input", type=Path,
|
||||
help="the JSONL input file with \"id\" and \"content\"")
|
||||
parser.add_argument(
|
||||
"--phase", required=True,
|
||||
help="the phase name for the archive directory")
|
||||
"archive_dir", type=Path,
|
||||
help="the destination archive directory")
|
||||
parser.add_argument(
|
||||
"--max-tokens", type=int, default=2048,
|
||||
help="the maximum output tokens per request (default 2048)")
|
||||
parser.add_argument(
|
||||
"--dry-run", action="store_true",
|
||||
help="validate and archive without calling the API")
|
||||
parser.add_argument(
|
||||
"--replace", action="store_true",
|
||||
help="replace an already existing archive directory")
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
@@ -256,19 +247,6 @@ def build_request(item: InputItem, system_prompt: str,
|
||||
}
|
||||
|
||||
|
||||
def build_arbitration_content(content: str, run1_text: str,
|
||||
run2_text: str) -> str:
|
||||
"""Build the arbitration user message for one item.
|
||||
|
||||
:param content: The original item content.
|
||||
:param run1_text: The run-1 output text.
|
||||
:param run2_text: The run-2 output text.
|
||||
:return: The user message text.
|
||||
"""
|
||||
return ARBITRATION_TEMPLATE.format(
|
||||
content=content, run1=run1_text, run2=run2_text)
|
||||
|
||||
|
||||
def submit_batch(client: anthropic.Anthropic,
|
||||
requests: list[dict[str, Any]]) -> str:
|
||||
"""Submit one message batch.
|
||||
@@ -314,6 +292,22 @@ def usage_to_dict(usage: Any) -> dict[str, Any]:
|
||||
if v is not None}
|
||||
|
||||
|
||||
def sum_usage(results: Results) -> dict[str, int]:
|
||||
"""Sum the token usage of every succeeded result.
|
||||
|
||||
:param results: The result records, keyed by item ID.
|
||||
:return: The summed integer usage fields.
|
||||
"""
|
||||
totals: dict[str, int] = {}
|
||||
for result in results.values():
|
||||
if result.usage is None:
|
||||
continue
|
||||
for key, value in result.usage.items():
|
||||
if isinstance(value, int):
|
||||
totals[key] = totals.get(key, 0) + value
|
||||
return totals
|
||||
|
||||
|
||||
def collect_results(client: anthropic.Anthropic,
|
||||
batch_id: str) -> Results:
|
||||
"""Collect the results of an ended batch.
|
||||
@@ -343,78 +337,25 @@ def find_failures(item_ids: list[str],
|
||||
if x not in results or results[x].is_failure]
|
||||
|
||||
|
||||
def split_by_agreement(
|
||||
items: list[InputItem], run1: Results, run2: Results,
|
||||
) -> tuple[list[str], list[str]]:
|
||||
"""Split the item IDs into agreed and disagreeing ones.
|
||||
def create_archive_dir(directory: Path, replace: bool) -> Path:
|
||||
"""Create the archive directory.
|
||||
|
||||
Two outputs agree when their texts are identical after strip().
|
||||
Only this directory is ever created or removed; no other
|
||||
directory is ever touched.
|
||||
|
||||
:param items: The input items.
|
||||
:param run1: The run-1 results, keyed by item ID.
|
||||
:param run2: The run-2 results, keyed by item ID.
|
||||
:return: A tuple of the agreed item IDs and the disagreeing item
|
||||
IDs, both in input order.
|
||||
"""
|
||||
agreed: list[str] = []
|
||||
disagreed: list[str] = []
|
||||
for item in items:
|
||||
text1: str | None = run1[item.id].text
|
||||
text2: str | None = run2[item.id].text
|
||||
assert text1 is not None and text2 is not None
|
||||
if text1.strip() == text2.strip():
|
||||
agreed.append(item.id)
|
||||
else:
|
||||
disagreed.append(item.id)
|
||||
return agreed, disagreed
|
||||
|
||||
|
||||
def build_final_records(
|
||||
items: list[InputItem], run1: Results, arbitration: Results,
|
||||
) -> list[dict[str, str]]:
|
||||
"""Assemble the final records, one per item, in input order.
|
||||
|
||||
An arbitrated item takes the arbitration output; an agreed item
|
||||
takes the agreed (stripped) run text.
|
||||
|
||||
:param items: The input items.
|
||||
:param run1: The run-1 results, keyed by item ID.
|
||||
:param arbitration: The arbitration results, keyed by item ID.
|
||||
:return: The final records with "id", "text", and "source".
|
||||
"""
|
||||
records: list[dict[str, str]] = []
|
||||
for item in items:
|
||||
text: str | None
|
||||
if item.id in arbitration:
|
||||
text = arbitration[item.id].text
|
||||
assert text is not None
|
||||
records.append({"id": item.id, "text": text,
|
||||
"source": Source.ARBITRATION})
|
||||
else:
|
||||
text = run1[item.id].text
|
||||
assert text is not None
|
||||
records.append({"id": item.id, "text": text.strip(),
|
||||
"source": Source.AGREED})
|
||||
return records
|
||||
|
||||
|
||||
def create_archive_dir(runs_root: Path, phase: str, prompt_path: Path,
|
||||
now: datetime) -> Path:
|
||||
"""Create the archive directory for this execution.
|
||||
|
||||
:param runs_root: The root directory of the run archives.
|
||||
:param phase: The phase name.
|
||||
:param prompt_path: The path of the prompt definition file.
|
||||
:param now: The local timestamp of this execution.
|
||||
:param directory: The destination archive directory.
|
||||
:param replace: Whether to remove an already existing archive
|
||||
directory before creating it.
|
||||
:return: The created archive directory.
|
||||
:raises FileExistsError: When the directory already exists.
|
||||
:raises FileExistsError: When the archive directory already
|
||||
exists and ``replace`` is False.
|
||||
"""
|
||||
stem: str = prompt_path.stem
|
||||
name: str = f"{now.strftime('%Y%m%d-%H%M')}-{stem}"
|
||||
directory: Path = runs_root / phase / name
|
||||
if directory.exists():
|
||||
raise FileExistsError(
|
||||
f"archive directory {directory} already exists")
|
||||
if not replace:
|
||||
raise FileExistsError(
|
||||
f"{directory} already exists; pass --replace to"
|
||||
" replace it")
|
||||
shutil.rmtree(directory)
|
||||
directory.mkdir(parents=True)
|
||||
return directory
|
||||
|
||||
@@ -446,17 +387,14 @@ def write_json(path: Path, data: dict[str, Any]) -> None:
|
||||
def write_meta(path: Path, meta: dict[str, Any]) -> None:
|
||||
"""Write the metadata to the ``meta.json`` file.
|
||||
|
||||
The ``BatchInfo`` values under ``batches`` are written as
|
||||
plain JSON objects.
|
||||
The ``BatchInfo`` value under ``batch`` is written as a plain
|
||||
JSON object.
|
||||
|
||||
:param path: The path of the ``meta.json`` file.
|
||||
:param meta: The metadata to write.
|
||||
:return: None.
|
||||
"""
|
||||
write_json(path, {
|
||||
**meta,
|
||||
"batches": {k: asdict(v)
|
||||
for k, v in meta["batches"].items()}})
|
||||
write_json(path, {**meta, "batch": asdict(meta["batch"])})
|
||||
|
||||
|
||||
def sha256_of(path: Path) -> str:
|
||||
@@ -477,13 +415,14 @@ def now_iso() -> str:
|
||||
return datetime.now().astimezone().isoformat(timespec="seconds")
|
||||
|
||||
|
||||
def execute_runs(
|
||||
def execute_run(
|
||||
client: anthropic.Anthropic, items: list[InputItem],
|
||||
system_prompt: str, max_tokens: int, meta: dict[str, Any],
|
||||
) -> tuple[Results, Results]:
|
||||
"""Submit the two identical runs and await their results.
|
||||
system_prompt: str, max_tokens: int,
|
||||
meta: dict[str, Any],
|
||||
) -> Results:
|
||||
"""Submit the batch of this run and await its results.
|
||||
|
||||
The batch IDs and timestamps are recorded into the metadata as an
|
||||
The batch ID and timestamps are recorded into the metadata as an
|
||||
observable side effect.
|
||||
|
||||
:param client: The Anthropic client.
|
||||
@@ -491,72 +430,22 @@ def execute_runs(
|
||||
:param system_prompt: The system prompt text.
|
||||
:param max_tokens: The maximum output tokens per request.
|
||||
:param meta: The metadata to record the batch bookkeeping into.
|
||||
:return: The results of run 1 and run 2, keyed by item ID.
|
||||
:return: The results of this run, keyed by item ID.
|
||||
"""
|
||||
requests: list[dict[str, Any]] = [
|
||||
build_request(x, system_prompt, max_tokens) for x in items]
|
||||
infos: dict[str, BatchInfo] = {}
|
||||
for run_name in ("run1", "run2"):
|
||||
info: BatchInfo = BatchInfo(
|
||||
batch_id=submit_batch(client, requests),
|
||||
submitted_at=now_iso())
|
||||
infos[run_name] = info
|
||||
meta["batches"][run_name] = info
|
||||
print(f"{run_name}: submitted batch {info.batch_id}",
|
||||
file=sys.stderr)
|
||||
batches: dict[str, Any] = poll_batches(
|
||||
client, [x.batch_id for x in infos.values()])
|
||||
for info in infos.values():
|
||||
info.ended_at = batches[info.batch_id].ended_at.isoformat()
|
||||
return (collect_results(client, infos["run1"].batch_id),
|
||||
collect_results(client, infos["run2"].batch_id))
|
||||
|
||||
|
||||
def execute_arbitration(
|
||||
client: anthropic.Anthropic, items: list[InputItem],
|
||||
disagreed: list[str], run1: Results, run2: Results,
|
||||
system_prompt: str, max_tokens: int, meta: dict[str, Any],
|
||||
) -> Results:
|
||||
"""Submit the arbitration batch and await its results.
|
||||
|
||||
The batch ID and timestamps are recorded into the metadata as an
|
||||
observable side effect.
|
||||
|
||||
:param client: The Anthropic client.
|
||||
:param items: The input items.
|
||||
:param disagreed: The disagreeing item IDs.
|
||||
:param run1: The run-1 results, keyed by item ID.
|
||||
:param run2: The run-2 results, keyed by item ID.
|
||||
:param system_prompt: The arbitration system prompt text.
|
||||
:param max_tokens: The maximum output tokens per request.
|
||||
:param meta: The metadata to record the batch bookkeeping into.
|
||||
:return: The arbitration results, keyed by item ID.
|
||||
"""
|
||||
content_by_id: dict[str, str] = {
|
||||
x.id: x.content for x in items}
|
||||
requests: list[dict[str, Any]] = []
|
||||
for item_id in disagreed:
|
||||
text1: str | None = run1[item_id].text
|
||||
text2: str | None = run2[item_id].text
|
||||
assert text1 is not None and text2 is not None
|
||||
requests.append(build_request(
|
||||
InputItem(id=item_id,
|
||||
content=build_arbitration_content(
|
||||
content_by_id[item_id], text1, text2)),
|
||||
system_prompt, max_tokens))
|
||||
info: BatchInfo = BatchInfo(
|
||||
batch_id=submit_batch(client, requests),
|
||||
submitted_at=now_iso())
|
||||
meta["batches"]["arbitration"] = info
|
||||
print(f"arbitration: submitted batch {info.batch_id}",
|
||||
file=sys.stderr)
|
||||
meta["batch"] = info
|
||||
print(f"submitted batch {info.batch_id}", file=sys.stderr)
|
||||
batches: dict[str, Any] = poll_batches(client, [info.batch_id])
|
||||
info.ended_at = batches[info.batch_id].ended_at.isoformat()
|
||||
return collect_results(client, info.batch_id)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
"""Run one LLM step end-to-end.
|
||||
"""Run one LLM definition file against one input and archive it.
|
||||
|
||||
:param argv: The command-line arguments, or None for ``sys.argv``.
|
||||
:return: The exit status: 0 on success, non-zero on failure.
|
||||
@@ -565,89 +454,56 @@ def main(argv: list[str] | None = None) -> int:
|
||||
try:
|
||||
items: list[InputItem] = load_items(args.input)
|
||||
prompt_text: str = args.prompt.read_text(encoding="utf-8")
|
||||
arbitration_text: str = args.arbitration_prompt.read_text(
|
||||
encoding="utf-8")
|
||||
except (OSError, InputFormatError) as error:
|
||||
print(f"error: {error}", file=sys.stderr)
|
||||
return 1
|
||||
try:
|
||||
run_dir: Path = create_archive_dir(
|
||||
args.runs_dir, args.phase, args.prompt,
|
||||
datetime.now())
|
||||
archive_dir: Path = create_archive_dir(
|
||||
args.archive_dir, args.replace)
|
||||
except FileExistsError as error:
|
||||
print(f"error: {error}", file=sys.stderr)
|
||||
return 1
|
||||
meta_path: Path = run_dir / "meta.json"
|
||||
(run_dir / "prompt.md").write_bytes(args.prompt.read_bytes())
|
||||
(run_dir / "arbitration_prompt.md").write_bytes(
|
||||
args.arbitration_prompt.read_bytes())
|
||||
meta_path: Path = archive_dir / "meta.json"
|
||||
(archive_dir / "prompt.md").write_bytes(args.prompt.read_bytes())
|
||||
meta: dict[str, Any] = {
|
||||
"script_version": SCRIPT_VERSION,
|
||||
"model": MODEL,
|
||||
"temperature": TEMPERATURE,
|
||||
"max_tokens": args.max_tokens,
|
||||
"thinking": THINKING,
|
||||
"max_tokens": args.max_tokens,
|
||||
"prompt_path": str(args.prompt),
|
||||
"prompt_sha256": sha256_of(args.prompt),
|
||||
"arbitration_prompt_path": str(args.arbitration_prompt),
|
||||
"arbitration_prompt_sha256": sha256_of(
|
||||
args.arbitration_prompt),
|
||||
"batches": {},
|
||||
"input_path": str(args.input),
|
||||
"input_sha256": sha256_of(args.input),
|
||||
"item_count": len(items),
|
||||
"agreed_count": None,
|
||||
"agreement_rate": None,
|
||||
"dry_run": args.dry_run,
|
||||
"script_version": SCRIPT_VERSION,
|
||||
"started_at": now_iso(),
|
||||
"batch": None,
|
||||
"usage": {},
|
||||
}
|
||||
if args.dry_run:
|
||||
write_meta(meta_path, meta)
|
||||
write_json(meta_path, meta)
|
||||
print(json.dumps(
|
||||
build_request(items[0], prompt_text, args.max_tokens),
|
||||
ensure_ascii=False, indent=2))
|
||||
print(f"dry run: archive created at {run_dir}",
|
||||
print(f"dry run: archive created at {archive_dir}",
|
||||
file=sys.stderr)
|
||||
return 0
|
||||
client: anthropic.Anthropic = anthropic.Anthropic(
|
||||
api_key=get_settings().ANTHROPIC_API_KEY)
|
||||
run1: Results
|
||||
run2: Results
|
||||
run1, run2 = execute_runs(
|
||||
results: Results = execute_run(
|
||||
client, items, prompt_text, args.max_tokens, meta)
|
||||
item_ids: list[str] = [x.id for x in items]
|
||||
write_jsonl(run_dir / "run1.jsonl",
|
||||
[run1[x].to_record() for x in item_ids if x in run1])
|
||||
write_jsonl(run_dir / "run2.jsonl",
|
||||
[run2[x].to_record() for x in item_ids if x in run2])
|
||||
failed: set[str] = (set(find_failures(item_ids, run1))
|
||||
| set(find_failures(item_ids, run2)))
|
||||
write_jsonl(
|
||||
archive_dir / "output.jsonl",
|
||||
[results[x].to_record() for x in item_ids if x in results])
|
||||
meta["usage"] = sum_usage(results)
|
||||
write_meta(meta_path, meta)
|
||||
failed: list[str] = find_failures(item_ids, results)
|
||||
if len(failed) > 0:
|
||||
write_meta(meta_path, meta)
|
||||
names: str = ", ".join(x for x in item_ids if x in failed)
|
||||
print(f"error: failed items: {names}", file=sys.stderr)
|
||||
return 1
|
||||
agreed: list[str]
|
||||
disagreed: list[str]
|
||||
agreed, disagreed = split_by_agreement(items, run1, run2)
|
||||
meta["agreed_count"] = len(agreed)
|
||||
meta["agreement_rate"] = len(agreed) / len(items)
|
||||
arbitration: Results = {}
|
||||
if len(disagreed) > 0:
|
||||
arbitration = execute_arbitration(
|
||||
client, items, disagreed, run1, run2, arbitration_text,
|
||||
args.max_tokens, meta)
|
||||
write_jsonl(run_dir / "arbitration.jsonl",
|
||||
[arbitration[x].to_record() for x in disagreed
|
||||
if x in arbitration])
|
||||
arb_failed: list[str] = find_failures(disagreed, arbitration)
|
||||
if len(arb_failed) > 0:
|
||||
write_meta(meta_path, meta)
|
||||
names = ", ".join(arb_failed)
|
||||
print(f"error: failed arbitration items: {names}",
|
||||
print(f"error: failed items: {', '.join(failed)}",
|
||||
file=sys.stderr)
|
||||
return 1
|
||||
write_jsonl(run_dir / "final.jsonl",
|
||||
build_final_records(items, run1, arbitration))
|
||||
write_meta(meta_path, meta)
|
||||
print(f"done: {len(items)} items, {len(agreed)} agreed,"
|
||||
f" {len(disagreed)} arbitrated; archived to {run_dir}",
|
||||
file=sys.stderr)
|
||||
print(f"done: {len(items)} items;"
|
||||
f" archived to {archive_dir}", file=sys.stderr)
|
||||
return 0
|
||||
|
||||
+131
-186
@@ -3,7 +3,7 @@
|
||||
# Authors:
|
||||
# imacat@mail.imacat.idv.tw (imacat), 2026/7/30
|
||||
# AI assistance: Claude Code (Anthropic)
|
||||
"""Unit tests for the run_llm batch runner module."""
|
||||
"""Unit tests for the run_llm batch executor module."""
|
||||
import io
|
||||
import json
|
||||
import tempfile
|
||||
@@ -180,62 +180,6 @@ class TestRequestBuilding(RunLLMTestCase):
|
||||
self.assertEqual(params["messages"],
|
||||
[{"role": "user", "content": "the lyrics"}])
|
||||
|
||||
def test_arbitration_content(self) -> None:
|
||||
"""Test the arbitration user message format."""
|
||||
content: str = run_llm.build_arbitration_content(
|
||||
"the item", "output one", "output two")
|
||||
self.assertEqual(content,
|
||||
"<item>\nthe item\n</item>\n"
|
||||
"<run1>\noutput one\n</run1>\n"
|
||||
"<run2>\noutput two\n</run2>")
|
||||
|
||||
|
||||
class TestAgreement(RunLLMTestCase):
|
||||
"""Test cases for the agreement computation."""
|
||||
|
||||
def test_split_by_agreement(self) -> None:
|
||||
"""Test splitting items into agreed and disagreeing ones."""
|
||||
items: list[run_llm.InputItem] = [
|
||||
run_llm.InputItem(id="a", content="one"),
|
||||
run_llm.InputItem(id="b", content="two"),
|
||||
run_llm.InputItem(id="c", content="three")]
|
||||
run1: run_llm.Results = {
|
||||
"a": run_llm.BatchResult(id="a", text="same\n"),
|
||||
"b": run_llm.BatchResult(id="b", text="left"),
|
||||
"c": run_llm.BatchResult(id="c", text=" padded ")}
|
||||
run2: run_llm.Results = {
|
||||
"a": run_llm.BatchResult(id="a", text="same"),
|
||||
"b": run_llm.BatchResult(id="b", text="right"),
|
||||
"c": run_llm.BatchResult(id="c", text="padded")}
|
||||
agreed: list[str]
|
||||
disagreed: list[str]
|
||||
agreed, disagreed = run_llm.split_by_agreement(
|
||||
items, run1, run2)
|
||||
self.assertEqual(agreed, ["a", "c"])
|
||||
self.assertEqual(disagreed, ["b"])
|
||||
|
||||
|
||||
class TestFinalRecords(RunLLMTestCase):
|
||||
"""Test cases for the final record assembly."""
|
||||
|
||||
def test_build_final_records(self) -> None:
|
||||
"""Test assembling final records from runs and arbitration."""
|
||||
items: list[run_llm.InputItem] = [
|
||||
run_llm.InputItem(id="a", content="one"),
|
||||
run_llm.InputItem(id="b", content="two")]
|
||||
run1: run_llm.Results = {
|
||||
"a": run_llm.BatchResult(id="a", text="agreed text\n"),
|
||||
"b": run_llm.BatchResult(id="b", text="left")}
|
||||
arbitration: run_llm.Results = {
|
||||
"b": run_llm.BatchResult(id="b",
|
||||
text="arbitrated text")}
|
||||
records: list[dict[str, str]] = run_llm.build_final_records(
|
||||
items, run1, arbitration)
|
||||
self.assertEqual(records, [
|
||||
{"id": "a", "text": "agreed text", "source": "agreed"},
|
||||
{"id": "b", "text": "arbitrated text",
|
||||
"source": "arbitration"}])
|
||||
|
||||
|
||||
class TestCollectResults(RunLLMTestCase):
|
||||
"""Test cases for the batch result collection."""
|
||||
@@ -266,6 +210,20 @@ class TestCollectResults(RunLLMTestCase):
|
||||
run_llm.find_failures(["a", "b", "c"], results),
|
||||
["b", "c"])
|
||||
|
||||
def test_sum_usage(self) -> None:
|
||||
"""Test summing the token usage across results."""
|
||||
results: run_llm.Results = {
|
||||
"a": run_llm.BatchResult(
|
||||
id="a", text="fine",
|
||||
usage={"input_tokens": 10, "output_tokens": 5}),
|
||||
"b": run_llm.BatchResult(
|
||||
id="b", text="fine",
|
||||
usage={"input_tokens": 3, "output_tokens": 2}),
|
||||
"c": run_llm.BatchResult(id="c", error="errored")}
|
||||
self.assertEqual(
|
||||
run_llm.sum_usage(results),
|
||||
{"input_tokens": 13, "output_tokens": 7})
|
||||
|
||||
|
||||
class TestArchive(RunLLMTestCase):
|
||||
"""Test cases for the archive directory handling."""
|
||||
@@ -275,23 +233,43 @@ class TestArchive(RunLLMTestCase):
|
||||
self.__dir: Path = self._make_temp_dir()
|
||||
|
||||
def test_create_archive_dir(self) -> None:
|
||||
"""Test the archive directory naming and creation."""
|
||||
now: datetime = datetime(2026, 7, 30, 20, 5)
|
||||
directory: Path = run_llm.create_archive_dir(
|
||||
self.__dir, "coding", Path("prompts/gender_v3.md"), now)
|
||||
"""Test the archive directory creation."""
|
||||
target: Path = self.__dir / "01-01-tag" / "run1"
|
||||
directory: Path = run_llm.create_archive_dir(target, False)
|
||||
self.assertTrue(directory.is_dir())
|
||||
self.assertEqual(
|
||||
directory,
|
||||
self.__dir / "coding" / "20260730-2005-gender_v3")
|
||||
self.assertEqual(directory, target)
|
||||
|
||||
def test_existing_archive_dir_rejected(self) -> None:
|
||||
"""Test that an existing archive directory is rejected."""
|
||||
now: datetime = datetime(2026, 7, 30, 20, 5)
|
||||
run_llm.create_archive_dir(
|
||||
self.__dir, "coding", Path("gender_v3.md"), now)
|
||||
def test_existing_archive_dir_rejected_without_replace(
|
||||
self) -> None:
|
||||
"""Test that an existing archive is rejected by default."""
|
||||
target: Path = self.__dir / "01-01-tag" / "run1"
|
||||
run_llm.create_archive_dir(target, False)
|
||||
with self.assertRaises(FileExistsError):
|
||||
run_llm.create_archive_dir(
|
||||
self.__dir, "coding", Path("gender_v3.md"), now)
|
||||
run_llm.create_archive_dir(target, False)
|
||||
|
||||
def test_existing_archive_dir_replaced(self) -> None:
|
||||
"""Test that --replace replaces an existing archive."""
|
||||
target: Path = self.__dir / "01-01-tag" / "run1"
|
||||
first: Path = run_llm.create_archive_dir(target, False)
|
||||
(first / "stale.txt").write_text("stale", encoding="utf-8")
|
||||
second: Path = run_llm.create_archive_dir(target, True)
|
||||
self.assertEqual(first, second)
|
||||
self.assertFalse((second / "stale.txt").exists())
|
||||
|
||||
def test_replace_leaves_sibling_dir_untouched(self) -> None:
|
||||
"""Test that replacing run2 does not touch run1."""
|
||||
run1: Path = run_llm.create_archive_dir(
|
||||
self.__dir / "01-01-tag" / "run1", False)
|
||||
(run1 / "output.jsonl").write_text(
|
||||
"run1 data", encoding="utf-8")
|
||||
run2: Path = run_llm.create_archive_dir(
|
||||
self.__dir / "01-01-tag" / "run2", False)
|
||||
(run2 / "stale.jsonl").write_text("stale", encoding="utf-8")
|
||||
run_llm.create_archive_dir(
|
||||
self.__dir / "01-01-tag" / "run2", True)
|
||||
self.assertEqual(
|
||||
(run1 / "output.jsonl").read_text(encoding="utf-8"),
|
||||
"run1 data")
|
||||
|
||||
def test_write_jsonl(self) -> None:
|
||||
"""Test writing records as JSON Lines."""
|
||||
@@ -313,54 +291,38 @@ class TestMainFlow(RunLLMTestCase):
|
||||
"""Create a temporary directory with the input files."""
|
||||
directory: Path = self._make_temp_dir()
|
||||
self.__runs: Path = directory / "runs"
|
||||
prompt: Path = directory / "task_v1.md"
|
||||
prompt.write_text("The task prompt.\n", encoding="utf-8")
|
||||
arbitration: Path = directory / "task_arbitration_v1.md"
|
||||
arbitration.write_text(
|
||||
"The arbitration prompt.\n", encoding="utf-8")
|
||||
self.__archive_dir: Path = self.__runs / "task_v1" / "run1"
|
||||
self.__prompt: Path = directory / "task_v1.md"
|
||||
self.__prompt.write_text(
|
||||
"The task prompt.\n", encoding="utf-8")
|
||||
self.__input: Path = directory / "items.jsonl"
|
||||
self.__input.write_text(
|
||||
'{"id": "a", "content": "first item"}\n'
|
||||
'{"id": "b", "content": "second item"}\n',
|
||||
encoding="utf-8")
|
||||
self.__argv: list[str] = [
|
||||
str(self.__runs),
|
||||
"--prompt", str(prompt),
|
||||
"--arbitration-prompt", str(arbitration),
|
||||
"--input", str(self.__input),
|
||||
"--phase", "coding"]
|
||||
str(self.__prompt), str(self.__input),
|
||||
str(self.__archive_dir)]
|
||||
self.__settings: config.Settings = config.Settings(
|
||||
SQLALCHEMY_DATABASE_URL="sqlite://",
|
||||
ANTHROPIC_API_KEY="test-key")
|
||||
config.set_settings(self.__settings)
|
||||
|
||||
@staticmethod
|
||||
def __make_client(run1: list[Any], run2: list[Any],
|
||||
arbitration: list[Any] | None = None) \
|
||||
-> mock.Mock:
|
||||
"""Create a mock Anthropic client serving canned batch results.
|
||||
def __make_client(entries: list[Any]) -> mock.Mock:
|
||||
"""Create a mock Anthropic client serving canned results.
|
||||
|
||||
:param run1: The result entries of the run-1 batch.
|
||||
:param run2: The result entries of the run-2 batch.
|
||||
:param arbitration: The result entries of the arbitration
|
||||
batch.
|
||||
:param entries: The result entries of the single run batch.
|
||||
:return: The mock client.
|
||||
"""
|
||||
client: mock.Mock = mock.Mock()
|
||||
batch_ids: list[str] = ["batch_run1", "batch_run2",
|
||||
"batch_arb"]
|
||||
client.messages.batches.create.side_effect = [
|
||||
mock.Mock(id=x) for x in batch_ids]
|
||||
client.messages.batches.create.return_value = mock.Mock(
|
||||
id="batch_run1")
|
||||
ended: mock.Mock = mock.Mock()
|
||||
ended.processing_status = "ended"
|
||||
ended.ended_at = datetime(2026, 7, 30, 20, 0)
|
||||
client.messages.batches.retrieve.return_value = ended
|
||||
results: dict[str, list[Any]] = {
|
||||
"batch_run1": run1, "batch_run2": run2,
|
||||
"batch_arb": arbitration if arbitration is not None
|
||||
else []}
|
||||
client.messages.batches.results.side_effect = (
|
||||
lambda batch_id: iter(results[batch_id]))
|
||||
client.messages.batches.results.return_value = iter(entries)
|
||||
return client
|
||||
|
||||
@staticmethod
|
||||
@@ -380,16 +342,6 @@ class TestMainFlow(RunLLMTestCase):
|
||||
status: int = run_llm.main(argv)
|
||||
return status, stdout.getvalue()
|
||||
|
||||
def __archive_dir(self) -> Path:
|
||||
"""Locate the single archive directory of the run.
|
||||
|
||||
:return: The archive directory.
|
||||
"""
|
||||
directories: list[Path] = list(
|
||||
(self.__runs / "coding").iterdir())
|
||||
self.assertEqual(len(directories), 1)
|
||||
return directories[0]
|
||||
|
||||
def test_dry_run(self) -> None:
|
||||
"""Test the dry-run behavior."""
|
||||
status: int
|
||||
@@ -397,114 +349,107 @@ class TestMainFlow(RunLLMTestCase):
|
||||
status, stdout = self.__run_main(
|
||||
self.__argv + ["--dry-run"])
|
||||
self.assertEqual(status, 0)
|
||||
run_dir: Path = self.__archive_dir()
|
||||
run_dir: Path = self.__archive_dir
|
||||
self.assertEqual((run_dir / "prompt.md").read_text(
|
||||
encoding="utf-8"), "The task prompt.\n")
|
||||
self.assertEqual(
|
||||
(run_dir / "arbitration_prompt.md").read_text(
|
||||
encoding="utf-8"), "The arbitration prompt.\n")
|
||||
meta: dict[str, Any] = json.loads(
|
||||
(run_dir / "meta.json").read_text(encoding="utf-8"))
|
||||
self.assertTrue(meta["dry_run"])
|
||||
self.assertEqual(meta["item_count"], 2)
|
||||
self.assertEqual(meta["model"], "claude-sonnet-4-6")
|
||||
self.assertFalse((run_dir / "run1.jsonl").exists())
|
||||
self.assertNotIn("run", meta)
|
||||
self.assertFalse((run_dir / "output.jsonl").exists())
|
||||
request: dict[str, Any] = json.loads(stdout)
|
||||
self.assertEqual(request["custom_id"], "a")
|
||||
self.assertEqual(request["params"]["system"],
|
||||
"The task prompt.\n")
|
||||
|
||||
def test_all_agreed_skips_arbitration(self) -> None:
|
||||
"""Test that full agreement skips the arbitration batch."""
|
||||
def test_run_produces_output_file(self) -> None:
|
||||
"""Test that a run submits one batch and writes output."""
|
||||
client: mock.Mock = self.__make_client(
|
||||
run1=[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b")],
|
||||
run2=[self._make_success_entry("a", "answer a\n"),
|
||||
self._make_success_entry("b", "answer b")])
|
||||
[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b")])
|
||||
status: int = self.__run_main(self.__argv, client)[0]
|
||||
self.assertEqual(status, 0)
|
||||
self.assertEqual(
|
||||
client.messages.batches.create.call_count, 2)
|
||||
run_dir: Path = self.__archive_dir()
|
||||
self.assertEqual(
|
||||
(run_dir / "arbitration.jsonl").read_text(
|
||||
encoding="utf-8"), "")
|
||||
final: list[dict[str, Any]] = [
|
||||
json.loads(x) for x in (run_dir / "final.jsonl")
|
||||
client.messages.batches.create.call_count, 1)
|
||||
run_dir: Path = self.__archive_dir
|
||||
self.assertTrue((run_dir / "output.jsonl").exists())
|
||||
output: list[dict[str, Any]] = [
|
||||
json.loads(x) for x in (run_dir / "output.jsonl")
|
||||
.read_text(encoding="utf-8").splitlines()]
|
||||
self.assertEqual(final, [
|
||||
{"id": "a", "text": "answer a", "source": "agreed"},
|
||||
{"id": "b", "text": "answer b", "source": "agreed"}])
|
||||
self.assertEqual(output[0]["text"], "answer a")
|
||||
meta: dict[str, Any] = json.loads(
|
||||
(run_dir / "meta.json").read_text(encoding="utf-8"))
|
||||
self.assertEqual(meta["agreed_count"], 2)
|
||||
self.assertEqual(meta["agreement_rate"], 1.0)
|
||||
self.assertNotIn("arbitration", meta["batches"])
|
||||
self.assertNotIn("run", meta)
|
||||
self.assertEqual(meta["batch"]["batch_id"], "batch_run1")
|
||||
self.assertEqual(meta["usage"],
|
||||
{"input_tokens": 20, "output_tokens": 10})
|
||||
|
||||
def test_disagreement_triggers_arbitration(self) -> None:
|
||||
"""Test that disagreeing items go through arbitration."""
|
||||
def test_existing_archive_rejected_without_replace(self) -> None:
|
||||
"""Test that an existing archive without --replace fails."""
|
||||
run_dir: Path = self.__archive_dir
|
||||
run_dir.mkdir(parents=True)
|
||||
(run_dir / "stale.jsonl").write_text(
|
||||
"stale", encoding="utf-8")
|
||||
status: int = self.__run_main(
|
||||
self.__argv + ["--dry-run"])[0]
|
||||
self.assertEqual(status, 1)
|
||||
self.assertTrue((run_dir / "stale.jsonl").exists())
|
||||
|
||||
def test_rerun_replaces_archive_with_flag(self) -> None:
|
||||
"""Test that --replace replaces a pre-existing archive."""
|
||||
run_dir: Path = self.__archive_dir
|
||||
run_dir.mkdir(parents=True)
|
||||
(run_dir / "stale.jsonl").write_text(
|
||||
"stale", encoding="utf-8")
|
||||
client: mock.Mock = self.__make_client(
|
||||
run1=[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b1")],
|
||||
run2=[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b2")],
|
||||
arbitration=[
|
||||
self._make_success_entry("b", "answer b final")])
|
||||
status: int = self.__run_main(self.__argv, client)[0]
|
||||
[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b")])
|
||||
status: int = self.__run_main(
|
||||
self.__argv + ["--replace"], client)[0]
|
||||
self.assertEqual(status, 0)
|
||||
self.assertFalse((run_dir / "stale.jsonl").exists())
|
||||
self.assertTrue((run_dir / "output.jsonl").exists())
|
||||
|
||||
def test_replace_leaves_sibling_dir_untouched(self) -> None:
|
||||
"""Test that replacing run2 does not affect run1's archive."""
|
||||
run1_dir: Path = self.__archive_dir
|
||||
run1_dir.mkdir(parents=True)
|
||||
(run1_dir / "output.jsonl").write_text(
|
||||
"run1 data", encoding="utf-8")
|
||||
run2_dir: Path = self.__runs / "task_v1" / "run2"
|
||||
run2_dir.mkdir(parents=True)
|
||||
(run2_dir / "output.jsonl").write_text(
|
||||
"stale run2 data", encoding="utf-8")
|
||||
client: mock.Mock = self.__make_client(
|
||||
[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b")])
|
||||
argv: list[str] = [
|
||||
str(self.__prompt), str(self.__input),
|
||||
str(run2_dir), "--replace"]
|
||||
status: int = self.__run_main(argv, client)[0]
|
||||
self.assertEqual(status, 0)
|
||||
self.assertEqual(
|
||||
client.messages.batches.create.call_count, 3)
|
||||
arb_call: mock.call = \
|
||||
client.messages.batches.create.call_args_list[2]
|
||||
arb_requests: list[dict[str, Any]] = \
|
||||
arb_call.kwargs["requests"]
|
||||
self.assertEqual(len(arb_requests), 1)
|
||||
self.assertEqual(arb_requests[0]["custom_id"], "b")
|
||||
self.assertEqual(arb_requests[0]["params"]["system"],
|
||||
"The arbitration prompt.\n")
|
||||
self.assertEqual(
|
||||
arb_requests[0]["params"]["messages"][0]["content"],
|
||||
"<item>\nsecond item\n</item>\n"
|
||||
"<run1>\nanswer b1\n</run1>\n"
|
||||
"<run2>\nanswer b2\n</run2>")
|
||||
run_dir: Path = self.__archive_dir()
|
||||
arb_lines: list[str] = (run_dir / "arbitration.jsonl") \
|
||||
.read_text(encoding="utf-8").splitlines()
|
||||
self.assertEqual(len(arb_lines), 1)
|
||||
self.assertEqual(json.loads(arb_lines[0])["text"],
|
||||
"answer b final")
|
||||
final: list[dict[str, Any]] = [
|
||||
json.loads(x) for x in (run_dir / "final.jsonl")
|
||||
.read_text(encoding="utf-8").splitlines()]
|
||||
self.assertEqual(final, [
|
||||
{"id": "a", "text": "answer a", "source": "agreed"},
|
||||
{"id": "b", "text": "answer b final",
|
||||
"source": "arbitration"}])
|
||||
meta: dict[str, Any] = json.loads(
|
||||
(run_dir / "meta.json").read_text(encoding="utf-8"))
|
||||
self.assertEqual(meta["agreed_count"], 1)
|
||||
self.assertEqual(meta["agreement_rate"], 0.5)
|
||||
self.assertEqual(meta["batches"]["arbitration"]["batch_id"],
|
||||
"batch_arb")
|
||||
(run1_dir / "output.jsonl").read_text(encoding="utf-8"),
|
||||
"run1 data")
|
||||
self.assertNotEqual(
|
||||
(run2_dir / "output.jsonl").read_text(encoding="utf-8"),
|
||||
"stale run2 data")
|
||||
|
||||
def test_run_failure_exits_non_zero(self) -> None:
|
||||
"""Test that a failed item aborts with a non-zero status."""
|
||||
client: mock.Mock = self.__make_client(
|
||||
run1=[self._make_success_entry("a", "answer a"),
|
||||
self._make_error_entry(
|
||||
"b", "invalid_request_error")],
|
||||
run2=[self._make_success_entry("a", "answer a"),
|
||||
self._make_success_entry("b", "answer b")])
|
||||
[self._make_success_entry("a", "answer a"),
|
||||
self._make_error_entry("b", "invalid_request_error")])
|
||||
status: int = self.__run_main(self.__argv, client)[0]
|
||||
self.assertEqual(status, 1)
|
||||
run_dir: Path = self.__archive_dir()
|
||||
self.assertTrue((run_dir / "run1.jsonl").exists())
|
||||
self.assertTrue((run_dir / "run2.jsonl").exists())
|
||||
run_dir: Path = self.__archive_dir
|
||||
self.assertTrue((run_dir / "output.jsonl").exists())
|
||||
self.assertTrue((run_dir / "meta.json").exists())
|
||||
self.assertFalse((run_dir / "final.jsonl").exists())
|
||||
run1_lines: list[str] = (run_dir / "run1.jsonl") \
|
||||
output_lines: list[str] = (run_dir / "output.jsonl") \
|
||||
.read_text(encoding="utf-8").splitlines()
|
||||
self.assertEqual(json.loads(run1_lines[1]),
|
||||
self.assertEqual(json.loads(output_lines[1]),
|
||||
{"id": "b",
|
||||
"error": "invalid_request_error"})
|
||||
|
||||
|
||||
Reference in New Issue
Block a user