From 1249385a994838a8ef0b0c3aa0d267c8676bf48a Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Tue, 12 Aug 2025 20:53:47 -0700 Subject: [PATCH] [Feat] GEMINI CLI - Add Token Counter for VertexAI Models (#13558) * add VertexAIModelInfo * working API call to vertex ai * add count_tokens MODE * _construct_url * test_vertex_ai_gemini_token_counting_with_contents --- litellm/llms/gemini/count_tokens/handler.py | 32 +++-- litellm/llms/vertex_ai/common_utils.py | 114 +++++++++++++++++- .../llms/vertex_ai/count_tokens/handler.py | 46 +++++++ litellm/proxy/proxy_config.yaml | 2 +- litellm/utils.py | 3 + .../test_proxy_token_counter.py | 89 +++++++++++++- 6 files changed, 272 insertions(+), 14 deletions(-) create mode 100644 litellm/llms/vertex_ai/count_tokens/handler.py diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index 1f13b0c3144..bcc8ab9553d 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union import httpx @@ -12,21 +12,36 @@ else: GenerateContentContentListUnionDict = Any class GoogleAIStudioTokenCounter: - def validate_environment( + + def _construct_url(self, model: str, api_base: Optional[str] = None) -> str: + """ + Construct the URL for the Google Gen AI Studio countTokens endpoint. + """ + base_url = api_base or "https://generativelanguage.googleapis.com" + return f"{base_url}/v1beta/models/{model}:countTokens" + + + async def validate_environment( self, + api_base: Optional[str] = None, api_key: Optional[str] = None, headers: Optional[Dict[str, Any]] = None, model: str = "", litellm_params: Optional[Dict[str, Any]] = None, - ): + ) -> Tuple[Dict[str, Any], str]: + """ + Returns a Tuple of headers and url for the Google Gen AI Studio countTokens endpoint. + """ from litellm.llms.gemini.google_genai.transformation import GoogleGenAIConfig - return GoogleGenAIConfig().validate_environment( + headers = GoogleGenAIConfig().validate_environment( api_key=api_key, headers=headers, model=model, litellm_params=litellm_params, ) + url = self._construct_url(model=model, api_base=api_base) + return headers, url async def acount_tokens( self, @@ -70,12 +85,11 @@ class GoogleAIStudioTokenCounter: Exception: For any other unexpected errors """ # Set up API base URL - base_url = api_base or "https://generativelanguage.googleapis.com" - url = f"{base_url}/v1beta/models/{model}:countTokens" - + # Prepare headers - headers = self.validate_environment( + headers, url = await self.validate_environment( api_key=api_key, + api_base=api_base, headers={}, model=model, litellm_params=kwargs, @@ -89,7 +103,7 @@ class GoogleAIStudioTokenCounter: async_httpx_client = get_async_httpx_client( llm_provider=LlmProviders.GEMINI, ) - + try: response = await async_httpx_client.post( url=url, diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index cceac0ea794..4931631d75d 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -7,8 +7,11 @@ import litellm from litellm import supports_response_schema, supports_system_messages, verbose_logger from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs +from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.types.llms.openai import AllMessageValues from litellm.types.llms.vertex_ai import PartType, Schema +from litellm.types.utils import TokenCountResponse class VertexAIError(BaseLLMException): @@ -63,7 +66,7 @@ def get_supports_response_schema( from typing import Literal, Optional all_gemini_url_modes = Literal[ - "chat", "embedding", "batch_embedding", "image_generation" + "chat", "embedding", "batch_embedding", "image_generation", "count_tokens" ] @@ -113,6 +116,12 @@ def _get_vertex_url( url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}" if model.isdigit(): url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}" + elif mode == "count_tokens": + endpoint = "countTokens" + if vertex_location == "global": + url = f"https://aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/global/publishers/google/models/{model}:{endpoint}" + else: + url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}" if not url or not endpoint: raise ValueError(f"Unable to get vertex url/endpoint for mode: {mode}") return url, endpoint @@ -148,10 +157,17 @@ def _get_gemini_url( url = "https://generativelanguage.googleapis.com/v1beta/{}:{}?key={}".format( _gemini_model_name, endpoint, gemini_api_key ) + elif mode == "count_tokens": + endpoint = "countTokens" + url = "https://generativelanguage.googleapis.com/v1beta/{}:{}?key={}".format( + _gemini_model_name, endpoint, gemini_api_key + ) elif mode == "image_generation": raise ValueError( "LiteLLM's `gemini/` route does not support image generation yet. Let us know if you need this feature by opening an issue at https://github.com/BerriAI/litellm/issues" ) + else: + raise ValueError(f"Unsupported mode: {mode}") return url, endpoint @@ -522,3 +538,99 @@ def is_global_only_vertex_model(model: str) -> bool: if supported_regions is None: return False return "global" in supported_regions + +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, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: Optional[str] = None, + 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]: + """ + Returns a list of models supported by this provider. + """ + raise NotImplementedError("Vertex AI models are not supported yet") + + @staticmethod + def get_api_key(api_key: Optional[str] = None) -> Optional[str]: + raise NotImplementedError("Vertex AI models are not supported yet") + + @staticmethod + def get_api_base( + api_base: Optional[str] = None, + ) -> Optional[str]: + raise NotImplementedError("Vertex AI models are not supported yet") + + + + @staticmethod + def get_base_model(model: str) -> Optional[str]: + """ + Returns the base model name from the given model name. + + Some providers like bedrock - can receive model=`invoke/anthropic.claude-3-opus-20240229-v1:0` or `converse/anthropic.claude-3-opus-20240229-v1:0` + This function will return `anthropic.claude-3-opus-20240229-v1:0` + """ + raise NotImplementedError("Vertex AI models are not supported yet") + + +class VertexAITokenCounter(BaseTokenCounter): + """Token counter implementation for Google AI Studio provider.""" + def should_use_token_counting_api( + 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, + messages: Optional[List[Dict[str, Any]]], + contents: Optional[List[Dict[str, Any]]], + deployment: Optional[Dict[str, Any]] = None, + request_model: str = "", + ) -> 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, + ) + + 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 \ No newline at end of file diff --git a/litellm/llms/vertex_ai/count_tokens/handler.py b/litellm/llms/vertex_ai/count_tokens/handler.py new file mode 100644 index 00000000000..d95c6801e57 --- /dev/null +++ b/litellm/llms/vertex_ai/count_tokens/handler.py @@ -0,0 +1,46 @@ +from typing import Any, Dict, Optional, Tuple + +from litellm.llms.gemini.count_tokens.handler import GoogleAIStudioTokenCounter +from litellm.llms.vertex_ai.vertex_llm_base import VertexBase + + +class VertexAITokenCounter(GoogleAIStudioTokenCounter, VertexBase): + async def validate_environment( + self, + api_base: Optional[str] = None, + api_key: Optional[str] = None, + headers: Optional[Dict[str, Any]] = None, + model: str = "", + litellm_params: Optional[Dict[str, Any]] = None, + ) -> Tuple[Dict[str, Any], str]: + """ + Returns a Tuple of headers and url for the Vertex AI countTokens endpoint. + """ + litellm_params = litellm_params or {} + vertex_credentials = self.get_vertex_ai_credentials(litellm_params=litellm_params) + vertex_project = self.get_vertex_ai_project(litellm_params=litellm_params) + vertex_location = self.get_vertex_ai_location(litellm_params=litellm_params) + should_use_v1beta1_features = self.is_using_v1beta1_features(litellm_params) + _auth_header, vertex_project = await self._ensure_access_token_async( + credentials=vertex_credentials, + project_id=vertex_project, + custom_llm_provider="vertex_ai", + ) + + auth_header, api_base = self._get_token_and_url( + model=model, + gemini_api_key=None, + auth_header=_auth_header, + vertex_project=vertex_project, + vertex_location=vertex_location, + vertex_credentials=vertex_credentials, + stream=False, + custom_llm_provider="vertex_ai", + api_base=None, + should_use_v1beta1_features=should_use_v1beta1_features, + mode="count_tokens", + ) + headers = { + "Authorization": f"Bearer {auth_header}", + } + return headers, api_base \ No newline at end of file diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 938c0fc49eb..755fe82118c 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -1,4 +1,4 @@ model_list: - model_name: vertex_ai/* litellm_params: - model: gemini/* + model: vertex_ai/* diff --git a/litellm/utils.py b/litellm/utils.py index 908844d06fe..fb4f1662a73 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7140,6 +7140,9 @@ class ProviderConfigManager: return litellm.OpenAIGPTConfig() elif LlmProviders.GEMINI == provider: return litellm.GeminiModelInfo() + elif LlmProviders.VERTEX_AI == provider: + from litellm.llms.vertex_ai.common_utils import VertexAIModelInfo + return VertexAIModelInfo() elif LlmProviders.LITELLM_PROXY == provider: return litellm.LiteLLMProxyChatConfig() elif LlmProviders.TOPAZ == provider: diff --git a/tests/proxy_unit_tests/test_proxy_token_counter.py b/tests/proxy_unit_tests/test_proxy_token_counter.py index e771c9330f4..fdce6fa3c84 100644 --- a/tests/proxy_unit_tests/test_proxy_token_counter.py +++ b/tests/proxy_unit_tests/test_proxy_token_counter.py @@ -26,11 +26,86 @@ verbose_proxy_logger.setLevel(level=logging.DEBUG) from litellm.proxy._types import TokenCountRequest from litellm.types.utils import TokenCountResponse +import json, tempfile from litellm import Router +def get_vertex_ai_creds_json() -> dict: + # Define the path to the vertex_key.json file + print("loading vertex ai credentials") + filepath = os.path.dirname(os.path.abspath(__file__)) + vertex_key_path = filepath + "/vertex_key.json" + # Read the existing content of the file or create an empty dictionary + try: + with open(vertex_key_path, "r") as file: + # Read the file content + print("Read vertexai file path") + content = file.read() + + # If the file is empty or not valid JSON, create an empty dictionary + if not content or not content.strip(): + service_account_key_data = {} + else: + # Attempt to load the existing JSON content + file.seek(0) + service_account_key_data = json.load(file) + except FileNotFoundError: + # If the file doesn't exist, create an empty dictionary + service_account_key_data = {} + + # Update the service_account_key_data with environment variables + private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") + private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") + private_key = private_key.replace("\\n", "\n") + service_account_key_data["private_key_id"] = private_key_id + service_account_key_data["private_key"] = private_key + + return service_account_key_data + + +def load_vertex_ai_credentials(): + # Define the path to the vertex_key.json file + print("loading vertex ai credentials") + filepath = os.path.dirname(os.path.abspath(__file__)) + vertex_key_path = filepath + "/vertex_key.json" + + # Read the existing content of the file or create an empty dictionary + try: + with open(vertex_key_path, "r") as file: + # Read the file content + print("Read vertexai file path") + content = file.read() + + # If the file is empty or not valid JSON, create an empty dictionary + if not content or not content.strip(): + service_account_key_data = {} + else: + # Attempt to load the existing JSON content + file.seek(0) + service_account_key_data = json.load(file) + except FileNotFoundError: + # If the file doesn't exist, create an empty dictionary + service_account_key_data = {} + + # Update the service_account_key_data with environment variables + private_key_id = os.environ.get("VERTEX_AI_PRIVATE_KEY_ID", "") + private_key = os.environ.get("VERTEX_AI_PRIVATE_KEY", "") + private_key = private_key.replace("\\n", "\n") + service_account_key_data["private_key_id"] = private_key_id + service_account_key_data["private_key"] = private_key + + # Create a temporary file + with tempfile.NamedTemporaryFile(mode="w+", delete=False) as temp_file: + # Write the updated content to the temporary files + json.dump(service_account_key_data, temp_file, indent=2) + + # Export the temporary file as GOOGLE_APPLICATION_CREDENTIALS + os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name) + + + @pytest.mark.asyncio async def test_vLLM_token_counting(): """ @@ -558,10 +633,12 @@ async def test_factory_registration(): @pytest.mark.asyncio -async def test_vertex_ai_gemini_token_counting_with_contents(): +@pytest.mark.parametrize("model_name", ["gemini-2.5-pro", "vertex-ai-gemini-2.5-pro"]) +async def test_vertex_ai_gemini_token_counting_with_contents(model_name): """ Test token counting for Vertex AI Gemini model using contents format with call_endpoint=True """ + load_vertex_ai_credentials() llm_router = Router( model_list=[ { @@ -569,7 +646,13 @@ async def test_vertex_ai_gemini_token_counting_with_contents(): "litellm_params": { "model": "gemini/gemini-2.5-pro", }, - } + }, + { + "model_name": "vertex-ai-gemini-2.5-pro", + "litellm_params": { + "model": "vertex_ai/gemini-2.5-pro", + }, + }, ] ) @@ -578,7 +661,7 @@ async def test_vertex_ai_gemini_token_counting_with_contents(): # Test with contents format and call_endpoint=True response = await token_counter( request=TokenCountRequest( - model="gemini-2.5-pro", + model=model_name, contents=[ { "parts": [