mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(chatgpt): enforce sideband policies and websocket credentials
This commit is contained in:
parent
8e5901e90a
commit
36727584a9
2 changed files with 138 additions and 25 deletions
|
|
@ -3,11 +3,13 @@ import hashlib
|
|||
import json
|
||||
import time
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
from fastapi import HTTPException, Request, Response, WebSocket
|
||||
from starlette.types import Message
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.chatgpt.codex import (
|
||||
CodexRealtimeCall,
|
||||
|
|
@ -52,10 +54,38 @@ async def read_codex_offer(request: Request) -> CodexRealtimeOffer:
|
|||
return CodexRealtimeOffer.model_validate(await request.json())
|
||||
|
||||
|
||||
async def create_codex_realtime_call(request: Request) -> Response:
|
||||
async def process_codex_request(
|
||||
request: Request,
|
||||
data: dict[str, object], # mutable-ok: common request processor enriches this dictionary
|
||||
auth: UserAPIKeyAuth,
|
||||
model: str,
|
||||
route_type: Literal["arealtime_calls", "_arealtime"],
|
||||
) -> dict[str, object]: # mutable-ok: common request processor returns enriched routing arguments
|
||||
from litellm.proxy import proxy_server as server
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
||||
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=route_type,
|
||||
)
|
||||
return processed
|
||||
|
||||
|
||||
async def create_codex_realtime_call(request: Request) -> Response:
|
||||
from litellm.proxy import proxy_server as server
|
||||
|
||||
try:
|
||||
offer: Final = await read_codex_offer(request)
|
||||
except ValueError as exc:
|
||||
|
|
@ -80,22 +110,7 @@ async def create_codex_realtime_call(request: Request) -> Response:
|
|||
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",
|
||||
)
|
||||
processed: Final = await process_codex_request(request, data, auth, model, "arealtime_calls")
|
||||
result: Final = await server.route_request(
|
||||
data=processed,
|
||||
route_type="arealtime_calls",
|
||||
|
|
@ -134,9 +149,16 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
|
|||
import litellm
|
||||
from litellm.proxy import proxy_server as server
|
||||
|
||||
protocols: Final = tuple(
|
||||
p.strip() for p in websocket.headers.get("sec-websocket-protocol", "").split(",") if p.strip()
|
||||
)
|
||||
alternate_key: Final = websocket.headers.get("api-key") or next(
|
||||
(p.removeprefix("openai-insecure-api-key.") for p in protocols if p.startswith("openai-insecure-api-key.")), ""
|
||||
)
|
||||
authorization: Final = websocket.headers.get("authorization") or f"Bearer {alternate_key}"
|
||||
try:
|
||||
try:
|
||||
call: Final = decode_call(token, websocket.headers.get("authorization", ""))
|
||||
call: Final = decode_call(token, authorization)
|
||||
await can_key_call_resolved_model(
|
||||
model=call.alias,
|
||||
llm_model_list=server.llm_model_list,
|
||||
|
|
@ -146,11 +168,47 @@ async def codex_realtime_sideband(websocket: WebSocket, token: str, auth: UserAP
|
|||
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
|
||||
|
||||
async def receive() -> Message:
|
||||
return { # mutable-ok: ASGI receive message
|
||||
"type": "http.request",
|
||||
"body": json.dumps({"model": call.alias}).encode(), # mutable-ok: JSON request serialization
|
||||
"more_body": False,
|
||||
}
|
||||
|
||||
request: Final = Request(
|
||||
{ # mutable-ok: Starlette stores request state in the ASGI scope
|
||||
**websocket.scope,
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": websocket.scope.get("path", "/v1/realtime"),
|
||||
},
|
||||
receive=receive,
|
||||
)
|
||||
data: Final = { # mutable-ok: common request processor enriches routing arguments
|
||||
**build_sideband_request(call),
|
||||
websocket=websocket,
|
||||
user_api_key_dict=auth,
|
||||
"model": call.alias,
|
||||
"websocket": websocket,
|
||||
"guardrails": [ # mutable-ok: guardrail processing expects a list
|
||||
name.strip() for name in websocket.query_params.get("guardrails", "").split(",") if name.strip()
|
||||
],
|
||||
}
|
||||
try:
|
||||
processed: Final = await process_codex_request(request, data, auth, call.alias, "_arealtime")
|
||||
except Exception:
|
||||
verbose_proxy_logger.exception("Realtime sideband pre-call rejected")
|
||||
await websocket.close(code=1008, reason="Realtime pre-call rejected")
|
||||
return
|
||||
await websocket.accept(
|
||||
subprotocol=next((p for p in protocols if not p.startswith("openai-insecure-api-key.")), None)
|
||||
)
|
||||
await litellm._arealtime( # pyright: ignore[reportPrivateUsage] # dispatch for an already authorized call
|
||||
**{ # mutable-ok: retain processed policy metadata while pinning the existing call's routing
|
||||
**processed,
|
||||
**build_sideband_request(call),
|
||||
"websocket": websocket,
|
||||
"user_api_key_dict": auth,
|
||||
}
|
||||
)
|
||||
finally:
|
||||
await release_or_invalidate_budget_reservation(budget_reservation=auth.budget_reservation)
|
||||
|
|
|
|||
|
|
@ -106,7 +106,8 @@ async def test_realtime_endpoint_rejects_untrusted_call_ids(monkeypatch, call_id
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("multipart", [False, True])
|
||||
async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart):
|
||||
@pytest.mark.parametrize("credential", ["authorization", "api-key", "subprotocol"])
|
||||
async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch, multipart, credential):
|
||||
import json
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
|
@ -147,6 +148,11 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
|
||||
async def common_processing_pre_call_logic(self, **kwargs):
|
||||
assert kwargs["user_api_key_dict"] is auth
|
||||
if kwargs["route_type"] == "_arealtime":
|
||||
assert self.data["model"] == "voice-alias"
|
||||
assert self.data["guardrails"] == ["query-guardrail"]
|
||||
assert await kwargs["request"].json() == {"model": "voice-alias"}
|
||||
return {**self.data, "metadata": {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}}, None
|
||||
return self.data, None
|
||||
|
||||
monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", Processor)
|
||||
|
|
@ -184,11 +190,20 @@ async def test_offer_exchange_wraps_call_and_filters_client_headers(monkeypatch,
|
|||
async def receive_ws():
|
||||
return {"type": "websocket.connect"}
|
||||
|
||||
websocket = WebSocket({"type": "websocket", "headers": [(b"authorization", b"Bearer owner")]}, receive_ws, send)
|
||||
credential_headers = {
|
||||
"authorization": [(b"authorization", b"Bearer owner")],
|
||||
"api-key": [(b"api-key", b"owner")],
|
||||
"subprotocol": [(b"sec-websocket-protocol", b"realtime, openai-insecure-api-key.owner")],
|
||||
}
|
||||
websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque",
|
||||
"query_string": b"guardrails=query-guardrail", "headers": credential_headers[credential]}, receive_ws, send)
|
||||
forward = AsyncMock()
|
||||
monkeypatch.setattr(litellm, "_arealtime", forward)
|
||||
await codex.codex_realtime_sideband(websocket, token, auth)
|
||||
assert sent[0]["type"] == "websocket.accept"
|
||||
if credential == "subprotocol":
|
||||
assert sent[0]["subprotocol"] == "realtime"
|
||||
assert forward.await_args.kwargs["metadata"] == {"guardrails": ["policy-guardrail"], "user_api_key_team_id": "team"}
|
||||
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
|
||||
|
|
@ -211,3 +226,43 @@ async def test_invalid_offers_fail_before_authentication(monkeypatch, body):
|
|||
await codex.create_codex_realtime_call(request)
|
||||
assert error.value.status_code == 400
|
||||
authenticate.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sideband_pre_call_block_prevents_upstream_connection(monkeypatch):
|
||||
from unittest.mock import AsyncMock
|
||||
from fastapi import WebSocket
|
||||
import litellm
|
||||
from litellm.proxy import common_request_processing
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.realtime_endpoints import call_sessions as codex
|
||||
|
||||
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 owner").hexdigest(), expires_at=time.time()+300)
|
||||
token = encode_call(call)
|
||||
sent = []
|
||||
|
||||
async def receive():
|
||||
return {"type": "websocket.connect"}
|
||||
|
||||
async def send(message):
|
||||
sent.append(message)
|
||||
|
||||
class BlockingProcessor:
|
||||
def __init__(self, data):
|
||||
assert data["model"] == "voice"
|
||||
|
||||
async def common_processing_pre_call_logic(self, **kwargs):
|
||||
assert kwargs["route_type"] == "_arealtime"
|
||||
raise HTTPException(403, "Policy blocked this call")
|
||||
|
||||
forward = AsyncMock()
|
||||
monkeypatch.setattr(litellm, "_arealtime", forward)
|
||||
monkeypatch.setattr(codex, "can_key_call_resolved_model", AsyncMock())
|
||||
monkeypatch.setattr(common_request_processing, "ProxyBaseLLMRequestProcessing", BlockingProcessor)
|
||||
websocket = WebSocket({"type": "websocket", "path": "/v1/live/opaque", "query_string": b"",
|
||||
"headers": [(b"authorization", b"Bearer owner")]}, receive, send)
|
||||
await codex.codex_realtime_sideband(websocket, token, UserAPIKeyAuth())
|
||||
forward.assert_not_called()
|
||||
assert sent == [{"type": "websocket.close", "code": 1008, "reason": "Realtime pre-call rejected"}]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue