mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #18871 from BerriAI/litellm_fix_test_count_tokens_caching
Fix :test_count_tokens_caching
This commit is contained in:
commit
0c7db97ad5
3 changed files with 32 additions and 23 deletions
|
|
@ -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)}",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue