mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): protect realtime identity and reserve hangup time
This commit is contained in:
parent
c360d187f5
commit
8ec259bdbf
6 changed files with 201 additions and 21 deletions
|
|
@ -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"))
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue