mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Handle not supported region for vertex ai count tokens - v1/messages/count_tokens (#18665)
* Handle not supported region for vertex ai count tokens * add unit test
This commit is contained in:
parent
762345172c
commit
18ca6b2add
3 changed files with 55 additions and 0 deletions
|
|
@ -92,6 +92,7 @@ model_list:
|
|||
model: vertex_ai/claude-3-5-sonnet-v2@20241022
|
||||
vertex_project: my-project
|
||||
vertex_location: us-east5
|
||||
vertex_count_tokens_location: us-east5 # Optional: Override location for token counting (count_tokens not available on global location)
|
||||
|
||||
- model_name: claude-bedrock
|
||||
litellm_params:
|
||||
|
|
|
|||
|
|
@ -941,9 +941,16 @@ class VertexAITokenCounter(BaseTokenCounter):
|
|||
vertex_project = count_tokens_params_request.get(
|
||||
"vertex_project"
|
||||
) or count_tokens_params_request.get("vertex_ai_project")
|
||||
|
||||
vertex_location = count_tokens_params_request.get(
|
||||
"vertex_location"
|
||||
) or count_tokens_params_request.get("vertex_ai_location")
|
||||
|
||||
# Count tokens not available on global location: https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens
|
||||
vertex_location = count_tokens_params_request.get(
|
||||
"vertex_count_tokens_location"
|
||||
) or vertex_location
|
||||
|
||||
vertex_credentials = count_tokens_params_request.get(
|
||||
"vertex_credentials"
|
||||
) or count_tokens_params_request.get("vertex_ai_credentials")
|
||||
|
|
|
|||
|
|
@ -1021,6 +1021,53 @@ async def test_vertex_ai_token_counter_routes_partner_models():
|
|||
assert result.tokenizer_type == "vertex_ai_partner_models"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_ai_token_counter_uses_count_tokens_location():
|
||||
"""
|
||||
Test that VertexAITokenCounter uses vertex_count_tokens_location to override
|
||||
vertex_location when counting tokens for partner models.
|
||||
|
||||
Count tokens API is not available on global location for partner models:
|
||||
https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens
|
||||
"""
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter
|
||||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
token_counter = VertexAITokenCounter()
|
||||
|
||||
# Mock the partner models handler
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.vertex_ai_partner_models.main.VertexAIPartnerModels.count_tokens"
|
||||
) as mock_partner_count_tokens:
|
||||
mock_partner_count_tokens.return_value = {
|
||||
"input_tokens": 42,
|
||||
"tokenizer_used": "vertex_ai_partner_models",
|
||||
}
|
||||
|
||||
# Test with vertex_count_tokens_location overriding vertex_location
|
||||
await token_counter.count_tokens(
|
||||
model_to_use="claude-3-5-sonnet-20241022",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
contents=None,
|
||||
deployment={
|
||||
"litellm_params": {
|
||||
"vertex_project": "test-project",
|
||||
"vertex_location": "global", # Original location (not supported for count_tokens)
|
||||
"vertex_count_tokens_location": "us-east5", # Override for count_tokens
|
||||
}
|
||||
},
|
||||
request_model="vertex_ai/claude-3-5-sonnet-20241022",
|
||||
)
|
||||
|
||||
# Verify the partner models handler was called with the overridden location
|
||||
assert mock_partner_count_tokens.called
|
||||
call_kwargs = mock_partner_count_tokens.call_args.kwargs
|
||||
assert call_kwargs["vertex_location"] == "us-east5"
|
||||
assert call_kwargs["vertex_project"] == "test-project"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_ai_token_counter_routes_gemini_models():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue