fix(xai): encode video ids and validate CDN download URLs

This commit is contained in:
hx 2026-09-09 14:59:29 +08:00
parent c25883d55a
commit 1d8942f70b
6 changed files with 197 additions and 17 deletions

View file

@ -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

View file

@ -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:

View file

@ -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,

View file

@ -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"

View file

@ -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 #
# =========================================================================== #

View file

@ -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
# =========================================================================== #