Merge pull request #36513 from Souravrajvi0/bugfix/video-edits-form-body-36487-8fc4

fix(proxy): parse form-encoded video edit/extension bodies after auth
This commit is contained in:
Mateo Wang 2026-08-24 12:03:10 -07:00 committed by GitHub
commit f818a48ae5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 138 additions and 21 deletions

View file

@ -2,7 +2,6 @@
from typing import Any, Final from typing import Any, Final
import orjson
from fastapi import APIRouter, Depends, File, Form, Request, Response, UploadFile from fastapi import APIRouter, Depends, File, Form, Request, Response, UploadFile
from fastapi.responses import ORJSONResponse from fastapi.responses import ORJSONResponse
@ -20,6 +19,7 @@ from litellm.proxy.video_endpoints.utils import (
encode_character_id_in_response, encode_character_id_in_response,
extract_model_from_target_model_names, extract_model_from_target_model_names,
get_custom_provider_from_data, get_custom_provider_from_data,
video_reference_to_id,
) )
from litellm.types.videos.utils import ( from litellm.types.videos.utils import (
decode_character_id_with_provider, decode_character_id_with_provider,
@ -451,9 +451,7 @@ async def video_remix(
version, version,
) )
# Read request body data: Final = await _read_request_body(request=request)
body: Final = await request.body()
data: Final = orjson.loads(body)
data["video_id"] = video_id data["video_id"] = video_id
decoded: Final = decode_video_id_with_provider(video_id) decoded: Final = decode_video_id_with_provider(video_id)
@ -760,15 +758,10 @@ async def video_edit(
version, version,
) )
body: Final = await request.body() data: Final = await _read_request_body(request=request)
data: Final = orjson.loads(body) data["video_id"] = video_reference_to_id(data.pop("video", None))
# Extract video_id from nested video object decoded: Final = decode_video_id_with_provider(data["video_id"])
video_ref: Final = data.pop("video", {})
video_id: Final = video_ref.get("id", "") if isinstance(video_ref, dict) else ""
data["video_id"] = video_id
decoded: Final = decode_video_id_with_provider(video_id)
provider_from_id: Final = decoded.get("custom_llm_provider") provider_from_id: Final = decoded.get("custom_llm_provider")
model_id_from_decoded: Final = decoded.get("model_id") model_id_from_decoded: Final = decoded.get("model_id")
@ -860,15 +853,10 @@ async def video_extension(
version, version,
) )
body: Final = await request.body() data: Final = await _read_request_body(request=request)
data: Final = orjson.loads(body) data["video_id"] = video_reference_to_id(data.pop("video", None))
# Extract video_id from nested video object decoded: Final = decode_video_id_with_provider(data["video_id"])
video_ref: Final = data.pop("video", {})
video_id: Final = video_ref.get("id", "") if isinstance(video_ref, dict) else ""
data["video_id"] = video_id
decoded: Final = decode_video_id_with_provider(video_id)
provider_from_id: Final = decoded.get("custom_llm_provider") provider_from_id: Final = decoded.get("custom_llm_provider")
model_id_from_decoded: Final = decoded.get("model_id") model_id_from_decoded: Final = decoded.get("model_id")

View file

@ -13,6 +13,18 @@ def extract_model_from_target_model_names(target_model_names: Any) -> str | None
return target_model_names[0] if target_model_names else None return target_model_names[0] if target_model_names else None
def video_reference_to_id(video_ref: object) -> str:
if isinstance(video_ref, dict):
return video_ref.get("id", "")
if not isinstance(video_ref, str):
return ""
try:
parsed_ref: Final = orjson.loads(video_ref)
except orjson.JSONDecodeError:
return video_ref
return parsed_ref.get("id", "") if isinstance(parsed_ref, dict) else video_ref
def get_custom_provider_from_data(data: dict[str, Any]) -> str | None: def get_custom_provider_from_data(data: dict[str, Any]) -> str | None:
custom_llm_provider: Final = data.get("custom_llm_provider") custom_llm_provider: Final = data.get("custom_llm_provider")
if custom_llm_provider: if custom_llm_provider:

View file

@ -372,6 +372,7 @@ async def test_content__model_encoded_id(harness):
async def call_edit( async def call_edit(
harness: Harness, *, body: Dict[str, Any], headers=None, query=None harness: Harness, *, body: Dict[str, Any], headers=None, query=None
): ):
harness.read_body.return_value = dict(body)
return await endpoints.video_edit( return await endpoints.video_edit(
request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)),
fastapi_response=Response(), fastapi_response=Response(),
@ -428,6 +429,27 @@ async def test_edit__missing_video_object_defaults_to_openai(harness):
assert "video" not in data assert "video" not in data
@pytest.mark.asyncio
async def test_edit__bare_string_video_id_from_form_field(harness):
await call_edit(harness, body={"prompt": "brighter", "video": "video_plain"})
assert harness.processor_data() == {
"prompt": "brighter",
"video_id": "video_plain",
"custom_llm_provider": "openai",
}
@pytest.mark.asyncio
async def test_edit__json_string_video_reference_from_form_field(harness):
await call_edit(
harness,
body={"prompt": "brighter", "video": orjson.dumps({"id": "video_plain"}).decode()},
)
assert harness.processor_data()["video_id"] == "video_plain"
# =========================================================================== # # =========================================================================== #
# GET /v1/videos - video_list # # GET /v1/videos - video_list #
# =========================================================================== # # =========================================================================== #
@ -471,6 +493,7 @@ async def test_list__provider_from_header(harness):
async def call_remix( async def call_remix(
harness: Harness, video_id: str, *, body, headers=None, query=None harness: Harness, video_id: str, *, body, headers=None, query=None
): ):
harness.read_body.return_value = dict(body)
return await endpoints.video_remix( return await endpoints.video_remix(
video_id=video_id, video_id=video_id,
request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)),
@ -629,6 +652,7 @@ async def test_get_character__plain_id_defaults_openai_no_encode(harness):
async def call_extension(harness: Harness, *, body, headers=None, query=None): async def call_extension(harness: Harness, *, body, headers=None, query=None):
harness.read_body.return_value = dict(body)
return await endpoints.video_extension( return await endpoints.video_extension(
request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)),
fastapi_response=Response(), fastapi_response=Response(),

