From a04321d2e119e501678e5efd465bd73b1b4cf4a7 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Mon, 29 Jun 2026 09:12:51 +0530 Subject: [PATCH] test(videos): add 1:1 test file scaffold for videos component paths (#30631) Keep only video test files and CI workflow entries; drop unrelated production code and non-video test changes from this branch. Co-authored-by: Cursor --- .github/workflows/test-unit-misc.yml | 1 + .../workflows/test-unit-proxy-endpoints.yml | 1 + .../proxy/video_endpoints/__init__.py | 0 .../proxy/video_endpoints/test_endpoints.py | 656 ++++++++++++++++++ .../proxy/video_endpoints/test_utils.py | 184 +++++ tests/test_litellm/videos/__init__.py | 0 tests/test_litellm/videos/test_main.py | 458 ++++++++++++ tests/test_litellm/videos/test_utils.py | 196 ++++++ 8 files changed, 1496 insertions(+) create mode 100644 tests/test_litellm/proxy/video_endpoints/__init__.py create mode 100644 tests/test_litellm/proxy/video_endpoints/test_endpoints.py create mode 100644 tests/test_litellm/proxy/video_endpoints/test_utils.py create mode 100644 tests/test_litellm/videos/__init__.py create mode 100644 tests/test_litellm/videos/test_main.py create mode 100644 tests/test_litellm/videos/test_utils.py diff --git a/.github/workflows/test-unit-misc.yml b/.github/workflows/test-unit-misc.yml index 133c135d97a..d411d996d8e 100644 --- a/.github/workflows/test-unit-misc.yml +++ b/.github/workflows/test-unit-misc.yml @@ -36,6 +36,7 @@ jobs: tests/test_litellm/passthrough tests/test_litellm/sandbox tests/test_litellm/vector_stores + tests/test_litellm/videos tests/test_litellm/test_*.py workers: 2 reruns: 2 diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index d4f00050596..23ca7a8f2f9 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -31,6 +31,7 @@ jobs: tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/openai_files_endpoint + tests/test_litellm/proxy/video_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/vector_store_endpoints diff --git a/tests/test_litellm/proxy/video_endpoints/__init__.py b/tests/test_litellm/proxy/video_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py new file mode 100644 index 00000000000..40a26fad3c3 --- /dev/null +++ b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py @@ -0,0 +1,656 @@ +""" +Routing-contract tests for litellm/proxy/video_endpoints/endpoints.py + +Unlike the batches layer, every video endpoint funnels into a single downstream +seam - ProxyBaseLLMRequestProcessing.base_process_llm_request - so there is no +provider-dispatch to assert. All of the video-specific, regression-worthy logic +runs *before* that call, while the endpoint assembles the `data` dict. Each test +therefore locks four things: + + 1. ROUTE_TYPE - the exact route_type each endpoint forwards + (avideo_generation/status/content/edit). Swapping two would + silently route requests to the wrong handler. + 2. DATA SHAPE - the entire `data` dict the processor is constructed with: + provider-precedence resolution, video_id passthrough/extraction, + model resolution from the decoded model_id, and file attachment. + 3. RESULT - base_process_llm_request's return value is propagated untouched + (except where the endpoint transforms it). + 4. OUTPUT SHAPE - video_content wraps raw bytes in a Response (video/mp4 + + Content-Disposition). + +Only true I/O boundaries are mocked (the downstream processor call, request body +parsing, file->bytes conversion, the provider-from-request readers, and the +router's model-id resolver). The id decode helpers and get_custom_provider_from_data +run for real, so the data assertions reflect production exactly. base_process is +patched with autospec so the real __init__ still stores self.data (captured via the +mock's call args), and a brand-new kwarg added to this layer surfaces as a failure. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict, Optional +from unittest.mock import AsyncMock, MagicMock, patch + +import orjson +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm.proxy.proxy_server as proxy_server +import litellm.proxy.video_endpoints.endpoints as endpoints +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.utils import ProxyLogging +from litellm.router import Router +from litellm.types.videos.utils import ( + encode_character_id_with_provider, + encode_video_id_with_provider, +) + +from fastapi import Response + +# --------------------------------------------------------------------------- # +# A real model-encoded video id: decodes (for real) to provider "azure", +# model_id VIDEO_MODEL_ID, original video id "video_orig123". The router's +# resolver maps that model_id to a model name; an unknown id resolves to None, +# so a wrong/hardcoded model_id cannot produce a plausible-looking result. +# --------------------------------------------------------------------------- # + +VIDEO_MODEL_ID = "deployment-123" +AZURE_VIDEO_ID = encode_video_id_with_provider("video_orig123", "azure", VIDEO_MODEL_ID) +# A real model-encoded character id: decodes to provider "azure", VIDEO_MODEL_ID, +# original character id "char_orig". Distinct from the video id so a test cannot +# pass by reusing the wrong constant. +AZURE_CHARACTER_ID = encode_character_id_with_provider( + "char_orig", "azure", VIDEO_MODEL_ID +) +RESOLVED_MODELS: Dict[str, str] = {VIDEO_MODEL_ID: "azure-sora"} + +# Sentinel propagated by base_process for the passthrough endpoints. +SENTINEL = object() + + +class FakeRequest: + """Minimal stand-in. headers/query_params are read by the provider readers + (mocked) and on the edit path the raw body is parsed for real via orjson.""" + + def __init__( + self, + headers: Optional[Dict[str, str]] = None, + query: Optional[Dict[str, str]] = None, + raw_body: bytes = b"{}", + ): + self.headers = headers or {} + self.query_params = query or {} + self._raw_body = raw_body + + async def body(self) -> bytes: + return self._raw_body + + +@dataclass +class Harness: + read_body: AsyncMock + batch_to_bytesio: AsyncMock + base_process: MagicMock + handle_exc: AsyncMock + provider_from_headers: MagicMock + provider_from_query: MagicMock + provider_from_body: AsyncMock + router: MagicMock + resolve_model: MagicMock + + def processor_data(self) -> Dict[str, Any]: + """The exact `data` dict the processor was constructed with.""" + assert self.base_process.call_count == 1 + return dict(self.base_process.call_args.args[0].data) + + def route_type(self) -> str: + return self.base_process.call_args.kwargs["route_type"] + + +@pytest.fixture +def harness(): + logging = MagicMock(spec=ProxyLogging) + + router = MagicMock(spec=Router) + resolve_model = MagicMock( + side_effect=lambda model_id: RESOLVED_MODELS.get(model_id) + ) + router.resolve_model_name_from_model_id = resolve_model + + read_body = AsyncMock(return_value={}) + batch_to_bytesio = AsyncMock(return_value=[b"filebytes"]) + handle_exc = AsyncMock(return_value=RuntimeError("handled")) + provider_from_headers = MagicMock(return_value=None) + provider_from_query = MagicMock(return_value=None) + provider_from_body = AsyncMock(return_value=None) + + with ExitStack() as stack: + base_process = stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "base_process_llm_request", + autospec=True, + ) + ) + base_process.return_value = SENTINEL + stack.enter_context( + patch.object( + ProxyBaseLLMRequestProcessing, + "_handle_llm_api_exception", + handle_exc, + ) + ) + stack.enter_context(patch.object(endpoints, "_read_request_body", read_body)) + stack.enter_context( + patch.object(endpoints, "batch_to_bytesio", batch_to_bytesio) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_headers", + provider_from_headers, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_query", + provider_from_query, + ) + ) + stack.enter_context( + patch.object( + endpoints, + "get_custom_llm_provider_from_request_body", + provider_from_body, + ) + ) + stack.enter_context(patch.object(proxy_server, "llm_router", router)) + stack.enter_context(patch.object(proxy_server, "proxy_logging_obj", logging)) + stack.enter_context(patch.object(proxy_server, "general_settings", {})) + stack.enter_context(patch.object(proxy_server, "proxy_config", MagicMock())) + stack.enter_context( + patch.object(proxy_server, "select_data_generator", MagicMock()) + ) + stack.enter_context(patch.object(proxy_server, "user_model", None)) + stack.enter_context(patch.object(proxy_server, "user_temperature", None)) + stack.enter_context(patch.object(proxy_server, "user_request_timeout", None)) + stack.enter_context(patch.object(proxy_server, "user_max_tokens", None)) + stack.enter_context(patch.object(proxy_server, "user_api_base", None)) + stack.enter_context(patch.object(proxy_server, "version", "test-version")) + + yield Harness( + read_body=read_body, + batch_to_bytesio=batch_to_bytesio, + base_process=base_process, + handle_exc=handle_exc, + provider_from_headers=provider_from_headers, + provider_from_query=provider_from_query, + provider_from_body=provider_from_body, + router=router, + resolve_model=resolve_model, + ) + + +def _user() -> UserAPIKeyAuth: + return UserAPIKeyAuth(api_key="sk-test") + + +# =========================================================================== # +# POST /v1/videos - video_generation # +# =========================================================================== # + + +async def call_generation( + harness: Harness, *, body: Dict[str, Any], input_reference=None +): + harness.read_body.return_value = body + return await endpoints.video_generation( + request=FakeRequest(), + fastapi_response=Response(), + input_reference=input_reference, + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_generation__route_type_data_and_no_provider_default(harness): + body = {"model": "sora-2", "prompt": "a sunset"} + + resp = await call_generation(harness, body=body) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_generation" + # generation does NOT resolve a provider; data is the body, untouched. A + # future default custom_llm_provider injection would break this row. + assert harness.processor_data() == {"model": "sora-2", "prompt": "a sunset"} + harness.batch_to_bytesio.assert_not_called() + + +@pytest.mark.asyncio +async def test_generation__input_reference_attached(harness): + body = {"model": "sora-2", "prompt": "a sunset"} + upload = MagicMock(name="upload_file") + + await call_generation(harness, body=body, input_reference=upload) + + harness.batch_to_bytesio.assert_called_once_with([upload]) + assert harness.processor_data() == { + "model": "sora-2", + "prompt": "a sunset", + "input_reference": b"filebytes", + } + + +@pytest.mark.asyncio +async def test_generation__exception_routed_through_handler(harness): + harness.base_process.side_effect = ValueError("provider boom") + + with pytest.raises(RuntimeError, match="handled"): + await call_generation(harness, body={"model": "sora-2"}) + + harness.handle_exc.assert_called_once() + assert harness.handle_exc.call_args.kwargs["e"].args[0] == "provider boom" + + +# =========================================================================== # +# GET /v1/videos/{video_id} - video_status # +# =========================================================================== # + + +async def call_status(harness: Harness, video_id: str, *, headers=None, query=None): + return await endpoints.video_status( + video_id=video_id, + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_status__model_encoded_id_full_contract(harness): + resp = await call_status(harness, AZURE_VIDEO_ID) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_status" + # provider comes from the decoded id; model_id resolved to a model name. + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +@pytest.mark.asyncio +async def test_status__plain_id_defaults_to_openai(harness): + await call_status(harness, "video_plain") + + # plain id -> nothing decoded, no header/query/body provider -> "openai". + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "video_id": "video_plain", + "custom_llm_provider": "openai", + } + + +@pytest.mark.asyncio +async def test_status__header_provider_beats_decoded_id(harness): + harness.provider_from_headers.return_value = "bedrock" + + await call_status(harness, AZURE_VIDEO_ID) + + data = harness.processor_data() + # header wins over the provider decoded from the id ... + assert data["custom_llm_provider"] == "bedrock" + # ... but the model is still resolved from the decoded model_id. + assert data["model"] == "azure-sora" + + +# =========================================================================== # +# GET /v1/videos/{video_id}/content - video_content # +# =========================================================================== # + + +async def call_content(harness: Harness, video_id: str, *, headers=None, query=None): + return await endpoints.video_content( + video_id=video_id, + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_content__wraps_raw_bytes_in_response(harness): + harness.base_process.return_value = b"VIDEOBYTES" + + resp = await call_content(harness, "video_plain") + + assert harness.route_type() == "avideo_content" + assert isinstance(resp, Response) + assert resp.body == b"VIDEOBYTES" + assert resp.media_type == "video/mp4" + assert ( + resp.headers["content-disposition"] + == "attachment; filename=video_video_plain.mp4" + ) + + +@pytest.mark.asyncio +async def test_content__plain_id_has_no_openai_default(harness): + """The high-value asymmetry vs video_status: content stops at the decoded + provider and never injects an 'openai' default, so a plain id leaves + custom_llm_provider unset. A copy-paste of status' fallback breaks this.""" + harness.base_process.return_value = b"x" + + await call_content(harness, "video_plain") + + assert harness.processor_data() == {"video_id": "video_plain"} + + +@pytest.mark.asyncio +async def test_content__model_encoded_id(harness): + harness.base_process.return_value = b"x" + + await call_content(harness, AZURE_VIDEO_ID) + + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +# =========================================================================== # +# POST /v1/videos/edits - video_edit # +# =========================================================================== # + + +async def call_edit( + harness: Harness, *, body: Dict[str, Any], headers=None, query=None +): + return await endpoints.video_edit( + request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_edit__extracts_nested_video_id_full_contract(harness): + resp = await call_edit( + harness, body={"prompt": "brighter", "video": {"id": AZURE_VIDEO_ID}} + ) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_edit" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + # nested video object is popped; its id becomes video_id; provider/model + # derived from the encoded id. + assert harness.processor_data() == { + "prompt": "brighter", + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +@pytest.mark.asyncio +async def test_edit__provider_from_body_data_for_plain_id(harness): + """For a plain id, get_custom_provider_from_data (run for real) pulls the + provider out of the request body before the 'openai' default.""" + await call_edit( + harness, + body={ + "prompt": "x", + "video": {"id": "video_plain"}, + "custom_llm_provider": "vertex_ai", + }, + ) + + data = harness.processor_data() + assert data["video_id"] == "video_plain" + assert data["custom_llm_provider"] == "vertex_ai" + harness.resolve_model.assert_not_called() + + +@pytest.mark.asyncio +async def test_edit__missing_video_object_defaults_to_openai(harness): + await call_edit(harness, body={"prompt": "x"}) + + data = harness.processor_data() + # no video object -> empty video_id; plain -> default provider. + assert data["video_id"] == "" + assert data["custom_llm_provider"] == "openai" + assert "video" not in data + + +# =========================================================================== # +# GET /v1/videos - video_list # +# =========================================================================== # + + +async def call_list(harness: Harness, *, headers=None, query=None): + return await endpoints.video_list( + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_list__query_params_and_no_provider(harness): + resp = await call_list(harness, query={"limit": "5"}) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_list" + # no provider anywhere -> custom_llm_provider stays absent (only set if truthy). + assert harness.processor_data() == {"query_params": {"limit": "5"}} + + +@pytest.mark.asyncio +async def test_list__provider_from_header(harness): + harness.provider_from_headers.return_value = "bedrock" + + await call_list(harness) + + assert harness.processor_data() == { + "query_params": {}, + "custom_llm_provider": "bedrock", + } + + +# =========================================================================== # +# POST /v1/videos/{video_id}/remix - video_remix # +# =========================================================================== # + + +async def call_remix( + harness: Harness, video_id: str, *, body, headers=None, query=None +): + return await endpoints.video_remix( + video_id=video_id, + request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_remix__model_encoded_id_full_contract(harness): + resp = await call_remix(harness, AZURE_VIDEO_ID, body={"prompt": "new colors"}) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_remix" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "prompt": "new colors", + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } + + +@pytest.mark.asyncio +async def test_remix__provider_from_body_data_not_request_body_reader(harness): + """remix resolves the provider from data.get('custom_llm_provider'), never + from the async request-body reader (unlike status/get_character). Setting + that reader to a sentinel and asserting it is untouched locks the difference.""" + harness.provider_from_body.return_value = "must-not-win" + + await call_remix( + harness, + "video_plain", + body={"prompt": "x", "custom_llm_provider": "vertex_ai"}, + ) + + harness.provider_from_body.assert_not_called() + data = harness.processor_data() + assert data["video_id"] == "video_plain" + assert data["custom_llm_provider"] == "vertex_ai" + + +@pytest.mark.asyncio +async def test_remix__plain_id_has_no_openai_default(harness): + await call_remix(harness, "video_plain", body={"prompt": "x"}) + + # like video_content, remix stops at provider_from_id with no 'openai' default. + assert harness.processor_data() == {"prompt": "x", "video_id": "video_plain"} + + +# =========================================================================== # +# POST /v1/videos/characters - video_create_character # +# =========================================================================== # + + +async def call_create_character(harness: Harness, *, body, video=None, name="my_char"): + harness.read_body.return_value = body + return await endpoints.video_create_character( + request=FakeRequest(), + fastapi_response=Response(), + video=video if video is not None else MagicMock(name="video_upload"), + name=name, + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_create_character__video_attached_default_provider_no_encode(harness): + upload = MagicMock(name="video_upload") + + resp = await call_create_character(harness, body={"prompt": "x"}, video=upload) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_create_character" + harness.batch_to_bytesio.assert_called_once_with([upload]) + # no target_model_names -> no model injected, no id re-encoding. + assert harness.processor_data() == { + "prompt": "x", + "video": b"filebytes", + "custom_llm_provider": "openai", + } + + +@pytest.mark.asyncio +async def test_create_character__target_model_sets_model_and_encodes_id(harness): + harness.base_process.return_value = {"id": "char_raw"} + + resp = await call_create_character( + harness, + body={"target_model_names": "azure-sora-model", "custom_llm_provider": "azure"}, + ) + + data = harness.processor_data() + assert data["model"] == "azure-sora-model" + assert data["custom_llm_provider"] == "azure" + # response id re-encoded with the resolved provider + model for the round-trip. + assert resp["id"] == encode_character_id_with_provider( + "char_raw", "azure", "azure-sora-model" + ) + + +# =========================================================================== # +# GET /v1/videos/characters/{character_id} - video_get_character # +# =========================================================================== # + + +async def call_get_character( + harness: Harness, character_id: str, *, headers=None, query=None +): + return await endpoints.video_get_character( + character_id=character_id, + request=FakeRequest(headers=headers, query=query), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_get_character__encoded_id_full_contract(harness): + harness.base_process.return_value = {"id": "char_raw2"} + + resp = await call_get_character(harness, AZURE_CHARACTER_ID) + + assert harness.route_type() == "avideo_get_character" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + # character_id decoded to its inner value; provider/model from the encoded id. + assert harness.processor_data() == { + "character_id": "char_orig", + "custom_llm_provider": "azure", + "model": "azure-sora", + } + # response id re-encoded for the client round-trip. + assert resp["id"] == encode_character_id_with_provider( + "char_raw2", "azure", VIDEO_MODEL_ID + ) + + +@pytest.mark.asyncio +async def test_get_character__plain_id_defaults_openai_no_encode(harness): + harness.base_process.return_value = {"id": "char_raw3"} + + resp = await call_get_character(harness, "char_plain") + + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "character_id": "char_plain", + "custom_llm_provider": "openai", + } + # id does not start with 'character_' -> returned untouched. + assert resp["id"] == "char_raw3" + + +# =========================================================================== # +# POST /v1/videos/extensions - video_extension # +# =========================================================================== # + + +async def call_extension(harness: Harness, *, body, headers=None, query=None): + return await endpoints.video_extension( + request=FakeRequest(headers=headers, query=query, raw_body=orjson.dumps(body)), + fastapi_response=Response(), + user_api_key_dict=_user(), + ) + + +@pytest.mark.asyncio +async def test_extension__extracts_nested_video_id_full_contract(harness): + resp = await call_extension( + harness, body={"prompt": "continue", "video": {"id": AZURE_VIDEO_ID}} + ) + + assert resp is SENTINEL + assert harness.route_type() == "avideo_extension" + harness.resolve_model.assert_called_once_with(VIDEO_MODEL_ID) + assert harness.processor_data() == { + "prompt": "continue", + "video_id": AZURE_VIDEO_ID, + "custom_llm_provider": "azure", + "model": "azure-sora", + } diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py new file mode 100644 index 00000000000..ae22ae233b5 --- /dev/null +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -0,0 +1,184 @@ +""" +Pure-logic contract tests for litellm/proxy/video_endpoints/utils.py + +Three helpers the video proxy endpoints lean on: + - extract_model_from_target_model_names: first model from a comma string / list + - get_custom_provider_from_data: provider precedence (top-level > extra_body) + - encode_character_id_in_response: re-encode a response id in place + +Every test asserts the exact result (or identity), so a mutation that flips a +branch, drops a strip/filter, or changes precedence fails. The only collaborator +is encode_character_id_with_provider, which runs for real; encoding assertions +are checked by the genuine decode round-trip. +""" + +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.proxy.video_endpoints.utils import ( + encode_character_id_in_response, + extract_model_from_target_model_names, + get_custom_provider_from_data, +) +from litellm.types.videos.utils import ( + decode_character_id_with_provider, + encode_character_id_with_provider, +) + +# =========================================================================== # +# extract_model_from_target_model_names +# =========================================================================== # + + +@pytest.mark.parametrize( + "value,expected", + [ + ("m1,m2,m3", "m1"), + (" a , b ", "a"), # leading/trailing whitespace stripped + (",, m1 ,,", "m1"), # empty tokens filtered out + ("solo", "solo"), # single token, no comma + ("", None), # empty string -> no tokens + (" , , ", None), # only separators/whitespace -> no tokens + (["x", "y"], "x"), # list -> first element + ([], None), # empty list + ], +) +def test_extract_model__str_and_list(value, expected): + assert extract_model_from_target_model_names(value) == expected + + +@pytest.mark.parametrize("value", [None, 123, {"a": 1}, 4.5]) +def test_extract_model__non_str_non_list_is_none(value): + assert extract_model_from_target_model_names(value) is None + + +# =========================================================================== # +# get_custom_provider_from_data +# =========================================================================== # + + +def test_provider__top_level_wins_over_extra_body(): + data = { + "custom_llm_provider": "azure", + "extra_body": {"custom_llm_provider": "openai"}, + } + assert get_custom_provider_from_data(data) == "azure" + + +@pytest.mark.parametrize("falsy", ["", None]) +def test_provider__falsy_top_level_falls_through_to_extra_body(falsy): + data = { + "custom_llm_provider": falsy, + "extra_body": {"custom_llm_provider": "vertex_ai"}, + } + assert get_custom_provider_from_data(data) == "vertex_ai" + + +def test_provider__from_extra_body_dict(): + assert ( + get_custom_provider_from_data( + {"extra_body": {"custom_llm_provider": "bedrock"}} + ) + == "bedrock" + ) + + +def test_provider__from_extra_body_json_string(): + data = {"extra_body": '{"custom_llm_provider": "gemini"}'} + assert get_custom_provider_from_data(data) == "gemini" + + +def test_provider__invalid_json_string_is_none(): + assert get_custom_provider_from_data({"extra_body": "not-json{"}) is None + + +def test_provider__json_string_parsing_to_non_dict_is_none(): + # parses to a list, not a dict -> no provider extracted. + assert get_custom_provider_from_data({"extra_body": "[1, 2]"}) is None + + +def test_provider__extra_body_provider_not_a_string_is_none(): + assert ( + get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) + is None + ) + + +@pytest.mark.parametrize( + "data", + [ + {}, + {"extra_body": {}}, + {"extra_body": 5}, # non-dict, non-str + {"extra_body": {"other": "x"}}, # dict without provider key + ], +) +def test_provider__no_provider_anywhere_is_none(data): + assert get_custom_provider_from_data(data) is None + + +# =========================================================================== # +# encode_character_id_in_response +# =========================================================================== # + + +class _Resp: + """Minimal response object exposing an `id` attribute.""" + + +def test_encode__dict_with_id_mutates_in_place_and_preserves_other_keys(): + response = {"id": "char_raw", "object": "character", "name": "hero"} + + out = encode_character_id_in_response(response, "azure", "model-1") + + assert out is response # same dict, mutated in place + assert out["object"] == "character" and out["name"] == "hero" + assert out["id"] == encode_character_id_with_provider( + "char_raw", "azure", "model-1" + ) + decoded = decode_character_id_with_provider(out["id"]) + assert decoded["custom_llm_provider"] == "azure" + assert decoded["model_id"] == "model-1" + assert decoded["character_id"] == "char_raw" + + +@pytest.mark.parametrize("response", [{}, {"id": ""}, {"id": None}]) +def test_encode__dict_without_usable_id_unchanged(response): + snapshot = dict(response) + out = encode_character_id_in_response(response, "azure", "model-1") + assert out == snapshot + + +def test_encode__object_with_str_id(): + resp = _Resp() + resp.id = "char_raw" + + out = encode_character_id_in_response(resp, "openai", None) + + assert out is resp + assert resp.id == encode_character_id_with_provider("char_raw", "openai", None) + decoded = decode_character_id_with_provider(resp.id) + assert decoded["custom_llm_provider"] == "openai" + assert decoded["character_id"] == "char_raw" + + +@pytest.mark.parametrize("bad_id", [None, 123, ""]) +def test_encode__object_non_str_or_empty_id_unchanged(bad_id): + resp = _Resp() + resp.id = bad_id + + out = encode_character_id_in_response(resp, "azure", "model-1") + + assert out is resp + assert resp.id == bad_id # untouched + + +def test_encode__object_without_id_attr_returned_unchanged(): + resp = _Resp() + out = encode_character_id_in_response(resp, "azure", "model-1") + assert out is resp + assert not hasattr(resp, "id") diff --git a/tests/test_litellm/videos/__init__.py b/tests/test_litellm/videos/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/videos/test_main.py b/tests/test_litellm/videos/test_main.py new file mode 100644 index 00000000000..a04a89ded99 --- /dev/null +++ b/tests/test_litellm/videos/test_main.py @@ -0,0 +1,458 @@ +""" +Dispatch-contract tests for litellm/videos/main.py + +Each public video operation is a pair: a sync `video_*` worker (decorated with +@client) that resolves the provider, fetches the provider config, logs, and then +forwards to exactly one `base_llm_http_handler.video_*_handler`; and an async +`avideo_*` wrapper that delegates to the sync worker in an executor. + +This file locks the contract of that layer so a regression fails loudly: + + 1. DISPATCH - the one correct handler fired and every sibling video handler + asserted NOT called. A copy-paste that calls the wrong handler + (e.g. remix -> edit) flips this. + 2. RESULT - the handler's return value is propagated by identity. + 3. PROVIDER - custom_llm_provider is decoded from an encoded video id when not + passed (status/content/remix/edit/extension), or defaults to + "openai" (list/create_character/get_character). This is the exact + surface of the historical "content defaulted to openai" bug. + 4. PAYLOAD - the provider config object and the operation's identifying args + (video_id/prompt/name/...) reach the handler; _is_async is False + on the sync path. + 5. SHORT-CIRCUIT - mock_response returns a typed object without any handler call. + 6. UNSUPPORTED - a None provider config raises before any handler fires. + 7. DELEGATION - avideo_* returns the sync worker's result untouched, sets + async_call=True, and pre-resolves the provider where it must. + +Seams mocked: the http handler (network), the provider-config registry lookup, +get_llm_provider, and the video-generation optional-param builders. The id decode +helper runs for real against genuinely-encoded ids, so the provider assertions +reflect production. +""" + +import os +import sys +from contextlib import ExitStack +from dataclasses import dataclass +from typing import Any, Dict +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.types.videos.main import CharacterObject, VideoObject +from litellm.types.videos.utils import encode_video_id_with_provider +from litellm.videos import main as videos_main + +# A real model-encoded video id: decodes (for real) to provider "azure". Used to +# prove the sync workers derive custom_llm_provider from the id, not a hardcode. +AZURE_VIDEO_ID = encode_video_id_with_provider("video_raw", "azure", "deployment-1") + +# The nine sync handlers on base_llm_http_handler. Dispatch tests assert exactly +# one fired and the other eight did not. +SYNC_HANDLERS = ( + "video_generation_handler", + "video_content_handler", + "video_remix_handler", + "video_create_character_handler", + "video_get_character_handler", + "video_edit_handler", + "video_extension_handler", + "video_list_handler", + "video_status_handler", +) + +GEN_OPTIONAL_PARAMS = {"seconds": "8", "size": "720x1280"} + + +@dataclass +class Seams: + handler: MagicMock + get_config: MagicMock + config: MagicMock + + def kwargs_of(self, handler_name: str) -> Dict[str, Any]: + method = getattr(self.handler, handler_name) + assert method.call_count == 1 + return dict(method.call_args.kwargs) + + def assert_only(self, handler_name: str) -> None: + for name in SYNC_HANDLERS: + method = getattr(self.handler, name) + if name == handler_name: + method.assert_called_once() + else: + method.assert_not_called() + + +@pytest.fixture +def seams(): + handler = MagicMock(spec=BaseLLMHTTPHandler) + config = MagicMock(name="provider_video_config") + get_config = MagicMock(return_value=config) + + with ExitStack() as stack: + stack.enter_context(patch.object(videos_main, "base_llm_http_handler", handler)) + stack.enter_context( + patch.object( + videos_main.ProviderConfigManager, + "get_provider_video_config", + get_config, + ) + ) + # video_generation resolves model+provider through get_llm_provider and + # builds optional params; mock those so the dispatch payload is deterministic. + stack.enter_context( + patch.object( + videos_main, + "get_llm_provider", + MagicMock(return_value=("sora-2", "openai", None, None)), + ) + ) + stack.enter_context( + patch.object( + videos_main.VideoGenerationRequestUtils, + "get_requested_video_generation_optional_param", + MagicMock(return_value={"seconds": "8"}), + ) + ) + stack.enter_context( + patch.object( + videos_main.VideoGenerationRequestUtils, + "get_optional_params_video_generation", + MagicMock(return_value=dict(GEN_OPTIONAL_PARAMS)), + ) + ) + yield Seams(handler=handler, get_config=get_config, config=config) + + +# =========================================================================== # +# Dispatch contract - one rich test per sync worker. +# =========================================================================== # + + +def test_video_generation__dispatch(seams): + result = videos_main.video_generation(prompt="a sunset", model="sora-2") + + seams.assert_only("video_generation_handler") + assert result is seams.handler.video_generation_handler.return_value + kw = seams.kwargs_of("video_generation_handler") + assert kw["model"] == "sora-2" + assert kw["prompt"] == "a sunset" + assert kw["custom_llm_provider"] == "openai" + assert kw["video_generation_provider_config"] is seams.config + assert kw["video_generation_optional_request_params"] == GEN_OPTIONAL_PARAMS + assert kw["_is_async"] is False + + +def test_video_status__dispatch_and_provider_from_id(seams): + result = videos_main.video_status(video_id=AZURE_VIDEO_ID) + + seams.assert_only("video_status_handler") + assert result is seams.handler.video_status_handler.return_value + kw = seams.kwargs_of("video_status_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["custom_llm_provider"] == "azure" # decoded from the id, not openai + assert kw["video_status_provider_config"] is seams.config + assert kw["_is_async"] is False + # provider config requested for the decoded provider, not a hardcode. + assert seams.get_config.call_args.kwargs["provider"] == litellm.LlmProviders.AZURE + + +def test_video_content__dispatch_and_provider_from_id(seams): + result = videos_main.video_content(video_id=AZURE_VIDEO_ID, variant="thumbnail") + + seams.assert_only("video_content_handler") + assert result is seams.handler.video_content_handler.return_value + kw = seams.kwargs_of("video_content_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["custom_llm_provider"] == "azure" + assert kw["variant"] == "thumbnail" + assert kw["video_content_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_content__plain_id_defaults_to_openai(seams): + videos_main.video_content(video_id="video_plain") + + assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai" + + +def test_video_remix__dispatch_and_provider_from_id(seams): + result = videos_main.video_remix(video_id=AZURE_VIDEO_ID, prompt="new colors") + + seams.assert_only("video_remix_handler") + assert result is seams.handler.video_remix_handler.return_value + kw = seams.kwargs_of("video_remix_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["prompt"] == "new colors" + assert kw["custom_llm_provider"] == "azure" + assert kw["video_remix_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_edit__dispatch_and_provider_from_id(seams): + result = videos_main.video_edit(video_id=AZURE_VIDEO_ID, prompt="brighter") + + seams.assert_only("video_edit_handler") + assert result is seams.handler.video_edit_handler.return_value + kw = seams.kwargs_of("video_edit_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["prompt"] == "brighter" + assert kw["custom_llm_provider"] == "azure" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_extension__dispatch_and_provider_from_id(seams): + result = videos_main.video_extension( + video_id=AZURE_VIDEO_ID, prompt="continue", seconds="5" + ) + + seams.assert_only("video_extension_handler") + assert result is seams.handler.video_extension_handler.return_value + kw = seams.kwargs_of("video_extension_handler") + assert kw["video_id"] == AZURE_VIDEO_ID + assert kw["prompt"] == "continue" + assert kw["seconds"] == "5" + assert kw["custom_llm_provider"] == "azure" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_list__dispatch_defaults_to_openai(seams): + result = videos_main.video_list(after="cur", limit=5, order="desc") + + seams.assert_only("video_list_handler") + assert result is seams.handler.video_list_handler.return_value + kw = seams.kwargs_of("video_list_handler") + assert kw["after"] == "cur" + assert kw["limit"] == 5 + assert kw["order"] == "desc" + assert kw["custom_llm_provider"] == "openai" + assert kw["video_list_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_create_character__dispatch_defaults_to_openai(seams): + video = MagicMock(name="video_upload") + result = videos_main.video_create_character(name="hero", video=video) + + seams.assert_only("video_create_character_handler") + assert result is seams.handler.video_create_character_handler.return_value + kw = seams.kwargs_of("video_create_character_handler") + assert kw["name"] == "hero" + assert kw["video"] is video + assert kw["custom_llm_provider"] == "openai" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_video_get_character__dispatch_defaults_to_openai(seams): + result = videos_main.video_get_character(character_id="char_1") + + seams.assert_only("video_get_character_handler") + assert result is seams.handler.video_get_character_handler.return_value + kw = seams.kwargs_of("video_get_character_handler") + assert kw["character_id"] == "char_1" + assert kw["custom_llm_provider"] == "openai" + assert kw["video_provider_config"] is seams.config + assert kw["_is_async"] is False + + +def test_explicit_provider_beats_decoded_id(seams): + """An explicit custom_llm_provider wins over the one encoded in the id.""" + videos_main.video_status(video_id=AZURE_VIDEO_ID, custom_llm_provider="vertex_ai") + + assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "vertex_ai" + + +# =========================================================================== # +# mock_response short-circuit - returns a typed object, no handler call. +# =========================================================================== # + + +def test_generation__mock_response_short_circuits(seams): + resp = videos_main.video_generation( + prompt="x", + model="sora-2", + mock_response={"id": "v1", "object": "video", "status": "queued"}, + ) + + assert isinstance(resp, VideoObject) + assert resp.id == "v1" + seams.handler.video_generation_handler.assert_not_called() + + +def test_list__mock_response_short_circuits(seams): + resp = videos_main.video_list( + mock_response=[{"id": "v1", "object": "video", "status": "completed"}] + ) + + assert isinstance(resp, list) + assert resp[0].id == "v1" + seams.handler.video_list_handler.assert_not_called() + + +def test_get_character__mock_response_short_circuits(seams): + resp = videos_main.video_get_character( + character_id="char_1", + mock_response={ + "id": "char_1", + "object": "character", + "created_at": 1, + "name": "hero", + }, + ) + + assert isinstance(resp, CharacterObject) + assert resp.id == "char_1" + seams.handler.video_get_character_handler.assert_not_called() + + +# =========================================================================== # +# Unsupported provider - a None provider config raises before any dispatch. +# =========================================================================== # + + +def test_unsupported_provider_raises_without_dispatch(seams): + seams.get_config.return_value = None + + with pytest.raises(Exception): + videos_main.video_status(video_id=AZURE_VIDEO_ID) + + seams.handler.video_status_handler.assert_not_called() + + +# =========================================================================== # +# Async-wrapper delegation - representative coverage. +# =========================================================================== # + + +@pytest.mark.asyncio +async def test_avideo_generation__delegates_with_async_flag(): + sentinel = VideoObject(id="v-async", object="video", status="queued") + with ( + patch.object( + videos_main, "video_generation", MagicMock(return_value=sentinel) + ) as sync, + patch.object( + litellm, + "get_llm_provider", + MagicMock(return_value=("sora-2", "openai", None, None)), + ), + ): + result = await videos_main.avideo_generation(prompt="x", model="sora-2") + + assert result is sentinel + assert sync.call_args.kwargs["async_call"] is True + assert sync.call_args.kwargs["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_avideo_status__delegates_untouched(): + sentinel = VideoObject(id="v-async", object="video", status="queued") + with patch.object( + videos_main, "video_status", MagicMock(return_value=sentinel) + ) as sync: + result = await videos_main.avideo_status(video_id="video_plain") + + assert result is sentinel + assert sync.call_args.kwargs["async_call"] is True + assert sync.call_args.kwargs["video_id"] == "video_plain" + + +@pytest.mark.asyncio +async def test_avideo_content__pre_decodes_provider_before_delegating(): + """avideo_content resolves the provider from the encoded id itself before + handing off, so the sync worker receives the decoded provider, not None.""" + sentinel = b"mp4-bytes" + with patch.object( + videos_main, "video_content", MagicMock(return_value=sentinel) + ) as sync: + result = await videos_main.avideo_content(video_id=AZURE_VIDEO_ID) + + assert result is sentinel + assert sync.call_args.kwargs["async_call"] is True + assert sync.call_args.kwargs["custom_llm_provider"] == "azure" + + +# =========================================================================== # +# Credential passthrough - DB/YAML model-config credentials the router injects +# via kwargs must reach the provider call for EVERY video handler, carried in +# litellm_params. Distinct per-field values catch a cross-wired field. +# =========================================================================== # + +DB_YAML_CREDS = { + "api_key": "sk-db-credential", + "api_base": "https://db-resource.test", + "api_version": "2024-12-31", + "vertex_project": "db-project-xyz", +} + +CREDENTIAL_OPERATIONS = [ + ( + "video_generation_handler", + lambda: videos_main.video_generation( + prompt="p", model="sora-2", **DB_YAML_CREDS + ), + ), + ( + "video_status_handler", + lambda: videos_main.video_status(video_id=AZURE_VIDEO_ID, **DB_YAML_CREDS), + ), + ( + "video_content_handler", + lambda: videos_main.video_content(video_id=AZURE_VIDEO_ID, **DB_YAML_CREDS), + ), + ( + "video_remix_handler", + lambda: videos_main.video_remix( + video_id=AZURE_VIDEO_ID, prompt="p", **DB_YAML_CREDS + ), + ), + ( + "video_edit_handler", + lambda: videos_main.video_edit( + video_id=AZURE_VIDEO_ID, prompt="p", **DB_YAML_CREDS + ), + ), + ( + "video_extension_handler", + lambda: videos_main.video_extension( + video_id=AZURE_VIDEO_ID, prompt="p", seconds="5", **DB_YAML_CREDS + ), + ), + ( + "video_list_handler", + lambda: videos_main.video_list(**DB_YAML_CREDS), + ), + ( + "video_create_character_handler", + lambda: videos_main.video_create_character( + name="hero", video=MagicMock(name="vid"), **DB_YAML_CREDS + ), + ), + ( + "video_get_character_handler", + lambda: videos_main.video_get_character(character_id="char_1", **DB_YAML_CREDS), + ), +] + + +@pytest.mark.parametrize( + "handler_name,invoke", + CREDENTIAL_OPERATIONS, + ids=[op[0] for op in CREDENTIAL_OPERATIONS], +) +def test_db_yaml_credentials_reach_every_handler(seams, handler_name, invoke): + invoke() + + litellm_params = seams.kwargs_of(handler_name)["litellm_params"] + assert litellm_params.get("api_key") == DB_YAML_CREDS["api_key"] + assert litellm_params.get("api_base") == DB_YAML_CREDS["api_base"] + assert litellm_params.get("api_version") == DB_YAML_CREDS["api_version"] + assert litellm_params.get("vertex_project") == DB_YAML_CREDS["vertex_project"] diff --git a/tests/test_litellm/videos/test_utils.py b/tests/test_litellm/videos/test_utils.py new file mode 100644 index 00000000000..09975829531 --- /dev/null +++ b/tests/test_litellm/videos/test_utils.py @@ -0,0 +1,196 @@ +""" +Pure-logic contract tests for litellm/videos/main.py's request utils +(litellm/videos/utils.py: VideoGenerationRequestUtils). + +These lock the exact param-shaping behavior so a mutation that drops a filter, +flips a precedence, or stops removing a key fails loudly. The only seam is the +provider config's map_openai_params (a provider boundary); filter_out_litellm_params +runs for real, so the "litellm-internal params get stripped" assertions reflect +production. Every test asserts the exact resulting dict, never "ran without error". +""" + +import os +import sys +from unittest.mock import MagicMock + + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm.videos.utils import VideoGenerationRequestUtils + +get_requested = ( + VideoGenerationRequestUtils.get_requested_video_generation_optional_param +) +get_optional = VideoGenerationRequestUtils.get_optional_params_video_generation + + +# =========================================================================== # +# get_requested_video_generation_optional_param +# +# Receives the caller's full local_vars; must return only the API-bound optional +# params. filter_out_litellm_params strips known internal keys for real; the +# values used below were chosen against the live set: seconds/size/user/foo_param/ +# vertex_project/extra/a/b survive, api_key/metadata/litellm_* are stripped. +# =========================================================================== # + + +def test_requested__drops_none_and_excluded_keys(): + result = get_requested( + { + "seconds": "8", + "size": None, # None -> dropped + "prompt": "a sunset", # excluded + "model": "sora-2", # excluded + "user": "u1", + } + ) + assert result == {"seconds": "8", "user": "u1"} + + +def test_requested__strips_litellm_internal_params(): + result = get_requested( + { + "seconds": "8", + "api_key": "sk-secret", + "metadata": {"x": 1}, + "litellm_call_id": "id-123", + } + ) + assert result == {"seconds": "8"} + + +def test_requested__timeout_always_removed(): + # timeout is NOT a litellm-internal param, so only the explicit pop removes it. + result = get_requested({"seconds": "8", "timeout": 30}) + assert result == {"seconds": "8"} + + +def test_requested__nested_kwargs_merge_and_override_base(): + result = get_requested( + {"seconds": "8", "kwargs": {"size": "720x1280", "seconds": "override"}} + ) + # nested kwargs win over the top-level base params on collision. + assert result == {"seconds": "override", "size": "720x1280"} + + +def test_requested__non_dict_kwargs_treated_as_empty(): + result = get_requested({"seconds": "8", "kwargs": "not-a-dict"}) + assert result == {"seconds": "8"} + + +def test_requested__none_input_returns_empty(): + assert get_requested(None) == {} + + +def test_requested__top_level_extra_body_spread_and_preserved(): + result = get_requested( + {"seconds": "8", "extra_body": {"vertex_project": "proj", "foo_param": "bar"}} + ) + # extra_body keys are both spread at top level AND kept under "extra_body". + assert result == { + "seconds": "8", + "vertex_project": "proj", + "foo_param": "bar", + "extra_body": {"vertex_project": "proj", "foo_param": "bar"}, + } + + +def test_requested__extra_body_kwargs_overrides_top_level(): + result = get_requested( + { + "extra_body": {"a": "top", "b": "top_b"}, + "kwargs": {"extra_body": {"a": "kw"}}, + } + ) + # kwargs' extra_body wins over the top-level extra_body on collision; the + # non-colliding top-level key survives. + assert result == { + "a": "kw", + "b": "top_b", + "extra_body": {"a": "kw", "b": "top_b"}, + } + + +def test_requested__extra_body_strips_litellm_internal_params(): + result = get_requested({"extra_body": {"api_key": "sk", "foo_param": "bar"}}) + # api_key filtered out of extra_body; only foo_param remains (and is spread). + assert result == {"foo_param": "bar", "extra_body": {"foo_param": "bar"}} + + +def test_requested__empty_extra_body_not_added(): + result = get_requested({"seconds": "8", "extra_body": {}}) + assert result == {"seconds": "8"} + assert "extra_body" not in result + + +# =========================================================================== # +# get_optional_params_video_generation +# +# Delegates mapping to the provider config (the seam) then folds extra_body in. +# =========================================================================== # + + +def _config(map_return): + config = MagicMock() + config.map_openai_params.return_value = map_return + return config + + +def test_optional__delegates_to_map_openai_params_with_drop_params(): + config = _config({"seconds": "8"}) + optional_params = {"seconds": "8"} + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params=optional_params, + ) + + assert result == {"seconds": "8"} + config.map_openai_params.assert_called_once_with( + video_create_optional_params=optional_params, + model="sora-2", + drop_params=litellm.drop_params, + ) + + +def test_optional__extra_body_overrides_mapped_and_is_removed(): + # mapped output carries a leftover extra_body that must be popped; the input + # extra_body overrides a colliding mapped key and is spread in. + config = _config({"seconds": "8", "size": "mapped", "extra_body": {"leftover": 1}}) + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params={ + "extra_body": {"size": "override", "extra": "x"} + }, + ) + + assert result == {"seconds": "8", "size": "override", "extra": "x"} + assert "extra_body" not in result + + +def test_optional__no_extra_body_returns_mapped_unchanged(): + config = _config({"seconds": "8"}) + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params={"seconds": "8"}, + ) + + assert result == {"seconds": "8"} + + +def test_optional__non_dict_extra_body_ignored(): + config = _config({"seconds": "8"}) + + result = get_optional( + model="sora-2", + video_generation_provider_config=config, + video_generation_optional_params={"seconds": "8", "extra_body": None}, + ) + + assert result == {"seconds": "8"}