fix(realtime): preserve transport settings and translation duration

This commit is contained in:
Emerson Gomes 2026-09-15 13:03:12 -05:00
parent ac4946477b
commit a290cc5a7c
No known key found for this signature in database
GPG key ID: D3DF28AB5D1B5E17
6 changed files with 58 additions and 28 deletions

View file

@ -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):

View file

@ -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):

View file

@ -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,
)

View file

@ -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()

View file

@ -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

View file

@ -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)