diff --git a/src/pact/backends/anthropic.py b/src/pact/backends/anthropic.py index c1cc1c1..a4d5fe7 100644 --- a/src/pact/backends/anthropic.py +++ b/src/pact/backends/anthropic.py @@ -1,6 +1,14 @@ """Anthropic backend — direct API calls with tool_choice schema enforcement. Reused from swarm with import path adaptation. + +Newer models (Claude Opus 5.5, Sonnet 5.5, Fable 5.1) reject forced tool +use with a 400. The backend starts with a forced tool_choice and, on that +400, switches the instance to tool_choice "auto" plus a system-prompt +instruction naming the tool. Strict tool use is not an option here: several +Pact schemas have dict[str, X] fields, and strict mode only accepts +additionalProperties: false. Pydantic validation and the correction retry +loop keep the output schema-valid in both modes. """ from __future__ import annotations @@ -24,13 +32,21 @@ "claude-opus-4-6": 32768, "claude-sonnet-4-5-20250929": 64000, "claude-haiku-4-5-20251001": 8192, + "claude-opus-5-5": 128000, } _DEFAULT_MAX_TOKENS_CAP = 32768 +_TOOL_CHOICE_UNSUPPORTED_MARKER = "tool_choice" + + class AnthropicBackend: """Backend using the Anthropic API with tool_choice for structured extraction.""" + # Flipped to False per instance the first time the model rejects a + # forced tool_choice; reset by set_model(). + _forced_tool_choice: bool = True + def __init__(self, budget: BudgetTracker, model: str = "claude-opus-4-6") -> None: try: import anthropic @@ -53,6 +69,7 @@ def __init__(self, budget: BudgetTracker, model: str = "claude-opus-4-6") -> Non def set_model(self, model: str) -> None: self._model = model + self._forced_tool_choice = True def _max_tokens_cap(self) -> int: return _MODEL_MAX_TOKENS.get(self._model, _DEFAULT_MAX_TOKENS_CAP) @@ -74,15 +91,22 @@ async def assess( cap = self._max_tokens_cap() current_max = min(max_tokens, cap) last_error: ValidationError | None = None + missed_tool_call = False for attempt in range(3): - # On retry after validation error, augment prompt with correction + # On retry, augment prompt with the correction for the last failure effective_prompt = prompt if last_error is not None: correction = self._format_validation_correction(last_error) effective_prompt = f"{prompt}\n\n{correction}" logger.info("Retrying %s with validation feedback (attempt %d)", schema.__name__, attempt + 1) + elif missed_tool_call: + effective_prompt = ( + f"{prompt}\n\n{self._missed_tool_call_correction(schema.__name__)}" + ) + logger.info("Retrying %s after response without a tool call (attempt %d)", + schema.__name__, attempt + 1) raw_input, stop_reason, in_tok, out_tok = await self._call_llm( schema, effective_prompt, system, current_max, @@ -90,10 +114,8 @@ async def assess( total_in += in_tok total_out += out_tok - if raw_input is None: - raise RuntimeError( - f"No tool_use block found for {schema.__name__}" - ) + if stop_reason == "refusal": + raise RuntimeError(f"Model refused to produce {schema.__name__}") if stop_reason == "max_tokens" and attempt < 2: new_max = min(current_max * 2, cap) @@ -101,6 +123,18 @@ async def assess( current_max = new_max continue + if raw_input is None: + # Under tool_choice "auto" the model can answer in text + # without calling the tool. + if attempt < 2: + missed_tool_call = True + last_error = None + continue + raise RuntimeError( + f"No tool_use block found for {schema.__name__}" + ) + missed_tool_call = False + raw_input = self._coerce_fields(raw_input) try: @@ -117,6 +151,22 @@ async def assess( raise RuntimeError(f"Failed to get valid {schema.__name__} after 3 attempts") + @staticmethod + def _missed_tool_call_correction(tool_name: str) -> str: + return ( + "IMPORTANT: Your previous response did not call the " + f"`{tool_name}` tool. Respond by calling `{tool_name}` exactly " + "once with your complete answer as its input." + ) + + @staticmethod + def _tool_instruction(tool_name: str) -> str: + return ( + f"Deliver your answer by calling the `{tool_name}` tool exactly " + "once, with the complete result as its input. Do not answer in " + "plain text." + ) + @staticmethod def _format_validation_correction(error: ValidationError) -> str: """Format a Pydantic ValidationError into a clear correction prompt. @@ -267,28 +317,78 @@ async def _stream_with_stall_detection( except Exception: pass - async with self._client.messages.stream( - model=self._model, - max_tokens=max_tokens, - system=system, - messages=[{"role": "user", "content": prompt}], - tools=[{ - "name": tool_name, - "description": tool_description, - "input_schema": tool_schema, - }], - tool_choice={"type": "tool", "name": tool_name}, - ) as stream: - aiter = stream.__aiter__() - while True: - try: - await asyncio.wait_for(aiter.__anext__(), timeout=stall_timeout) - except StopAsyncIteration: - break - except asyncio.TimeoutError: - raise asyncio.TimeoutError() + return await self._stream_tool_call( + tool_name, tool_schema, tool_description, + system, prompt, max_tokens, stall_timeout, + ) - return await stream.get_final_message() + async def _stream_tool_call( + self, + tool_name: str, + tool_schema: dict, + tool_description: str, + system: str | list[dict], + user_content: str | list[dict], + max_tokens: int, + stall_timeout: float, + ): + """Stream one tool-extraction request and return the final message. + + Sends a forced tool_choice until the model rejects it, then retries + with tool_choice "auto" and a system instruction naming the tool. + Raises asyncio.TimeoutError if no event arrives within stall_timeout. + """ + import anthropic + + tool = { + "name": tool_name, + "description": tool_description, + "input_schema": tool_schema, + } + # Request-local: a concurrent request may flip the instance flag + # while this one is in flight. + forced = self._forced_tool_choice + while True: + if forced: + request_system = system + tool_choice = {"type": "tool", "name": tool_name} + else: + instruction = self._tool_instruction(tool_name) + if isinstance(system, str): + request_system = f"{system}\n\n{instruction}" + else: + request_system = [*system, {"type": "text", "text": instruction}] + tool_choice = {"type": "auto", "disable_parallel_tool_use": True} + + try: + async with self._client.messages.stream( + model=self._model, + max_tokens=max_tokens, + system=request_system, + messages=[{"role": "user", "content": user_content}], + tools=[tool], + tool_choice=tool_choice, + ) as stream: + aiter = stream.__aiter__() + while True: + try: + await asyncio.wait_for(aiter.__anext__(), timeout=stall_timeout) + except StopAsyncIteration: + break + except asyncio.TimeoutError: + raise asyncio.TimeoutError() + + return await stream.get_final_message() + except anthropic.BadRequestError as exc: + if not forced or _TOOL_CHOICE_UNSUPPORTED_MARKER not in str(exc): + raise + if self._forced_tool_choice: + logger.info( + "%s rejected forced tool_choice; using tool_choice auto", + self._model, + ) + self._forced_tool_choice = False + forced = False # ── Prompt caching helpers ────────────────────────────────────────── @@ -360,28 +460,10 @@ async def _call_llm_cached( user_content = self._build_user_blocks(cache_prefix, prompt) try: - async with self._client.messages.stream( - model=self._model, - max_tokens=max_tokens, - system=system_blocks, - messages=[{"role": "user", "content": user_content}], - tools=[{ - "name": tool_name, - "description": schema.__doc__ or f"Extract {tool_name}", - "input_schema": tool_schema, - }], - tool_choice={"type": "tool", "name": tool_name}, - ) as stream: - aiter = stream.__aiter__() - while True: - try: - await asyncio.wait_for(aiter.__anext__(), timeout=stall_timeout) - except StopAsyncIteration: - break - except asyncio.TimeoutError: - raise asyncio.TimeoutError() - - message = await stream.get_final_message() + message = await self._stream_tool_call( + tool_name, tool_schema, schema.__doc__ or f"Extract {tool_name}", + system_blocks, user_content, max_tokens, stall_timeout, + ) except asyncio.TimeoutError: logger.error( "Anthropic API stalled (no progress for %.0fs) for %s", @@ -430,6 +512,7 @@ async def assess_with_cache( cap = self._max_tokens_cap() current_max = min(max_tokens, cap) last_error: ValidationError | None = None + missed_tool_call = False for attempt in range(3): effective_prompt = prompt @@ -438,6 +521,12 @@ async def assess_with_cache( effective_prompt = f"{prompt}\n\n{correction}" logger.info("Retrying %s with validation feedback (attempt %d)", schema.__name__, attempt + 1) + elif missed_tool_call: + effective_prompt = ( + f"{prompt}\n\n{self._missed_tool_call_correction(schema.__name__)}" + ) + logger.info("Retrying %s after response without a tool call (attempt %d)", + schema.__name__, attempt + 1) raw_input, stop_reason, in_tok, out_tok = await self._call_llm_cached( schema, effective_prompt, system, cache_prefix, current_max, @@ -445,10 +534,8 @@ async def assess_with_cache( total_in += in_tok total_out += out_tok - if raw_input is None: - raise RuntimeError( - f"No tool_use block found for {schema.__name__}" - ) + if stop_reason == "refusal": + raise RuntimeError(f"Model refused to produce {schema.__name__}") if stop_reason == "max_tokens" and attempt < 2: new_max = min(current_max * 2, cap) @@ -456,6 +543,18 @@ async def assess_with_cache( current_max = new_max continue + if raw_input is None: + # Under tool_choice "auto" the model can answer in text + # without calling the tool. + if attempt < 2: + missed_tool_call = True + last_error = None + continue + raise RuntimeError( + f"No tool_use block found for {schema.__name__}" + ) + missed_tool_call = False + raw_input = self._coerce_fields(raw_input) try: diff --git a/src/pact/budget.py b/src/pact/budget.py index 2c2fb0a..82cca0d 100644 --- a/src/pact/budget.py +++ b/src/pact/budget.py @@ -26,22 +26,32 @@ "claude-sonnet-4": (3.00, 15.00), "claude-sonnet-4-5-20250929": (3.00, 15.00), "claude-sonnet-4-6": (3.00, 15.00), + "claude-sonnet-5": (2.00, 10.00), + "claude-sonnet-5-5": (2.00, 10.00), "claude-opus-4": (15.00, 75.00), "claude-opus-4-1": (15.00, 75.00), "claude-opus-4-5": (5.00, 25.00), "claude-opus-4-6": (5.00, 25.00), - # OpenAI + "claude-opus-4-7": (5.00, 25.00), + "claude-opus-4-8": (5.00, 25.00), + "claude-opus-5": (5.00, 25.00), + "claude-opus-5-5": (4.00, 20.00), + "claude-fable-5": (10.00, 50.00), + "claude-fable-5-1": (10.00, 50.00), + "claude-mythos-5": (10.00, 50.00), + "claude-mythos-5-1": (10.00, 50.00), + # OpenAI — https://developers.openai.com/api/docs/pricing "gpt-4o": (2.50, 10.00), "gpt-4o-mini": (0.15, 0.60), "gpt-4-turbo": (10.00, 30.00), - "o3": (10.00, 40.00), + "o3": (2.00, 8.00), "o3-mini": (1.10, 4.40), - # Google Gemini + # Google Gemini — https://ai.google.dev/gemini-api/docs/pricing "gemini-2.5-pro": (1.25, 10.00), - "gemini-2.5-flash": (0.15, 0.60), - "gemini-2.5-flash-lite": (0.075, 0.30), - "gemini-3-pro-preview": (1.25, 10.00), - "gemini-3-flash-preview": (0.15, 0.60), + "gemini-2.5-flash": (0.30, 2.50), + "gemini-2.5-flash-lite": (0.10, 0.40), + "gemini-3.1-pro-preview": (2.00, 12.00), + "gemini-3-flash-preview": (0.50, 3.00), } # Active pricing table — starts as defaults, can be updated diff --git a/tests/test_anthropic_tool_choice.py b/tests/test_anthropic_tool_choice.py new file mode 100644 index 0000000..5913c33 --- /dev/null +++ b/tests/test_anthropic_tool_choice.py @@ -0,0 +1,223 @@ +"""Tests for AnthropicBackend's tool_choice fallback. + +Claude Opus 5.5 (and Sonnet 5.5, Fable 5.1) return a 400 for a forced +tool_choice. The backend switches to tool_choice "auto" plus a system +instruction and retries when a response comes back without a tool call. +""" + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import anthropic +import pytest +from pydantic import BaseModel + +from pact.backends.anthropic import AnthropicBackend +from pact.budget import BudgetTracker + + +class SimpleSchema(BaseModel): + """Simple test schema.""" + name: str + value: int + + +def _make_backend() -> AnthropicBackend: + backend = AnthropicBackend.__new__(AnthropicBackend) + backend._model = "claude-opus-5-5" + budget = BudgetTracker(per_project_cap=100.0) + budget.set_model_pricing("claude-opus-5-5") + backend._budget = budget + backend._client = MagicMock() + return backend + + +def _bad_request(message: str) -> anthropic.BadRequestError: + response = MagicMock(status_code=400, headers={}) + return anthropic.BadRequestError(message, response=response, body=None) + + +class _FakeStream: + def __init__(self, message): + self._message = message + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc): + return False + + def __aiter__(self): + return self + + async def __anext__(self): + raise StopAsyncIteration + + async def get_final_message(self): + return self._message + + +def _tool_message(tool_name: str, tool_input: dict): + return SimpleNamespace( + content=[SimpleNamespace(type="tool_use", name=tool_name, input=tool_input)], + stop_reason="tool_use", + usage=SimpleNamespace( + input_tokens=10, output_tokens=5, + cache_creation_input_tokens=0, cache_read_input_tokens=0, + ), + ) + + +def _stream_factory(outcomes: list): + """Return a messages.stream stand-in that raises or yields each outcome in turn.""" + outcomes = list(outcomes) + + def stream(**kwargs): + outcome = outcomes.pop(0) + if isinstance(outcome, Exception): + raise outcome + return _FakeStream(outcome) + + return MagicMock(side_effect=stream) + + +_TOOL_CHOICE_400 = _bad_request( + 'tool_choice: type "tool" and "any" are not supported for this model.' +) + + +class TestForcedToolChoiceFallback: + async def test_forced_rejection_retries_with_auto(self): + backend = _make_backend() + backend._client.messages.stream = _stream_factory([ + _TOOL_CHOICE_400, + _tool_message("SimpleSchema", {"name": "a", "value": 1}), + ]) + + result, _, _ = await backend.assess(SimpleSchema, "prompt", "system prompt") + + assert result.value == 1 + calls = backend._client.messages.stream.call_args_list + assert calls[0].kwargs["tool_choice"] == {"type": "tool", "name": "SimpleSchema"} + assert calls[1].kwargs["tool_choice"]["type"] == "auto" + assert calls[1].kwargs["system"].startswith("system prompt") + assert "`SimpleSchema` tool" in calls[1].kwargs["system"] + + async def test_auto_mode_persists_for_later_calls(self): + backend = _make_backend() + backend._client.messages.stream = _stream_factory([ + _TOOL_CHOICE_400, + _tool_message("SimpleSchema", {"name": "a", "value": 1}), + _tool_message("SimpleSchema", {"name": "b", "value": 2}), + ]) + + await backend.assess(SimpleSchema, "prompt", "system") + await backend.assess(SimpleSchema, "prompt", "system") + + calls = backend._client.messages.stream.call_args_list + assert len(calls) == 3 + assert calls[2].kwargs["tool_choice"]["type"] == "auto" + + async def test_set_model_restores_forced_tool_choice(self): + backend = _make_backend() + backend._forced_tool_choice = False + backend.set_model("claude-opus-4-8") + assert backend._forced_tool_choice is True + + async def test_cached_path_appends_instruction_block(self): + backend = _make_backend() + backend._client.messages.stream = _stream_factory([ + _TOOL_CHOICE_400, + _tool_message("SimpleSchema", {"name": "a", "value": 1}), + ]) + + await backend.assess_with_cache( + SimpleSchema, "prompt", "system prompt", cache_prefix="x" * 400, + ) + + system = backend._client.messages.stream.call_args_list[1].kwargs["system"] + assert system[0]["text"] == "system prompt" + assert system[0]["cache_control"] == {"type": "ephemeral"} + assert "`SimpleSchema` tool" in system[-1]["text"] + + async def test_other_bad_request_is_raised(self): + backend = _make_backend() + backend._client.messages.stream = _stream_factory([ + _bad_request("max_tokens: too large"), + ]) + + with pytest.raises(anthropic.BadRequestError): + await backend.assess(SimpleSchema, "prompt", "system") + assert backend._forced_tool_choice is True + + async def test_concurrent_forced_rejections_both_fall_back(self): + # Both requests go out forced. The first 400 flips the instance to + # auto before the second 400 lands; the second must still retry. + backend = _make_backend() + both_sent = asyncio.Event() + forced_sent = 0 + + class _RejectingStream(_FakeStream): + async def __aenter__(self): + nonlocal forced_sent + forced_sent += 1 + if forced_sent == 2: + both_sent.set() + await both_sent.wait() + raise _TOOL_CHOICE_400 + + def stream(**kwargs): + if kwargs["tool_choice"]["type"] == "tool": + return _RejectingStream(None) + return _FakeStream(_tool_message("SimpleSchema", {"name": "a", "value": 1})) + + backend._client.messages.stream = MagicMock(side_effect=stream) + + results = await asyncio.gather( + backend.assess(SimpleSchema, "prompt", "system"), + backend.assess(SimpleSchema, "prompt", "system"), + ) + + assert [r.value for r, _, _ in results] == [1, 1] + assert forced_sent == 2 + assert backend._forced_tool_choice is False + + +class TestMissedToolCall: + async def test_retries_with_correction_when_no_tool_call(self): + backend = _make_backend() + backend._call_llm = AsyncMock(side_effect=[ + (None, "end_turn", 100, 50), + ({"name": "a", "value": 1}, "tool_use", 100, 50), + ]) + + result, in_tok, _ = await backend.assess(SimpleSchema, "prompt", "system") + + assert result.value == 1 + assert in_tok == 200 + retry_prompt = backend._call_llm.call_args_list[1][0][1] + assert retry_prompt.startswith("prompt") + assert "did not call the `SimpleSchema` tool" in retry_prompt + + async def test_cached_retries_with_correction_when_no_tool_call(self): + backend = _make_backend() + backend._call_llm_cached = AsyncMock(side_effect=[ + (None, "end_turn", 100, 50), + ({"name": "a", "value": 1}, "tool_use", 100, 50), + ]) + + result, _, _ = await backend.assess_with_cache(SimpleSchema, "prompt", "system") + + assert result.value == 1 + assert "did not call" in backend._call_llm_cached.call_args_list[1][0][1] + + async def test_refusal_raises_without_retry(self): + backend = _make_backend() + backend._call_llm = AsyncMock(return_value=(None, "refusal", 100, 0)) + + with pytest.raises(RuntimeError, match="refused"): + await backend.assess(SimpleSchema, "prompt", "system") + assert backend._call_llm.call_count == 1 diff --git a/tests/test_budget.py b/tests/test_budget.py index b881d22..9846fca 100644 --- a/tests/test_budget.py +++ b/tests/test_budget.py @@ -15,6 +15,46 @@ def test_exact_match(self): assert inp == 5.00 assert out == 25.00 + def test_opus_4_8(self): + inp, out = pricing_for_model("claude-opus-4-8") + assert inp == 5.00 + assert out == 25.00 + + def test_opus_5(self): + inp, out = pricing_for_model("claude-opus-5") + assert inp == 5.00 + assert out == 25.00 + + def test_opus_4_7(self): + inp, out = pricing_for_model("claude-opus-4-7") + assert inp == 5.00 + assert out == 25.00 + + def test_sonnet_5_5(self): + inp, out = pricing_for_model("claude-sonnet-5-5") + assert inp == 2.00 + assert out == 10.00 + + def test_fable_5_1(self): + inp, out = pricing_for_model("claude-fable-5-1") + assert inp == 10.00 + assert out == 50.00 + + def test_o3(self): + inp, out = pricing_for_model("o3") + assert inp == 2.00 + assert out == 8.00 + + def test_gemini_3_flash_preview(self): + inp, out = pricing_for_model("gemini-3-flash-preview") + assert inp == 0.50 + assert out == 3.00 + + def test_opus_5_5(self): + inp, out = pricing_for_model("claude-opus-5-5") + assert inp == 4.00 + assert out == 20.00 + def test_haiku(self): inp, out = pricing_for_model("claude-haiku-4-5-20251001") assert inp == 1.00 diff --git a/tests/test_gemini_backend.py b/tests/test_gemini_backend.py index a3c11c1..af51ca9 100644 --- a/tests/test_gemini_backend.py +++ b/tests/test_gemini_backend.py @@ -155,8 +155,8 @@ def test_gemini_models_in_pricing(self): from pact.budget import pricing_for_model inp, out = pricing_for_model("gemini-2.5-flash") - assert inp == 0.15 - assert out == 0.60 + assert inp == 0.30 + assert out == 2.50 def test_gemini_pro_pricing(self): from pact.budget import pricing_for_model diff --git a/tests/test_validation_retry.py b/tests/test_validation_retry.py index 4840d3d..0e227e4 100644 --- a/tests/test_validation_retry.py +++ b/tests/test_validation_retry.py @@ -247,7 +247,7 @@ async def test_token_counts_accumulate_across_retries(self): assert out_tok == 125 async def test_no_tool_use_block_raises_runtime_error(self): - """When _call_llm returns None, raises RuntimeError immediately.""" + """When every attempt returns no tool call, raises RuntimeError.""" backend = _make_backend() backend._call_llm = AsyncMock(return_value=( @@ -256,6 +256,7 @@ async def test_no_tool_use_block_raises_runtime_error(self): with pytest.raises(RuntimeError, match="No tool_use block found"): await backend.assess(SimpleSchema, "prompt", "system") + assert backend._call_llm.call_count == 3 async def test_original_prompt_preserved_in_correction(self): """Correction is appended to the original prompt, not replacing it."""