mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(vertex_ai): exclude Gemini Google Search grounding tokens from input token billing (#33742)
* fix(vertex_ai): exclude Google Search grounding tokens from Gemini input token billing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): stub get_configured_token_limits on mocked routers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Krrish Dholakia <krrishdholakia@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
f759c75466
commit
07e07e6e2b
4 changed files with 136 additions and 1 deletions
|
|
@ -1744,6 +1744,30 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
return non_thinking_tokens == usage_metadata.get("totalTokenCount", 0)
|
||||
|
||||
@staticmethod
|
||||
def _response_has_search_grounding(
|
||||
completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage],
|
||||
) -> bool:
|
||||
"""
|
||||
Whether the response used Grounding with Google Search, detected via
|
||||
groundingMetadata.webSearchQueries (an actual web search was performed).
|
||||
|
||||
Google bills grounding-with-Google-Search retrieved tokens separately (a per-request /
|
||||
per-query search fee) and excludes them from input token billing, unlike URL context /
|
||||
File Search / code execution whose tool-use tokens are charged at the input token rate.
|
||||
URL context also emits groundingMetadata (with groundingChunks but no webSearchQueries),
|
||||
so presence of groundingMetadata alone is not a sufficient signal.
|
||||
See https://ai.google.dev/gemini-api/docs/pricing and
|
||||
https://github.com/BerriAI/litellm/discussions/33198
|
||||
"""
|
||||
if "candidates" not in completion_response:
|
||||
return False
|
||||
for candidate in completion_response["candidates"] or []:
|
||||
grounding_metadata, _, _, _ = VertexGeminiConfig._extract_candidate_metadata(candidate)
|
||||
if VertexGeminiConfig._calculate_web_search_requests(grounding_metadata):
|
||||
return True
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def _calculate_usage(
|
||||
completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage],
|
||||
|
|
@ -1899,12 +1923,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
tool_use_tokens=tool_use_prompt_tokens,
|
||||
)
|
||||
|
||||
billable_tool_use_prompt_tokens = (
|
||||
0
|
||||
if VertexGeminiConfig._response_has_search_grounding(completion_response)
|
||||
else (tool_use_prompt_tokens or 0)
|
||||
)
|
||||
|
||||
completion_tokens = response_tokens or completion_response["usageMetadata"].get("candidatesTokenCount", 0)
|
||||
if not VertexGeminiConfig.is_candidate_token_count_inclusive(usage_metadata) and reasoning_tokens:
|
||||
completion_tokens = reasoning_tokens + completion_tokens
|
||||
## GET USAGE ##
|
||||
usage = Usage(
|
||||
prompt_tokens=usage_metadata.get("promptTokenCount", 0) + (tool_use_prompt_tokens or 0),
|
||||
prompt_tokens=usage_metadata.get("promptTokenCount", 0) + billable_tool_use_prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=usage_metadata.get("totalTokenCount", 0),
|
||||
prompt_tokens_details=prompt_tokens_details,
|
||||
|
|
|
|||
|
|
@ -547,6 +547,103 @@ def test_vertex_ai_non_grounded_usage_omits_tool_use_tokens():
|
|||
assert not hasattr(usage.prompt_tokens_details, "tool_use_tokens")
|
||||
|
||||
|
||||
def test_response_has_search_grounding_detection():
|
||||
"""
|
||||
Only groundingMetadata.webSearchQueries signals an actual Google Search. URL context also
|
||||
emits groundingMetadata (groundingChunks but no webSearchQueries) and must not be treated
|
||||
as search grounding.
|
||||
"""
|
||||
assert (
|
||||
VertexGeminiConfig._response_has_search_grounding(
|
||||
{"candidates": [{"groundingMetadata": {"webSearchQueries": ["latest nobel physics"]}}]}
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert (
|
||||
VertexGeminiConfig._response_has_search_grounding(
|
||||
{
|
||||
"candidates": [
|
||||
{
|
||||
"urlContextMetadata": {"urlMetadata": []},
|
||||
"groundingMetadata": {
|
||||
"groundingChunks": [{"web": {"uri": "https://example.com", "title": "Example"}}]
|
||||
},
|
||||
}
|
||||
]
|
||||
}
|
||||
)
|
||||
is False
|
||||
)
|
||||
assert (
|
||||
VertexGeminiConfig._response_has_search_grounding({"candidates": [{"groundingMetadata": {"webSearchQueries": []}}]})
|
||||
is False
|
||||
)
|
||||
assert VertexGeminiConfig._response_has_search_grounding({"candidates": []}) is False
|
||||
assert VertexGeminiConfig._response_has_search_grounding({}) is False
|
||||
|
||||
|
||||
def test_vertex_ai_search_grounding_tool_use_tokens_excluded_from_prompt_tokens():
|
||||
"""
|
||||
Grounding with Google Search retrieved tokens are not billed at the input token rate
|
||||
(Google charges a separate per-request / per-query search fee), so toolUsePromptTokenCount
|
||||
must be surfaced on prompt_tokens_details.tool_use_tokens but excluded from prompt_tokens.
|
||||
See https://ai.google.dev/gemini-api/docs/pricing and
|
||||
https://github.com/BerriAI/litellm/discussions/33198
|
||||
"""
|
||||
v = VertexGeminiConfig()
|
||||
completion_response = {
|
||||
"candidates": [{"groundingMetadata": {"webSearchQueries": ["latest nobel physics"]}}],
|
||||
"usageMetadata": UsageMetadata(
|
||||
promptTokenCount=19,
|
||||
candidatesTokenCount=304,
|
||||
thoughtsTokenCount=122,
|
||||
toolUsePromptTokenCount=142,
|
||||
totalTokenCount=587,
|
||||
),
|
||||
}
|
||||
|
||||
usage = v._calculate_usage(completion_response=completion_response)
|
||||
|
||||
assert usage.prompt_tokens == 19
|
||||
assert usage.completion_tokens == 304 + 122
|
||||
assert usage.total_tokens == 587
|
||||
assert usage.prompt_tokens_details.tool_use_tokens == 142
|
||||
assert usage.total_tokens - usage.prompt_tokens - usage.completion_tokens == 142
|
||||
|
||||
|
||||
def test_vertex_ai_url_context_tool_use_tokens_billed_as_input_tokens():
|
||||
"""
|
||||
URL context / File Search / code execution tool-use tokens are billed as input tokens, so
|
||||
toolUsePromptTokenCount is folded into prompt_tokens when the response is not search grounded.
|
||||
"""
|
||||
v = VertexGeminiConfig()
|
||||
completion_response = {
|
||||
"candidates": [
|
||||
{
|
||||
"urlContextMetadata": {"urlMetadata": []},
|
||||
"groundingMetadata": {
|
||||
"groundingChunks": [{"web": {"uri": "https://example.com", "title": "Example"}}]
|
||||
},
|
||||
}
|
||||
],
|
||||
"usageMetadata": UsageMetadata(
|
||||
promptTokenCount=19,
|
||||
candidatesTokenCount=304,
|
||||
thoughtsTokenCount=122,
|
||||
toolUsePromptTokenCount=142,
|
||||
totalTokenCount=587,
|
||||
),
|
||||
}
|
||||
|
||||
usage = v._calculate_usage(completion_response=completion_response)
|
||||
|
||||
assert usage.prompt_tokens == 19 + 142
|
||||
assert usage.completion_tokens == 304 + 122
|
||||
assert usage.total_tokens == 587
|
||||
assert usage.prompt_tokens_details.tool_use_tokens == 142
|
||||
assert usage.total_tokens - usage.prompt_tokens - usage.completion_tokens == 0
|
||||
|
||||
|
||||
def test_streaming_chunk_includes_reasoning_tokens():
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
|
|
|
|||
|
|
@ -1727,6 +1727,7 @@ class TestModelInfoEndpoint:
|
|||
"gpt-3.5-turbo",
|
||||
]
|
||||
mock_router.get_model_access_groups.return_value = {}
|
||||
mock_router.get_configured_token_limits.return_value = (None, None)
|
||||
mock_get_key_models.return_value = ["gpt-4", "claude-3"]
|
||||
mock_get_team_models.return_value = ["gpt-3.5-turbo"]
|
||||
mock_get_complete_models.return_value = [
|
||||
|
|
@ -1812,6 +1813,7 @@ class TestModelInfoEndpoint:
|
|||
# Setup mocks
|
||||
mock_router.get_model_names.return_value = ["team-model-1"]
|
||||
mock_router.get_model_access_groups.return_value = {}
|
||||
mock_router.get_configured_token_limits.return_value = (None, None)
|
||||
mock_get_key_models.return_value = []
|
||||
mock_get_team_models.return_value = ["team-model-1"]
|
||||
mock_get_complete_models.return_value = ["team-model-1"]
|
||||
|
|
|
|||
|
|
@ -725,6 +725,7 @@ async def test_v1_models_translates_team_model_for_access_group_key(monkeypatch)
|
|||
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
|
||||
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
|
||||
router.get_fully_blocked_model_names.return_value = set()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.model_list = [team_dep]
|
||||
router.get_model_list.return_value = [team_dep]
|
||||
|
||||
|
|
@ -766,6 +767,7 @@ async def test_v1_models_keeps_internal_names_when_public_name_flag_disabled(
|
|||
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
|
||||
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
|
||||
router.get_fully_blocked_model_names.return_value = set()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.model_list = [team_dep]
|
||||
router.get_model_list.return_value = [team_dep]
|
||||
|
||||
|
|
@ -800,6 +802,7 @@ async def test_v1_models_translates_team_model_with_metadata(monkeypatch):
|
|||
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
|
||||
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
|
||||
router.get_fully_blocked_model_names.return_value = set()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.model_list = [team_dep]
|
||||
router.get_model_list.return_value = [team_dep]
|
||||
router.get_model_group_info.return_value = None
|
||||
|
|
@ -845,6 +848,7 @@ async def test_v1_models_metadata_fallbacks_use_internal_routing_key(monkeypatch
|
|||
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
|
||||
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
|
||||
router.get_fully_blocked_model_names.return_value = set()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.model_list = [team_dep]
|
||||
router.get_model_list.return_value = [team_dep]
|
||||
# Fallbacks are keyed on the internal routing name, as the router stores them.
|
||||
|
|
@ -901,6 +905,7 @@ async def test_v1_models_metadata_does_not_leak_other_team_fallbacks(monkeypatch
|
|||
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
|
||||
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
|
||||
router.get_fully_blocked_model_names.return_value = set()
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
router.model_list = [team_x, team_y]
|
||||
router.get_model_list.return_value = [team_x, team_y]
|
||||
router.fallbacks = [
|
||||
|
|
@ -1155,6 +1160,7 @@ def test_translate_team_model_names_for_listing_respects_legacy_flag():
|
|||
def _public_named_router(*team_rows: dict) -> MagicMock:
|
||||
router = MagicMock()
|
||||
router.get_model_list.return_value = list(team_rows)
|
||||
router.get_configured_token_limits.return_value = (None, None)
|
||||
return router
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue