diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index a750b5e985..2c53457736 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -27,6 +27,7 @@ class VertexAIError(BaseLLMException): class VertexAIModelRoute(str, Enum): """Enum for Vertex AI model routing""" + PARTNER_MODELS = "partner_models" GEMINI = "gemini" GEMMA = "gemma" @@ -34,27 +35,29 @@ class VertexAIModelRoute(str, Enum): NON_GEMINI = "non_gemini" -def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None) -> VertexAIModelRoute: +def get_vertex_ai_model_route( + model: str, litellm_params: Optional[dict] = None +) -> VertexAIModelRoute: """ Determine which handler to use for a Vertex AI model based on the model name. - + Args: model: The model name (e.g., "llama3-405b", "gemini-pro", "gemma/gemma-3-12b-it", "openai/gpt-oss-120b") litellm_params: Optional litellm parameters dict that may contain base_model for routing - + Returns: VertexAIModelRoute: The route enum indicating which handler should be used - + Examples: >>> get_vertex_ai_model_route("llama3-405b") VertexAIModelRoute.PARTNER_MODELS - + >>> get_vertex_ai_model_route("gemini-pro") VertexAIModelRoute.GEMINI - + >>> get_vertex_ai_model_route("gemma/gemma-3-12b-it") VertexAIModelRoute.GEMMA - + >>> get_vertex_ai_model_route("openai/gpt-oss-120b") VertexAIModelRoute.MODEL_GARDEN """ @@ -66,23 +69,23 @@ def get_vertex_ai_model_route(model: str, litellm_params: Optional[dict] = None) if litellm_params and litellm_params.get("base_model") is not None: if "gemini" in litellm_params["base_model"]: return VertexAIModelRoute.GEMINI - + # Check for partner models (llama, mistral, claude, etc.) if VertexAIPartnerModels.is_vertex_partner_model(model=model): return VertexAIModelRoute.PARTNER_MODELS - + # Check for gemma models if "gemma/" in model: return VertexAIModelRoute.GEMMA - + # Check for model garden openai models if "openai" in model: return VertexAIModelRoute.MODEL_GARDEN - + # Check for gemini models if "gemini" in model: return VertexAIModelRoute.GEMINI - + # Default to non-gemini (legacy vertex models like chat-bison, text-bison, etc.) return VertexAIModelRoute.NON_GEMINI @@ -253,8 +256,10 @@ def _check_text_in_content(parts: List[PartType]) -> bool: def _fix_enum_empty_strings(schema, depth=0): """Fix empty strings in enum values by replacing them with None. Gemini doesn't accept empty strings in enums.""" if depth > DEFAULT_MAX_RECURSE_DEPTH: - raise ValueError(f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while processing schema.") - + raise ValueError( + f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while processing schema." + ) + if "enum" in schema and isinstance(schema["enum"], list): schema["enum"] = [None if value == "" else value for value in schema["enum"]] @@ -529,19 +534,18 @@ def _convert_vertex_datetime_to_openai_datetime(vertex_datetime: str) -> int: def _convert_schema_types(schema, depth=0): """ Convert type arrays and lowercase types for Vertex AI compatibility. - - Transforms OpenAI-style schemas to Vertex AI format by converting type arrays + + Transforms OpenAI-style schemas to Vertex AI format by converting type arrays like ["string", "number"] to anyOf format and converting all types to uppercase. """ if depth > DEFAULT_MAX_RECURSE_DEPTH: raise ValueError( f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while processing schema. Please check the schema for excessive nesting." ) - + if not isinstance(schema, dict): return - # Handle type field if "type" in schema: type_val = schema["type"] @@ -553,7 +557,7 @@ def _convert_schema_types(schema, depth=0): schema["type"] = type_val[0] elif isinstance(type_val, str): schema["type"] = type_val - + # Recursively process nested properties, items, and anyOf for key in ["properties", "items", "anyOf"]: if key in schema: @@ -567,6 +571,7 @@ def _convert_schema_types(schema, depth=0): for anyof_schema in value: _convert_schema_types(anyof_schema, depth + 1) + def get_vertex_project_id_from_url(url: str) -> Optional[str]: """ Get the vertex project id from the url @@ -665,17 +670,18 @@ def is_global_only_vertex_model(model: str) -> bool: return False return "global" in supported_regions -class VertexAIModelInfo(BaseLLMModelInfo): + +class VertexAIModelInfo(BaseLLMModelInfo): def get_token_counter(self) -> Optional[BaseTokenCounter]: """ Factory method to create a token counter for this provider. - + Returns: Optional TokenCounterInterface implementation for this provider, or None if token counting is not supported. """ return VertexAITokenCounter() - + def validate_environment( self, headers: dict, @@ -687,7 +693,7 @@ class VertexAIModelInfo(BaseLLMModelInfo): api_base: Optional[str] = None, ) -> dict: raise NotImplementedError("Vertex AI models are not supported yet") - + def get_models( self, api_key: Optional[str] = None, api_base: Optional[str] = None ) -> List[str]: @@ -706,8 +712,6 @@ class VertexAIModelInfo(BaseLLMModelInfo): ) -> Optional[str]: raise NotImplementedError("Vertex AI models are not supported yet") - - @staticmethod def get_base_model(model: str) -> Optional[str]: """ @@ -721,13 +725,15 @@ class VertexAIModelInfo(BaseLLMModelInfo): class VertexAITokenCounter(BaseTokenCounter): """Token counter implementation for Google AI Studio provider.""" + def should_use_token_counting_api( - self, + self, custom_llm_provider: Optional[str] = None, ) -> bool: from litellm.types.utils import LlmProviders + return custom_llm_provider == LlmProviders.VERTEX_AI.value - + async def count_tokens( self, model_to_use: str, @@ -738,25 +744,68 @@ class VertexAITokenCounter(BaseTokenCounter): ) -> Optional[TokenCountResponse]: import copy - from litellm.llms.vertex_ai.count_tokens.handler import VertexAITokenCounter - deployment = deployment or {} - count_tokens_params_request = copy.deepcopy(deployment.get("litellm_params", {})) - count_tokens_params = { - "model": model_to_use, - "contents": contents, - } - count_tokens_params_request.update(count_tokens_params) - result = await VertexAITokenCounter().acount_tokens( - **count_tokens_params_request, + from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( + VertexAIPartnerModels, ) - - if result is not None: - return TokenCountResponse( - total_tokens=result.get("totalTokens", 0), - request_model=request_model, - model_used=model_to_use, - tokenizer_type=result.get("tokenizer_used", ""), - original_response=result, + + deployment = deployment or {} + count_tokens_params_request = copy.deepcopy( + deployment.get("litellm_params", {}) + ) + + # Check if this is a partner model (Claude, Mistral, etc.) + if VertexAIPartnerModels.is_vertex_partner_model(model_to_use): + # Use partner models token counter + partner_models_handler = VertexAIPartnerModels() + + # Extract vertex-specific params from litellm_params + 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") + vertex_credentials = count_tokens_params_request.get( + "vertex_credentials" + ) or count_tokens_params_request.get("vertex_ai_credentials") + + result = await partner_models_handler.count_tokens( + model=model_to_use, + messages=messages or [], + litellm_params=count_tokens_params_request, + vertex_project=vertex_project, + vertex_location=vertex_location, + vertex_credentials=vertex_credentials, ) - - return None \ No newline at end of file + + if result is not None: + return TokenCountResponse( + total_tokens=result.get("input_tokens", 0), + request_model=request_model, + model_used=model_to_use, + tokenizer_type=result.get("tokenizer_used", ""), + original_response=result, + ) + else: + # Use standard Vertex AI (Gemini) token counter + from litellm.llms.vertex_ai.count_tokens.handler import VertexAITokenCounter + + count_tokens_params = { + "model": model_to_use, + "contents": contents, + } + count_tokens_params_request.update(count_tokens_params) + result = await VertexAITokenCounter().acount_tokens( + **count_tokens_params_request, + ) + + if result is not None: + return TokenCountResponse( + total_tokens=result.get("totalTokens", 0), + request_model=request_model, + model_used=model_to_use, + tokenizer_type=result.get("tokenizer_used", ""), + original_response=result, + ) + + return None diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py new file mode 100644 index 0000000000..008957c87a --- /dev/null +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py @@ -0,0 +1 @@ +# Count tokens handler for Vertex AI Partner Models (Anthropic, Mistral, etc.) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py new file mode 100644 index 0000000000..0f073d0226 --- /dev/null +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py @@ -0,0 +1,157 @@ +""" +Token counter for Vertex AI Partner Models (Anthropic Claude, Mistral, etc.) + +This handler provides token counting for partner models hosted on Vertex AI. +Unlike Gemini models which use Google's token counting API, partner models use +their respective publisher-specific count-tokens endpoints. +""" +from typing import Any, Dict, Optional + +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.llms.vertex_ai.vertex_llm_base import VertexBase +from litellm.types.llms.vertex_ai import VertexPartnerProvider + + +class VertexAIPartnerModelsTokenCounter(VertexBase): + """ + Token counter for Vertex AI Partner Models. + + Handles token counting for models like Claude (Anthropic), Mistral, etc. + that are available through Vertex AI's partner model program. + """ + + def _get_publisher_for_model(self, model: str) -> str: + """ + Determine the publisher name for the given model. + + Args: + model: The model name (e.g., "claude-3-5-sonnet-20241022") + + Returns: + Publisher name to use in the Vertex AI endpoint URL + + Raises: + ValueError: If the model is not a recognized partner model + """ + if "claude" in model: + return "anthropic" + elif "mistral" in model or "codestral" in model: + return "mistralai" + elif "llama" in model or "meta/" in model: + return "meta" + else: + raise ValueError(f"Unknown partner model: {model}") + + def _build_count_tokens_endpoint( + self, + model: str, + project_id: str, + vertex_location: str, + api_base: Optional[str] = None, + ) -> str: + """ + Build the count-tokens endpoint URL for a partner model. + + Args: + model: The model name + project_id: Google Cloud project ID + vertex_location: Vertex AI location (e.g., "us-east5") + api_base: Optional custom API base URL + + Returns: + Full endpoint URL for the count-tokens API + """ + publisher = self._get_publisher_for_model(model) + + # Use custom api_base if provided, otherwise construct default + if api_base: + base_url = api_base + else: + base_url = f"https://{vertex_location}-aiplatform.googleapis.com" + + # Construct the count-tokens endpoint + # Format: /v1/projects/{project}/locations/{location}/publishers/{publisher}/models/count-tokens:rawPredict + endpoint = ( + f"{base_url}/v1/projects/{project_id}/locations/{vertex_location}/" + f"publishers/{publisher}/models/count-tokens:rawPredict" + ) + + return endpoint + + async def handle_count_tokens_request( + self, + model: str, + request_data: Dict[str, Any], + litellm_params: Dict[str, Any], + ) -> Dict[str, Any]: + """ + Handle token counting request for a Vertex AI partner model. + + Args: + model: The model name + request_data: Request payload (Anthropic Messages API format) + litellm_params: LiteLLM parameters containing credentials, project, location + + Returns: + Dict containing token count information + + Raises: + ValueError: If required parameters are missing or invalid + """ + # Validate request + if "messages" not in request_data: + raise ValueError("messages required for token counting") + + # Extract Vertex AI credentials and settings + vertex_credentials = self.get_vertex_ai_credentials(litellm_params) + vertex_project = self.get_vertex_ai_project(litellm_params) + vertex_location = self.get_vertex_ai_location(litellm_params) + + # Get access token and resolved project ID + access_token, project_id = await self._ensure_access_token_async( + credentials=vertex_credentials, + project_id=vertex_project, + custom_llm_provider="vertex_ai", + ) + + # Build the endpoint URL + endpoint_url = self._build_count_tokens_endpoint( + model=model, + project_id=project_id, + vertex_location=vertex_location or "us-central1", + api_base=litellm_params.get("api_base"), + ) + + # Prepare headers + headers = {"Authorization": f"Bearer {access_token}"} + + # Get async HTTP client + from litellm import LlmProviders + + async_client = get_async_httpx_client(llm_provider=LlmProviders.VERTEX_AI) + + # Make the request + # Note: Partner models (especially Claude) accept Anthropic Messages API format directly + response = await async_client.post( + endpoint_url, + headers=headers, + json=request_data, + timeout=30.0, + ) + + # Check for errors + if response.status_code != 200: + error_text = response.text + raise ValueError( + f"Token counting request failed with status {response.status_code}: {error_text}" + ) + + # Parse response + result = response.json() + + # Return token count + # Vertex AI Anthropic returns: {"input_tokens": 123} + return { + "input_tokens": result.get("input_tokens", 0), + "tokenizer_used": "vertex_ai_partner_models", + } diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index ea29970f0a..85b1a6bc0d 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -28,6 +28,7 @@ class VertexAIError(Exception): self.message ) # Call the base class constructor with the parameters it needs + class PartnerModelPrefixes(str, Enum): META_PREFIX = "meta/" DEEPSEEK_PREFIX = "deepseek-ai" @@ -64,7 +65,7 @@ class VertexAIPartnerModels(VertexBase): ): return True return False - + @staticmethod def should_use_openai_handler(model: str): OPENAI_LIKE_VERTEX_PROVIDERS = [ @@ -258,3 +259,77 @@ class VertexAIPartnerModels(VertexBase): if hasattr(e, "status_code"): raise e raise VertexAIError(status_code=500, message=str(e)) + + async def count_tokens( + self, + model: str, + messages: list, + litellm_params: dict, + vertex_project=None, + vertex_location=None, + vertex_credentials=None, + ): + """ + Count tokens for Vertex AI partner models (Anthropic Claude, Mistral, etc.) + + Args: + model: The model name (e.g., "claude-3-5-sonnet-20241022") + messages: List of messages in Anthropic Messages API format + litellm_params: LiteLLM parameters dict + vertex_project: Optional Google Cloud project ID + vertex_location: Optional Vertex AI location + vertex_credentials: Optional Vertex AI credentials + + Returns: + Dict containing token count information + """ + try: + import vertexai + except Exception as e: + raise VertexAIError( + status_code=400, + message=f"""vertexai import failed please run `pip install -U "google-cloud-aiplatform>=1.38"`. Got error: {e}""", + ) + + if not ( + hasattr(vertexai, "preview") or hasattr(vertexai.preview, "language_models") + ): + raise VertexAIError( + status_code=400, + message="""Upgrade vertex ai. Run `pip install "google-cloud-aiplatform>=1.38"`""", + ) + + try: + from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler import ( + VertexAIPartnerModelsTokenCounter, + ) + + # Prepare request data in Anthropic Messages API format + request_data = { + "model": model, + "messages": messages, + } + + # Prepare litellm_params with credentials + _litellm_params = litellm_params.copy() + if vertex_project: + _litellm_params["vertex_project"] = vertex_project + if vertex_location: + _litellm_params["vertex_location"] = vertex_location + if vertex_credentials: + _litellm_params["vertex_credentials"] = vertex_credentials + + # Call the token counter + token_counter = VertexAIPartnerModelsTokenCounter() + result = await token_counter.handle_count_tokens_request( + model=model, + request_data=request_data, + litellm_params=_litellm_params, + ) + + return result + + except Exception as e: + if hasattr(e, "status_code"): + raise e + raise VertexAIError(status_code=500, message=str(e)) diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index 534569f15a..41132a65f2 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -743,3 +743,79 @@ async def test_bedrock_count_tokens_endpoint(): await mock_count_tokens_handler( request_data, {}, "anthropic.claude-3-sonnet-20240229-v1:0" ) + + +@pytest.mark.asyncio +async def test_vertex_ai_anthropic_token_counting(): + """ + Unit test for Vertex AI Anthropic token counting with mocked API calls. + + This tests the token counting implementation for Vertex AI partner models + without making actual API calls. Mocks at the handler level to test the full flow. + """ + from unittest.mock import AsyncMock, patch, MagicMock + + # Mock the Vertex AI partner models token counter response + mock_token_response = { + "input_tokens": 15, + "tokenizer_used": "vertex_ai_partner_models", + } + + llm_router = Router( + model_list=[ + { + "model_name": "vertex_ai/claude-3-5-sonnet-20241022", + "litellm_params": { + "model": "vertex_ai/claude-3-5-sonnet-20241022", + "vertex_project": "test-project", + "vertex_location": "us-east5", + }, + } + ] + ) + + setattr(litellm.proxy.proxy_server, "llm_router", llm_router) + + # Mock the lower level handler method + with patch( + "litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler.VertexAIPartnerModelsTokenCounter.handle_count_tokens_request" + ) as mock_handle_count_tokens: + mock_handle_count_tokens.return_value = mock_token_response + + # Test with messages format and call_endpoint=True + response = await token_counter( + request=TokenCountRequest( + model="vertex_ai/claude-3-5-sonnet-20241022", + messages=[ + { + "role": "user", + "content": "Hello Claude on Vertex AI! How are you?", + } + ], + ), + call_endpoint=True, + ) + + # Validate that handle_count_tokens_request was called + assert mock_handle_count_tokens.called + + # Verify the call arguments + call_args = mock_handle_count_tokens.call_args + assert call_args is not None + assert call_args.kwargs["model"] == "claude-3-5-sonnet-20241022" + assert "messages" in call_args.kwargs["request_data"] + assert ( + call_args.kwargs["request_data"]["messages"][0]["content"] + == "Hello Claude on Vertex AI! How are you?" + ) + + # Validate response structure + assert response.model_used == "claude-3-5-sonnet-20241022" + assert response.request_model == "vertex_ai/claude-3-5-sonnet-20241022" + assert response.total_tokens == 15 + assert response.tokenizer_type == "vertex_ai_partner_models" + + # Validate original response contains input_tokens + assert response.original_response is not None + assert "input_tokens" in response.original_response + assert response.original_response["input_tokens"] == 15 diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index 02dac0a93d..db457999a0 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -856,4 +856,121 @@ def test_get_token_url(): assert "v1beta1" not in url assert "/v1/" in url - pass \ No newline at end of file + pass + + +@pytest.mark.asyncio +async def test_vertex_ai_token_counter_routes_partner_models(): + """ + Test that VertexAITokenCounter correctly routes partner models (Claude, Mistral, etc.) + to the partner models token counter instead of the Gemini token counter. + """ + from unittest.mock import AsyncMock, 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 a Claude model (partner model) + result = 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": "us-east5", + } + }, + request_model="vertex_ai/claude-3-5-sonnet-20241022", + ) + + # Verify partner models handler was called + assert mock_partner_count_tokens.called + assert result is not None + assert isinstance(result, TokenCountResponse) + assert result.total_tokens == 42 + assert result.tokenizer_type == "vertex_ai_partner_models" + + +@pytest.mark.asyncio +async def test_vertex_ai_token_counter_routes_gemini_models(): + """ + Test that VertexAITokenCounter correctly routes Gemini models + to the Gemini token counter (not partner models). + """ + from unittest.mock import AsyncMock, patch + from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter + from litellm.types.utils import TokenCountResponse + + token_counter = VertexAITokenCounter() + + # Mock the Gemini handler (different import path) + with patch( + "litellm.llms.vertex_ai.count_tokens.handler.VertexAITokenCounter.acount_tokens" + ) as mock_gemini_count_tokens: + mock_gemini_count_tokens.return_value = { + "totalTokens": 50, + "tokenizer_used": "gemini", + } + + # Test with a Gemini model (not a partner model) + result = await token_counter.count_tokens( + model_to_use="gemini-1.5-pro", + messages=[{"role": "user", "content": "Hello"}], + contents=None, + deployment={ + "litellm_params": { + "vertex_project": "test-project", + "vertex_location": "us-central1", + } + }, + request_model="vertex_ai/gemini-1.5-pro", + ) + + # Verify Gemini handler was called + assert mock_gemini_count_tokens.called + assert result is not None + assert isinstance(result, TokenCountResponse) + assert result.total_tokens == 50 + + +@pytest.mark.asyncio +async def test_vertex_ai_partner_model_detection(): + """ + Test that VertexAIPartnerModels.is_vertex_partner_model correctly identifies + partner models (Claude, Mistral, Llama, etc.). + """ + from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( + VertexAIPartnerModels, + ) + + # Test Claude models (should be detected as partner model) + assert VertexAIPartnerModels.is_vertex_partner_model("claude-3-5-sonnet-20241022") + assert VertexAIPartnerModels.is_vertex_partner_model("claude-3-opus-20240229") + assert VertexAIPartnerModels.is_vertex_partner_model("claude-3-haiku-20240307") + + # Test Mistral models + assert VertexAIPartnerModels.is_vertex_partner_model("mistral-large-2407") + assert VertexAIPartnerModels.is_vertex_partner_model("mistral-7b-instruct-v0.3") + + # Test Meta/Llama models + assert VertexAIPartnerModels.is_vertex_partner_model("meta/llama-3.1-405b") + + # Test Gemini models (should NOT be detected as partner model) + assert not VertexAIPartnerModels.is_vertex_partner_model("gemini-1.5-pro") + assert not VertexAIPartnerModels.is_vertex_partner_model("gemini-1.0-pro") + assert not VertexAIPartnerModels.is_vertex_partner_model("gemini-pro-vision") + + # Test other non-partner models + assert not VertexAIPartnerModels.is_vertex_partner_model("text-bison-001") + assert not VertexAIPartnerModels.is_vertex_partner_model("chat-bison-001") \ No newline at end of file