[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
This commit is contained in:
Ishaan Jaff 2025-08-12 20:53:47 -07:00 • committed by GitHub
parent 74ddd21067
commit 1249385a99
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 272 additions and 14 deletions

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -1,4 +1,4 @@
model_list:
- model_name: vertex_ai/*
litellm_params:
model: gemini/*
model: vertex_ai/*

View file

@ -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:

View file

@ -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": [