diff --git a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py index 2d02b152c61..ed61707f99e 100644 --- a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py +++ b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py @@ -3,27 +3,100 @@ Bedrock Token Counter implementation using the CountTokens API. """ from collections.abc import Mapping, Sequence +from dataclasses import dataclass from typing import Any, Final +from pydantic import JsonValue +from typing_extensions import assert_never + from litellm._logging import verbose_logger from litellm.llms.base_llm.base_utils import BaseTokenCounter from litellm.llms.bedrock.common_utils import BedrockError, get_bedrock_base_model from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler -from litellm.types.utils import LlmProviders, TokenCountResponse +from litellm.llms.bedrock.count_tokens.mantle_handler import BedrockMantleCountTokensHandler +from litellm.llms.bedrock_mantle.common_utils import is_mantle_claude_model +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.types.utils import LiteLLMPydanticObjectBase, LlmProviders, TokenCountResponse + +RUNTIME_TOKENIZER_TYPE: Final = "bedrock_api" +MANTLE_TOKENIZER_TYPE: Final = "bedrock_mantle_api" + + +class _CountTokensReply(LiteLLMPydanticObjectBase): + input_tokens: int + + +@dataclass(frozen=True, slots=True) +class CountedTokens: + input_tokens: int + original_response: Mapping[str, JsonValue] + tokenizer_type: str + + +@dataclass(frozen=True, slots=True) +class CountTokensFailure: + status_code: int + message: str + tokenizer_type: str + + +CountTokensOutcome = CountedTokens | CountTokensFailure + + +def _runtime_rejected_claude_model(outcome: CountTokensOutcome, resolved_model: str) -> bool: + return ( + isinstance(outcome, CountTokensFailure) + and outcome.status_code == 400 + and is_mantle_claude_model(resolved_model) + ) class BedrockTokenCounter(BaseTokenCounter): - """Token counter implementation for AWS Bedrock provider using the CountTokens API.""" + """Token counter for AWS Bedrock: bedrock-runtime CountTokens first, and Anthropic's count_tokens + on bedrock-mantle for the Claude models bedrock-runtime answers 400 for""" + + def __init__( + self, + runtime_handler: BedrockCountTokensHandler | None = None, + mantle_handler: BedrockMantleCountTokensHandler | None = None, + client: AsyncHTTPHandler | None = None, + ) -> None: + self._runtime_handler: Final = runtime_handler or BedrockCountTokensHandler() + self._mantle_handler: Final = mantle_handler or BedrockMantleCountTokensHandler() + self._client: Final = client def should_use_token_counting_api( self, custom_llm_provider: str | None = None, ) -> bool: - """ - Returns True if we should use the Bedrock CountTokens API for token counting. - """ return custom_llm_provider == LlmProviders.BEDROCK.value + async def _count_with( + self, + handler: BedrockCountTokensHandler | BedrockMantleCountTokensHandler, + tokenizer_type: str, + request_data: dict[str, object], + litellm_params: dict[str, object], + resolved_model: str, + ) -> CountTokensOutcome: + try: + result: Final = await handler.handle_count_tokens_request( + request_data=request_data, + litellm_params=litellm_params, + resolved_model=resolved_model, + client=self._client, + ) + reply: Final = _CountTokensReply.model_validate(result) + except BedrockError as e: + verbose_logger.debug( + "%s CountTokens API error: status=%s, message=%s", tokenizer_type, e.status_code, e.message + ) + return CountTokensFailure(status_code=e.status_code, message=e.message, tokenizer_type=tokenizer_type) + except Exception as e: + verbose_logger.debug("Error calling %s CountTokens API: %s", tokenizer_type, e) + return CountTokensFailure(status_code=500, message=str(e), tokenizer_type=tokenizer_type) + return CountedTokens(input_tokens=reply.input_tokens, original_response=result, tokenizer_type=tokenizer_type) + async def count_tokens( self, model_to_use: str, @@ -34,81 +107,52 @@ class BedrockTokenCounter(BaseTokenCounter): tools: Sequence[Mapping[str, object]] | None = None, system: object | None = None, ) -> TokenCountResponse | None: - """ - Count tokens using AWS Bedrock's CountTokens API. - - This method calls the existing BedrockCountTokensHandler to make an API call - to Bedrock's token counting endpoint, bypassing the local tiktoken-based counting. - - Args: - model_to_use: The model identifier - messages: The messages to count tokens for - contents: Alternative content format (not used for Bedrock) - deployment: Deployment configuration containing litellm_params - request_model: The original request model name - - Returns: - TokenCountResponse with token count, or None if counting fails - """ if not messages: return None - deployment = deployment or {} - litellm_params: Final = deployment.get("litellm_params", {}) - - # Build request data in the format expected by BedrockCountTokensHandler + litellm_params: Final = (deployment or {}).get("litellm_params", {}) request_data: Final[dict[str, object]] = { "model": model_to_use, "messages": messages, + **({"tools": tools} if tools else {}), + **({"system": system} if system else {}), } - - if tools: - request_data["tools"] = tools - - if system: - request_data["system"] = system - - # Get the resolved model (strip prefixes like bedrock/, converse/, etc.) resolved_model: Final = get_bedrock_base_model(model_to_use) - try: - handler: Final = BedrockCountTokensHandler() - result: Final = await handler.handle_count_tokens_request( - request_data=request_data, - litellm_params=litellm_params, - resolved_model=resolved_model, + runtime: Final = await self._count_with( + self._runtime_handler, RUNTIME_TOKENIZER_TYPE, request_data, litellm_params, resolved_model + ) + outcome: Final = ( + await self._count_with( + self._mantle_handler, MANTLE_TOKENIZER_TYPE, request_data, litellm_params, resolved_model ) - - # Transform response to TokenCountResponse - if result is not None: + if _runtime_rejected_claude_model(runtime, resolved_model) + else runtime + ) + match outcome: + case CountedTokens(): return TokenCountResponse( - total_tokens=result.get("input_tokens", 0), + total_tokens=outcome.input_tokens, request_model=request_model, model_used=model_to_use, - tokenizer_type="bedrock_api", - original_response=result, + tokenizer_type=outcome.tokenizer_type, + original_response=dict(outcome.original_response), ) - except BedrockError as e: - verbose_logger.warning("Bedrock CountTokens API error: status=%s, message=%s", e.status_code, e.message) - return TokenCountResponse( - total_tokens=0, - request_model=request_model, - model_used=model_to_use, - tokenizer_type="bedrock_api", - error=True, - error_message=e.message, - status_code=e.status_code, - ) - except Exception as e: - verbose_logger.warning("Error calling Bedrock CountTokens API: %s", e) - return TokenCountResponse( - total_tokens=0, - request_model=request_model, - model_used=model_to_use, - tokenizer_type="bedrock_api", - error=True, - error_message=str(e), - status_code=500, - ) - - return None + case CountTokensFailure(): + verbose_logger.warning( + "%s CountTokens API error: status=%s, message=%s", + outcome.tokenizer_type, + outcome.status_code, + outcome.message, + ) + return TokenCountResponse( + total_tokens=0, + request_model=request_model, + model_used=model_to_use, + tokenizer_type=outcome.tokenizer_type, + error=True, + error_message=outcome.message, + status_code=outcome.status_code, + ) + case _: + assert_never(outcome) diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py index d7fb510f057..be8243e48ef 100644 --- a/litellm/llms/bedrock/count_tokens/handler.py +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -125,7 +125,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig): raise except httpx.HTTPStatusError as e: # HTTP errors - preserve the actual status code - verbose_logger.error("HTTP error in CountTokens handler: %s", e) + verbose_logger.debug("HTTP error in CountTokens handler: %s", e) raise BedrockError( status_code=e.response.status_code, message=e.response.text, diff --git a/litellm/llms/bedrock/count_tokens/mantle_handler.py b/litellm/llms/bedrock/count_tokens/mantle_handler.py new file mode 100644 index 00000000000..23c3dc54c95 --- /dev/null +++ b/litellm/llms/bedrock/count_tokens/mantle_handler.py @@ -0,0 +1,99 @@ +from typing import Final + +import httpx +from pydantic import JsonValue, TypeAdapter +from typing_extensions import NotRequired, ReadOnly, TypedDict + +import litellm +from litellm._logging import verbose_logger +from litellm.llms.anthropic.count_tokens.transformation import AnthropicCountTokensConfig +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing +from litellm.llms.bedrock.common_utils import BedrockError, build_mantle_messages_url +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client + +MANTLE_COUNT_TOKENS_SUFFIX: Final = "/count_tokens" +MANTLE_ANTHROPIC_VERSION: Final = "2023-06-01" + + +class MantleCountTokensRequest(TypedDict): + messages: ReadOnly[list[dict[str, JsonValue]]] + system: ReadOnly[NotRequired[JsonValue]] + tools: ReadOnly[NotRequired[list[dict[str, JsonValue]]]] + + +_COUNT_REQUEST: Final = TypeAdapter(MantleCountTokensRequest) +_COUNT_RESPONSE: Final = TypeAdapter(dict[str, JsonValue]) + + +class BedrockMantleCountTokensHandler(BaseAWSLLM): + """Counts tokens through Anthropic's count_tokens on the bedrock-mantle endpoint. + + Claude models that Bedrock offers only through cross-region inference answer 400 on + bedrock-runtime's CountTokens; AWS documents Mantle's /anthropic/v1/messages/count_tokens + as the way to count them, with the base model id and the deployment's AWS credentials + """ + + def __init__(self, anthropic_config: AnthropicCountTokensConfig | None = None) -> None: + super().__init__() + self._anthropic_config: Final = anthropic_config or AnthropicCountTokensConfig() + + def get_mantle_count_tokens_endpoint(self, aws_region_name: str) -> str: + messages_url: Final = build_mantle_messages_url( + api_base=None, aws_bedrock_runtime_endpoint=None, region=aws_region_name + ) + return f"{messages_url}{MANTLE_COUNT_TOKENS_SUFFIX}" + + async def handle_count_tokens_request( + self, + request_data: dict[str, object], + litellm_params: dict[str, object], + resolved_model: str, + client: AsyncHTTPHandler | None = None, + ) -> dict[str, JsonValue]: + try: + request: Final = _COUNT_REQUEST.validate_python(request_data) + aws_region_name: Final = self._get_aws_region_name( + optional_params=litellm_params, model=resolved_model, model_id=None + ) + endpoint_url: Final = self.get_mantle_count_tokens_endpoint(aws_region_name) + body: Final = self._anthropic_config.transform_request_to_count_tokens( + model=resolved_model, + messages=request["messages"], + tools=request.get("tools"), + system=request.get("system"), + ) + verbose_logger.debug("Making bedrock-mantle count_tokens request to: %s", endpoint_url) + api_key: Final = litellm_params.get("api_key") + signed_headers, signed_body = await run_aws_signing( + self._sign_request, + service_name="bedrock", + headers={"Content-Type": "application/json", "anthropic-version": MANTLE_ANTHROPIC_VERSION}, + optional_params=litellm_params, + request_data=body, + api_base=endpoint_url, + model=resolved_model, + api_key=api_key if isinstance(api_key, str) else None, + ) + async_client: Final = client or get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK) + response: Final = await async_client.post( + endpoint_url, headers=signed_headers, data=signed_body, timeout=30.0 + ) + if response.status_code != 200: + raise BedrockError( + status_code=response.status_code, + message=response.text, + headers=response.headers, + response=response, + ) + return _COUNT_RESPONSE.validate_json(response.content) + except BedrockError: + raise + except httpx.HTTPStatusError as e: + raise BedrockError( + status_code=e.response.status_code, + message=e.response.text, + headers=e.response.headers, + response=e.response, + ) + except Exception as e: + raise BedrockError(status_code=500, message=f"bedrock-mantle count_tokens processing error: {e}") diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index 87d2ee07818..8638bc57615 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -191,26 +191,12 @@ async def google_count_tokens(request: Request, model_name: str): request=token_request, call_endpoint=True, ) - if token_response is not None: - # cast the response to the well known format - original_response: Final[dict] = token_response.original_response or {} - if original_response: - return TokenCountDetailsResponse( - totalTokens=original_response.get("totalTokens", 0), - promptTokensDetails=original_response.get("promptTokensDetails", []), - ) - else: - return TokenCountDetailsResponse( - totalTokens=token_response.total_tokens or 0, - promptTokensDetails=[], - ) - - ######################################################### - # Return the response in the well known format - ######################################################### + if token_response is None: + return TokenCountDetailsResponse(totalTokens=0, promptTokensDetails=[]) + original_response: Final[dict] = token_response.original_response or {} return TokenCountDetailsResponse( - totalTokens=0, - promptTokensDetails=[], + totalTokens=original_response.get("totalTokens") or token_response.total_tokens or 0, + promptTokensDetails=original_response.get("promptTokensDetails", []), ) diff --git a/tests/e2e/llm_translation/test_messages_count_tokens_bedrock_e2e.py b/tests/e2e/llm_translation/test_messages_count_tokens_bedrock_e2e.py new file mode 100644 index 00000000000..ee5cd54e4e0 --- /dev/null +++ b/tests/e2e/llm_translation/test_messages_count_tokens_bedrock_e2e.py @@ -0,0 +1,70 @@ +"""Live e2e: `/v1/messages/count_tokens` on a Bedrock Claude model that bedrock-runtime +cannot count. + +Claude Opus 4.8 is offered only through cross-region inference, and bedrock-runtime's +CountTokens answers 400 for it. The proxy then has to count through bedrock-mantle's +Anthropic count_tokens, and the answer must sit within a few percent of what `/v1/messages` +bills as `usage.input_tokens`. The local tokenizer fallback undercounts these models by +about 40%, so this is the line that proves the real count is served. Both calls go through +the real Anthropic SDK, the client customers count with +""" + +from __future__ import annotations + +from typing import Final + +import pytest +from anthropic.types import MessageParam +from e2e_config import unique_marker +from e2e_metadata import Domain, Provider, Route, Subject, meta +from lifecycle import ResourceManager +from models import LiteLLMParamsBody +from proxy_client import ProxyClient +from sdk_clients import NO_PROXY_CACHE, SdkClients + +pytestmark = pytest.mark.e2e + +CROSS_REGION_ONLY_CLAUDE_BACKEND: Final = "bedrock/global.anthropic.claude-opus-4-8" +COUNT_TOLERANCE: Final = 0.05 + + +def _register(proxy: ProxyClient, resources: ResourceManager, backend: str) -> tuple[str, str]: + model: Final = f"e2e-count-tokens-bedrock-{unique_marker()}" + model_id: Final = proxy.create_model( + model, + LiteLLMParamsBody( + model=backend, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + return model, resources.key() + + +class TestBedrockMessagesCountTokens: + @meta( + Subject( + domain=Domain.LLM_TRANSLATION, + route=Route.COUNT_TOKENS, + providers=(Provider.BEDROCK,), + models=(CROSS_REGION_ONLY_CLAUDE_BACKEND,), + ) + ) + def test_count_matches_billed_input_tokens_for_a_model_bedrock_runtime_cannot_count( + self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients + ) -> None: + model, key = _register(proxy, resources, CROSS_REGION_ONLY_CLAUDE_BACKEND) + client: Final = sdk.anthropic(key) + prompt: Final = f"{unique_marker()} " + "The quick brown fox jumps over the lazy dog. " * 40 + message: Final[MessageParam] = {"role": "user", "content": prompt} + + counted: Final = client.messages.count_tokens(model=model, messages=[message]) + answered: Final = client.messages.create( + model=model, max_tokens=1, messages=[message], extra_body=NO_PROXY_CACHE + ) + + billed: Final = answered.usage.input_tokens + assert billed, answered.usage + assert abs(counted.input_tokens - billed) <= billed * COUNT_TOLERANCE, (counted, answered.usage) diff --git a/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py new file mode 100644 index 00000000000..f9d4f7f91bb --- /dev/null +++ b/tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py @@ -0,0 +1,1072 @@ +import asyncio +import json +import re +import signal +import socket +import threading +import uuid +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import ExitStack +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final + +import anthropic +import httpx +import psutil +import pytest +import yaml +from anthropic.types import MessageCountTokensToolParam, MessageParam +from integration._support.bedrock_runtime_peer import answer, marker_of, target_of +from integration._support.bedrock_runtime_peer import respond as runtime_generation +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 pydantic import JsonValue, TypeAdapter + +pytestmark = pytest.mark.timeout(2 * graceful_stop_seconds() + 120) + +_OPUS: Final = "global.anthropic.claude-opus-4-8" +_OPUS_BASE: Final = "anthropic.claude-opus-4-8" +_SONNET: Final = "anthropic.claude-sonnet-4-6" +_NOVA: Final = "amazon.nova-lite-v1:0" +_REGION: Final = "us-east-1" +_OWNED_OPUS: Final = "mantle-opus" +_OWNED_SONNET: Final = "mantle-sonnet" +_OWNED_NOVA: Final = "mantle-nova" +_OWNED_MODELS: Final = MappingProxyType({_OWNED_OPUS: _OPUS, _OWNED_SONNET: _SONNET, _OWNED_NOVA: _NOVA}) +_ACCESS_KEY: Final = "AKIAINTEGRATIONMANTLE" +_SECRET_KEY: Final = "integration-mantle-secret" +_SIGV4_SCOPE: Final = f"/{_REGION}/bedrock/aws4_request" +_MANTLE_TARGET: Final = "/anthropic/v1/messages/count_tokens" +_MANTLE_VERSION: Final = "2023-06-01" +_MANTLE_COUNT: Final = 4242 +_RUNTIME_COUNT: Final = 1345 +_UNSUPPORTED: Final = "The provided model doesn't support counting tokens." +_REJECTION: Final = "scripted mantle rejection" +_COUNT_TARGET: Final = re.compile(r"^/model/(.+)/count-tokens$") +_INVOKE_TARGET: Final = re.compile(r"^/model/(.+)/invoke$") +_SCRIPTED_STATUS: Final = re.compile(r"status=(\d{3})") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_JSON_LIST: Final = TypeAdapter(list[JsonValue]) + +_SDK_MESSAGES: Final[list[MessageParam]] = [{"role": "user", "content": "Count this message"}] +_MESSAGES: Final = _JSON_LIST.validate_json(json.dumps(_SDK_MESSAGES)) +_SYSTEM: Final = "You are a terse assistant that answers in one sentence" +_SYSTEM_BLOCKS: Final[list[JsonValue]] = [ + {"type": "text", "text": "You are a terse assistant"}, + {"type": "text", "text": "Answer in one sentence"}, +] +_SDK_TOOLS: Final[list[MessageCountTokensToolParam]] = [ + { + "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"], + }, + } +] +_TOOLS: Final = _JSON_LIST.validate_json(json.dumps(_SDK_TOOLS)) +_GEMINI_BODY: Final[dict[str, JsonValue]] = {"contents": [{"role": "user", "parts": [{"text": "Count this"}]}]} +_GEMINI_MESSAGES: Final[list[JsonValue]] = [{"role": "user", "content": "Count this"}] + + +def _mantle_body(**fields: JsonValue) -> dict[str, JsonValue]: + return {"model": _OPUS_BASE, "messages": _MESSAGES, **fields} + + +_MANTLE_BARE: Final = _mantle_body() +_MANTLE_FULL: Final = _mantle_body(system=_SYSTEM, tools=_TOOLS) + + +def _json_reply(status: int, payload: Mapping[str, JsonValue]) -> Reply: + return Reply(status=status, body=json.dumps(payload).encode()) + + +def _runtime_count(request: Request, model: str) -> Reply: + scripted: Final = _SCRIPTED_STATUS.search(request.body.decode(errors="replace")) + if scripted is not None: + return _json_reply(int(scripted.group(1)), {"message": f"scripted {scripted.group(1)}"}) + if "sonnet" in model: + return _json_reply(200, {"inputTokens": _RUNTIME_COUNT}) + return _json_reply(400, {"message": _UNSUPPORTED}) + + +def _invoke_reply(marker: str) -> Reply: + return _json_reply( + 200, + { + "id": f"msg_{marker}", + "type": "message", + "role": "assistant", + "model": _OPUS_BASE, + "content": [{"type": "text", "text": answer(marker)}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 5, "output_tokens": 3}, + }, + ) + + +def _runtime(request: Request) -> Reply: + target: Final = target_of(request) + counted: Final = _COUNT_TARGET.match(target) + if counted is not None: + return _runtime_count(request, counted.group(1)) + if _INVOKE_TARGET.match(target): + return _invoke_reply(marker_of(request)) + return runtime_generation(request) + + +def _mantle_counted(_request: Request) -> Reply: + return _json_reply(200, {"input_tokens": _MANTLE_COUNT}) + + +def _rejected(status: int) -> Reply: + return _json_reply(status, {"type": "error", "error": {"type": "invalid_request_error", "message": _REJECTION}}) + + +def _rejecting(status: int) -> Callable[[Request], Reply]: + def count(_request: Request) -> Reply: + return _rejected(status) + + 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)) + ) + return _mantle_counted(request) if accepted else _rejected(400) + + +def _mantle(count: Callable[[Request], Reply] = _mantle_counted) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.target == _MANTLE_TARGET: + return count(request) + return _json_reply(404, {"error": f"unscripted mantle target {request.target}"}) + + return respond + + +def _mantle_environment(port: int) -> Mapping[str, str]: + return {"BEDROCK_MANTLE_API_BASE": f"http://127.0.0.1:{port}"} + + +_INHERITED_BEARER: Final = ("AWS_BEARER_TOKEN_BEDROCK",) + + +def _reserved_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +def _closed_port_url() -> str: + return f"http://127.0.0.1:{_reserved_port()}" + + +@pytest.fixture(scope="module") +def mantle_port() -> int: + return _reserved_port() + + +@pytest.fixture(scope="module") +def counting_proxy(tmp_path_factory: pytest.TempPathFactory, mantle_port: int) -> Iterator[OwnedProxy]: + with ( + gateway_from_environment() as gateway, + owned_proxy_process( + gateway, + tmp_path_factory.mktemp("mantle-count"), + _mantle_environment(mantle_port), + workers=1, + remove_environment=_INHERITED_BEARER, + ) as owned, + ): + yield owned + + +def _litellm_params(model: str, api_base: str) -> dict[str, JsonValue]: + return { + "model": f"bedrock/{model}", + "api_base": api_base, + "aws_access_key_id": _ACCESS_KEY, + "aws_secret_access_key": _SECRET_KEY, + "aws_region_name": _REGION, + } + + +def _owned_config(path: Path, runtime_url: str, settings: Mapping[str, JsonValue]) -> Path: + config: Final = object_value(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + litellm_settings: Final = object_value(config["litellm_settings"]) + path.write_text( + yaml.safe_dump( + { + **config, + "model_list": [ + {"model_name": name, "litellm_params": _litellm_params(model, runtime_url)} + for name, model in _OWNED_MODELS.items() + ], + "litellm_settings": {**litellm_settings, **settings}, + } + ) + ) + return path + + +def _deployment(scenario: Scenario, api_base: str, model: str = _OPUS) -> str: + return scenario.model(model_info=None, api_key=None, **_litellm_params(model, api_base)) + + +def _bare(model: str) -> dict[str, JsonValue]: + return {"model": model, "messages": _MESSAGES} + + +def _full(model: str) -> dict[str, JsonValue]: + return {**_bare(model), "system": _SYSTEM, "tools": _TOOLS} + + +def _count(gateway: Gateway, body: Mapping[str, JsonValue]) -> httpx.Response: + return gateway.request("POST", "/v1/messages/count_tokens", body) + + +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 ("bedrock_api", "bedrock_mantle_api"), response.text + assert isinstance(total, int) and total > 0 and total not in (_MANTLE_COUNT, _RUNTIME_COUNT), response.text + return total + + +def _assert_sigv4(request: Request) -> None: + authorization: Final = request.headers.get("authorization", "") + assert authorization.startswith(f"AWS4-HMAC-SHA256 Credential={_ACCESS_KEY}/"), request.headers + assert _SIGV4_SCOPE in authorization, authorization + + +def _mantle_requests(requests: Sequence[Request]) -> tuple[dict[str, JsonValue], ...]: + for request in requests: + assert (request.method, request.target) == ("POST", _MANTLE_TARGET), request.target + assert request.headers["anthropic-version"] == _MANTLE_VERSION, request.headers + assert request.headers["content-type"] == "application/json", request.headers + _assert_sigv4(request) + return tuple(_JSON_OBJECT.validate_json(request.body) for request in requests) + + +def _mantle_bodies(wire: Wire) -> tuple[dict[str, JsonValue], ...]: + return _mantle_requests(wire.drain()) + + +def _runtime_count_targets(wire: Wire) -> tuple[str, ...]: + counts: Final = tuple(request for request in wire.drain() if _COUNT_TARGET.match(target_of(request))) + for request in counts: + assert request.method == "POST", request.method + _assert_sigv4(request) + return tuple(target_of(request) for request in counts) + + +def _clients(stack: ExitStack, base_url: str, count: int, timeout: float = 30) -> tuple[httpx.Client, ...]: + return tuple( + stack.enter_context(httpx.Client(base_url=base_url, timeout=timeout, 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") + + +def _generated_then_counted( + client: httpx.Client, key: str, model: str, body: Mapping[str, JsonValue] +) -> tuple[int, int, JsonValue]: + generated: Final = client.post( + "/v1/messages", + json={ + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": f"Generate before counting marker-{uuid.uuid4().hex}"}], + }, + headers={"Authorization": f"Bearer {key}"}, + ) + return generated.status_code, *_counted_on(client, key, body) + + +def _counted_or_dropped(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue] | None: + try: + return _counted_on(client, key, body) + except httpx.TransportError: + return None + + +def _probed_then_counted_or_dropped( + client: httpx.Client, key: str, body: Mapping[str, JsonValue], probed: SimpleQueue[int] +) -> tuple[int, tuple[int, JsonValue] | None]: + port: Final = _local_port(client) + probed.put(port) + return port, _counted_or_dropped(client, key, body) + + +def _local_port(client: httpx.Client) -> int: + with client.stream("GET", "/health/liveliness") as response: + port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) + response.read() + assert response.status_code == 200, response.text + return port + + +def _accepted_client_ports(pid: int, proxy_port: int) -> frozenset[int]: + return frozenset( + connection.raddr.port + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.raddr and connection.laddr.port == proxy_port + ) + + +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 _mantle_counted(request) + + return hold + + +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) + + +@pytest.mark.parametrize("system", [_SYSTEM, _SYSTEM_BLOCKS], ids=["string", "blocks"]) +def test_messages_count_tokens_counts_through_mantle_when_the_runtime_cannot( + counting_proxy: OwnedProxy, mantle_port: int, system: JsonValue +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "system": system, "tools": _TOOLS}) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert _runtime_count_targets(runtime) == (f"/model/{_OPUS_BASE}/count-tokens",) + assert _mantle_bodies(mantle) == (_mantle_body(system=system, tools=_TOOLS),) + + +def test_messages_count_tokens_without_system_or_tools_sends_a_bare_body_to_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_utils_token_counter_call_endpoint_reports_the_mantle_tokenizer( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", "/utils/token_counter", _full(model), 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_FULL,) + + +def test_gemini_count_tokens_route_reports_the_mantle_total(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY) + assert _payload(response) == {"totalTokens": _MANTLE_COUNT, "promptTokensDetails": []}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == ({"model": _OPUS_BASE, "messages": _GEMINI_MESSAGES},) + + +def test_gemini_count_tokens_route_reports_the_runtime_total_for_a_model_bedrock_counts( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url, _SONNET) + response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY) + assert _payload(response) == {"totalTokens": _RUNTIME_COUNT, "promptTokensDetails": []}, response.text + assert _runtime_count_targets(runtime) == (f"/model/{_SONNET}/count-tokens",) + assert mantle.drain() == () + + +def test_responses_input_tokens_counts_through_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), 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": "Count this message"} + ) + 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_BARE,) + + +def test_anthropic_sdk_count_tokens_counts_through_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), 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=_SDK_MESSAGES, system=_SYSTEM, tools=_SDK_TOOLS + ) + assert counted.input_tokens == _MANTLE_COUNT, counted + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) + + +def test_async_anthropic_sdk_count_tokens_counts_through_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), 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=_SDK_MESSAGES, system=_SYSTEM) + ) + assert counted.input_tokens == _MANTLE_COUNT, counted + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_mantle_body(system=_SYSTEM),) + + +def test_model_the_runtime_counts_never_reaches_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url, _SONNET) + assert _payload(_count(gateway, _full(model))) == {"input_tokens": _RUNTIME_COUNT} + detailed: Final = gateway.request( + "POST", "/utils/token_counter", _full(model), params={"call_endpoint": "true"} + ) + payload: Final = _payload(detailed) + assert (payload["total_tokens"], payload["tokenizer_type"]) == (_RUNTIME_COUNT, "bedrock_api"), detailed.text + assert _runtime_count_targets(runtime) == (f"/model/{_SONNET}/count-tokens",) * 2 + assert mantle.drain() == () + + +def test_non_claude_model_the_runtime_cannot_count_falls_back_locally_without_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url, _NOVA) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + assert _runtime_count_targets(runtime) == (f"/model/{_NOVA}/count-tokens",) + assert mantle.drain() == () + + +def test_runtime_403_on_a_claude_model_never_reaches_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + body: Final[dict[str, JsonValue]] = { + "model": model, + "messages": [{"role": "user", "content": "status=403 Count this message"}], + } + response: Final = _count(gateway, body) + assert _payload(response) == {"input_tokens": _local_count(gateway, body)}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert mantle.drain() == () + + +def test_unreachable_runtime_falls_back_locally_without_mantle(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with wire_server(_mantle(), port=mantle_port) as mantle, gateway.scenario() as scenario: + model: Final = _deployment(scenario, _closed_port_url()) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + assert mantle.drain() == () + + +def test_bedrock_passthrough_count_tokens_still_answers_the_runtime_rejection( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request("POST", "/bedrock/v1/messages/count_tokens", _full(model)) + assert response.status_code == 400, response.text + assert _UNSUPPORTED in response.text and "input_tokens" not in response.text, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert mantle.drain() == () + + +def test_responses_input_tokens_with_instructions_still_counts_locally( + 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": "Count this message", "instructions": "Be terse"}, + ) + payload: Final = _payload(response) + 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 + + +@pytest.mark.parametrize("status", [400, 403, 404, 500, 503]) +def test_messages_count_tokens_falls_back_locally_when_mantle_errors( + counting_proxy: OwnedProxy, mantle_port: int, status: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(_rejecting(status)), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) + + +@pytest.mark.parametrize( + "body", + [b'{"inputTokens": 7}', b"not json at all", b"{}"], + ids=["runtime_key", "not_json", "empty_object"], +) +def test_messages_count_tokens_falls_back_locally_when_mantle_answers_without_input_tokens( + counting_proxy: OwnedProxy, mantle_port: int, body: bytes +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(lambda _request: Reply(body=body)), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_messages_count_tokens_falls_back_locally_when_mantle_is_unreachable(counting_proxy: OwnedProxy) -> None: + gateway: Final = counting_proxy.gateway + with wire_server(_runtime) as runtime, gateway.scenario() as scenario: + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + + +def test_messages_count_tokens_falls_back_locally_when_mantle_never_answers( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + + def hold(request: Request) -> Reply: + held.put(request.target) + release.wait(timeout=120) + return Reply(drop_connection=True) + + with ExitStack() as stack: + (client,) = _clients(stack, str(gateway.client.base_url), 1, timeout=120) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=1)) + runtime: Final = stack.enter_context(wire_server(_runtime)) + mantle: Final = stack.enter_context(wire_server(_mantle(hold), port=mantle_port)) + stack.callback(release.set) + scenario: Final = stack.enter_context(gateway.scenario()) + model: Final = _deployment(scenario, runtime.url) + body: Final = _full(model) + local: Final = _local_count(gateway, body) + future: Final = pool.submit(_counted_on, client, gateway.key, body) + eventually(held.qsize, lambda size: size == 1, seconds=30) + assert gateway.request("GET", "/health/liveliness").status_code == 200 + assert _local_count(gateway, body) == local + assert not future.done() + assert future.result(timeout=120) == (200, local) + assert len(_runtime_count_targets(runtime)) == 1 + assert len(mantle.drain()) == 1 + + +def test_messages_count_tokens_falls_back_locally_when_mantle_rejects_a_non_text_system( + 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, {**_bare(model), "system": 5}) + assert _payload(response) == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_mantle_body(system=5),) + + +@pytest.mark.parametrize( + "tools", + [5, "", "x" * 5120, ["get_weather"]], + ids=["int", "empty_string", "5kb_string", "list_of_strings"], +) +def test_messages_count_tokens_rejects_malformed_tools_without_calling_either_peer( + counting_proxy: OwnedProxy, mantle_port: int, tools: JsonValue +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + refused: Final = _count(gateway, {**_bare(model), "tools": tools}) + assert 400 <= refused.status_code < 600, refused.text + assert "input_tokens" not in refused.text, refused.text + assert runtime.drain() == () and mantle.drain() == () + assert _generated_then_counted(gateway.client, gateway.key, model, _bare(model)) == (200, 200, _MANTLE_COUNT) + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +@pytest.mark.parametrize("fields", [{}, {"messages": []}], ids=["missing", "empty"]) +def test_messages_count_tokens_without_messages_is_rejected( + counting_proxy: OwnedProxy, mantle_port: int, fields: dict[str, JsonValue] +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {"model": model, "system": _SYSTEM, "tools": _TOOLS, **fields}) + assert response.status_code == 400, response.text + assert "messages parameter is required" in response.text, response.text + assert runtime.drain() == () and mantle.drain() == () + + +def test_messages_count_tokens_unauthenticated_request_never_reaches_either_peer( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", "/v1/messages/count_tokens", _full(model), key="sk-not-a-key-this-proxy-issued" + ) + assert response.status_code == 401, response.text + assert "input_tokens" not in response.text, response.text + assert runtime.drain() == () and mantle.drain() == () + + +def test_messages_count_tokens_forwards_a_5kb_system_verbatim_to_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + system: Final = "Answer in one sentence. " * 214 + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "system": system}) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_mantle_body(system=system),) + + +def test_messages_count_tokens_duplicate_system_and_tools_keys_forward_one_value_each_to_mantle( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + fields: Final = f'"system": {json.dumps(_SYSTEM)}, "tools": {json.dumps(_TOOLS)}' + response: Final = gateway.client.post( + "/v1/messages/count_tokens", + content=f'{{"model": "{model}", "messages": {json.dumps(_MESSAGES)}, {fields}, {fields}}}', + 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 + (sent,) = mantle.drain() + assert _mantle_requests((sent,)) == (_MANTLE_FULL,) + assert (sent.body.count(b'"system"'), sent.body.count(b'"tools"')) == (1, 1), sent.body + + +def test_messages_count_tokens_leaves_empty_tools_and_system_out_of_the_mantle_body( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = _count(gateway, {**_bare(model), "tools": [], "system": ""}) + assert _payload(response) == {"input_tokens": _MANTLE_COUNT}, response.text + assert len(_runtime_count_targets(runtime)) == 1 + assert _mantle_bodies(mantle) == (_MANTLE_BARE,) + + +def test_messages_count_tokens_repeated_request_reaches_both_peers_each_time( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + answers: Final = tuple(_payload(_count(gateway, _bare(model))) 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_BARE,) * 2 + + +def test_disabled_token_counter_surfaces_the_mantle_error_instead_of_counting_locally( + 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.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(_rejecting(403)), port=mantle_port) as refusing: + refused: Final = _count(owned.gateway, _full(_OWNED_OPUS)) + assert refused.status_code == 403, refused.text + assert _REJECTION in refused.text and "input_tokens" not in refused.text, refused.text + assert _mantle_bodies(refusing) == (_MANTLE_FULL,) + with wire_server(_mantle(), port=mantle_port) as counting: + assert _payload(_count(owned.gateway, _full(_OWNED_OPUS))) == {"input_tokens": _MANTLE_COUNT} + assert _mantle_bodies(counting) == (_MANTLE_FULL,) + assert _payload(_count(owned.gateway, _full(_OWNED_SONNET))) == {"input_tokens": _RUNTIME_COUNT} + rejected: Final = _count(owned.gateway, _full(_OWNED_NOVA)) + assert rejected.status_code == 400, rejected.text + assert _UNSUPPORTED in rejected.text and "input_tokens" not in rejected.text, rejected.text + assert counting.drain() == () + assert len(_runtime_count_targets(runtime)) == 4 + + +def test_mantle_outage_between_concurrent_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 = _full(model) + local: Final = _local_count(gateway, body) + + def generate_then_count(client: httpx.Client) -> tuple[int, int, JsonValue]: + return _generated_then_counted(client, gateway.key, model, body) + + def count_only(client: httpx.Client) -> tuple[int, JsonValue]: + return _counted_on(client, gateway.key, body) + + with wire_server(_mantle(), port=mantle_port) as mantle: + assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _MANTLE_COUNT),) * len(clients) + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) * len(clients) + outage: Final = tuple(pool.map(count_only, clients)) + assert outage == ((200, local),) * len(clients) + with wire_server(_mantle(), port=mantle_port) as revived: + assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _MANTLE_COUNT),) * len(clients) + assert _mantle_bodies(revived) == (_MANTLE_FULL,) * len(clients) + assert len(_runtime_count_targets(runtime)) == 3 * len(clients) + + +def test_slow_mantle_holds_concurrent_counts_without_stalling_the_proxy( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + with ExitStack() as stack: + clients: Final = _clients(stack, str(gateway.client.base_url), 6) + runtime: Final = stack.enter_context(wire_server(_runtime)) + mantle: Final = stack.enter_context(wire_server(_mantle(_holding(held, release, 20)), port=mantle_port)) + scenario: Final = stack.enter_context(gateway.scenario()) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + stack.callback(release.set) + model: Final = _deployment(scenario, runtime.url) + futures: Final = tuple( + pool.submit(_generated_then_counted, client, gateway.key, model, _full(model)) 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, _full(model)) > 0 + assert not any(future.done() for future in futures) + release.set() + assert tuple(future.result(timeout=30) for future in futures) == ((200, 200, _MANTLE_COUNT),) * len(clients) + assert _mantle_bodies(mantle) == (_MANTLE_FULL,) * len(clients) + assert len(_runtime_count_targets(runtime)) == len(clients) + + +def test_worker_sigkill_mid_burst_leaves_the_sibling_counting( + gateway: Gateway, mantle_port: int, tmp_path: Path +) -> None: + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + with ExitStack() as stack: + runtime: Final = stack.enter_context(wire_server(_runtime)) + mantle: Final = stack.enter_context(wire_server(_mantle(_holding(held, release, 60)), port=mantle_port)) + config: Final = _owned_config(tmp_path / "worker-kill.yaml", runtime.url, {}) + owned: Final = stack.enter_context( + owned_proxy_process( + gateway, + tmp_path, + _mantle_environment(mantle_port), + config=config, + workers=2, + remove_environment=_INHERITED_BEARER, + ) + ) + body: Final = _full(_OWNED_OPUS) + proxy_url: Final = owned.gateway.client.base_url + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + clients: Final = _clients(stack, str(proxy_url), 12) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + stack.callback(release.set) + probed: Final[SimpleQueue[int]] = SimpleQueue() + futures: Final[tuple[Future[tuple[int, tuple[int, JsonValue] | None]], ...]] = tuple( + pool.submit(_probed_then_counted_or_dropped, client, owned.gateway.key, body, probed) for client in clients + ) + eventually(held.qsize, lambda size: size == len(clients), seconds=30) + ports: Final = frozenset(probed.get_nowait() for _ in clients) + shares: Final = {pid: _accepted_client_ports(pid, proxy_url.port or 0) & ports for pid in workers} + assert sum(map(len, shares.values())) == len(clients), shares + victim: Final = min((pid for pid in workers if shares[pid]), key=lambda pid: len(shares[pid])) + psutil.Process(victim).send_signal(signal.SIGKILL) + release.set() + for port, result in (future.result(timeout=60) for future in futures): + assert result == (None if port in shares[victim] else (200, _MANTLE_COUNT)), (port, result, shares) + second_wave: Final = _clients(stack, str(proxy_url), 6) + assert tuple(_counted_on(client, owned.gateway.key, body) for client in second_wave) == ( + (200, _MANTLE_COUNT), + ) * len(second_wave) + assert len(_mantle_bodies(mantle)) == len(clients) + len(second_wave) + eventually(lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda started: started >= 3, 120) + assert owned.process.poll() is None + + +def _streamed_text(text: str) -> str: + events: Final = tuple( + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in text.splitlines() + if line.startswith("data: {") + ) + return "".join(_delta_content(event) for event in events) + + +def _delta_content(event: dict[str, JsonValue]) -> str: + choices: Final = event.get("choices") + if not isinstance(choices, list) or not choices: + return "" + delta: Final = object_value(choices[0]).get("delta") + if not isinstance(delta, dict): + return "" + content: Final = delta.get("content") + return content if isinstance(content, str) else "" + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_chat_completions_on_the_same_deployment_still_generate( + counting_proxy: OwnedProxy, mantle_port: int, stream: bool +) -> None: + gateway: Final = counting_proxy.gateway + marker: Final = uuid.uuid4().hex + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": f"chat control marker-{marker}"}], + "stream": stream, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + generated: Final = _streamed_text(response.text) if stream else response.text + assert answer(marker) in generated, response.text + assert not stream or response.text.rstrip().endswith("data: [DONE]"), response.text + (sent,) = runtime.drain() + assert target_of(sent) == f"/model/{_OPUS}/{'converse-stream' if stream else 'converse'}", sent.target + assert mantle.drain() == () + + +def test_messages_endpoint_on_the_same_deployment_still_generates(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + marker: Final = uuid.uuid4().hex + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": f"marker-{marker}"}]}, + ) + assert response.status_code == 200, response.text + assert answer(marker) in response.text, response.text + (sent,) = runtime.drain() + assert target_of(sent) == f"/model/{_OPUS}/invoke", sent.target + assert mantle.drain() == () + + +def test_responses_endpoint_on_the_same_deployment_still_generates( + counting_proxy: OwnedProxy, mantle_port: int +) -> None: + gateway: Final = counting_proxy.gateway + marker: Final = uuid.uuid4().hex + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": f"marker-{marker}"}) + assert response.status_code == 200, response.text + assert answer(marker) in response.text, response.text + (sent,) = runtime.drain() + assert target_of(sent) == f"/model/{_OPUS}/converse", sent.target + assert mantle.drain() == () + + +def test_utils_token_counter_local_mode_never_calls_either_peer(counting_proxy: OwnedProxy, mantle_port: int) -> None: + gateway: Final = counting_proxy.gateway + with ( + wire_server(_runtime) as runtime, + wire_server(_mantle(), port=mantle_port) as mantle, + gateway.scenario() as scenario, + ): + model: Final = _deployment(scenario, runtime.url) + assert _local_count(gateway, _full(model)) > 0 + assert runtime.drain() == () and mantle.drain() == () diff --git a/tests/unit/llms/bedrock/count_tokens/test_bedrock_mantle_count_tokens_handler.py b/tests/unit/llms/bedrock/count_tokens/test_bedrock_mantle_count_tokens_handler.py new file mode 100644 index 00000000000..71b451f834c --- /dev/null +++ b/tests/unit/llms/bedrock/count_tokens/test_bedrock_mantle_count_tokens_handler.py @@ -0,0 +1,109 @@ +import json +from collections.abc import Mapping +from typing import Final + +import httpx +import pytest + +from litellm.llms.bedrock.common_utils import BedrockError +from litellm.llms.bedrock.count_tokens.mantle_handler import BedrockMantleCountTokensHandler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + +MANTLE_COUNT_URL: Final = "https://bedrock-mantle.us-east-1.api.aws/anthropic/v1/messages/count_tokens" +LITELLM_PARAMS: Final = { + "aws_access_key_id": "AKIATESTACCESSKEY", + "aws_secret_access_key": "test-secret", + "aws_region_name": "us-east-1", +} +REQUEST: Final[dict[str, object]] = { + "model": "global.anthropic.claude-opus-4-8", + "messages": [{"role": "user", "content": "The quick brown fox"}], + "system": "You are a terse assistant.", + "tools": [{"name": "get_weather", "input_schema": {"type": "object", "properties": {}}}], +} + + +class _MantleEndpoint: + def __init__(self, status_code: int, body: Mapping[str, object]) -> None: + self.status_code: Final = status_code + self.body: Final = body + self.requests: tuple[httpx.Request, ...] = () + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests = (*self.requests, request) + return httpx.Response(self.status_code, json=dict(self.body), request=request) + + def only_request(self) -> httpx.Request: + assert len(self.requests) == 1, self.requests + return self.requests[0] + + +def _client(endpoint: _MantleEndpoint) -> AsyncHTTPHandler: + return AsyncHTTPHandler(transport=httpx.MockTransport(endpoint)) + + +@pytest.fixture(autouse=True) +def _sigv4_only(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + + +@pytest.mark.asyncio +async def test_counts_on_mantle_with_the_base_model_and_the_deployment_credentials() -> None: + mantle: Final = _MantleEndpoint(200, {"input_tokens": 2177}) + + result: Final = await BedrockMantleCountTokensHandler().handle_count_tokens_request( + request_data=dict(REQUEST), + litellm_params=dict(LITELLM_PARAMS), + resolved_model="anthropic.claude-opus-4-8", + client=_client(mantle), + ) + + assert result == {"input_tokens": 2177} + posted: Final = mantle.only_request() + assert str(posted.url) == MANTLE_COUNT_URL + assert posted.headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "us-east-1/bedrock/aws4_request" in posted.headers["Authorization"] + assert posted.headers["anthropic-version"] == "2023-06-01" + assert json.loads(posted.content) == { + "model": "anthropic.claude-opus-4-8", + "messages": REQUEST["messages"], + "system": REQUEST["system"], + "tools": REQUEST["tools"], + } + + +@pytest.mark.asyncio +async def test_bedrock_mantle_api_base_env_names_the_host(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("BEDROCK_MANTLE_API_BASE", "https://vpce-abc.bedrock-mantle.us-east-1.vpce.api.aws") + mantle: Final = _MantleEndpoint(200, {"input_tokens": 3}) + + await BedrockMantleCountTokensHandler().handle_count_tokens_request( + request_data=dict(REQUEST), + litellm_params=dict(LITELLM_PARAMS), + resolved_model="anthropic.claude-opus-4-8", + client=_client(mantle), + ) + + assert ( + str(mantle.only_request().url) + == "https://vpce-abc.bedrock-mantle.us-east-1.vpce.api.aws/anthropic/v1/messages/count_tokens" + ) + + +@pytest.mark.asyncio +async def test_non_200_answers_raise_bedrock_error_with_mantle_status() -> None: + mantle: Final = _MantleEndpoint( + 404, {"type": "error", "error": {"type": "not_found_error", "message": "does not exist"}} + ) + + with pytest.raises(BedrockError) as raised: + await BedrockMantleCountTokensHandler().handle_count_tokens_request( + request_data=dict(REQUEST), + litellm_params=dict(LITELLM_PARAMS), + resolved_model="anthropic.claude-sonnet-5-5", + client=_client(mantle), + ) + + assert raised.value.status_code == 404 + assert "does not exist" in raised.value.message diff --git a/tests/unit/llms/bedrock/count_tokens/test_bedrock_token_counter.py b/tests/unit/llms/bedrock/count_tokens/test_bedrock_token_counter.py new file mode 100644 index 00000000000..a3dd8021f73 --- /dev/null +++ b/tests/unit/llms/bedrock/count_tokens/test_bedrock_token_counter.py @@ -0,0 +1,140 @@ +from collections.abc import Mapping +from typing import Final + +import httpx +import pytest +from pydantic import JsonValue, TypeAdapter + +from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + +RUNTIME_HOST: Final = "bedrock-runtime.us-east-1.amazonaws.com" +MANTLE_HOST: Final = "bedrock-mantle.us-east-1.api.aws" +UNSUPPORTED: Final = {"message": "The provided model doesn't support counting tokens."} +DEPLOYMENT: Final = { + "litellm_params": { + "aws_access_key_id": "AKIATESTACCESSKEY", + "aws_secret_access_key": "test-secret", + "aws_region_name": "us-east-1", + } +} +MESSAGES: Final = [{"role": "user", "content": "The quick brown fox jumps over the lazy dog."}] +_JSON_BODY: Final = TypeAdapter(dict[str, JsonValue]) + + +class _Bedrock: + def __init__(self, runtime: tuple[int, Mapping[str, object]], mantle: tuple[int, Mapping[str, object]]) -> None: + self.runtime: Final = runtime + self.mantle: Final = mantle + self.requests: tuple[httpx.Request, ...] = () + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests = (*self.requests, request) + status_code, body = self.runtime if request.url.host == RUNTIME_HOST else self.mantle + return httpx.Response(status_code, json=dict(body), request=request) + + def posted_hosts(self) -> tuple[str, ...]: + return tuple(request.url.host for request in self.requests) + + +def _counter(bedrock: _Bedrock) -> BedrockTokenCounter: + return BedrockTokenCounter(client=AsyncHTTPHandler(transport=httpx.MockTransport(bedrock))) + + +@pytest.fixture(autouse=True) +def _sigv4_only(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + monkeypatch.delenv("BEDROCK_MANTLE_API_BASE", raising=False) + + +@pytest.mark.asyncio +async def test_claude_model_bedrock_runtime_cannot_count_is_counted_on_mantle() -> None: + bedrock: Final = _Bedrock(runtime=(400, UNSUPPORTED), mantle=(200, {"input_tokens": 2177})) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use="global.anthropic.claude-opus-4-8", + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model="claude-opus-4-8", + system="You are a terse assistant.", + ) + + assert result is not None + assert result.error is False + assert result.total_tokens == 2177 + assert result.tokenizer_type == "bedrock_mantle_api" + assert result.original_response == {"input_tokens": 2177} + assert bedrock.posted_hosts() == (RUNTIME_HOST, MANTLE_HOST) + mantle_body: Final = _JSON_BODY.validate_json(bedrock.requests[1].content) + assert mantle_body["model"] == "anthropic.claude-opus-4-8" + assert mantle_body["system"] == "You are a terse assistant." + + +@pytest.mark.asyncio +async def test_bedrock_runtime_count_is_kept_when_it_answers() -> None: + bedrock: Final = _Bedrock(runtime=(200, {"inputTokens": 1353}), mantle=(200, {"input_tokens": 1336})) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use="global.anthropic.claude-sonnet-4-6", + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model="claude-sonnet-4-6", + ) + + assert result is not None + assert result.total_tokens == 1353 + assert result.tokenizer_type == "bedrock_api" + assert bedrock.posted_hosts() == (RUNTIME_HOST,) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model_to_use", "runtime"), + ( + ("global.anthropic.claude-opus-4-8", (403, {"Message": "not authorized to perform: bedrock:CountTokens"})), + ("amazon.nova-pro-v1:0", (400, UNSUPPORTED)), + ), +) +async def test_other_bedrock_runtime_failures_are_not_retried_on_mantle( + model_to_use: str, runtime: tuple[int, Mapping[str, object]] +) -> None: + bedrock: Final = _Bedrock(runtime=runtime, mantle=(200, {"input_tokens": 2177})) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use=model_to_use, + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model=model_to_use, + ) + + assert result is not None + assert result.error is True + assert result.status_code == runtime[0] + assert result.tokenizer_type == "bedrock_api" + assert bedrock.posted_hosts() == (RUNTIME_HOST,) + + +@pytest.mark.asyncio +async def test_mantle_failure_is_reported_with_its_status() -> None: + bedrock: Final = _Bedrock( + runtime=(400, UNSUPPORTED), + mantle=(404, {"type": "error", "error": {"type": "not_found_error", "message": "does not exist"}}), + ) + + result: Final = await _counter(bedrock).count_tokens( + model_to_use="global.anthropic.claude-sonnet-5-5", + messages=MESSAGES, + contents=None, + deployment=DEPLOYMENT, + request_model="claude-sonnet-5-5", + ) + + assert result is not None + assert result.error is True + assert result.status_code == 404 + assert result.tokenizer_type == "bedrock_mantle_api" + assert "does not exist" in (result.error_message or "") + assert bedrock.posted_hosts() == (RUNTIME_HOST, MANTLE_HOST) diff --git a/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py b/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py index e4cd7d9dfa8..451792109de 100644 --- a/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py +++ b/tests/unit/proxy/google_endpoints/test_google_api_endpoints.py @@ -3,6 +3,7 @@ Test to verify the Google GenAI proxy API endpoints """ +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -179,3 +180,30 @@ def test_google_count_tokens_unchanged(): assert response.status_code == 200 body = response.json() assert body["totalTokens"] == 7 + + +def test_google_count_tokens_uses_normalized_total_when_provider_shape_differs(): + """A provider that counts in Anthropic shape (Bedrock, Anthropic) carries no totalTokens key, so the normalized total must win over the raw response's missing field.""" + try: + client: Final = _build_test_client() + except ImportError as e: + pytest.skip(f"Skipping test due to missing dependency: {e}") + + fake_response: Final = MagicMock() + fake_response.original_response = {"input_tokens": 2167} + fake_response.total_tokens = 2167 + + with patch( + "litellm.proxy.proxy_server.token_counter", + new_callable=AsyncMock, + return_value=fake_response, + ): + response: Final = client.post( + "/v1beta/models/claude-opus-4-8:countTokens", + json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]}, + ) + + assert response.status_code == 200 + body: Final = response.json() + assert body["totalTokens"] == 2167 + assert body["promptTokensDetails"] == [] diff --git a/tests/unit/proxy/test_proxy_token_counter.py b/tests/unit/proxy/test_proxy_token_counter.py index e7a32816464..e4c4ddce64f 100644 --- a/tests/unit/proxy/test_proxy_token_counter.py +++ b/tests/unit/proxy/test_proxy_token_counter.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest from dotenv import load_dotenv +from pydantic import JsonValue load_dotenv() @@ -20,6 +21,7 @@ from litellm import Router from litellm.llms.bedrock.common_utils import BedrockError from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import ProxyException, TokenCountRequest from litellm.proxy.anthropic_endpoints.endpoints import ( count_tokens as anthropic_count_tokens, @@ -753,41 +755,45 @@ def test_vertex_ai_partner_models_token_counting_endpoint(vertex_location): ) +class _FailingRuntimeHandler(BedrockCountTokensHandler): + def __init__(self, error: Exception) -> None: + super().__init__() + self._error = error + + async def handle_count_tokens_request( + self, + request_data: dict[str, object], + litellm_params: dict[str, object], + resolved_model: str, + client: AsyncHTTPHandler | None = None, + ) -> dict[str, JsonValue]: + raise self._error + + @pytest.mark.asyncio async def test_bedrock_token_counter_error_propagation_bedrock_error(): """ Test that BedrockTokenCounter properly returns error response when BedrockError is raised. Verifies that the status code and error message are preserved. """ - counter = BedrockTokenCounter() + counter = BedrockTokenCounter( + runtime_handler=_FailingRuntimeHandler(BedrockError(status_code=429, message="Rate limit exceeded")) + ) - # Mock the handler to raise BedrockError with specific status code - with patch.object( - counter, "count_tokens", wraps=counter.count_tokens - ) as mock_count: - # We need to patch at the handler level - with patch( - "litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler" - ) as MockHandler: - mock_handler_instance = MockHandler.return_value - mock_handler_instance.handle_count_tokens_request = AsyncMock( - side_effect=BedrockError(status_code=429, message="Rate limit exceeded") - ) + result = await counter.count_tokens( + model_to_use="anthropic.claude-3-sonnet", + messages=[{"role": "user", "content": "hello"}], + contents=None, + deployment={"litellm_params": {}}, + request_model="bedrock/anthropic.claude-3-sonnet", + ) - result = await counter.count_tokens( - model_to_use="anthropic.claude-3-sonnet", - messages=[{"role": "user", "content": "hello"}], - contents=None, - deployment={"litellm_params": {}}, - request_model="bedrock/anthropic.claude-3-sonnet", - ) - - assert result is not None - assert result.error is True - assert result.status_code == 429 - assert "Rate limit exceeded" in result.error_message - assert result.tokenizer_type == "bedrock_api" - assert result.total_tokens == 0 + assert result is not None + assert result.error is True + assert result.status_code == 429 + assert "Rate limit exceeded" in result.error_message + assert result.tokenizer_type == "bedrock_api" + assert result.total_tokens == 0 @pytest.mark.asyncio @@ -795,28 +801,20 @@ async def test_bedrock_token_counter_error_propagation_generic_exception(): """ Test that BedrockTokenCounter returns error response with 500 status for generic exceptions. """ - counter = BedrockTokenCounter() + counter = BedrockTokenCounter(runtime_handler=_FailingRuntimeHandler(Exception("Unexpected error"))) - with patch( - "litellm.llms.bedrock.count_tokens.bedrock_token_counter.BedrockCountTokensHandler" - ) as MockHandler: - mock_handler_instance = MockHandler.return_value - mock_handler_instance.handle_count_tokens_request = AsyncMock( - side_effect=Exception("Unexpected error") - ) + result = await counter.count_tokens( + model_to_use="anthropic.claude-3-sonnet", + messages=[{"role": "user", "content": "hello"}], + contents=None, + deployment={"litellm_params": {}}, + request_model="bedrock/anthropic.claude-3-sonnet", + ) - result = await counter.count_tokens( - model_to_use="anthropic.claude-3-sonnet", - messages=[{"role": "user", "content": "hello"}], - contents=None, - deployment={"litellm_params": {}}, - request_model="bedrock/anthropic.claude-3-sonnet", - ) - - assert result is not None - assert result.error is True - assert result.status_code == 500 - assert "Unexpected error" in result.error_message + assert result is not None + assert result.error is True + assert result.status_code == 500 + assert "Unexpected error" in result.error_message @pytest.mark.asyncio