mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(bedrock): count tokens on bedrock-mantle when bedrock-runtime cannot count a Claude model (#45317)
This commit is contained in:
parent
3822947b0d
commit
577d74c1ce
10 changed files with 1682 additions and 136 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
99
litellm/llms/bedrock/count_tokens/mantle_handler.py
Normal file
99
litellm/llms/bedrock/count_tokens/mantle_handler.py
Normal 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}")
|
||||
|
|
@ -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", []),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
1072
tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py
Normal file
1072
tests/integration/providers/test_bedrock_mantle_count_tokens_wire.py
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"] == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue