Merge pull request #18871 from BerriAI/litellm_fix_test_count_tokens_caching

Fix :test_count_tokens_caching
This commit is contained in:
Sameer Kankute 2026-01-10 01:20:55 +05:30 • committed by GitHub
commit 0c7db97ad5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 32 additions and 23 deletions

View file

@ -6,10 +6,9 @@ Simplified handler leveraging existing LiteLLM Bedrock infrastructure.
from typing import Any, Dict
from fastapi import HTTPException
import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
@ -97,9 +96,9 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
if response.status_code != 200:
error_text = response.text
verbose_logger.error(f"AWS Bedrock error: {error_text}")
raise HTTPException(
status_code=400,
detail={"error": f"AWS Bedrock error: {error_text}"},
raise BedrockError(
status_code=response.status_code,
message=f"AWS Bedrock error: {error_text}",
)
bedrock_response = response.json()
@ -115,12 +114,12 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
return final_response
except HTTPException:
# Re-raise HTTP exceptions as-is
except BedrockError:
# Re-raise Bedrock exceptions as-is
raise
except Exception as e:
verbose_logger.error(f"Error in CountTokens handler: {str(e)}")
raise HTTPException(
raise BedrockError(
status_code=500,
detail={"error": f"CountTokens processing error: {str(e)}"},
message=f"CountTokens processing error: {str(e)}",
)

View file

@ -776,6 +776,7 @@ async def handle_bedrock_count_tokens(
- /v1/messages/count_tokens
- /v1/messages/count-tokens
"""
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler
from litellm.proxy.proxy_server import llm_router
@ -822,6 +823,12 @@ async def handle_bedrock_count_tokens(
return result
except BedrockError as e:
# Convert BedrockError to HTTPException for FastAPI
verbose_proxy_logger.error(f"BedrockError in handle_bedrock_count_tokens: {str(e)}")
raise HTTPException(
status_code=e.status_code, detail={"error": e.message}
)
except HTTPException:
# Re-raise HTTP exceptions as-is
raise

View file

@ -1,4 +1,5 @@
import pytest
import sys
from unittest.mock import MagicMock, patch, AsyncMock
from litellm.proxy.utils import count_tokens_with_anthropic_api, _anthropic_async_clients
@ -14,29 +15,31 @@ async def test_count_tokens_caching():
messages = [{"role": "user", "content": "hello"}]
model = "claude-3-opus-20240229"
# Mock anthropic
with patch("anthropic.AsyncAnthropic") as mock_cls:
mock_client = MagicMock()
mock_cls.return_value = mock_client
# Mock response
mock_response = MagicMock()
mock_response.input_tokens = 10
# Setup async return for count_tokens
mock_client.beta.messages.count_tokens = AsyncMock(return_value=mock_response)
# Create a mock anthropic module
mock_anthropic = MagicMock()
mock_client = MagicMock()
mock_anthropic.AsyncAnthropic.return_value = mock_client
# Mock response
mock_response = MagicMock()
mock_response.input_tokens = 10
# Setup async return for count_tokens
mock_client.beta.messages.count_tokens = AsyncMock(return_value=mock_response)
# Patch sys.modules to ensure our mock is used when anthropic is imported
with patch.dict(sys.modules, {"anthropic": mock_anthropic}):
# First call
with patch.dict("os.environ", {"ANTHROPIC_API_KEY": api_key}):
await count_tokens_with_anthropic_api(model, messages)
assert api_key in _anthropic_async_clients
assert _anthropic_async_clients[api_key] == mock_client
mock_cls.assert_called_once() # Should be called once
mock_anthropic.AsyncAnthropic.assert_called_once() # Should be called once
# Second call
with patch.dict("os.environ", {"ANTHROPIC_API_KEY": api_key}):
await count_tokens_with_anthropic_api(model, messages)
# Should still be called once (cached)
mock_cls.assert_called_once()
mock_anthropic.AsyncAnthropic.assert_called_once()