diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 51f9f1e9fa4..3351d34b16f 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -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): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 2b86b9bbe5c..b57ce8ea2e8 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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 diff --git a/litellm/llms/vertex_ai/realtime/transformation.py b/litellm/llms/vertex_ai/realtime/transformation.py index 9fed6d52f0e..c14dd535145 100644 --- a/litellm/llms/vertex_ai/realtime/transformation.py +++ b/litellm/llms/vertex_ai/realtime/transformation.py @@ -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" # ------------------------------------------------------------------ diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 83cacf6c4d9..0e79cbfda13 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -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 diff --git a/tests/integration/_support/upstream.py b/tests/integration/_support/upstream.py index ae87768cd06..0c7dab90fcc 100644 --- a/tests/integration/_support/upstream.py +++ b/tests/integration/_support/upstream.py @@ -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" diff --git a/tests/integration/providers/test_vertex_realtime_multi_region_host.py b/tests/integration/providers/test_vertex_realtime_multi_region_host.py new file mode 100644 index 00000000000..df279aa2ced --- /dev/null +++ b/tests/integration/providers/test_vertex_realtime_multi_region_host.py @@ -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 diff --git a/tests/unit/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py index f5ae65e9898..b6b0db69162 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -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 diff --git a/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py b/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py index 14b3bdb48a1..d6986a2697a 100644 --- a/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py +++ b/tests/unit/llms/vertex_ai/realtime/test_vertex_ai_realtime_transformation.py @@ -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" diff --git a/tests/unit/realtime_api/test_main.py b/tests/unit/realtime_api/test_main.py index 25f116f5991..a692f38ec04 100644 --- a/tests/unit/realtime_api/test_main.py +++ b/tests/unit/realtime_api/test_main.py @@ -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"