From 284cbc18ca0c69046b3bae6527fba00d3a743f97 Mon Sep 17 00:00:00 2001 From: jibanez-staticduo Date: Wed, 9 Sep 2026 08:43:31 +0200 Subject: [PATCH] refactor(chatgpt): keep HTTP handling inside the proxy --- litellm/llms/chatgpt/codex.py | 196 ++++------------ litellm/proxy/proxy_server.py | 4 +- .../proxy/realtime_endpoints/call_sessions.py | 156 ++++++++++++ litellm/proxy/realtime_endpoints/endpoints.py | 2 +- tests/test_litellm/llms/chatgpt/test_codex.py | 222 ++---------------- .../realtime_endpoints/test_call_sessions.py | 213 +++++++++++++++++ 6 files changed, 429 insertions(+), 364 deletions(-) create mode 100644 litellm/proxy/realtime_endpoints/call_sessions.py create mode 100644 tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py diff --git a/litellm/llms/chatgpt/codex.py b/litellm/llms/chatgpt/codex.py index 669b56d4371..ad801dcdcdb 100644 --- a/litellm/llms/chatgpt/codex.py +++ b/litellm/llms/chatgpt/codex.py @@ -1,21 +1,11 @@ -import base64 -import hashlib -import json -import time -from types import MappingProxyType +from collections.abc import Mapping from typing import Final from urllib.parse import urlsplit import httpx -from fastapi import HTTPException, Request, Response, WebSocket from pydantic import BaseModel, Field +from typing_extensions import ReadOnly, TypedDict -from litellm.llms.base_llm.chat.transformation import BaseLLMException -from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.proxy.auth.auth_checks import can_key_call_resolved_model -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper -from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation from litellm.types.realtime import RealtimeQueryParams, RealtimeSessionConfig @@ -36,153 +26,49 @@ class ChatGPTCallRouting(BaseModel): model: str -def encode_call(call: CodexRealtimeCall) -> str: - encrypted: Final = encrypt_value_helper(call.model_dump_json()) - return "rtc_litellm_" + base64.urlsafe_b64encode(encrypted.encode()).decode().rstrip("=") +class CodexSidebandRequest(TypedDict): + model: ReadOnly[str] + chatgpt_realtime_call_id: ReadOnly[str] + query_params: ReadOnly[RealtimeQueryParams] -def decode_call(token: str, authorization: str) -> CodexRealtimeCall: - try: - if not token.startswith("rtc_litellm_"): - raise ValueError("Invalid call prefix") - encoded: Final = token.removeprefix("rtc_litellm_") - encrypted: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True) - plaintext: Final = decrypt_value_helper(encrypted.decode(), key="codex_realtime_call") - call: Final = CodexRealtimeCall.model_validate_json(plaintext or "") - except (ValueError, TypeError, UnicodeError) as exc: - raise HTTPException(403, "Invalid realtime call") from exc - if call.expires_at < time.time() or call.owner != hashlib.sha256(authorization.encode()).hexdigest(): - raise HTTPException(403, "Invalid or expired realtime call") - return call +def build_call_request( + offer: CodexRealtimeOffer, query: Mapping[str, str], headers: Mapping[str, str] +) -> dict[str, object]: # mutable-ok: proxy processor enriches the request dictionary + return { # mutable-ok: proxy processor enriches the request dictionary + "model": offer.session.model, + "sdp_body": offer.sdp.encode(), + "session": offer.session.model_dump(exclude_none=True), + "openai_ephemeral_key": "", + "extra_query": { # mutable-ok: router request parameters + key: value for key, value in query.items() if key in ("intent", "architecture") + }, + "extra_headers": { # mutable-ok: router request headers + key: value + for key, value in headers.items() + if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") + }, + } -async def read_codex_offer(request: Request) -> CodexRealtimeOffer: - if request.headers.get("content-type", "").startswith("multipart/form-data"): - form: Final = await request.form() - return CodexRealtimeOffer.model_validate( - MappingProxyType({"sdp": form.get("sdp"), "session": json.loads(str(form.get("session", "{}")))}) - ) - return CodexRealtimeOffer.model_validate(await request.json()) - - -async def create_codex_realtime_call(request: Request) -> Response: - from litellm.proxy import proxy_server as server - from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing - - try: - offer: Final = await read_codex_offer(request) - except ValueError as exc: - raise HTTPException(400, "Invalid realtime offer: expected sdp and session") from exc - model: Final = offer.session.model - if not model: - raise HTTPException(400, "session.model is required") - auth: Final = await user_api_key_auth( - request=request, - api_key=request.headers.get("authorization", ""), - azure_api_key_header="", - anthropic_api_key_header=None, - google_ai_studio_api_key_header=None, - azure_apim_header=None, - custom_litellm_key_header=None, +def parse_call_response(response: httpx.Response, alias: str, owner: str, expires_at: float) -> CodexRealtimeCall: + routing_data: Final = response.extensions.get("chatgpt_realtime") + if not routing_data: + raise ValueError("Direct call signaling requires a ChatGPT deployment") + routing: Final = ChatGPTCallRouting.model_validate(routing_data) + call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1] + return CodexRealtimeCall( + call_id=call_id, + model=routing.model, + alias=alias, + owner=owner, + expires_at=expires_at, ) - try: - await can_key_call_resolved_model( - model=model, - llm_model_list=server.llm_model_list, - valid_token=auth, - llm_router=server.llm_router, - ) - data: Final = { # mutable-ok: proxy processor enriches request data - "model": model, - "sdp_body": offer.sdp.encode(), - "session": offer.session.model_dump(exclude_none=True), - "openai_ephemeral_key": "", - "extra_query": { # mutable-ok: router request parameters - key: value for key, value in request.query_params.items() if key in ("intent", "architecture") - }, - "extra_headers": { # mutable-ok: router request headers - key: value - for key, value in request.headers.items() - if key in ("openai-alpha", "openai-beta", "x-session-id", "x-oai-attestation") - }, - } - processor: Final = ProxyBaseLLMRequestProcessing(data=data) - processed, _ = await processor.common_processing_pre_call_logic( - request=request, - general_settings=server.general_settings, - user_api_key_dict=auth, - version=server.version, - proxy_logging_obj=server.proxy_logging_obj, - proxy_config=server.proxy_config, - user_model=server.user_model, - user_temperature=server.user_temperature, - user_request_timeout=server.user_request_timeout, - user_max_tokens=server.user_max_tokens, - user_api_base=server.user_api_base, - model=model, - route_type="arealtime_calls", - ) - result: Final = await server.route_request( - data=processed, - route_type="arealtime_calls", - llm_router=server.llm_router, - user_model=server.user_model, - ) - try: - response: Final = await result - except BaseLLMException as exc: - raise HTTPException(exc.status_code, str(exc)) from exc - if not isinstance(response, httpx.Response): - raise HTTPException(502, "Invalid realtime signaling response") - routing_data: Final = response.extensions.get("chatgpt_realtime") - if response.is_error: - return Response(response.content, status_code=response.status_code, media_type="application/json") - if not routing_data: - raise HTTPException(400, "Direct call signaling requires a ChatGPT deployment") - routing: Final = ChatGPTCallRouting.model_validate(routing_data) - call_id: Final = urlsplit(response.headers.get("location", "")).path.rstrip("/").rsplit("/", 1)[-1] - call: Final = CodexRealtimeCall( - call_id=call_id, - model=routing.model, - alias=model, - owner=hashlib.sha256(request.headers.get("authorization", "").encode()).hexdigest(), - expires_at=time.time() + 3600, - ) - token: Final = encode_call(call) - return Response( - response.content, - status_code=response.status_code, - media_type="application/sdp", - headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), - ) - finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) -async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None: - import litellm - from litellm.proxy import proxy_server as server - - try: - try: - call: Final = decode_call(token, websocket.headers.get("authorization", "")) - await can_key_call_resolved_model( - model=call.alias, - llm_model_list=server.llm_model_list, - valid_token=auth, - llm_router=server.llm_router, - ) - except (HTTPException, ProxyException): - await websocket.close(code=1008, reason="Invalid realtime call") - return - await websocket.accept() - query: Final[RealtimeQueryParams] = {"model": call.model} - await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # internal proxy entrypoint for an already authorized call - model=f"chatgpt/{call.model}", - websocket=websocket, - chatgpt_realtime_call_id=call.call_id, - query_params=query, - user_api_key_dict=auth, - ) - finally: - await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) +def build_sideband_request(call: CodexRealtimeCall) -> CodexSidebandRequest: + return { + "model": f"chatgpt/{call.model}", + "chatgpt_realtime_call_id": call.call_id, + "query_params": {"model": call.model}, + } diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9d55cf21d81..7fcf89904c2 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -11598,7 +11598,7 @@ async def codex_live_sideband_endpoint( call_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ) -> None: - from litellm.llms.chatgpt.codex import codex_realtime_sideband + from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband await codex_realtime_sideband(websocket, call_id, user_api_key_dict) @@ -11620,7 +11620,7 @@ async def realtime_websocket_endpoint( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth_websocket), ): if call_id is not None: - from litellm.llms.chatgpt.codex import codex_realtime_sideband + from litellm.proxy.realtime_endpoints.call_sessions import codex_realtime_sideband await codex_realtime_sideband(websocket, call_id, user_api_key_dict) return diff --git a/litellm/proxy/realtime_endpoints/call_sessions.py b/litellm/proxy/realtime_endpoints/call_sessions.py new file mode 100644 index 00000000000..202f919beb6 --- /dev/null +++ b/litellm/proxy/realtime_endpoints/call_sessions.py @@ -0,0 +1,156 @@ +import base64 +import hashlib +import json +import time +from types import MappingProxyType +from typing import Final + +import httpx +from fastapi import HTTPException, Request, Response, WebSocket + +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.chatgpt.codex import ( + CodexRealtimeCall, + CodexRealtimeOffer, + build_call_request, + build_sideband_request, + parse_call_response, +) +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import can_key_call_resolved_model +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.spend_tracking.budget_reservation import release_or_invalidate_budget_reservation + + +def encode_call(call: CodexRealtimeCall) -> str: + encrypted: Final = encrypt_value_helper(call.model_dump_json()) + return "rtc_litellm_" + base64.urlsafe_b64encode(encrypted.encode()).decode().rstrip("=") + + +def decode_call(token: str, authorization: str) -> CodexRealtimeCall: + try: + if not token.startswith("rtc_litellm_"): + raise ValueError("Invalid call prefix") + encoded: Final = token.removeprefix("rtc_litellm_") + encrypted: Final = base64.b64decode(encoded + "=" * (-len(encoded) % 4), altchars=b"-_", validate=True) + plaintext: Final = decrypt_value_helper(encrypted.decode(), key="codex_realtime_call") + call: Final = CodexRealtimeCall.model_validate_json(plaintext or "") + except (ValueError, TypeError, UnicodeError) as exc: + raise HTTPException(403, "Invalid realtime call") from exc + if call.expires_at < time.time() or call.owner != hashlib.sha256(authorization.encode()).hexdigest(): + raise HTTPException(403, "Invalid or expired realtime call") + return call + + +async def read_codex_offer(request: Request) -> CodexRealtimeOffer: + if request.headers.get("content-type", "").startswith("multipart/form-data"): + form: Final = await request.form() + return CodexRealtimeOffer.model_validate( + MappingProxyType({"sdp": form.get("sdp"), "session": json.loads(str(form.get("session", "{}")))}) + ) + return CodexRealtimeOffer.model_validate(await request.json()) + + +async def create_codex_realtime_call(request: Request) -> Response: + from litellm.proxy import proxy_server as server + from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing + + try: + offer: Final = await read_codex_offer(request) + except ValueError as exc: + raise HTTPException(400, "Invalid realtime offer: expected sdp and session") from exc + model: Final = offer.session.model + if not model: + raise HTTPException(400, "session.model is required") + auth: Final = await user_api_key_auth( + request=request, + api_key=request.headers.get("authorization", ""), + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + custom_litellm_key_header=None, + ) + try: + await can_key_call_resolved_model( + model=model, + llm_model_list=server.llm_model_list, + valid_token=auth, + llm_router=server.llm_router, + ) + data: Final = build_call_request(offer, request.query_params, request.headers) + processor: Final = ProxyBaseLLMRequestProcessing(data=data) + processed, _ = await processor.common_processing_pre_call_logic( + request=request, + general_settings=server.general_settings, + user_api_key_dict=auth, + version=server.version, + proxy_logging_obj=server.proxy_logging_obj, + proxy_config=server.proxy_config, + user_model=server.user_model, + user_temperature=server.user_temperature, + user_request_timeout=server.user_request_timeout, + user_max_tokens=server.user_max_tokens, + user_api_base=server.user_api_base, + model=model, + route_type="arealtime_calls", + ) + result: Final = await server.route_request( + data=processed, + route_type="arealtime_calls", + llm_router=server.llm_router, + user_model=server.user_model, + ) + try: + response: Final = await result + except BaseLLMException as exc: + raise HTTPException(exc.status_code, str(exc)) from exc + if not isinstance(response, httpx.Response): + raise HTTPException(502, "Invalid realtime signaling response") + if response.is_error: + return Response(response.content, status_code=response.status_code, media_type="application/json") + try: + call: Final = parse_call_response( + response, + alias=model, + owner=hashlib.sha256(request.headers.get("authorization", "").encode()).hexdigest(), + expires_at=time.time() + 3600, + ) + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + token: Final = encode_call(call) + return Response( + response.content, + status_code=response.status_code, + media_type="application/sdp", + headers=MappingProxyType({"Location": f"/v1/realtime/calls/{token}"}), + ) + finally: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) + + +async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAPIKeyAuth) -> None: + import litellm + from litellm.proxy import proxy_server as server + + try: + try: + call: Final = decode_call(token, websocket.headers.get("authorization", "")) + await can_key_call_resolved_model( + model=call.alias, + llm_model_list=server.llm_model_list, + valid_token=auth, + llm_router=server.llm_router, + ) + except (HTTPException, ProxyException): + await websocket.close(code=1008, reason="Invalid realtime call") + return + await websocket.accept() + await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call + **build_sideband_request(call), + websocket=websocket, + user_api_key_dict=auth, + ) + finally: + await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation) diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index 76843d68353..32c6abdb7b6 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -376,7 +376,7 @@ async def proxy_realtime_calls( fastapi_response: Response, ) -> Response: if request.headers.get("content-type", "").split(";", 1)[0] in ("application/json", "multipart/form-data"): - from litellm.llms.chatgpt.codex import create_codex_realtime_call + from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call return await create_codex_realtime_call(request) diff --git a/tests/test_litellm/llms/chatgpt/test_codex.py b/tests/test_litellm/llms/chatgpt/test_codex.py index 583b81cd507..078c8d85ce3 100644 --- a/tests/test_litellm/llms/chatgpt/test_codex.py +++ b/tests/test_litellm/llms/chatgpt/test_codex.py @@ -1,212 +1,22 @@ -import hashlib -import time - +import httpx import pytest -from fastapi import HTTPException, WebSocket -from litellm.proxy._types import ProxyException, UserAPIKeyAuth -from litellm.llms.chatgpt import codex -from litellm.llms.chatgpt.codex import CodexRealtimeCall, decode_call, encode_call +from litellm.llms.chatgpt.codex import build_sideband_request, parse_call_response -def test_sideband_token_binds_owner_and_model(monkeypatch): - monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") - call = CodexRealtimeCall( - call_id="rtc_test", - model="gpt-live-1-codex", - alias="gpt-live-1-codex", - owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), - expires_at=time.time() + 300, - ) - token = encode_call(call) - assert "/" not in token - assert decode_call(token, "Bearer test-owner") == call - with pytest.raises(HTTPException) as error: - decode_call(token, "Bearer different-owner") - assert error.value.status_code == 403 - with pytest.raises(HTTPException): - decode_call(token[:30] + "tampered" + token[30:], "Bearer test-owner") +@pytest.mark.parametrize("location", ["", "/v1/realtime/calls/foreign-id"]) +def test_signaling_rejects_invalid_upstream_call_id(location): + response = httpx.Response(201, headers={"Location": location}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + with pytest.raises(ValueError): + parse_call_response(response, "voice", "owner", 1000) -def test_sideband_rejects_expired_token(monkeypatch): - monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") - call = CodexRealtimeCall( - call_id="rtc_test", - model="gpt-realtime-1.5", - alias="gpt-realtime-1.5", - owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), - expires_at=time.time() - 1, - ) - with pytest.raises(HTTPException): - decode_call(encode_call(call), "Bearer test-owner") - - -@pytest.mark.parametrize("token", ["", "rtc_other", "rtc_litellm_%%%%", "rtc_litellm_a"]) -def test_sideband_rejects_malformed_tokens(token): - with pytest.raises(HTTPException): - decode_call(token, "Bearer test-owner") - - -@pytest.mark.asyncio -async def test_sideband_rejects_revoked_model_access(monkeypatch): - monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") - call = CodexRealtimeCall( - call_id="rtc_test", - model="gpt-live-1-codex", - alias="voice", - owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), - expires_at=time.time() + 300, - ) - sent = [] - - async def receive(): - return {"type": "websocket.connect"} - - async def send(message): - sent.append(message) - - async def deny_model(**kwargs): - raise ProxyException("Model access revoked", "auth_error", "model", 403) - - monkeypatch.setattr(codex, "can_key_call_resolved_model", deny_model) - websocket = WebSocket( - {"type": "websocket", "headers": [(b"authorization", b"Bearer test-owner")]}, receive, send - ) - await codex.codex_realtime_sideband(websocket, encode_call(call), UserAPIKeyAuth()) - assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Invalid realtime call"}] - - -@pytest.mark.asyncio -@pytest.mark.parametrize("call_id", ["rtc_raw", "", "rtc_litellm_invalid"]) -async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id): - from unittest.mock import AsyncMock - from fastapi import WebSocket - from litellm.proxy import proxy_server as server - from litellm.proxy._types import UserAPIKeyAuth - - sent = [] - - async def receive(): - return {"type": "websocket.connect"} - - async def send(message): - sent.append(message) - - route = AsyncMock() - monkeypatch.setattr(server, "route_request", route) - websocket = WebSocket({"type": "websocket", "headers": [], "query_string": b""}, receive, send) - await server.realtime_websocket_endpoint( - websocket, model="gpt-realtime-1.5", call_id=call_id, - intent=None, guardrails=None, user_api_key_dict=UserAPIKeyAuth() - ) - assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Invalid realtime call"}] - route.assert_not_called() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("multipart", [False, True]) -async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart): - import json - from unittest.mock import AsyncMock - - import httpx - from fastapi import Request, WebSocket - from litellm.proxy import common_request_processing, proxy_server - from litellm.proxy._types import UserAPIKeyAuth - from litellm.llms.chatgpt import codex - import litellm - - monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") - session = {"model": "voice-alias", "audio": {"output": {"voice": "sol"}}} - if multipart: - body_request = httpx.Request("POST", "http://test/v1/realtime/calls", files={ - "sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session)) - }) - else: - body_request = httpx.Request("POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session}) - body = body_request.read() - - async def receive(): - return {"type": "http.request", "body": body, "more_body": False} - - request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", - "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", - "headers": [(b"content-type", body_request.headers["content-type"].encode()), - (b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"), - (b"x-untrusted", b"bad")]}, receive) - auth = UserAPIKeyAuth() - authenticate = AsyncMock(return_value=auth) - authorize = AsyncMock() - monkeypatch.setattr(codex, "user_api_key_auth", authenticate) - monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize) - - class Processor: - def __init__(self, data): - self.data = data - - async def common_processing_pre_call_logic(self, **kwargs): - assert kwargs["user_api_key_dict"] is auth - return self.data, None - - monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor) - - async def route(**kwargs): - data = kwargs["data"] - assert data["sdp_body"] == b"v=0\r\n" - assert data["session"] == session - assert data["extra_headers"] == {"openai-alpha": "quicksilver=v2"} - assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} - - async def respond(): - return httpx.Response(201, content=b"v=0\r\nanswer", headers={"Location": "/v1/realtime/calls/rtc_private"}, - extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) - return respond() - - monkeypatch.setattr(proxy_server, "route_request", route) - response = await codex.create_codex_realtime_call(request) - assert response.status_code == 201 - assert response.body == b"v=0\r\nanswer" - token = response.headers["location"].rsplit("/", 1)[-1] - call = codex.decode_call(token, "Bearer owner") - assert call.call_id == "rtc_private" - assert call.alias == "voice-alias" - assert call.model == "gpt-live-1-codex" - assert "rtc_private" not in token - assert time.time() < call.expires_at < time.time() + 3601 - authorize.assert_awaited_once() - - sent = [] - - async def send(message): - sent.append(message) - - async def receive_ws(): - return {"type": "websocket.connect"} - - websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send) - forward = AsyncMock() - monkeypatch.setattr(litellm, "_arealtime", forward) - await codex.codex_realtime_sideband(websocket, token, auth) - assert sent[0]["type"] == "websocket.accept" - assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" - assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" - assert authorize.await_count == 2 - - -@pytest.mark.asyncio -@pytest.mark.parametrize("body", [b"not json", b'{}', b'{"sdp":"v=0","session":{}}']) -async def test_invalid_offers_fail_before_authentication(monkeypatch, body): - from unittest.mock import AsyncMock - from fastapi import Request - from litellm.llms.chatgpt import codex - - async def receive(): - return {"type": "http.request", "body": body} - - request = Request({"type": "http", "headers": [(b"content-type", b"application/json")]}, receive) - authenticate = AsyncMock() - monkeypatch.setattr(codex, "user_api_key_auth", authenticate) - with pytest.raises(HTTPException) as error: - await codex.create_codex_realtime_call(request) - assert error.value.status_code == 400 - authenticate.assert_not_called() +def test_signaling_preserves_selected_model_for_sideband(): + response = httpx.Response(201, headers={"Location": "/v1/realtime/calls/rtc_provider"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + call = parse_call_response(response, "voice", "owner", 1000) + request = build_sideband_request(call) + assert request["model"] == "chatgpt/gpt-live-1-codex" + assert request["chatgpt_realtime_call_id"] == "rtc_provider" + assert request["query_params"] == {"model": "gpt-live-1-codex"} diff --git a/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py new file mode 100644 index 00000000000..b7217964116 --- /dev/null +++ b/tests/test_litellm/proxy/realtime_endpoints/test_call_sessions.py @@ -0,0 +1,213 @@ +import hashlib +import time + +import pytest +from fastapi import HTTPException, WebSocket + +from litellm.proxy._types import ProxyException, UserAPIKeyAuth +from litellm.proxy.realtime_endpoints import call_sessions as codex +from litellm.llms.chatgpt.codex import CodexRealtimeCall +from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call + + +def test_sideband_token_binds_owner_and_model(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="gpt-live-1-codex", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + token = encode_call(call) + assert "/" not in token + assert decode_call(token, "Bearer test-owner") == call + with pytest.raises(HTTPException) as error: + decode_call(token, "Bearer different-owner") + assert error.value.status_code == 403 + with pytest.raises(HTTPException): + decode_call(token[:30] + "tampered" + token[30:], "Bearer test-owner") + + +def test_sideband_rejects_expired_token(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-realtime-1.5", + alias="gpt-realtime-1.5", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() - 1, + ) + with pytest.raises(HTTPException): + decode_call(encode_call(call), "Bearer test-owner") + + +@pytest.mark.parametrize("token", ["", "rtc_other", "rtc_litellm_%%%%", "rtc_litellm_a"]) +def test_sideband_rejects_malformed_tokens(token): + with pytest.raises(HTTPException): + decode_call(token, "Bearer test-owner") + + +@pytest.mark.asyncio +async def test_sideband_rejects_revoked_model_access(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + call = CodexRealtimeCall( + call_id="rtc_test", + model="gpt-live-1-codex", + alias="voice", + owner=hashlib.sha256(b"Bearer test-owner").hexdigest(), + expires_at=time.time() + 300, + ) + sent = [] + + async def receive(): + return {"type": "websocket.connect"} + + async def send(message): + sent.append(message) + + async def deny_model(**kwargs): + raise ProxyException("Model access revoked", "auth_error", "model", 403) + + monkeypatch.setattr(codex, "can_key_call_resolved_model", deny_model) + websocket = WebSocket( + {"type": "websocket", "headers": [(b"authorization", b"Bearer test-owner")]}, receive, send + ) + await codex.codex_realtime_sideband(websocket, encode_call(call), UserAPIKeyAuth()) + assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Invalid realtime call"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("call_id", ["rtc_raw", "", "rtc_litellm_invalid"]) +async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id): + from unittest.mock import AsyncMock + from fastapi import WebSocket + from litellm.proxy import proxy_server as server + from litellm.proxy._types import UserAPIKeyAuth + + sent = [] + + async def receive(): + return {"type": "websocket.connect"} + + async def send(message): + sent.append(message) + + route = AsyncMock() + monkeypatch.setattr(server, "route_request", route) + websocket = WebSocket({"type": "websocket", "headers": [], "query_string": b""}, receive, send) + await server.realtime_websocket_endpoint( + websocket, model="gpt-realtime-1.5", call_id=call_id, + intent=None, guardrails=None, user_api_key_dict=UserAPIKeyAuth() + ) + assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Invalid realtime call"}] + route.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("multipart", [False, True]) +async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart): + import json + from unittest.mock import AsyncMock + + import httpx + from fastapi import Request, WebSocket + from litellm.proxy import common_request_processing, proxy_server + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.realtime_endpoints import call_sessions as codex + import litellm + + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-salt-for-codex-realtime") + session = {"model": "voice-alias", "audio": {"output": {"voice": "sol"}}} + if multipart: + body_request = httpx.Request("POST", "http://test/v1/realtime/calls", files={ + "sdp": (None, "v=0\r\n"), "session": (None, json.dumps(session)) + }) + else: + body_request = httpx.Request("POST", "http://test/v1/realtime/calls", json={"sdp": "v=0\r\n", "session": session}) + body = body_request.read() + + async def receive(): + return {"type": "http.request", "body": body, "more_body": False} + + request = Request({"type": "http", "method": "POST", "path": "/v1/realtime/calls", + "query_string": b"intent=quicksilver&architecture=avas&untrusted=bad", + "headers": [(b"content-type", body_request.headers["content-type"].encode()), + (b"authorization", b"Bearer owner"), (b"openai-alpha", b"quicksilver=v2"), + (b"x-untrusted", b"bad")]}, receive) + auth = UserAPIKeyAuth() + authenticate = AsyncMock(return_value=auth) + authorize = AsyncMock() + monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + monkeypatch.setattr(codex, "can_key_call_resolved_model", authorize) + + class Processor: + def __init__(self, data): + self.data = data + + async def common_processing_pre_call_logic(self, **kwargs): + assert kwargs["user_api_key_dict"] is auth + return self.data, None + + monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor) + + async def route(**kwargs): + data = kwargs["data"] + assert data["sdp_body"] == b"v=0\r\n" + assert data["session"] == session + assert data["extra_headers"] == {"openai-alpha": "quicksilver=v2"} + assert data["extra_query"] == {"intent": "quicksilver", "architecture": "avas"} + + async def respond(): + return httpx.Response(201, content=b"v=0\r\nanswer", headers={"Location": "/v1/realtime/calls/rtc_private"}, + extensions={"chatgpt_realtime": {"model": "gpt-live-1-codex"}}) + return respond() + + monkeypatch.setattr(proxy_server, "route_request", route) + response = await codex.create_codex_realtime_call(request) + assert response.status_code == 201 + assert response.body == b"v=0\r\nanswer" + token = response.headers["location"].rsplit("/", 1)[-1] + call = codex.decode_call(token, "Bearer owner") + assert call.call_id == "rtc_private" + assert call.alias == "voice-alias" + assert call.model == "gpt-live-1-codex" + assert "rtc_private" not in token + assert time.time() < call.expires_at < time.time() + 3601 + authorize.assert_awaited_once() + + sent = [] + + async def send(message): + sent.append(message) + + async def receive_ws(): + return {"type": "websocket.connect"} + + websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send) + forward = AsyncMock() + monkeypatch.setattr(litellm, "_arealtime", forward) + await codex.codex_realtime_sideband(websocket, token, auth) + assert sent[0]["type"] == "websocket.accept" + assert forward.await_args.kwargs["chatgpt_realtime_call_id"] == "rtc_private" + assert forward.await_args.kwargs["model"] == "chatgpt/gpt-live-1-codex" + assert authorize.await_count == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("body", [b"not json", b'{}', b'{"sdp":"v=0","session":{}}']) +async def test_invalid_offers_fail_before_authentication(monkeypatch, body): + from unittest.mock import AsyncMock + from fastapi import Request + from litellm.proxy.realtime_endpoints import call_sessions as codex + + async def receive(): + return {"type": "http.request", "body": body} + + request = Request({"type": "http", "headers": [(b"content-type", b"application/json")]}, receive) + authenticate = AsyncMock() + monkeypatch.setattr(codex, "user_api_key_auth", authenticate) + with pytest.raises(HTTPException) as error: + await codex.create_codex_realtime_call(request) + assert error.value.status_code == 400 + authenticate.assert_not_called()