fix(openai/realtime): drop model from upstream URL for intent=transcription (#43854)

* fix(openai/realtime): drop model from upstream URL for intent=transcription

* fix(openai/realtime): keep model on the upstream URL for non-OpenAI hosts

A custom api_base on the openai provider can be a gateway that routes on
the model query param, so the transcription model drop now applies only
to OpenAI's own hosts. xAI's URL is unchanged again

* test(integration): audit realtime transcription upstream URL on a custom api_base

Twenty-two integration cells cover transcription and conversation sessions on every realtime route, the OpenAI SDK sync and async clients, malformed and duplicated query params, unauthenticated and refused upgrades, an unknown model, idempotent spend logging, a twenty-session burst, an upstream outage mid burst, a worker kill, and a proxy restart, each asserting the exact query pairs the upstream received. The scripted upstream now records every websocket upgrade as an observation

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
michelligabriele 2026-10-05 21:14:56 +02:00 • committed by GitHub
parent 45e7be1abc
commit 5dd77eb8ca
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 718 additions and 4 deletions

View file

@ -19,6 +19,7 @@ from ....litellm_core_utils.realtime_streaming import (
client_sent_openai_beta_realtime_header,
)
from ....llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
from ..common_utils import is_openai_backed_api_base
from ..openai import OpenAIChatCompletion
@ -84,18 +85,27 @@ class OpenAIRealtime(OpenAIChatCompletion):
def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str:
"""
Construct the backend websocket URL with all query parameters (including 'model').
Construct the backend websocket URL with the client's query parameters.
`model` is left out for `intent=transcription` on OpenAI's own hosts: OpenAI
reads `?model=` as selecting a conversation session and rejects transcription
sessions with `invalid_model`. The transcription model is applied to the session
instead (see `force_transcription_model`), mirroring the Azure GA handler. Any
other `api_base` keeps `model`, since an OpenAI-compatible gateway may route on it.
"""
from httpx import URL
drops_model: Final = query_params.get("intent") == "transcription" and is_openai_backed_api_base(api_base)
api_base = api_base.replace("https://", "wss://")
api_base = api_base.replace("http://", "ws://")
url = URL(api_base)
# Set the correct path
url = url.copy_with(path="/v1/realtime")
# Include all query parameters including 'model'
if query_params:
url = url.copy_with(params=query_params)
upstream_params: Final = tuple(
(key, value) for key, value in query_params.items() if not (drops_model and key == "model")
)
if upstream_params:
url = url.copy_with(params=upstream_params)
return str(url)
def _make_event_normalizer(self) -> RealtimeEventNormalizer | None:

View file

@ -385,6 +385,14 @@ class Provider:
async def realtime(self, websocket: WebSocket) -> None:
scenario_id: Final = websocket.headers.get("authorization", "").removeprefix("Bearer ")
self.observations.put(
Observation(
websocket.url.path,
websocket.headers.get("authorization", ""),
{"query": [[key, value] for key, value in websocket.query_params.multi_items()]},
"WEBSOCKET",
)
)
response: Final = self.scenario_store.get(scenario_id)
if not isinstance(response, RealtimeResponse):
await websocket.close(code=4404)

View file

@ -0,0 +1,644 @@
"""Upstream websocket URL for realtime sessions on a custom ``api_base``.
``model`` leaves the upstream URL only for ``intent=transcription`` on OpenAI's own hosts. Every row here
runs against the scripted upstream on 127.0.0.1, where that gate is off, so the rows pin the forwarding
that must not move: ``model`` and ``intent`` reach the upstream exactly as the proxy resolved them, on
every route alias, through the OpenAI SDK and raw websockets, under malformed query strings, and while
the upstream or a proxy worker dies.
"""
from __future__ import annotations
import asyncio
import json
import os
import uuid
from collections.abc import AsyncIterator, Callable
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import psutil
import pytest
import websockets
from openai import AsyncOpenAI, OpenAI
from pydantic import JsonValue
from websockets.asyncio.client import ClientConnection
from websockets.exceptions import ConnectionClosed, InvalidStatus
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_proxy_process, owned_upstream, stop_root_process
from tests.integration._support.upstream import 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)
TRANSCRIBE_MODEL: Final = "gpt-live-transcribe"
SECOND_TRANSCRIBE_MODEL: Final = "gpt-live-transcribe-second"
CONVERSATION_MODEL: Final = "gpt-realtime-2"
WHISPER_DEFAULT: Final = "gpt-realtime-whisper"
XAI_MODEL: Final = "grok-4-1-fast-non-reasoning"
TRANSCRIPTION: Final = "transcription"
FIVE_KB: Final = "x" * 5120
BURST: Final = 20
OPEN_SESSIONS: Final = 6
WORKERS: Final = 2
INPUT_TOKENS: Final = 7
OUTPUT_TOKENS: Final = 5
USAGE: Final[dict[str, JsonValue]] = {
"total_tokens": INPUT_TOKENS + OUTPUT_TOKENS,
"input_tokens": INPUT_TOKENS,
"output_tokens": OUTPUT_TOKENS,
"input_token_details": {"text_tokens": INPUT_TOKENS, "audio_tokens": 0, "cached_tokens": 0},
"output_token_details": {"text_tokens": OUTPUT_TOKENS, "audio_tokens": 0},
}
TRANSCRIPTION_PAIRS: Final = [["model", TRANSCRIBE_MODEL], ["intent", TRANSCRIPTION]]
TRANSCRIPTION_QUERY: Final = f"intent={TRANSCRIPTION}"
UNAUTHENTICATED_STATUS: Final = 403
UPSTREAM_REFUSAL_CLOSE: Final = 1008
UNKNOWN_MODEL_CLOSE: Final = 1011
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], ...]
refused: int | None
@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)
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() -> str:
return os.environ["INTEGRATION_PROXY_URL"].rstrip("/")
def _done_event() -> RealtimeResponse:
return RealtimeResponse(
content_type="application/x-realtime",
events=(
{
"type": "response.done",
"event_id": "evt_$UNIQUE_ID",
"response": {
"id": "resp_$UNIQUE_ID",
"object": "realtime.response",
"status": "completed",
"output": [],
"usage": USAGE,
},
},
),
)
def _scripted(scenario: Scenario) -> ScenarioHandle:
handle: Final = register_scenario(f"realtime-url-{uuid.uuid4().hex[:12]}", _done_event())
scenario.cleanups.callback(delete_scenario, handle)
return handle
def _deployment(
gateway: Gateway,
scenario: Scenario,
scenario_id: str,
*,
model: str = f"openai/{TRANSCRIBE_MODEL}",
api_base: str | None = None,
) -> str:
base: Final = gateway.upstream_url.rstrip("/") if api_base is None else api_base
return scenario.model(model=model, api_key=scenario_id, api_base=base)
def _named_deployment(creator: Gateway, scenario: Scenario, name: str, model: str, scenario_id: str) -> str:
created: Final = creator.post(
"/model/new",
{
"model_name": name,
"litellm_params": {"model": model, "api_key": scenario_id, "api_base": creator.upstream_url.rstrip("/")},
"model_info": {},
},
)
scenario.cleanups.callback(scenario.delete_model, string_value(object_value(created["model_info"])["id"]))
return name
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(json.dumps({"type": "response.create"}))
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, path: str, query: str, key: str | None, *, turns: int = 0) -> Session:
headers: Final = {} if key is None else {"Authorization": f"Bearer {key}"}
try:
async with websockets.connect(f"{ws_base}{path}?{query}", additional_headers=headers) as socket:
return Session(await asyncio.wait_for(_collect(socket, turns), 60), None)
except InvalidStatus as refusal:
return Session((), refusal.response.status_code)
def _run(path: str, query: str, key: str | None, *, turns: int = 0, ws_base: str | None = None) -> Session:
return asyncio.run(_session(_ws_base(_proxy_url()) if ws_base is None else ws_base, path, query, key, turns=turns))
def _upgrades(gateway: Gateway, scenario_id: str) -> tuple[dict[str, JsonValue], ...]:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
observed: Final = tuple(map(object_value, upstream.get("/__observations").json()["requests"]))
return tuple(
request
for request in observed
if request["method"] == "WEBSOCKET" and request["authorization"] == f"Bearer {scenario_id}"
)
def _pairs(upgrade: dict[str, JsonValue]) -> JsonValue:
return object_value(upgrade["body"])["query"]
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 _assert_transcription_reached_upstream(gateway: Gateway, scenario_id: str, session: Session) -> None:
assert session.types == ("session.created",), session
assert session.session_model == TRANSCRIBE_MODEL, session
upgrades: Final = _upgrades(gateway, scenario_id)
assert [upgrade["path"] for upgrade in upgrades] == ["/v1/realtime"], upgrades
assert [_pairs(upgrade) for upgrade in upgrades] == [TRANSCRIPTION_PAIRS], upgrades
@pytest.mark.parametrize("path", ["/v1/realtime", "/realtime", "/openai/v1/realtime"])
def test_transcription_session_forwards_model_and_intent_to_a_custom_api_base(gateway: Gateway, path: str) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id)
session: Final = _run(path, f"model={model}&{TRANSCRIPTION_QUERY}", key)
_assert_transcription_reached_upstream(gateway, handle.scenario_id, session)
rows: Final = _spend_rows(key, 1)
assert rows[0]["call_type"] == "_arealtime", rows
def test_openai_sdk_async_transcription_session_reaches_a_custom_api_base(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id)
async def connect() -> dict[str, JsonValue]:
client: Final = AsyncOpenAI(
api_key=key, base_url=f"{_proxy_url()}/v1", websocket_base_url=f"{_ws_base(_proxy_url())}/v1"
)
async with client.realtime.connect(model=model, extra_query={"intent": TRANSCRIPTION}) as connection:
return JSON_OBJECT.validate_python((await connection.recv()).model_dump())
created: Final = asyncio.run(connect())
_assert_transcription_reached_upstream(gateway, handle.scenario_id, Session((created,), None))
def test_openai_sdk_sync_transcription_session_reaches_a_custom_api_base(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id)
client: Final = OpenAI(
api_key=key, base_url=f"{_proxy_url()}/v1", websocket_base_url=f"{_ws_base(_proxy_url())}/v1"
)
with client.realtime.connect(model=model, extra_query={"intent": TRANSCRIPTION}) as connection:
created: Final = JSON_OBJECT.validate_python(connection.recv().model_dump())
_assert_transcription_reached_upstream(gateway, handle.scenario_id, Session((created,), None))
def test_conversation_session_forwards_only_model_and_bills_the_scripted_usage(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id, model=f"openai/{CONVERSATION_MODEL}")
session: Final = _run("/v1/realtime", f"model={model}", key, turns=1)
assert session.types == ("session.created", "response.done"), session
assert session.session_model == CONVERSATION_MODEL, session
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert [_pairs(upgrade) for upgrade in upgrades] == [[["model", CONVERSATION_MODEL]]], upgrades
rows: Final = _spend_rows(key, 1)
assert rows[0]["call_type"] == "_arealtime", rows
assert rows[0]["prompt_tokens"] == INPUT_TOKENS, rows
assert rows[0]["completion_tokens"] == OUTPUT_TOKENS, rows
def test_intent_without_model_routes_to_the_whisper_default_and_forwards_intent_only(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
_named_deployment(gateway, scenario, WHISPER_DEFAULT, f"openai/{WHISPER_DEFAULT}", handle.scenario_id)
session: Final = _run("/v1/realtime", TRANSCRIPTION_QUERY, key)
assert session.types == ("session.created",), session
assert session.session_model == "", session
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert [_pairs(upgrade) for upgrade in upgrades] == [[["intent", TRANSCRIPTION]]], upgrades
def test_xai_transcription_session_keeps_model_on_a_custom_api_base(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id, model=f"xai/{XAI_MODEL}")
session: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", key)
assert session.types == ("session.created",), session
assert session.session_model == XAI_MODEL, session
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert [_pairs(upgrade) for upgrade in upgrades] == [[["model", XAI_MODEL], ["intent", TRANSCRIPTION]]], (
upgrades
)
def test_api_base_with_a_version_path_still_upgrades_at_v1_realtime(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id, api_base=f"{gateway.upstream_url}/v1")
session: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", key)
_assert_transcription_reached_upstream(gateway, handle.scenario_id, session)
@pytest.mark.parametrize(
("intent_query", "forwarded_intent"),
[
(f"intent={TRANSCRIPTION}&intent={TRANSCRIPTION}", TRANSCRIPTION),
("intent=", ""),
("intent=1", "1"),
(f"intent={FIVE_KB}", FIVE_KB),
(f"intent={TRANSCRIPTION}&intent=other", "other"),
],
ids=["twice_same_value", "empty", "integer", "five_kilobytes", "two_values"],
)
def test_malformed_intent_is_forwarded_as_the_proxy_resolved_it(
gateway: Gateway, intent_query: str, forwarded_intent: str
) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id)
session: Final = _run("/v1/realtime", f"model={model}&{intent_query}", key)
assert session.types == ("session.created",), session
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert [_pairs(upgrade) for upgrade in upgrades] == [
[["model", TRANSCRIBE_MODEL], ["intent", forwarded_intent]]
]
def test_model_given_twice_resolves_to_the_last_deployment(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
first: Final = _deployment(gateway, scenario, handle.scenario_id)
second: Final = _deployment(gateway, scenario, handle.scenario_id, model=f"openai/{SECOND_TRANSCRIBE_MODEL}")
session: Final = _run("/v1/realtime", f"model={first}&model={second}&{TRANSCRIPTION_QUERY}", key)
assert session.types == ("session.created",), session
assert session.session_model == SECOND_TRANSCRIBE_MODEL, session
upgrades: Final = _upgrades(gateway, handle.scenario_id)
expected: Final = [[["model", SECOND_TRANSCRIBE_MODEL], ["intent", TRANSCRIPTION]]]
assert [_pairs(upgrade) for upgrade in upgrades] == expected, upgrades
def test_unauthenticated_upgrade_is_refused_and_the_next_key_still_connects(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
model: Final = _deployment(gateway, scenario, handle.scenario_id)
refused: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", None)
assert refused.refused == UNAUTHENTICATED_STATUS, refused
assert _upgrades(gateway, handle.scenario_id) == (), "the upstream saw an unauthenticated upgrade"
session: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", scenario.key())
_assert_transcription_reached_upstream(gateway, handle.scenario_id, session)
def test_upstream_handshake_refusal_reaches_the_client_and_the_next_key_still_connects(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
unknown: Final = f"realtime-url-unknown-{uuid.uuid4().hex[:12]}"
key: Final = scenario.key()
refused_model: Final = _deployment(gateway, scenario, unknown)
refused: Final = _run("/v1/realtime", f"model={refused_model}&{TRANSCRIPTION_QUERY}", key)
assert refused.types == ("error", "closed"), refused
assert refused.close_code == UPSTREAM_REFUSAL_CLOSE, refused
assert [_pairs(upgrade) for upgrade in _upgrades(gateway, unknown)] == [TRANSCRIPTION_PAIRS]
handle: Final = _scripted(scenario)
model: Final = _deployment(gateway, scenario, handle.scenario_id)
session: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", scenario.key())
_assert_transcription_reached_upstream(gateway, handle.scenario_id, session)
def test_unknown_model_is_rejected_before_any_upstream_upgrade(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id)
missing: Final = f"realtime-url-missing-{uuid.uuid4().hex[:12]}"
rejected: Final = _run("/v1/realtime", f"model={missing}&{TRANSCRIPTION_QUERY}", key)
assert rejected.types == ("error", "closed"), rejected
assert rejected.close_code == UNKNOWN_MODEL_CLOSE, rejected
assert "Invalid model" in string_value(object_value(rejected.events[0]["error"])["message"]), rejected
session: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", key)
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert [_pairs(upgrade) for upgrade in upgrades] == [TRANSCRIPTION_PAIRS], upgrades
assert session.types == ("session.created",), session
def test_repeated_conversation_sessions_write_one_spend_row_each(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id, model=f"openai/{CONVERSATION_MODEL}")
sessions: Final = tuple(_run("/v1/realtime", f"model={model}", key, turns=1) for _ in range(3))
assert [session.types for session in sessions] == [("session.created", "response.done")] * 3, sessions
response_ids: Final = frozenset(session.response_ids[0] for session in sessions)
assert len(response_ids) == 3, sessions
rows: Final = _spend_rows(key, 3)
assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows
async def _burst(ws_base: str, model: str, key: str) -> tuple[Session, ...]:
query: Final = f"model={model}&{TRANSCRIPTION_QUERY}"
return tuple(await asyncio.gather(*(_session(ws_base, "/v1/realtime", query, key) for _ in range(BURST))))
def test_concurrent_transcription_sessions_each_reach_the_upstream_with_model_and_intent(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, handle.scenario_id)
sessions: Final = asyncio.run(_burst(_ws_base(_proxy_url()), model, key))
assert [session.types for session in sessions] == [("session.created",)] * BURST, sessions
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert [_pairs(upgrade) for upgrade in upgrades] == [TRANSCRIPTION_PAIRS] * BURST, upgrades
rows: Final = _spend_rows(key, BURST)
assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows
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, query: str, key: str, opened: asyncio.Queue[str]) -> Session:
async with websockets.connect(
f"{ws_base}/v1/realtime?{query}", additional_headers={"Authorization": f"Bearer {key}"}
) 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)]), None)
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()
query: Final = f"model={model}&{TRANSCRIPTION_QUERY}"
holders: Final = tuple(asyncio.ensure_future(_hold_until_closed(ws_base, query, key, opened)) for _ in range(BURST))
opened_sessions: Final = await asyncio.wait_for(_drain(opened, BURST), 60)
assert len(opened_sessions) == 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_transcription_session_and_recovers(
gateway: Gateway, tmp_path: Path, record_property: RecordProperty
) -> None:
with gateway.scenario() as scenario, owned_upstream(tmp_path) as slot:
scenario_id: Final = f"realtime-url-outage-{uuid.uuid4().hex[:12]}"
register_scenario(scenario_id, _done_event(), control_url=slot.url)
key: Final = scenario.key()
model: Final = _deployment(gateway, scenario, scenario_id, api_base=slot.url)
held: Final = asyncio.run(_burst_through_outage(_ws_base(_proxy_url()), _proxy_url(), 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")] * 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(scenario_id, _done_event(), control_url=slot.url)
recovered: Final = _run("/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", key)
assert recovered.types == ("session.created",), recovered
assert recovered.session_model == TRANSCRIBE_MODEL, recovered
rows: Final = _spend_rows(key, BURST + 1)
assert {str(row["call_type"]) for row in rows} == {"_arealtime"}, rows
def _spawned_worker(child: psutil.Process) -> bool:
try:
return any("multiprocessing.spawn" in part for part in child.cmdline())
except (psutil.NoSuchProcess, psutil.AccessDenied):
return False
def _workers(root: psutil.Process) -> tuple[psutil.Process, ...]:
return tuple(
sorted((child for child in root.children() if _spawned_worker(child)), key=lambda process: process.pid)
)
@dataclass(frozen=True, slots=True)
class KillOutcome:
served: int
closed: int
killed_pid: int
async def _one_turn_or_close(socket: ClientConnection) -> bool:
try:
await socket.send(json.dumps({"type": "response.create"}))
async for message in socket:
if JSON_OBJECT.validate_json(message).get("type") == "response.done":
return True
except ConnectionClosed:
return False
raise AssertionError("session ended without response.done or a close frame")
async def _first_frames(sockets: tuple[ClientConnection, ...]) -> tuple[dict[str, JsonValue], ...]:
return tuple([JSON_OBJECT.validate_json(await socket.recv()) for socket in sockets])
async def _sessions_through_worker_kill(ws_base: str, model: str, key: str, root: psutil.Process) -> KillOutcome:
query: Final = f"model={model}&{TRANSCRIPTION_QUERY}"
headers: Final = {"Authorization": f"Bearer {key}"}
sockets: Final = tuple(
[
await websockets.connect(f"{ws_base}/v1/realtime?{query}", additional_headers=headers)
for _ in range(OPEN_SESSIONS)
]
)
try:
created: Final = await asyncio.wait_for(_first_frames(sockets), 60)
assert [event.get("type") for event in created] == ["session.created"] * OPEN_SESSIONS, created
workers: Final = _workers(root)
assert len(workers) == WORKERS, [process.pid for process in workers]
victim: Final = workers[0]
victim.kill()
await asyncio.to_thread(victim.wait, 10)
served: Final = await asyncio.wait_for(asyncio.gather(*(_one_turn_or_close(socket) for socket in sockets)), 60)
return KillOutcome(served.count(True), served.count(False), victim.pid)
finally:
await asyncio.gather(*(socket.close() for socket in sockets))
async def _open_sessions(ws_base: str, model: str, key: str) -> tuple[ClientConnection, ...]:
query: Final = f"model={model}&{TRANSCRIPTION_QUERY}"
headers: Final = {"Authorization": f"Bearer {key}"}
sockets: Final = tuple(
[
await websockets.connect(f"{ws_base}/v1/realtime?{query}", additional_headers=headers)
for _ in range(OPEN_SESSIONS)
]
)
created: Final = await asyncio.wait_for(_first_frames(sockets), 60)
assert [event.get("type") for event in created] == ["session.created"] * OPEN_SESSIONS, created
return sockets
async def _await_closes(sockets: tuple[ClientConnection, ...]) -> tuple[int, ...]:
async def one(socket: ClientConnection) -> int:
try:
unexpected: Final = await socket.recv()
except ConnectionClosed as closed:
return _close_code(closed)
raise AssertionError(f"the proxy shutdown did not close the session: {unexpected!r}")
return tuple(await asyncio.wait_for(asyncio.gather(*(one(socket) for socket in sockets)), 90))
async def _sessions_through_proxy_shutdown(
ws_base: str, model: str, key: str, shutdown: Callable[[], bool]
) -> tuple[int, ...]:
sockets: Final = await _open_sessions(ws_base, model, key)
stopped: Final = await asyncio.to_thread(shutdown)
assert stopped, "the owned proxy root did not stop within the graceful window"
return await _await_closes(sockets)
@pytest.mark.timeout(480)
def test_worker_kill_then_proxy_restart_keep_transcription_sessions_serving(
gateway: Gateway, tmp_path: Path, record_property: RecordProperty
) -> None:
overrides: Final = {"DATABASE_URL": os.environ["DATABASE_URL"]}
with gateway.scenario() as scenario:
handle: Final = _scripted(scenario)
key: Final = scenario.key()
with owned_proxy_process(gateway, tmp_path, overrides, workers=WORKERS) as owned:
owned_url: Final = str(owned.gateway.client.base_url).rstrip("/")
model: Final = _named_deployment(
owned.gateway,
scenario,
f"realtime-url-owned-{uuid.uuid4().hex[:8]}",
f"openai/{TRANSCRIBE_MODEL}",
handle.scenario_id,
)
root: Final = psutil.Process(owned.process.pid)
outcome: Final = asyncio.run(_sessions_through_worker_kill(_ws_base(owned_url), model, key, root))
record_property(
"worker_kill", {"served": outcome.served, "closed": outcome.closed, "killed": outcome.killed_pid}
)
assert outcome.served + outcome.closed == OPEN_SESSIONS, outcome
with httpx.Client(base_url=owned_url, timeout=15, trust_env=False) as fresh:
readiness: Final = fresh.get("/health/readiness")
assert readiness.status_code == 200, readiness.text
respawned: Final = eventually(
lambda: tuple(process.pid for process in _workers(root)),
lambda pids: len(pids) == WORKERS and outcome.killed_pid not in pids,
seconds=60,
)
record_property("worker_pids_after_respawn", respawned)
after_kill: Final = _run(
"/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", key, ws_base=_ws_base(owned_url)
)
assert after_kill.types == ("session.created",), after_kill
codes: Final = asyncio.run(
_sessions_through_proxy_shutdown(
_ws_base(owned_url), model, key, lambda: stop_root_process(owned.process)
)
)
record_property("close_codes_during_proxy_shutdown", sorted(codes))
assert len(codes) == OPEN_SESSIONS, codes
with owned_proxy_process(gateway, tmp_path, overrides, workers=WORKERS) as restarted:
restarted_url: Final = str(restarted.gateway.client.base_url).rstrip("/")
recovered: Final = _run(
"/v1/realtime", f"model={model}&{TRANSCRIPTION_QUERY}", key, ws_base=_ws_base(restarted_url)
)
assert recovered.types == ("session.created",), recovered
assert recovered.session_model == TRANSCRIBE_MODEL, recovered
upgrades: Final = _upgrades(gateway, handle.scenario_id)
assert len(upgrades) == OPEN_SESSIONS * 2 + 2, upgrades
assert {json.dumps(_pairs(upgrade)) for upgrade in upgrades} == {json.dumps(TRANSCRIPTION_PAIRS)}, upgrades

View file

@ -599,3 +599,55 @@ async def test_arealtime_openai_forwards_the_intent_query_param_to_the_upstream_
litellm_logging_obj=FakeLogging(),
)
assert connect.url == "wss://api.openai.com/v1/realtime?model=gpt-realtime&intent=chat"
@pytest.mark.parametrize(
("model", "api_base", "client_query_params", "expected_backend_url"),
[
(
"openai/gpt-live-transcribe",
"https://api.openai.com/",
{"model": "my-transcribe-alias", "intent": "transcription"},
"wss://api.openai.com/v1/realtime?intent=transcription",
),
(
"openai/gpt-live-transcribe",
"https://eu.api.openai.com/",
{"model": "my-transcribe-alias", "intent": "transcription"},
"wss://eu.api.openai.com/v1/realtime?intent=transcription",
),
(
"openai/gpt-live-transcribe",
"https://api.openai.com/",
{"model": "my-realtime-alias"},
"wss://api.openai.com/v1/realtime?model=gpt-live-transcribe",
),
(
"openai/gateway-transcribe-alias",
"https://gateway.example/",
{"model": "my-transcribe-alias", "intent": "transcription"},
"wss://gateway.example/v1/realtime?model=gateway-transcribe-alias&intent=transcription",
),
(
"xai/grok-voice",
"https://api.x.ai/v1",
{"model": "my-voice-alias", "intent": "transcription"},
"wss://api.x.ai/v1/realtime?model=grok-voice&intent=transcription",
),
],
)
@pytest.mark.asyncio
async def test_arealtime_drops_model_from_the_upstream_url_only_for_transcription_sessions_on_openai_hosts(
model, api_base, client_query_params, expected_backend_url
):
connect: Final = _ConnectThatStopsAfterCapturingTheUrl()
with patch("websockets.connect", connect):
await realtime_main._arealtime.__wrapped__(
model=model,
websocket=_ClosableGaClientWebSocket(),
api_base=api_base,
api_key="fake-key",
litellm_logging_obj=FakeLogging(),
query_params=client_query_params,
)
assert connect.url == expected_backend_url