From 272a48d880c0347da312896285dd22059f008c64 Mon Sep 17 00:00:00 2001 From: Raghav Jhavar <156360524+raghav-stripe@users.noreply.github.com> Date: Tue, 13 Jan 2026 19:53:38 -0500 Subject: [PATCH] [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 --- .../exceptions/__init__.py | 19 + .../exceptions/exception_mapping_utils.py | 168 +++++++ .../exceptions/exceptions.py | 41 ++ .../count_tokens/bedrock_token_counter.py | 26 +- litellm/llms/bedrock/count_tokens/handler.py | 11 +- .../proxy/anthropic_endpoints/endpoints.py | 14 +- litellm/proxy/proxy_server.py | 26 +- litellm/types/utils.py | 6 + .../test_proxy_token_counter.py | 441 +++++++++++++++++- .../test_exception_mapping_utils.py | 185 ++++++++ 10 files changed, 920 insertions(+), 17 deletions(-) create mode 100644 litellm/anthropic_interface/exceptions/__init__.py create mode 100644 litellm/anthropic_interface/exceptions/exception_mapping_utils.py create mode 100644 litellm/anthropic_interface/exceptions/exceptions.py create mode 100644 tests/test_litellm/anthropic_interface/exceptions/test_exception_mapping_utils.py diff --git a/litellm/anthropic_interface/exceptions/__init__.py b/litellm/anthropic_interface/exceptions/__init__.py new file mode 100644 index 00000000000..875b09e3da3 --- /dev/null +++ b/litellm/anthropic_interface/exceptions/__init__.py @@ -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", +] diff --git a/litellm/anthropic_interface/exceptions/exception_mapping_utils.py b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py new file mode 100644 index 00000000000..b8a5079a4eb --- /dev/null +++ b/litellm/anthropic_interface/exceptions/exception_mapping_utils.py @@ -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, + ) diff --git a/litellm/anthropic_interface/exceptions/exceptions.py b/litellm/anthropic_interface/exceptions/exceptions.py new file mode 100644 index 00000000000..984390fa702 --- /dev/null +++ b/litellm/anthropic_interface/exceptions/exceptions.py @@ -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 diff --git a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py index b680bd046ef..54f8a8dbd65 100644 --- a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py +++ b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py @@ -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 diff --git a/litellm/llms/bedrock/count_tokens/handler.py b/litellm/llms/bedrock/count_tokens/handler.py index e8366165b65..9d2be6cca89 100644 --- a/litellm/llms/bedrock/count_tokens/handler.py +++ b/litellm/llms/bedrock/count_tokens/handler.py @@ -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( diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 7de6b7fccfc..0033deb0766 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -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( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1dfaf78cb5b..64d2c2a3a1b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 891826787ef..1e685f88086 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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): diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index a8486358f71..8e1057bb8e1 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -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") \ No newline at end of file + 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 + + diff --git a/tests/test_litellm/anthropic_interface/exceptions/test_exception_mapping_utils.py b/tests/test_litellm/anthropic_interface/exceptions/test_exception_mapping_utils.py new file mode 100644 index 00000000000..6f69a54b5d5 --- /dev/null +++ b/tests/test_litellm/anthropic_interface/exceptions/test_exception_mapping_utils.py @@ -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"]'