refactor(chatgpt): keep HTTP handling inside the proxy

This commit is contained in:
jibanez-staticduo 2026-09-09 08:43:31 +02:00
parent 58a5657784
commit 284cbc18ca
No known key found for this signature in database
6 changed files with 429 additions and 364 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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