From d34ad3528191e2d338b6f8f982cf97482ef52bf3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 8 Oct 2026 22:34:25 -0700 Subject: [PATCH] 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> --- .../mid_conversation_system.py | 16 +- .../anthropic/count_tokens/transformation.py | 40 +- .../providers/_count_tokens_system_lift.py | 159 +++++ ...anthropic_count_tokens_system_lift_wire.py | 572 ++++++++++++++++++ .../test_bedrock_mantle_count_tokens_wire.py | 368 ++++++++++- ...rompt_templates_mid_conversation_system.py | 21 + ...t_anthropic_count_tokens_transformation.py | 112 ++++ 7 files changed, 1254 insertions(+), 34 deletions(-) create mode 100644 tests/integration/providers/_count_tokens_system_lift.py create mode 100644 tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py diff --git a/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py b/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py index 2169dfcad39..cb4bbf00c85 100644 --- a/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py +++ b/litellm/litellm_core_utils/prompt_templates/mid_conversation_system.py @@ -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] = { diff --git a/litellm/llms/anthropic/count_tokens/transformation.py b/litellm/llms/anthropic/count_tokens/transformation.py index e745e7dcd19..8e2688d4360 100644 --- a/litellm/llms/anthropic/count_tokens/transformation.py +++ b/litellm/llms/anthropic/count_tokens/transformation.py @@ -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} diff --git a/tests/integration/providers/_count_tokens_system_lift.py b/tests/integration/providers/_count_tokens_system_lift.py new file mode 100644 index 00000000000..56619d1c963 --- /dev/null +++ b/tests/integration/providers/_count_tokens_system_lift.py @@ -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), + ) diff --git a/tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py b/tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py new file mode 100644 index 00000000000..13091ccd14a --- /dev/null +++ b/tests/integration/providers/test_anthropic_count_tokens_system_lift_wire.py @@ -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 diff --git a/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py index f9d4f7f91bb..933fc49c149 100644 --- a/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py @@ -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]) diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py index d1e23a17747..9e3704a9df0 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_mid_conversation_system.py @@ -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"}}, + ) diff --git a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py index 2a31ac75d1d..9af7497bc82 100644 --- a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py +++ b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py @@ -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?"}], + }