fix(vertex_ai): dial the multi-region Live API host for realtime sessions (#45166)

* fix(vertex_ai): dial the multi-region Live API host for realtime sessions

A realtime deployment with vertex_location us or eu dialed
{location}-aiplatform.googleapis.com, which Vertex does not serve for a
multi-region, so the WebSocket handshake came back 404. The URL builder now
resolves its host through the shared Vertex host resolver, which already
maps us and eu to aiplatform.{geo}.rep.googleapis.com and leaves regions
and global unchanged.

* test(vertex_ai): annotate the realtime multi-region URL tests

* fix(vertex_ai): skip the TLS argument when the realtime api_base is plain ws://

The Vertex realtime session and health-check connects passed the shared
TLS context to every websocket dial, so an api_base the transformation
already maps from http:// to ws:// failed with 'ssl argument is
incompatible with a ws:// URI' before any frame was sent. The OpenAI
realtime handler already skips the argument for ws://; the shared helper
now does the same for the Vertex path

* test(integration): cover the Vertex realtime multi-region host through the proxy

Fourteen cells on the scripted Gemini Live upstream: api_base overrides for
us, eu, us-central1 and global complete a turn; malformed locations are
refused before any dial; a missing location defaults to us-central1; the
realtime health check handshakes the override host and reports a malformed
location without dialing; a refused handshake reaches the client and the
next session connects; twelve concurrent sessions across locations each
reach the upstream once; an upstream outage closes every open session and
the proxy recovers

* test(realtime): type the health-check helpers and flatten the burst order

* test(http_handler): type the realtime TLS test signatures
This commit is contained in:
Mateo Wang 2026-10-07 17:39:38 -07:00 • committed by GitHub
parent bf9f59d76b
commit 70e6a2ff24
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 628 additions and 28 deletions

View file

@ -450,6 +450,18 @@ def get_shared_realtime_ssl_context() -> bool | str | ssl.SSLContext:
return _shared_realtime_ssl_context
def realtime_ssl_for_url(url: str) -> bool | str | ssl.SSLContext | None:
if url.startswith("ws://"):
return None
shared: Final = get_shared_realtime_ssl_context()
if shared is not False:
return shared
unverified: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
unverified.check_hostname = False
unverified.verify_mode = ssl.CERT_NONE
return unverified
def mask_sensitive_info(error_message):
# Find the start of the key parameter
if isinstance(error_message, str):

View file

@ -197,7 +197,7 @@ def _rust_responses_websocket_enabled(
return decision(context) is not Decision.PYTHON
from .http_handler import get_shared_realtime_ssl_context
from .http_handler import get_shared_realtime_ssl_context, realtime_ssl_for_url
if TYPE_CHECKING:
from aiohttp import ClientSession
@ -258,7 +258,7 @@ class _WebsocketsModule(Protocol):
*,
additional_headers: Mapping[str, str],
max_size: int | None,
ssl: bool | str | ssl.SSLContext,
ssl: bool | str | ssl.SSLContext | None,
open_timeout: float,
) -> Awaitable["ClientConnection"]: ...
@ -6257,7 +6257,7 @@ class BaseLLMHTTPHandler:
websockets_module: _WebsocketsModule,
url: str,
headers: dict,
ssl_context: bool | str | ssl.SSLContext,
ssl_context: bool | str | ssl.SSLContext | None,
*,
open_timeout: float = 8.0,
max_attempts: int = 3,
@ -6329,12 +6329,7 @@ class BaseLLMHTTPHandler:
)
try:
ssl_context = get_shared_realtime_ssl_context()
if url.startswith("wss://") and ssl_context is False:
# Keep TLS for wss:// while honoring SSL_VERIFY=False semantics.
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
ssl_context.check_hostname = False
ssl_context.verify_mode = ssl.CERT_NONE
ssl_context: Final = realtime_ssl_for_url(url)
provider_backend: Final = await provider_config.open_backend(url, headers)
backend_ws: Final = (
provider_backend

View file

@ -5,9 +5,14 @@ Extends GeminiRealtimeConfig but adapts the WSS URL and auth header for the
Vertex AI endpoint instead of Google AI Studio.
URL pattern:
wss://{location}-aiplatform.googleapis.com/ws/
wss://{vertex host for the location}/ws/
google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent
The host is the one ``litellm.llms.vertex_ai.common_utils.get_vertex_base_url``
resolves for the location: ``{region}-aiplatform.googleapis.com`` for a region,
``aiplatform.{geo}.rep.googleapis.com`` for the ``us`` / ``eu`` multi-regions,
and ``aiplatform.googleapis.com`` for ``global``.
Auth: OAuth2 Bearer token (not an API key).
"""
@ -21,6 +26,7 @@ from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import (
VertexChirpRealtimeConfig,
is_vertex_speech_to_text_model,
)
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
@ -63,12 +69,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
base = base.replace("https://", "wss://").replace("http://", "ws://")
return f"{base}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
location: Final = self._location
if location == "global":
host = "aiplatform.googleapis.com"
else:
host = f"{location}-aiplatform.googleapis.com"
host: Final = get_vertex_base_url(self._location).removeprefix("https://")
return f"wss://{host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
# ------------------------------------------------------------------

View file

@ -36,7 +36,7 @@ from ..litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from ..llms.azure.common_utils import get_azure_ad_token
from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_protocol_for_client
from ..llms.bedrock.realtime.handler import BedrockRealtime
from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from ..llms.custom_httpx.http_handler import realtime_ssl_for_url
from ..llms.openai.realtime.handler import OpenAIRealtime
from ..llms.vertex_ai.audio_transcription.realtime_transformation import is_vertex_speech_to_text_model
from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig, vertex_realtime_config
@ -721,23 +721,21 @@ async def realtime_health_check(
location=resolved_location,
)
url = vertex_realtime_config.get_complete_url(api_base=resolved_api_base, model=model)
vertex_ssl_context: Final = get_shared_realtime_ssl_context()
headers: Final = vertex_realtime_config.validate_environment(headers={}, model=model, api_key=None)
async with websockets.connect(
url,
additional_headers=headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=vertex_ssl_context,
ssl=realtime_ssl_for_url(url),
):
return True
else:
raise ValueError(f"Unsupported model: {model}")
ssl_context: Final = get_shared_realtime_ssl_context()
async with websockets.connect(
url,
additional_headers=auth_headers,
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
ssl=ssl_context,
ssl=realtime_ssl_for_url(url),
):
return True

View file

@ -440,6 +440,37 @@ class Provider:
while pending:
await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id))
async def gemini_live(self, websocket: WebSocket) -> None:
authorization: Final = websocket.headers.get("authorization", "")
scenario_id: Final = websocket.headers.get("x-goog-user-project", "")
self.observations.put(
Observation(
websocket.url.path,
authorization,
{"host": websocket.headers.get("host", "")},
"WEBSOCKET",
scenario_id,
)
)
response: Final = self.scenario_store.get(scenario_id)
if not isinstance(response, RealtimeResponse):
await websocket.close(code=4404)
return
await websocket.accept()
pending: Final = deque(response.events)
async for message in websocket.iter_json():
payload: Final = JSON_OBJECT.validate_python(message)
self.observations.put(
Observation(websocket.url.path, authorization, payload, "WEBSOCKET_FRAME", scenario_id)
)
if "setup" in payload:
await websocket.send_json({"setupComplete": {}})
continue
if not _GEMINI_LIVE_TRIGGERS.intersection(payload):
continue
while pending:
await websocket.send_json(_rendered_realtime_event(pending.popleft(), scenario_id))
@staticmethod
def _response(response: StoredResponse, scenario_id: str) -> Response:
unique_id: Final = f"{scenario_id}-{uuid.uuid4().hex[:8]}"
@ -539,6 +570,7 @@ class Provider:
WebSocketRoute("/openai/v1/realtime", self.realtime),
WebSocketRoute("/openai/realtime", self.realtime),
WebSocketRoute("/v1/asr/realtime", self.muse_realtime),
WebSocketRoute(GEMINI_LIVE_PATH, self.gemini_live),
]
)
@ -562,6 +594,8 @@ def _interaction_body(interaction_id: str, state: InteractionState) -> dict[str,
_REALTIME_TRIGGERS: Final = frozenset({"response.create", "input_audio_buffer.commit"})
_GEMINI_LIVE_TRIGGERS: Final = frozenset({"clientContent", "realtimeInput"})
GEMINI_LIVE_PATH: Final = "/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
_TRANSCRIPTION_UPDATE_REFUSED: Final = "Passing a realtime session update to a transcription session is not allowed."
_REALTIME_UPDATE_REFUSED: Final = "Passing a transcription session update to a realtime session is not allowed."
_NESTED_TURN_DETECTION_TYPE: Final = "session.audio.input.turn_detection.type"

View file

@ -0,0 +1,468 @@
"""Vertex AI Live sessions resolve their host through the shared Vertex resolver.
``VertexAIRealtimeConfig.get_complete_url`` dials ``aiplatform.{us,eu}.rep.googleapis.com`` for the
multi-regions, ``{region}-aiplatform.googleapis.com`` for a region and ``aiplatform.googleapis.com`` for
``global``, and a malformed location fails the shared validator before any socket opens. Every row here
runs against the scripted upstream on 127.0.0.1: an ``api_base`` override pins the path and the host the
upstream sees, the setup frame it receives names the location the proxy resolved, and the malformed rows
pin the validator's answer, which the proxy gives without dialing anything.
"""
from __future__ import annotations
import asyncio
import json
import uuid
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from hashlib import sha256
from itertools import chain, repeat
from pathlib import Path
from typing import Final
import httpx
import pytest
import websockets
from pydantic import JsonValue
from websockets.asyncio.client import ClientConnection
from websockets.exceptions import ConnectionClosed
from tests.integration._support.client import JSON_OBJECT, Gateway, Scenario, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.process import owned_upstream
from tests.integration._support.upstream import GEMINI_LIVE_PATH, ScenarioHandle, delete_scenario, register_scenario
from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse
RecordProperty = Callable[[str, object], None]
pytestmark: Final = pytest.mark.timeout(180)
LIVE_MODEL: Final = "gemini-3.8-live"
LOCATIONS: Final = ("us", "eu", "us-central1", "global")
MALFORMED_LOCATIONS: Final = ("US", "us/evil")
DEFAULT_LOCATION: Final = "us-central1"
INVALID_LOCATION: Final = "Invalid vertex_location format"
HANDSHAKE_REFUSED: Final = "Upstream realtime handshake rejected with HTTP 403"
INTERNAL_CLOSE: Final = 1011
REFUSAL_CLOSE: Final = 1008
SESSIONS_PER_LOCATION: Final = 3
OUTAGE_BURST: Final = 20
INPUT_TOKENS: Final = 7
OUTPUT_TOKENS: Final = 5
CONVERGENCE_SECONDS: Final = 30
SPEND_SQL: Final = 'SELECT call_type, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE api_key = %s'
@dataclass(frozen=True, slots=True)
class Session:
events: tuple[dict[str, JsonValue], ...]
@property
def types(self) -> tuple[str, ...]:
return tuple(string_value(event["type"]) for event in self.events)
@property
def close_code(self) -> int | None:
closes: Final = tuple(event for event in self.events if event["type"] == "closed")
return _integer(closes[-1]["code"]) if closes else None
@property
def session_model(self) -> str:
return string_value(object_value(self.events[0]["session"])["model"])
@property
def response_ids(self) -> tuple[str, ...]:
done: Final = tuple(event for event in self.events if event["type"] == "response.done")
return tuple(string_value(object_value(event["response"])["id"]) for event in done)
@property
def error_messages(self) -> tuple[str, ...]:
errors: Final = tuple(event for event in self.events if event["type"] == "error")
return tuple(string_value(object_value(event["error"])["message"]) for event in errors)
@property
def text(self) -> str:
deltas: Final = tuple(event for event in self.events if event["type"] == "response.output_text.delta")
return "".join(string_value(event["delta"]) for event in deltas)
def completed_turn(self, scenario_id: str) -> bool:
return (
self.types[0] == "session.created"
and self.types[-1] == "response.done"
and self.session_model == LIVE_MODEL
and self.text == f"scripted {scenario_id}"
and len(self.response_ids) == 1
)
def _integer(value: JsonValue) -> int:
assert isinstance(value, int), value
return value
def _ws_base(http_url: str) -> str:
return http_url.replace("https://", "wss://").replace("http://", "ws://")
def _proxy_url(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
def _upstream_url(gateway: Gateway) -> str:
return gateway.upstream_url.rstrip("/")
def _scripted_turn() -> RealtimeResponse:
return RealtimeResponse(
content_type="application/x-realtime",
events=(
{"serverContent": {"modelTurn": {"parts": [{"text": "scripted $REQUEST_ID"}]}}},
{
"serverContent": {"turnComplete": True},
"usageMetadata": {
"promptTokenCount": INPUT_TOKENS,
"responseTokenCount": OUTPUT_TOKENS,
"totalTokenCount": INPUT_TOKENS + OUTPUT_TOKENS,
},
},
),
)
def _scripted(scenario: Scenario) -> ScenarioHandle:
handle: Final = register_scenario(f"vertex-live-{uuid.uuid4().hex[:12]}", _scripted_turn())
scenario.cleanups.callback(delete_scenario, handle)
return handle
def _credentials(upstream_url: str) -> str:
return json.dumps(
{
"type": "external_account",
"audience": "synthetic-vertex-audience",
"subject_token_type": "urn:ietf:params:oauth:token-type:jwt",
"token_url": f"{upstream_url}/_oauth/token",
"credential_source": {"url": f"{upstream_url}/health"},
}
)
def _deployment(
gateway: Gateway, scenario: Scenario, project: str, *, location: str | None, api_base: str | None
) -> str:
name: Final = f"vertex-live-{uuid.uuid4().hex[:12]}"
litellm_params: Final[dict[str, JsonValue]] = {
"model": f"vertex_ai/{LIVE_MODEL}",
"vertex_project": project,
"vertex_credentials": _credentials(_upstream_url(gateway)),
**({} if location is None else {"vertex_location": location}),
**({} if api_base is None else {"api_base": api_base}),
}
created: Final = gateway.post(
"/model/new", {"model_name": name, "litellm_params": litellm_params, "model_info": {"mode": "realtime"}}
)
scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"]))
return name
def _model_path(project: str, location: str) -> str:
return f"projects/{project}/locations/{location}/publishers/google/models/{LIVE_MODEL}"
def _user_turn() -> str:
return json.dumps(
{
"type": "conversation.item.create",
"item": {
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": "say the scripted line"}],
},
}
)
def _close_code(closed: ConnectionClosed) -> int:
return 1006 if closed.rcvd is None else closed.rcvd.code
async def _frames(socket: ClientConnection, turns: int) -> AsyncIterator[dict[str, JsonValue]]:
try:
first: Final = JSON_OBJECT.validate_json(await socket.recv())
yield first
if first.get("type") != "session.created":
async for message in socket:
yield JSON_OBJECT.validate_json(message)
return
for _ in range(turns):
await socket.send(_user_turn())
async for message in socket:
event: Final = JSON_OBJECT.validate_json(message)
yield event
if event.get("type") == "response.done":
break
except ConnectionClosed as closed:
yield {"type": "closed", "code": _close_code(closed)}
async def _collect(socket: ClientConnection, turns: int) -> tuple[dict[str, JsonValue], ...]:
return tuple([frame async for frame in _frames(socket, turns)])
async def _session(ws_base: str, model: str, key: str, *, turns: int = 1) -> Session:
headers: Final = {"Authorization": f"Bearer {key}"}
async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket:
return Session(await asyncio.wait_for(_collect(socket, turns), 60))
def _run(gateway: Gateway, model: str, key: str, *, turns: int = 1) -> Session:
return asyncio.run(_session(_ws_base(_proxy_url(gateway)), model, key, turns=turns))
def _observations(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
return tuple(map(object_value, upstream.get("/__observations").json()["requests"]))
def _for_project(observed: tuple[dict[str, JsonValue], ...], project: str) -> tuple[dict[str, JsonValue], ...]:
return tuple(request for request in observed if request["api_key"] == project)
def _upgrade_hosts(observed: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]:
upgrades: Final = tuple(request for request in observed if request["method"] == "WEBSOCKET")
assert all(request["path"] == GEMINI_LIVE_PATH for request in upgrades), upgrades
return tuple(string_value(object_value(request["body"])["host"]) for request in upgrades)
def _setup_models(observed: tuple[dict[str, JsonValue], ...]) -> tuple[str, ...]:
frames: Final = tuple(
object_value(request["body"]) for request in observed if request["method"] == "WEBSOCKET_FRAME"
)
return tuple(string_value(object_value(frame["setup"])["model"]) for frame in frames if "setup" in frame)
def _authority(gateway: Gateway) -> str:
return _upstream_url(gateway).removeprefix("http://")
def _spend_rows(key: str, count: int) -> list[dict[str, JsonValue]]:
return eventually(
lambda: read_rows(SPEND_SQL, (sha256(key.encode()).hexdigest(),)),
lambda rows: len(rows) == count,
seconds=70,
)
def _billed(rows: list[dict[str, JsonValue]]) -> list[tuple[JsonValue, JsonValue, JsonValue]]:
return [(row["call_type"], row["prompt_tokens"], row["completion_tokens"]) for row in rows]
def _health(gateway: Gateway, model: str) -> httpx.Response:
return gateway.request("GET", "/health", params={"model": model})
def _probed_count(response: httpx.Response) -> int:
body: Final = JSON_OBJECT.validate_json(response.content)
return _integer(body.get("healthy_count", 0)) + _integer(body.get("unhealthy_count", 0))
def _converged_health(gateway: Gateway, model: str) -> httpx.Response:
return eventually(
lambda: _health(gateway, model), lambda response: _probed_count(response) == 1, seconds=CONVERGENCE_SECONDS
)
def _assert_scripted_turn_reached_upstream(gateway: Gateway, project: str, location: str, session: Session) -> None:
assert session.completed_turn(project), session
observed: Final = _for_project(_observations(gateway), project)
assert _upgrade_hosts(observed) == (_authority(gateway),), observed
assert _setup_models(observed) == (_model_path(project, location),), observed
@pytest.mark.parametrize("location", LOCATIONS)
def test_api_base_override_completes_a_turn_for_every_location(gateway: Gateway, location: str) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(
gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway)
)
session: Final = _run(gateway, model, key)
_assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, location, session)
assert _billed(_spend_rows(key, 1)) == [("_arealtime", INPUT_TOKENS, OUTPUT_TOKENS)]
@pytest.mark.parametrize("location", MALFORMED_LOCATIONS)
def test_malformed_location_is_refused_by_the_validator_before_any_dial(gateway: Gateway, location: str) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
malformed: Final = _deployment(gateway, scenario, handle.scenario_id, location=location, api_base=None)
refused: Final = _run(gateway, malformed, key)
assert refused.types == ("error", "closed"), refused
assert refused.error_messages == (INVALID_LOCATION,), refused
assert refused.close_code == INTERNAL_CLOSE, refused
model: Final = _deployment(
gateway, scenario, handle.scenario_id, location="us", api_base=_upstream_url(gateway)
)
session: Final = _run(gateway, model, key)
_assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, "us", session)
@pytest.mark.parametrize("location", [None, ""], ids=["omitted", "empty"])
def test_missing_location_defaults_to_us_central1(gateway: Gateway, location: str | None) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(
gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway)
)
session: Final = _run(gateway, model, key)
_assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, DEFAULT_LOCATION, session)
@pytest.mark.parametrize("location", ["us", "eu"])
def test_health_check_handshakes_with_the_override_host(gateway: Gateway, location: str) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
model: Final = _deployment(
gateway, scenario, handle.scenario_id, location=location, api_base=_upstream_url(gateway)
)
health: Final = _converged_health(gateway, model)
assert health.status_code == 200, health.text
body: Final = JSON_OBJECT.validate_json(health.content)
assert (body["healthy_count"], body["unhealthy_count"]) == (1, 0), health.text
observed: Final = _for_project(_observations(gateway), handle.scenario_id)
assert _upgrade_hosts(observed) == (_authority(gateway),), observed
assert _setup_models(observed) == (), observed
def test_health_check_reports_a_malformed_location_without_dialing(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
model: Final = _deployment(gateway, scenario, handle.scenario_id, location="US", api_base=None)
health: Final = _converged_health(gateway, model)
assert health.status_code == 503, health.text
body: Final = JSON_OBJECT.validate_json(health.content)
assert (body["healthy_count"], body["unhealthy_count"]) == (0, 1), health.text
unhealthy: Final = body["unhealthy_endpoints"]
assert isinstance(unhealthy, list) and len(unhealthy) == 1, health.text
assert INVALID_LOCATION in string_value(object_value(unhealthy[0])["error"]), health.text
assert _for_project(_observations(gateway), handle.scenario_id) == (), "the upstream saw a dial"
def test_upstream_handshake_refusal_reaches_the_client_and_the_next_session_connects(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
unregistered: Final = f"vertex-live-unknown-{uuid.uuid4().hex[:12]}"
refused_model: Final = _deployment(
gateway, scenario, unregistered, location="us", api_base=_upstream_url(gateway)
)
refused: Final = _run(gateway, refused_model, key)
assert refused.types == ("error", "closed"), refused
assert refused.error_messages == (HANDSHAKE_REFUSED,), refused
assert refused.close_code == REFUSAL_CLOSE, refused
assert _upgrade_hosts(_for_project(_observations(gateway), unregistered)) == (_authority(gateway),)
handle: Final = _scripted(scenario)
model: Final = _deployment(
gateway, scenario, handle.scenario_id, location="us", api_base=_upstream_url(gateway)
)
session: Final = _run(gateway, model, key)
_assert_scripted_turn_reached_upstream(gateway, handle.scenario_id, "us", session)
chat: Final = gateway.chat(scenario.model(), key=key)
assert string_value(chat["object"]) == "chat.completion", chat
async def _burst(ws_base: str, models: tuple[str, ...], key: str) -> tuple[Session, ...]:
return tuple(await asyncio.gather(*(_session(ws_base, model, key) for model in models)))
def test_concurrent_sessions_across_locations_each_reach_the_upstream_once(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
key: Final = scenario.key()
handles: Final = {location: _scripted(scenario) for location in LOCATIONS}
models: Final = {
location: _deployment(
gateway, scenario, handles[location].scenario_id, location=location, api_base=_upstream_url(gateway)
)
for location in LOCATIONS
}
order: Final = tuple(chain.from_iterable(repeat(location, SESSIONS_PER_LOCATION) for location in LOCATIONS))
sessions: Final = asyncio.run(
_burst(_ws_base(_proxy_url(gateway)), tuple(models[location] for location in order), key)
)
assert all(
session.completed_turn(handles[location].scenario_id) for location, session in zip(order, sessions)
), sessions
assert len({session.response_ids[0] for session in sessions}) == len(order), sessions
observed: Final = _observations(gateway)
for location, handle in handles.items():
mine: Final = _for_project(observed, handle.scenario_id)
assert _upgrade_hosts(mine) == (_authority(gateway),) * SESSIONS_PER_LOCATION, mine
assert _setup_models(mine) == (_model_path(handle.scenario_id, location),) * SESSIONS_PER_LOCATION, mine
assert _billed(_spend_rows(key, len(order))) == [("_arealtime", INPUT_TOKENS, OUTPUT_TOKENS)] * len(order)
async def _frames_until_closed(socket: ClientConnection) -> AsyncIterator[dict[str, JsonValue]]:
try:
async for message in socket:
yield JSON_OBJECT.validate_json(message)
except ConnectionClosed as closed:
yield {"type": "closed", "code": _close_code(closed)}
async def _hold_until_closed(ws_base: str, model: str, key: str, opened: asyncio.Queue[str]) -> Session:
headers: Final = {"Authorization": f"Bearer {key}"}
async with websockets.connect(f"{ws_base}/v1/realtime?model={model}", additional_headers=headers) as socket:
created: Final = JSON_OBJECT.validate_json(await socket.recv())
assert created.get("type") == "session.created", created
await opened.put(string_value(object_value(created["session"])["id"]))
return Session(tuple([frame async for frame in _frames_until_closed(socket)]))
def _relays_the_upstream_close(session: Session) -> bool:
return f"upstream websocket closed with code {session.close_code}" in session.error_messages[0]
async def _drain(opened: asyncio.Queue[str], count: int) -> tuple[str, ...]:
return tuple([await opened.get() for _ in range(count)])
async def _burst_through_outage(
ws_base: str, proxy_url: str, model: str, key: str, stop_upstream: Callable[[], None]
) -> tuple[Session, ...]:
opened: Final[asyncio.Queue[str]] = asyncio.Queue()
holders: Final = tuple(
asyncio.ensure_future(_hold_until_closed(ws_base, model, key, opened)) for _ in range(OUTAGE_BURST)
)
opened_sessions: Final = await asyncio.wait_for(_drain(opened, OUTAGE_BURST), 60)
assert len(opened_sessions) == OUTAGE_BURST, opened_sessions
await asyncio.to_thread(stop_upstream)
async with httpx.AsyncClient(base_url=proxy_url, timeout=15, trust_env=False) as client:
liveliness: Final = await client.get("/health/liveliness")
assert liveliness.status_code == 200, liveliness.text
return tuple(await asyncio.wait_for(asyncio.gather(*holders), 90))
@pytest.mark.timeout(240)
def test_upstream_outage_closes_every_open_session_and_recovers(
gateway: Gateway, tmp_path: Path, record_property: RecordProperty
) -> None:
with gateway.scenario() as scenario, owned_upstream(tmp_path) as slot:
project: Final = f"vertex-live-outage-{uuid.uuid4().hex[:12]}"
register_scenario(project, _scripted_turn(), control_url=slot.url)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, project, location="us", api_base=slot.url)
held: Final = asyncio.run(
_burst_through_outage(_ws_base(_proxy_url(gateway)), _proxy_url(gateway), model, key, slot.stop)
)
record_property("close_codes_during_upstream_outage", sorted(session.close_code or 0 for session in held))
assert [session.types for session in held] == [("error", "closed")] * OUTAGE_BURST, held
assert all(_relays_the_upstream_close(session) for session in held), held
assert len({session.close_code for session in held}) == 1, held
slot.start()
register_scenario(project, _scripted_turn(), control_url=slot.url)
recovered: Final = _run(gateway, model, key)
assert recovered.completed_turn(project), recovered
rows: Final = _spend_rows(key, OUTAGE_BURST + 1)
assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows

View file

@ -24,7 +24,9 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
MaskedHTTPStatusError,
get_httpx_client,
get_shared_realtime_ssl_context,
get_ssl_configuration,
realtime_ssl_for_url,
)
from litellm.types.llms.custom_http import VerifyTypes
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
@ -1493,7 +1495,6 @@ async def test_connection_error_retry_forwards_content(method: str):
await handler.close()
@pytest.fixture
def forward_proxy_server():
"""Plain HTTP forward proxy that records the absolute URIs it is asked to fetch."""
@ -1624,9 +1625,7 @@ def private_ca_tls_upstream(tmp_path: pathlib.Path):
ca_pem.write_bytes(cert.public_bytes(serialization.Encoding.PEM))
key_pem = tmp_path / "key.pem"
key_pem.write_bytes(
key.private_bytes(
serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()
)
key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption())
)
class OkTlsHandler(BaseHTTPRequestHandler):
@ -1801,8 +1800,11 @@ async def test_bounded_get_preserves_sdk_redirect_auth_and_query_handling(respx_
handler = AsyncHTTPHandler()
try:
response = await handler.get(
"https://example.com/spec.json?original=1", max_response_bytes=100, follow_redirects=True,
headers={"Authorization": "Bearer sentinel", "Accept-Encoding": "gzip"}, timeout=2.0,
"https://example.com/spec.json?original=1",
max_response_bytes=100,
follow_redirects=True,
headers={"Authorization": "Bearer sentinel", "Accept-Encoding": "gzip"},
timeout=2.0,
)
finally:
await handler.close()
@ -1888,6 +1890,7 @@ def _vcr_outcome_gate(request, vcr):
yield
record_vcr_outcome(request, vcr)
@pytest.fixture(scope="function")
def isolate_litellm_state():
"""
@ -1940,6 +1943,7 @@ def isolate_litellm_state():
setattr(litellm, attr, original_value)
_invalidate_model_cost_lowercase_map()
_SCALAR_DEFAULTS = {
"num_retries": getattr(litellm, "num_retries", None),
"num_retries_per_request": getattr(litellm, "num_retries_per_request", None),
@ -1960,6 +1964,7 @@ _SCALAR_DEFAULTS = {
"api_key": getattr(litellm, "api_key", None),
}
@pytest.fixture(scope="module")
def setup_and_teardown():
"""
@ -1982,12 +1987,14 @@ def setup_and_teardown():
litellm.in_memory_llm_clients_cache.flush_cache()
yield
_SERVER_DELAY_S = 5
_PER_REQUEST_TIMEOUT_S = 1.0
_CLIENT_DEFAULT_TIMEOUT_S = 60.0
class _SlowHandler(BaseHTTPRequestHandler):
def do_POST(self):
time.sleep(_SERVER_DELAY_S)
@ -2001,6 +2008,7 @@ class _SlowHandler(BaseHTTPRequestHandler):
def log_message(self, *args):
pass
@pytest.mark.usefixtures("_vcr_outcome_gate", "isolate_litellm_state", "setup_and_teardown")
def test_post_delay_exceeds_per_request_timeout_raises():
server = ThreadingHTTPServer(("127.0.0.1", 0), _SlowHandler)
@ -2022,3 +2030,24 @@ def test_post_delay_exceeds_per_request_timeout_raises():
handler.close()
server.shutdown()
server.server_close()
def test_realtime_ssl_for_url_sends_no_tls_argument_for_a_plain_ws_endpoint() -> None:
assert (
realtime_ssl_for_url("ws://127.0.0.1:8080/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent")
is None
)
def test_realtime_ssl_for_url_keeps_the_shared_context_for_wss_endpoints() -> None:
shared: Final = get_shared_realtime_ssl_context()
assert isinstance(shared, ssl.SSLContext)
assert realtime_ssl_for_url("wss://aiplatform.us.rep.googleapis.com/ws") is shared
def test_realtime_ssl_for_url_turns_ssl_verify_false_into_an_unverified_tls_context(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr("litellm.llms.custom_httpx.http_handler._shared_realtime_ssl_context", False)
selected: Final = realtime_ssl_for_url("wss://aiplatform.us.rep.googleapis.com/ws")
assert isinstance(selected, ssl.SSLContext)
assert selected.verify_mode == ssl.CERT_NONE
assert selected.check_hostname is False

View file

@ -2,7 +2,7 @@
Unit tests for VertexAIRealtimeConfig.
Validates:
- URL construction (regional and global)
- URL construction (regional, multi-region and global)
- Auth headers (Bearer token + project header)
- Session setup message format
- Full text-in / text-out round-trip via RealTimeStreaming with a mocked
@ -10,12 +10,14 @@ Validates:
"""
import json
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
import websockets.exceptions # registers websockets.exceptions on the websockets namespace
import litellm
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
# ---------------------------------------------------------------------------
@ -45,6 +47,26 @@ def test_get_complete_url_global():
)
@pytest.mark.parametrize("location", ["us", "eu"])
def test_get_complete_url_multi_region_uses_rep_host(location: str):
cfg: Final = VertexAIRealtimeConfig(access_token="tok", project="my-proj", location=location)
url: Final = cfg.get_complete_url(api_base=None, model="gemini-3.8-live")
# Google documents the multi-region Vertex endpoints as aiplatform.{us,eu}.rep.googleapis.com
# (https://docs.cloud.google.com/vertex-ai/generative-ai/docs/learn/locations, read 2026-10-07)
assert url == (
f"wss://aiplatform.{location}.rep.googleapis.com"
"/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
)
@pytest.mark.parametrize("location", ["us", "eu", "global", "us-central1", "europe-west4"])
def test_get_complete_url_host_matches_shared_vertex_host(location: str):
cfg: Final = VertexAIRealtimeConfig(access_token="tok", project="my-proj", location=location)
url: Final = cfg.get_complete_url(api_base=None, model="gemini-3.8-live")
shared_host: Final = get_vertex_base_url(location).removeprefix("https://")
assert url == f"wss://{shared_host}/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
def test_get_complete_url_custom_api_base():
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"

View file

@ -658,6 +658,47 @@ async def test_arealtime_drops_model_from_the_upstream_url_only_for_transcriptio
assert connect.url == expected_backend_url
async def _vertex_health_check_connect_for(monkeypatch: pytest.MonkeyPatch, api_base: str | None) -> _CapturingConnect:
async def fake_token_resolver(
credentials: object, project_id: str | None, custom_llm_provider: str
) -> tuple[str, str]:
return "access-token", project_id or ""
monkeypatch.setattr(realtime_main, "vertex_access_token_resolver", fake_token_resolver)
connect: Final = _CapturingConnect()
with patch("websockets.connect", connect):
assert await realtime_main._realtime_health_check(
model="gemini-3.8-live",
custom_llm_provider="vertex_ai",
api_key=None,
api_base=api_base,
model_params={
"vertex_project": "proj-1",
"vertex_credentials": "fake-credentials",
"vertex_location": "us",
},
)
return connect
@pytest.mark.asyncio
async def test_vertex_health_check_sends_no_tls_argument_to_a_plain_ws_api_base(
monkeypatch: pytest.MonkeyPatch,
) -> None:
connect: Final = await _vertex_health_check_connect_for(monkeypatch, "http://127.0.0.1:8080")
assert connect.url == "ws://127.0.0.1:8080/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
assert connect.kwargs["ssl"] is None
@pytest.mark.asyncio
async def test_vertex_health_check_keeps_tls_for_the_multi_region_host(monkeypatch: pytest.MonkeyPatch) -> None:
connect: Final = await _vertex_health_check_connect_for(monkeypatch, None)
assert connect.url == (
"wss://aiplatform.us.rep.googleapis.com/ws/google.cloud.aiplatform.v1.LlmBidiService/BidiGenerateContent"
)
assert connect.kwargs["ssl"] is not None
BLOCKED_PHRASE = "XSECRETBLOCKTESTPHRASEX"