From 8b47acaef0591fe6554990f41811eaaa619eedc9 Mon Sep 17 00:00:00 2001 From: Marty Sullivan Date: Mon, 7 Sep 2026 00:33:43 -0400 Subject: [PATCH] fix(proxy): enforce model access on the Vertex AI Live websocket passthrough /vertex_ai/live authenticates the key at connect time, but the model a session actually runs is named in the client's setup frame, which arrives afterwards. ?model= is optional on the route, and when it is absent the synthesised auth body carries an empty model, so get_model_from_request returns None and the key/team model allowlists are never applied. Any valid virtual key could then name an arbitrary publisher model and have the gateway mint its own service-account token and forward it. Authorize each client frame that names a model before it reaches the upstream, through can_key_call_resolved_model, the same check the routed /v1/realtime path runs. A refused frame is not forwarded: the client socket closes 1008 with the denial, and the upstream socket closes with it. The denial travels back out as a ProxyException that keeps the underlying check's error type and status, so a refused session reaches post_call_failure_hook, the same way this handler already reports an upstream rejection. Returning normally left the unconditional success hooks running, which recorded a denied session as a successful one. Closing the upstream moved to the caller for the same reason: awaiting it inside the frame relay lets the upstream reader finish first and cancel the refusing task mid-close, dropping the denial The route's own coverage drives the real relay rather than a stand-in for it, so a setup frame naming a model the key cannot call is proven never to reach the upstream socket and the denial is read back off post_call_failure_hook. That checks what the session does instead of which arguments the route passed, which needs no module global set and no collaborator swapped out The model is authorized as the router model group the frame names, whichever addressing carried it. The Live SDK wraps whatever the caller typed as models/, callers send the full projects/

