mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
commit
f818a48ae5
5 changed files with 138 additions and 21 deletions
|
|
@ -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")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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:
|
||||||
|
|
|
||||||
|
|
@ -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(),
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
# =========================================================================== #
|
# =========================================================================== #
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue