fix(bedrock): count tokens on bedrock-mantle when bedrock-runtime cannot count a Claude model (#45317)

This commit is contained in:
devin-ai-integration[bot] 2026-10-08 15:26:49 -07:00 • committed by GitHub
parent 3822947b0d
commit 577d74c1ce
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 1682 additions and 136 deletions

View file

@ -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)

View file

@ -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,

View file

@ -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}")

View file

@ -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", []),
)

View file

@ -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)

File diff suppressed because it is too large Load diff

View file

@ -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

View file

@ -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)

View file

@ -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"] == []

View file

@ -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