/locations//publishers/google/models/ path because a bare name connects but bills nothing, and vertex_ai/ is how the same model is named everywhere else in LiteLLM, so comparing the raw string refuses every addressing but one. The allowlists are keyed by group, and so is get_model_access_groups, which returns nothing at all for a prefixed name, so a key entitled through an access group needs the same normalization. A model that resolves to no group is authorized as it arrived rather than stripped down to its last segment. Stripping would make this a bypass: a wildcard key would reach an arbitrary unserved publisher model through any prefix, and a full resource path would carry a project of the caller's choosing with it. One deployment lookup backs both the gate and the rewriter, which is where the prefixed forms were already understood --- .../llm_passthrough_endpoints.py | 126 ++++- .../pass_through_endpoints.py | 83 +++- .../test_pass_through_endpoints.py | 466 +++++++++++++++++- 3 files changed, 654 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 2ea46b740a8..a3575b17b8c 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -37,7 +37,7 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client from litellm.llms.vertex_ai.vertex_llm_base import VertexBase from litellm.passthrough.main import AsyncPassthroughStreamingResponse from litellm.proxy._types import * -from litellm.proxy.auth.auth_checks import enforced_model_allowlists +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model, enforced_model_allowlists from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( @@ -59,6 +59,7 @@ from litellm.proxy.common_utils.sse_keepalive import ( from litellm.proxy.pass_through_endpoints.common_utils import get_litellm_virtual_key from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( HttpPassThroughEndpointHelpers, + WebsocketClientFrameGate, create_pass_through_route, create_websocket_passthrough_route, websocket_passthrough_request, @@ -86,6 +87,7 @@ from .passthrough_endpoint_router import PassthroughEndpointRouter if TYPE_CHECKING: from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig from litellm.router import Router + from litellm.types.router import DeploymentTypedDict ProxyConfig = _ProxyConfig # rebind-ok: conditional type alias else: @@ -2693,6 +2695,10 @@ VERTEX_LIVE_UNCONFIGURED_CLOSE_REASON: Final = ( "Vertex AI auth failed: set a use_in_pass_through vertex model, default_vertex_config, or DEFAULT_VERTEXAI_* env" ) +VERTEX_LIVE_UNVERIFIED_MODEL_CLOSE_REASON: Final = ( + "Could not check your key's access to the model this setup frame names; ask a proxy admin to read the logs" +) + VERTEX_PUBLISHER_MODEL_PREFIX: Final = "publishers/google/models/" VERTEX_PUBLISHERS_SEGMENT: Final = "publishers/" @@ -2717,6 +2723,107 @@ def _get_llm_router() -> Router | None: return llm_router +def _get_llm_model_list() -> list | None: # mutable-ok: mirrors proxy_server.llm_model_list's own type + from litellm.proxy.proxy_server import llm_model_list + + return llm_model_list + + +def _may_carry_vertex_live_setup(frame_data: str | bytes) -> bool: + """ + Whether a frame is worth parsing: audio and video frames stream continuously and are large, and one that + cannot spell ``setup`` cannot carry one. A backslash means the key may be escaped, which only a parse settles + """ + if isinstance(frame_data, bytes): + return b'"setup"' in frame_data or b"\\" in frame_data + return '"setup"' in frame_data or "\\" in frame_data + + +def _vertex_live_setup_model(frame_data: str | bytes) -> str | None: + """ + The model a Vertex AI Live ``setup`` frame names, which is the only place BidiGenerateContent carries one + """ + if not _may_carry_vertex_live_setup(frame_data): + return None + try: + frame: Final = json.loads(frame_data) + except (json.JSONDecodeError, UnicodeDecodeError): + return None + setup: Final = frame.get("setup") if isinstance(frame, dict) else None + model: Final = setup.get("model") if isinstance(setup, dict) else None + return model if isinstance(model, str) and model else None + + +def _vertex_live_deployment(setup_model: str, llm_router: Router | None) -> DeploymentTypedDict | None: + """ + The router deployment whose group a ``setup`` frame's model names, whichever addressing the client used. + + The Live SDK wraps whatever the caller typed as ``models/``, and clients also send LiteLLM ids and + full Vertex resource paths, so the group is read from the last segment as well as from the whole string + """ + if llm_router is None: + return None + candidates: Final = (setup_model, setup_model.rsplit("/", 1)[-1]) + return next( + ( + deployment + for deployment in (llm_router.get_model_list() or ()) + if deployment.get("model_name") in candidates + ), + None, + ) + + +def _build_vertex_live_client_frame_gate( + llm_model_list: list | None, # mutable-ok: matches can_key_call_resolved_model's own type + llm_router: Router | None, +) -> WebsocketClientFrameGate: + """ + Apply the key's, team's and project's model access to the model a client's ``setup`` frame names. + + ``?model=`` is optional on this route and the Live protocol carries the real model in the first client + frame, so connect-time auth has no model to check and every frame that names one has to be authorized + here instead. + + The allowlists and the access groups behind them are keyed by model group, and every addressing a client + sends carries that group inside a prefix, so the group is what gets authorized. A frame naming no group + is still checked as it arrived, which is the strictest reading available for a name nothing recognises + """ + + async def gate(frame_data: str | bytes, valid_token: UserAPIKeyAuth, /) -> ProxyException | None: + model: Final = _vertex_live_setup_model(frame_data) + if model is None: + return None + deployment: Final = _vertex_live_deployment(model, llm_router) + try: + await can_key_call_resolved_model( + model=model if deployment is None else deployment["model_name"], + llm_model_list=llm_model_list, + valid_token=valid_token, + llm_router=llm_router, + ) + except ProxyException as denial: + # A close frame carries 123 bytes, so the model leads: the tail of the allowlist dump is what a + # caller can most afford to lose + return ProxyException( + message=f"{model}: {denial.message}", + type=denial.type, + param=denial.param, + code=denial.code, + ) + except Exception: # noqa: BLE001 # fail closed: a gate that cannot answer refuses, and says so in band + verbose_proxy_logger.exception("Vertex AI Live passthrough: model access check failed for %s", model) + return ProxyException( + message=VERTEX_LIVE_UNVERIFIED_MODEL_CLOSE_REASON, + type=ProxyErrorTypes.internal_server_error, + param="model", + code=500, + ) + return None + + return gate + + def _resolve_vertex_live_credentials( vertex_project: str | None, vertex_location: str | None, @@ -2775,17 +2882,8 @@ def _resolve_alias_to_upstream_model(setup_model: str, llm_router: Router | None """ The Live SDK wraps whatever the caller typed as ``models/``, so a gateway alias arrives prefixed """ - if llm_router is None: - return setup_model - candidates: Final = (setup_model, setup_model.rsplit("/", 1)[-1]) - upstream: Final = next( - ( - deployment["litellm_params"].get("model") - for deployment in (llm_router.get_model_list() or ()) - if deployment.get("model_name") in candidates - ), - None, - ) + deployment: Final = _vertex_live_deployment(setup_model, llm_router) + upstream: Final = None if deployment is None else deployment["litellm_params"].get("model") if upstream is None: return setup_model try: @@ -2899,6 +2997,10 @@ async def vertex_ai_live_websocket_passthrough( vertex_location=resolved_location, llm_router=_get_llm_router(), ), + client_frame_gate=_build_vertex_live_client_frame_gate( + llm_model_list=_get_llm_model_list(), + llm_router=_get_llm_router(), + ), ) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index ea4ede7e513..826d0641dfd 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -9,7 +9,7 @@ from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequenc from dataclasses import dataclass from datetime import datetime from itertools import groupby -from typing import TYPE_CHECKING, Any, Final, TypedDict, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, TypedDict, cast from urllib.parse import urlencode, urlparse import httpx @@ -27,7 +27,7 @@ from fastapi import ( from fastapi.responses import StreamingResponse from starlette.datastructures import UploadFile as StarletteUploadFile from starlette.websockets import WebSocketState -from websockets.asyncio.client import connect +from websockets.asyncio.client import ClientConnection, connect from websockets.exceptions import ( ConnectionClosedError, ConnectionClosedOK, @@ -2132,6 +2132,39 @@ def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None: return upstream_close +class WebsocketClientFrameGate(Protocol): + """Authorize one client frame, returning the denial to report or ``None`` to let it through. + + A websocket passthrough authenticates before any frame arrives, so a protocol that names its model + inside a frame has no model to authorize at connect time. The frame, not the extracted model, is the + argument because where a model sits is part of each provider's wire protocol + """ + + async def __call__(self, frame_data: str | bytes, valid_token: UserAPIKeyAuth, /) -> ProxyException | None: ... + + +async def _relay_client_frame( + websocket: WebSocket, + upstream_ws: ClientConnection, + checked_frame: str | bytes, + forwarded_frame: str | bytes, + gate: WebsocketClientFrameGate | None, + valid_token: UserAPIKeyAuth, +) -> ProxyException | None: + """Forward one client frame upstream, or refuse it and hand back the denial that ended the session. + + The denial travels out so the caller reports a failed session instead of leaving the success hooks to + record a refusal as a success, and tearing the upstream down is the caller's job: awaiting that close + here lets the upstream reader finish first, which cancels this task mid-close and loses the denial + """ + refusal: Final = None if gate is None else await gate(checked_frame, valid_token) + if refusal is None: + await upstream_ws.send(forwarded_frame) + return None + await websocket.close(code=CloseCode.POLICY_VIOLATION, reason=_truncated_close_reason(refusal.message)) + return refusal + + async def websocket_passthrough_request( websocket: WebSocket, target: str, @@ -2142,6 +2175,7 @@ async def websocket_passthrough_request( cost_per_request: float | None = None, accept_websocket: bool = True, setup_model_rewriter: Callable[[str], str] | None = None, + client_frame_gate: WebsocketClientFrameGate | None = None, ): """ WebSocket passthrough request handler. @@ -2155,6 +2189,7 @@ async def websocket_passthrough_request( endpoint: The endpoint path (for logging purposes) cost_per_request: Optional field - cost per request to the target endpoint setup_model_rewriter: Optional rewrite of the setup frame's model before it reaches the upstream + client_frame_gate: Optional per-frame authorization, applied before a frame reaches the upstream """ from litellm.litellm_core_utils.litellm_logging import Logging from litellm.proxy.proxy_server import proxy_config, proxy_logging_obj @@ -2284,8 +2319,8 @@ async def websocket_passthrough_request( "WebSocket passthrough (%s): Upstream connection established successfully", endpoint ) - async def forward_client_to_upstream() -> None: - """Forward messages from client to upstream WebSocket""" + async def forward_client_to_upstream() -> ProxyException | None: + """Forward messages from client to upstream WebSocket, returning any denial that ended it""" try: while True: message = await websocket.receive() @@ -2352,9 +2387,27 @@ async def websocket_passthrough_request( ) # Not a JSON message or doesn't contain setup data - await upstream_ws.send(_rewrite_vertex_live_setup_model(text_data, setup_model_rewriter)) + refusal = await _relay_client_frame( + websocket, + upstream_ws, + text_data, + _rewrite_vertex_live_setup_model(text_data, setup_model_rewriter), + client_frame_gate, + user_api_key_dict, + ) + if refusal is not None: + return refusal elif bytes_data is not None: - await upstream_ws.send(bytes_data) + refusal = await _relay_client_frame( + websocket, + upstream_ws, + bytes_data, + bytes_data, + client_frame_gate, + user_api_key_dict, + ) + if refusal is not None: + return refusal except asyncio.CancelledError: raise except Exception: @@ -2456,13 +2509,29 @@ async def websocket_passthrough_request( except asyncio.CancelledError: pass + task_results: Final = tuple(task.result() for task in done if task.exception() is None) + + # A refusal outranks any exception it caused: closing the client socket is what makes the + # upstream reader's next send fail, so raising that instead would report the wrong cause + client_frame_refusal: Final = next( + (result for result in task_results if isinstance(result, ProxyException)), None + ) + if client_frame_refusal is not None: + await upstream_ws.close() + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=client_frame_refusal, + request_data={**kwargs, "litellm_logging_obj": logging_obj}, # mutable-ok: the hook pops keys + ) + return + # Check for exceptions in completed tasks for task in done: exception = task.exception() if exception is not None: raise exception - upstream_close: Final = _upstream_close_to_relay(task.result() for task in done) + upstream_close: Final = _upstream_close_to_relay(task_results) if upstream_close is not None and _client_socket_is_open(websocket): await websocket.close( code=upstream_close.code, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index d57bed430c1..1780e468154 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -32,7 +32,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( ) from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, ) @@ -4869,6 +4869,39 @@ class RecordingUpstreamWebSocket: raise StopAsyncIteration +class ClosableUpstreamWebSocket: + """An upstream whose ``recv`` ends when ``close`` does, and whose close keeps yielding after that. + + A real close handshake outlives the ``recv`` it unblocks, so the reader task finishes while the closing + task is still suspended. That is the interleaving that leaves a refusal stranded in a cancelled task + """ + + CLOSE_HANDSHAKE_YIELDS = 5 + + def __init__(self): + self._closing = asyncio.Event() + self.send = AsyncMock() + self.close = AsyncMock(side_effect=self._close) + + async def _close(self): + self._closing.set() + for _ in range(self.CLOSE_HANDSHAKE_YIELDS): + await asyncio.sleep(0) + + async def recv(self, decode: bool = True): + from websockets.exceptions import ConnectionClosedOK + from websockets.frames import Close + + await self._closing.wait() + raise ConnectionClosedOK(rcvd=Close(1000, ""), sent=Close(1000, ""), rcvd_then_sent=True) + + def __aiter__(self): + return self + + async def __anext__(self): + raise StopAsyncIteration + + def _client_websocket(receive): from starlette.websockets import WebSocketState @@ -4904,7 +4937,7 @@ def _patched_websocket_passthrough_environment(upstream_ws): mock_worker.ensure_initialized_and_enqueue = MagicMock( side_effect=lambda async_coroutine: async_coroutine.close() ) - yield + yield SimpleNamespace(proxy_logging=mock_proxy_logging, logging_worker=mock_worker) async def _pending_receive(): @@ -5030,6 +5063,435 @@ 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" +VERTEX_LIVE_TARGET = "wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent" + + +VERTEX_LIVE_GROUP = "gemini-live-2.5-flash-native-audio" + +VERTEX_LIVE_ALIAS_GROUP = "gemini-live-native-audio" + +VERTEX_LIVE_CLIENT_ADDRESSINGS = ( + "{model}", + "models/{model}", + "vertex_ai/{model}", + "publishers/google/models/{model}", + "projects/proj-db/locations/global/publishers/google/models/{model}", +) + + +def _vertex_live_router(): + return litellm.Router( + model_list=[ + { + "model_name": "gemini-live-2.5-flash", + "litellm_params": {"model": "vertex_ai/gemini-live-2.5-flash"}, + }, + { + "model_name": VERTEX_LIVE_GROUP, + "litellm_params": {"model": f"vertex_ai/{VERTEX_LIVE_GROUP}"}, + }, + { + "model_name": VERTEX_LIVE_ALIAS_GROUP, + "litellm_params": {"model": f"vertex_ai/{VERTEX_LIVE_GROUP}"}, + }, + { + "model_name": "gemini-3-pro-live", + "litellm_params": {"model": "vertex_ai/gemini-3-pro-live"}, + }, + ] + ) + + +async def _run_vertex_live_gated_passthrough(client_frames, valid_token): + """Drive /vertex_ai/live with the real frame gate, returning the client and upstream sockets""" + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _build_vertex_live_client_frame_gate, + _build_vertex_live_setup_model_rewriter, + ) + + llm_router = _vertex_live_router() + upstream_ws = RecordingUpstreamWebSocket() + websocket = _client_websocket( + AsyncMock(side_effect=[*client_frames, {"type": "websocket.disconnect"}]), + ) + + with _patched_websocket_passthrough_environment(upstream_ws): + await websocket_passthrough_request( + websocket=websocket, + target=VERTEX_LIVE_TARGET, + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=valid_token, + 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=llm_router, + ), + client_frame_gate=_build_vertex_live_client_frame_gate( + llm_model_list=llm_router.get_model_list(), + llm_router=llm_router, + ), + ) + + return websocket, upstream_ws + + +def _setup_frame(model: str) -> dict: + return { + "type": "websocket.receive", + "text": json.dumps({"setup": {"model": model, "generationConfig": {"responseModalities": ["TEXT"]}}}), + } + + +def _policy_close_reason(websocket) -> str | None: + return next( + (call.kwargs.get("reason") for call in websocket.close.await_args_list if call.kwargs.get("code") == 1008), + None, + ) + + +@pytest.mark.asyncio +async def test_vertex_live_setup_frame_naming_a_model_the_key_cannot_call_is_refused(): + """ + ?model= is optional on this route, so a key that names an unlisted model only in the setup frame used to + reach Vertex with the gateway's own service-account credentials + """ + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame("totally-unlisted-model-abc")], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + upstream_ws.send.assert_not_awaited() + reason = _policy_close_reason(websocket) + assert reason is not None and "totally-unlisted-model-abc" in reason + upstream_ws.close.assert_awaited() + + +@pytest.mark.asyncio +async def test_vertex_live_setup_frame_naming_a_permitted_model_still_reaches_vertex(): + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame("gemini-live-2.5-flash")], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + upstream_ws.send.assert_awaited_once() + sent_setup = json.loads(upstream_ws.send.await_args.args[0])["setup"] + assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash" + assert _policy_close_reason(websocket) is None + + +@pytest.mark.asyncio +async def test_vertex_live_setup_frame_from_an_unrestricted_key_still_reaches_vertex(): + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame("gemini-3-pro-live")], + UserAPIKeyAuth(token="hashed", models=[]), + ) + + upstream_ws.send.assert_awaited_once() + assert _policy_close_reason(websocket) is None + + +@pytest.mark.asyncio +async def test_vertex_live_setup_frame_is_checked_against_the_team_for_an_all_team_models_key(): + """ + ``all-team-models`` skips the key allowlist, so the team's own restriction is the only thing left to apply + """ + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame("gemini-3-pro-live")], + UserAPIKeyAuth( + token="hashed", + models=["all-team-models"], + team_id="team-1", + team_models=["gemini-live-2.5-flash"], + ), + ) + + upstream_ws.send.assert_not_awaited() + reason = _policy_close_reason(websocket) + assert reason is not None and "team not allowed to access model" in reason + + +@pytest.mark.asyncio +async def test_vertex_live_setup_frame_is_gated_when_it_arrives_as_binary(): + """A client that sends its setup frame as bytes must not skip the check the text path applies""" + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [ + { + "type": "websocket.receive", + "bytes": json.dumps({"setup": {"model": "totally-unlisted-model-abc"}}).encode(), + } + ], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + upstream_ws.send.assert_not_awaited() + assert _policy_close_reason(websocket) is not None + + +@pytest.mark.asyncio +async def test_vertex_live_setup_frame_with_an_escaped_key_is_still_gated(): + """Large audio frames skip the parse on a substring test, so a JSON-escaped key must not slip past it""" + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [{"type": "websocket.receive", "text": '{"\\u0073etup": {"model": "totally-unlisted-model-abc"}}'}], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + upstream_ws.send.assert_not_awaited() + assert _policy_close_reason(websocket) is not None + + +@pytest.mark.asyncio +async def test_vertex_live_non_setup_frames_are_forwarded_without_a_model_check(): + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [ + _setup_frame("gemini-live-2.5-flash"), + {"type": "websocket.receive", "text": json.dumps({"realtimeInput": {"audio": {"data": "AAAA"}}})}, + ], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + assert upstream_ws.send.await_count == 2 + assert _policy_close_reason(websocket) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("addressing", VERTEX_LIVE_CLIENT_ADDRESSINGS) +@pytest.mark.parametrize("group", [VERTEX_LIVE_GROUP, VERTEX_LIVE_ALIAS_GROUP]) +async def test_vertex_live_setup_frame_authorizes_a_permitted_group_in_every_addressing(addressing, group): + """ + The Live SDK wraps the caller's model as ``models/`` and the docs tell callers to send the full + resource path, so a key holding the group has to work whichever addressing its client actually sends + """ + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame(addressing.format(model=group))], + UserAPIKeyAuth(token="hashed", models=[group]), + ) + + upstream_ws.send.assert_awaited_once() + sent_setup = json.loads(upstream_ws.send.await_args.args[0])["setup"] + assert sent_setup["model"].startswith("projects/") + assert _policy_close_reason(websocket) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("addressing", VERTEX_LIVE_CLIENT_ADDRESSINGS) +async def test_vertex_live_setup_frame_refuses_an_unpermitted_group_in_every_addressing(addressing): + """Normalizing the addressing must not become a way to reach a group the key does not hold""" + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame(addressing.format(model="gemini-3-pro-live"))], + UserAPIKeyAuth(token="hashed", models=[VERTEX_LIVE_GROUP]), + ) + + upstream_ws.send.assert_not_awaited() + reason = _policy_close_reason(websocket) + assert reason is not None and "gemini-3-pro-live" in reason + + +@pytest.mark.asyncio +@pytest.mark.parametrize("addressing", VERTEX_LIVE_CLIENT_ADDRESSINGS) +async def test_vertex_live_setup_frame_naming_no_group_at_all_is_refused_in_every_addressing(addressing): + """A name that resolves to no group is authorized as it arrived, so an unknown model cannot ride a prefix in""" + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame(addressing.format(model="totally-unlisted-model-abc"))], + UserAPIKeyAuth(token="hashed", models=[VERTEX_LIVE_GROUP]), + ) + + upstream_ws.send.assert_not_awaited() + assert _policy_close_reason(websocket) is not None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("addressing", VERTEX_LIVE_CLIENT_ADDRESSINGS[1:]) +async def test_vertex_live_setup_frame_does_not_strip_a_prefix_off_a_model_no_group_serves(addressing): + """ + Stripping the addressing off a name the router does not serve would let a wildcard key reach an arbitrary + publisher model, and a full path would carry a project of the caller's choosing with it + """ + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame(addressing.format(model="gemini-4-pro-live-unserved"))], + UserAPIKeyAuth(token="hashed", models=["gemini-*"]), + ) + + upstream_ws.send.assert_not_awaited() + assert _policy_close_reason(websocket) is not None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("addressing", VERTEX_LIVE_CLIENT_ADDRESSINGS) +async def test_vertex_live_setup_frame_is_checked_against_the_team_in_every_addressing(addressing): + """The team allowlist is keyed by group too, so the addressing must not decide whether the team is checked""" + websocket, upstream_ws = await _run_vertex_live_gated_passthrough( + [_setup_frame(addressing.format(model="gemini-3-pro-live"))], + UserAPIKeyAuth( + token="hashed", + models=["all-team-models"], + team_id="team-1", + team_models=[VERTEX_LIVE_GROUP], + ), + ) + + upstream_ws.send.assert_not_awaited() + reason = _policy_close_reason(websocket) + assert reason is not None and "team not allowed to access model" in reason + + +@pytest.mark.asyncio +@pytest.mark.parametrize("addressing", VERTEX_LIVE_CLIENT_ADDRESSINGS) +async def test_vertex_live_setup_frame_resolves_an_access_group_grant_in_every_addressing(addressing): + """ + ``get_model_access_groups`` is keyed by group as well, and returns nothing for a prefixed name, so a key + entitled through an access group rather than a model name needs the same normalization + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _build_vertex_live_client_frame_gate, + ) + + llm_router = litellm.Router( + model_list=[ + { + "model_name": VERTEX_LIVE_GROUP, + "litellm_params": {"model": f"vertex_ai/{VERTEX_LIVE_GROUP}"}, + "model_info": {"access_groups": ["live-team"]}, + }, + ] + ) + gate = _build_vertex_live_client_frame_gate( + llm_model_list=llm_router.get_model_list(), + llm_router=llm_router, + ) + + frame = json.dumps({"setup": {"model": addressing.format(model=VERTEX_LIVE_GROUP)}}) + denial = await gate(frame, UserAPIKeyAuth(token="hashed", models=["live-team"])) + + assert denial is None + + +async def _run_vertex_live_route(client_frames, valid_token): + """Drive the /vertex_ai/live route, so the gate under test is whichever one the route itself wires""" + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + vertex_ai_live_websocket_passthrough, + vertex_llm_base, + ) + + upstream_ws = RecordingUpstreamWebSocket() + websocket = _client_websocket( + AsyncMock(side_effect=[*client_frames, {"type": "websocket.disconnect"}]), + ) + + with ( + _patched_websocket_passthrough_environment(upstream_ws) as environment, + patch.object( # test-quality-ok: an OAuth exchange with Google, so this is the boundary being faked + vertex_llm_base, + "_ensure_access_token_async", + AsyncMock(return_value=("token-abc", "proj-db")), + ), + ): + await vertex_ai_live_websocket_passthrough( + websocket=websocket, + model="gemini-live-2.5-flash", + vertex_project="proj-db", + vertex_location="global", + user_api_key_dict=valid_token, + ) + + return websocket, upstream_ws, environment + + +@pytest.mark.asyncio +async def test_vertex_live_route_gates_the_setup_frame_even_when_the_query_model_is_permitted(): + """ + ?model= is all connect-time auth sees, so a key that passes it and then names a different model in its setup + frame reaches Vertex on the gateway's own credentials unless the route gates the frames as well + """ + websocket, upstream_ws, environment = await _run_vertex_live_route( + [_setup_frame("gemini-3-pro-live")], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + upstream_ws.send.assert_not_awaited() + reason = _policy_close_reason(websocket) + assert reason is not None and "gemini-3-pro-live" in reason + refusal = environment.proxy_logging.post_call_failure_hook.await_args.kwargs["original_exception"] + assert refusal.code == "403" + assert refusal.type == ProxyErrorTypes.key_model_access_denied + + +@pytest.mark.asyncio +async def test_vertex_live_route_forwards_a_setup_frame_the_key_is_allowed_to_run(): + websocket, upstream_ws, _ = await _run_vertex_live_route( + [_setup_frame("gemini-live-2.5-flash")], + UserAPIKeyAuth(token="hashed", models=["gemini-live-2.5-flash"]), + ) + + upstream_ws.send.assert_awaited_once() + assert _policy_close_reason(websocket) is None + + +async def _run_passthrough_with_frame_gate(gate): + upstream_ws = ClosableUpstreamWebSocket() + websocket = _client_websocket( + AsyncMock( + side_effect=[ + {"type": "websocket.receive", "text": json.dumps({"setup": {"model": "gemini-live-2.5-flash"}})}, + {"type": "websocket.disconnect"}, + ] + ), + ) + + with _patched_websocket_passthrough_environment(upstream_ws) as environment: + await websocket_passthrough_request( + websocket=websocket, + target=VERTEX_LIVE_TARGET, + custom_headers={"Authorization": "Bearer token"}, + user_api_key_dict=UserAPIKeyAuth(token="hashed"), + forward_headers=False, + endpoint="/vertex_ai/live", + accept_websocket=False, + client_frame_gate=gate, + ) + + return websocket, upstream_ws, environment + + +@pytest.mark.asyncio +async def test_websocket_passthrough_reports_a_refused_client_frame_as_a_failure(): + """A denied session logged through the success hooks tells spend tracking and audit that it succeeded""" + denial = ProxyException( + message="gemini-live-2.5-flash: key not allowed to access model", + type=ProxyErrorTypes.key_model_access_denied, + param="model", + code=403, + ) + + async def refusing_gate(frame_data, valid_token): + return denial + + websocket, upstream_ws, environment = await _run_passthrough_with_frame_gate(refusing_gate) + + assert _policy_close_reason(websocket) == "gemini-live-2.5-flash: key not allowed to access model" + upstream_ws.close.assert_awaited() + environment.proxy_logging.post_call_success_hook.assert_not_awaited() + environment.logging_worker.ensure_initialized_and_enqueue.assert_not_called() + environment.proxy_logging.post_call_failure_hook.assert_awaited_once() + failure_kwargs = environment.proxy_logging.post_call_failure_hook.await_args.kwargs + assert failure_kwargs["original_exception"] is denial + assert failure_kwargs["request_data"]["litellm_logging_obj"].litellm_call_id + + +@pytest.mark.asyncio +async def test_websocket_passthrough_still_reports_an_allowed_session_as_a_success(): + async def permitting_gate(frame_data, valid_token): + return None + + _, _, environment = await _run_passthrough_with_frame_gate(permitting_gate) + + environment.proxy_logging.post_call_failure_hook.assert_not_awaited() + environment.proxy_logging.post_call_success_hook.assert_awaited_once() + environment.logging_worker.ensure_initialized_and_enqueue.assert_called_once() + + @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):