Merge pull request #40915 from BerriAI/litellm_internal_copy_37075

fix(vertex-live): bill Gemini Live sessions end to end (internal copy of #37075)
This commit is contained in:
Mateo Wang 2026-09-15 02:06:49 -07:00 committed by GitHub
commit 3ed6c19b8d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 1019 additions and 303 deletions

View file

@ -115,6 +115,19 @@ def _parse_setup(session_configuration_request: str) -> BidiGenerateContentSetup
return envelope.get("setup", empty_setup)
def _grounding_metadata_from_frame(frame: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
"""Read ``serverContent.groundingMetadata`` off the frame that carries the turn's usage.
Live reports grounding in the server frames rather than in ``usageMetadata``, and it emits both
on the same frame, so the per-query charge is countable at the point usage is built.
"""
server_content: Final = frame.get("serverContent")
if not isinstance(server_content, Mapping):
return ()
metadata: Final = server_content.get("groundingMetadata")
return (metadata,) if isinstance(metadata, Mapping) else ()
# Google bills Live transcription at an estimated 25 audio tokens/sec of input and
# 175 text tokens/min of output (ai.google.dev/gemini-api/docs/pricing).
GEMINI_LIVE_TRANSCRIBE_AUDIO_TOKENS_PER_SECOND: Final = 25
@ -323,7 +336,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
)
elif key == "input_audio_transcription" and value is not None:
optional_params["inputAudioTranscription"] = {}
elif key == "turn_detection":
elif key == "turn_detection" and value is not None:
value_typed = cast(OpenAIRealtimeTurnDetection, value)
if (
isinstance(value_typed, dict)
@ -1049,6 +1062,11 @@ class GeminiRealtimeConfig(BaseRealtimeConfig):
{**cast(dict, message), "usageMetadata": resolved_usage_metadata},
),
)
grounding_metadata: Final = _grounding_metadata_from_frame(message)
if grounding_metadata:
VertexGeminiConfig._set_grounding_usage_counters( # pyright: ignore[reportPrivateUsage] # shared with the chat path; no public alias exists yet
_chat_completion_usage, grounding_metadata
)
else:
_chat_completion_usage = get_empty_usage()

View file

@ -5,18 +5,105 @@ Handles cost tracking and logging for Vertex AI Live API WebSocket passthrough e
Supports different modalities: text, audio, video, and web search.
"""
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import Any, Final
from itertools import chain, pairwise
from types import MappingProxyType
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,
)
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
PassThroughEndpointLoggingTypedDict,
)
from litellm.types.utils import LlmProviders, ModelResponse, Usage
from litellm.utils import get_model_info
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
CostBreakdown,
LlmProviders,
ModelResponse,
PromptTokensDetailsWrapper,
Usage,
)
_NO_GROUNDING: Final = GroundingRequests(web_search_requests=None, google_maps_grounding_requests=None)
_AGGREGATED_FIELDS: Final = frozenset(
{
"promptTokenCount",
"candidatesTokenCount",
"totalTokenCount",
"toolUsePromptTokenCount",
"promptTokensDetails",
"candidatesTokensDetails",
}
)
def _detail_entries(raw: object) -> tuple[Mapping[str, object], ...]:
"""Narrow one turn's ``*TokensDetails`` value to the entries that are actually shaped like one."""
return tuple(entry for entry in raw if isinstance(entry, Mapping)) if isinstance(raw, Sequence) else ()
def _grounding_metadata(websocket_messages: Sequence[object]) -> tuple[Mapping[str, object], ...]:
"""Collect every ``serverContent.groundingMetadata`` a session emitted.
Live reports grounding in the server frames, never in ``usageMetadata``, so the per-query
charge has to be counted here rather than derived from the token totals.
"""
return tuple(
metadata
for message in websocket_messages
if isinstance(message, Mapping)
for server_content in (message.get("serverContent"),)
if isinstance(server_content, Mapping)
for metadata in (server_content.get("groundingMetadata"),)
if isinstance(metadata, Mapping)
)
def _turns(websocket_messages: Sequence[object]) -> tuple[tuple[object, ...], ...]:
"""Split a session at every ``usageMetadata`` frame; frames after the last one never got their usage."""
closes: Final = tuple(
index + 1
for index, message in enumerate(websocket_messages)
if isinstance(message, Mapping) and isinstance(message.get("usageMetadata"), dict)
)
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",
"tool_usage_cost",
"cache_read_cost",
"cache_creation_cost",
"reasoning_cost",
"original_cost",
"discount_amount",
"margin_fixed_amount",
"margin_total_amount",
]
def _summed(breakdowns: Sequence[CostBreakdown], field: _SummedField) -> float | None:
values: Final = tuple(value for breakdown in breakdowns if (value := breakdown.get(field)) is not None)
return sum(values) if values else None
class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
@ -48,186 +135,110 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""Return the LLM provider name."""
return LlmProviders.VERTEX_AI
@staticmethod
def _resolve_detail_counts(
details: Sequence[Mapping[str, object]],
declared_total: object,
) -> tuple[tuple[str, int], ...]:
"""
Pair each of one turn's ``*TokensDetails`` entries with its token count.
Live sometimes names the modality that carries the rest of a turn without a
``tokenCount``, and reading the absent key as zero drops those tokens from the
breakdown, so real audio ends up priced as text. A lone unpriced entry therefore takes
whatever the turn's declared count leaves over. Two or more cannot be told apart, so
they are left out and the cost calculator charges the remainder as text.
"""
priced: Final = tuple(
(str(detail.get("modality", "TEXT")), count)
for detail in details
if isinstance(count := detail.get("tokenCount"), int)
)
unpriced: Final = tuple(
str(detail.get("modality", "TEXT")) for detail in details if not isinstance(detail.get("tokenCount"), int)
)
if len(unpriced) != 1 or not isinstance(declared_total, int):
return priced
residual: Final = declared_total - sum(count for _, count in priced)
return priced if residual <= 0 else (*priced, (unpriced[0], residual))
@staticmethod
def _sum_by_modality(counts: Sequence[tuple[str, int]]) -> Mapping[str, int]:
"""Total the (modality, tokenCount) pairs of one or more turns per modality."""
return MappingProxyType({modality: sum(c for m, c in counts if m == modality) for modality, _ in counts})
@staticmethod
def _merged_modality_totals(
snapshots: Sequence[Mapping[str, object]],
count_key: str,
details_key: str,
) -> Mapping[str, int]:
"""Total every turn's per-modality counts, so the breakdown adds up the way the totals do."""
return VertexAILivePassthroughLoggingHandler._sum_by_modality(
tuple(
chain.from_iterable(
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
_detail_entries(snapshot.get(details_key)), snapshot.get(count_key)
)
for snapshot in snapshots
)
)
)
@staticmethod
def _extract_usage_metadata_from_websocket_messages(
websocket_messages: list[dict],
websocket_messages: Sequence[object],
) -> dict | None:
"""
Extract and aggregate usage metadata from a list of WebSocket messages.
Live emits one ``usageMetadata`` per turn and Google charges per turn for every token in
the session context window, which is the current turn's tokens plus all accumulated
tokens from previous turns, so the turns add up rather than restating each other. See
the Live API note under https://cloud.google.com/vertex-ai/generative-ai/pricing.
Args:
websocket_messages: List of WebSocket messages from the Live API
Returns:
Dictionary containing aggregated usage metadata, or None if not found
"""
all_usage_metadata: Final = []
snapshots: Final = tuple(
metadata
for message in websocket_messages
if isinstance(message, Mapping)
for metadata in (message.get("usageMetadata"),)
if isinstance(metadata, dict)
)
# Collect all usage metadata messages
for message in websocket_messages:
if isinstance(message, dict) and "usageMetadata" in message:
all_usage_metadata.append(message["usageMetadata"])
if not all_usage_metadata:
if not snapshots:
return None
# If only one usage metadata, return it as-is
if len(all_usage_metadata) == 1:
return all_usage_metadata[0]
# Aggregate multiple usage metadata messages
aggregated: Final[dict[str, Any]] = {
"promptTokenCount": 0,
"candidatesTokenCount": 0,
"totalTokenCount": 0,
"promptTokensDetails": [],
"candidatesTokensDetails": [],
prompt_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals(
snapshots, "promptTokenCount", "promptTokensDetails"
)
candidate_totals: Final = VertexAILivePassthroughLoggingHandler._merged_modality_totals(
snapshots, "candidatesTokenCount", "candidatesTokensDetails"
)
return {
**{key: value for key, value in snapshots[0].items() if key not in _AGGREGATED_FIELDS},
"promptTokenCount": sum(snapshot.get("promptTokenCount", 0) for snapshot in snapshots),
"candidatesTokenCount": sum(snapshot.get("candidatesTokenCount", 0) for snapshot in snapshots),
"totalTokenCount": sum(snapshot.get("totalTokenCount", 0) for snapshot in snapshots),
"toolUsePromptTokenCount": sum(snapshot.get("toolUsePromptTokenCount", 0) for snapshot in snapshots),
"promptTokensDetails": [
{"modality": modality, "tokenCount": count} for modality, count in prompt_totals.items() if count > 0
],
"candidatesTokensDetails": [
{"modality": modality, "tokenCount": count} for modality, count in candidate_totals.items() if count > 0
],
}
# Aggregate token counts
for usage in all_usage_metadata:
aggregated["promptTokenCount"] += usage.get("promptTokenCount", 0)
aggregated["candidatesTokenCount"] += usage.get("candidatesTokenCount", 0)
aggregated["totalTokenCount"] += usage.get("totalTokenCount", 0)
# Aggregate token details by modality
modality_totals: Final = {}
for usage in all_usage_metadata:
# Process prompt tokens details
for detail in usage.get("promptTokensDetails", []):
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality not in modality_totals:
modality_totals[modality] = {"prompt": 0, "candidate": 0}
modality_totals[modality]["prompt"] += token_count
# Process candidate tokens details
for detail in usage.get("candidatesTokensDetails", []):
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality not in modality_totals:
modality_totals[modality] = {"prompt": 0, "candidate": 0}
modality_totals[modality]["candidate"] += token_count
# Convert aggregated modality totals back to details format
for modality, totals in modality_totals.items():
if totals["prompt"] > 0:
aggregated["promptTokensDetails"].append({"modality": modality, "tokenCount": totals["prompt"]})
if totals["candidate"] > 0:
aggregated["candidatesTokensDetails"].append({"modality": modality, "tokenCount": totals["candidate"]})
# Add any additional fields from the first usage metadata
first_usage: Final = all_usage_metadata[0]
for key, value in first_usage.items():
if key not in aggregated:
aggregated[key] = value
return aggregated
@staticmethod
def _calculate_live_api_cost(
model: str,
usage_metadata: dict,
custom_llm_provider: str = "vertex_ai",
) -> float:
"""
Calculate cost for Vertex AI Live API based on usage metadata.
Args:
model: The model name (e.g., "gemini-2.0-flash-live-preview-04-09")
usage_metadata: Usage metadata from the Live API response
custom_llm_provider: The LLM provider (default: "vertex_ai")
Returns:
Total cost in USD
"""
try:
# Get model pricing information
model_info: Final = get_model_info(model=model, custom_llm_provider=custom_llm_provider)
verbose_proxy_logger.debug("Vertex AI Live API model info for '%s': %s", model, model_info)
# Check if pricing info is available
if not model_info or not model_info.get("input_cost_per_token"):
verbose_proxy_logger.error("No pricing info found for %s in local model pricing database", model)
return 0.0
total_cost = 0.0
# Extract token counts from usage metadata
prompt_token_count: Final = usage_metadata.get("promptTokenCount", 0)
candidates_token_count: Final = usage_metadata.get("candidatesTokenCount", 0)
# Calculate base text token costs
input_cost_per_token: Final = model_info.get("input_cost_per_token", 0.0)
output_cost_per_token: Final = model_info.get("output_cost_per_token", 0.0)
total_cost += prompt_token_count * input_cost_per_token
total_cost += candidates_token_count * output_cost_per_token
# Handle modality-specific costs if present
prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", [])
candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", [])
# Process prompt tokens by modality
for detail in prompt_tokens_details:
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality == "AUDIO":
audio_cost_per_token = model_info.get("input_cost_per_audio_token", 0.0)
total_cost += token_count * audio_cost_per_token
elif modality == "VIDEO":
# Video tokens are typically per second, but we'll treat as per token for now
video_cost_per_token = model_info.get("input_cost_per_video_per_second", 0.0)
total_cost += token_count * video_cost_per_token
# TEXT tokens are already handled above
# Process candidate tokens by modality
for detail in candidates_tokens_details:
modality = detail.get("modality", "TEXT")
token_count = detail.get("tokenCount", 0)
if modality == "AUDIO":
audio_cost_per_token = model_info.get("output_cost_per_audio_token", 0.0)
total_cost += token_count * audio_cost_per_token
elif modality == "VIDEO":
# Video tokens are typically per second, but we'll treat as per token for now
video_cost_per_token = model_info.get("output_cost_per_video_per_second", 0.0)
total_cost += token_count * video_cost_per_token
# TEXT tokens are already handled above
# Handle web search costs if present
tool_use_prompt_token_count: Final = usage_metadata.get("toolUsePromptTokenCount", 0)
if tool_use_prompt_token_count > 0:
# Web search typically has a fixed cost per request
web_search_cost: Final = model_info.get("web_search_cost_per_request", 0.0)
if isinstance(web_search_cost, (int, float)) and web_search_cost > 0:
total_cost += web_search_cost
else:
# Fallback to token-based pricing for tool use
total_cost += tool_use_prompt_token_count * input_cost_per_token
verbose_proxy_logger.debug(
f"Vertex AI Live API cost calculation - Model: {model}, "
f"Prompt tokens: {prompt_token_count}, "
f"Candidate tokens: {candidates_token_count}, "
f"Total cost: ${total_cost:.6f}"
)
return total_cost
except Exception as e:
verbose_proxy_logger.error("Error calculating Vertex AI Live API cost: %s", e)
return 0.0
@staticmethod
def _create_usage_object_from_metadata(
usage_metadata: dict,
model: str,
grounding_requests: GroundingRequests = _NO_GROUNDING,
) -> Usage:
"""
Create a LiteLLM Usage object from Live API usage metadata.
@ -235,48 +246,124 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
Args:
usage_metadata: Usage metadata from the Live API response
model: The model name
grounding_requests: The Search and Maps grounding requests summed over the session's
turns, matching the per-turn charge
Returns:
LiteLLM Usage object
"""
prompt_tokens: Final = usage_metadata.get("promptTokenCount", 0)
completion_tokens: Final = usage_metadata.get("candidatesTokenCount", 0)
total_tokens: Final = usage_metadata.get("totalTokenCount", 0)
prompt_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality(
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
_detail_entries(usage_metadata.get("promptTokensDetails")), usage_metadata.get("promptTokenCount")
)
)
candidates_by_modality: Final = VertexAILivePassthroughLoggingHandler._sum_by_modality(
VertexAILivePassthroughLoggingHandler._resolve_detail_counts(
_detail_entries(usage_metadata.get("candidatesTokensDetails")),
usage_metadata.get("candidatesTokenCount"),
)
)
# Create modality-specific token details if available
prompt_tokens_details: Final = usage_metadata.get("promptTokensDetails", [])
candidates_tokens_details: Final = usage_metadata.get("candidatesTokensDetails", [])
# Extract text tokens from details
text_prompt_tokens = 0
text_completion_tokens = 0
for detail in prompt_tokens_details:
if detail.get("modality") == "TEXT":
text_prompt_tokens = detail.get("tokenCount", 0)
break
for detail in candidates_tokens_details:
if detail.get("modality") == "TEXT":
text_completion_tokens = detail.get("tokenCount", 0)
break
# If no text tokens found in details, use total counts
if text_prompt_tokens == 0:
text_prompt_tokens = prompt_tokens
if text_completion_tokens == 0:
text_completion_tokens = completion_tokens
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())
return Usage(
prompt_tokens=text_prompt_tokens,
completion_tokens=text_completion_tokens,
total_tokens=total_tokens,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=usage_metadata.get("totalTokenCount", 0) or (prompt_tokens + completion_tokens),
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=prompt_by_modality.get("TEXT"),
audio_tokens=prompt_by_modality.get("AUDIO"),
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"),
audio_tokens=candidates_by_modality.get("AUDIO"),
image_tokens=candidates_by_modality.get("IMAGE"),
video_tokens=candidates_by_modality.get("VIDEO"),
),
)
def _session_usage(self, websocket_messages: Sequence[object], model: str) -> Usage | None:
usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages)
if usage_metadata is None:
return None
return self._create_usage_object_from_metadata(
usage_metadata=usage_metadata,
grounding_requests=_session_grounding_requests(websocket_messages),
model=model,
)
def _turn_cost(
self,
turn: Sequence[object],
model: str,
logging_obj: LiteLLMLoggingObj,
) -> tuple[float, CostBreakdown] | None:
usage: Final = self._session_usage(turn, model)
if usage is None:
return None
cost: Final = logging_obj._response_cost_calculator( # pyright: ignore[reportPrivateUsage] # the call's own calculator keeps custom pricing and the deployment's region in step with the spend row
result=ModelResponse(model=model, usage=usage),
litellm_model_name=model,
)
if cost is None:
return None
breakdown: Final = logging_obj.cost_breakdown
return None if breakdown is None else (cost, breakdown)
def _session_cost(
self,
websocket_messages: Sequence[object],
model: str,
logging_obj: LiteLLMLoggingObj,
) -> float | None:
"""Price each turn on its own tokens and grounding, so two grounded turns pay the query fee twice.
The fixed cost margin is a flat per-request fee, so the session's single spend row carries it once
rather than once per turn.
"""
turn_costs: Final = tuple(self._turn_cost(turn, model, logging_obj) for turn in _turns(websocket_messages))
priced: Final = tuple(turn_cost for turn_cost in turn_costs if turn_cost is not None)
if not priced or len(priced) != len(turn_costs):
return None
breakdowns: Final = tuple(breakdown for _, breakdown in priced)
first: Final = breakdowns[0]
fixed_margin: Final = first.get("margin_fixed_amount") or 0.0
duplicated_fixed_margin: Final = fixed_margin * (len(priced) - 1)
total_cost: Final = sum(cost for cost, _ in priced) - duplicated_fixed_margin
summed_margin_total: Final = _summed(breakdowns, "margin_total_amount")
margin_total_amount: Final = (
None if summed_margin_total is None else summed_margin_total - duplicated_fixed_margin
)
logging_obj.set_cost_breakdown(
input_cost=_summed(breakdowns, "input_cost") or 0.0,
output_cost=_summed(breakdowns, "output_cost") or 0.0,
total_cost=total_cost,
cost_for_built_in_tools_cost_usd_dollar=_summed(breakdowns, "tool_usage_cost") or 0.0,
original_cost=_summed(breakdowns, "original_cost"),
discount_percent=first.get("discount_percent"),
discount_amount=_summed(breakdowns, "discount_amount"),
margin_percent=first.get("margin_percent"),
margin_fixed_amount=first.get("margin_fixed_amount"),
margin_total_amount=margin_total_amount,
cache_read_cost=_summed(breakdowns, "cache_read_cost"),
cache_creation_cost=_summed(breakdowns, "cache_creation_cost"),
reasoning_cost=_summed(breakdowns, "reasoning_cost"),
service_tier=first.get("service_tier"),
data_residency=first.get("data_residency"),
vertex_location=first.get("vertex_location"),
)
return total_cost
def vertex_ai_live_passthrough_handler(
self,
websocket_messages: list[dict],
logging_obj,
websocket_messages: Sequence[object],
logging_obj: LiteLLMLoggingObj,
url_route: str,
start_time: datetime,
end_time: datetime,
@ -300,34 +387,25 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
"""
try:
# Extract model from request body or kwargs
model: Final = kwargs.get("model", "gemini-2.0-flash-live-preview-04-09")
requested_model: Final = kwargs.get("model")
model: Final = (
requested_model if isinstance(requested_model, str) else "gemini-2.0-flash-live-preview-04-09"
)
custom_llm_provider: Final = kwargs.get("custom_llm_provider", "vertex_ai")
verbose_proxy_logger.debug(
"Vertex AI Live API model: %s, custom_llm_provider: %s", model, custom_llm_provider
)
# Extract usage metadata from WebSocket messages
usage_metadata: Final = self._extract_usage_metadata_from_websocket_messages(websocket_messages)
usage: Final = self._session_usage(websocket_messages, model)
if not usage_metadata:
if usage is None:
verbose_proxy_logger.warning("No usage metadata found in Vertex AI Live API WebSocket messages")
return {
"result": None,
"kwargs": kwargs,
}
# Calculate cost using Live API specific pricing
response_cost: Final = self._calculate_live_api_cost(
model=model,
usage_metadata=usage_metadata,
custom_llm_provider=custom_llm_provider,
)
# Create Usage object for standard LiteLLM logging
usage: Final = self._create_usage_object_from_metadata(
usage_metadata=usage_metadata,
model=model,
)
response_cost: Final = self._session_cost(websocket_messages, model, logging_obj)
# Create a mock ModelResponse for standard logging
litellm_model_response: Final = ModelResponse(
@ -338,9 +416,9 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
usage=usage,
choices=[],
)
if response_cost is not None:
litellm_model_response._hidden_params["response_cost"] = response_cost # pyright: ignore[reportPrivateUsage] # the logger reads the cost off the response's hidden params; the constructor's hidden_params kwarg is reset by pydantic
# Update kwargs with cost information
kwargs["response_cost"] = response_cost
kwargs["model"] = model
kwargs["custom_llm_provider"] = custom_llm_provider
@ -348,12 +426,15 @@ class VertexAILivePassthroughLoggingHandler(BasePassthroughLoggingHandler):
import re
allowed_pattern: Final = re.compile(r"^[A-Za-z0-9._\-:]+$")
safe_model: Final = model if isinstance(model, str) and allowed_pattern.match(model) else "[REDACTED]"
safe_model: Final = model if allowed_pattern.match(model) else "[REDACTED]"
verbose_proxy_logger.debug(
f"Vertex AI Live API passthrough cost tracking - "
f"Model: {safe_model}, Cost: ${response_cost:.6f}, "
f"Prompt tokens: {usage.prompt_tokens}, "
f"Completion tokens: {usage.completion_tokens}"
"Vertex AI Live API passthrough cost tracking - Model: %s, "
"Prompt tokens: %s %s, Completion tokens: %s %s",
safe_model,
usage.prompt_tokens,
usage.prompt_tokens_details,
usage.completion_tokens,
usage.completion_tokens_details,
)
return {

View file

@ -2090,6 +2090,22 @@ def _rewrite_vertex_live_setup_model(text_data: str, setup_model_rewriter: Calla
return json.dumps({**message, "setup": {**setup, "model": rewritten_model}}) # mutable-ok: one-shot json payload
def _resolved_vertex_live_setup(
setup_data: Mapping[str, object], setup_model_rewriter: Callable[[str], str] | None
) -> Mapping[str, object]:
"""
Give the model extractor the same fully qualified path the upstream will receive.
Clients may name a bare gateway alias, which the rewriter turns into a ``projects/...`` path before
it reaches Vertex. The extractor only reads a path containing ``/models/``, so running it on the raw
frame logs the session as ``unknown`` at no cost, which is precisely the supported client form
"""
setup_model: Final = setup_data.get("model")
if setup_model_rewriter is None or not isinstance(setup_model, str):
return setup_data
return {**setup_data, "model": setup_model_rewriter(setup_model)}
def _truncated_close_reason(reason: str) -> str:
"""
Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character
@ -2314,7 +2330,9 @@ async def websocket_passthrough_request(
setup_data,
)
if isinstance(setup_data, dict) and "model" in setup_data:
extracted_model = _extract_model_from_vertex_ai_setup(setup_data)
extracted_model = _extract_model_from_vertex_ai_setup(
_resolved_vertex_live_setup(setup_data, setup_model_rewriter)
)
if extracted_model:
kwargs["model"] = extracted_model
kwargs["custom_llm_provider"] = "vertex_ai-language-models"

View file

@ -2831,6 +2831,22 @@ class LiteLLMCompletionResponsesConfig:
if cache_write_tokens is not None
else MappingProxyType({})
)
# The cost path reads the grounding counters off the input details, and a realtime
# session's usage is rebuilt from its own response.done, so dropping them here bills
# no per-query grounding fee at all.
grounding_request_counts: Final[Mapping[str, int]] = MappingProxyType(
{
counter: count
for counter, count in (
("web_search_requests", getattr(prompt_details, "web_search_requests", None)),
(
"google_maps_grounding_requests",
getattr(prompt_details, "google_maps_grounding_requests", None),
),
)
if count is not None
}
)
response_usage.input_tokens_details = InputTokensDetails(
cached_tokens=prompt_details.cached_tokens if prompt_details.cached_tokens is not None else 0,
text_tokens=prompt_details.text_tokens,
@ -2839,6 +2855,7 @@ class LiteLLMCompletionResponsesConfig:
cached_tokens_details if isinstance(cached_tokens_details, CachedTokensDetails) else None
),
**cache_write_extra,
**grounding_request_counts,
)
# Translate completion_tokens_details to output_tokens_details

