Merge pull request #37602 from BerriAI/litellm_vertex_live_db_credentials

fix(passthrough): resolve vertex live credentials from db model deployments
This commit is contained in:
Mateo Wang 2026-08-20 10:24:50 -07:00 committed by GitHub
commit 0904d5836e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 1018 additions and 44 deletions

View file

@ -247,6 +247,9 @@ REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS: Final = float(
os.getenv("REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS", "20.0")
)
# RFC 6455 caps the close frame payload at 125 bytes, 2 of which carry the status code
WEBSOCKET_CLOSE_REASON_MAX_BYTES: Final = 123
# SSL/TLS cipher configuration for faster handshakes
# Strategy: Strongly prefer fast modern ciphers, but allow fallback to commonly supported ones
# This balances performance with broad compatibility

View file

@ -9,8 +9,9 @@ Use litellm with Anthropic SDK, Vertex AI SDK, Cohere SDK, etc.
import json
import os
import re
from collections.abc import Callable
from types import MappingProxyType
from typing import Annotated, Any, Final, cast
from typing import TYPE_CHECKING, Annotated, Any, Final, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
@ -55,11 +56,15 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
from .passthrough_endpoint_router import PassthroughEndpointRouter
if TYPE_CHECKING:
from litellm.router import Router
vertex_llm_base: Final = VertexBase()
router: Final = APIRouter()
openai_passthrough_router: Final = APIRouter()
@ -2373,6 +2378,112 @@ async def cursor_proxy_route(
return received_value
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_PUBLISHER_MODEL_PREFIX: Final = "publishers/google/models/"
VERTEX_PUBLISHERS_SEGMENT: Final = "publishers/"
def _vertex_publisher_model_suffix(model: str) -> str:
"""
Turn whatever the client named into the ``publishers/<publisher>/models/<id>`` tail of a Vertex resource name.
Clients send bare ids, LiteLLM ids (``vertex_ai/gemini-live-2.5-flash``), and the Live SDK's ``models/<id>``,
and a publisher model id never contains a slash, so anything ahead of the last one is addressing, not identity
"""
publishers_at: Final = model.find(VERTEX_PUBLISHERS_SEGMENT)
if publishers_at != -1:
return model[publishers_at:]
return f"{VERTEX_PUBLISHER_MODEL_PREFIX}{model.rsplit('/', 1)[-1]}"
def _get_llm_router() -> "Router | None":
from litellm.proxy.proxy_server import llm_router
return llm_router
def _resolve_vertex_live_credentials(
vertex_project: str | None,
vertex_location: str | None,
model: str | None,
) -> VertexPassThroughCredentials | None:
"""
Resolution order: an explicit project/location registration, then ``default_vertex_config`` (which the proxy
fills from the ``DEFAULT_VERTEXAI_*`` env vars whenever the yaml leaves it out), then any DB model entry
flagged ``use_in_pass_through``.
DB entries come last on purpose: an operator who set a global default already said which project
pass-through traffic should bill to, and this route silently ignoring that would be the worse surprise
"""
keyed: Final = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,
)
if keyed is not None and keyed.vertex_project is not None:
return keyed
from_deployments: Final = passthrough_endpoint_router.get_vertex_credentials_from_router_deployments(model=model)
if from_deployments is not None:
return from_deployments
if keyed is not None:
return keyed
passthrough_endpoint_router.set_default_vertex_config()
return passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,
)
def _build_vertex_live_setup_model_rewriter(
vertex_project: str | None,
vertex_location: str | None,
llm_router: "Router | None",
) -> Callable[[str], str] | None:
"""
Rewrite the ``setup`` frame's model into the full Vertex resource path the Live API requires.
Clients address the gateway the way they address LiteLLM (bare id or model alias); Vertex reads anything
that is not a ``projects/...`` path as a project name and closes the socket
"""
if vertex_project is None or vertex_location is None:
return None
def rewrite(setup_model: str) -> str:
if setup_model.startswith("projects/"):
return setup_model
aliased: Final = _resolve_alias_to_upstream_model(setup_model, llm_router)
return f"projects/{vertex_project}/locations/{vertex_location}/{_vertex_publisher_model_suffix(aliased)}"
return rewrite
def _resolve_alias_to_upstream_model(setup_model: str, llm_router: "Router | None") -> str:
"""
The Live SDK wraps whatever the caller typed as ``models/<name>``, 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,
)
if upstream is None:
return setup_model
try:
_, provider, _, _ = litellm.get_llm_provider(model=upstream)
except litellm.exceptions.BadRequestError:
return upstream
return upstream.removeprefix(f"{provider}/")
async def vertex_ai_live_websocket_passthrough(
websocket: WebSocket,
model: str | None = None,
@ -2396,51 +2507,38 @@ async def vertex_ai_live_websocket_passthrough(
await websocket.accept()
incoming_headers: Final = dict(websocket.headers)
vertex_credentials_config = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,
vertex_credentials_config: Final = _resolve_vertex_live_credentials(
vertex_project=vertex_project,
vertex_location=vertex_location,
model=model,
)
if vertex_credentials_config is None:
# Attempt to load defaults from environment/config if not already initialised
passthrough_endpoint_router.set_default_vertex_config()
vertex_credentials_config = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
location=vertex_location,
)
resolved_project = vertex_project
resolved_location: str | None = vertex_location
credentials_value: str | None = None
if vertex_credentials_config is not None:
resolved_project = resolved_project or vertex_credentials_config.vertex_project
temp_location: Final = resolved_location or vertex_credentials_config.vertex_location
# Ensure resolved_location is a string
if isinstance(temp_location, dict) or temp_location is not None:
resolved_location = str(temp_location)
else:
resolved_location = None
credentials_value = (
str(vertex_credentials_config.vertex_credentials)
if vertex_credentials_config.vertex_credentials is not None
else None
)
configured_project: Final = vertex_project or (
vertex_credentials_config.vertex_project if vertex_credentials_config is not None else None
)
configured_location: Final = vertex_location or (
vertex_credentials_config.vertex_location if vertex_credentials_config is not None else None
)
credentials_value: Final = (
vertex_credentials_config.vertex_credentials if vertex_credentials_config is not None else None
)
try:
resolved_location = resolved_location or (vertex_llm_base.get_default_vertex_location())
if model:
resolved_location = vertex_llm_base.get_vertex_region(
vertex_region=resolved_location,
resolved_location: Final = (
vertex_llm_base.get_vertex_region(
vertex_region=configured_location or vertex_llm_base.get_default_vertex_location(),
model=model,
)
if model
else configured_location or vertex_llm_base.get_default_vertex_location()
)
(
access_token,
resolved_project,
) = await vertex_llm_base._ensure_access_token_async(
credentials=credentials_value,
project_id=resolved_project,
project_id=configured_project,
custom_llm_provider="vertex_ai_beta",
)
except Exception as e:
@ -2453,7 +2551,7 @@ async def vertex_ai_live_websocket_passthrough(
request_data={},
)
if websocket.client_state != WebSocketState.DISCONNECTED:
await websocket.close(code=1011, reason="Vertex AI authentication failed")
await websocket.close(code=1011, reason=VERTEX_LIVE_UNCONFIGURED_CLOSE_REASON)
return
host_location: Final = resolved_location or vertex_llm_base.get_default_vertex_location()
@ -2485,6 +2583,11 @@ async def vertex_ai_live_websocket_passthrough(
forward_headers=False,
endpoint="/vertex_ai/live",
accept_websocket=False,
setup_model_rewriter=_build_vertex_live_setup_model_rewriter(
vertex_project=resolved_project,
vertex_location=resolved_location,
llm_router=_get_llm_router(),
),
)

View file

@ -5,7 +5,7 @@ import json
import posixpath
import traceback
from base64 import b64encode
from collections.abc import AsyncGenerator, Callable, Mapping
from collections.abc import AsyncGenerator, Callable, Iterable, Mapping
from datetime import datetime
from itertools import groupby
from typing import Any, Final, TypedDict, cast
@ -32,11 +32,15 @@ from websockets.exceptions import (
ConnectionClosedOK,
InvalidStatus,
)
from websockets.frames import Close, CloseCode
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.constants import (
MAXIMUM_TRACEBACK_LINES_TO_LOG,
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
@ -1890,6 +1894,72 @@ def create_websocket_passthrough_route(
return websocket_endpoint_func
def _rewrite_vertex_live_setup_model(text_data: str, setup_model_rewriter: Callable[[str], str] | None) -> str:
"""
Rewrite the model of a Vertex AI Live ``setup`` frame, leaving every other frame byte-identical
"""
if setup_model_rewriter is None:
return text_data
try:
message: Final = json.loads(text_data)
except json.JSONDecodeError:
return text_data
if not isinstance(message, dict):
return text_data
setup: Final = message.get("setup")
if not isinstance(setup, dict):
return text_data
setup_model: Final = setup.get("model")
if not isinstance(setup_model, str):
return text_data
rewritten_model: Final = setup_model_rewriter(setup_model)
if rewritten_model == setup_model:
return text_data
return json.dumps({**message, "setup": {**setup, "model": rewritten_model}}) # mutable-ok: one-shot json payload
def _truncated_close_reason(reason: str) -> str:
"""
Fit a close reason inside the byte budget a WebSocket close frame allows, without splitting a character
"""
encoded: Final = reason.encode("utf-8")
if len(encoded) <= WEBSOCKET_CLOSE_REASON_MAX_BYTES:
return reason
return encoded[:WEBSOCKET_CLOSE_REASON_MAX_BYTES].decode("utf-8", errors="ignore")
SENDABLE_CLOSE_CODES: Final = frozenset(CloseCode) - frozenset(
{CloseCode.NO_STATUS_RCVD, CloseCode.ABNORMAL_CLOSURE, CloseCode.TLS_HANDSHAKE}
)
def _client_socket_is_open(websocket: WebSocket) -> bool:
"""
Starlette tracks the two halves separately and raises on a second close, so both have to still be live
"""
return (
websocket.client_state != WebSocketState.DISCONNECTED
and websocket.application_state != WebSocketState.DISCONNECTED
)
def _upstream_close_to_relay(task_results: Iterable[object]) -> Close | None:
"""
The upstream close worth telling the client about: anything other than a plain, reasonless normal close.
Codes outside ``SENDABLE_CLOSE_CODES`` and the private range never travel on the wire (1006 for a socket that
died without a close frame, 1005 for one that sent no code), so relaying them would build an invalid frame
"""
upstream_close: Final = next((result for result in task_results if isinstance(result, Close)), None)
if upstream_close is None:
return None
if upstream_close.code == 1000 and upstream_close.reason == "":
return None
if upstream_close.code not in SENDABLE_CLOSE_CODES and not 3000 <= upstream_close.code < 5000:
return None
return upstream_close
async def websocket_passthrough_request(
websocket: WebSocket,
target: str,
@ -1899,6 +1969,7 @@ async def websocket_passthrough_request(
endpoint: str | None = None,
cost_per_request: float | None = None,
accept_websocket: bool = True,
setup_model_rewriter: Callable[[str], str] | None = None,
):
"""
WebSocket passthrough request handler.
@ -1911,6 +1982,7 @@ async def websocket_passthrough_request(
forward_headers: Whether to forward incoming headers
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
"""
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.proxy_server import proxy_logging_obj
@ -2100,7 +2172,7 @@ async def websocket_passthrough_request(
)
# Not a JSON message or doesn't contain setup data
await upstream_ws.send(text_data)
await upstream_ws.send(_rewrite_vertex_live_setup_model(text_data, setup_model_rewriter))
elif bytes_data is not None:
await upstream_ws.send(bytes_data)
except asyncio.CancelledError:
@ -2111,8 +2183,8 @@ async def websocket_passthrough_request(
)
await upstream_ws.close()
async def forward_upstream_to_client() -> None:
"""Forward messages from upstream to client WebSocket"""
async def forward_upstream_to_client() -> Close | None:
"""Forward messages from upstream to client WebSocket, returning the upstream's close frame"""
try:
# Wait for the first response from upstream
raw_response = await upstream_ws.recv(decode=False)
@ -2177,6 +2249,7 @@ async def websocket_passthrough_request(
except (ConnectionClosedOK, ConnectionClosedError) as e:
verbose_proxy_logger.debug("Upstream WebSocket connection closed: %s", e)
return e.rcvd
except asyncio.CancelledError:
verbose_proxy_logger.debug("asyncio.CancelledError in forward_upstream_to_client")
raise
@ -2209,6 +2282,13 @@ async def websocket_passthrough_request(
if exception is not None:
raise exception
upstream_close: Final = _upstream_close_to_relay(task.result() for task in done)
if upstream_close is not None and _client_socket_is_open(websocket):
await websocket.close(
code=upstream_close.code,
reason=_truncated_close_reason(upstream_close.reason),
)
end_time: Final = datetime.now()
# Update passthrough logging payload with response data
@ -2294,7 +2374,7 @@ async def websocket_passthrough_request(
),
)
if websocket.client_state != WebSocketState.DISCONNECTED:
if _client_socket_is_open(websocket):
await websocket.close(
code=getattr(exc, "status_code", 1011),
reason="Upstream connection rejected",
@ -2322,10 +2402,10 @@ async def websocket_passthrough_request(
),
)
if websocket.client_state != WebSocketState.DISCONNECTED:
if _client_socket_is_open(websocket):
await websocket.close(code=1011, reason="WebSocket passthrough error")
finally:
if websocket.client_state != WebSocketState.DISCONNECTED:
if _client_socket_is_open(websocket):
await websocket.close()

View file

@ -1,3 +1,4 @@
import json
from collections.abc import Callable
from typing import TYPE_CHECKING, Final
@ -10,7 +11,7 @@ from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
from litellm.types.router import LiteLLMParamsTypedDict
from litellm.types.router import DeploymentTypedDict, LiteLLMParamsTypedDict
if TYPE_CHECKING:
from litellm.router import Router
@ -27,6 +28,15 @@ def _get_str_value(values: dict[str, object] | None, key: str) -> str | None:
return value if isinstance(value, str) else None
def _credential_identity(credentials: VERTEX_CREDENTIALS_TYPES | None) -> str | None:
"""
A hashable stand-in for a credential, so two deployments can be compared for holding the same one
"""
if isinstance(credentials, dict):
return json.dumps(credentials, sort_keys=True)
return credentials
class PassthroughEndpointRouter:
"""
Use this class to Get credentials for pass-through endpoints
@ -120,6 +130,86 @@ class PassthroughEndpointRouter:
return None
return provider
def get_vertex_credentials_from_router_deployments(self, model: str | None) -> VertexPassThroughCredentials | None:
"""
Resolve vertex pass-through credentials from the live router deployments flagged ``use_in_pass_through``.
``deployment_key_to_vertex_credentials`` is only reachable when the caller names a project and location,
which WebSocket clients never do, so DB-stored deployments need this lookup to be usable at all.
With no model to go on, only deployments that agree on a project, a location, and a credential answer:
guessing between two Vertex projects would mint a token for one and later send the other one's model name
"""
llm_router: Final = self.llm_router_getter()
if llm_router is None:
return None
resolved: Final = tuple(
(deployment, credentials)
for deployment in (llm_router.get_model_list() or ())
if (credentials := self._resolve_vertex_deployment_credentials(deployment["litellm_params"])) is not None
)
matched: Final = next(
(
credentials
for deployment, credentials in resolved
if model is not None and self._deployment_matches_model(deployment, model)
),
None,
)
if matched is not None:
return matched
targets: Final = frozenset(
(
credentials.vertex_project,
credentials.vertex_location,
_credential_identity(credentials.vertex_credentials),
)
for _, credentials in resolved
)
if len(targets) != 1:
return None
return resolved[0][1]
def _resolve_vertex_deployment_credentials(
self, litellm_params: LiteLLMParamsTypedDict
) -> VertexPassThroughCredentials | None:
if litellm_params.get("use_in_pass_through") is not True:
return None
if self._get_deployment_provider(litellm_params) != "vertex_ai":
return None
credential_name: Final = litellm_params.get("litellm_credential_name")
credential_values: Final = (
CredentialAccessor.get_credential_values(credential_name) if credential_name is not None else None
)
vertex_project: Final = _get_str_value(credential_values, "vertex_project") or litellm_params.get(
"vertex_project"
)
vertex_location: Final = _get_str_value(credential_values, "vertex_location") or litellm_params.get(
"vertex_location"
)
stored_credentials: Final = (
credential_values.get("vertex_credentials") if credential_values is not None else None
)
vertex_credentials: Final = (
stored_credentials if isinstance(stored_credentials, (str, dict)) else None
) or litellm_params.get("vertex_credentials")
if vertex_project is None or vertex_location is None:
return None
return VertexPassThroughCredentials(
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
)
@staticmethod
def _deployment_matches_model(deployment: DeploymentTypedDict, model: str) -> bool:
upstream_model: Final = deployment["litellm_params"].get("model")
return model in (
deployment.get("model_name"),
upstream_model,
upstream_model.split("/", 1)[-1] if upstream_model is not None else None,
)
def _get_vertex_env_vars(self) -> VertexPassThroughCredentials:
"""
Helper to get vertex pass through config from environment variables

View file

@ -3887,3 +3887,211 @@ class TestComprehendMedicalProxyRoute:
user_api_key_dict=Mock(),
)
assert exc_info.value.status_code == 400
LIVE_RESOURCE_PATH = "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash"
class TestVertexAILiveWebsocketPassthrough:
def _websocket(self):
from starlette.websockets import WebSocketState
websocket = MagicMock()
websocket.accept = AsyncMock()
websocket.close = AsyncMock()
websocket.headers = {}
websocket.client_state = WebSocketState.CONNECTED
return websocket
def _clear_vertex_env(self, monkeypatch):
monkeypatch.delenv("DEFAULT_VERTEXAI_PROJECT", raising=False)
monkeypatch.delenv("DEFAULT_VERTEXAI_LOCATION", raising=False)
monkeypatch.delenv("DEFAULT_GOOGLE_APPLICATION_CREDENTIALS", raising=False)
@pytest.mark.asyncio
async def test_uses_db_deployment_credentials_without_query_params(self, monkeypatch):
from litellm.proxy.pass_through_endpoints import (
llm_passthrough_endpoints as passthrough_module,
)
llm_router = litellm.Router(
model_list=[
{
"model_name": "gemini-live",
"litellm_params": {
"model": "vertex_ai/gemini-live-2.5-flash",
"use_in_pass_through": True,
"vertex_project": "proj-db",
"vertex_location": "global",
"vertex_credentials": '{"type": "service_account"}',
},
}
]
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
passthrough_module.passthrough_endpoint_router, "default_vertex_config", None
)
self._clear_vertex_env(monkeypatch)
websocket = self._websocket()
ensure_token = AsyncMock(return_value=("token-abc", "proj-db"))
ws_passthrough = AsyncMock()
with (
patch.object(passthrough_module.vertex_llm_base, "_ensure_access_token_async", ensure_token),
patch.object(passthrough_module, "websocket_passthrough_request", ws_passthrough),
):
await passthrough_module.vertex_ai_live_websocket_passthrough(
websocket=websocket,
user_api_key_dict=UserAPIKeyAuth(),
)
ensure_token.assert_awaited_once_with(
credentials='{"type": "service_account"}',
project_id="proj-db",
custom_llm_provider="vertex_ai_beta",
)
passthrough_kwargs = ws_passthrough.await_args.kwargs
assert passthrough_kwargs["target"] == (
"wss://aiplatform.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
)
assert passthrough_kwargs["custom_headers"]["Authorization"] == "Bearer token-abc"
rewriter = passthrough_kwargs["setup_model_rewriter"]
assert rewriter("gemini-live") == (
"projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash"
)
websocket.close.assert_not_awaited()
@pytest.mark.parametrize(
"setup_model, expected",
[
("gemini-live-2.5-flash", LIVE_RESOURCE_PATH),
("models/gemini-live-2.5-flash", LIVE_RESOURCE_PATH),
("vertex_ai/gemini-live-2.5-flash", LIVE_RESOURCE_PATH),
("gemini-live", LIVE_RESOURCE_PATH),
("models/gemini-live", LIVE_RESOURCE_PATH),
(
"publishers/meta/models/llama-3.3-70b-instruct-maas",
"projects/proj-db/locations/global/publishers/meta/models/llama-3.3-70b-instruct-maas",
),
(
"projects/other/locations/us-central1/publishers/google/models/gemini-2.0-flash",
"projects/other/locations/us-central1/publishers/google/models/gemini-2.0-flash",
),
],
)
def test_setup_model_rewriter_normalises_the_forms_clients_send(self, setup_model, expected):
from litellm.proxy.pass_through_endpoints import (
llm_passthrough_endpoints as passthrough_module,
)
llm_router = litellm.Router(
model_list=[
{
"model_name": "gemini-live",
"litellm_params": {
"model": "vertex_ai/gemini-live-2.5-flash",
"use_in_pass_through": True,
"vertex_project": "proj-db",
"vertex_location": "global",
},
}
]
)
rewriter = passthrough_module._build_vertex_live_setup_model_rewriter(
vertex_project="proj-db",
vertex_location="global",
llm_router=llm_router,
)
assert rewriter is not None
assert rewriter(setup_model) == expected
@pytest.mark.asyncio
async def test_default_vertex_config_outranks_db_deployment(self, monkeypatch):
from litellm.proxy.pass_through_endpoints import (
llm_passthrough_endpoints as passthrough_module,
)
from litellm.types.passthrough_endpoints.vertex_ai import (
VertexPassThroughCredentials,
)
llm_router = litellm.Router(
model_list=[
{
"model_name": "gemini-live",
"litellm_params": {
"model": "vertex_ai/gemini-live-2.5-flash",
"use_in_pass_through": True,
"vertex_project": "proj-db",
"vertex_location": "global",
"vertex_credentials": '{"type": "db_account"}',
},
}
]
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
monkeypatch.setattr(
passthrough_module.passthrough_endpoint_router,
"default_vertex_config",
VertexPassThroughCredentials(
vertex_project="proj-env",
vertex_location="global",
vertex_credentials='{"type": "env_account"}',
),
)
self._clear_vertex_env(monkeypatch)
websocket = self._websocket()
ensure_token = AsyncMock(return_value=("token-abc", "proj-env"))
ws_passthrough = AsyncMock()
with (
patch.object(passthrough_module.vertex_llm_base, "_ensure_access_token_async", ensure_token),
patch.object(passthrough_module, "websocket_passthrough_request", ws_passthrough),
):
await passthrough_module.vertex_ai_live_websocket_passthrough(
websocket=websocket,
model="gemini-live",
user_api_key_dict=UserAPIKeyAuth(),
)
ensure_token.assert_awaited_once_with(
credentials='{"type": "env_account"}',
project_id="proj-env",
custom_llm_provider="vertex_ai_beta",
)
rewriter = ws_passthrough.await_args.kwargs["setup_model_rewriter"]
assert rewriter("gemini-live") == (
"projects/proj-env/locations/global/publishers/google/models/gemini-live-2.5-flash"
)
@pytest.mark.asyncio
async def test_credential_failure_close_names_configuration_options(self, monkeypatch):
from litellm.proxy.pass_through_endpoints import (
llm_passthrough_endpoints as passthrough_module,
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr(
passthrough_module.passthrough_endpoint_router, "default_vertex_config", None
)
self._clear_vertex_env(monkeypatch)
websocket = self._websocket()
ensure_token = AsyncMock(side_effect=Exception("Unable to find your credentials"))
with (
patch.object(passthrough_module.vertex_llm_base, "_ensure_access_token_async", ensure_token),
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
):
mock_proxy_logging.post_call_failure_hook = AsyncMock()
await passthrough_module.vertex_ai_live_websocket_passthrough(
websocket=websocket,
user_api_key_dict=UserAPIKeyAuth(),
)
close_kwargs = websocket.close.await_args.kwargs
assert close_kwargs["code"] == 1011
assert "use_in_pass_through" in close_kwargs["reason"]
assert "default_vertex_config" in close_kwargs["reason"]
assert len(close_kwargs["reason"].encode("utf-8")) <= 123

View file

@ -4994,6 +4994,287 @@ async def test_websocket_passthrough_forwards_non_ascii_first_frame():
assert all(call.kwargs.get("code") != 1011 for call in websocket.close.await_args_list)
class ClosingUpstreamWebSocket:
def __init__(self, close_exc: Exception):
self._close_exc = close_exc
self.close = AsyncMock()
self.send = AsyncMock()
async def recv(self, decode: bool = True):
raise self._close_exc
def __aiter__(self):
return self
async def __anext__(self):
raise StopAsyncIteration
class RecordingUpstreamWebSocket:
def __init__(self):
self.close = AsyncMock()
self.send = AsyncMock()
async def recv(self, decode: bool = True):
await asyncio.Event().wait()
def __aiter__(self):
return self
async def __anext__(self):
raise StopAsyncIteration
def _client_websocket(receive):
from starlette.websockets import WebSocketState
websocket = MagicMock()
websocket.accept = AsyncMock()
websocket.send_text = AsyncMock()
websocket.send_bytes = AsyncMock()
websocket.receive = receive
websocket.headers = {}
websocket.client_state = WebSocketState.CONNECTED
websocket.application_state = WebSocketState.CONNECTED
def _mark_closed(*args, **kwargs):
websocket.application_state = WebSocketState.DISCONNECTED
websocket.close = AsyncMock(side_effect=_mark_closed)
return websocket
@contextmanager
def _patched_websocket_passthrough_environment(upstream_ws):
with (
patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging,
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect",
return_value=FakeUpstreamConnect(upstream_ws),
),
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.GLOBAL_LOGGING_WORKER"
) as mock_worker,
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_success_hook = AsyncMock()
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_worker.ensure_initialized_and_enqueue = MagicMock(
side_effect=lambda async_coroutine: async_coroutine.close()
)
yield
async def _pending_receive():
await asyncio.Event().wait()
@pytest.mark.asyncio
async def test_websocket_passthrough_relays_upstream_policy_close_to_client():
from websockets.exceptions import ConnectionClosedError
from websockets.frames import Close
upstream_reason = "Publisher Model `projects/p/locations/global/publishers/google/models/nope` was not found"
upstream_ws = ClosingUpstreamWebSocket(
ConnectionClosedError(
rcvd=Close(1008, upstream_reason),
sent=Close(1008, ""),
rcvd_then_sent=True,
)
)
websocket = _client_websocket(_pending_receive)
with _patched_websocket_passthrough_environment(upstream_ws):
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,
)
websocket.close.assert_awaited_once_with(code=1008, reason=upstream_reason)
@pytest.mark.asyncio
async def test_websocket_passthrough_keeps_normal_upstream_close_normal():
from websockets.exceptions import ConnectionClosedOK
from websockets.frames import Close
upstream_ws = ClosingUpstreamWebSocket(
ConnectionClosedOK(rcvd=Close(1000, ""), sent=Close(1000, ""), rcvd_then_sent=True)
)
websocket = _client_websocket(_pending_receive)
with _patched_websocket_passthrough_environment(upstream_ws):
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,
)
websocket.close.assert_awaited_once_with()
async def _run_setup_rewrite_passthrough(setup_model: str, llm_router) -> str:
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": setup_model, "generationConfig": {"responseModalities": ["TEXT"]}}})
websocket = _client_websocket(
AsyncMock(
side_effect=[
{"type": "websocket.receive", "text": setup_frame},
{"type": "websocket.disconnect"},
]
)
)
with _patched_websocket_passthrough_environment(upstream_ws):
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=llm_router,
),
)
upstream_ws.send.assert_awaited_once()
return upstream_ws.send.await_args.args[0]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"setup_model",
["gemini-live-2.5-flash", "publishers/google/models/gemini-live-2.5-flash"],
)
async def test_websocket_passthrough_rewrites_setup_model_to_full_resource(setup_model):
sent_frame = await _run_setup_rewrite_passthrough(setup_model, llm_router=None)
sent_setup = json.loads(sent_frame)["setup"]
assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/gemini-live-2.5-flash"
assert sent_setup["generationConfig"] == {"responseModalities": ["TEXT"]}
@pytest.mark.asyncio
async def test_websocket_passthrough_rewrites_gateway_alias_setup_model():
llm_router = litellm.Router(
model_list=[
{
"model_name": "gemini-live",
"litellm_params": {"model": "vertex_ai/gemini-live-2.5-flash"},
}
]
)
sent_frame = await _run_setup_rewrite_passthrough("gemini-live", llm_router=llm_router)
sent_setup = json.loads(sent_frame)["setup"]
assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/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):
from websockets.exceptions import ConnectionClosedError
from websockets.frames import Close
rcvd = {
None: None,
"abnormal": Close(1006, "connection died"),
"no_status": Close(1005, ""),
}[rcvd_close]
upstream_ws = ClosingUpstreamWebSocket(
ConnectionClosedError(rcvd=rcvd, sent=None, rcvd_then_sent=None)
)
websocket = _client_websocket(_pending_receive)
with _patched_websocket_passthrough_environment(upstream_ws):
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,
)
websocket.close.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_websocket_passthrough_rewrites_alias_of_unrecognised_upstream_model():
llm_router = MagicMock()
llm_router.get_model_list.return_value = [
{"model_name": "gemini-live", "litellm_params": {"model": "self-hosted-live-endpoint"}}
]
sent_frame = await _run_setup_rewrite_passthrough("gemini-live", llm_router=llm_router)
sent_setup = json.loads(sent_frame)["setup"]
assert sent_setup["model"] == "projects/proj-db/locations/global/publishers/google/models/self-hosted-live-endpoint"
@pytest.mark.asyncio
async def test_websocket_passthrough_leaves_full_resource_setup_model_untouched():
full_resource = "projects/other/locations/us-central1/publishers/google/models/gemini-2.0-flash-live-preview-04-09"
sent_frame = await _run_setup_rewrite_passthrough(full_resource, llm_router=None)
assert json.loads(sent_frame)["setup"]["model"] == full_resource
assert sent_frame == json.dumps(
{"setup": {"model": full_resource, "generationConfig": {"responseModalities": ["TEXT"]}}}
)
@pytest.mark.asyncio
async def test_websocket_passthrough_does_not_close_twice_when_success_logging_fails():
from websockets.exceptions import ConnectionClosedError
from websockets.frames import Close
upstream_reason = "Publisher Model `projects/p/locations/global/publishers/google/models/nope` was not found"
upstream_ws = ClosingUpstreamWebSocket(
ConnectionClosedError(rcvd=Close(1008, upstream_reason), sent=Close(1008, ""), rcvd_then_sent=True)
)
websocket = _client_websocket(_pending_receive)
with (
_patched_websocket_passthrough_environment(upstream_ws),
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints."
"GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=RuntimeError("logging worker down"),
),
):
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,
)
websocket.close.assert_awaited_once_with(code=1008, reason=upstream_reason)
def _passthrough_kwargs_for_reservation(
user_api_key_dict: UserAPIKeyAuth, parsed_body: Optional[dict] = None
) -> dict:

