mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
[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:
parent
74ddd21067
commit
1249385a99
6 changed files with 272 additions and 14 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
46
litellm/llms/vertex_ai/count_tokens/handler.py
Normal file
46
litellm/llms/vertex_ai/count_tokens/handler.py
Normal 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
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
model_list:
|
||||
- model_name: vertex_ai/*
|
||||
litellm_params:
|
||||
model: gemini/*
|
||||
model: vertex_ai/*
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue