mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
refactor(chatgpt): keep HTTP handling inside the proxy
This commit is contained in:
parent
58a5657784
commit
284cbc18ca
6 changed files with 429 additions and 364 deletions
|
|
@ -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},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
156
litellm/proxy/realtime_endpoints/call_sessions.py
Normal file
156
litellm/proxy/realtime_endpoints/call_sessions.py
Normal 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)
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
Loading…
Add table
Reference in a new issue