mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
[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:
parent
70987ffdc8
commit
272a48d880
10 changed files with 920 additions and 17 deletions
19
litellm/anthropic_interface/exceptions/__init__.py
Normal file
19
litellm/anthropic_interface/exceptions/__init__.py
Normal 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",
|
||||
]
|
||||
|
|
@ -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,
|
||||
)
|
||||
41
litellm/anthropic_interface/exceptions/exceptions.py
Normal file
41
litellm/anthropic_interface/exceptions/exceptions.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]'
|
||||
Loading…
Add table
Reference in a new issue