fix(anthropic): count a leading system run through count_tokens' system parameter (#45463)

* fix(anthropic): count a leading system run through count_tokens' system parameter

Anthropic's count_tokens rejects role "system" at the head of messages, so
/v1/responses/input_tokens with instructions, and /v1/messages/count_tokens
with a system-role message, fell back to the local tokenizer. The shared
Anthropic count_tokens transformation now lifts the leading run of system
messages into the top-level system parameter, the way the chat path sends
it, after any system the caller set. Anthropic direct, Azure AI Anthropic,
and Bedrock Mantle share that transformation.

* test(integration): cover the count_tokens leading-system lift across Anthropic, Azure AI and Bedrock Mantle

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-08 22:34:25 -07:00 • committed by GitHub
parent 72ab736863
commit d34ad35281
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 1254 additions and 34 deletions

View file

@ -31,7 +31,7 @@ Anthropic wire shape is built later by ``anthropic_messages_pt``.
from collections.abc import Iterator, Mapping, Sequence
from itertools import chain, groupby
from typing import Final, Literal, TypeAlias
from typing import Final, Literal, TypeAlias, TypeVar
from litellm.types.llms.anthropic import AnthropicMessagesSystemMessageParam, AnthropicSystemMessageContent
from litellm.types.llms.openai import (
@ -55,6 +55,7 @@ _RENDERED_ASSISTANT_PART_TYPES: Final = frozenset({"text", "server_tool_use"})
_THINKING_BLOCK_TYPES: Final = frozenset({"thinking", "redacted_thinking"})
_MessageKind: TypeAlias = Literal["system", "tool", "user", "other"]
_Message: Final = TypeVar("_Message")
_TextPart: TypeAlias = tuple[str, ChatCompletionCachedContent | None]
@ -96,8 +97,8 @@ def _kind(message: object) -> _MessageKind:
def split_leading_system_run(
messages: Sequence[AllMessageValues],
) -> tuple[tuple[AllMessageValues, ...], tuple[AllMessageValues, ...]]:
messages: Sequence[_Message],
) -> tuple[tuple[_Message, ...], tuple[_Message, ...]]:
"""Split ``messages`` into the leading run of system messages and everything after it."""
leading_count: Final = next(
(index for index, message in enumerate(messages) if not is_system_message(message)),
@ -160,9 +161,16 @@ def _anthropic_text_block(part: _TextPart) -> AnthropicSystemMessageContent:
return cached
def anthropic_system_blocks(run: Sequence[object]) -> tuple[AnthropicSystemMessageContent, ...]:
"""The top-level ``system`` blocks for a run of system messages: every non-empty text part in order,
each keeping its ``cache_control``, which is the shape the chat path sends for the leading run."""
parts: Final = chain.from_iterable(_text_parts(message) for message in run)
return tuple(_anthropic_text_block(part) for part in parts)
def anthropic_system_messages(message: object) -> tuple[AnthropicMessagesSystemMessageParam, ...]:
"""The Anthropic wire message for a system message, or nothing when it carries no text."""
blocks: Final = tuple(_anthropic_text_block(part) for part in _text_parts(message))
blocks: Final = anthropic_system_blocks((message,))
if not blocks:
return ()
wire: Final[AnthropicMessagesSystemMessageParam] = {

View file

@ -11,11 +11,16 @@ from typing import Final
from pydantic import JsonValue, TypeAdapter
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
anthropic_system_blocks,
split_leading_system_run,
)
from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers
from litellm.llms.anthropic.wif import resolve_anthropic_base
from litellm.types.llms.openai import ChatCompletionImageObject
_COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue])
_SYSTEM_BLOCKS: Final = TypeAdapter(list[JsonValue])
_IMAGE_BLOCK: Final = TypeAdapter(ChatCompletionImageObject)
COUNT_TOKEN_OPTION_NAMES: Final = ("thinking", "tool_choice", "output_config")
@ -48,6 +53,27 @@ def _count_content(content: JsonValue) -> JsonValue:
return [_count_block(block) for block in content] if isinstance(content, list) else content
def _lift_leading_system(
messages: Sequence[Mapping[str, JsonValue]], system: JsonValue
) -> tuple[tuple[Mapping[str, JsonValue], ...], JsonValue]:
"""Move the leading run of system-role messages into the top-level ``system`` parameter.
count_tokens only takes the initial system prompt there and answers 400 on ``role: "system"``
at the head of ``messages``; the chat path sends the same run as ``system``. A caller's own
``system`` keeps its place ahead of the lifted blocks, and a ``system`` that is neither text
nor a block list is left as sent, messages included, for the provider to judge.
"""
leading, conversation = split_leading_system_run(messages)
if not leading or not (system is None or isinstance(system, (str, list))):
return tuple(messages), system
lifted: Final = _SYSTEM_BLOCKS.validate_python(list(anthropic_system_blocks(leading)))
if isinstance(system, list):
return conversation, [*system, *lifted]
if isinstance(system, str) and system:
return conversation, [{"type": "text", "text": system}, *lifted]
return conversation, lifted or system
class AnthropicCountTokensConfig:
"""
Configuration and transformation logic for Anthropic CountTokens API.
@ -85,16 +111,24 @@ class AnthropicCountTokensConfig:
"""
Transform request to Anthropic CountTokens format.
Includes optional system and tools fields for accurate token counting.
Includes optional system and tools fields for accurate token counting; a leading run of
system-role messages is counted through ``system``, the only place count_tokens accepts it.
"""
options: Final[Mapping[str, JsonValue]] = optional_params or MappingProxyType({})
counted_messages, counted_system = _lift_leading_system(messages, system)
return _COUNT_REQUEST.validate_python(
MappingProxyType(
{
"model": model,
"messages": [{**message, "content": _count_content(message["content"])} for message in messages],
"messages": [
{**message, "content": _count_content(message["content"])} for message in counted_messages
],
**MappingProxyType(
{key: value for key, value in (("system", system), ("tools", tools)) if value is not None}
{
key: value
for key, value in (("system", counted_system), ("tools", tools))
if value is not None
}
),
**MappingProxyType(
{key: value for key, value in options.items() if key in COUNT_TOKEN_OPTION_NAMES}

View file

@ -0,0 +1,159 @@
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
import anthropic
import httpx
import openai
from integration._support.client import Gateway
from pydantic import JsonValue, TypeAdapter
JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
TOKEN_COUNTING_BETA: Final = "token-counting-2024-11-01"
ANTHROPIC_VERSION: Final = "2023-06-01"
REJECTION: Final = 'messages.0: Unexpected role "system". The Messages API accepts a top-level `system` parameter'
USER_TEXT: Final = "Count this message"
USER: Final[dict[str, JsonValue]] = {"role": "user", "content": USER_TEXT}
ASSISTANT: Final[dict[str, JsonValue]] = {"role": "assistant", "content": "One."}
FOLLOW_UP: Final[dict[str, JsonValue]] = {"role": "user", "content": "Again"}
INSTRUCTION: Final = "You are a terse assistant"
REMINDER: Final = "Answer in one sentence"
LEADING: Final[dict[str, JsonValue]] = {"role": "system", "content": INSTRUCTION}
LIFTED: Final[dict[str, JsonValue]] = {"type": "text", "text": INSTRUCTION}
MID_SYSTEM: Final[dict[str, JsonValue]] = {"role": "system", "content": REMINDER}
EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
ONE_HOUR: Final[dict[str, JsonValue]] = {"type": "ephemeral", "ttl": "1h"}
IMAGE_PART: Final[dict[str, JsonValue]] = {
"type": "image_url",
"image_url": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
},
}
FIVE_KB: Final = "Answer in one sentence. " * 214
CALLER_SYSTEM: Final = "Prefer metric units"
CALLER_BLOCKS: Final[list[JsonValue]] = [
{"type": "text", "text": CALLER_SYSTEM},
{"type": "text", "text": "Never guess", "cache_control": EPHEMERAL},
]
TOOLS: Final[list[JsonValue]] = [
{
"name": "get_weather",
"description": "Look up the current weather for a city",
"input_schema": {
"type": "object",
"properties": {"city": {"type": "string", "description": "City to look up"}},
"required": ["city"],
},
}
]
@dataclass(frozen=True, slots=True)
class LiftCase:
messages: tuple[dict[str, JsonValue], ...]
system: JsonValue | None
lifted: list[JsonValue] | None
LIFT_CASES: Final[Mapping[str, LiftCase]] = MappingProxyType(
{
"string": LiftCase((LEADING, USER), None, [LIFTED]),
"run_with_cache_control": LiftCase(
(
{"role": "system", "content": INSTRUCTION, "cache_control": ONE_HOUR},
{
"role": "system",
"content": [
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
{"type": "text", "text": ""},
IMAGE_PART,
],
},
USER,
),
None,
[
{"type": "text", "text": INSTRUCTION, "cache_control": ONE_HOUR},
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
],
),
"caller_system_string_first": LiftCase(
(LEADING, USER), CALLER_SYSTEM, [{"type": "text", "text": CALLER_SYSTEM}, LIFTED]
),
"caller_system_blocks_first": LiftCase((LEADING, USER), CALLER_BLOCKS, [*CALLER_BLOCKS, LIFTED]),
"caller_empty_system_dropped": LiftCase((LEADING, USER), "", [LIFTED]),
"empty_content_dropped": LiftCase(({"role": "system", "content": ""}, USER), None, None),
"image_only_content_dropped": LiftCase(({"role": "system", "content": [IMAGE_PART]}, USER), None, None),
"integer_content_dropped": LiftCase(({"role": "system", "content": 5}, USER), None, None),
"integer_text_part_dropped": LiftCase(
({"role": "system", "content": [{"type": "text", "text": 7}]}, USER), None, None
),
"five_kb": LiftCase(({"role": "system", "content": FIVE_KB}, USER), None, [{"type": "text", "text": FIVE_KB}]),
}
)
STRING_CASE: Final = LIFT_CASES["string"]
def count_request(model: str, case: LiftCase) -> dict[str, JsonValue]:
return {"model": model, "messages": list(case.messages), **({} if case.system is None else {"system": case.system})}
def expected_count_body(model: str, case: LiftCase) -> dict[str, JsonValue]:
conversation: Final[list[JsonValue]] = [message for message in case.messages if message["role"] != "system"]
return {"model": model, "messages": conversation, **({} if case.lifted is None else {"system": case.lifted})}
def _opening_message(message: JsonValue) -> bool:
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
def _conversation_message(message: JsonValue) -> bool:
return isinstance(message, dict) and message.get("role") in ("user", "assistant", "system")
def _anthropic_tool(tool: JsonValue) -> bool:
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
def accepts_count_body(body: Mapping[str, JsonValue]) -> bool:
# Anthropic, Azure AI Foundry and Bedrock Mantle count_tokens verdicts observed live on 2026-10-09: an empty
# messages list and a system role at messages[0] answer 400 invalid_request_error, a later system role counts
messages: Final = body.get("messages")
tools: Final = body.get("tools", [])
return (
isinstance(messages, list)
and len(messages) > 0
and _opening_message(messages[0])
and all(map(_conversation_message, messages))
and isinstance(body.get("system", ""), (str, list))
and isinstance(tools, list)
and all(map(_anthropic_tool, tools))
)
def anthropic_client(gateway: Gateway) -> anthropic.Anthropic:
return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
def async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic:
return anthropic.AsyncAnthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0)
def openai_client(gateway: Gateway) -> openai.OpenAI:
return openai.OpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.Client(trust_env=False),
)
def async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
return openai.AsyncOpenAI(
base_url=f"{gateway.client.base_url}/v1",
api_key=gateway.key,
max_retries=0,
http_client=httpx.AsyncClient(trust_env=False),
)

View file

@ -0,0 +1,572 @@
import asyncio
import json
import socket
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import AbstractContextManager, ExitStack
from dataclasses import dataclass
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final, cast
import httpcore
import httpx
import psutil
import pytest
from anthropic.types import MessageParam
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._count_tokens_system_lift import (
ANTHROPIC_VERSION,
ASSISTANT,
EPHEMERAL,
FOLLOW_UP,
IMAGE_PART,
INSTRUCTION,
JSON_OBJECT,
LEADING,
LIFT_CASES,
LIFTED,
MID_SYSTEM,
REJECTION,
REMINDER,
STRING_CASE,
TOKEN_COUNTING_BETA,
TOOLS,
USER,
USER_TEXT,
LiftCase,
accepts_count_body,
anthropic_client,
async_anthropic_client,
async_openai_client,
count_request,
expected_count_body,
openai_client,
)
from pydantic import JsonValue, TypeAdapter
_MODEL: Final = "claude-opus-5-5"
_API_KEY: Final = "synthetic-count-tokens-key"
_COUNT: Final = 3131
_PROVIDER_TOKENIZERS: Final = frozenset({"anthropic_api", "azure_ai_anthropic_api"})
_SDK_MESSAGES: Final = cast(
list[MessageParam], [LEADING, USER]
) # cast-ok: the SDK types reject the role the proxy lifts
_STRING_BODY: Final = expected_count_body(_MODEL, STRING_CASE)
_PROXY_MODULE: Final = "integration._support.proxy"
_PROBES_PER_ROUND: Final = 8
_CLIENT_ADDRESS: Final = TypeAdapter(tuple[str, int])
_ARGUMENTS: Final = TypeAdapter(tuple[str, ...])
_NAME: Final = TypeAdapter(str)
@dataclass(frozen=True, slots=True)
class _Provider:
prefix: str
target: str
tokenizer: str
azure: bool
@dataclass(frozen=True, slots=True)
class _Deployment:
provider: _Provider
port: int
model: str
_ANTHROPIC: Final = _Provider("anthropic", "/v1/messages/count_tokens", "anthropic_api", False)
_AZURE: Final = _Provider("azure_ai", "/anthropic/v1/messages/count_tokens", "azure_ai_anthropic_api", True)
_PROVIDERS: Final = MappingProxyType({"anthropic": _ANTHROPIC, "azure_ai": _AZURE})
@pytest.fixture(params=_PROVIDERS.keys())
def provider(request: pytest.FixtureRequest) -> _Provider:
name: Final[object] = request.param # pyright: ignore[reportAny] # pytest types the fixture param as Any
return _PROVIDERS[_NAME.validate_python(name)]
def _json_reply(status: int, payload: Mapping[str, JsonValue]) -> Reply:
return Reply(status=status, body=json.dumps(payload).encode())
def _rejected(status: int) -> Reply:
return _json_reply(status, {"type": "error", "error": {"type": "invalid_request_error", "message": REJECTION}})
def _counted(request: Request) -> Reply:
accepted: Final = accepts_count_body(JSON_OBJECT.validate_json(request.body))
return _json_reply(200, {"input_tokens": _COUNT}) if accepted else _rejected(400)
def _rejecting(status: int) -> Callable[[Request], Reply]:
def count(_request: Request) -> Reply:
return _rejected(status)
return count
def _holding(held: SimpleQueue[str], release: threading.Event, seconds: float) -> Callable[[Request], Reply]:
def hold(request: Request) -> Reply:
held.put(request.target)
assert release.wait(timeout=seconds), "Held count was never released"
return _counted(request)
return hold
def _peer(provider: _Provider, count: Callable[[Request], Reply] = _counted) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if request.target == provider.target:
return count(request)
return _json_reply(404, {"error": f"unscripted target {request.target}"})
return respond
def _listening(deployment: _Deployment, count: Callable[[Request], Reply] = _counted) -> AbstractContextManager[Wire]:
return wire_server(_peer(deployment.provider, count), port=deployment.port)
def _reserved_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return _CLIENT_ADDRESS.validate_python(reserve.getsockname())[1]
def _cmdline(process: psutil.Process) -> tuple[str, ...]:
try:
return _ARGUMENTS.validate_python(process.cmdline())
except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess):
return ()
def _serves(cmdline: tuple[str, ...], proxy_port: int) -> bool:
if _PROXY_MODULE not in cmdline or "--port" not in cmdline:
return False
return cmdline[cmdline.index("--port") + 1] == str(proxy_port)
def _listens(process: psutil.Process, proxy_port: int) -> bool:
return any(
connection.status == psutil.CONN_LISTEN and connection.laddr and connection.laddr.port == proxy_port
for connection in process.net_connections(kind="tcp")
)
def _proxy_workers(proxy_port: int) -> frozenset[int]:
(master,) = tuple(process for process in psutil.process_iter() if _serves(_cmdline(process), proxy_port))
spawned: Final = frozenset(child.pid for child in master.children() if _listens(child, proxy_port))
return spawned or frozenset({master.pid})
def _holder(workers: frozenset[int], client_port: int) -> int | None:
def holds(pid: int) -> bool:
return any(
connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == client_port
for connection in psutil.Process(pid).net_connections(kind="tcp")
)
return next((pid for pid in sorted(workers) if holds(pid)), None)
def _client_port(response: httpx.Response) -> int:
stream: Final[object] = response.extensions["network_stream"] # pyright: ignore[reportAny] # httpx types extensions as Any
assert isinstance(stream, httpcore.NetworkStream), stream
return _CLIENT_ADDRESS.validate_python(stream.get_extra_info("client_addr"))[1]
def _probe(gateway: Gateway, body: Mapping[str, JsonValue], workers: frozenset[int]) -> tuple[int | None, JsonValue]:
with httpx.Client(base_url=str(gateway.client.base_url), timeout=30, trust_env=False) as client:
response: Final = client.post(
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {gateway.key}"}
)
return _holder(workers, _client_port(response)), JSON_OBJECT.validate_json(response.content).get("input_tokens")
def _round(
gateway: Gateway, body: Mapping[str, JsonValue], workers: frozenset[int]
) -> tuple[tuple[int | None, JsonValue], ...]:
def probe(_index: int) -> tuple[int | None, JsonValue]:
return _probe(gateway, body, workers)
with ThreadPoolExecutor(max_workers=_PROBES_PER_ROUND) as pool:
return tuple(pool.map(probe, range(_PROBES_PER_ROUND)))
def _settled_on_every_worker(gateway: Gateway, body: Mapping[str, JsonValue]) -> None:
proxy_port: Final = gateway.client.base_url.port
assert proxy_port is not None, gateway.client.base_url
workers: Final = _proxy_workers(proxy_port)
eventually(
lambda: _round(gateway, body, workers),
lambda observed: (
frozenset(pid for pid, _ in observed) == workers and all(count == _COUNT for _, count in observed)
),
seconds=60,
)
def _deploy(gateway: Gateway, scenario: Scenario, provider: _Provider) -> _Deployment:
port: Final = _reserved_port()
model: Final = scenario.model(
model_info=None, model=f"{provider.prefix}/{_MODEL}", api_base=f"http://127.0.0.1:{port}", api_key=_API_KEY
)
with wire_server(_peer(provider), port=port):
_settled_on_every_worker(gateway, {"model": model, "messages": [USER]})
return _Deployment(provider, port, model)
@pytest.fixture(scope="module")
def deployments() -> Iterator[Mapping[str, _Deployment]]:
with gateway_from_environment() as gateway, gateway.scenario() as scenario:
yield MappingProxyType({chosen.prefix: _deploy(gateway, scenario, chosen) for chosen in (_ANTHROPIC, _AZURE)})
@pytest.fixture
def deployment(provider: _Provider, deployments: Mapping[str, _Deployment]) -> _Deployment:
return deployments[provider.prefix]
def _count(gateway: Gateway, body: Mapping[str, JsonValue], key: str | None = None) -> httpx.Response:
return gateway.request("POST", "/v1/messages/count_tokens", body, key=key)
def _payload(response: httpx.Response) -> dict[str, JsonValue]:
assert response.status_code == 200, response.text
return JSON_OBJECT.validate_json(response.content)
def _local_count(gateway: Gateway, body: Mapping[str, JsonValue]) -> int:
response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "false"})
payload: Final = _payload(response)
total: Final = payload["total_tokens"]
assert payload["tokenizer_type"] not in _PROVIDER_TOKENIZERS, response.text
assert isinstance(total, int) and total > 0 and total != _COUNT, response.text
return total
def _bodies(wire: Wire, provider: _Provider) -> tuple[dict[str, JsonValue], ...]:
received: Final = wire.drain()
for request in received:
assert (request.method, request.target) == ("POST", provider.target), request.target
assert request.headers["anthropic-version"] == ANTHROPIC_VERSION, request.headers
assert TOKEN_COUNTING_BETA in request.headers["anthropic-beta"], request.headers
assert request.headers["content-type"] == "application/json", request.headers
assert request.headers["x-api-key"] == _API_KEY, request.headers
assert (request.headers.get("api-key") == _API_KEY) is provider.azure, request.headers
return tuple(JSON_OBJECT.validate_json(request.body) for request in received)
def _clients(stack: ExitStack, base_url: str, count: int) -> tuple[httpx.Client, ...]:
return tuple(
stack.enter_context(httpx.Client(base_url=base_url, timeout=30, trust_env=False)) for _ in range(count)
)
def _counted_on(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue]:
response: Final = client.post(
"/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {key}"}
)
return response.status_code, JSON_OBJECT.validate_json(response.content).get("input_tokens")
@pytest.mark.parametrize("case", LIFT_CASES.values(), ids=LIFT_CASES.keys())
def test_messages_count_tokens_lifts_the_leading_system_run(
deployment: _Deployment, gateway: Gateway, case: LiftCase
) -> None:
with _listening(deployment) as peer:
response: Final = _count(gateway, count_request(deployment.model, case))
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (expected_count_body(_MODEL, case),)
def test_anthropic_sdk_count_tokens_lifts_the_leading_system_message(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
counted: Final = anthropic_client(gateway).messages.count_tokens(model=deployment.model, messages=_SDK_MESSAGES)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_async_anthropic_sdk_count_tokens_lifts_the_leading_system_message(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
counted: Final = asyncio.run(
async_anthropic_client(gateway).messages.count_tokens(model=deployment.model, messages=_SDK_MESSAGES)
)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_utils_token_counter_call_endpoint_counts_a_leading_system_through_the_provider(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = gateway.request(
"POST",
"/utils/token_counter",
count_request(deployment.model, STRING_CASE),
params={"call_endpoint": "true"},
)
payload: Final = _payload(response)
expected_tokenizer: Final = deployment.provider.tokenizer
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_COUNT, expected_tokenizer), response.text
assert payload["original_response"] == {"input_tokens": _COUNT}, response.text
assert (payload["request_model"], payload["model_used"]) == (deployment.model, _MODEL), response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_utils_token_counter_local_mode_never_calls_the_peer_for_a_leading_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
assert _local_count(gateway, count_request(deployment.model, STRING_CASE)) > 0
assert peer.drain() == ()
def test_responses_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": deployment.model, "input": USER_TEXT, "instructions": INSTRUCTION},
)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_responses_input_tokens_lifts_instructions_ahead_of_a_leading_system_item(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": deployment.model, "input": [MID_SYSTEM, USER], "instructions": INSTRUCTION},
)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (
{"model": _MODEL, "messages": [USER], "system": [LIFTED, {"type": "text", "text": REMINDER}]},
)
def test_openai_sdk_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
counted: Final = openai_client(gateway).responses.input_tokens.count(
model=deployment.model, input=USER_TEXT, instructions=INSTRUCTION
)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_async_openai_sdk_input_tokens_lifts_instructions(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
counted: Final = asyncio.run(
async_openai_client(gateway).responses.input_tokens.count(
model=deployment.model, input=USER_TEXT, instructions=INSTRUCTION
)
)
assert counted.input_tokens == _COUNT, counted
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_messages_count_tokens_forwards_tools_beside_the_lifted_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = _count(gateway, {**count_request(deployment.model, STRING_CASE), "tools": TOOLS})
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == ({**_STRING_BODY, "tools": TOOLS},)
def test_messages_count_tokens_keeps_a_mid_conversation_system_in_place(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
body: Final[dict[str, JsonValue]] = {
"model": deployment.model,
"messages": [LEADING, USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP],
}
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (
{"model": _MODEL, "messages": [USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP], "system": [LIFTED]},
)
def test_messages_count_tokens_falls_back_locally_when_every_message_is_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
body: Final[dict[str, JsonValue]] = {"model": deployment.model, "messages": [LEADING]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert _bodies(peer, deployment.provider) == ({"model": _MODEL, "messages": [], "system": [LIFTED]},)
def test_messages_count_tokens_leaves_the_request_untouched_for_a_non_text_system(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
local: Final = _local_count(gateway, count_request(deployment.model, STRING_CASE))
response: Final = _count(gateway, {**count_request(deployment.model, STRING_CASE), "system": 5})
assert _payload(response) == {"input_tokens": local}, response.text
assert _bodies(peer, deployment.provider) == ({"model": _MODEL, "messages": [LEADING, USER], "system": 5},)
def test_messages_count_tokens_repeated_request_lifts_each_time(deployment: _Deployment, gateway: Gateway) -> None:
with _listening(deployment) as peer:
answers: Final = tuple(
_payload(_count(gateway, count_request(deployment.model, STRING_CASE))) for _ in range(2)
)
assert answers == ({"input_tokens": _COUNT},) * 2
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * 2
def test_messages_count_tokens_answers_a_leading_system_without_content_before_any_peer_call(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
body: Final[dict[str, JsonValue]] = {"model": deployment.model, "messages": [{"role": "system"}, USER]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert peer.drain() == ()
follow_up: Final = _count(gateway, count_request(deployment.model, STRING_CASE))
assert _payload(follow_up) == {"input_tokens": _COUNT}, follow_up.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_messages_count_tokens_duplicate_messages_key_lifts_the_last_value(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
first: Final = json.dumps([USER])
last: Final = json.dumps([LEADING, USER])
response: Final = gateway.client.post(
"/v1/messages/count_tokens",
content=f'{{"model": "{deployment.model}", "messages": {first}, "messages": {last}}}',
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
)
assert _payload(response) == {"input_tokens": _COUNT}, response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
@pytest.mark.parametrize("status", [400, 403, 404, 500, 503])
def test_messages_count_tokens_falls_back_locally_when_the_peer_rejects_the_lifted_body(
deployment: _Deployment, gateway: Gateway, status: int
) -> None:
with _listening(deployment, _rejecting(status)) as peer:
body: Final = count_request(deployment.model, STRING_CASE)
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert _bodies(peer, deployment.provider) == (_STRING_BODY,)
def test_messages_count_tokens_unauthenticated_request_never_reaches_the_peer(
deployment: _Deployment, gateway: Gateway
) -> None:
with _listening(deployment) as peer:
response: Final = _count(gateway, count_request(deployment.model, STRING_CASE), key="sk-not-a-key")
assert response.status_code == 401, response.text
assert peer.drain() == ()
def test_peer_outage_between_concurrent_waves_falls_back_then_recovers(
deployment: _Deployment, gateway: Gateway
) -> None:
with ExitStack() as stack:
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
body: Final = count_request(deployment.model, STRING_CASE)
local: Final = _local_count(gateway, body)
def count(client: httpx.Client) -> tuple[int, JsonValue]:
return _counted_on(client, gateway.key, body)
with _listening(deployment) as peer:
assert tuple(pool.map(count, clients)) == ((200, _COUNT),) * len(clients)
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * len(clients)
assert tuple(pool.map(count, clients)) == ((200, local),) * len(clients)
with _listening(deployment) as revived:
assert tuple(pool.map(count, clients)) == ((200, _COUNT),) * len(clients)
assert _bodies(revived, deployment.provider) == (_STRING_BODY,) * len(clients)
def test_slow_peer_holds_concurrent_lifted_counts_without_stalling_the_proxy(
deployment: _Deployment, gateway: Gateway
) -> None:
held: Final[SimpleQueue[str]] = SimpleQueue()
release: Final = threading.Event()
with ExitStack() as stack:
clients: Final = _clients(stack, str(gateway.client.base_url), 6)
peer: Final = stack.enter_context(_listening(deployment, _holding(held, release, 20)))
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
stack.callback(release.set)
body: Final = count_request(deployment.model, STRING_CASE)
futures: Final = tuple(pool.submit(_counted_on, client, gateway.key, body) for client in clients)
eventually(held.qsize, lambda size: size == len(clients), seconds=30)
assert gateway.request("GET", "/health/liveliness").status_code == 200
assert _local_count(gateway, body) > 0
assert not any(future.done() for future in futures)
release.set()
assert tuple(future.result(timeout=30) for future in futures) == ((200, _COUNT),) * len(clients)
assert _bodies(peer, deployment.provider) == (_STRING_BODY,) * len(clients)
def test_chat_completions_on_the_same_deployment_keeps_a_mid_conversation_system_in_place(
deployments: Mapping[str, _Deployment], gateway: Gateway
) -> None:
identity: Final = f"msg_{uuid.uuid4().hex}"
anthropic_deployment: Final = deployments[_ANTHROPIC.prefix]
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", "/v1/messages"), request.target
return _json_reply(
200,
{
"id": identity,
"type": "message",
"role": "assistant",
"model": _MODEL,
"content": [{"type": "text", "text": "done"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 12, "output_tokens": 3},
},
)
with wire_server(respond, port=anthropic_deployment.port) as peer:
reminder: Final[dict[str, JsonValue]] = {
"role": "system",
"content": [
{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL},
{"type": "text", "text": ""},
IMAGE_PART,
],
}
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": anthropic_deployment.model,
"max_tokens": 16,
"messages": [LEADING, USER, reminder, ASSISTANT, FOLLOW_UP],
},
)
assert response.status_code == 200, response.text
(sent,) = peer.drain()
body: Final = JSON_OBJECT.validate_json(sent.body)
assert body["system"] == [LIFTED], body
assert body["messages"] == [
{"role": "user", "content": [{"type": "text", "text": USER_TEXT}]},
{"role": "system", "content": [{"type": "text", "text": REMINDER, "cache_control": EPHEMERAL}]},
{"role": "assistant", "content": [{"type": "text", "text": "One."}]},
{"role": "user", "content": [{"type": "text", "text": "Again"}]},
], body

View file

@ -11,7 +11,7 @@ from contextlib import ExitStack
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
from typing import Final, cast
import anthropic
import httpx
@ -24,6 +24,25 @@ from integration._support.bedrock_runtime_peer import respond as runtime_generat
from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment, object_value
from integration._support.process import OwnedProxy, graceful_stop_seconds, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from integration.providers._count_tokens_system_lift import (
ASSISTANT,
FOLLOW_UP,
INSTRUCTION,
LEADING,
LIFT_CASES,
LIFTED,
MID_SYSTEM,
REMINDER,
STRING_CASE,
USER,
USER_TEXT,
LiftCase,
accepts_count_body,
async_openai_client,
count_request,
expected_count_body,
openai_client,
)
from pydantic import JsonValue, TypeAdapter
pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120)
@ -81,6 +100,10 @@ def _mantle_body(**fields: JsonValue) -> dict[str, JsonValue]:
_MANTLE_BARE: Final = _mantle_body()
_MANTLE_LIFTED: Final = _mantle_body(system=[LIFTED])
_LIFT_SDK_MESSAGES: Final = cast(
list[MessageParam], [LEADING, USER]
) # cast-ok: the SDK types reject the role the proxy lifts
_MANTLE_FULL: Final = _mantle_body(system=_SYSTEM, tools=_TOOLS)
@ -138,25 +161,8 @@ def _rejecting(status: int) -> Callable[[Request], Reply]:
return count
def _anthropic_message(message: JsonValue) -> bool:
return isinstance(message, dict) and message.get("role") in ("user", "assistant")
def _anthropic_tool(tool: JsonValue) -> bool:
return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict)
def _strict(request: Request) -> Reply:
body: Final = _JSON_OBJECT.validate_json(request.body)
messages: Final = body.get("messages")
tools: Final = body.get("tools", [])
accepted: Final = (
isinstance(messages, list)
and all(map(_anthropic_message, messages))
and isinstance(body.get("system", ""), (str, list))
and isinstance(tools, list)
and all(map(_anthropic_tool, tools))
)
accepted: Final = accepts_count_body(_JSON_OBJECT.validate_json(request.body))
return _mantle_counted(request) if accepted else _rejected(400)
@ -577,7 +583,7 @@ def test_bedrock_passthrough_count_tokens_still_answers_the_runtime_rejection(
assert mantle.drain() == ()
def test_responses_input_tokens_with_instructions_still_counts_locally(
def test_responses_input_tokens_with_instructions_counts_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
@ -592,14 +598,322 @@ def test_responses_input_tokens_with_instructions_still_counts_locally(
"/v1/responses/input_tokens",
{"model": model, "input": "Count this message", "instructions": "Be terse"},
)
payload: Final = _payload(response)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
(sent,) = _mantle_bodies(mantle)
messages: Final = sent["messages"]
assert isinstance(messages, list) and messages[0] == {"role": "system", "content": "Be terse"}, sent
assert "system" not in sent, sent
local: Final = _local_count(gateway, {"model": model, "messages": messages})
assert payload == {"object": "response.input_tokens", "input_tokens": local}, response.text
assert _mantle_bodies(mantle) == (_mantle_body(system=[{"type": "text", "text": "Be terse"}]),)
@pytest.mark.parametrize("case", LIFT_CASES.values(), ids=LIFT_CASES.keys())
def test_messages_count_tokens_lifts_the_leading_system_run_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int, case: LiftCase
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = _count(gateway, count_request(model, case))
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (expected_count_body(_OPUS_BASE, case),)
def test_anthropic_sdk_count_tokens_lifts_the_leading_system_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = _anthropic_client(gateway).messages.count_tokens(model=model, messages=_LIFT_SDK_MESSAGES)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_async_anthropic_sdk_count_tokens_lifts_the_leading_system_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = asyncio.run(
_async_anthropic_client(gateway).messages.count_tokens(model=model, messages=_LIFT_SDK_MESSAGES)
)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_utils_token_counter_call_endpoint_counts_a_leading_system_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = gateway.request(
"POST", "/utils/token_counter", count_request(model, STRING_CASE), params={"call_endpoint": "true"}
)
payload: Final = _payload(response)
assert (payload["total_tokens"], payload["tokenizer_type"]) == (_MANTLE_COUNT, "bedrock_mantle_api")
assert payload["original_response"] == {"input_tokens": _MANTLE_COUNT}, response.text
assert (payload["request_model"], payload["model_used"]) == (model, _OPUS), response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_openai_sdk_input_tokens_lifts_instructions_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = openai_client(gateway).responses.input_tokens.count(
model=model, input=USER_TEXT, instructions=INSTRUCTION
)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_async_openai_sdk_input_tokens_lifts_instructions_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
counted: Final = asyncio.run(
async_openai_client(gateway).responses.input_tokens.count(
model=model, input=USER_TEXT, instructions=INSTRUCTION
)
)
assert counted.input_tokens == _MANTLE_COUNT, counted
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_responses_input_tokens_lifts_instructions_ahead_of_a_leading_system_item_through_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = gateway.request(
"POST",
"/v1/responses/input_tokens",
{"model": model, "input": [MID_SYSTEM, USER], "instructions": INSTRUCTION},
)
assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(system=[LIFTED, {"type": "text", "text": REMINDER}]),)
def test_messages_count_tokens_forwards_tools_beside_the_lifted_system_to_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
response: Final = _count(gateway, {**count_request(model, STRING_CASE), "tools": _TOOLS})
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(system=[LIFTED], tools=_TOOLS),)
def test_messages_count_tokens_keeps_a_mid_conversation_system_in_place_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
body: Final[dict[str, JsonValue]] = {
"model": model,
"messages": [LEADING, USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP],
}
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (
_mantle_body(messages=[USER, MID_SYSTEM, ASSISTANT, FOLLOW_UP], system=[LIFTED]),
)
def test_messages_count_tokens_falls_back_locally_when_every_message_is_system_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [LEADING]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(messages=[], system=[LIFTED]),)
def test_messages_count_tokens_leaves_a_leading_system_in_place_beside_a_non_text_system_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
local: Final = _local_count(gateway, count_request(model, STRING_CASE))
response: Final = _count(gateway, {**count_request(model, STRING_CASE), "system": 5})
assert _payload(response) == {"input_tokens": local}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_mantle_body(messages=[LEADING, USER], system=5),)
def test_messages_count_tokens_repeated_request_lifts_the_leading_system_each_time_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
answers: Final = tuple(_payload(_count(gateway, count_request(model, STRING_CASE))) for _ in range(2))
assert answers == ({"input_tokens": _MANTLE_COUNT},) * 2
assert _runtime_count_targets(runtime) == (f"/model/{_OPUS_BASE}/count-tokens",) * 2
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,) * 2
def test_messages_count_tokens_answers_a_leading_system_without_content_before_any_mantle_call(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
body: Final[dict[str, JsonValue]] = {"model": model, "messages": [{"role": "system"}, USER]}
local: Final = _local_count(gateway, body)
response: Final = _count(gateway, body)
assert _payload(response) == {"input_tokens": local}, response.text
assert mantle.drain() == ()
follow_up: Final = _count(gateway, count_request(model, STRING_CASE))
assert _payload(follow_up) == {"input_tokens": _MANTLE_COUNT}, follow_up.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_messages_count_tokens_duplicate_messages_key_lifts_the_last_value_for_mantle(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with (
wire_server(_runtime) as runtime,
wire_server(_mantle(_strict), port=mantle_port) as mantle,
gateway.scenario() as scenario,
):
model: Final = _deployment(scenario, runtime.url)
first: Final = json.dumps([USER])
last: Final = json.dumps([LEADING, USER])
response: Final = gateway.client.post(
"/v1/messages/count_tokens",
content=f'{{"model": "{model}", "messages": {first}, "messages": {last}}}',
headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"},
)
assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text
assert len(_runtime_count_targets(runtime)) == 1
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,)
def test_disabled_token_counter_counts_a_leading_system_through_mantle(
gateway: Gateway, mantle_port: int, tmp_path: Path
) -> None:
with ExitStack() as stack:
runtime: Final = stack.enter_context(wire_server(_runtime))
config: Final = _owned_config(
tmp_path / "disabled-token-counter-lift.yaml", runtime.url, {"disable_token_counter": True}
)
owned: Final = stack.enter_context(
owned_proxy_process(
gateway,
tmp_path,
_mantle_environment(mantle_port),
config=config,
workers=2,
remove_environment=_INHERITED_BEARER,
)
)
with wire_server(_mantle(_strict), port=mantle_port) as strict:
counted: Final = _count(owned.gateway, count_request(_OWNED_OPUS, STRING_CASE))
assert _payload(counted) == {"input_tokens": _MANTLE_COUNT}, counted.text
assert _mantle_bodies(strict) == (_MANTLE_LIFTED,)
assert len(_runtime_count_targets(runtime)) == 1
def test_mantle_outage_between_concurrent_lifted_waves_falls_back_then_recovers(
counting_proxy: OwnedProxy, mantle_port: int
) -> None:
gateway: Final = counting_proxy.gateway
with ExitStack() as stack:
clients: Final = _clients(stack, str(gateway.client.base_url), 8)
pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients)))
runtime: Final = stack.enter_context(wire_server(_runtime))
scenario: Final = stack.enter_context(gateway.scenario())
model: Final = _deployment(scenario, runtime.url)
body: Final = count_request(model, STRING_CASE)
local: Final = _local_count(gateway, body)
def count(client: httpx.Client) -> tuple[int, JsonValue]:
return _counted_on(client, gateway.key, body)
with wire_server(_mantle(_strict), port=mantle_port) as mantle:
assert tuple(pool.map(count, clients)) == ((200, _MANTLE_COUNT),) * len(clients)
assert _mantle_bodies(mantle) == (_MANTLE_LIFTED,) * len(clients)
assert tuple(pool.map(count, clients)) == ((200, local),) * len(clients)
with wire_server(_mantle(_strict), port=mantle_port) as revived:
assert tuple(pool.map(count, clients)) == ((200, _MANTLE_COUNT),) * len(clients)
assert _mantle_bodies(revived) == (_MANTLE_LIFTED,) * len(clients)
assert len(_runtime_count_targets(runtime)) == 3 * len(clients)
@pytest.mark.parametrize("status", [400, 403, 404, 500, 503])

View file

@ -11,6 +11,7 @@ import litellm
from litellm.litellm_core_utils.prompt_templates.common_utils import encrypted_reasoning_signature
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import (
CONVERTED_SYSTEM_NOTE,
anthropic_system_blocks,
place_mid_conversation_system,
split_leading_system_run,
)
@ -390,3 +391,23 @@ def test_flagged_placement_keeps_a_system_before_an_assistant_turn_whose_empty_t
)
assert _roles(placed) == ["user", "system", "assistant", "user"]
def test_anthropic_system_blocks_keeps_text_parts_with_their_cache_control_and_drops_the_rest():
run = [
{"role": "system", "content": "one", "cache_control": {"type": "ephemeral", "ttl": "1h"}},
{
"role": "system",
"content": [
{"type": "text", "text": "two", "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": ""},
{"type": "image_url", "image_url": {"url": "data:image/png;base64,aW1hZ2U="}},
],
},
{"role": "system", "content": ""},
]
assert anthropic_system_blocks(run) == (
{"type": "text", "text": "one", "cache_control": {"type": "ephemeral", "ttl": "1h"}},
{"type": "text", "text": "two", "cache_control": {"type": "ephemeral"}},
)

View file

@ -329,3 +329,115 @@ async def test_remote_image_fetch_keeps_counting_handler_event_loop_responsive(
]}],
}
assert messages == original
@pytest.mark.parametrize(
"config_type", (AnthropicCountTokensConfig, AzureAIAnthropicCountTokensConfig)
)
def test_count_lifts_the_leading_system_run_into_system(
config_type: type[AnthropicCountTokensConfig],
) -> None:
"""A Responses ``instructions`` arrives as a leading system-role message. count_tokens answers 400
on that role at the head of ``messages`` and only takes the initial prompt in ``system``, so the
leading run moves there with its cache_control, empty text dropped, and a later reminder stays."""
cache_control: Final[dict[str, JsonValue]] = {"type": "ephemeral"}
messages: Final[list[dict[str, JsonValue]]] = [
{"role": "system", "content": "Be terse", "cache_control": cache_control},
{"role": "system", "content": [{"type": "text", "text": "Answer in French"}, {"type": "text", "text": ""}]},
{"role": "user", "content": "Hello, how are you?"},
{"role": "system", "content": "later reminder"},
{"role": "assistant", "content": "Bonjour."},
]
original: Final = deepcopy(messages)
result: Final = config_type().transform_request_to_count_tokens(model="claude-opus-5-5", messages=messages)
assert result == {
"model": "claude-opus-5-5",
"system": [
{"type": "text", "text": "Be terse", "cache_control": cache_control},
{"type": "text", "text": "Answer in French"},
],
"messages": [
{"role": "user", "content": "Hello, how are you?"},
{"role": "system", "content": "later reminder"},
{"role": "assistant", "content": "Bonjour."},
],
}
assert messages == original
@pytest.mark.parametrize(
("system", "expected_system"),
(
(None, [{"type": "text", "text": "Be terse"}]),
("", [{"type": "text", "text": "Be terse"}]),
("Answer in French", [{"type": "text", "text": "Answer in French"}, {"type": "text", "text": "Be terse"}]),
(
[{"type": "text", "text": "Answer in French", "cache_control": {"type": "ephemeral"}}],
[
{"type": "text", "text": "Answer in French", "cache_control": {"type": "ephemeral"}},
{"type": "text", "text": "Be terse"},
],
),
),
ids=["absent", "empty", "string", "blocks"],
)
def test_count_keeps_the_callers_system_ahead_of_the_lifted_run(
system: JsonValue, expected_system: list[dict[str, JsonValue]]
) -> None:
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
model="claude-opus-5-5",
messages=[{"role": "system", "content": "Be terse"}, {"role": "user", "content": "hi"}],
system=system,
)
assert result == {
"model": "claude-opus-5-5",
"system": expected_system,
"messages": [{"role": "user", "content": "hi"}],
}
def test_count_leaves_a_non_text_system_and_its_messages_as_sent() -> None:
"""A malformed ``system`` is the provider's to reject, so nothing is rearranged around it."""
messages: Final[list[dict[str, JsonValue]]] = [
{"role": "system", "content": "Be terse"},
{"role": "user", "content": "hi"},
]
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
model="claude-opus-5-5", messages=messages, system=5
)
assert result == {"model": "claude-opus-5-5", "system": 5, "messages": messages}
def test_count_drops_a_leading_system_message_without_text() -> None:
result: Final = AnthropicCountTokensConfig().transform_request_to_count_tokens(
model="claude-opus-5-5",
messages=[{"role": "system", "content": ""}, {"role": "user", "content": "hi"}],
)
assert result == {"model": "claude-opus-5-5", "messages": [{"role": "user", "content": "hi"}]}
@pytest.mark.asyncio
async def test_handler_sends_the_leading_system_run_as_system_not_as_a_message(httpx_transport_clients):
"""The wire body is what the provider judges: ``system`` carries the prompt and no message has
``role: "system"``, so a Responses ``instructions`` is counted by Anthropic instead of 400ing."""
with respx.mock:
route = respx.post("https://gateway.example/v1/messages/count_tokens").mock(
return_value=httpx.Response(200, json={"input_tokens": 21})
)
result = await AnthropicCountTokensHandler().handle_count_tokens_request(
model="claude-opus-5-5",
messages=[{"role": "system", "content": "Be terse"}, {"role": "user", "content": "Hello, how are you?"}],
auth_header={"x-api-key": "sk-ant-api03-test-key"},
api_base="https://gateway.example",
)
assert result == {"input_tokens": 21}
assert TypeAdapter(dict[str, JsonValue]).validate_json(route.calls.last.request.content) == {
"model": "claude-opus-5-5",
"system": [{"type": "text", "text": "Be terse"}],
"messages": [{"role": "user", "content": "Hello, how are you?"}],
}