View file

@ -172,3 +172,212 @@ def test_returns_none_when_no_router_and_no_env():
passthrough_router = _passthrough_router(None)
assert passthrough_router.get_credentials(custom_llm_provider="openai", region_name=None) is None
def _vertex_credential(name: str, values: dict) -> CredentialItem:
return CredentialItem(credential_name=name, credential_values=values, credential_info={})
def _vertex_deployment(model_name: str, model: str, **litellm_params) -> dict:
return {
"model_name": model_name,
"litellm_params": {"model": model, "use_in_pass_through": True, **litellm_params},
}
def test_vertex_deployment_resolves_via_named_credential():
CredentialAccessor.upsert_credentials(
[
_vertex_credential(
"cred_gcp",
{
"vertex_project": "proj-db",
"vertex_location": "global",
"vertex_credentials": '{"type": "service_account"}',
},
)
]
)
llm_router = litellm.Router(
model_list=[
_vertex_deployment(
"gemini-live", "vertex_ai/gemini-live-2.5-flash", litellm_credential_name="cred_gcp"
)
]
)
passthrough_router = _passthrough_router(llm_router)
resolved = passthrough_router.get_vertex_credentials_from_router_deployments(model=None)
assert resolved is not None
assert resolved.vertex_project == "proj-db"
assert resolved.vertex_location == "global"
assert resolved.vertex_credentials == '{"type": "service_account"}'
def test_vertex_deployment_resolves_from_inline_litellm_params():
llm_router = litellm.Router(
model_list=[
_vertex_deployment(
"gemini-live",
"vertex_ai/gemini-live-2.5-flash",
vertex_project="proj-inline",
vertex_location="us-east4",
vertex_credentials='{"type": "service_account", "project_id": "proj-inline"}',
)
]
)
passthrough_router = _passthrough_router(llm_router)
resolved = passthrough_router.get_vertex_credentials_from_router_deployments(model=None)
assert resolved is not None
assert resolved.vertex_project == "proj-inline"
assert resolved.vertex_location == "us-east4"
assert resolved.vertex_credentials == '{"type": "service_account", "project_id": "proj-inline"}'
def _two_vertex_deployments_router() -> litellm.Router:
return litellm.Router(
model_list=[
_vertex_deployment(
"gemini-flash",
"vertex_ai/gemini-2.5-flash",
vertex_project="proj-first",
vertex_location="us-central1",
),
_vertex_deployment(
"gemini-live",
"vertex_ai/gemini-live-2.5-flash",
vertex_project="proj-live",
vertex_location="global",
),
]
)
def test_vertex_model_hint_prefers_matching_deployment():
passthrough_router = _passthrough_router(_two_vertex_deployments_router())
by_alias = passthrough_router.get_vertex_credentials_from_router_deployments(model="gemini-live")
by_upstream_id = passthrough_router.get_vertex_credentials_from_router_deployments(
model="gemini-live-2.5-flash"
)
assert by_alias is not None and by_alias.vertex_project == "proj-live"
assert by_upstream_id is not None and by_upstream_id.vertex_project == "proj-live"
def test_vertex_without_usable_hint_refuses_to_guess_between_projects():
passthrough_router = _passthrough_router(_two_vertex_deployments_router())
assert passthrough_router.get_vertex_credentials_from_router_deployments(model="unknown-model") is None
assert passthrough_router.get_vertex_credentials_from_router_deployments(model=None) is None
def test_vertex_without_hint_falls_back_when_deployments_share_a_target():
llm_router = litellm.Router(
model_list=[
_vertex_deployment(
"gemini-flash", "vertex_ai/gemini-2.5-flash", vertex_project="proj-one", vertex_location="global"
),
_vertex_deployment(
"gemini-live", "vertex_ai/gemini-live-2.5-flash", vertex_project="proj-one", vertex_location="global"
),
]
)
passthrough_router = _passthrough_router(llm_router)
resolved = passthrough_router.get_vertex_credentials_from_router_deployments(model=None)
assert resolved is not None and resolved.vertex_project == "proj-one"
def test_vertex_without_hint_refuses_to_guess_between_service_accounts():
llm_router = litellm.Router(
model_list=[
_vertex_deployment(
"gemini-flash",
"vertex_ai/gemini-2.5-flash",
vertex_project="proj-one",
vertex_location="global",
vertex_credentials='{"client_email": "flash@proj-one.iam"}',
),
_vertex_deployment(
"gemini-live",
"vertex_ai/gemini-live-2.5-flash",
vertex_project="proj-one",
vertex_location="global",
vertex_credentials='{"client_email": "live@proj-one.iam"}',
),
]
)
passthrough_router = _passthrough_router(llm_router)
assert passthrough_router.get_vertex_credentials_from_router_deployments(model=None) is None
def test_vertex_named_credential_keeps_dict_service_account():
service_account = {"type": "service_account", "client_email": "live@proj-db.iam"}
CredentialAccessor.upsert_credentials(
[
_vertex_credential(
"cred_gcp_dict",
{
"vertex_project": "proj-db",
"vertex_location": "global",
"vertex_credentials": service_account,
},
)
]
)
llm_router = litellm.Router(
model_list=[
_vertex_deployment(
"gemini-live", "vertex_ai/gemini-live-2.5-flash", litellm_credential_name="cred_gcp_dict"
)
]
)
passthrough_router = _passthrough_router(llm_router)
resolved = passthrough_router.get_vertex_credentials_from_router_deployments(model=None)
assert resolved is not None
assert resolved.vertex_credentials == service_account
def test_no_flagged_vertex_deployment_returns_none():
llm_router = litellm.Router(
model_list=[
{
"model_name": "gemini-live",
"litellm_params": {
"model": "vertex_ai/gemini-live-2.5-flash",
"vertex_project": "proj-unflagged",
"vertex_location": "global",
},
},
_flagged_deployment("openai/gpt-4o", api_key="sk-flagged"),
]
)
passthrough_router = _passthrough_router(llm_router)
assert passthrough_router.get_vertex_credentials_from_router_deployments(model=None) is None
assert _passthrough_router(None).get_vertex_credentials_from_router_deployments(model=None) is None
def test_vertex_deployment_with_deleted_credential_is_skipped(monkeypatch):
CredentialAccessor.upsert_credentials(
[_vertex_credential("cred_gone", {"vertex_project": "proj-db", "vertex_location": "global"})]
)
llm_router = litellm.Router(
model_list=[
_vertex_deployment(
"gemini-live", "vertex_ai/gemini-live-2.5-flash", litellm_credential_name="cred_gone"
)
]
)
passthrough_router = _passthrough_router(llm_router)
monkeypatch.setattr(litellm, "credential_list", [])
assert passthrough_router.get_vertex_credentials_from_router_deployments(model=None) is None