fix: guard vertex live passthrough provider lookup and close-code relay

This commit is contained in:
mateo-berri 2026-08-20 02:15:17 -07:00
parent 021a09b156
commit b9d977aeee
3 changed files with 56 additions and 5 deletions

View file

@ -2397,8 +2397,8 @@ def _resolve_vertex_live_credentials(
model: str | None,
) -> VertexPassThroughCredentials | None:
"""
Resolution order: credentials registered for the requested project/location, then any DB model entry
flagged ``use_in_pass_through``, then ``default_vertex_config`` and the ``DEFAULT_VERTEXAI_*`` env vars
Resolution order: an explicit project/location registration or ``default_vertex_config``, then any DB model
entry flagged ``use_in_pass_through``, then the ``DEFAULT_VERTEXAI_*`` env vars
"""
keyed: Final = passthrough_endpoint_router.get_vertex_credentials(
project_id=vertex_project,
@ -2457,7 +2457,10 @@ def _resolve_alias_to_upstream_model(setup_model: str, llm_router: "Router | Non
)
if upstream is None:
return setup_model
_, provider, _, _ = litellm.get_llm_provider(model=upstream)
try:
_, provider, _, _ = litellm.get_llm_provider(model=upstream)
except litellm.exceptions.BadRequestError:
return upstream
return upstream.removeprefix(f"{provider}/")

View file

@ -32,7 +32,7 @@ from websockets.exceptions import (
ConnectionClosedOK,
InvalidStatus,
)
from websockets.frames import Close
from websockets.frames import EXTERNAL_CLOSE_CODES, Close
import litellm
from litellm._logging import verbose_proxy_logger
@ -1930,13 +1930,18 @@ def _truncated_close_reason(reason: str) -> str:
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
The upstream close worth telling the client about: anything other than a plain, reasonless normal close.
Codes outside ``EXTERNAL_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 EXTERNAL_CLOSE_CODES and not 3000 <= upstream_close.code < 5000:
return None
return upstream_close

View file

@ -5188,6 +5188,49 @@ 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.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"