diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 3193b72a7d9..624190a0b61 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -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, diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 5adc5b76990..95e8e6561f1 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -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, diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 8c6bdefedae..79c5f3ea549 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -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"] diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py index 25e84fb59a7..577af3dcffc 100644 --- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py +++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py @@ -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