View file

@ -1,8 +1,9 @@
""" """
Pure-logic contract tests for litellm/proxy/video_endpoints/utils.py Pure-logic contract tests for litellm/proxy/video_endpoints/utils.py
Three helpers the video proxy endpoints lean on: Four helpers the video proxy endpoints lean on:
- extract_model_from_target_model_names: first model from a comma string / list - extract_model_from_target_model_names: first model from a comma string / list
- video_reference_to_id: normalize a video reference (dict / bare id / JSON string) to an id
- get_custom_provider_from_data: provider precedence (top-level > extra_body) - get_custom_provider_from_data: provider precedence (top-level > extra_body)
- encode_character_id_in_response: re-encode a response id in place - encode_character_id_in_response: re-encode a response id in place
@ -20,6 +21,7 @@ from litellm.proxy.video_endpoints.utils import (
encode_character_id_in_response, encode_character_id_in_response,
extract_model_from_target_model_names, extract_model_from_target_model_names,
get_custom_provider_from_data, get_custom_provider_from_data,
video_reference_to_id,
) )
from litellm.types.videos.utils import ( from litellm.types.videos.utils import (
decode_character_id_with_provider, decode_character_id_with_provider,
@ -53,6 +55,31 @@ def test_extract_model__non_str_non_list_is_none(value):
assert extract_model_from_target_model_names(value) is None assert extract_model_from_target_model_names(value) is None
# =========================================================================== #
# video_reference_to_id
# =========================================================================== #
@pytest.mark.parametrize(
"video_ref,expected",
[
({"id": "video_123"}, "video_123"), # dict reference -> its id
({"id": ""}, ""), # dict with empty id
({}, ""), # dict missing id -> default empty
({"other": "x"}, ""), # dict without id key
("video_123", "video_123"), # bare id string (not valid JSON) -> itself
('{"id": "video_9"}', "video_9"), # JSON-encoded dict -> its id
('{"other": 1}', ""), # JSON-encoded dict without id -> empty
("[1, 2]", "[1, 2]"), # JSON parses to non-dict -> original string
(None, ""), # non-str, non-dict
(123, ""), # non-str, non-dict
(["video_123"], ""), # list is neither dict nor str
],
)
def test_video_reference_to_id(video_ref, expected):
assert video_reference_to_id(video_ref) == expected
# =========================================================================== # # =========================================================================== #
# get_custom_provider_from_data # get_custom_provider_from_data
# =========================================================================== # # =========================================================================== #

View file

@ -2316,6 +2316,72 @@ def test_edit_and_extension_support_custom_provider_from_extra_body(
assert captured_data["custom_llm_provider"] == "vertex_ai" assert captured_data["custom_llm_provider"] == "vertex_ai"
@pytest.mark.parametrize(
"handler_name, path, form",
[
(
"video_edit",
"/v1/videos/edits",
{"model": "my-video-model", "prompt": "brighter", "video": "video_123"},
),
(
"video_extension",
"/v1/videos/extensions",
{"model": "my-video-model", "prompt": "continue", "seconds": "4", "video": "video_123"},
),
],
)
@pytest.mark.asyncio
async def test_edit_and_extension_read_cached_body_after_auth_consumes_stream(
handler_name, path, form
):
from urllib.parse import urlencode
from fastapi import Response
from starlette.requests import Request
import litellm.proxy.video_endpoints.endpoints as endpoints
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
body = urlencode(form).encode()
stream = {"sent": False}
async def receive():
if stream["sent"]:
return {"type": "http.request", "body": b"", "more_body": False}
stream["sent"] = True
return {"type": "http.request", "body": body, "more_body": False}
request = Request(
{
"type": "http",
"method": "POST",
"path": path,
"headers": [
(b"content-type", b"application/x-www-form-urlencoded"),
(b"content-length", str(len(body)).encode()),
],
"query_string": b"",
},
receive,
)
await _read_request_body(request=request)
handler = getattr(endpoints, handler_name)
with pytest.raises(ProxyException) as exc_info:
await handler(
request=request,
fastapi_response=Response(),
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234"),
)
message = str(exc_info.value)
assert "Stream consumed" not in message
assert "my-video-model" in message
@pytest.mark.parametrize("endpoint", ["/v1/videos/edits", "/v1/videos/extensions"]) @pytest.mark.parametrize("endpoint", ["/v1/videos/edits", "/v1/videos/extensions"])
def test_edit_and_extension_route_with_encoded_video_ids( def test_edit_and_extension_route_with_encoded_video_ids(
video_proxy_test_client, endpoint video_proxy_test_client, endpoint