View file

@ -1183,6 +1183,10 @@ class ResponseAPILoggingUtils:
response_api_usage.input_tokens_details, "cached_tokens_details", None
),
cache_write_tokens=getattr(response_api_usage.input_tokens_details, "cache_write_tokens", None),
web_search_requests=getattr(response_api_usage.input_tokens_details, "web_search_requests", None),
google_maps_grounding_requests=getattr(
response_api_usage.input_tokens_details, "google_maps_grounding_requests", None
),
)
completion_tokens_details: CompletionTokensDetailsWrapper | None = None
output_tokens_details: Final[OutputTokensDetails | None] = getattr(

View file

@ -6,12 +6,15 @@ including the logging handler, cost tracking, and WebSocket message processing.
"""
import json
from collections.abc import Sequence
from datetime import datetime
from unittest.mock import AsyncMock, Mock, patch, MagicMock
from typing import Dict, List, Any, Optional
import pytest
import httpx
import litellm
from typing_extensions import NotRequired, ReadOnly, TypedDict
# Add the parent directory to the system path
@ -22,10 +25,16 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import LlmProviders
from litellm.types.utils import CostBreakdown, LlmProviders, Usage
from litellm.proxy._types import UserAPIKeyAuth
class _LiveTurn(TypedDict):
prompt: ReadOnly[tuple[int, int]]
candidates: ReadOnly[tuple[int, int]]
candidate_audio_token_count_missing: NotRequired[ReadOnly[bool]]
class TestVertexAILivePassthroughLoggingHandler:
"""Test the Vertex AI Live Passthrough Logging Handler"""
@ -39,6 +48,7 @@ class TestVertexAILivePassthroughLoggingHandler:
"""Create a mock logging object"""
mock = MagicMock(spec=LiteLLMLoggingObj)
mock.model_call_details = {}
mock._response_cost_calculator.return_value = None
return mock
@pytest.fixture
@ -201,88 +211,490 @@ class TestVertexAILivePassthroughLoggingHandler:
assert text_prompt["tokenCount"] == 10
assert audio_prompt["tokenCount"] == 10
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
)
def test_calculate_cost_basic(self, mock_get_model_info, handler):
"""Test basic cost calculation"""
mock_get_model_info.return_value = {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
}
def test_usage_carries_every_modality(self, handler):
"""Regression: the Usage object reported only TEXT, so audio and image billed as nothing.
prompt_tokens must be the full count and the details must name each modality,
because the cost calculator prices audio and image from *_tokens_details.
"""
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
}
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
# The cost calculation may include additional factors, so we check it's reasonable
expected_min_cost = (100 * 0.000001) + (50 * 0.000002)
assert cost >= expected_min_cost
assert cost > 0
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
)
def test_calculate_cost_with_audio(self, mock_get_model_info, handler):
"""Test cost calculation with audio tokens"""
mock_get_model_info.return_value = {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
"input_cost_per_audio_token": 0.0001,
"output_cost_per_audio_token": 0.0002,
}
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
"promptTokenCount": 1300,
"candidatesTokenCount": 124,
"totalTokenCount": 1424,
"promptTokensDetails": [
{"modality": "TEXT", "tokenCount": 80},
{"modality": "AUDIO", "tokenCount": 20},
{"modality": "TEXT", "tokenCount": 13},
{"modality": "AUDIO", "tokenCount": 127},
{"modality": "IMAGE", "tokenCount": 1160},
],
"candidatesTokensDetails": [
{"modality": "TEXT", "tokenCount": 30},
{"modality": "AUDIO", "tokenCount": 20},
{"modality": "TEXT", "tokenCount": 29},
{"modality": "AUDIO", "tokenCount": 95},
],
}
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
usage = handler._create_usage_object_from_metadata(
usage_metadata=usage_metadata, model="gemini-live-2.5-flash"
)
# Should include both text and audio costs
assert cost > 0
assert cost > (100 * 0.000001) + (
50 * 0.000002
) # Should be higher due to audio
assert usage.prompt_tokens == 1300, "the full prompt count must survive, not just its text share"
assert usage.completion_tokens == 124
assert usage.prompt_tokens_details.text_tokens == 13
assert usage.prompt_tokens_details.audio_tokens == 127
assert usage.prompt_tokens_details.image_tokens == 1160
assert usage.completion_tokens_details.text_tokens == 29
assert usage.completion_tokens_details.audio_tokens == 95
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
def test_usage_sums_repeated_modality_entries(self, handler):
"""A modality can appear more than once across aggregated turns; sum, don't overwrite."""
usage = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": 40,
"candidatesTokenCount": 0,
"promptTokensDetails": [
{"modality": "IMAGE", "tokenCount": 10},
{"modality": "IMAGE", "tokenCount": 25},
{"modality": "TEXT", "tokenCount": 5},
],
},
model="gemini-live-2.5-flash",
)
assert usage.prompt_tokens_details.image_tokens == 35
assert usage.prompt_tokens_details.text_tokens == 5
NATIVE_AUDIO_MODEL = "gemini-live-2.5-flash-preview-native-audio-09-2025"
# A four-turn native-audio session. Google charges per turn for the whole session context
# window, so the prompt side repeats the accumulated audio while the candidates side reports
# only that turn's own response. The last turn names AUDIO and omits its tokenCount, which is
# the shape Live really emits at the end of a spoken answer.
AUDIO_SESSION: tuple[_LiveTurn, ...] = (
{"prompt": (14, 122), "candidates": (8, 20)},
{"prompt": (21, 182), "candidates": (5, 50)},
{"prompt": (24, 203), "candidates": (13, 27)},
{"prompt": (24, 203), "candidates": (0, 3), "candidate_audio_token_count_missing": True},
)
def test_calculate_cost_with_web_search(self, mock_get_model_info, handler):
"""Test cost calculation with web search (tool use)"""
mock_get_model_info.return_value = {
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
"web_search_cost_per_request": 0.01,
}
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
"toolUsePromptTokenCount": 10,
}
@staticmethod
def _live_messages(turns: Sequence[_LiveTurn]) -> list[dict[str, object]]:
"""Wrap (text, audio) prompt/candidate pairs as the server messages a Live session emits."""
return [{"type": "session.created", "session": {"id": "s"}}] + [
{
"type": "response.done",
"usageMetadata": {
"promptTokenCount": sum(turn["prompt"]),
"candidatesTokenCount": sum(turn["candidates"]),
"totalTokenCount": sum(turn["prompt"]) + sum(turn["candidates"]),
"promptTokensDetails": [
{"modality": "TEXT", "tokenCount": turn["prompt"][0]},
{"modality": "AUDIO", "tokenCount": turn["prompt"][1]},
],
"candidatesTokensDetails": (
[{"modality": "AUDIO"}]
if turn.get("candidate_audio_token_count_missing")
else [
{"modality": "TEXT", "tokenCount": turn["candidates"][0]},
{"modality": "AUDIO", "tokenCount": turn["candidates"][1]},
]
),
},
}
for turn in turns
]
cost = handler._calculate_live_api_cost("gemini-1.5-pro", usage_metadata)
@staticmethod
def _session_usage(
handler: VertexAILivePassthroughLoggingHandler,
mock_logging_obj: MagicMock,
messages: list[dict[str, object]],
model: str,
) -> Usage:
result = handler.vertex_ai_live_passthrough_handler(
websocket_messages=messages,
logging_obj=mock_logging_obj,
url_route="/vertex_ai/live",
start_time=datetime.now(),
end_time=datetime.now(),
request_body={},
model=model,
)
assert result["result"] is not None, "the handler must produce a usage-bearing response to bill"
return result["result"].usage
# Should include web search cost
expected_base_cost = (100 * 0.000001) + (50 * 0.000002)
# The web search cost might be handled differently, so just check it's reasonable
assert cost >= expected_base_cost
assert cost > 0
@classmethod
def _session_cost(
cls,
handler: VertexAILivePassthroughLoggingHandler,
mock_logging_obj: MagicMock,
messages: list[dict[str, object]],
model: str,
) -> float:
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse
usage = cls._session_usage(handler, mock_logging_obj, messages, model)
return completion_cost(
completion_response=ModelResponse(
id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[]
),
model=f"vertex_ai/{model}",
custom_llm_provider="vertex_ai",
call_type="acompletion",
)
@classmethod
def _expected_session_cost(cls, turns: Sequence[_LiveTurn]) -> float:
from litellm.utils import get_model_info
info = get_model_info(model=cls.NATIVE_AUDIO_MODEL, custom_llm_provider="vertex_ai")
return (
sum(turn["prompt"][0] for turn in turns) * info["input_cost_per_token"]
+ sum(turn["prompt"][1] for turn in turns) * info["input_cost_per_audio_token"]
+ sum(turn["candidates"][0] for turn in turns) * info["output_cost_per_token"]
+ sum(turn["candidates"][1] for turn in turns) * info["output_cost_per_audio_token"]
)
def test_every_turn_of_a_session_is_billed(self, handler, mock_logging_obj):
"""Google charges per turn for the whole context window, so every turn adds to the bill.
Billing one snapshot instead gives away all the other turns: on this session the
largest single turn is well under the session total, and its share of the audio is
priced 6x the text rate, so the gap is money rather than rounding.
"""
turns = self.AUDIO_SESSION[:3]
cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
widest_single_turn = max(self._expected_session_cost([turn]) for turn in turns)
assert cost > widest_single_turn, "billing one snapshot drops every other turn of the session"
def test_audio_named_without_a_token_count_bills_at_the_audio_rate(self, handler, mock_logging_obj):
"""Live can name the modality carrying the rest of a turn and omit its tokenCount.
Reading the absent key as zero left those tokens inside candidatesTokenCount but outside
the breakdown, so the calculator charged real speech at the text output rate. At this
entry's rates the last turn's 3 audio tokens are $0.0000360 rather than $0.0000060.
"""
turns = self.AUDIO_SESSION
usage = self._session_usage(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
assert usage.completion_tokens_details.audio_tokens == 100, "the unpriced entry takes the turn's residual"
assert usage.completion_tokens_details.text_tokens == 26
assert usage.completion_tokens == 126
cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
assert cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
TOOL_USE_PER_TURN = (100, 250, 400)
def _grounded_messages(self):
"""The three-turn session again, with each turn's own toolUsePromptTokenCount attached."""
messages = self._live_messages(self.AUDIO_SESSION[:3])
head, turns = messages[0], messages[1:]
return [head] + [
{**message, "usageMetadata": {**message["usageMetadata"], "toolUsePromptTokenCount": tool_use}}
for message, tool_use in zip(turns, self.TOOL_USE_PER_TURN)
]
def test_server_side_tool_use_prompt_tokens_are_summed_over_the_session(self, handler, mock_logging_obj):
"""toolUsePromptTokenCount rode the unknown-key pass-through, so it took the first turn only.
Every other total beside it is summed across the session, and the first turn is the
smallest number in the series, so a grounded session logged far fewer tool-use tokens
than it used. This session's turns are deliberately distinct, so 750 can only come from
summing: first-turn selection gives 100, last-turn or max gives 400.
"""
grounded = self._grounded_messages()
usage = self._session_usage(handler, mock_logging_obj, grounded, self.NATIVE_AUDIO_MODEL)
assert usage.prompt_tokens_details.tool_use_tokens == sum(self.TOOL_USE_PER_TURN)
@staticmethod
def _grounding_frame(metadata: dict[str, object]) -> dict[str, object]:
"""One server frame carrying grounding metadata, the way Live reports it."""
return {"type": "response.done", "serverContent": {"groundingMetadata": metadata}}
def test_web_grounding_is_counted_so_it_can_be_billed(self, handler, mock_logging_obj):
"""Live reports grounding in the server frames and never in usageMetadata.
Nothing read those frames, so web_search_requests stayed unset and the cost path's only
trigger for the per-query grounding charge never fired. Google bills a grounded Live
prompt on top of its tokens, so the whole fee was missing from the bill.
"""
messages = [
self._grounding_frame(
{
"webSearchQueries": ["who won the 2026 world cup final"],
"groundingChunks": [{"web": {"uri": "https://example.com"}}],
}
),
*self._live_messages(self.AUDIO_SESSION[:1]),
]
usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL)
assert usage.prompt_tokens_details.web_search_requests == 1, "a grounded turn must report its query"
assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None
def test_maps_grounding_is_counted_under_its_own_sku(self, handler, mock_logging_obj):
"""Maps grounding is a separate SKU from web search, so it needs its own counter.
A maps-only turn carries grounding chunks but no webSearchQueries, so counting queries
alone would report nothing and bill nothing.
"""
messages = [
self._grounding_frame({"groundingChunks": [{"maps": {"placeId": "abc123"}}]}),
*self._live_messages(self.AUDIO_SESSION[:1]),
]
usage = self._session_usage(handler, mock_logging_obj, messages, self.NATIVE_AUDIO_MODEL)
assert usage.prompt_tokens_details.google_maps_grounding_requests == 1
assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None
def test_an_ungrounded_session_reports_no_grounding(self, handler, mock_logging_obj):
"""The counters must stay absent when no tool ran, or every session pays a grounding fee."""
usage = self._session_usage(
handler, mock_logging_obj, self._live_messages(self.AUDIO_SESSION[:1]), self.NATIVE_AUDIO_MODEL
)
assert getattr(usage.prompt_tokens_details, "web_search_requests", None) is None
assert getattr(usage.prompt_tokens_details, "google_maps_grounding_requests", None) is None
def test_grounding_adds_its_query_fee_to_the_session_bill(self, handler, mock_logging_obj):
"""The counter only matters if it reaches the bill, so assert against the cost, not the field.
Same tokens either way: the difference between the two sessions is the grounding fee alone.
"""
turns = self.AUDIO_SESSION[:1]
plain = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
grounded = self._session_cost(
handler,
mock_logging_obj,
[self._grounding_frame({"webSearchQueries": ["q"]}), *self._live_messages(turns)],
self.NATIVE_AUDIO_MODEL,
)
assert grounded > plain, "a grounded session must cost more than the same tokens ungrounded"
def _priced_logging_obj(self) -> LiteLLMLoggingObj:
"""A real logging object, since the session's price is handed to it turn by turn."""
logging_obj = LiteLLMLoggingObj(
model=self.NATIVE_AUDIO_MODEL,
messages=[],
stream=True,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="live-session",
function_id="live",
)
logging_obj.update_environment_variables(
model=self.NATIVE_AUDIO_MODEL,
user="u",
optional_params={},
litellm_params={},
call_type="pass_through_endpoint",
)
logging_obj.model_call_details["custom_llm_provider"] = "vertex_ai"
return logging_obj
def _billed_session(
self, handler: VertexAILivePassthroughLoggingHandler, messages: list[dict[str, object]]
) -> tuple[float, CostBreakdown]:
logging_obj = self._priced_logging_obj()
result = handler.vertex_ai_live_passthrough_handler(
websocket_messages=messages,
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",
)
assert result["result"] is not None, "the handler must produce a usage-bearing response to bill"
assert logging_obj.cost_breakdown is not None, "the session's price must reach the logging object"
return result["result"]._hidden_params["response_cost"], logging_obj.cost_breakdown
def test_each_grounded_turn_pays_its_own_query_fee(self, handler):
"""Google charges the grounding fee per grounded prompt, not per session.
Summing the session into one usage collapsed two grounded turns into one query, so the
second question was answered for free. The bill now grows by one fee per grounded turn.
"""
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
plain_cost, _ = self._billed_session(handler, [head, turn, turn])
one_cost, one_breakdown = self._billed_session(handler, [head, grounding, turn, turn])
two_cost, two_breakdown = self._billed_session(handler, [head, grounding, turn, grounding, turn])
fee = one_cost - plain_cost
assert fee > 0, "a grounded turn must cost more than the same tokens ungrounded"
assert two_cost - plain_cost == pytest.approx(2 * fee), "two grounded turns must pay the fee twice"
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.
Pricing each turn on its own applied the fixed margin per turn, so a two-turn session paid it
twice. The session now carries the fixed margin once no matter how many turns it billed.
"""
head, turn = self._live_messages(self.AUDIO_SESSION[:1])
grounding = self._grounding_frame({"webSearchQueries": ["q"]})
messages = [head, grounding, turn, grounding, turn]
plain_cost, _ = self._billed_session(handler, messages)
fixed_amount = 0.01
with patch.object(litellm, "cost_margin_config", {"vertex_ai": {"fixed_amount": fixed_amount}}):
margined_cost, breakdown = self._billed_session(handler, messages)
assert margined_cost - plain_cost == pytest.approx(
fixed_amount
), "a two-turn session must add the fixed margin once, not once per billed turn"
assert breakdown["margin_fixed_amount"] == pytest.approx(fixed_amount)
assert breakdown["margin_total_amount"] == pytest.approx(fixed_amount)
def test_reporting_tool_use_tokens_does_not_move_the_bill(self, handler, mock_logging_obj):
"""Deliberate boundary: these tokens are reported here, and priced nowhere.
generic_cost_per_token reads the input bill out of prompt_tokens_details, and falls
back to prompt_tokens only when the details carry no text or a cache hit overlaps them,
so adding tool-use tokens to prompt_tokens is worth nothing on an ordinary Live turn and
over-charges against the cache-overlap correction when it is not. Pricing them belongs
in the shared input-cost path, beside the modality terms that already read the details.
"""
turns = self.AUDIO_SESSION[:3]
plain_cost = self._session_cost(handler, mock_logging_obj, self._live_messages(turns), self.NATIVE_AUDIO_MODEL)
grounded_cost = self._session_cost(
handler, mock_logging_obj, self._grounded_messages(), self.NATIVE_AUDIO_MODEL
)
assert plain_cost == pytest.approx(self._expected_session_cost(turns), rel=1e-9)
assert grounded_cost == pytest.approx(plain_cost, rel=1e-9), "reporting tool use must not move the bill"
def test_a_malformed_details_entry_does_not_cost_the_whole_session(self, handler, mock_logging_obj):
"""A ``*TokensDetails`` value that is not a list of objects must not take the session down.
The handler's only error path returns no result at all, so one odd frame used to throw
while reading it and the whole session billed nothing. The good turns still bill.
"""
turns = self.AUDIO_SESSION[:3]
messages = self._live_messages(turns)
mangled = [dict(message) for message in messages]
mangled[1]["usageMetadata"] = {**mangled[1]["usageMetadata"], "promptTokensDetails": "TEXT"}
usage = self._session_usage(handler, mock_logging_obj, mangled, self.NATIVE_AUDIO_MODEL)
surviving = turns[1:]
assert usage.prompt_tokens_details.audio_tokens == sum(turn["prompt"][1] for turn in surviving)
assert usage.prompt_tokens_details.text_tokens == sum(turn["prompt"][0] for turn in surviving)
assert usage.prompt_tokens == sum(sum(turn["prompt"]) for turn in turns), "the totals still cover every turn"
direct = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": 40,
"candidatesTokenCount": 12,
"promptTokensDetails": [{"modality": "AUDIO", "tokenCount": 40}, "AUDIO"],
"candidatesTokensDetails": {"modality": "TEXT", "tokenCount": 12},
},
model=self.NATIVE_AUDIO_MODEL,
)
assert direct.prompt_tokens_details.audio_tokens == 40, "the well-formed entry beside a bad one still counts"
assert direct.completion_tokens == 12
@pytest.mark.parametrize(
"label,prompt_details,candidate_details",
[
("text only", [("TEXT", 6)], [("TEXT", 2)]),
("audio in", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 18)]),
("image in", [("TEXT", 10), ("IMAGE", 258)], [("TEXT", 24)]),
("frames in", [("TEXT", 11), ("IMAGE", 1032)], [("TEXT", 26)]),
("audio both ways", [("TEXT", 13), ("AUDIO", 127)], [("TEXT", 29), ("AUDIO", 95)]),
],
)
def test_live_session_bills_each_modality_at_its_own_rate(self, handler, label, prompt_details, candidate_details):
"""Every payload here is a real Vertex Live session's usageMetadata.
Before the fix these billed the text share only, from 1x (text) to 55x under.
The expected amount is derived from the entry's own rates rather than hardcoded,
so this stays correct as prices move, and it is asserted exactly, so dropping a
modality and double-charging one both fail.
"""
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse
from litellm.utils import get_model_info
model = self.NATIVE_AUDIO_MODEL
info = get_model_info(model=model, custom_llm_provider="vertex_ai")
text_in = info["input_cost_per_token"]
audio_in = info.get("input_cost_per_audio_token") or text_in
image_in = info.get("input_cost_per_image_token") or text_in
text_out = info["output_cost_per_token"]
audio_out = info.get("output_cost_per_audio_token") or text_out
rate_in = {"TEXT": text_in, "AUDIO": audio_in, "IMAGE": image_in}
rate_out = {"TEXT": text_out, "AUDIO": audio_out}
expected = sum(c * rate_in[m] for m, c in prompt_details) + sum(c * rate_out[m] for m, c in candidate_details)
usage = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": sum(c for _, c in prompt_details),
"candidatesTokenCount": sum(c for _, c in candidate_details),
"promptTokensDetails": [{"modality": m, "tokenCount": c} for m, c in prompt_details],
"candidatesTokensDetails": [{"modality": m, "tokenCount": c} for m, c in candidate_details],
},
model=model,
)
cost = completion_cost(
completion_response=ModelResponse(
id="x", object="chat.completion", created=0, model=model, usage=usage, choices=[]
),
model=f"vertex_ai/{model}",
custom_llm_provider="vertex_ai",
call_type="acompletion",
)
assert cost == pytest.approx(expected, rel=1e-9), label
text_only = sum(c for m, c in prompt_details if m == "TEXT") * text_in + sum(
c for m, c in candidate_details if m == "TEXT"
) * text_out
if any(m != "TEXT" for m, _ in prompt_details + candidate_details) and audio_in != text_in:
assert cost > text_only, f"{label}: non-text modalities must add cost"
def test_vertex_ai_live_passthrough_handler_integration(
self, handler, mock_logging_obj, sample_websocket_messages
@ -376,6 +788,7 @@ class TestVertexAILivePassthroughIntegration:
"""Create a mock logging object"""
mock = MagicMock(spec=LiteLLMLoggingObj)
mock.model_call_details = {}
mock._response_cost_calculator.return_value = None
return mock
@patch(
@ -509,6 +922,7 @@ class TestVertexAILivePassthroughErrorHandling:
"""Create a mock logging object"""
mock = MagicMock(spec=LiteLLMLoggingObj)
mock.model_call_details = {}
mock._response_cost_calculator.return_value = None
return mock
def test_invalid_websocket_messages_format(self):
@ -540,25 +954,24 @@ class TestVertexAILivePassthroughErrorHandling:
result = handler._extract_usage_metadata_from_websocket_messages(messages)
assert result is None
@patch(
"litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_ai_live_passthrough_logging_handler.get_model_info"
)
def test_cost_calculation_with_missing_model_info(self, mock_get_model_info):
"""Test cost calculation when model info is missing"""
def test_usage_without_modality_details(self):
"""Older payloads carry only the totals; fall back to them rather than reporting zero."""
handler = VertexAILivePassthroughLoggingHandler()
# Mock missing model info
mock_get_model_info.return_value = {}
usage = handler._create_usage_object_from_metadata(
usage_metadata={
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
},
model="unknown-model",
)
usage_metadata = {
"promptTokenCount": 100,
"candidatesTokenCount": 50,
"totalTokenCount": 150,
}
# Should not raise an exception, should return 0 or handle gracefully
cost = handler._calculate_live_api_cost("unknown-model", usage_metadata)
assert cost == 0.0
assert usage.prompt_tokens == 100
assert usage.completion_tokens == 50
assert usage.total_tokens == 150
assert usage.prompt_tokens_details.audio_tokens is None
assert usage.prompt_tokens_details.image_tokens is None
def test_handler_with_none_websocket_messages(self, mock_logging_obj):
"""Test handler with None websocket messages"""

View file

@ -1,4 +1,6 @@
import json
from collections.abc import Mapping
from typing import cast
from unittest.mock import MagicMock
import pytest
@ -6,6 +8,7 @@ import pytest
import litellm
from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig
from litellm.types.llms.gemini import BidiGenerateContentServerMessage
def test_gemini_realtime_transformation_session_created():
@ -2178,3 +2181,71 @@ def test_unbilled_usage_on_session_close_flushes_trailing_audio(patch_gemini_tra
}
assert usage == expected
assert config.unbilled_usage_on_session_close("gemini-3.5-transcribe-live") is None
def _grounded_live_frame(grounding_metadata: Mapping[str, object] | None) -> Mapping[str, object]:
"""One Live server frame. Grounding metadata and usageMetadata arrive together, as Vertex sends them."""
from typing import Final
server_content: Final = {
"turnComplete": True,
**({} if grounding_metadata is None else {"groundingMetadata": grounding_metadata}),
}
return {
"serverContent": server_content,
"usageMetadata": {
"promptTokenCount": 19,
"candidatesTokenCount": 157,
"totalTokenCount": 176,
"promptTokensDetails": ({"modality": "TEXT", "tokenCount": 19},),
"candidatesTokensDetails": ({"modality": "AUDIO", "tokenCount": 157},),
},
}
def _response_done_input_details(message: Mapping[str, object]) -> Mapping[str, object]:
"""The ``input_tokens_details`` a ``response.done`` event carries, read off the emitted event."""
from typing import Final
config: Final = GeminiRealtimeConfig()
event: Final = config.transform_response_done_event(
message=cast( # cast-ok: a test fixture stands in for the server frame TypedDict
BidiGenerateContentServerMessage, message
),
current_response_id="resp_grounding",
current_conversation_id="conv_grounding",
output_items=None,
)
usage: Final = event["response"]["usage"]
assert usage, "response.done must carry a usage object"
return usage.get("input_tokens_details") or {}
def test_gemini_realtime_response_done_counts_web_grounding():
"""Regression: Live reports grounding in the server frames and never in usageMetadata.
Nothing read those frames on the realtime path, so web_search_requests stayed unset and the
cost path's only trigger for Google's per-query grounding charge never fired.
The counter is read off the emitted event, which is what the cost path is handed, so this covers
the grounding read and the usage bridge that carries it together
"""
input_details = _response_done_input_details(
_grounded_live_frame(
{
"webSearchQueries": ["who won the 2026 world cup final"],
"groundingChunks": [{"web": {"uri": "https://example.com"}}],
}
)
)
assert input_details.get("web_search_requests") == 1, "a grounded turn must report its query"
assert input_details.get("text_tokens") == 19, "the modality breakdown must survive alongside it"
def test_gemini_realtime_response_done_reports_no_grounding_when_none_ran():
"""The counter must stay unset on an ordinary turn, or every session pays a grounding fee."""
input_details = _response_done_input_details(_grounded_live_frame(None))
assert input_details.get("web_search_requests") is None
assert input_details.get("google_maps_grounding_requests") is None

View file

@ -5030,6 +5030,100 @@ async def test_websocket_passthrough_rewrites_gateway_alias_setup_model():
assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash"
@pytest.mark.parametrize(
"setup_model",
["gemini-live-2.5-flash", "models/gemini-live-2.5-flash", "publishers/google/models/gemini-live-2.5-flash"],
)
def test_vertex_live_setup_model_resolves_before_extraction(setup_model):
"""A bare gateway alias left the session logged as ``unknown`` at zero cost.
The model was read off the raw client frame, and the extractor only yields a name when the string
already contains ``/models/``. The rewriter qualifies it a few lines later for the upstream, so a
client that addressed the gateway the documented way, by alias, logged no model and therefore
resolved no cost-map entry. Resolving first is what puts the real name on the logging object.
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_build_vertex_live_setup_model_rewriter,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_extract_model_from_vertex_ai_setup,
_resolved_vertex_live_setup,
)
rewriter = _build_vertex_live_setup_model_rewriter(
vertex_project="proj-db", vertex_location="global", llm_router=None
)
setup_data = {"model": setup_model}
resolved = _extract_model_from_vertex_ai_setup(_resolved_vertex_live_setup(setup_data, rewriter))
assert resolved == "gemini-live-2.5-flash", "an unresolved setup model logs the session as 'unknown'"
@pytest.mark.asyncio
async def test_websocket_passthrough_logs_a_bare_alias_setup_model():
"""End to end through the relay: a bare alias must reach the logging object as a real model name.
This is the call-site half of the fix. The helper tests above pass even if extraction moves back
before the rewrite, so this one drives the real websocket relay and asserts on what got logged,
which is the name the cost map is looked up by. An unbilled session logs ``unknown``.
"""
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
_build_vertex_live_setup_model_rewriter,
)
upstream_ws = RecordingUpstreamWebSocket()
setup_frame = json.dumps({"setup": {"model": "gemini-live-2.5-flash"}})
websocket = _client_websocket(
AsyncMock(
side_effect=[
{"type": "websocket.receive", "text": setup_frame},
{"type": "websocket.disconnect"},
]
)
)
built = []
real_logging = litellm.litellm_core_utils.litellm_logging.Logging
def _capture(*args, **kwargs):
obj = real_logging(*args, **kwargs)
built.append(obj)
return obj
with _patched_websocket_passthrough_environment(upstream_ws):
with patch("litellm.litellm_core_utils.litellm_logging.Logging", side_effect=_capture):
await websocket_passthrough_request(
websocket=websocket,
target="wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent",
custom_headers={"Authorization": "Bearer token"},
user_api_key_dict=UserAPIKeyAuth(),
forward_headers=False,
endpoint="/vertex_ai/live",
accept_websocket=False,
setup_model_rewriter=_build_vertex_live_setup_model_rewriter(
vertex_project="proj-db", vertex_location="global", llm_router=None
),
)
assert built, "the relay should have built a logging object"
assert built[0].model == "gemini-live-2.5-flash", "a bare alias must not log as 'unknown'"
def test_vertex_live_setup_resolution_is_inert_without_a_rewriter():
"""Non-Live passthrough routes pass no rewriter, so the frame must be handed over untouched."""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_extract_model_from_vertex_ai_setup,
_resolved_vertex_live_setup,
)
setup_data = {"model": "projects/p/locations/global/publishers/google/models/gemini-live-2.5-flash"}
assert _resolved_vertex_live_setup(setup_data, None) is setup_data
assert _extract_model_from_vertex_ai_setup(_resolved_vertex_live_setup(setup_data, None)) == (
"gemini-live-2.5-flash"
)
@pytest.mark.asyncio
@pytest.mark.parametrize("rcvd_close", [None, "abnormal", "no_status"])
async def test_websocket_passthrough_does_not_relay_unsendable_upstream_close(rcvd_close):