fix(realtime): normalize offer media types and close parsed forms

This commit is contained in:
jibanez-staticduo 2026-09-11 11:58:29 +02:00
parent 464c1eb2bc
commit 45cd83c9c6
No known key found for this signature in database
3 changed files with 102 additions and 9 deletions

View file

@ -12,6 +12,7 @@ from typing import Final, Literal
import httpx
from fastapi import HTTPException, Request, Response, WebSocket
from pydantic import TypeAdapter
from starlette.formparsers import MultiPartException, MultiPartParser
from starlette.types import Message
from litellm._logging import verbose_proxy_logger
@ -44,6 +45,9 @@ 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.common_utils.http_parsing_utils import (
_normalize_media_type, # pyright: ignore[reportPrivateUsage] # reuse the shared HTTP media-type normalization contract
)
from litellm.proxy.hooks.parallel_request_limiter import (
_PROXY_MaxParallelRequestsHandler, # pyright: ignore[reportPrivateUsage] # existing built-in limiter has no public alias
)
@ -209,7 +213,14 @@ def decode_call(token: str, authorization: str) -> CodexRealtimeCall:
async def read_codex_offer(request: Request) -> CodexRealtimeOffer:
if request.headers.get("content-type", "").startswith("multipart/form-data"):
content_type: Final = request.headers.get("content-type", "")
if _normalize_media_type(content_type) == "multipart/form-data":
if content_type.split(";", 1)[0] != "multipart/form-data" and not await request.form():
try:
request._form = await MultiPartParser(request.headers, request.stream()).parse() # pyright: ignore[reportPrivateUsage] # Starlette exposes no setter for its shared form cache
request.scope.pop("parsed_body", None)
except MultiPartException as exc:
raise HTTPException(400, "Invalid realtime multipart offer") from exc
form: Final = await request.form()
return CodexRealtimeOffer.model_validate(
MappingProxyType({"sdp": form.get("sdp"), "session": json.loads(str(form.get("session", "{}")))})
@ -257,8 +268,11 @@ async def process_codex_request(
async def create_codex_realtime_call(request: Request) -> Response:
with isolated_request_stash():
return await _create_codex_realtime_call(request)
try:
with isolated_request_stash():
return await _create_codex_realtime_call(request)
finally:
await request.close()
async def _create_codex_realtime_call(request: Request) -> Response:

View file

@ -16,7 +16,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.http_parsing_utils import (
_normalize_media_type, # pyright: ignore[reportPrivateUsage] # reuse the shared HTTP media-type normalization contract
_read_request_body,
)
from litellm.proxy.common_utils.openai_error_payload import (
error_status_code,
openai_error_param,
@ -375,7 +378,7 @@ async def proxy_realtime_calls(
request: Request,
fastapi_response: Response,
) -> Response:
if request.headers.get("content-type", "").split(";", 1)[0] in ("application/json", "multipart/form-data"):
if _normalize_media_type(request.headers.get("content-type", "")) in ("application/json", "multipart/form-data"):
from litellm.proxy.realtime_endpoints.call_sessions import create_codex_realtime_call
return await create_codex_realtime_call(request)

View file

@ -11,15 +11,75 @@ from litellm.llms.chatgpt.codex import CodexRealtimeCall
from litellm.proxy.realtime_endpoints.call_sessions import decode_call, encode_call
@pytest.mark.asyncio
@pytest.mark.parametrize("malformed", [False, True])
async def test_mixed_case_offer_preserves_boundary_metadata_and_closes_extra_files(monkeypatch, malformed):
from fastapi import Request
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
boundary = "AbCdEf123"
fields = {
"sdp": "v=0",
"session": "invalid" if malformed else '{"model":"voice"}',
"metadata": '{"policy":"keep"}',
"extra_policy": "keep",
}
body = (
"".join(
f'--{boundary}\r\nContent-Disposition: form-data; name="{name}"\r\n\r\n{value}\r\n'
for name, value in fields.items()
)
+ f'--{boundary}\r\nContent-Disposition: form-data; name="extra_file"; filename="test.txt"\r\n\r\nextra\r\n--{boundary}--\r\n'
).encode()
async def receive():
return {"type": "http.request", "body": body, "more_body": False}
request = Request(
{"type": "http", "headers": [(b"content-type", f'Multipart/Form-Data; boundary="{boundary}"'.encode())]},
receive,
)
assert not await request.form()
if malformed:
with pytest.raises(HTTPException) as error:
await codex.create_codex_realtime_call(request)
assert error.value.status_code == 400
assert (await request.form())["extra_file"].file.closed
return
first = await codex.read_codex_offer(request)
second = await codex.read_codex_offer(request)
assert first == second
parsed = await _read_request_body(request)
assert parsed["metadata"] == {"policy": "keep"}
assert parsed["extra_policy"] == "keep"
assert not parsed["extra_file"].file.closed
async def deny_auth(**kwargs):
assert kwargs["request"] is request
auth_form = await request.form()
assert auth_form["extra_policy"] == "keep"
assert await auth_form["extra_file"].read() == b"extra"
assert not auth_form["extra_file"].file.closed
raise HTTPException(403, "policy denied")
monkeypatch.setattr(codex, "user_api_key_auth", deny_auth)
with pytest.raises(HTTPException, match="policy denied"):
await codex.create_codex_realtime_call(request)
assert parsed["extra_file"].file.closed
assert request.headers["content-type"] == f'Multipart/Form-Data; boundary="{boundary}"'
@pytest.mark.asyncio
@pytest.mark.parametrize("multipart", [False, True])
@pytest.mark.parametrize("policy", ["budget", "personal_models"])
async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypatch, multipart, policy):
@pytest.mark.parametrize("mixed_case", [False, True])
@pytest.mark.parametrize("pre_read", [False, True])
async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypatch, multipart, policy, mixed_case, pre_read):
import json
from unittest.mock import AsyncMock, MagicMock
import httpx
from fastapi import Request
from fastapi import Request, Response
import litellm
from litellm.exceptions import BudgetExceededError
@ -27,6 +87,7 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth.auth_checks import common_checks
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.realtime_endpoints.endpoints import proxy_realtime_calls
session = {"model": "forbidden-voice"}
payload = (
@ -36,8 +97,16 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa
)
outbound = httpx.Request("POST", "http://localhost/v1/realtime/calls", **payload)
body = outbound.read()
content_type = outbound.headers["content-type"]
if mixed_case:
content_type = content_type.replace("multipart/form-data", "Multipart/Form-Data").replace(
"application/json", "Application/JSON"
)
receives = []
async def receive():
receives.append(True)
assert len(receives) == 1
return {"type": "http.request", "body": body, "more_body": False}
request = Request(
@ -48,7 +117,7 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa
"query_string": b"model=query-decoy&policy=keep",
"client": ("127.0.0.7", 1234),
"headers": [
*((key.lower(), value) for key, value in outbound.headers.raw),
(b"content-type", content_type.encode()),
(b"x-policy-key", b"Bearer test-key"),
(b"x-custom-policy", b"preserved"),
(b"x-litellm-model", b"header-decoy"),
@ -59,12 +128,17 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa
token = UserAPIKeyAuth(token="test-key", user_id="personal-user", model_max_budget={"forbidden-voice": 0})
budget = AsyncMock(side_effect=BudgetExceededError(current_cost=1, max_budget=0))
upstream = AsyncMock()
original_request = request
async def custom_auth(request: Request, api_key: str):
assert request is original_request
assert request.headers["content-type"] == content_type
assert api_key == "test-key"
assert request.headers["x-custom-policy"] == "preserved"
assert request.query_params["policy"] == "keep"
assert request.client.host == "127.0.0.7"
if multipart:
assert (await request.form())["model"] == "body-decoy"
parsed = await _read_request_body(request)
assert parsed["model"] == "body-decoy"
assert isinstance(parsed["session"], str) is multipart
@ -95,8 +169,10 @@ async def test_offer_auth_enforces_session_model_policy_before_upstream(monkeypa
monkeypatch.setattr(server, "model_max_budget_limiter", SimpleNamespace(is_key_within_model_budget=budget))
monkeypatch.setattr(server, "route_request", upstream)
monkeypatch.setattr(litellm, "enable_post_custom_auth_checks", True, raising=False)
if pre_read:
await _read_request_body(request)
with pytest.raises(ProxyException) as denied:
await codex.create_codex_realtime_call(request)
await proxy_realtime_calls(request, Response())
if policy == "personal_models":
assert "user not allowed to access model" in str(denied.value)
assert "forbidden-voice" in str(denied.value)