diff --git a/litellm/llms/xai/videos/transformation.py b/litellm/llms/xai/videos/transformation.py index bbb207fcc56..b871f4a6011 100644 --- a/litellm/llms/xai/videos/transformation.py +++ b/litellm/llms/xai/videos/transformation.py @@ -7,6 +7,7 @@ from httpx._types import RequestFiles import litellm from litellm.constants import XAI_API_BASE from litellm.exceptions import AuthenticationError +from litellm.litellm_core_utils.url_utils import async_safe_get, encode_url_path_segment, safe_get from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, @@ -248,6 +249,13 @@ class XAIVideoConfig(BaseVideoConfig): video_obj._hidden_params["video_url"] = None return video_obj + def _video_resource_url(self, api_base: str, video_id: str) -> str: + encoded_video_id: Final = encode_url_path_segment( + extract_original_video_id(video_id), + field_name="video_id", + ) + return f"{self._v1_root(api_base)}/videos/{encoded_video_id}" + def transform_video_status_retrieve_request( self, video_id: str, @@ -255,8 +263,7 @@ class XAIVideoConfig(BaseVideoConfig): litellm_params: GenericLiteLLMParams, headers: dict, ) -> tuple[str, dict]: - original_id: Final = extract_original_video_id(video_id) - return f"{self._v1_root(api_base)}/videos/{original_id}", {} + return self._video_resource_url(api_base, video_id), {} def transform_video_status_retrieve_response( self, @@ -303,8 +310,7 @@ class XAIVideoConfig(BaseVideoConfig): headers: dict, variant: str | None = None, ) -> tuple[str, dict]: - original_id: Final = extract_original_video_id(video_id) - return f"{self._v1_root(api_base)}/videos/{original_id}", {} + return self._video_resource_url(api_base, video_id), {} def _video_cdn_url(self, raw_response: httpx.Response) -> str | None: content_type: Final = (raw_response.headers.get("content-type") or "").lower() @@ -328,7 +334,7 @@ class XAIVideoConfig(BaseVideoConfig): if url is None: return raw_response.content httpx_client: Final[HTTPHandler] = _get_httpx_client() - video_response: Final = httpx_client.get(url) + video_response: Final = safe_get(httpx_client, url) video_response.raise_for_status() return video_response.content @@ -343,7 +349,7 @@ class XAIVideoConfig(BaseVideoConfig): async_httpx_client: Final[AsyncHTTPHandler] = get_async_httpx_client( llm_provider=litellm.LlmProviders.XAI, ) - video_response: Final = await async_httpx_client.get(url) + video_response: Final = await async_safe_get(async_httpx_client, url) video_response.raise_for_status() return video_response.content diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 9cc1f033326..6f13d88341f 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -20,6 +20,7 @@ from litellm.proxy.video_endpoints.utils import ( encode_character_id_in_response, extract_model_from_target_model_names, get_custom_provider_from_data, + infer_video_provider_from_model, resolve_video_request_model, video_reference_to_id, ) @@ -253,15 +254,12 @@ async def video_status( provider_from_id: Final = decoded.get("custom_llm_provider") model_id_from_decoded: Final = decoded.get("model_id") - custom_llm_provider: Final = ( + explicit_provider: Final = ( get_custom_llm_provider_from_request_headers(request=request) or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) or provider_from_id - or "openai" ) - if custom_llm_provider: - data["custom_llm_provider"] = custom_llm_provider resolved_model: Final = resolve_video_request_model( model_id_from_decoded=model_id_from_decoded, @@ -271,6 +269,12 @@ async def video_status( if resolved_model: data["model"] = resolved_model + custom_llm_provider: Final = ( + explicit_provider or infer_video_provider_from_model(resolved_model) or "openai" + ) + if custom_llm_provider: + data["custom_llm_provider"] = custom_llm_provider + # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index af0e9548568..7d0357371d8 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -9,6 +9,15 @@ class VideoModelIdResolver(Protocol): def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: ... +def infer_video_provider_from_model(model: str | None) -> str | None: + if not isinstance(model, str) or not model: + return None + unprefixed: Final = model.split("/", 1)[-1] + if unprefixed.startswith("grok-imagine-video"): + return "xai" + return None + + def resolve_video_request_model( *, model_id_from_decoded: str | None, diff --git a/tests/test_litellm/llms/xai/test_xai_video_generation.py b/tests/test_litellm/llms/xai/test_xai_video_generation.py index fee57425264..261ce2e88c0 100644 --- a/tests/test_litellm/llms/xai/test_xai_video_generation.py +++ b/tests/test_litellm/llms/xai/test_xai_video_generation.py @@ -3,6 +3,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +from litellm.litellm_core_utils.url_utils import SSRFError from litellm.llms.xai.videos.transformation import XAIVideoConfig from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders @@ -141,6 +142,50 @@ def test_content_request_is_get_status(): assert params == {} +@pytest.mark.parametrize("video_id", ["..", ""]) +@pytest.mark.parametrize( + "transform_name", + ["transform_video_status_retrieve_request", "transform_video_content_request"], +) +def test_video_id_rejects_empty_and_dot_path_segments(video_id, transform_name): + transform = getattr(XAIVideoConfig(), transform_name) + with pytest.raises(ValueError, match="video_id"): + transform( + video_id=video_id, + api_base="https://api.x.ai/v1", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + +@pytest.mark.parametrize( + "transform_name", + ["transform_video_status_retrieve_request", "transform_video_content_request"], +) +def test_video_id_parent_path_is_one_encoded_segment(transform_name): + transform = getattr(XAIVideoConfig(), transform_name) + url, params = transform( + video_id="../models", + api_base="https://api.x.ai/v1", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert url == "https://api.x.ai/v1/videos/..%2Fmodels" + assert "/videos/../" not in url + assert params == {} + + +def test_video_id_is_percent_encoded_as_one_segment(): + url, params = XAIVideoConfig().transform_video_status_retrieve_request( + video_id="req-123?x=1#frag", + api_base="https://api.x.ai/v1", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert url == "https://api.x.ai/v1/videos/req-123%3Fx%3D1%23frag" + assert params == {} + + def test_content_response_fetches_cdn_via_shared_client(): status = httpx.Response( 200, @@ -151,14 +196,16 @@ def test_content_response_fetches_cdn_via_shared_client(): video_resp = MagicMock() video_resp.content = b"mp4-bytes" video_resp.raise_for_status.return_value = None - cdn.get.return_value = video_resp with patch( "litellm.llms.xai.videos.transformation._get_httpx_client", return_value=cdn, - ): + ), patch( + "litellm.llms.xai.videos.transformation.safe_get", + return_value=video_resp, + ) as safe_get: body = XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) assert body == b"mp4-bytes" - cdn.get.assert_called_once_with("https://vidgen.x.ai/x.mp4") + safe_get.assert_called_once_with(cdn, "https://vidgen.x.ai/x.mp4") @pytest.mark.asyncio @@ -172,19 +219,21 @@ async def test_async_content_response_does_not_use_sync_client(): video_resp = MagicMock() video_resp.content = b"async-mp4" video_resp.raise_for_status.return_value = None - async_client.get = AsyncMock(return_value=video_resp) with patch( "litellm.llms.xai.videos.transformation.get_async_httpx_client", return_value=async_client, ), patch( "litellm.llms.xai.videos.transformation._get_httpx_client", - ) as sync_client: + ) as sync_client, patch( + "litellm.llms.xai.videos.transformation.async_safe_get", + new=AsyncMock(return_value=video_resp), + ) as async_safe_get: body = await XAIVideoConfig().async_transform_video_content_response( status, logging_obj=MagicMock() ) assert body == b"async-mp4" sync_client.assert_not_called() - async_client.get.assert_awaited_once_with("https://vidgen.x.ai/x.mp4") + async_safe_get.assert_awaited_once_with(async_client, "https://vidgen.x.ai/x.mp4") def test_content_response_raises_when_status_has_no_url(): @@ -204,3 +253,85 @@ def test_content_response_returns_raw_bytes_when_not_json(): content=b"already-mp4", ) assert XAIVideoConfig().transform_video_content_response(raw, logging_obj=MagicMock()) == b"already-mp4" + + +def test_content_response_rejects_internal_cdn_host(): + status = httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"status": "done", "video": {"url": "http://127.0.0.1/secret.mp4"}}, + ) + + def boom(*args, **kwargs): + raise AssertionError("unsafe CDN fetch must not run") + + with patch( + "litellm.llms.xai.videos.transformation._get_httpx_client", + return_value=MagicMock(get=boom), + ): + with pytest.raises((SSRFError, ValueError)): + XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) + + +def test_content_response_fetches_public_cdn_via_safe_get(): + status = httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}}, + ) + video_resp = MagicMock() + video_resp.content = b"safe-mp4" + video_resp.raise_for_status.return_value = None + with patch( + "litellm.llms.xai.videos.transformation.safe_get", + return_value=video_resp, + create=True, + ) as safe_get: + body = XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) + assert body == b"safe-mp4" + assert safe_get.call_count == 1 + assert safe_get.call_args.args[1] == "https://vidgen.x.ai/x.mp4" + + +@pytest.mark.asyncio +async def test_async_content_response_rejects_internal_cdn_host(): + status = httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"status": "done", "video": {"url": "http://127.0.0.1/secret.mp4"}}, + ) + + async def boom(*args, **kwargs): + raise AssertionError("unsafe async CDN fetch must not run") + + with patch( + "litellm.llms.xai.videos.transformation.get_async_httpx_client", + return_value=MagicMock(get=boom), + ): + with pytest.raises((SSRFError, ValueError)): + await XAIVideoConfig().async_transform_video_content_response( + status, logging_obj=MagicMock() + ) + + +@pytest.mark.asyncio +async def test_async_content_response_fetches_public_cdn_via_async_safe_get(): + status = httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}}, + ) + video_resp = MagicMock() + video_resp.content = b"async-safe-mp4" + video_resp.raise_for_status.return_value = None + with patch( + "litellm.llms.xai.videos.transformation.async_safe_get", + new=AsyncMock(return_value=video_resp), + create=True, + ) as async_safe_get: + body = await XAIVideoConfig().async_transform_video_content_response( + status, logging_obj=MagicMock() + ) + assert body == b"async-safe-mp4" + async_safe_get.assert_awaited() + assert async_safe_get.call_args.args[1] == "https://vidgen.x.ai/x.mp4" diff --git a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py index eeb59e31988..78877b84f6c 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py @@ -338,11 +338,25 @@ async def test_status__query_model_on_plain_id(harness): harness.resolve_model.assert_not_called() assert harness.processor_data() == { "video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd", - "custom_llm_provider": "openai", + "custom_llm_provider": "xai", "model": "grok-imagine-video-1.5", } +@pytest.mark.asyncio +async def test_status__query_model_grok_imagine_does_not_default_openai_before_inference( + harness, +): + await call_status( + harness, + "video_plain_xai", + query={"model": "grok-imagine-video"}, + ) + + assert harness.processor_data()["custom_llm_provider"] == "xai" + assert harness.processor_data()["model"] == "grok-imagine-video" + + # =========================================================================== # # GET /v1/videos/{video_id}/content - video_content # # =========================================================================== # diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py index 5944fd7a394..af8af6cdcf2 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_utils.py +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -21,6 +21,7 @@ from litellm.proxy.video_endpoints.utils import ( encode_character_id_in_response, extract_model_from_target_model_names, get_custom_provider_from_data, + infer_video_provider_from_model, resolve_video_request_model, video_reference_to_id, ) @@ -75,6 +76,21 @@ def test_resolve_video_request_model__query_model_on_plain_id(): ) +@pytest.mark.parametrize( + "model,expected", + [ + ("grok-imagine-video", "xai"), + ("grok-imagine-video-1.5", "xai"), + ("xai/grok-imagine-video", "xai"), + ("sora-2", None), + (None, None), + ("", None), + ], +) +def test_infer_video_provider_from_model(model, expected): + assert infer_video_provider_from_model(model) == expected + + # =========================================================================== # # extract_model_from_target_model_names # =========================================================================== #