diff --git a/README.md b/README.md index 4806593..5c71b79 100644 --- a/README.md +++ b/README.md @@ -89,6 +89,10 @@ rejects any polished turn that alters a number, dollar amount, or percentage (or balloons the text), falling back to the deterministic version and telling you how many turns it kept. Captions are never sent to the model. +The polish uses `claude-sonnet-5-5` at low effort by default. Pick another +model with `--model`; it must accept the effort parameter (current Opus and +Sonnet models do, Haiku 4.5 does not). + ## Develop ```bash diff --git a/pyproject.toml b/pyproject.toml index e81c801..9b1f29e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -10,8 +10,8 @@ keywords = ["transcript", "captions", "subtitles", "vtt", "srt", "zoom", "youtub dependencies = ["pyyaml>=6.0"] [project.optional-dependencies] -llm = ["anthropic>=0.40"] -dev = ["pytest>=8.0"] +llm = ["anthropic>=1.0"] +dev = ["pytest>=8.0", "anthropic>=1.0"] [project.scripts] clean-transcript = "transcript_tools.cli:main" diff --git a/tests/test_cli.py b/tests/test_cli.py new file mode 100644 index 0000000..4c46fe4 --- /dev/null +++ b/tests/test_cli.py @@ -0,0 +1,10 @@ +from transcript_tools.cli import _build_parser + + +def test_default_llm_model(): + assert _build_parser().parse_args(["x.vtt"]).model == "claude-sonnet-5-5" + + +def test_model_override(): + args = _build_parser().parse_args(["x.vtt", "--model", "claude-opus-5-5"]) + assert args.model == "claude-opus-5-5" diff --git a/tests/test_llm.py b/tests/test_llm.py index 9804cc8..e2fbc03 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -1,8 +1,14 @@ +import json +from types import SimpleNamespace + +import pytest + from transcript_tools.llm import ( guard, numbers, numbers_preserved, polish_paragraphs, + polish_turn, ) @@ -12,7 +18,9 @@ def test_numbers_extraction(): def test_numbers_preserved(): - assert numbers_preserved("it was 80.3% of $2,000", "It was 80.3% of $2,000 exactly.") + assert numbers_preserved( + "it was 80.3% of $2,000", "It was 80.3% of $2,000 exactly." + ) assert not numbers_preserved("80.3%", "80.4%") @@ -61,3 +69,120 @@ def boom(_): out = polish_paragraphs(paras, model="x", polish_fn=boom) assert out == paras # original preserved on error + + +class _FakeMessages: + def __init__(self, resp): + self.resp = resp + self.calls = [] + + def create(self, **kwargs): + self.calls.append(kwargs) + return self.resp + + +def _fake_client(stop_reason="end_turn", text="We got 80.3%."): + content = [ + SimpleNamespace(type="thinking", thinking="", signature="sig"), + SimpleNamespace(type="text", text=text), + ] + resp = SimpleNamespace(stop_reason=stop_reason, content=content) + return SimpleNamespace(messages=_FakeMessages(resp)) + + +def test_polish_turn_request_shape(): + client = _fake_client() + out = polish_turn("um we got 80.3%", model="claude-sonnet-5-5", client=client) + assert out == "We got 80.3%." # thinking blocks are skipped + (kwargs,) = client.messages.calls + assert kwargs["model"] == "claude-sonnet-5-5" + assert kwargs["max_tokens"] >= 16000 + assert kwargs["output_config"] == {"effort": "low"} + # Rejected (400) or unsupported on current Claude models. + for banned in ("temperature", "top_p", "top_k", "thinking", "tool_choice"): + assert banned not in kwargs + # No assistant prefill: the only message is the user's turn. + assert [m["role"] for m in kwargs["messages"]] == ["user"] + + +# Every stop_reason in anthropic.types.StopReason (anthropic 1.0-1.9). +STOP_REASONS = [ + "end_turn", + "max_tokens", + "stop_sequence", + "tool_use", + "pause_turn", + "refusal", + "model_context_window_exceeded", +] + + +@pytest.mark.parametrize("stop_reason", STOP_REASONS) +def test_only_complete_responses_are_used(stop_reason): + paras = [("Max", "um we got, 80.3% on $2,000")] + client = _fake_client(stop_reason=stop_reason, text="We got 80.3% on $2,000.") + skipped = [] + out = polish_paragraphs( + paras, + model="claude-sonnet-5-5", + polish_fn=lambda t: polish_turn(t, model="claude-sonnet-5-5", client=client), + on_skip=lambda s, w: skipped.append((s, w)), + ) + if stop_reason == "end_turn": + assert out == [("Max", "We got 80.3% on $2,000.")] + assert not skipped + else: + # A refusal or cut-off keeps the deterministic turn, even when the + # partial text would pass the number guard. + assert out == paras + assert skipped == [ + ("Max", f"error: incomplete response (stop_reason={stop_reason})") + ] + + +@pytest.mark.parametrize("stop_reason", ["end_turn", "refusal"]) +def test_polish_turn_against_real_sdk(stop_reason): + """The installed SDK sends the request shape and parses the reply (no network).""" + anthropic = pytest.importorskip("anthropic") + import httpx2 # the HTTP layer of anthropic>=1.0 + + sent = [] + + def handler(request): + sent.append(json.loads(request.content)) + return httpx2.Response( + 200, + json={ + "id": "msg_test", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-5-5", + "content": [ + {"type": "thinking", "thinking": "", "signature": "sig"}, + {"type": "text", "text": "We got 80.3%."}, + ], + "stop_reason": stop_reason, + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 5}, + }, + ) + + client = anthropic.Anthropic( + api_key="test-key", + max_retries=0, + http_client=anthropic.DefaultHttpxClient( + transport=httpx2.MockTransport(handler) + ), + ) + if stop_reason == "end_turn": + out = polish_turn("um we got 80.3%", model="claude-sonnet-5-5", client=client) + assert out == "We got 80.3%." + else: + with pytest.raises(RuntimeError, match="stop_reason=refusal"): + polish_turn("um we got 80.3%", model="claude-sonnet-5-5", client=client) + (body,) = sent + assert body["model"] == "claude-sonnet-5-5" + assert body["max_tokens"] == 16000 + assert body["output_config"] == {"effort": "low"} + assert not {"temperature", "top_p", "top_k", "thinking", "tool_choice"} & set(body) + assert [m["role"] for m in body["messages"]] == ["user"] diff --git a/transcript_tools/cli.py b/transcript_tools/cli.py index 10b35c2..defa4da 100644 --- a/transcript_tools/cli.py +++ b/transcript_tools/cli.py @@ -56,8 +56,9 @@ def _build_parser() -> argparse.ArgumentParser: ) p.add_argument( "--model", - default="claude-haiku-4-5-20251001", - help="model id for --llm (default: %(default)s)", + default="claude-sonnet-5-5", + help="model id for --llm; must accept the effort parameter " + "(default: %(default)s)", ) return p diff --git a/transcript_tools/llm.py b/transcript_tools/llm.py index d66c47b..52b7893 100644 --- a/transcript_tools/llm.py +++ b/transcript_tools/llm.py @@ -45,13 +45,23 @@ def _client(api_key: str | None): def polish_turn(text: str, *, model: str, client) -> str: - """Polish one turn via the Anthropic API. Returns raw model text.""" + """Polish one turn via the Anthropic API. Returns raw model text. + + Raises if the response is incomplete (a refusal, or a ``max_tokens`` + cut-off): partial text could still pass the number guard, so the caller + must keep the deterministic turn instead. + """ resp = client.messages.create( model=model, - max_tokens=4096, + # Thinking counts toward max_tokens, so leave room beyond the reply. + max_tokens=16000, + # Per-turn cleanup: low effort skips thinking on most turns. + output_config={"effort": "low"}, system=_SYSTEM, messages=[{"role": "user", "content": text}], ) + if resp.stop_reason != "end_turn": + raise RuntimeError(f"incomplete response (stop_reason={resp.stop_reason})") parts = [b.text for b in resp.content if getattr(b, "type", None) == "text"] return "".join(parts).strip()