diff --git a/litellm/llms/chatgpt/realtime.py b/litellm/llms/chatgpt/realtime.py index 27ce6cb3966..559a6ef18f6 100644 --- a/litellm/llms/chatgpt/realtime.py +++ b/litellm/llms/chatgpt/realtime.py @@ -5,7 +5,6 @@ from typing import TYPE_CHECKING, Final from httpx import URL, QueryParams from pydantic import TypeAdapter -from websockets.exceptions import ConnectionClosed from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES from litellm.llms.openai.realtime.handler import OpenAIRealtime @@ -114,6 +113,8 @@ class ChatGPTRealtime(OpenAIRealtime): ) async def close_call(self, connection: "ClientConnection", model: str, api_base: str) -> None: + from websockets.exceptions import ConnectionClosed + if realtime_endpoint(model) == "live": try: await connection.send('{"type":"session.close"}') @@ -124,7 +125,8 @@ class ChatGPTRealtime(OpenAIRealtime): await self.hangup_call(api_base) async def hangup_call(self, api_base: str) -> None: - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + from litellm.types.utils import LlmProviders base: Final = URL(api_base) url: Final = base.copy_with( @@ -132,12 +134,9 @@ class ChatGPTRealtime(OpenAIRealtime): path=f"{base.path.rstrip('/')}/realtime/calls/{self._call_id}/hangup", params=tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")), ) - client: Final = AsyncHTTPHandler() - try: - response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) - response.raise_for_status() - finally: - await client.close() + client: Final = get_async_httpx_client(llm_provider=LlmProviders.CHATGPT) + response: Final = await client.post(str(url), headers=self._profile_headers, data=b"", timeout=10) + response.raise_for_status() @staticmethod def get_api_base(api_base: str | None = None) -> str: @@ -182,7 +181,9 @@ class ChatGPTRealtime(OpenAIRealtime): base.copy_with( scheme="wss" if base.scheme in ("https", "wss") else "ws", path=f"{base.path.rstrip('/')}/{endpoint}", - params=query_params, + params=QueryParams(TypeAdapter(Mapping[str, str | None]).validate_python(query_params)).merge( + tuple((key, value) for key, value in self._extra_query.items() if key not in ("model", "call_id")) + ), ) ) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 974f52b4268..44099efedeb 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6330,6 +6330,9 @@ class BaseLLMHTTPHandler: Uses provider_config (BaseRealtimeHTTPConfig) for URL construction and header auth when available; falls back to the legacy OpenAI-style defaults. """ + from litellm.llms.chatgpt.common_utils import without_oauth_identity_headers + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + if client is None or not isinstance(client, AsyncHTTPHandler): async_httpx_client = get_async_httpx_client( llm_provider=litellm.LlmProviders.OPENAI, @@ -6355,7 +6358,11 @@ class BaseLLMHTTPHandler: } if extra_headers: - headers.update(extra_headers) + headers.update( + without_oauth_identity_headers(extra_headers) + if isinstance(provider_config, ChatGPTRealtimeHTTPConfig) + else extra_headers + ) logging_obj.pre_call( input=request_data, diff --git a/litellm/proxy/realtime_endpoints/call_supervision.py b/litellm/proxy/realtime_endpoints/call_supervision.py index e750ce61d6a..ae04582a36b 100644 --- a/litellm/proxy/realtime_endpoints/call_supervision.py +++ b/litellm/proxy/realtime_endpoints/call_supervision.py @@ -120,12 +120,19 @@ class CallSupervisor: try: if not self._terminal: deadline: Final = asyncio.get_running_loop().time() + self._termination_timeout + primary_deadline: Final = ( + deadline - self._termination_timeout / 2 + if self._terminal_usage_required and self._force_close_call is not None + else deadline + ) try: - await asyncio.wait_for(self._close_call(), timeout=self._termination_timeout) + await asyncio.wait_for( + self._close_call(), timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time()) + ) self._close_confirmed = True except Exception: # noqa: BLE001 # provider exceptions can contain credentials verbose_proxy_logger.error("Realtime observer could not terminate upstream call") - await self._drain(reader, timeout=max(0.0, deadline - asyncio.get_running_loop().time())) + await self._drain(reader, timeout=max(0.0, primary_deadline - asyncio.get_running_loop().time())) if self._terminal_usage_required and not self._terminal and self._force_close_call is not None: remaining: Final = max(0.0, deadline - asyncio.get_running_loop().time()) try: diff --git a/tests/test_litellm/llms/chatgpt/test_realtime.py b/tests/test_litellm/llms/chatgpt/test_realtime.py index 6db166e956b..b02b3a551f2 100644 --- a/tests/test_litellm/llms/chatgpt/test_realtime.py +++ b/tests/test_litellm/llms/chatgpt/test_realtime.py @@ -1,4 +1,5 @@ import json +import sys from contextlib import nullcontext from types import SimpleNamespace from unittest.mock import AsyncMock, patch @@ -14,13 +15,14 @@ from litellm.types.router import GenericLiteLLMParams @pytest.mark.asyncio @pytest.mark.parametrize("failure", ["closed", "network"]) -@pytest.mark.parametrize( - "hangup_status, expectation", [(200, nullcontext()), (503, pytest.raises(httpx.HTTPStatusError))] -) -async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, expectation, chatgpt_tokens): +@pytest.mark.parametrize("hangup_status", [200, 503]) +async def test_live_closed_observer_uses_independent_hangup(failure, hangup_status, chatgpt_tokens, monkeypatch): from websockets.exceptions import ConnectionClosedOK from websockets.frames import Close + from litellm.caching.llm_caching_handler import LLMClientCache + + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", LLMClientCache()) handler = ChatGPTRealtime( GenericLiteLLMParams( chatgpt_realtime_call_id="rtc_live_closed", @@ -46,15 +48,21 @@ async def test_live_closed_observer_uses_independent_hangup(failure, hangup_stat return httpx.Response(hangup_status) client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) - with patch("httpx.AsyncClient", return_value=client): - with expectation: - await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1") - assert len(requests) == 1 + try: + with patch("httpx.AsyncClient", return_value=client) as create_client: + for _ in range(2): + with pytest.raises(httpx.HTTPStatusError) if hangup_status == 503 else nullcontext(): + await handler.close_call(connection, "gpt-live-1-codex", "https://gateway.example/v1") + assert not client.is_closed + create_client.assert_called_once() + finally: + await client.aclose() + assert len(requests) == 2 assert requests[0].method == "POST" assert str(requests[0].url) == "https://gateway.example/v1/realtime/calls/rtc_live_closed/hangup?gateway=tenant" assert requests[0].headers["x-gateway-token"] == "test-only" assert requests[0].headers["Authorization"] == "Bearer test-token-default" - assert client.is_closed + assert requests[0].extensions["timeout"]["read"] == 10 @pytest.mark.asyncio @@ -305,6 +313,63 @@ def test_realtime_uses_platform_endpoint_with_oauth_headers(model, endpoint, cha assert headers["openai-alpha"] == "quicksilver=v2" +@pytest.mark.parametrize("endpoint", ["live", "realtime"]) +def test_new_realtime_session_preserves_gateway_query(endpoint, chatgpt_tokens, local_model_cost_map): + model = "gpt-live-1-codex" if endpoint == "live" else "gpt-realtime-1.5" + handler = ChatGPTRealtime( + GenericLiteLLMParams( + chatgpt_token_dir=chatgpt_tokens, + chatgpt_realtime_client_query={"intent": "conversation", "architecture": "client-architecture"}, + extra_query={ + "gateway_token": "opaque +/& value", + "intent": "gateway-intent", + "architecture": "gateway-architecture", + "model": "other-model", + "call_id": "rtc_other", + }, + ), + {}, + ) + url = httpx.URL(handler._construct_url("https://gateway.example/v1", {"model": model, "intent": "query-intent"})) + assert url.path == f"/v1/{endpoint}" + assert dict(url.params) == { + "model": model, + "gateway_token": "opaque +/& value", + "intent": "gateway-intent", + "architecture": "gateway-architecture", + } + + +@pytest.mark.asyncio +async def test_openai_http_call_does_not_require_websockets(monkeypatch): + monkeypatch.delitem(sys.modules, "litellm.llms.chatgpt.realtime", raising=False) + for name in tuple(sys.modules): + if name == "websockets" or name.startswith("websockets."): + monkeypatch.delitem(sys.modules, name) + monkeypatch.setitem(sys.modules, "websockets", None) + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(201, text="v=0\r\n") + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + response = await litellm.arealtime_calls( + model="openai/gpt-realtime-1.5", + openai_ephemeral_key="test-only", + sdp_body=b"v=0\r\n", + api_key="test-only", + client=client, + ) + assert response.status_code == 201 + assert len(requests) == 1 + assert requests[0].url.path == "/v1/realtime/calls" + finally: + await client.client.aclose() + + @pytest.mark.parametrize("endpoint", ["live", "realtime"]) @pytest.mark.parametrize("call_id", [None, "rtc_metadata"]) def test_realtime_routes_new_models_using_registered_metadata(endpoint, call_id, chatgpt_tokens, local_model_cost_map): diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index cea2d439198..29a58b788f0 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -3620,3 +3620,62 @@ def test_image_edit_handler_keeps_the_sync_transform(): assert config.transform_calls == ["sync"] assert captured["body"] == {"transformed_by": "sync"} assert response.data[0].b64_json == "sync" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("endpoint", ["client_secrets", "transcription_sessions"]) +@pytest.mark.parametrize("provider", ["chatgpt", "openai"]) +@pytest.mark.parametrize("authorization_header", ["Authorization", "aUtHoRiZaTiOn"]) +async def test_realtime_http_sessions_preserve_provider_identity( + endpoint, provider, authorization_header, tmp_path, monkeypatch +): + import time + + from litellm.llms.chatgpt.realtime import ChatGPTRealtimeHTTPConfig + from litellm.llms.openai.realtime.http_transformation import OpenAIRealtimeHTTPConfig + from litellm.types.router import GenericLiteLLMParams + + monkeypatch.setenv("CHATGPT_TOKEN_DIR", str(tmp_path)) + monkeypatch.setenv("CHATGPT_AUTH_FILE", "auth.json") + (tmp_path / "auth.json").write_text( + json.dumps({"access_token": "test-resolved", "account_id": "test-selected", "expires_at": time.time() + 3600}) + ) + config = ChatGPTRealtimeHTTPConfig(GenericLiteLLMParams()) if provider == "chatgpt" else OpenAIRealtimeHTTPConfig() + requests = [] + + def respond(request): + requests.append(request) + return httpx.Response(200, json={"id": "session-test"}) + + client = AsyncHTTPHandler() + client.client = httpx.AsyncClient(transport=httpx.MockTransport(respond)) + try: + response = await BaseLLMHTTPHandler()._async_realtime_session_post( + endpoint=endpoint, + api_base="https://gateway.example/v1", + api_key="test-openai", + request_data={"session": {"model": "gpt-realtime-1.5"}}, + logging_obj=Mock(), + timeout=5, + provider_config=config, + model="gpt-realtime-1.5", + extra_headers={ + authorization_header: "Bearer test-override", + "CHATGPT-ACCOUNT-ID": "test-other-account", + "x-gateway-route": "required", + }, + client=client, + ) + assert response.status_code == 200 + assert not client.client.is_closed + finally: + await client.client.aclose() + assert len(requests) == 1 + assert requests[0].url.path == f"/v1/realtime/{endpoint}" + assert requests[0].headers["x-gateway-route"] == "required" + if provider == "chatgpt": + assert requests[0].headers.get_list("authorization") == ["Bearer test-resolved"] + assert requests[0].headers.get_list("chatgpt-account-id") == ["test-selected"] + else: + assert requests[0].headers.get_list("authorization")[-1] == "Bearer test-override" + assert requests[0].headers["chatgpt-account-id"] == "test-other-account" diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py index 7d0d614f31d..c14597a012a 100644 --- a/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_supervision.py @@ -30,6 +30,47 @@ class Socket: self.closed = True +@pytest.mark.asyncio +@pytest.mark.parametrize("stalled_step", ["close", "drain"]) +async def test_live_initial_close_reserves_time_for_independent_hangup(stalled_step): + socket = Socket() + logger = MagicMock(spec=Logging) + logger.model_call_details = {} + close_cancelled = asyncio.Event() + + async def close(): + if stalled_step == "close": + try: + await asyncio.Event().wait() + finally: + close_cancelled.set() + + async def force_close(): + await socket.messages.put({"type": "session.closed", "usage": {"audio_duration_ms": 1000}}) + + force = AsyncMock(side_effect=force_close) + sink = Sink(logger) + supervisor = CallSupervisor( + socket, + sink, + logger, + UserAPIKeyAuth(), + close, + force_close_call=force, + drain_timeout=1, + termination_timeout=0.08, + ) + await socket.messages.put({"type": "session.started"}) + await supervisor.start() + await asyncio.wait_for(supervisor.close(), timeout=0.5) + force.assert_awaited_once() + assert close_cancelled.is_set() == (stalled_step == "close") + assert any(event["type"] == "session.closed" for event in sink.events) + assert not logger.model_call_details.get("realtime_usage_incomplete") + assert sink.logs == 1 + assert socket.closed + + @pytest.mark.asyncio @pytest.mark.parametrize("fallback", ["terminal", "no_terminal", "timeout"]) async def test_live_unacknowledged_close_uses_bounded_independent_hangup(monkeypatch, fallback):