mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
fix(xai): encode video ids and validate CDN download URLs
This commit is contained in:
parent
c25883d55a
commit
1d8942f70b
6 changed files with 197 additions and 17 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 #
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue