Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
203 changes: 151 additions & 52 deletions src/pact/backends/anthropic.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand All @@ -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)
Expand All @@ -74,33 +91,50 @@ 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,
)
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)
if new_max > current_max:
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:
Expand All @@ -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.
Expand Down Expand Up @@ -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 ──────────────────────────────────────────

Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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
Expand All @@ -438,24 +521,40 @@ 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,
)
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)
if new_max > current_max:
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:
Expand Down
24 changes: 17 additions & 7 deletions src/pact/budget.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading
Loading