diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index 5e9e04ae54f..47ef8369e56 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -438,7 +438,11 @@ class RealTimeStreaming: if not self._is_translation_session: return self._capture_translation_output_format(event_obj) - if event_obj.get("type") != "session.output_audio.delta": + if event_obj.get("type") not in ( + "session.output_audio.delta", + "response.output_audio.delta", + "response.audio.delta", + ): return delta: Final = event_obj.get("delta") if not isinstance(delta, str): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 052b9ee9e87..c8f2bb1af57 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -6469,17 +6469,7 @@ class BaseLLMHTTPHandler: cast_to=httpx.Response, body=request_data, ) - response_headers: Final = { # mutable-ok: httpx requires a concrete response-header mapping - key: value # mutable-ok: transport headers are materialized after filtering - for key, value in raw_response.headers.items() # mutable-ok: transport headers are materialized - if key.lower() not in ("content-encoding", "content-length", "transfer-encoding") - } - return httpx.Response( - status_code=raw_response.status_code, - headers=response_headers, - content=raw_response.content, - request=httpx.Request("POST", f"{normalized_api_base}/realtime/client_secrets"), - ) + return self._decoded_realtime_sdk_response(raw_response) finally: if owns_client: await openai_client.close() @@ -6564,11 +6554,12 @@ class BaseLLMHTTPHandler: key: str(value) for key, value in (extra_headers or {}).items() }, ) - return await configured_client.post( + raw_response: Final = await configured_client.post( "/realtime/translations/client_secrets", cast_to=httpx.Response, body=request_data, ) + return self._decoded_realtime_sdk_response(raw_response) finally: if owns_client: await openai_client.close() @@ -6706,11 +6697,12 @@ class BaseLLMHTTPHandler: }, }, ) - return await configured_client.post( + translation_response: Final = await configured_client.post( "/realtime/translations/calls", cast_to=httpx.Response, content=sdp_text.encode("utf-8"), ) + return self._decoded_realtime_sdk_response(translation_response) realtime_session_data: Final = cast( # cast-ok: endpoint validation produced an OpenAI realtime session RealtimeSessionCreateRequestParam, session_data, @@ -6724,16 +6716,25 @@ class BaseLLMHTTPHandler: extra_headers=sdk_extra_headers, timeout=timeout, ) - return httpx.Response( - status_code=raw_response.status_code, - headers=raw_response.headers, - content=raw_response.content, - request=httpx.Request("POST", f"{normalized_api_base}/realtime/calls"), - ) + return self._decoded_realtime_sdk_response(raw_response.http_response) finally: if owns_client: await openai_client.close() + @staticmethod + def _decoded_realtime_sdk_response(response: httpx.Response) -> httpx.Response: + headers: Final = { # mutable-ok: httpx accepts a concrete response header mapping + key: value + for key, value in response.headers.items() + if key.lower() not in ("content-encoding", "content-length", "transfer-encoding") + } + return httpx.Response( + status_code=response.status_code, + headers=headers, + content=response.content, + request=response.request, + ) + @staticmethod def _get_realtime_async_http_client(client: object | None) -> AsyncHTTPHandler: if isinstance(client, AsyncHTTPHandler): diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index df2d9ca4c75..d3fc39f6f74 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -174,15 +174,14 @@ class OpenAIRealtime(OpenAIChatCompletion): ) -> AbstractAsyncContextManager[object]: import websockets - if realtime_mode == "translation" or client is None: + if realtime_mode == "translation" or not isinstance(client, AsyncOpenAI): return websockets.connect( url, additional_headers=headers, max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, ssl=ssl_config, + **({"open_timeout": timeout} if timeout is not None else {}), ) - if not isinstance(client, AsyncOpenAI): - raise TypeError("client must be an AsyncOpenAI instance") openai_client: Final = client model_query: Final = query_params.get("model") extra_query: Final = { # mutable-ok: OpenAI SDK accepts a mutable query-parameter mapping @@ -195,6 +194,8 @@ class OpenAIRealtime(OpenAIChatCompletion): extra_headers=headers, websocket_connection_options={ # mutable-ok: OpenAI SDK forwards a mutable options mapping "max_size": REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, + **({"ssl": ssl_config} if url.startswith("wss://") else {}), + **({"open_timeout": timeout} if timeout is not None else {}), }, max_retries=0, ) diff --git a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py index 38b27895a35..33493ffe6c1 100644 --- a/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py +++ b/tests/test_litellm/litellm_core_utils/test_realtime_streaming.py @@ -2910,7 +2910,8 @@ def test_store_message_skips_pydantic_for_unlogged_audio_delta(): assert streaming.messages == [] -def test_translation_audio_duration_is_finalized_once(): +@pytest.mark.parametrize("event_type", ["session.output_audio.delta", "response.output_audio.delta", "response.audio.delta"]) +def test_translation_audio_duration_is_finalized_once(event_type: str): import base64 streaming = RealTimeStreaming( @@ -2921,7 +2922,7 @@ def test_translation_audio_duration_is_finalized_once(): translation_session=True, ) payload = base64.b64encode(bytes(48000)).decode() - streaming._capture_translation_output_audio({"type": "session.output_audio.delta", "delta": payload}) + streaming._capture_translation_output_audio({"type": event_type, "delta": payload}) streaming._finalize_translation_usage() streaming._finalize_translation_usage() diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py index ebf93df7191..425298c3d9a 100644 --- a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py @@ -295,7 +295,7 @@ async def test_async_realtime_uses_max_size_parameter(): called_kwargs = sdk_client.realtime.connect.call_args.kwargs connection_options = called_kwargs["websocket_connection_options"] assert connection_options["max_size"] is REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES - assert "ssl" not in connection_options + assert connection_options["ssl"] is not None mock_realtime_streaming.assert_called_once() mock_streaming_instance.bidirectional_forward.assert_awaited_once() @@ -453,3 +453,22 @@ async def test_translation_websocket_uses_direct_transport(): connect.assert_called_once() assert connect.call_args.args[0] == expected_url assert streaming.call_args.kwargs["translation_session"] is True + + +@pytest.mark.parametrize("sdk_client", [True, False]) +def test_connection_manager_preserves_transport_settings(sdk_client: bool): + import ssl + from litellm.llms.openai.realtime.handler import OpenAIRealtime + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + client = make_realtime_sdk_client() if sdk_client else MagicMock(spec=AsyncHTTPHandler) + ssl_config = ssl.create_default_context() + with patch("websockets.connect") as connect: + OpenAIRealtime()._create_connection_manager( + api_base="https://example.com", api_key="test", model="gpt-realtime-2.1", + query_params={"model": "gpt-realtime-2.1"}, headers={}, timeout=7.0, + realtime_mode="realtime", ssl_config=ssl_config, client=client, url="wss://example.com/v1/realtime", + ) + options = client.realtime.connect.call_args.kwargs["websocket_connection_options"] if sdk_client else connect.call_args.kwargs + assert options["ssl"] is ssl_config + assert options["open_timeout"] == 7.0 diff --git a/tests/test_litellm/llms/openai/realtime/test_translation.py b/tests/test_litellm/llms/openai/realtime/test_translation.py index 6eb2a618b93..a329200b2f7 100644 --- a/tests/test_litellm/llms/openai/realtime/test_translation.py +++ b/tests/test_litellm/llms/openai/realtime/test_translation.py @@ -209,7 +209,7 @@ async def test_translation_client_secret_uses_openai_sdk_custom_post(): async def send_response(request: httpx.Request) -> httpx.Response: assert request.url.path == "/v1/realtime/translations/client_secrets" assert json.loads(request.content)["session"]["audio"]["output"]["language"] == "es" - return httpx.Response(200, json={"value": "ek_translation"}) + return httpx.Response(200, content=gzip.compress(b'{"value":"ek_translation"}'), headers={"content-encoding": "gzip"}) http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response)) openai_client = AsyncOpenAI(api_key="sk-test", base_url="https://example.com/v1", http_client=http_client) @@ -233,6 +233,8 @@ async def test_translation_client_secret_uses_openai_sdk_custom_post(): assert response.status_code == 200 assert response.json() == {"value": "ek_translation"} + assert "content-encoding" not in response.headers + assert int(response.headers["content-length"]) == len(response.content) @pytest.mark.asyncio @@ -242,7 +244,7 @@ async def test_translation_calls_use_openai_sdk_custom_post(): body = await request.aread() assert request.headers["content-type"] == "application/sdp" assert body == b"v=0\r\n" - return httpx.Response(201, content=b"v=0\r\n") + return httpx.Response(201, content=gzip.compress(b"v=0\r\n"), headers={"content-encoding": "gzip"}) http_client = httpx.AsyncClient(transport=httpx.MockTransport(send_response)) openai_client = AsyncOpenAI(api_key="ek_test", base_url="https://example.com/v1", http_client=http_client) @@ -263,3 +265,5 @@ async def test_translation_calls_use_openai_sdk_custom_post(): assert response.status_code == 201 assert response.text == "v=0\r\n" + assert "content-encoding" not in response.headers + assert int(response.headers["content-length"]) == len(response.content)