mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(realtime): normalize offer media types and close parsed forms
This commit is contained in:
parent
464c1eb2bc
commit
45cd83c9c6
3 changed files with 102 additions and 9 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue