[bug fix] do not fallback to token counter if disable_token_counter is enabled (#19041)

* do not fallback to token counter if disable_token_counter is enabled, and return errors instead

* add exceptions and exception utils to map the same as /v1/chat/completions

* use safe_json_loads
This commit is contained in:
Raghav Jhavar 2026-01-13 19:53:38 -05:00 • committed by GitHub
parent 70987ffdc8
commit 272a48d880
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 920 additions and 17 deletions

View file

@ -0,0 +1,19 @@
"""Anthropic error format utilities."""
from .exception_mapping_utils import (
ANTHROPIC_ERROR_TYPE_MAP,
AnthropicExceptionMapping,
)
from .exceptions import (
AnthropicErrorDetail,
AnthropicErrorResponse,
AnthropicErrorType,
)
__all__ = [
"AnthropicErrorType",
"AnthropicErrorDetail",
"AnthropicErrorResponse",
"ANTHROPIC_ERROR_TYPE_MAP",
"AnthropicExceptionMapping",
]

View file

@ -0,0 +1,168 @@
"""
Utilities for mapping exceptions to Anthropic error format.
Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthropic response format.
"""
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from typing import Dict, Optional
from .exceptions import AnthropicErrorResponse, AnthropicErrorType
# HTTP status code -> Anthropic error type
# Source: https://docs.anthropic.com/en/api/errors
ANTHROPIC_ERROR_TYPE_MAP: Dict[int, AnthropicErrorType] = {
400: "invalid_request_error",
401: "authentication_error",
403: "permission_error",
404: "not_found_error",
413: "request_too_large",
429: "rate_limit_error",
500: "api_error",
529: "overloaded_error",
}
class AnthropicExceptionMapping:
"""
Helper class for mapping exceptions to Anthropic error format.
Similar pattern to ExceptionCheckers in litellm_core_utils/exception_mapping_utils.py
"""
@staticmethod
def get_error_type(status_code: int) -> AnthropicErrorType:
"""Map HTTP status code to Anthropic error type."""
return ANTHROPIC_ERROR_TYPE_MAP.get(status_code, "api_error")
@staticmethod
def create_error_response(
status_code: int,
message: str,
request_id: Optional[str] = None,
) -> AnthropicErrorResponse:
"""
Create an Anthropic-formatted error response dict.
Anthropic error format:
{
"type": "error",
"error": {"type": "...", "message": "..."},
"request_id": "req_..."
}
"""
error_type = AnthropicExceptionMapping.get_error_type(status_code)
response: AnthropicErrorResponse = {
"type": "error",
"error": {
"type": error_type,
"message": message,
},
}
if request_id:
response["request_id"] = request_id
return response
@staticmethod
def extract_error_message(raw_message: str) -> str:
"""
Extract error message from various provider response formats.
Handles:
- Bedrock: {"detail": {"message": "..."}}
- AWS: {"Message": "..."}
- Generic: {"message": "..."}
- Plain strings
"""
parsed = safe_json_loads(raw_message)
if isinstance(parsed, dict):
# Bedrock format
if "detail" in parsed and isinstance(parsed["detail"], dict):
return parsed["detail"].get("message", raw_message)
# AWS/generic format
return parsed.get("Message") or parsed.get("message") or raw_message
return raw_message
@staticmethod
def _is_anthropic_error_dict(parsed: dict) -> bool:
"""
Check if a parsed dict is in Anthropic error format.
Anthropic error format:
{
"type": "error",
"error": {"type": "...", "message": "..."}
}
"""
return (
parsed.get("type") == "error"
and isinstance(parsed.get("error"), dict)
and "type" in parsed["error"]
and "message" in parsed["error"]
)
@staticmethod
def _extract_message_from_dict(parsed: dict, raw_message: str) -> str:
"""
Extract error message from a parsed provider-specific dict.
Handles:
- Bedrock: {"detail": {"message": "..."}}
- AWS: {"Message": "..."}
- Generic: {"message": "..."}
"""
# Bedrock format
if "detail" in parsed and isinstance(parsed["detail"], dict):
return parsed["detail"].get("message", raw_message)
# AWS/generic format
return parsed.get("Message") or parsed.get("message") or raw_message
@staticmethod
def transform_to_anthropic_error(
status_code: int,
raw_message: str,
request_id: Optional[str] = None,
) -> AnthropicErrorResponse:
"""
Transform an error message to Anthropic format.
- If already in Anthropic format: passthrough unchanged
- Otherwise: extract message and create Anthropic error
Parses JSON only once for efficiency.
Args:
status_code: HTTP status code
raw_message: Raw error message (may be JSON string or plain text)
request_id: Optional request ID to include
Returns:
AnthropicErrorResponse dict
"""
# Try to parse as JSON once
parsed: Optional[dict] = safe_json_loads(raw_message)
if not isinstance(parsed, dict):
parsed = None
# If parsed and already in Anthropic format - passthrough
if parsed is not None and AnthropicExceptionMapping._is_anthropic_error_dict(parsed):
# Optionally add request_id if provided and not present
if request_id and "request_id" not in parsed:
parsed["request_id"] = request_id
return parsed # type: ignore
# Extract message - use parsed dict if available, otherwise raw string
if parsed is not None:
message = AnthropicExceptionMapping._extract_message_from_dict(parsed, raw_message)
else:
message = raw_message
return AnthropicExceptionMapping.create_error_response(
status_code=status_code,
message=message,
request_id=request_id,
)

View file

@ -0,0 +1,41 @@
"""Anthropic error format type definitions."""
from typing_extensions import Literal, Required, TypedDict
# Known Anthropic error types
# Source: https://docs.anthropic.com/en/api/errors
AnthropicErrorType = Literal[
"invalid_request_error",
"authentication_error",
"permission_error",
"not_found_error",
"request_too_large",
"rate_limit_error",
"api_error",
"overloaded_error",
]
class AnthropicErrorDetail(TypedDict):
"""Inner error detail in Anthropic format."""
type: AnthropicErrorType
message: str
class AnthropicErrorResponse(TypedDict, total=False):
"""
Anthropic-formatted error response.
Format:
{
"type": "error",
"error": {"type": "...", "message": "..."},
"request_id": "req_..." # optional
}
"""
type: Required[Literal["error"]]
error: Required[AnthropicErrorDetail]
request_id: str

View file

@ -6,7 +6,7 @@ from typing import Any, Dict, List, Optional
from litellm._logging import verbose_logger
from litellm.llms.base_llm.base_utils import BaseTokenCounter
from litellm.llms.bedrock.common_utils import get_bedrock_base_model
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
@ -79,9 +79,31 @@ class BedrockTokenCounter(BaseTokenCounter):
tokenizer_type="bedrock_api",
original_response=result,
)
except BedrockError as e:
verbose_logger.warning(
f"Bedrock CountTokens API error: status={e.status_code}, message={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(
f"Error calling Bedrock CountTokens API: {e}, falling back to default tokenizer"
f"Error calling Bedrock CountTokens API: {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

View file

@ -6,6 +6,8 @@ Simplified handler leveraging existing LiteLLM Bedrock infrastructure.
from typing import Any, Dict
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.common_utils import BedrockError
@ -98,7 +100,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
verbose_logger.error(f"AWS Bedrock error: {error_text}")
raise BedrockError(
status_code=response.status_code,
message=f"AWS Bedrock error: {error_text}",
message=error_text,
)
bedrock_response = response.json()
@ -117,6 +119,13 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
except BedrockError:
# Re-raise Bedrock exceptions as-is
raise
except httpx.HTTPStatusError as e:
# HTTP errors - preserve the actual status code
verbose_logger.error(f"HTTP error in CountTokens handler: {str(e)}")
raise BedrockError(
status_code=e.response.status_code,
message=e.response.text,
)
except Exception as e:
verbose_logger.error(f"Error in CountTokens handler: {str(e)}")
raise BedrockError(

View file

@ -2,13 +2,13 @@
Unified /v1/messages endpoint - (Anthropic Spec)
"""
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from litellm._logging import verbose_proxy_logger
from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
create_response,
@ -221,6 +221,16 @@ async def count_tokens(
except HTTPException:
raise
except ProxyException as e:
status_code = int(e.code) if e.code and e.code.isdigit() else 500
detail = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=status_code,
raw_message=e.message,
)
raise HTTPException(
status_code=status_code,
detail=detail,
)
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.anthropic_endpoints.count_tokens(): Exception occurred - {}".format(

View file

@ -7067,9 +7067,33 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False)
#########################################################
# Transfrom the Response to the well known format
#########################################################
if result is not None:
if result is not None and result.error is True:
# If disable_token_counter is enabled, raise HTTP error
if litellm.disable_token_counter is True:
raise ProxyException(
message=result.error_message or "Token counting failed",
type="token_counting_error",
param="model",
code=result.status_code or 500,
)
# Otherwise, log warning and fall back to local counter
verbose_proxy_logger.warning(
f"Provider token counting failed ({result.status_code}): {result.error_message}. "
"Falling back to local tokenizer."
)
else:
# Success - return the result
return result
# Check if token counter is disabled before fallback
if litellm.disable_token_counter is True:
raise ProxyException(
message="Token counting is disabled and no provider API result available",
type="token_counting_disabled",
param="model",
code=503,
)
# Default LiteLLM token counting
custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None
if model_info is not None:

View file

@ -3089,6 +3089,12 @@ class TokenCountResponse(LiteLLMPydanticObjectBase):
"""
Original Response from upstream API call - if an API call was made for token counting
"""
error: bool = False
error_message: Optional[str] = None
"""
HTTP status code from the token counting API (e.g., 200 for success, 429 for rate limit, 400 for bad request)
"""
status_code: Optional[int] = None
class CustomHuggingfaceTokenizer(TypedDict):

View file

@ -2,30 +2,42 @@
# 1. Generate a Key, and use it to make a call
import sys, os
import json
import logging
import os
import sys
import tempfile
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from dotenv import load_dotenv
load_dotenv()
import os
# this file is to test litellm/proxy
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import pytest, logging
from fastapi import HTTPException, Request
import litellm
from litellm.proxy.proxy_server import token_counter
from litellm import Router
from litellm._logging import verbose_proxy_logger
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.proxy._types import ProxyException, TokenCountRequest
from litellm.proxy.anthropic_endpoints.endpoints import (
count_tokens as anthropic_count_tokens,
)
from litellm.proxy.proxy_server import token_counter
from litellm.types.utils import TokenCountResponse
verbose_proxy_logger.setLevel(level=logging.DEBUG)
from litellm.proxy._types import TokenCountRequest
import json, tempfile
from litellm import Router
def get_vertex_ai_creds_json() -> dict:
# Define the path to the vertex_key.json file
@ -839,4 +851,411 @@ def test_vertex_ai_partner_models_token_counting_endpoint(vertex_location):
if vertex_location == "global":
assert endpoint.startswith("https://aiplatform.googleapis.com")
else:
assert endpoint.startswith(f"https://{vertex_location}-aiplatform.googleapis.com")
assert endpoint.startswith(f"https://{vertex_location}-aiplatform.googleapis.com")
@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()
# 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",
)
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
async def test_bedrock_token_counter_error_propagation_generic_exception():
"""
Test that BedrockTokenCounter returns error response with 500 status for generic exceptions.
"""
counter = BedrockTokenCounter()
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",
)
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
async def test_bedrock_handler_httpx_error_status_code_propagation():
"""
Test that BedrockCountTokensHandler properly extracts status code from httpx.HTTPStatusError.
"""
handler = BedrockCountTokensHandler()
# Create a mock httpx response with 403 status
mock_response = MagicMock()
mock_response.status_code = 403
mock_response.text = "Forbidden - Invalid credentials"
# Create HTTPStatusError
http_error = httpx.HTTPStatusError(
message="Client error '403 Forbidden'",
request=MagicMock(),
response=mock_response,
)
with patch.object(handler, "validate_count_tokens_request"):
with patch.object(handler, "_get_aws_region_name", return_value="us-west-2"):
with patch.object(
handler, "transform_anthropic_to_bedrock_count_tokens", return_value={}
):
with patch.object(
handler,
"get_bedrock_count_tokens_endpoint",
return_value="https://example.com",
):
with patch.object(handler, "_sign_request", return_value=({}, "{}")):
with patch(
"litellm.llms.bedrock.count_tokens.handler.get_async_httpx_client"
) as mock_client:
mock_async_client = AsyncMock()
mock_async_client.post = AsyncMock(side_effect=http_error)
mock_client.return_value = mock_async_client
with pytest.raises(BedrockError) as exc_info:
await handler.handle_count_tokens_request(
request_data={
"model": "test",
"messages": [
{"role": "user", "content": "hello"}
],
},
litellm_params={},
resolved_model="anthropic.claude-3-sonnet",
)
assert exc_info.value.status_code == 403
# Message should be the raw response text
assert exc_info.value.message == "Forbidden - Invalid credentials"
@pytest.mark.asyncio
async def test_proxy_token_counter_error_raises_exception_when_disabled():
"""
Test that proxy token_counter raises ProxyException when disable_token_counter=True
and provider returns an error response.
"""
# Create error response
error_response = TokenCountResponse(
total_tokens=0,
request_model="bedrock/anthropic.claude-3-sonnet",
model_used="anthropic.claude-3-sonnet",
tokenizer_type="bedrock_api",
error=True,
error_message="Rate limit exceeded",
status_code=429,
)
# Create mock router that returns a deployment
mock_deployment = {
"litellm_params": {
"model": "bedrock/anthropic.claude-3-sonnet",
},
"model_info": {},
}
mock_router = MagicMock()
mock_router.async_get_available_deployment = AsyncMock(return_value=mock_deployment)
setattr(litellm.proxy.proxy_server, "llm_router", mock_router)
# Save original value and function
original_disable = litellm.disable_token_counter
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
try:
litellm.disable_token_counter = True
# Create a mock counter that returns an error response
mock_counter = MagicMock(spec=BedrockTokenCounter)
mock_counter.should_use_token_counting_api.return_value = True
mock_counter.count_tokens = AsyncMock(return_value=error_response)
# Replace the function directly
def mock_get_provider_token_counter(deployment, model_to_use):
return (mock_counter, "anthropic.claude-3-sonnet", "bedrock")
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
with pytest.raises(ProxyException) as exc_info:
await token_counter(
request=TokenCountRequest(
model="claude-bedrock",
messages=[{"role": "user", "content": "hello"}],
),
call_endpoint=True,
)
assert exc_info.value.code == "429"
assert "Rate limit exceeded" in exc_info.value.message
finally:
litellm.disable_token_counter = original_disable
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
@pytest.mark.asyncio
async def test_proxy_token_counter_error_falls_back_when_enabled():
"""
Test that proxy token_counter falls back to local tokenizer when disable_token_counter=False
and provider returns an error response.
"""
# Create error response
error_response = TokenCountResponse(
total_tokens=0,
request_model="bedrock/anthropic.claude-3-sonnet",
model_used="anthropic.claude-3-sonnet",
tokenizer_type="bedrock_api",
error=True,
error_message="Rate limit exceeded",
status_code=429,
)
# Create mock router that returns a deployment
mock_deployment = {
"litellm_params": {
"model": "bedrock/anthropic.claude-3-sonnet",
},
"model_info": {},
}
mock_router = MagicMock()
mock_router.async_get_available_deployment = AsyncMock(return_value=mock_deployment)
setattr(litellm.proxy.proxy_server, "llm_router", mock_router)
# Save original value and function
original_disable = litellm.disable_token_counter
original_get_provider_token_counter = litellm.proxy.proxy_server._get_provider_token_counter
try:
litellm.disable_token_counter = False
# Create a mock counter that returns an error response
mock_counter = MagicMock(spec=BedrockTokenCounter)
mock_counter.should_use_token_counting_api.return_value = True
mock_counter.count_tokens = AsyncMock(return_value=error_response)
# Replace the function directly
def mock_get_provider_token_counter(deployment, model_to_use):
return (mock_counter, "anthropic.claude-3-sonnet", "bedrock")
litellm.proxy.proxy_server._get_provider_token_counter = mock_get_provider_token_counter
# Should not raise, should fall back to local tokenizer
result = await token_counter(
request=TokenCountRequest(
model="claude-bedrock",
messages=[{"role": "user", "content": "hello"}],
),
call_endpoint=True,
)
# Should have used the fallback tokenizer
assert result.error is False
assert result.total_tokens > 0
assert result.tokenizer_type != "bedrock_api"
finally:
litellm.disable_token_counter = original_disable
litellm.proxy.proxy_server._get_provider_token_counter = original_get_provider_token_counter
@pytest.mark.asyncio
async def test_anthropic_endpoint_returns_anthropic_error_format():
"""
Test that /v1/messages/count_tokens returns errors in Anthropic format.
"""
import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints
import litellm.proxy.proxy_server as proxy_server
# Mock request object
mock_request = MagicMock(spec=Request)
mock_request_data = {
"model": "claude-bedrock",
"messages": [{"role": "user", "content": "Hello!"}],
}
async def mock_read_request_body(request):
return mock_request_data
mock_user_api_key_dict = MagicMock()
original_read_request_body = anthropic_endpoints._read_request_body
anthropic_endpoints._read_request_body = mock_read_request_body
original_token_counter = proxy_server.token_counter
# Mock token_counter to raise ProxyException with Bedrock-style error
async def mock_token_counter_error(request, call_endpoint=False):
raise ProxyException(
message='{"detail":{"message":"Input is too long for requested model."}}',
type="token_counting_error",
param="model",
code=400,
)
proxy_server.token_counter = mock_token_counter_error
try:
with pytest.raises(HTTPException) as exc_info:
await anthropic_count_tokens(mock_request, mock_user_api_key_dict)
# Verify HTTP status code is correct
assert exc_info.value.status_code == 400
# Verify error is in Anthropic format
detail = exc_info.value.detail
assert detail["type"] == "error"
assert detail["error"]["type"] == "invalid_request_error"
assert detail["error"]["message"] == "Input is too long for requested model."
finally:
anthropic_endpoints._read_request_body = original_read_request_body
proxy_server.token_counter = original_token_counter
@pytest.mark.asyncio
async def test_anthropic_endpoint_403_permission_error_format():
"""
Test that 403 errors are returned as permission_error in Anthropic format.
"""
import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints
import litellm.proxy.proxy_server as proxy_server
mock_request = MagicMock(spec=Request)
mock_request_data = {
"model": "claude-bedrock",
"messages": [{"role": "user", "content": "Hello!"}],
}
async def mock_read_request_body(request):
return mock_request_data
mock_user_api_key_dict = MagicMock()
original_read_request_body = anthropic_endpoints._read_request_body
anthropic_endpoints._read_request_body = mock_read_request_body
original_token_counter = proxy_server.token_counter
# Mock token_counter to raise ProxyException with 403 error
async def mock_token_counter_error(request, call_endpoint=False):
raise ProxyException(
message='{"Message":"Bearer Token has expired"}',
type="token_counting_error",
param="model",
code=403,
)
proxy_server.token_counter = mock_token_counter_error
try:
with pytest.raises(HTTPException) as exc_info:
await anthropic_count_tokens(mock_request, mock_user_api_key_dict)
assert exc_info.value.status_code == 403
detail = exc_info.value.detail
assert detail["type"] == "error"
assert detail["error"]["type"] == "permission_error"
assert detail["error"]["message"] == "Bearer Token has expired"
finally:
anthropic_endpoints._read_request_body = original_read_request_body
proxy_server.token_counter = original_token_counter
@pytest.mark.asyncio
async def test_anthropic_endpoint_429_rate_limit_error_format():
"""
Test that 429 errors are returned as rate_limit_error in Anthropic format.
"""
import litellm.proxy.anthropic_endpoints.endpoints as anthropic_endpoints
import litellm.proxy.proxy_server as proxy_server
mock_request = MagicMock(spec=Request)
mock_request_data = {
"model": "claude-bedrock",
"messages": [{"role": "user", "content": "Hello!"}],
}
async def mock_read_request_body(request):
return mock_request_data
mock_user_api_key_dict = MagicMock()
original_read_request_body = anthropic_endpoints._read_request_body
anthropic_endpoints._read_request_body = mock_read_request_body
original_token_counter = proxy_server.token_counter
# Mock token_counter to raise ProxyException with 429 error
async def mock_token_counter_error(request, call_endpoint=False):
raise ProxyException(
message="Rate limit exceeded",
type="token_counting_error",
param="model",
code=429,
)
proxy_server.token_counter = mock_token_counter_error
try:
with pytest.raises(HTTPException) as exc_info:
await anthropic_count_tokens(mock_request, mock_user_api_key_dict)
assert exc_info.value.status_code == 429
detail = exc_info.value.detail
assert detail["type"] == "error"
assert detail["error"]["type"] == "rate_limit_error"
assert detail["error"]["message"] == "Rate limit exceeded"
finally:
anthropic_endpoints._read_request_body = original_read_request_body
proxy_server.token_counter = original_token_counter

View file

@ -0,0 +1,185 @@
"""
Tests for AnthropicExceptionMapping class in litellm/anthropic_interface/exceptions/exception_mapping_utils.py
"""
import json
from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping
class TestCreateErrorResponse:
"""Tests for AnthropicExceptionMapping.create_error_response()"""
def test_400_invalid_request_error(self):
"""Test 400 maps to invalid_request_error."""
response = AnthropicExceptionMapping.create_error_response(400, "Invalid request")
assert response["type"] == "error"
assert response["error"]["type"] == "invalid_request_error"
assert response["error"]["message"] == "Invalid request"
assert "request_id" not in response
def test_401_authentication_error(self):
"""Test 401 maps to authentication_error."""
response = AnthropicExceptionMapping.create_error_response(401, "Unauthorized")
assert response["error"]["type"] == "authentication_error"
def test_403_permission_error(self):
"""Test 403 maps to permission_error."""
response = AnthropicExceptionMapping.create_error_response(403, "Forbidden")
assert response["error"]["type"] == "permission_error"
def test_404_not_found_error(self):
"""Test 404 maps to not_found_error."""
response = AnthropicExceptionMapping.create_error_response(404, "Not found")
assert response["error"]["type"] == "not_found_error"
def test_429_rate_limit_error(self):
"""Test 429 maps to rate_limit_error."""
response = AnthropicExceptionMapping.create_error_response(429, "Rate limit exceeded")
assert response["error"]["type"] == "rate_limit_error"
def test_500_api_error(self):
"""Test 500 maps to api_error."""
response = AnthropicExceptionMapping.create_error_response(500, "Internal error")
assert response["error"]["type"] == "api_error"
def test_with_request_id(self):
"""Test request_id is included when provided."""
response = AnthropicExceptionMapping.create_error_response(400, "Error", request_id="req_123")
assert response["request_id"] == "req_123"
def test_unknown_status_defaults_to_api_error(self):
"""Test unknown status code defaults to api_error."""
response = AnthropicExceptionMapping.create_error_response(418, "I'm a teapot")
assert response["error"]["type"] == "api_error"
class TestExtractErrorMessage:
"""Tests for AnthropicExceptionMapping.extract_error_message()"""
def test_bedrock_format(self):
"""Test extraction from Bedrock format: {"detail": {"message": "..."}}"""
bedrock_msg = '{"detail":{"message":"Input is too long for requested model."}}'
assert AnthropicExceptionMapping.extract_error_message(bedrock_msg) == "Input is too long for requested model."
def test_aws_message_format(self):
"""Test extraction from AWS format: {"Message": "..."}"""
msg = '{"Message":"Bearer Token has expired"}'
assert AnthropicExceptionMapping.extract_error_message(msg) == "Bearer Token has expired"
def test_generic_message_format(self):
"""Test extraction from generic format: {"message": "..."}"""
msg = '{"message":"Some error occurred"}'
assert AnthropicExceptionMapping.extract_error_message(msg) == "Some error occurred"
def test_plain_string(self):
"""Test plain string is returned as-is."""
assert AnthropicExceptionMapping.extract_error_message("Plain error message") == "Plain error message"
def test_invalid_json(self):
"""Test invalid JSON is returned as-is."""
assert AnthropicExceptionMapping.extract_error_message("Not JSON {invalid}") == "Not JSON {invalid}"
def test_empty_dict(self):
"""Test empty dict returns original string."""
assert AnthropicExceptionMapping.extract_error_message("{}") == "{}"
class TestTransformToAnthropicError:
"""Tests for AnthropicExceptionMapping.transform_to_anthropic_error()"""
def test_passthrough_anthropic_error(self):
"""Test that Anthropic errors pass through unchanged."""
anthropic_error = {
"type": "error",
"error": {"type": "rate_limit_error", "message": "Rate limited"}
}
raw = json.dumps(anthropic_error)
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=429,
raw_message=raw,
)
assert result["type"] == "error"
assert result["error"]["type"] == "rate_limit_error"
assert result["error"]["message"] == "Rate limited"
def test_passthrough_preserves_existing_request_id(self):
"""Test that existing request_id in Anthropic error is preserved."""
anthropic_error = {
"type": "error",
"error": {"type": "api_error", "message": "Server error"},
"request_id": "req_existing"
}
raw = json.dumps(anthropic_error)
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=500,
raw_message=raw,
request_id="req_new", # Should not override existing
)
assert result["request_id"] == "req_existing"
def test_passthrough_adds_request_id_if_missing(self):
"""Test that request_id is added to Anthropic error if missing."""
anthropic_error = {
"type": "error",
"error": {"type": "api_error", "message": "Server error"}
}
raw = json.dumps(anthropic_error)
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=500,
raw_message=raw,
request_id="req_123",
)
assert result["request_id"] == "req_123"
def test_translates_bedrock_error(self):
"""Test that Bedrock errors are translated to Anthropic format."""
bedrock_error = json.dumps({"detail": {"message": "Access denied"}})
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=403,
raw_message=bedrock_error,
)
assert result["type"] == "error"
assert result["error"]["type"] == "permission_error"
assert result["error"]["message"] == "Access denied"
def test_translates_aws_error(self):
"""Test that AWS errors are translated to Anthropic format."""
aws_error = json.dumps({"Message": "Resource not found"})
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=404,
raw_message=aws_error,
)
assert result["type"] == "error"
assert result["error"]["type"] == "not_found_error"
assert result["error"]["message"] == "Resource not found"
def test_handles_plain_string(self):
"""Test that plain string errors are wrapped in Anthropic format."""
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=400,
raw_message="Invalid request parameters",
)
assert result["type"] == "error"
assert result["error"]["type"] == "invalid_request_error"
assert result["error"]["message"] == "Invalid request parameters"
def test_handles_generic_message_json(self):
"""Test that generic {"message": "..."} JSON is translated."""
generic_error = json.dumps({"message": "Something went wrong"})
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=500,
raw_message=generic_error,
)
assert result["type"] == "error"
assert result["error"]["type"] == "api_error"
assert result["error"]["message"] == "Something went wrong"
def test_handles_non_dict_json(self):
"""Test that non-dict JSON (e.g., array) is treated as plain string."""
result = AnthropicExceptionMapping.transform_to_anthropic_error(
status_code=400,
raw_message='["error1", "error2"]',
)
assert result["type"] == "error"
assert result["error"]["message"] == '["error1", "error2"]'