mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(realtime): preserve transport settings and translation duration
This commit is contained in:
parent
ac4946477b
commit
a290cc5a7c
6 changed files with 58 additions and 28 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue