mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(vertex_ai): surface Gemini grounding toolUsePromptTokenCount in Usage (#33533)
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
0223383d94
commit
4cfc987f56
4 changed files with 71 additions and 8 deletions
|
|
@ -1731,18 +1731,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"""
|
||||
Check if the candidate token count is inclusive of the thinking token count
|
||||
|
||||
if prompttokencount + candidatesTokenCount == totalTokenCount, then the candidate token count is inclusive of the thinking token count
|
||||
if promptTokenCount + candidatesTokenCount + toolUsePromptTokenCount == totalTokenCount, then the candidate token count is inclusive of the thinking token count
|
||||
|
||||
else the candidate token count is exclusive of the thinking token count
|
||||
|
||||
Addresses - https://github.com/BerriAI/litellm/pull/10141#discussion_r2052272035
|
||||
"""
|
||||
if usage_metadata.get("promptTokenCount", 0) + usage_metadata.get(
|
||||
"candidatesTokenCount", 0
|
||||
) == usage_metadata.get("totalTokenCount", 0):
|
||||
return True
|
||||
else:
|
||||
return False
|
||||
non_thinking_tokens = (
|
||||
usage_metadata.get("promptTokenCount", 0)
|
||||
+ usage_metadata.get("candidatesTokenCount", 0)
|
||||
+ usage_metadata.get("toolUsePromptTokenCount", 0)
|
||||
)
|
||||
return non_thinking_tokens == usage_metadata.get("totalTokenCount", 0)
|
||||
|
||||
@staticmethod
|
||||
def _calculate_usage(
|
||||
|
|
@ -1888,12 +1888,15 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
response_tokens_details = CompletionTokensDetailsWrapper()
|
||||
response_tokens_details.reasoning_tokens = reasoning_tokens
|
||||
|
||||
tool_use_prompt_tokens = usage_metadata.get("toolUsePromptTokenCount") or None
|
||||
|
||||
prompt_tokens_details = PromptTokensDetailsWrapper(
|
||||
cached_tokens=cached_tokens,
|
||||
audio_tokens=prompt_audio_tokens,
|
||||
text_tokens=prompt_text_tokens,
|
||||
image_tokens=prompt_image_tokens,
|
||||
video_tokens=prompt_video_tokens,
|
||||
tool_use_tokens=tool_use_prompt_tokens,
|
||||
)
|
||||
|
||||
completion_tokens = response_tokens or completion_response["usageMetadata"].get("candidatesTokenCount", 0)
|
||||
|
|
@ -1901,7 +1904,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
completion_tokens = reasoning_tokens + completion_tokens
|
||||
## GET USAGE ##
|
||||
usage = Usage(
|
||||
prompt_tokens=usage_metadata.get("promptTokenCount", 0),
|
||||
prompt_tokens=usage_metadata.get("promptTokenCount", 0) + (tool_use_prompt_tokens or 0),
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=usage_metadata.get("totalTokenCount", 0),
|
||||
prompt_tokens_details=prompt_tokens_details,
|
||||
|
|
|
|||
|
|
@ -299,6 +299,8 @@ class UsageMetadata(TypedDict, total=False):
|
|||
candidatesTokenCount: int
|
||||
responseTokenCount: int
|
||||
cachedContentTokenCount: int
|
||||
toolUsePromptTokenCount: int
|
||||
toolUsePromptTokensDetails: List[PromptTokensDetails]
|
||||
promptTokensDetails: List[PromptTokensDetails]
|
||||
cacheTokensDetails: List[PromptTokensDetails]
|
||||
thoughtsTokenCount: int
|
||||
|
|
|
|||
|
|
@ -1474,6 +1474,9 @@ class PromptTokensDetailsWrapper(
|
|||
web_search_requests: Optional[int] = None
|
||||
"""Number of web search requests made by the tool call. Used for Anthropic to calculate web search cost."""
|
||||
|
||||
tool_use_tokens: Optional[int] = None
|
||||
"""Prompt tokens consumed by server-side tool use (e.g. Gemini grounding via googleSearch)."""
|
||||
|
||||
character_count: Optional[int] = None
|
||||
"""Character count sent to the model. Used for Vertex AI multimodal embeddings."""
|
||||
|
||||
|
|
@ -1504,6 +1507,8 @@ class PromptTokensDetailsWrapper(
|
|||
del self.audio_length_seconds
|
||||
if self.web_search_requests is None:
|
||||
del self.web_search_requests
|
||||
if self.tool_use_tokens is None:
|
||||
del self.tool_use_tokens
|
||||
if self.cache_creation_tokens is None:
|
||||
del self.cache_creation_tokens
|
||||
if self.cache_creation_token_details is None:
|
||||
|
|
|
|||
|
|
@ -474,6 +474,22 @@ def test_vertex_ai_empty_content():
|
|||
reasoning_tokens=5,
|
||||
),
|
||||
),
|
||||
(
|
||||
UsageMetadata(
|
||||
promptTokenCount=4647,
|
||||
candidatesTokenCount=1495,
|
||||
totalTokenCount=29426,
|
||||
thoughtsTokenCount=10785,
|
||||
toolUsePromptTokenCount=12499,
|
||||
),
|
||||
False,
|
||||
Usage(
|
||||
prompt_tokens=17146,
|
||||
completion_tokens=12280,
|
||||
total_tokens=29426,
|
||||
reasoning_tokens=10785,
|
||||
),
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_vertex_ai_candidate_token_count_inclusive(
|
||||
|
|
@ -494,6 +510,43 @@ def test_vertex_ai_candidate_token_count_inclusive(
|
|||
assert usage.total_tokens == expected_usage.total_tokens
|
||||
|
||||
|
||||
def test_vertex_ai_grounded_usage_surfaces_tool_use_tokens():
|
||||
"""
|
||||
Grounded Gemini requests (googleSearch) return toolUsePromptTokenCount as part of totalTokenCount.
|
||||
Regression for https://github.com/BerriAI/litellm/issues/33530: it must be folded into
|
||||
prompt_tokens (so prompt_tokens + completion_tokens == total_tokens) and surfaced on
|
||||
prompt_tokens_details.tool_use_tokens.
|
||||
"""
|
||||
v = VertexGeminiConfig()
|
||||
usage_metadata = UsageMetadata(
|
||||
promptTokenCount=4647,
|
||||
candidatesTokenCount=1495,
|
||||
totalTokenCount=29426,
|
||||
thoughtsTokenCount=10785,
|
||||
toolUsePromptTokenCount=12499,
|
||||
)
|
||||
|
||||
usage = v._calculate_usage(completion_response={"usageMetadata": usage_metadata})
|
||||
|
||||
assert usage.prompt_tokens + usage.completion_tokens == usage.total_tokens
|
||||
assert usage.prompt_tokens_details.tool_use_tokens == 12499
|
||||
|
||||
|
||||
def test_vertex_ai_non_grounded_usage_omits_tool_use_tokens():
|
||||
"""Non-grounded responses must not surface a tool_use_tokens field on prompt_tokens_details."""
|
||||
v = VertexGeminiConfig()
|
||||
usage_metadata = UsageMetadata(
|
||||
promptTokenCount=10,
|
||||
candidatesTokenCount=10,
|
||||
totalTokenCount=20,
|
||||
)
|
||||
|
||||
usage = v._calculate_usage(completion_response={"usageMetadata": usage_metadata})
|
||||
|
||||
assert usage.prompt_tokens == 10
|
||||
assert not hasattr(usage.prompt_tokens_details, "tool_use_tokens")
|
||||
|
||||
|
||||
def test_streaming_chunk_includes_reasoning_tokens():
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue