From 0770f663c1c9300474c52d170aea81986c08e78b Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 15 Sep 2026 01:35:36 -0700 Subject: [PATCH] fix(vertex-live): report repeated search queries once per grounded turn The session usage collapsed duplicate query strings across turns while the price was per turn, so two turns asking the same question paid two fees yet reported web_search_requests 1. Sum each turn's grounding requests so the counter matches the bill; duplicates within one turn still collapse. --- ...tex_ai_live_passthrough_logging_handler.py | 36 +++++++++++-------- .../test_vertex_ai_live_passthrough.py | 30 ++++++++++++++++ 2 files changed, 52 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py index 57ec960b032..0ff02b29c58 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_ai_live_passthrough_logging_handler.py @@ -13,6 +13,7 @@ from typing import Final, Literal, TypeAlias from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.vertex_ai.gemini.grounding_requests import GroundingRequests, calculate_grounding_requests from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import ( BasePassthroughLoggingHandler, ) @@ -28,6 +29,8 @@ from litellm.types.utils import ( Usage, ) +_NO_GROUNDING: Final = GroundingRequests(web_search_requests=None, google_maps_grounding_requests=None) + _AGGREGATED_FIELDS: Final = frozenset( { "promptTokenCount", @@ -72,6 +75,18 @@ def _turns(websocket_messages: Sequence[object]) -> tuple[tuple[object, ...], .. return tuple(tuple(websocket_messages[start:end]) for start, end in pairwise((0, *closes))) +def _session_grounding_requests(websocket_messages: Sequence[object]) -> GroundingRequests: + per_turn: Final = tuple( + calculate_grounding_requests(_grounding_metadata(turn)) for turn in _turns(websocket_messages) + ) + web_search_requests: Final = sum(requests.web_search_requests or 0 for requests in per_turn) + google_maps_grounding_requests: Final = sum(requests.google_maps_grounding_requests or 0 for requests in per_turn) + return GroundingRequests( + web_search_requests=web_search_requests or None, + google_maps_grounding_requests=google_maps_grounding_requests or None, + ) + + _SummedField: TypeAlias = Literal[ "input_cost", "output_cost", @@ -223,7 +238,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): def _create_usage_object_from_metadata( usage_metadata: dict, model: str, - grounding_metadata: Sequence[Mapping[str, object]] = (), + grounding_requests: GroundingRequests = _NO_GROUNDING, ) -> Usage: """ Create a LiteLLM Usage object from Live API usage metadata. @@ -231,8 +246,8 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): Args: usage_metadata: Usage metadata from the Live API response model: The model name - grounding_metadata: Every ``serverContent.groundingMetadata`` the session emitted, so - Search and Maps grounding carry their per-query charge + grounding_requests: The Search and Maps grounding requests summed over the session's + turns, matching the per-turn charge Returns: LiteLLM Usage object @@ -252,7 +267,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0) or sum(prompt_by_modality.values()) completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0) or sum(candidates_by_modality.values()) - usage: Final = Usage( + return Usage( prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, total_tokens=usage_metadata.get("totalTokenCount", 0) or (prompt_tokens + completion_tokens), @@ -262,6 +277,8 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): image_tokens=prompt_by_modality.get("IMAGE"), video_tokens=prompt_by_modality.get("VIDEO"), tool_use_tokens=usage_metadata.get("toolUsePromptTokenCount") or None, + web_search_requests=grounding_requests.web_search_requests, + google_maps_grounding_requests=grounding_requests.google_maps_grounding_requests, ), completion_tokens_details=CompletionTokensDetailsWrapper( text_tokens=candidates_by_modality.get("TEXT"), @@ -270,15 +287,6 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): video_tokens=candidates_by_modality.get("VIDEO"), ), ) - if grounding_metadata: - from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( - VertexGeminiConfig, - ) - - VertexGeminiConfig._set_grounding_usage_counters( # pyright: ignore[reportPrivateUsage] # shared with the chat path; no public alias exists yet - usage, grounding_metadata - ) - return usage def _session_usage(self, websocket_messages: Sequence[object], model: str) -> Usage | None: usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages) @@ -286,7 +294,7 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler): return None return self._create_usage_object_from_metadata( usage_metadata=usage_metadata, - grounding_metadata=_grounding_metadata(websocket_messages), + grounding_requests=_session_grounding_requests(websocket_messages), model=model, ) diff --git a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py index 70c2fda369b..8b3dc436b8f 100644 --- a/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py +++ b/tests/pass_through_unit_tests/test_vertex_ai_live_passthrough.py @@ -534,6 +534,36 @@ class TestVertexAILivePassthroughLoggingHandler: assert two_breakdown["total_cost"] == pytest.approx(two_cost) assert two_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"]) + def test_a_query_repeated_across_turns_is_reported_once_per_turn(self, handler): + """The reported query count must agree with the bill, which charges every grounded turn. + + The session usage collapsed duplicate query strings across turns while the price was + per turn, so two turns asking the same question paid two fees yet reported one query. + Duplicates within one turn still collapse, since that turn ran one search. + """ + head, turn = self._live_messages(self.AUDIO_SESSION[:1]) + grounding = self._grounding_frame({"webSearchQueries": ["q"]}) + logging_obj = self._priced_logging_obj() + + result = handler.vertex_ai_live_passthrough_handler( + websocket_messages=[head, grounding, turn, grounding, turn], + logging_obj=logging_obj, + url_route="/vertex_ai/live", + start_time=datetime.now(), + end_time=datetime.now(), + request_body={}, + model=self.NATIVE_AUDIO_MODEL, + custom_llm_provider="vertex_ai", + ) + _, one_breakdown = self._billed_session(handler, [head, grounding, turn]) + repeated_within_turn = handler._session_usage( + [head, self._grounding_frame({"webSearchQueries": ["q", "q"]}), turn], self.NATIVE_AUDIO_MODEL + ) + + assert result["result"].usage.prompt_tokens_details.web_search_requests == 2 + assert logging_obj.cost_breakdown["tool_usage_cost"] == pytest.approx(2 * one_breakdown["tool_usage_cost"]) + assert repeated_within_turn.prompt_tokens_details.web_search_requests == 1 + def test_the_fixed_cost_margin_is_charged_once_per_session(self, handler): """A fixed cost margin is a flat per-request fee, and a Live session is one spend row.