From 39b3d4346d66865590f0ac6fb16c44a6016f7916 Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Sun, 27 Sep 2026 17:48:31 +0800 Subject: [PATCH 1/6] feat(xai): add grok imagine image generation, edit, and video download Co-authored-by: Cursor --- litellm/images/main.py | 3 +- litellm/llms/xai/image_edit/__init__.py | 3 + litellm/llms/xai/image_edit/transformation.py | 251 +++++++++++ litellm/llms/xai/image_generation/__init__.py | 11 + .../xai/image_generation/transformation.py | 200 +++++++++ litellm/llms/xai/videos/__init__.py | 3 + litellm/llms/xai/videos/transformation.py | 420 ++++++++++++++++++ ...odel_prices_and_context_window_backup.json | 36 +- .../provider_endpoints_support_backup.json | 6 +- litellm/proxy/auth/auth_utils.py | 1 + litellm/proxy/image_endpoints/endpoints.py | 104 +++-- litellm/proxy/video_endpoints/endpoints.py | 56 ++- litellm/proxy/video_endpoints/utils.py | 84 +++- litellm/types/videos/main.py | 1 + litellm/types/videos/utils.py | 10 +- litellm/utils.py | 14 + litellm/videos/main.py | 60 ++- model_prices_and_context_window.json | 36 +- provider_endpoints_support.json | 6 +- .../proxy/auth/test_auth_utils.py | 19 + .../proxy/image_endpoints/test_endpoints.py | 48 ++ .../proxy/video_endpoints/test_endpoints.py | 84 ++++ .../proxy/video_endpoints/test_utils.py | 95 +++- tests/unit/llms/xai/test_xai_image_edit.py | 167 +++++++ .../llms/xai/test_xai_image_generation.py | 129 ++++++ .../llms/xai/test_xai_video_generation.py | 337 ++++++++++++++ tests/unit/videos/test_main.py | 24 + 27 files changed, 2094 insertions(+), 114 deletions(-) create mode 100644 litellm/llms/xai/image_edit/__init__.py create mode 100644 litellm/llms/xai/image_edit/transformation.py create mode 100644 litellm/llms/xai/image_generation/__init__.py create mode 100644 litellm/llms/xai/image_generation/transformation.py create mode 100644 litellm/llms/xai/videos/__init__.py create mode 100644 litellm/llms/xai/videos/transformation.py create mode 100644 tests/unit/llms/xai/test_xai_image_edit.py create mode 100644 tests/unit/llms/xai/test_xai_image_generation.py create mode 100644 tests/unit/llms/xai/test_xai_video_generation.py diff --git a/litellm/images/main.py b/litellm/images/main.py index 7dc68dafecc..fc1ab906f2d 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -384,6 +384,7 @@ def image_generation( litellm.LlmProviders.QWENCLOUD, litellm.LlmProviders.QWEN_AI_PLATFORM, litellm.LlmProviders.EDENAI, + litellm.LlmProviders.XAI, ): if image_generation_config is None: raise ValueError(f"image generation config is not supported for {custom_llm_provider}") @@ -393,7 +394,7 @@ def image_generation( litellm_params_dict["api_base"] = _api_base return llm_http_handler.image_generation_handler( - api_key=api_key, + api_key=api_key or dynamic_api_key, model=model, prompt=prompt, image_generation_provider_config=image_generation_config, diff --git a/litellm/llms/xai/image_edit/__init__.py b/litellm/llms/xai/image_edit/__init__.py new file mode 100644 index 00000000000..79664e14d6b --- /dev/null +++ b/litellm/llms/xai/image_edit/__init__.py @@ -0,0 +1,3 @@ +from .transformation import XAIImageEditConfig + +__all__ = ["XAIImageEditConfig"] # mutable-ok: provider JSON body and base-class dict signature diff --git a/litellm/llms/xai/image_edit/transformation.py b/litellm/llms/xai/image_edit/transformation.py new file mode 100644 index 00000000000..496a0c7e65d --- /dev/null +++ b/litellm/llms/xai/image_edit/transformation.py @@ -0,0 +1,251 @@ +import base64 +from io import BufferedReader, BytesIO +from typing import TYPE_CHECKING, Final + +import httpx +from httpx._types import RequestFiles + +from litellm.constants import XAI_API_BASE +from litellm.exceptions import AuthenticationError +from litellm.images.utils import ImageEditRequestUtils +from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig +from litellm.llms.xai.common_utils import XAIModelInfo +from litellm.secret_managers.main import get_secret_str +from litellm.types.images.main import ImageEditOptionalRequestParams +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import FileTypes, ImageObject, ImageResponse + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +_SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature + "1024x1024": "1:1", + "1792x1024": "16:9", + "1024x1792": "9:16", + "1536x1024": "3:2", + "1024x1536": "2:3", + "1280x720": "16:9", + "720x1280": "9:16", + "1920x1080": "16:9", + "1080x1920": "9:16", +} +_XAI_NATIVE_PARAMS: Final = frozenset({"aspect_ratio", "n", "resolution"}) + + +def _read_seekable(image: BytesIO | BufferedReader) -> bytes: + current_pos: Final = image.tell() + image.seek(0) + data: Final = image.read() + image.seek(current_pos) + return data + + +class XAIImageEditConfig(BaseImageEditConfig): + def get_supported_openai_params( + self, model: str + ) -> list: # mutable-ok: provider JSON body and base-class dict signature + return ["n", "response_format", "size", "user"] # mutable-ok: provider JSON body and base-class dict signature + + def map_openai_params( + self, + image_edit_optional_params: ImageEditOptionalRequestParams, + model: str, + drop_params: bool, + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + supported: Final = frozenset(self.get_supported_openai_params(model)) + allowed: Final = supported | _XAI_NATIVE_PARAMS + raw = image_edit_optional_params + incoming: Final = dict(raw) # mutable-ok: provider JSON body and base-class dict signature + unknown: Final = tuple(key for key in incoming if key not in allowed) + if unknown and not drop_params: + raise ValueError( + f"Parameter {unknown[0]} is not supported for model {model}. " + f"Supported parameters are {sorted(allowed)}. " + "Set drop_params=True to drop unsupported parameters." + ) + + pairs = ((key, value) for key, value in incoming.items() if key in allowed) + mapped: Final = dict(pairs) # mutable-ok: provider JSON body and base-class dict signature + size: Final = mapped.get("size") + aspect_ratio: Final = mapped.get("aspect_ratio") or ( + _SIZE_TO_ASPECT_RATIO.get(str(size), "1:1") if size else None + ) + n: Final = mapped.get("n") + resolution: Final = mapped.get("resolution") + aspect = {"aspect_ratio": aspect_ratio} if aspect_ratio is not None else None # mutable-ok: provider JSON body + count = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body + quality = {"resolution": resolution} if resolution is not None else None # mutable-ok: provider JSON body + aspect_ratio_field: Final = aspect or {} # mutable-ok: provider JSON body and base-class dict signature + n_field: Final = count or {} # mutable-ok: provider JSON body and base-class dict signature + resolution_field: Final = quality or {} # mutable-ok: provider JSON body and base-class dict signature + return { # mutable-ok: provider JSON body and base-class dict signature + **aspect_ratio_field, + **n_field, + **resolution_field, + } # mutable-ok: provider JSON body and base-class dict signature + + def use_multipart_form_data(self) -> bool: + return False + + def get_complete_url( + self, + model: str, + api_base: str | None, + litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> str: + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + api_key: Final = litellm_params.get("api_key") if isinstance(litellm_params, dict) else None + resolved_base: Final = ( + XAIOAuthAuthenticator().get_api_base() + if should_use_xai_oauth(litellm_params) and not XAIModelInfo.get_api_key(api_key) + else (api_base or get_secret_str("XAI_API_BASE") or get_secret_str("XAI_OAUTH_API_BASE") or XAI_API_BASE) + ) + base: Final = (resolved_base or XAI_API_BASE).rstrip("/") + if base.endswith("/v1"): + return f"{base}/images/edits" + return f"{base}/v1/images/edits" + + def validate_environment( + self, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + model: str, + api_key: str | None = None, + litellm_params: dict | None = None, # mutable-ok: provider JSON body and base-class dict signature + api_base: str | None = None, + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + params: Final = litellm_params or {} # mutable-ok: provider JSON body and base-class dict signature + dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key) + if should_use_xai_oauth(params) and not dynamic_api_key: + try: + headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().get_access_token()}" + except XAIOAuthError as exc: + raise AuthenticationError( + model=model, + llm_provider="xai", + message=str(exc), + ) from exc + else: + if not dynamic_api_key: + raise AuthenticationError( + model=model, + llm_provider="xai", + message=( + "Missing xAI credentials for image edit. Pass api_key / XAI_API_KEY, or set use_xai_oauth=True." + ), + ) + headers["Authorization"] = f"Bearer {dynamic_api_key}" + + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + return headers + + def transform_image_edit_request( + self, + model: str, + prompt: str | None, + image: FileTypes | None, + image_edit_optional_request_params: dict, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> tuple[dict, RequestFiles]: # mutable-ok: provider JSON body and base-class dict signature + if image is None: + raise ValueError("xAI image edit requires at least one reference image.") + + image_payloads: Final = tuple(self._to_image_url(item) for item in self._as_image_list(image)) + if not image_payloads: + raise ValueError("xAI image edit requires at least one reference image.") + + n: Final = image_edit_optional_request_params.get("n") + prompt_body = {"prompt": prompt} if prompt is not None else None # mutable-ok: provider JSON body + many = {"images": list(image_payloads)} # mutable-ok: provider JSON body + one = image_payloads[0] + image_body = {"image": one} if len(image_payloads) == 1 else many # mutable-ok: provider JSON body + prompt_field: Final = prompt_body or {} # mutable-ok: provider JSON body and base-class dict signature + image_field: Final = image_body # mutable-ok: provider JSON body and base-class dict signature + request: Final[dict[str, object]] = { # mutable-ok: provider JSON body and base-class dict signature + "model": XAIModelInfo.get_base_model(model) or model, + **prompt_field, + **image_field, + **{ # mutable-ok: provider JSON body and base-class dict signature + key: image_edit_optional_request_params[key] + for key in ("aspect_ratio", "resolution") + if image_edit_optional_request_params.get(key) is not None + }, + **({"n": int(n)} if n is not None else {}), # mutable-ok: provider JSON body and base-class dict signature + } + return request, [] # mutable-ok: provider JSON body and base-class dict signature + + def transform_image_edit_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + ) -> ImageResponse: + try: + response_data: Final = raw_response.json() + except Exception: + raise self.get_error_class( + error_message=raw_response.text, + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + images: Final = tuple( + ImageObject( + url=item.get("url"), + b64_json=item.get("b64_json") or item.get("b64"), + ) + for item in response_data.get("data") or () + if isinstance(item, dict) + ) + if not images: + raise self.get_error_class( + error_message=f"xAI image edit returned no image data: {response_data}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + return ImageResponse(data=list(images)) # mutable-ok: provider JSON body and base-class dict signature + + def _as_image_list( + self, image: FileTypes | list[FileTypes] + ) -> tuple[FileTypes, ...]: # mutable-ok: provider JSON body and base-class dict signature + if isinstance(image, list): + return tuple(item for item in image if item is not None) + return (image,) + + def _to_image_url( + self, image: FileTypes + ) -> dict[str, str]: # mutable-ok: provider JSON body and base-class dict signature + if isinstance(image, str): + return {"url": image} # mutable-ok: provider JSON body and base-class dict signature + if isinstance(image, dict): + if image.get("url"): + return {"url": str(image["url"])} # mutable-ok: provider JSON body and base-class dict signature + if image.get("file_id"): + file_id = str(image["file_id"]) + return {"file_id": file_id} # mutable-ok: provider JSON body and base-class dict signature + + mime: Final = ImageEditRequestUtils.get_image_content_type(image) + encoded: Final = base64.b64encode(self._read_all_bytes(image)).decode("utf-8") + return {"url": f"data:{mime};base64,{encoded}"} # mutable-ok: provider JSON body and base-class dict signature + + def _read_all_bytes(self, image: FileTypes) -> bytes: + if isinstance(image, bytes): + return image + if isinstance(image, bytearray): + return bytes(image) + if isinstance(image, (BytesIO, BufferedReader)): + return _read_seekable(image) + if hasattr(image, "read"): + raw: Final = image.read() + if isinstance(raw, str): + return raw.encode("utf-8") + return bytes(raw) + raise ValueError(f"Unsupported image input type for xAI image edit: {type(image)}") diff --git a/litellm/llms/xai/image_generation/__init__.py b/litellm/llms/xai/image_generation/__init__.py new file mode 100644 index 00000000000..b4d59285586 --- /dev/null +++ b/litellm/llms/xai/image_generation/__init__.py @@ -0,0 +1,11 @@ +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) + +from .transformation import XAIImageGenerationConfig + +__all__ = ("XAIImageGenerationConfig", "get_xai_image_generation_config") + + +def get_xai_image_generation_config(model: str) -> BaseImageGenerationConfig: + return XAIImageGenerationConfig() diff --git a/litellm/llms/xai/image_generation/transformation.py b/litellm/llms/xai/image_generation/transformation.py new file mode 100644 index 00000000000..e0290748218 --- /dev/null +++ b/litellm/llms/xai/image_generation/transformation.py @@ -0,0 +1,200 @@ +from typing import TYPE_CHECKING, Final + +import httpx + +from litellm.constants import XAI_API_BASE +from litellm.exceptions import AuthenticationError +from litellm.llms.base_llm.image_generation.transformation import ( + BaseImageGenerationConfig, +) +from litellm.llms.xai.common_utils import XAIModelInfo +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIImageGenerationOptionalParams, +) +from litellm.types.utils import ImageObject, ImageResponse + +if TYPE_CHECKING: + import tiktoken + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +_SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature + "1024x1024": "1:1", + "1792x1024": "16:9", + "1024x1792": "9:16", + "1536x1024": "3:2", + "1024x1536": "2:3", + "1280x720": "16:9", + "720x1280": "9:16", + "1920x1080": "16:9", + "1080x1920": "9:16", +} +_XAI_NATIVE_PARAMS: Final = frozenset({"aspect_ratio", "n"}) + + +class XAIImageGenerationConfig(BaseImageGenerationConfig): + def get_supported_openai_params( + self, model: str + ) -> list[OpenAIImageGenerationOptionalParams]: # mutable-ok: provider JSON body and base-class dict signature + return ["n", "response_format", "size", "user"] # mutable-ok: provider JSON body and base-class dict signature + + def map_openai_params( + self, + non_default_params: dict, # mutable-ok: provider JSON body and base-class dict signature + optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature + model: str, + drop_params: bool, + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + supported_params: Final = frozenset(self.get_supported_openai_params(model)) + allowed: Final = supported_params | _XAI_NATIVE_PARAMS + unknown: Final = tuple(key for key in non_default_params if key not in optional_params and key not in allowed) + if unknown and not drop_params: + raise ValueError( + f"Parameter {unknown[0]} is not supported for model {model}. " + f"Supported parameters are {sorted(allowed)}. " + "Set drop_params=True to drop unsupported parameters." + ) + + pairs = ((k, v) for k, v in non_default_params.items() if k in allowed) + native: Final = dict(pairs) # mutable-ok: provider JSON body and base-class dict signature + merged: Final = {**optional_params, **native} # mutable-ok: provider JSON body and base-class dict signature + size: Final = merged.get("size") + aspect_ratio: Final = merged.get("aspect_ratio") or ( + _SIZE_TO_ASPECT_RATIO.get(str(size), "1:1") if size else None + ) + n: Final = merged.get("n") + aspect = {"aspect_ratio": aspect_ratio} if aspect_ratio is not None else None # mutable-ok: provider JSON body + count = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body + aspect_ratio_field: Final = aspect or {} # mutable-ok: provider JSON body and base-class dict signature + n_field: Final = count or {} # mutable-ok: provider JSON body and base-class dict signature + return {**aspect_ratio_field, **n_field} # mutable-ok: provider JSON body and base-class dict signature + + def get_complete_url( + self, + api_base: str | None, + api_key: str | None, + model: str, + optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature + stream: bool | None = None, + ) -> str: + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + resolved_base: Final = ( + XAIOAuthAuthenticator().get_api_base() + if should_use_xai_oauth(litellm_params) and not XAIModelInfo.get_api_key(api_key) + else (api_base or get_secret_str("XAI_API_BASE") or get_secret_str("XAI_OAUTH_API_BASE") or XAI_API_BASE) + ) + base: Final = (resolved_base or XAI_API_BASE).rstrip("/") + if base.endswith("/v1"): + return f"{base}/images/generations" + return f"{base}/v1/images/generations" + + def validate_environment( + self, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + model: str, + messages: list[AllMessageValues], # mutable-ok: provider JSON body and base-class dict signature + optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key) + if should_use_xai_oauth(litellm_params) and not dynamic_api_key: + try: + headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().get_access_token()}" + except XAIOAuthError as exc: + raise AuthenticationError( + model=model, + llm_provider="xai", + message=str(exc), + ) from exc + else: + if not dynamic_api_key: + raise AuthenticationError( + model=model, + llm_provider="xai", + message=( + "Missing xAI credentials for image generation. " + "Pass api_key / XAI_API_KEY, or set use_xai_oauth=True." + ), + ) + headers["Authorization"] = f"Bearer {dynamic_api_key}" + + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + return headers + + def transform_image_generation_request( + self, + model: str, + prompt: str, + optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + n: Final = optional_params.get("n") + requested_ratio = optional_params.get("aspect_ratio") + aspect = {"aspect_ratio": requested_ratio} if requested_ratio is not None else None # fmt: skip # mutable-ok: provider JSON body + count = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body + return { # mutable-ok: provider JSON body and base-class dict signature + "model": XAIModelInfo.get_base_model(model) or model, + "prompt": prompt, + **(aspect or {}), # mutable-ok: provider JSON body and base-class dict signature + **(count or {}), # mutable-ok: provider JSON body and base-class dict signature + } + + def transform_image_generation_response( + self, + model: str, + raw_response: httpx.Response, + model_response: ImageResponse, + logging_obj: "LiteLLMLoggingObj", + request_data: dict, # mutable-ok: provider JSON body and base-class dict signature + optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature + encoding: "tiktoken.Encoding | None", + api_key: str | None = None, + json_mode: bool | None = None, + ) -> ImageResponse: + try: + response_data: Final = raw_response.json() + except Exception: + raise self.get_error_class( + error_message=raw_response.text, + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + logging_obj.post_call( + input=request_data.get("prompt", ""), + api_key=api_key, + additional_args={"complete_input_dict": request_data}, # mutable-ok: provider JSON body + original_response=response_data, + ) + + images: Final = tuple( + ImageObject( + url=item.get("url"), + b64_json=item.get("b64_json") or item.get("b64"), + ) + for item in response_data.get("data") or () + if isinstance(item, dict) + ) + if not images: + raise self.get_error_class( + error_message=f"xAI image generation returned no image data: {response_data}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + model_response.data = list(images) # mutable-ok: provider JSON body and base-class dict signature + return model_response diff --git a/litellm/llms/xai/videos/__init__.py b/litellm/llms/xai/videos/__init__.py new file mode 100644 index 00000000000..45650657059 --- /dev/null +++ b/litellm/llms/xai/videos/__init__.py @@ -0,0 +1,3 @@ +from .transformation import XAIVideoConfig + +__all__ = ["XAIVideoConfig"] # mutable-ok: provider JSON body and base-class dict signature diff --git a/litellm/llms/xai/videos/transformation.py b/litellm/llms/xai/videos/transformation.py new file mode 100644 index 00000000000..ef3786947f5 --- /dev/null +++ b/litellm/llms/xai/videos/transformation.py @@ -0,0 +1,420 @@ +import time +from typing import TYPE_CHECKING, Final + +import httpx +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, + HTTPHandler, + _get_httpx_client, + get_async_httpx_client, +) +from litellm.llms.xai.common_utils import XAIModelInfo +from litellm.secret_managers.main import get_secret_str +from litellm.types.router import GenericLiteLLMParams +from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject +from litellm.types.videos.utils import ( + encode_video_id_with_provider, + extract_original_video_id, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +_DROPPED = frozenset(("seconds", "size", "input_reference", "user", "extra_headers", "model")) +_SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature + "1024x1024": "1:1", + "1792x1024": "16:9", + "1024x1792": "9:16", + "1280x720": "16:9", + "720x1280": "9:16", + "1920x1080": "16:9", + "1080x1920": "9:16", +} + + +def _duration_from_seconds(seconds: object) -> int: + try: + return int(seconds) if seconds is not None else 6 + except (TypeError, ValueError): + return 6 + + +_STATUS_MAP: Final = { # mutable-ok: provider JSON body and base-class dict signature + "done": "completed", + "completed": "completed", + "succeeded": "completed", + "failed": "failed", + "expired": "failed", + "pending": "processing", + "processing": "processing", + "in_progress": "processing", +} + + +class XAIVideoConfig(BaseVideoConfig): + def get_supported_openai_params( + self, model: str + ) -> list: # mutable-ok: provider JSON body and base-class dict signature + return [ # mutable-ok: provider JSON body and base-class dict signature + "model", + "prompt", + "input_reference", + "seconds", + "size", + "user", + "extra_headers", + ] + + def map_openai_params( + self, + video_create_optional_params: VideoCreateOptionalRequestParams, + model: str, + drop_params: bool, + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + raw = video_create_optional_params + incoming: Final = dict(raw) # mutable-ok: provider JSON body and base-class dict signature + size: Final = incoming.get("size") + seconds = incoming.get("seconds") + use_duration = "seconds" in incoming and "duration" not in incoming + duration_value = _duration_from_seconds(seconds) + duration = {"duration": duration_value} if use_duration else None # mutable-ok: provider JSON body + mapped_ratio = incoming.get("aspect_ratio") or _SIZE_TO_ASPECT_RATIO.get(str(size), "16:9") + use_ratio = bool(size) and "aspect_ratio" not in incoming + ratio = {"aspect_ratio": mapped_ratio} if use_ratio else None # mutable-ok: provider JSON body + image_ref = incoming.get("image") or incoming.get("input_reference") + use_image = bool(incoming.get("input_reference")) and "image" not in incoming + image = {"image": image_ref} if use_image else None # mutable-ok: provider JSON body + return { # mutable-ok: provider JSON body and base-class dict signature + **{ # mutable-ok: provider JSON body and base-class dict signature + key: value for key, value in incoming.items() if key not in _DROPPED + }, + **(duration or {}), # mutable-ok: provider JSON body and base-class dict signature + **(ratio or {}), # mutable-ok: provider JSON body and base-class dict signature + **(image or {}), # mutable-ok: provider JSON body and base-class dict signature + } + + def _resolve_api_base( + self, + api_base: str | None, + api_key: str | None, + litellm_params: GenericLiteLLMParams + | dict + | None, # mutable-ok: provider JSON body and base-class dict signature + ) -> str: + from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth + + params: Final = ( + litellm_params.model_dump() + if isinstance(litellm_params, GenericLiteLLMParams) + else (litellm_params or {}) # mutable-ok: provider JSON body and base-class dict signature + ) + if should_use_xai_oauth(params) and not XAIModelInfo.get_api_key(api_key): + return XAIOAuthAuthenticator().get_api_base().rstrip("/") + + resolved: Final = ( + api_base + or (params.get("api_base") if isinstance(params, dict) else None) + or get_secret_str("XAI_API_BASE") + or get_secret_str("XAI_OAUTH_API_BASE") + or XAI_API_BASE + ) + return str(resolved).rstrip("/") + + def _v1_root(self, api_base: str) -> str: + base: Final = api_base.rstrip("/") + if base.endswith("/v1"): + return base + return f"{base}/v1" + + def validate_environment( + self, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + model: str, + api_key: str | None = None, + litellm_params: GenericLiteLLMParams | None = None, + ) -> dict: # mutable-ok: provider JSON body and base-class dict signature + from litellm.llms.xai.oauth import ( + XAIOAuthAuthenticator, + XAIOAuthError, + should_use_xai_oauth, + ) + + dumped = litellm_params.model_dump() if litellm_params is not None else None + params: Final = dumped or {} # mutable-ok: provider JSON body and base-class dict signature + resolved_api_key: Final = api_key or (litellm_params.api_key if litellm_params else None) + dynamic_api_key: Final = XAIModelInfo.get_api_key(resolved_api_key) + if should_use_xai_oauth(params) and not dynamic_api_key: + try: + headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().get_access_token()}" + except XAIOAuthError as exc: + raise AuthenticationError( + model=model or "xai-video", + llm_provider="xai", + message=str(exc), + ) from exc + else: + if not dynamic_api_key: + raise AuthenticationError( + model=model or "xai-video", + llm_provider="xai", + message=( + "Missing xAI credentials for video generation. " + "Pass api_key / XAI_API_KEY, or set use_xai_oauth=True." + ), + ) + headers["Authorization"] = f"Bearer {dynamic_api_key}" + + if "content-type" not in headers and "Content-Type" not in headers: + headers["Content-Type"] = "application/json" + return headers + + def get_complete_url( + self, + model: str, + api_base: str | None, + litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> str: + resolved: Final = self._resolve_api_base( + api_base=api_base, + api_key=litellm_params.get("api_key") if litellm_params else None, + litellm_params=litellm_params, + ) + if not model: + return self._v1_root(resolved) + return f"{self._v1_root(resolved)}/videos/generations" + + def transform_video_create_request( + self, + model: str, + prompt: str, + api_base: str, + video_create_optional_request_params: dict, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> tuple[dict, RequestFiles, str]: # mutable-ok: provider JSON body and base-class dict signature + copied: Final = { # mutable-ok: provider JSON body and base-class dict signature + key: video_create_optional_request_params[key] + for key in ( + "image", + "images", + "duration", + "resolution_name", + "aspect_ratio", + "size", + ) + if video_create_optional_request_params.get(key) is not None + } + prompt_body = {"prompt": prompt} if prompt else None # mutable-ok: provider JSON body + duration_body = {"duration": 6} if "duration" not in copied else None # mutable-ok: provider JSON body + return ( + { # mutable-ok: provider JSON body and base-class dict signature + "model": XAIModelInfo.get_base_model(model) or model, + **(prompt_body or {}), # mutable-ok: provider JSON body and base-class dict signature + **copied, + **(duration_body or {}), # mutable-ok: provider JSON body and base-class dict signature + }, + [], # mutable-ok: provider JSON body and base-class dict signature + api_base, + ) + + def transform_video_create_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + custom_llm_provider: str | None = None, + request_data: dict | None = None, # mutable-ok: provider JSON body and base-class dict signature + ) -> VideoObject: + response_data: Final = raw_response.json() + request_id: Final = response_data.get("request_id") or response_data.get("id") + if not request_id: + raise ValueError(f"xAI video generation response missing request_id: {response_data}") + + usage: Final = response_data.get("usage") or {} # mutable-ok: provider JSON body and base-class dict signature + video_obj: Final = VideoObject( + id=str(request_id), + object="video", + status="processing", + created_at=int(time.time()), + model=XAIModelInfo.get_base_model(model) or model, + progress=0, + ) + if custom_llm_provider: + video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model) + usage_body = usage if isinstance(usage, dict) else None + video_obj.usage = usage_body or {} # mutable-ok: provider JSON body and base-class dict signature + 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, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature + return self._video_resource_url( + api_base, video_id + ), {} # mutable-ok: provider JSON body and base-class dict signature + + def transform_video_status_retrieve_response( + self, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + custom_llm_provider: str | None = None, + ) -> VideoObject: + response_data: Final = raw_response.json() + status_raw: Final = str(response_data.get("status") or "processing").lower() + status: Final = _STATUS_MAP.get(status_raw, status_raw) + video_body = response_data.get("video") + video_meta: Final = video_body or {} # mutable-ok: provider JSON body and base-class dict signature + video_url: Final = video_meta.get("url") if isinstance(video_meta, dict) else None + seconds: Final = ( + str(video_meta.get("duration")) + if isinstance(video_meta, dict) and video_meta.get("duration") is not None + else None + ) + request_id: Final = ( + response_data.get("request_id") + or response_data.get("id") + or (video_url.split("/")[-1].replace(".mp4", "") if video_url else "unknown") + ) + video_obj: Final = VideoObject( + id=str(request_id), + object="video", + status=status, + created_at=response_data.get("created_at") or int(time.time()), + completed_at=int(time.time()) if status == "completed" else None, + model=response_data.get("model"), + progress=response_data.get("progress"), + seconds=seconds, + usage=response_data.get("usage") + if isinstance(response_data.get("usage"), dict) + else {}, # mutable-ok: provider JSON body and base-class dict signature + ) + video_obj._hidden_params["video_url"] = video_url + if custom_llm_provider and video_obj.id and video_obj.id != "unknown": + video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, response_data.get("model")) + return video_obj + + def transform_video_content_request( + self, + video_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + variant: str | None = None, + ) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature + return self._video_resource_url( + api_base, video_id + ), {} # mutable-ok: provider JSON body and base-class dict signature + + def _video_cdn_url(self, raw_response: httpx.Response) -> str | None: + content_type: Final = (raw_response.headers.get("content-type") or "").lower() + if "application/json" not in content_type and raw_response.content[:1] != b"{": + return None + payload: Final = raw_response.json() + if not isinstance(payload, dict): + return None + video_meta: Final = payload.get("video") or {} # mutable-ok: provider JSON body and base-class dict signature + url: Final = video_meta.get("url") if isinstance(video_meta, dict) else None + if isinstance(url, str) and url: + return url + raise ValueError(f"xAI video not ready for download (status={payload.get('status')}): {payload}") + + def transform_video_content_response( + self, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + ) -> bytes: + url: Final = self._video_cdn_url(raw_response) + if url is None: + return raw_response.content + httpx_client: Final[HTTPHandler] = _get_httpx_client() + video_response: Final = safe_get(httpx_client, url) + video_response.raise_for_status() + return video_response.content + + async def async_transform_video_content_response( + self, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + ) -> bytes: + url: Final = self._video_cdn_url(raw_response) + if url is None: + return raw_response.content + async_httpx_client: Final[AsyncHTTPHandler] = get_async_httpx_client( + llm_provider=litellm.LlmProviders.XAI, + ) + video_response: Final = await async_safe_get(async_httpx_client, url) + video_response.raise_for_status() + return video_response.content + + def transform_video_remix_request( + self, + video_id: str, + prompt: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + extra_body: dict[str, object] | None = None, # mutable-ok: provider JSON body and base-class dict signature + ) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature + raise NotImplementedError("Video remix is not supported by xAI Imagine API") + + def transform_video_remix_response( + self, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + custom_llm_provider: str | None = None, + ) -> VideoObject: + raise NotImplementedError("Video remix is not supported by xAI Imagine API") + + def transform_video_list_request( + self, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + after: str | None = None, + limit: int | None = None, + order: str | None = None, + extra_query: dict[str, object] | None = None, # mutable-ok: provider JSON body and base-class dict signature + ) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature + raise NotImplementedError("Video listing is not supported by xAI Imagine API") + + def transform_video_list_response( + self, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + custom_llm_provider: str | None = None, + ) -> dict[str, str]: # mutable-ok: provider JSON body and base-class dict signature + raise NotImplementedError("Video listing is not supported by xAI Imagine API") + + def transform_video_delete_request( + self, + video_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, # mutable-ok: provider JSON body and base-class dict signature + ) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature + raise NotImplementedError("Video delete is not supported by xAI Imagine API") + + def transform_video_delete_response( + self, + raw_response: httpx.Response, + logging_obj: "LiteLLMLoggingObj", + ) -> VideoObject: + raise NotImplementedError("Video delete is not supported by xAI Imagine API") diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 09fc442e5a7..082470d23c8 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -63368,7 +63368,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63384,7 +63385,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63401,7 +63403,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63418,7 +63421,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63435,7 +63439,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63451,7 +63456,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63467,7 +63473,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63485,6 +63492,9 @@ "output_cost_per_second_480p": 0.05, "output_cost_per_second_720p": 0.07, "source": "https://docs.x.ai/docs/models/grok-imagine-video", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63503,6 +63513,9 @@ "output_cost_per_second_480p": 0.08, "output_cost_per_second_720p": 0.14, "source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63521,6 +63534,9 @@ "output_cost_per_second_480p": 0.08, "output_cost_per_second_720p": 0.14, "source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63539,6 +63555,9 @@ "output_cost_per_second_480p": 0.08, "output_cost_per_second_720p": 0.14, "source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63582,7 +63601,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", diff --git a/litellm/provider_endpoints_support_backup.json b/litellm/provider_endpoints_support_backup.json index 1fcb7600a5e..0c913a36dde 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -2360,7 +2360,8 @@ "messages": true, "responses": true, "embeddings": false, - "image_generations": false, + "image_generations": true, + "image_edits": true, "audio_transcriptions": false, "audio_speech": false, "moderations": false, @@ -2368,7 +2369,8 @@ "rerank": false, "a2a": true, "interactions": true, - "realtime": true + "realtime": true, + "video_generations": true } }, "xinference": { diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c0123ae45a3..dec0d266874 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1691,6 +1691,7 @@ _MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS: Final = ( "/batches", "/skills", "/evals", + "/videos", ) _MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS: Final = ( "/files", diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index b9580ba3948..674165962bf 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -1,11 +1,13 @@ import asyncio import io -from collections.abc import Sequence +from collections.abc import Mapping, Sequence +from types import MappingProxyType from typing import Final, get_type_hints import orjson -from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, UploadFile, status +from fastapi import APIRouter, Depends, HTTPException, Request, Response, status from fastapi.responses import ORJSONResponse +from starlette.datastructures import FormData, UploadFile import litellm from litellm.litellm_core_utils.prompt_templates.common_utils import ( @@ -19,6 +21,7 @@ from litellm.proxy.common_request_processing import ( resolve_litellm_call_id, ) from litellm.proxy.common_utils.http_parsing_utils import ( + _is_form_content_type, coerce_numeric_form_fields, numeric_form_fields, ) @@ -39,6 +42,8 @@ IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(Image IMAGE_ARRAY_FIELD: Final = "image[]" MASK_ARRAY_FIELD: Final = "mask[]" BRACKETED_FILE_FIELDS: Final = frozenset({IMAGE_ARRAY_FIELD, MASK_ARRAY_FIELD}) +IMAGE_EDIT_FILE_FIELDS: Final = MappingProxyType({"image": IMAGE_ARRAY_FIELD, "mask": MASK_ARRAY_FIELD}) +IMAGE_REFERENCE_PREFIXES: Final = ("http://", "https://", "data:image/") async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO: @@ -63,6 +68,58 @@ async def batch_to_bytesio( return [await uploadfile_to_bytesio(u) for u in uploads] +async def _image_edit_part(value: object, field: str) -> io.BytesIO | str: + if isinstance(value, UploadFile): + return await uploadfile_to_bytesio(value) + if isinstance(value, str) and value.startswith(IMAGE_REFERENCE_PREFIXES): + return value + raise HTTPException( + status_code=422, + detail=f"'{field}' must be a multipart file upload, an http(s) URL, or a data:image URI.", + ) + + +def _image_edit_values(form: FormData | None, body: Mapping[str, object], name: str) -> tuple[object, ...]: + if form is not None: + return tuple(form.getlist(name)) + raw: Final = body.get(name) + if raw is None: + return () + return tuple(raw) if isinstance(raw, list) else (raw,) + + +async def _image_edit_field( + form: FormData | None, body: Mapping[str, object], field: str, alias: str +) -> list[io.BytesIO | str] | str | None: + values: Final = _image_edit_values(form, body, field) + alias_values: Final = _image_edit_values(form, body, alias) + if values and alias_values: + raise HTTPException(status_code=422, detail=f"Cannot specify both '{field}' and '{alias}'") + parts: Final = tuple([await _image_edit_part(value, field) for value in values or alias_values]) + if not parts: + return None + if len(parts) == 1 and isinstance(parts[0], str): + return parts[0] + return list(parts) # mutable-ok: provider image edit handlers take a list of image parts + + +async def image_edit_assets(request: Request, body: Mapping[str, object]) -> Mapping[str, object]: + """ + Collect ``image`` / ``mask`` (or their ``[]`` aliases) from a multipart form or JSON body. + + Each part is an uploaded file, an http(s) URL, or a ``data:image/`` URI; any other string is rejected so it can + never be treated as a filesystem path downstream. + """ + form: Final = await request.form() if _is_form_content_type(request.headers.get("content-type", "")) else None + return MappingProxyType( + { + field: parts + for field, alias in IMAGE_EDIT_FILE_FIELDS.items() + if (parts := await _image_edit_field(form, body, field, alias)) is not None + } + ) + + @router.post( "/v1/images/generations", dependencies=[Depends(user_api_key_auth)], @@ -247,10 +304,6 @@ async def image_edit_api( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), - image: list[UploadFile] | None = File(None), - image_array: list[UploadFile] | None = File(None, alias=IMAGE_ARRAY_FIELD), - mask: list[UploadFile] | None = File(None), - mask_array: list[UploadFile] | None = File(None, alias=MASK_ARRAY_FIELD), model: str | None = None, ): """ @@ -266,20 +319,6 @@ async def image_edit_api( -F 'prompt=Create a studio ghibli image of this' ``` """ - if image is not None and image_array is not None: - raise HTTPException(status_code=422, detail="Cannot specify both 'image' and 'image[]'") - if mask is not None and mask_array is not None: - raise HTTPException(status_code=422, detail="Cannot specify both 'mask' and 'mask[]'") - if image is None and image_array is not None: - image = image_array - if mask is None and mask_array is not None: - mask = mask_array - - # if image is None: - # raise HTTPException(status_code=422, detail="Field required: image") - # Note: Image is optional for some models (e.g., Bedrock Stability style-transfer) - # The validation will be done at the model level if image is truly required - from litellm.proxy.proxy_server import ( _read_request_body, general_settings, @@ -298,27 +337,14 @@ async def image_edit_api( ######################################################### # Read request body and convert UploadFiles to BytesIO ######################################################### + parsed_body: Final = coerce_numeric_form_fields( + parsed_body=await _read_request_body(request=request), + numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS, + ) data: Final = { - key: value - for key, value in coerce_numeric_form_fields( - parsed_body=await _read_request_body(request=request), - numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS, - ).items() - if key not in BRACKETED_FILE_FIELDS + **{key: value for key, value in parsed_body.items() if key not in BRACKETED_FILE_FIELDS}, + **await image_edit_assets(request, parsed_body), } - image_files: Final = await batch_to_bytesio(image) - mask_files: Final = await batch_to_bytesio(mask) - if image_files: - data["image"] = image_files - if mask_files: - data["mask"] = mask_files - - for _field in ("image", "mask"): - if _field in data and isinstance(data[_field], str): - raise HTTPException( - status_code=422, - detail=f"'{_field}' must be provided as a multipart file upload, not a string.", - ) # Ensure prompt exists in data (default to None for models that don't require it) if "prompt" not in data: diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index fe966c2e31a..4bddb528e0c 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -17,9 +17,14 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio from litellm.proxy.video_endpoints.utils import ( + assert_video_owner, 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, + stamp_video_owner, + video_owner_from_key, video_reference_to_id, ) from litellm.types.videos.utils import ( @@ -115,7 +120,19 @@ async def video_generation( version=version, ) else: - return generated + return _stamp_generated_video_owner( + generated, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key) + ) + + +def _stamp_generated_video_owner(result: object, owner: str | None) -> object: + video_id: Final = getattr(result, "id", None) + if isinstance(video_id, str): + result.id = stamp_video_owner(video_id, owner) # mutable-ok: response object id is rewritten before return + return result + if isinstance(result, dict) and isinstance(result.get("id"), str): + result["id"] = stamp_video_owner(result["id"], owner) # mutable-ok: JSON response body + return result @router.get( @@ -249,6 +266,8 @@ async def video_status( version, ) + assert_video_owner(video_id, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key)) + # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} @@ -256,23 +275,25 @@ 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" ) + + resolved_model: Final = resolve_video_request_model( + model_id_from_decoded=model_id_from_decoded, + query_model=request.query_params.get("model"), + llm_router=llm_router, + ) + 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 - # Resolve model_name from model_id if available - # This allows the router to automatically inject litellm_params from the model config - if model_id_from_decoded and llm_router: - resolved_model: Final = llm_router.resolve_model_name_from_model_id(model_id_from_decoded) - if resolved_model: - data["model"] = resolved_model - # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: @@ -350,6 +371,8 @@ async def video_content( version, ) + assert_video_owner(video_id, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key)) + # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} @@ -366,12 +389,13 @@ async def video_content( if custom_llm_provider: data["custom_llm_provider"] = custom_llm_provider - # Resolve model_name from model_id if available - # This allows the router to automatically inject litellm_params from the model config - if model_id_from_decoded and llm_router: - resolved_model: Final = llm_router.resolve_model_name_from_model_id(model_id_from_decoded) - if resolved_model: - data["model"] = resolved_model + resolved_content_model: Final = resolve_video_request_model( + model_id_from_decoded=model_id_from_decoded, + query_model=request.query_params.get("model"), + llm_router=llm_router, + ) + if resolved_content_model: + data["model"] = resolved_content_model # 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 a38226cc253..baa2aea81c3 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -1,16 +1,80 @@ -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Final, Protocol import orjson -from litellm.types.videos.utils import encode_character_id_with_provider +from litellm.proxy._types import ProxyException +from litellm.types.videos.utils import ( + decode_video_id_with_provider, + encode_character_id_with_provider, + encode_video_id_with_provider, +) -def extract_model_from_target_model_names(target_model_names: Any) -> str | None: - if isinstance(target_model_names, str): - target_model_names = [m.strip() for m in target_model_names.split(",") if m.strip()] - elif not isinstance(target_model_names, list): +class VideoModelIdResolver(Protocol): + def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: ... + + +def video_owner_from_key(token: str | None, api_key: str | None) -> str | None: + return token or api_key + + +def assert_video_owner(video_id: str, owner: str | None) -> None: + recorded: Final = decode_video_id_with_provider(video_id).get("owner") + if recorded and recorded != owner: + raise ProxyException( + message="Video does not belong to this API key", + type="permission_error", + param="video_id", + code=403, + ) + + +def stamp_video_owner(video_id: str, owner: str | None) -> str: + if not owner: + return video_id + decoded: Final = decode_video_id_with_provider(video_id) + provider: Final = decoded.get("custom_llm_provider") + raw_id: Final = decoded.get("video_id") + if not provider or not raw_id or decoded.get("owner"): + return video_id + return encode_video_id_with_provider(raw_id, provider, decoded.get("model_id"), owner) + + +def infer_video_provider_from_model(model: str | None) -> str | None: + if not isinstance(model, str) or not model: return None - return target_model_names[0] if target_model_names else 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, + query_model: str | None, + llm_router: VideoModelIdResolver | None, +) -> str | None: + if model_id_from_decoded: + if llm_router is not None: + resolved: Final = llm_router.resolve_model_name_from_model_id(model_id_from_decoded) + if isinstance(resolved, str) and resolved: + return resolved + return model_id_from_decoded + if isinstance(query_model, str) and query_model: + return query_model + return None + + +def extract_model_from_target_model_names(target_model_names: object) -> str | None: + if isinstance(target_model_names, str): + names: Final = tuple(m.strip() for m in target_model_names.split(",") if m.strip()) + return names[0] if names else None + if isinstance(target_model_names, Sequence) and not isinstance(target_model_names, (str, bytes)): + first: Final = target_model_names[0] if target_model_names else None + return first if isinstance(first, str) else None + return None def video_reference_to_id(video_ref: object) -> str: @@ -25,9 +89,9 @@ def video_reference_to_id(video_ref: object) -> str: return parsed_ref.get("id", "") if isinstance(parsed_ref, dict) else video_ref -def get_custom_provider_from_data(data: dict[str, Any]) -> str | None: +def get_custom_provider_from_data(data: Mapping[str, object]) -> str | None: custom_llm_provider: Final = data.get("custom_llm_provider") - if custom_llm_provider: + if isinstance(custom_llm_provider, str) and custom_llm_provider: return custom_llm_provider extra_body = data.get("extra_body") @@ -47,7 +111,7 @@ def get_custom_provider_from_data(data: dict[str, Any]) -> str | None: return None -def encode_character_id_in_response(response: Any, custom_llm_provider: str, model_id: str | None) -> Any: +def encode_character_id_in_response(response: object, custom_llm_provider: str, model_id: str | None) -> object: if isinstance(response, dict) and response.get("id"): response["id"] = encode_character_id_with_provider( character_id=response["id"], diff --git a/litellm/types/videos/main.py b/litellm/types/videos/main.py index f4369fd95af..f876df2d765 100644 --- a/litellm/types/videos/main.py +++ b/litellm/types/videos/main.py @@ -102,6 +102,7 @@ class DecodedVideoId(TypedDict, total=False): custom_llm_provider: str | None model_id: str | None video_id: str + owner: ReadOnly[str | None] class CharacterObject(BaseModel): diff --git a/litellm/types/videos/utils.py b/litellm/types/videos/utils.py index b23b2269543..e07f71b3240 100644 --- a/litellm/types/videos/utils.py +++ b/litellm/types/videos/utils.py @@ -35,7 +35,9 @@ def _add_base64_padding(value: str) -> str: return value -def encode_video_id_with_provider(video_id: str, provider: str, model_id: str | None = None) -> str: +def encode_video_id_with_provider( + video_id: str, provider: str, model_id: str | None = None, owner: str | None = None +) -> str: """Encode provider and model_id into video_id using base64.""" if not provider or not video_id: return video_id @@ -50,6 +52,8 @@ def encode_video_id_with_provider(video_id: str, provider: str, model_id: str | # ID is not encoded (even if it starts with video_), so encode it assembled_id = str(SpecialEnums.LITELLM_MANAGED_VIDEO_COMPLETE_STR.value).format(provider, model_id or "", video_id) + if owner: + assembled_id = f"{assembled_id};owner:{owner}" base64_encoded_id: Final[str] = base64.b64encode(assembled_id.encode("utf-8")).decode("utf-8") @@ -89,20 +93,24 @@ def decode_video_id_with_provider(encoded_video_id: str) -> DecodedVideoId: custom_llm_provider = None model_id = None decoded_video_id = encoded_video_id + owner = None if len(parts) >= 3: custom_llm_provider_part: Final = parts[0] model_id_part: Final = parts[1] video_id_part: Final = parts[2] + owner_part: Final = next((part for part in parts[3:] if part.startswith("owner:")), None) custom_llm_provider = custom_llm_provider_part.replace("litellm:custom_llm_provider:", "") model_id = model_id_part.replace("model_id:", "") decoded_video_id = video_id_part.replace("video_id:", "") + owner = owner_part.removeprefix("owner:") if owner_part else None return DecodedVideoId( custom_llm_provider=custom_llm_provider, model_id=model_id, video_id=decoded_video_id, + owner=owner, ) except Exception as e: verbose_logger.debug("Error decoding video_id '%s': %s", encoded_video_id, e) diff --git a/litellm/utils.py b/litellm/utils.py index d45b29c0f16..685952ff793 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9614,6 +9614,12 @@ class ProviderConfigManager: return get_modelscope_image_generation_config(model) elif LlmProviders.EDENAI == provider: return litellm.EdenAIImageGenerationConfig() + elif LlmProviders.XAI == provider: + from litellm.llms.xai.image_generation import ( + get_xai_image_generation_config, + ) + + return get_xai_image_generation_config(model) return None @staticmethod @@ -9645,6 +9651,10 @@ class ProviderConfigManager: from litellm.llms.fal_ai.videos.transformation import FalAIVideoConfig return FalAIVideoConfig() + elif LlmProviders.XAI == provider: + from litellm.llms.xai.videos.transformation import XAIVideoConfig + + return XAIVideoConfig() elif LlmProviders.HOSTED_VLLM == provider: from litellm.llms.hosted_vllm.videos import get_hosted_vllm_video_config @@ -9783,6 +9793,10 @@ class ProviderConfigManager: ) return get_openrouter_image_edit_config(model) + elif LlmProviders.XAI == provider: + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + return XAIImageEditConfig() return None @staticmethod diff --git a/litellm/videos/main.py b/litellm/videos/main.py index 445435a30fa..c01d1343c91 100644 --- a/litellm/videos/main.py +++ b/litellm/videos/main.py @@ -30,6 +30,51 @@ from litellm.videos.utils import VideoGenerationRequestUtils llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler() +def _litellm_provider_from_cost_entry(info: object) -> str | None: + if isinstance(info, dict): + catalog_provider: Final = info.get("litellm_provider") + if isinstance(catalog_provider, str) and catalog_provider: + return catalog_provider + return None + + +def _provider_from_prefixed_model(model: str) -> str | None: + if "/" not in model: + return None + try: + _, provider, _, _ = get_llm_provider(model=model) + except Exception: + return None + return provider + + +def _custom_llm_provider_from_model(model: str) -> str | None: + return ( + _provider_from_prefixed_model(model) + or _litellm_provider_from_cost_entry(litellm.model_cost.get(model)) + or _litellm_provider_from_cost_entry(litellm.model_cost.get(f"xai/{model}")) + or ("xai" if model.startswith("grok-imagine-video") else None) + ) + + +def _provider_for_video_id( + video_id: str, + custom_llm_provider: str | None, + model: object | None = None, +) -> str: + if custom_llm_provider is not None: + return custom_llm_provider + decoded: Final = decode_video_id_with_provider(video_id) + from_id: Final = decoded.get("custom_llm_provider") + if from_id: + return from_id + if isinstance(model, str) and model: + from_model: Final = _custom_llm_provider_from_model(model) + if from_model: + return from_model + return "openai" + + ##### Video Generation ####################### @client async def avideo_generation( @@ -317,10 +362,7 @@ def video_content( litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None) _is_async: Final = kwargs.pop("async_call", False) is True - # Try to decode provider from video_id if not explicitly provided - if custom_llm_provider is None: - decoded: Final = decode_video_id_with_provider(video_id) - custom_llm_provider = decoded.get("custom_llm_provider") or "openai" + custom_llm_provider = _provider_for_video_id(video_id, custom_llm_provider, kwargs.get("model")) # get llm provider logic litellm_params: Final = GenericLiteLLMParams(**kwargs) @@ -412,10 +454,7 @@ async def avideo_content( loop: Final = asyncio.get_event_loop() kwargs["async_call"] = True - # Try to decode provider from video_id if not explicitly provided - if custom_llm_provider is None: - decoded: Final = decode_video_id_with_provider(video_id) - custom_llm_provider = decoded.get("custom_llm_provider") or "openai" + custom_llm_provider = _provider_for_video_id(video_id, custom_llm_provider, kwargs.get("model")) func: Final = partial( video_content, @@ -1019,10 +1058,7 @@ def video_status( response: Final = VideoObject(**mock_response) return response - # Try to decode provider from video_id if not explicitly provided - if custom_llm_provider is None: - decoded: Final = decode_video_id_with_provider(video_id) - custom_llm_provider = decoded.get("custom_llm_provider") or "openai" + custom_llm_provider = _provider_for_video_id(video_id, custom_llm_provider, kwargs.get("model")) # get llm provider logic litellm_params: Final = GenericLiteLLMParams(**kwargs) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 09fc442e5a7..082470d23c8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -63368,7 +63368,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63384,7 +63385,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63401,7 +63403,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63418,7 +63421,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63435,7 +63439,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63451,7 +63456,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63467,7 +63473,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", @@ -63485,6 +63492,9 @@ "output_cost_per_second_480p": 0.05, "output_cost_per_second_720p": 0.07, "source": "https://docs.x.ai/docs/models/grok-imagine-video", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63503,6 +63513,9 @@ "output_cost_per_second_480p": 0.08, "output_cost_per_second_720p": 0.14, "source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63521,6 +63534,9 @@ "output_cost_per_second_480p": 0.08, "output_cost_per_second_720p": 0.14, "source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63539,6 +63555,9 @@ "output_cost_per_second_480p": 0.08, "output_cost_per_second_720p": 0.14, "source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5", + "supported_endpoints": [ + "/v1/videos" + ], "supported_modalities": [ "text", "image", @@ -63582,7 +63601,8 @@ "mode": "image_generation", "source": "https://docs.x.ai/docs/models", "supported_endpoints": [ - "/v1/images/generations" + "/v1/images/generations", + "/v1/images/edits" ], "supported_modalities": [ "text", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 790a050a878..ea4ecd05ef3 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2706,7 +2706,8 @@ "messages": true, "responses": true, "embeddings": false, - "image_generations": false, + "image_generations": true, + "image_edits": true, "audio_transcriptions": false, "audio_speech": false, "moderations": false, @@ -2714,7 +2715,8 @@ "rerank": false, "a2a": true, "interactions": true, - "realtime": true + "realtime": true, + "video_generations": true } }, "xinference": { diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 83ac56c4c85..53ad8e922b5 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -1022,6 +1022,25 @@ def test_get_model_from_request_includes_fine_tuning_target_model_query(): ) +def test_get_model_from_request_includes_video_query_model_for_plain_id(): + result = get_model_from_request( + request_data={"video_id": "plain-xai-id"}, + route="/v1/videos/{video_id}", + request_query_params={"model": "grok-imagine-video-1.5"}, + ) + assert result == "grok-imagine-video-1.5" + + +def test_get_model_from_request_includes_video_query_model_on_content_route(): + result = get_model_from_request( + request_data={"video_id": "plain-xai-id"}, + route="/v1/videos/{video_id}/content", + request_query_params={"model": "restricted-xai-model"}, + request_headers={"x-litellm-model": "also-restricted"}, + ) + assert result == ["restricted-xai-model", "also-restricted"] + + def test_get_model_from_request_extracts_video_id_model(): from litellm.types.videos.utils import encode_video_id_with_provider diff --git a/tests/test_litellm/proxy/image_endpoints/test_endpoints.py b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py index ad0901e9eee..bea993a364a 100644 --- a/tests/test_litellm/proxy/image_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/image_endpoints/test_endpoints.py @@ -222,6 +222,54 @@ def test_image_edit_multipart_n_that_is_not_a_number_is_left_alone(monkeypatch): assert captured["n"] == "two" +def test_image_edit_http_url_is_accepted(monkeypatch): + captured: Dict[str, Any] = {} + + response = _image_edit_client(monkeypatch, captured).post( + "/v1/images/edits", + data={ + "model": "grok-imagine-image", + "prompt": "make it night", + "image": "https://imgen.x.ai/source.jpeg", + }, + ) + + assert response.status_code == 200 + assert captured["image"] == "https://imgen.x.ai/source.jpeg" + + +def test_image_edit_data_uri_is_accepted(monkeypatch): + captured: Dict[str, Any] = {} + + response = _image_edit_client(monkeypatch, captured).post( + "/v1/images/edits", + data={ + "model": "grok-imagine-image", + "prompt": "make it red", + "image": "data:image/jpeg;base64,abc", + }, + ) + + assert response.status_code == 200 + assert captured["image"] == "data:image/jpeg;base64,abc" + + +def test_image_edit_plain_string_image_is_rejected(monkeypatch): + captured: Dict[str, Any] = {} + + response = _image_edit_client(monkeypatch, captured).post( + "/v1/images/edits", + data={ + "model": "grok-imagine-image", + "prompt": "make it red", + "image": "not-a-url-or-file", + }, + ) + + assert response.status_code == 422 + assert "multipart file" in response.json()["detail"] + + @pytest.mark.asyncio async def test_a_model_the_router_cannot_serve_answers_an_openai_typed_error(monkeypatch: pytest.MonkeyPatch): """A bare HTTPException carries no type or param, so the tail used to ship the diff --git a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py index c5996f95f54..78877b84f6c 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py @@ -309,6 +309,54 @@ async def test_status__header_provider_beats_decoded_id(harness): assert data["model"] == "azure-sora" +@pytest.mark.asyncio +async def test_status__resolve_fail_keeps_decoded_model_id(harness): + encoded = encode_video_id_with_provider( + "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + "xai", + "grok-imagine-video-1.5", + ) + + await call_status(harness, encoded) + + harness.resolve_model.assert_called_once_with("grok-imagine-video-1.5") + assert harness.processor_data() == { + "video_id": encoded, + "custom_llm_provider": "xai", + "model": "grok-imagine-video-1.5", + } + + +@pytest.mark.asyncio +async def test_status__query_model_on_plain_id(harness): + await call_status( + harness, + "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + query={"model": "grok-imagine-video-1.5"}, + ) + + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + "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 # # =========================================================================== # @@ -365,6 +413,42 @@ async def test_content__model_encoded_id(harness): } +@pytest.mark.asyncio +async def test_content__query_model_on_plain_id(harness): + harness.base_process.return_value = b"x" + + await call_content( + harness, + "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + query={"model": "grok-imagine-video-1.5"}, + ) + + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + "model": "grok-imagine-video-1.5", + } + + +@pytest.mark.asyncio +async def test_content__resolve_fail_keeps_decoded_model_id(harness): + harness.base_process.return_value = b"x" + encoded = encode_video_id_with_provider( + "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + "xai", + "grok-imagine-video-1.5", + ) + + await call_content(harness, encoded) + + harness.resolve_model.assert_called_once_with("grok-imagine-video-1.5") + assert harness.processor_data() == { + "video_id": encoded, + "custom_llm_provider": "xai", + "model": "grok-imagine-video-1.5", + } + + # =========================================================================== # # POST /v1/videos/edits - video_edit # # =========================================================================== # diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py index 9a2c208c075..3fb523944aa 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_utils.py +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -13,21 +13,86 @@ is encode_character_id_with_provider, which runs for real; encoding assertions are checked by the genuine decode round-trip. """ - import pytest - +from litellm.proxy._types import ProxyException from litellm.proxy.video_endpoints.utils import ( + assert_video_owner, 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, + stamp_video_owner, video_reference_to_id, ) from litellm.types.videos.utils import ( decode_character_id_with_provider, encode_character_id_with_provider, + encode_video_id_with_provider, ) +# =========================================================================== # +# resolve_video_request_model +# =========================================================================== # + + +class _Resolver: + def __init__(self, mapping: dict[str, str | None]): + self.mapping = mapping + + def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: + return self.mapping.get(model_id) if model_id else None + + +def test_resolve_video_request_model__router_hit(): + assert ( + resolve_video_request_model( + model_id_from_decoded="deployment-123", + query_model="ignored", + llm_router=_Resolver({"deployment-123": "azure-sora"}), + ) + == "azure-sora" + ) + + +def test_resolve_video_request_model__keeps_decoded_id_when_router_misses(): + assert ( + resolve_video_request_model( + model_id_from_decoded="grok-imagine-video-1.5", + query_model=None, + llm_router=_Resolver({}), + ) + == "grok-imagine-video-1.5" + ) + + +def test_resolve_video_request_model__query_model_on_plain_id(): + assert ( + resolve_video_request_model( + model_id_from_decoded=None, + query_model="grok-imagine-video-1.5", + llm_router=None, + ) + == "grok-imagine-video-1.5" + ) + + +@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 # =========================================================================== # @@ -103,12 +168,7 @@ def test_provider__falsy_top_level_falls_through_to_extra_body(falsy): def test_provider__from_extra_body_dict(): - assert ( - get_custom_provider_from_data( - {"extra_body": {"custom_llm_provider": "bedrock"}} - ) - == "bedrock" - ) + assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": "bedrock"}}) == "bedrock" def test_provider__from_extra_body_json_string(): @@ -126,10 +186,7 @@ def test_provider__json_string_parsing_to_non_dict_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 - ) + assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) is None @pytest.mark.parametrize( @@ -161,9 +218,7 @@ def test_encode__dict_with_id_mutates_in_place_and_preserves_other_keys(): 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" - ) + 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" @@ -201,6 +256,16 @@ def test_encode__object_non_str_or_empty_id_unchanged(bad_id): assert resp.id == bad_id # untouched +def test_stamp_and_assert_video_owner_round_trip(): + encoded = encode_video_id_with_provider("req_1", "xai", "grok-imagine-video") + stamped = stamp_video_owner(encoded, "key-hash") + assert stamped != encoded + assert_video_owner(stamped, "key-hash") + with pytest.raises(ProxyException): + assert_video_owner(stamped, "other-key") + assert_video_owner(encoded, "other-key") + + def test_encode__object_without_id_attr_returned_unchanged(): resp = _Resp() out = encode_character_id_in_response(resp, "azure", "model-1") diff --git a/tests/unit/llms/xai/test_xai_image_edit.py b/tests/unit/llms/xai/test_xai_image_edit.py new file mode 100644 index 00000000000..575cbf459b6 --- /dev/null +++ b/tests/unit/llms/xai/test_xai_image_edit.py @@ -0,0 +1,167 @@ +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + + +def test_provider_config_manager_returns_xai_image_edit_config(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + config = ProviderConfigManager.get_provider_image_edit_config( + model="grok-imagine-image", + provider=LlmProviders.XAI, + ) + assert isinstance(config, XAIImageEditConfig) + + +def test_get_complete_url_default(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + url = XAIImageEditConfig().get_complete_url( + model="grok-imagine-image", + api_base=None, + litellm_params={}, + ) + assert url.endswith("/v1/images/edits") + + +def test_uses_json_not_multipart(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + assert XAIImageEditConfig().use_multipart_form_data() is False + + +def test_map_size_to_aspect_ratio(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + mapped = XAIImageEditConfig().map_openai_params( + image_edit_optional_params={"size": "1024x1792", "n": 1}, + model="grok-imagine-image", + drop_params=True, + ) + assert mapped["aspect_ratio"] == "9:16" + assert mapped["n"] == 1 + assert "size" not in mapped + + +def test_validate_environment_requires_credentials(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + with pytest.raises(Exception, match="Missing xAI credentials"): + XAIImageEditConfig().validate_environment( + headers={}, + model="grok-imagine-image", + api_key=None, + litellm_params={}, + ) + + +def test_validate_environment_oauth_injects_bearer(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + with patch( + "litellm.llms.xai.oauth.XAIOAuthAuthenticator.get_access_token", + return_value="oauth-token", + ): + headers = XAIImageEditConfig().validate_environment( + headers={}, + model="grok-imagine-image", + api_key=None, + litellm_params={"use_xai_oauth": True}, + ) + assert headers["Authorization"] == "Bearer oauth-token" + + +def test_map_string_n_to_int(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + mapped = XAIImageEditConfig().map_openai_params( + image_edit_optional_params={"n": "1"}, + model="grok-imagine-image", + drop_params=True, + ) + assert mapped["n"] == 1 + assert isinstance(mapped["n"], int) + + +def test_transform_string_n_to_int(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + data, _ = XAIImageEditConfig().transform_image_edit_request( + model="xai/grok-imagine-image", + prompt="make the cube red", + image="https://imgen.x.ai/source.jpeg", + image_edit_optional_request_params={"n": "1"}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert data["n"] == 1 + assert isinstance(data["n"], int) + + +def test_transform_bytes_to_data_uri_and_response(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + config = XAIImageEditConfig() + data, files = config.transform_image_edit_request( + model="xai/grok-imagine-image", + prompt="make the cube red", + image=b"\xff\xd8\xfffakejpeg", + image_edit_optional_request_params={"aspect_ratio": "1:1"}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert files == [] + assert data["model"] == "grok-imagine-image" + assert data["prompt"] == "make the cube red" + assert data["aspect_ratio"] == "1:1" + assert data["image"]["url"].startswith("data:image/jpeg;base64,") + + raw = httpx.Response( + 200, + json={"data": [{"url": "https://imgen.x.ai/edited.jpeg", "mime_type": "image/jpeg"}]}, + ) + response = config.transform_image_edit_response( + model="grok-imagine-image", + raw_response=raw, + logging_obj=MagicMock(), + ) + assert response.data is not None + assert response.data[0].url == "https://imgen.x.ai/edited.jpeg" + + +def test_transform_http_url_passthrough(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + data, files = XAIImageEditConfig().transform_image_edit_request( + model="grok-imagine-image", + prompt="make it night", + image="https://imgen.x.ai/source.jpeg", + image_edit_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert files == [] + assert data["image"] == {"url": "https://imgen.x.ai/source.jpeg"} + + +def test_transform_multiple_images_uses_images_array(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + data, _ = XAIImageEditConfig().transform_image_edit_request( + model="grok-imagine-image", + prompt="combine styles", + image=["https://imgen.x.ai/a.jpeg", "https://imgen.x.ai/b.jpeg"], + image_edit_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert "image" not in data + assert data["images"] == [ + {"url": "https://imgen.x.ai/a.jpeg"}, + {"url": "https://imgen.x.ai/b.jpeg"}, + ] diff --git a/tests/unit/llms/xai/test_xai_image_generation.py b/tests/unit/llms/xai/test_xai_image_generation.py new file mode 100644 index 00000000000..a617388f34c --- /dev/null +++ b/tests/unit/llms/xai/test_xai_image_generation.py @@ -0,0 +1,129 @@ +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from litellm.llms.xai.image_generation.transformation import XAIImageGenerationConfig +from litellm.types.utils import ImageResponse, LlmProviders +from litellm.utils import ProviderConfigManager + + +def test_provider_config_manager_returns_xai_image_config(): + config = ProviderConfigManager.get_provider_image_generation_config( + model="grok-imagine-image", + provider=LlmProviders.XAI, + ) + assert isinstance(config, XAIImageGenerationConfig) + + +def test_map_size_to_aspect_ratio(): + mapped = XAIImageGenerationConfig().map_openai_params( + non_default_params={"size": "1024x1792"}, + optional_params={}, + model="grok-imagine-image", + drop_params=True, + ) + assert mapped["aspect_ratio"] == "9:16" + assert "size" not in mapped + + +def test_get_complete_url_default(): + url = XAIImageGenerationConfig().get_complete_url( + api_base=None, + api_key=None, + model="grok-imagine-image", + optional_params={}, + litellm_params={}, + ) + assert url.endswith("/v1/images/generations") + + +def test_validate_environment_requires_credentials(): + with pytest.raises(Exception, match="Missing xAI credentials"): + XAIImageGenerationConfig().validate_environment( + headers={}, + model="grok-imagine-image", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + +def test_validate_environment_oauth_injects_bearer(): + with patch( + "litellm.llms.xai.oauth.XAIOAuthAuthenticator.get_access_token", + return_value="oauth-token", + ): + headers = XAIImageGenerationConfig().validate_environment( + headers={}, + model="grok-imagine-image", + messages=[], + optional_params={}, + litellm_params={"use_xai_oauth": True}, + api_key=None, + ) + assert headers["Authorization"] == "Bearer oauth-token" + + +def test_map_string_n_to_int(): + mapped = XAIImageGenerationConfig().map_openai_params( + non_default_params={"n": "1"}, + optional_params={}, + model="grok-imagine-image", + drop_params=True, + ) + assert mapped["n"] == 1 + assert isinstance(mapped["n"], int) + + +def test_transform_string_n_to_int(): + request = XAIImageGenerationConfig().transform_image_generation_request( + model="xai/grok-imagine-image", + prompt="a red apple", + optional_params={"n": "1"}, + litellm_params={}, + headers={}, + ) + assert request["n"] == 1 + assert isinstance(request["n"], int) + + +def test_transform_request_and_response(): + config = XAIImageGenerationConfig() + request = config.transform_image_generation_request( + model="xai/grok-imagine-image", + prompt="a red apple", + optional_params={"aspect_ratio": "1:1"}, + litellm_params={}, + headers={}, + ) + assert request == { + "model": "grok-imagine-image", + "prompt": "a red apple", + "aspect_ratio": "1:1", + } + + raw = httpx.Response( + 200, + json={ + "data": [ + { + "url": "https://imgen.x.ai/example.jpeg", + "mime_type": "image/jpeg", + } + ] + }, + ) + response = config.transform_image_generation_response( + model="grok-imagine-image", + raw_response=raw, + model_response=ImageResponse(), + logging_obj=MagicMock(), + request_data=request, + optional_params={}, + litellm_params={}, + encoding=None, + ) + assert response.data is not None + assert response.data[0].url == "https://imgen.x.ai/example.jpeg" diff --git a/tests/unit/llms/xai/test_xai_video_generation.py b/tests/unit/llms/xai/test_xai_video_generation.py new file mode 100644 index 00000000000..261ce2e88c0 --- /dev/null +++ b/tests/unit/llms/xai/test_xai_video_generation.py @@ -0,0 +1,337 @@ +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 +from litellm.utils import ProviderConfigManager + + +def test_provider_config_manager_returns_xai_video_config(): + config = ProviderConfigManager.get_provider_video_config( + model="grok-imagine-video", + provider=LlmProviders.XAI, + ) + assert isinstance(config, XAIVideoConfig) + + +def test_map_seconds_and_size(): + mapped = XAIVideoConfig().map_openai_params( + video_create_optional_params={"seconds": "10", "size": "1280x720"}, + model="grok-imagine-video", + drop_params=True, + ) + assert mapped["duration"] == 10 + assert mapped["aspect_ratio"] == "16:9" + + +def test_map_seconds_nine_is_not_clamped_to_six(): + mapped = XAIVideoConfig().map_openai_params( + video_create_optional_params={"seconds": "9"}, + model="grok-imagine-video-1.5", + drop_params=True, + ) + assert mapped["duration"] == 9 + assert "seconds" not in mapped + + +def test_get_complete_url_create_and_status_root(): + config = XAIVideoConfig() + create_url = config.get_complete_url( + model="grok-imagine-video", + api_base="https://api.x.ai/v1", + litellm_params={}, + ) + assert create_url == "https://api.x.ai/v1/videos/generations" + + status_root = config.get_complete_url( + model="", + api_base="https://api.x.ai/v1", + litellm_params={}, + ) + assert status_root == "https://api.x.ai/v1" + + +def test_validate_environment_oauth_injects_bearer(): + with patch( + "litellm.llms.xai.oauth.XAIOAuthAuthenticator.get_access_token", + return_value="oauth-token", + ): + headers = XAIVideoConfig().validate_environment( + headers={}, + model="grok-imagine-video", + api_key=None, + litellm_params=GenericLiteLLMParams(use_xai_oauth=True), + ) + assert headers["Authorization"] == "Bearer oauth-token" + + +def test_transform_create_and_status_response(): + config = XAIVideoConfig() + data, files, api_base = config.transform_video_create_request( + model="xai/grok-imagine-video", + prompt="a cat walking", + api_base="https://api.x.ai/v1/videos/generations", + video_create_optional_request_params={"duration": 6}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert data["model"] == "grok-imagine-video" + assert data["prompt"] == "a cat walking" + assert data["duration"] == 6 + assert files == [] + + created = config.transform_video_create_response( + model="grok-imagine-video", + raw_response=httpx.Response(200, json={"request_id": "req-123"}), + logging_obj=MagicMock(), + custom_llm_provider="xai", + ) + assert created.status == "processing" + assert created.id + + status_url, params = config.transform_video_status_retrieve_request( + video_id=created.id, + api_base="https://api.x.ai/v1", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert status_url.endswith("/videos/req-123") + assert params == {} + + status = config.transform_video_status_retrieve_response( + raw_response=httpx.Response( + 200, + json={ + "status": "done", + "request_id": "req-123", + "video": {"url": "https://vidgen.x.ai/x.mp4", "duration": 6}, + "progress": 100, + "model": "grok-imagine-video", + }, + ), + logging_obj=MagicMock(), + custom_llm_provider="xai", + ) + assert status.status == "completed" + assert status.seconds == "6" + assert status._hidden_params.get("video_url") == "https://vidgen.x.ai/x.mp4" + + +def test_validate_environment_requires_credentials(): + with pytest.raises(Exception, match="Missing xAI credentials"): + XAIVideoConfig().validate_environment( + headers={}, + model="grok-imagine-video", + api_key=None, + litellm_params=GenericLiteLLMParams(), + ) + + +def test_content_request_is_get_status(): + url, params = XAIVideoConfig().transform_video_content_request( + video_id="req-123", + api_base="https://api.x.ai/v1", + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert url == "https://api.x.ai/v1/videos/req-123" + 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, + headers={"content-type": "application/json"}, + json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}}, + ) + cdn = MagicMock() + video_resp = MagicMock() + video_resp.content = b"mp4-bytes" + video_resp.raise_for_status.return_value = None + 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" + safe_get.assert_called_once_with(cdn, "https://vidgen.x.ai/x.mp4") + + +@pytest.mark.asyncio +async def test_async_content_response_does_not_use_sync_client(): + status = httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}}, + ) + async_client = MagicMock() + video_resp = MagicMock() + video_resp.content = b"async-mp4" + video_resp.raise_for_status.return_value = None + 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, 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_safe_get.assert_awaited_once_with(async_client, "https://vidgen.x.ai/x.mp4") + + +def test_content_response_raises_when_status_has_no_url(): + status = httpx.Response( + 200, + headers={"content-type": "application/json"}, + json={"status": "pending"}, + ) + with pytest.raises(ValueError, match="not ready"): + XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) + + +def test_content_response_returns_raw_bytes_when_not_json(): + raw = httpx.Response( + 200, + headers={"content-type": "video/mp4"}, + 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/unit/videos/test_main.py b/tests/unit/videos/test_main.py index 22e1e5c05eb..88e659490f6 100644 --- a/tests/unit/videos/test_main.py +++ b/tests/unit/videos/test_main.py @@ -178,6 +178,15 @@ def test_video_content__plain_id_defaults_to_openai(seams): assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai" +def test_video_content__plain_id_with_grok_model_uses_xai(seams): + videos_main.video_content( + video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd", + model="grok-imagine-video-1.5", + ) + + assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "xai" + + def test_video_remix__dispatch_and_provider_from_id(seams): result = videos_main.video_remix(video_id=AZURE_VIDEO_ID, prompt="new colors") @@ -377,6 +386,21 @@ async def test_avideo_content__pre_decodes_provider_before_delegating(): assert sync.call_args.kwargs["custom_llm_provider"] == "azure" +@pytest.mark.asyncio +async def test_avideo_content__plain_id_with_grok_model_uses_xai(): + 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="9b444cea-aaaa-bbbb-cccc-dddddddddddd", + model="grok-imagine-video-1.5", + ) + + assert result is sentinel + assert sync.call_args.kwargs["custom_llm_provider"] == "xai" + + # =========================================================================== # # Credential passthrough - DB/YAML model-config credentials the router injects # via kwargs must reach the provider call for EVERY video handler, carried in From 0722d6efdc0dfd52360d840b06751a2e1e321728 Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Sun, 27 Sep 2026 18:15:56 +0800 Subject: [PATCH 2/6] fix(proxy): stop embedding the caller's key in video ids The owner stamp added to every provider's video id was base64, not encrypted, and fell back to the raw bearer credential whenever the auth object had no hashed token (no master key, custom auth). It also did not bind anything: the provider request id can be decoded out of a stamped id and replayed as a plain id. Video ids go back to being the same bearer capability they are upstream Co-authored-by: Cursor --- litellm/proxy/video_endpoints/endpoints.py | 21 +----------- litellm/proxy/video_endpoints/utils.py | 33 +------------------ litellm/types/videos/main.py | 1 - litellm/types/videos/utils.py | 10 +----- .../proxy/video_endpoints/test_utils.py | 14 -------- 5 files changed, 3 insertions(+), 76 deletions(-) diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 4bddb528e0c..125bb30f86c 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -17,14 +17,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio from litellm.proxy.video_endpoints.utils import ( - assert_video_owner, 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, - stamp_video_owner, - video_owner_from_key, video_reference_to_id, ) from litellm.types.videos.utils import ( @@ -120,19 +117,7 @@ async def video_generation( version=version, ) else: - return _stamp_generated_video_owner( - generated, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key) - ) - - -def _stamp_generated_video_owner(result: object, owner: str | None) -> object: - video_id: Final = getattr(result, "id", None) - if isinstance(video_id, str): - result.id = stamp_video_owner(video_id, owner) # mutable-ok: response object id is rewritten before return - return result - if isinstance(result, dict) and isinstance(result.get("id"), str): - result["id"] = stamp_video_owner(result["id"], owner) # mutable-ok: JSON response body - return result + return generated @router.get( @@ -266,8 +251,6 @@ async def video_status( version, ) - assert_video_owner(video_id, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key)) - # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} @@ -371,8 +354,6 @@ async def video_content( version, ) - assert_video_owner(video_id, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key)) - # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index baa2aea81c3..9795e925863 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -3,44 +3,13 @@ from typing import Final, Protocol import orjson -from litellm.proxy._types import ProxyException -from litellm.types.videos.utils import ( - decode_video_id_with_provider, - encode_character_id_with_provider, - encode_video_id_with_provider, -) +from litellm.types.videos.utils import encode_character_id_with_provider class VideoModelIdResolver(Protocol): def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: ... -def video_owner_from_key(token: str | None, api_key: str | None) -> str | None: - return token or api_key - - -def assert_video_owner(video_id: str, owner: str | None) -> None: - recorded: Final = decode_video_id_with_provider(video_id).get("owner") - if recorded and recorded != owner: - raise ProxyException( - message="Video does not belong to this API key", - type="permission_error", - param="video_id", - code=403, - ) - - -def stamp_video_owner(video_id: str, owner: str | None) -> str: - if not owner: - return video_id - decoded: Final = decode_video_id_with_provider(video_id) - provider: Final = decoded.get("custom_llm_provider") - raw_id: Final = decoded.get("video_id") - if not provider or not raw_id or decoded.get("owner"): - return video_id - return encode_video_id_with_provider(raw_id, provider, decoded.get("model_id"), owner) - - def infer_video_provider_from_model(model: str | None) -> str | None: if not isinstance(model, str) or not model: return None diff --git a/litellm/types/videos/main.py b/litellm/types/videos/main.py index f876df2d765..f4369fd95af 100644 --- a/litellm/types/videos/main.py +++ b/litellm/types/videos/main.py @@ -102,7 +102,6 @@ class DecodedVideoId(TypedDict, total=False): custom_llm_provider: str | None model_id: str | None video_id: str - owner: ReadOnly[str | None] class CharacterObject(BaseModel): diff --git a/litellm/types/videos/utils.py b/litellm/types/videos/utils.py index e07f71b3240..b23b2269543 100644 --- a/litellm/types/videos/utils.py +++ b/litellm/types/videos/utils.py @@ -35,9 +35,7 @@ def _add_base64_padding(value: str) -> str: return value -def encode_video_id_with_provider( - video_id: str, provider: str, model_id: str | None = None, owner: str | None = None -) -> str: +def encode_video_id_with_provider(video_id: str, provider: str, model_id: str | None = None) -> str: """Encode provider and model_id into video_id using base64.""" if not provider or not video_id: return video_id @@ -52,8 +50,6 @@ def encode_video_id_with_provider( # ID is not encoded (even if it starts with video_), so encode it assembled_id = str(SpecialEnums.LITELLM_MANAGED_VIDEO_COMPLETE_STR.value).format(provider, model_id or "", video_id) - if owner: - assembled_id = f"{assembled_id};owner:{owner}" base64_encoded_id: Final[str] = base64.b64encode(assembled_id.encode("utf-8")).decode("utf-8") @@ -93,24 +89,20 @@ def decode_video_id_with_provider(encoded_video_id: str) -> DecodedVideoId: custom_llm_provider = None model_id = None decoded_video_id = encoded_video_id - owner = None if len(parts) >= 3: custom_llm_provider_part: Final = parts[0] model_id_part: Final = parts[1] video_id_part: Final = parts[2] - owner_part: Final = next((part for part in parts[3:] if part.startswith("owner:")), None) custom_llm_provider = custom_llm_provider_part.replace("litellm:custom_llm_provider:", "") model_id = model_id_part.replace("model_id:", "") decoded_video_id = video_id_part.replace("video_id:", "") - owner = owner_part.removeprefix("owner:") if owner_part else None return DecodedVideoId( custom_llm_provider=custom_llm_provider, model_id=model_id, video_id=decoded_video_id, - owner=owner, ) except Exception as e: verbose_logger.debug("Error decoding video_id '%s': %s", encoded_video_id, e) diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py index 3fb523944aa..4f7e4d5523e 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_utils.py +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -15,21 +15,17 @@ are checked by the genuine decode round-trip. import pytest -from litellm.proxy._types import ProxyException from litellm.proxy.video_endpoints.utils import ( - assert_video_owner, 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, - stamp_video_owner, video_reference_to_id, ) from litellm.types.videos.utils import ( decode_character_id_with_provider, encode_character_id_with_provider, - encode_video_id_with_provider, ) # =========================================================================== # @@ -256,16 +252,6 @@ def test_encode__object_non_str_or_empty_id_unchanged(bad_id): assert resp.id == bad_id # untouched -def test_stamp_and_assert_video_owner_round_trip(): - encoded = encode_video_id_with_provider("req_1", "xai", "grok-imagine-video") - stamped = stamp_video_owner(encoded, "key-hash") - assert stamped != encoded - assert_video_owner(stamped, "key-hash") - with pytest.raises(ProxyException): - assert_video_owner(stamped, "other-key") - assert_video_owner(encoded, "other-key") - - def test_encode__object_without_id_attr_returned_unchanged(): resp = _Resp() out = encode_character_id_in_response(resp, "azure", "model-1") From 09022930a3e22ac893f345aea5ca9e41fc3a5613 Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Sun, 27 Sep 2026 18:16:04 +0800 Subject: [PATCH 3/6] fix(xai): resolve plain video id providers from the requested model Video status forced an "openai" default whenever no provider was explicit, so a plain xAI request id with ?model=xai/grok-imagine-video was sent to OpenAI on the direct (non-router) path. The proxy now only defaults to openai when no model is known, and the SDK resolves the provider from the model with get_llm_provider instead of the grok-imagine name heuristics that lived in both the proxy and the SDK Co-authored-by: Cursor --- litellm/proxy/video_endpoints/endpoints.py | 3 +- litellm/proxy/video_endpoints/utils.py | 9 ---- litellm/videos/main.py | 42 ++++--------------- .../proxy/video_endpoints/test_endpoints.py | 22 ++-------- .../proxy/video_endpoints/test_utils.py | 34 +++++++-------- tests/unit/videos/test_main.py | 41 ++++++++++++++++-- 6 files changed, 65 insertions(+), 86 deletions(-) diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 125bb30f86c..540860c1a5b 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -20,7 +20,6 @@ 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, ) @@ -273,7 +272,7 @@ 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" + custom_llm_provider: Final = explicit_provider or (None if resolved_model else "openai") if custom_llm_provider: data["custom_llm_provider"] = custom_llm_provider diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index 9795e925863..6e305bdff2a 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -10,15 +10,6 @@ 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/litellm/videos/main.py b/litellm/videos/main.py index c01d1343c91..32a56ad645a 100644 --- a/litellm/videos/main.py +++ b/litellm/videos/main.py @@ -30,51 +30,25 @@ from litellm.videos.utils import VideoGenerationRequestUtils llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler() -def _litellm_provider_from_cost_entry(info: object) -> str | None: - if isinstance(info, dict): - catalog_provider: Final = info.get("litellm_provider") - if isinstance(catalog_provider, str) and catalog_provider: - return catalog_provider - return None - - -def _provider_from_prefixed_model(model: str) -> str | None: - if "/" not in model: +def _provider_from_model(model: object) -> str | None: + if not isinstance(model, str) or not model: return None try: _, provider, _, _ = get_llm_provider(model=model) - except Exception: + except litellm.BadRequestError: return None return provider -def _custom_llm_provider_from_model(model: str) -> str | None: +def _provider_for_video_id(video_id: str, custom_llm_provider: str | None, model: object) -> str: return ( - _provider_from_prefixed_model(model) - or _litellm_provider_from_cost_entry(litellm.model_cost.get(model)) - or _litellm_provider_from_cost_entry(litellm.model_cost.get(f"xai/{model}")) - or ("xai" if model.startswith("grok-imagine-video") else None) + custom_llm_provider + or decode_video_id_with_provider(video_id).get("custom_llm_provider") + or _provider_from_model(model) + or "openai" ) -def _provider_for_video_id( - video_id: str, - custom_llm_provider: str | None, - model: object | None = None, -) -> str: - if custom_llm_provider is not None: - return custom_llm_provider - decoded: Final = decode_video_id_with_provider(video_id) - from_id: Final = decoded.get("custom_llm_provider") - if from_id: - return from_id - if isinstance(model, str) and model: - from_model: Final = _custom_llm_provider_from_model(model) - if from_model: - return from_model - return "openai" - - ##### Video Generation ####################### @client async def avideo_generation( diff --git a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py index 78877b84f6c..861b68e4953 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py @@ -328,35 +328,21 @@ async def test_status__resolve_fail_keeps_decoded_model_id(harness): @pytest.mark.asyncio -async def test_status__query_model_on_plain_id(harness): +async def test_status__query_model_on_plain_id_leaves_provider_to_the_model(harness): + # an "openai" default here would override the provider the router / SDK derive from the model await call_status( harness, "9b444cea-aaaa-bbbb-cccc-dddddddddddd", - query={"model": "grok-imagine-video-1.5"}, + query={"model": "xai/grok-imagine-video-1.5"}, ) harness.resolve_model.assert_not_called() assert harness.processor_data() == { "video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd", - "custom_llm_provider": "xai", - "model": "grok-imagine-video-1.5", + "model": "xai/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 4f7e4d5523e..5944fd7a394 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_utils.py +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -13,13 +13,14 @@ is encode_character_id_with_provider, which runs for real; encoding assertions are checked by the genuine decode round-trip. """ + import pytest + 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, ) @@ -74,21 +75,6 @@ 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 # =========================================================================== # @@ -164,7 +150,12 @@ def test_provider__falsy_top_level_falls_through_to_extra_body(falsy): def test_provider__from_extra_body_dict(): - assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": "bedrock"}}) == "bedrock" + assert ( + get_custom_provider_from_data( + {"extra_body": {"custom_llm_provider": "bedrock"}} + ) + == "bedrock" + ) def test_provider__from_extra_body_json_string(): @@ -182,7 +173,10 @@ def test_provider__json_string_parsing_to_non_dict_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 + assert ( + get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) + is None + ) @pytest.mark.parametrize( @@ -214,7 +208,9 @@ def test_encode__dict_with_id_mutates_in_place_and_preserves_other_keys(): 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") + 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" diff --git a/tests/unit/videos/test_main.py b/tests/unit/videos/test_main.py index 88e659490f6..6325caa805d 100644 --- a/tests/unit/videos/test_main.py +++ b/tests/unit/videos/test_main.py @@ -39,6 +39,7 @@ import pytest import litellm +from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider 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 @@ -178,15 +179,47 @@ def test_video_content__plain_id_defaults_to_openai(seams): assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai" -def test_video_content__plain_id_with_grok_model_uses_xai(seams): +@pytest.fixture +def real_provider_resolution(seams): + with patch.object(videos_main, "get_llm_provider", get_llm_provider): + yield seams + + +def test_video_content__plain_id_with_xai_model_uses_xai(real_provider_resolution): + seams = real_provider_resolution videos_main.video_content( video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd", - model="grok-imagine-video-1.5", + model="xai/grok-imagine-video-1.5", ) assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "xai" +def test_video_status__plain_id_with_xai_model_uses_xai(real_provider_resolution): + seams = real_provider_resolution + videos_main.video_status( + video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd", + model="xai/grok-imagine-video-1.5", + ) + + assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "xai" + assert seams.get_config.call_args.kwargs["provider"] == litellm.LlmProviders.XAI + + +def test_video_status__plain_id_with_unresolvable_model_defaults_to_openai(real_provider_resolution): + seams = real_provider_resolution + videos_main.video_status(video_id="video_plain", model="not-a-known-model") + + assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "openai" + + +def test_video_status__decoded_provider_beats_model(real_provider_resolution): + seams = real_provider_resolution + videos_main.video_status(video_id=AZURE_VIDEO_ID, model="xai/grok-imagine-video-1.5") + + assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "azure" + + def test_video_remix__dispatch_and_provider_from_id(seams): result = videos_main.video_remix(video_id=AZURE_VIDEO_ID, prompt="new colors") @@ -387,14 +420,14 @@ async def test_avideo_content__pre_decodes_provider_before_delegating(): @pytest.mark.asyncio -async def test_avideo_content__plain_id_with_grok_model_uses_xai(): +async def test_avideo_content__plain_id_with_xai_model_uses_xai(): 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="9b444cea-aaaa-bbbb-cccc-dddddddddddd", - model="grok-imagine-video-1.5", + model="xai/grok-imagine-video-1.5", ) assert result is sentinel From b210c6bf4cc6e5695257503085a37bb98dcb64e3 Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Sun, 27 Sep 2026 18:16:11 +0800 Subject: [PATCH 4/6] fix(xai): clear type-discipline violations in imagine transforms Four `# mutable-ok` comments had drifted off the lines they were meant to cover and suppressed nothing (LIT013 is frozen at 0), and the new locals were missing `Final` (LIT010). Build the optional image params with one filtered comprehension instead of chained `x or {}` dicts Co-authored-by: Cursor --- litellm/llms/xai/image_edit/transformation.py | 38 +++++++++---------- .../xai/image_generation/transformation.py | 15 +++----- litellm/llms/xai/videos/transformation.py | 38 +++++++++---------- 3 files changed, 41 insertions(+), 50 deletions(-) diff --git a/litellm/llms/xai/image_edit/transformation.py b/litellm/llms/xai/image_edit/transformation.py index 496a0c7e65d..b20382a7acf 100644 --- a/litellm/llms/xai/image_edit/transformation.py +++ b/litellm/llms/xai/image_edit/transformation.py @@ -54,7 +54,7 @@ class XAIImageEditConfig(BaseImageEditConfig): ) -> dict: # mutable-ok: provider JSON body and base-class dict signature supported: Final = frozenset(self.get_supported_openai_params(model)) allowed: Final = supported | _XAI_NATIVE_PARAMS - raw = image_edit_optional_params + raw: Final = image_edit_optional_params incoming: Final = dict(raw) # mutable-ok: provider JSON body and base-class dict signature unknown: Final = tuple(key for key in incoming if key not in allowed) if unknown and not drop_params: @@ -64,7 +64,7 @@ class XAIImageEditConfig(BaseImageEditConfig): "Set drop_params=True to drop unsupported parameters." ) - pairs = ((key, value) for key, value in incoming.items() if key in allowed) + pairs: Final = ((key, value) for key, value in incoming.items() if key in allowed) mapped: Final = dict(pairs) # mutable-ok: provider JSON body and base-class dict signature size: Final = mapped.get("size") aspect_ratio: Final = mapped.get("aspect_ratio") or ( @@ -72,17 +72,12 @@ class XAIImageEditConfig(BaseImageEditConfig): ) n: Final = mapped.get("n") resolution: Final = mapped.get("resolution") - aspect = {"aspect_ratio": aspect_ratio} if aspect_ratio is not None else None # mutable-ok: provider JSON body - count = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body - quality = {"resolution": resolution} if resolution is not None else None # mutable-ok: provider JSON body - aspect_ratio_field: Final = aspect or {} # mutable-ok: provider JSON body and base-class dict signature - n_field: Final = count or {} # mutable-ok: provider JSON body and base-class dict signature - resolution_field: Final = quality or {} # mutable-ok: provider JSON body and base-class dict signature - return { # mutable-ok: provider JSON body and base-class dict signature - **aspect_ratio_field, - **n_field, - **resolution_field, - } # mutable-ok: provider JSON body and base-class dict signature + fields: Final = ( + ("aspect_ratio", aspect_ratio), + ("n", int(n) if n is not None else None), + ("resolution", resolution), + ) + return {key: value for key, value in fields if value is not None} # mutable-ok: base class returns a dict def use_multipart_form_data(self) -> bool: return False @@ -163,12 +158,12 @@ class XAIImageEditConfig(BaseImageEditConfig): raise ValueError("xAI image edit requires at least one reference image.") n: Final = image_edit_optional_request_params.get("n") - prompt_body = {"prompt": prompt} if prompt is not None else None # mutable-ok: provider JSON body - many = {"images": list(image_payloads)} # mutable-ok: provider JSON body - one = image_payloads[0] - image_body = {"image": one} if len(image_payloads) == 1 else many # mutable-ok: provider JSON body + prompt_body: Final = {"prompt": prompt} if prompt is not None else None # mutable-ok: provider JSON body + many: Final = {"images": list(image_payloads)} # mutable-ok: provider JSON body + one: Final = image_payloads[0] + image_body: Final = {"image": one} if len(image_payloads) == 1 else many # mutable-ok: provider JSON body prompt_field: Final = prompt_body or {} # mutable-ok: provider JSON body and base-class dict signature - image_field: Final = image_body # mutable-ok: provider JSON body and base-class dict signature + image_field: Final = image_body request: Final[dict[str, object]] = { # mutable-ok: provider JSON body and base-class dict signature "model": XAIModelInfo.get_base_model(model) or model, **prompt_field, @@ -214,8 +209,9 @@ class XAIImageEditConfig(BaseImageEditConfig): return ImageResponse(data=list(images)) # mutable-ok: provider JSON body and base-class dict signature def _as_image_list( - self, image: FileTypes | list[FileTypes] - ) -> tuple[FileTypes, ...]: # mutable-ok: provider JSON body and base-class dict signature + self, + image: FileTypes | list[FileTypes], # mutable-ok: the proxy passes multi-image edits as a list of uploads + ) -> tuple[FileTypes, ...]: if isinstance(image, list): return tuple(item for item in image if item is not None) return (image,) @@ -229,7 +225,7 @@ class XAIImageEditConfig(BaseImageEditConfig): if image.get("url"): return {"url": str(image["url"])} # mutable-ok: provider JSON body and base-class dict signature if image.get("file_id"): - file_id = str(image["file_id"]) + file_id: Final = str(image["file_id"]) return {"file_id": file_id} # mutable-ok: provider JSON body and base-class dict signature mime: Final = ImageEditRequestUtils.get_image_content_type(image) diff --git a/litellm/llms/xai/image_generation/transformation.py b/litellm/llms/xai/image_generation/transformation.py index e0290748218..5db907fe809 100644 --- a/litellm/llms/xai/image_generation/transformation.py +++ b/litellm/llms/xai/image_generation/transformation.py @@ -57,7 +57,7 @@ class XAIImageGenerationConfig(BaseImageGenerationConfig): "Set drop_params=True to drop unsupported parameters." ) - pairs = ((k, v) for k, v in non_default_params.items() if k in allowed) + pairs: Final = ((k, v) for k, v in non_default_params.items() if k in allowed) native: Final = dict(pairs) # mutable-ok: provider JSON body and base-class dict signature merged: Final = {**optional_params, **native} # mutable-ok: provider JSON body and base-class dict signature size: Final = merged.get("size") @@ -65,11 +65,8 @@ class XAIImageGenerationConfig(BaseImageGenerationConfig): _SIZE_TO_ASPECT_RATIO.get(str(size), "1:1") if size else None ) n: Final = merged.get("n") - aspect = {"aspect_ratio": aspect_ratio} if aspect_ratio is not None else None # mutable-ok: provider JSON body - count = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body - aspect_ratio_field: Final = aspect or {} # mutable-ok: provider JSON body and base-class dict signature - n_field: Final = count or {} # mutable-ok: provider JSON body and base-class dict signature - return {**aspect_ratio_field, **n_field} # mutable-ok: provider JSON body and base-class dict signature + fields: Final = (("aspect_ratio", aspect_ratio), ("n", int(n) if n is not None else None)) + return {key: value for key, value in fields if value is not None} # mutable-ok: base class returns a dict def get_complete_url( self, @@ -143,9 +140,9 @@ class XAIImageGenerationConfig(BaseImageGenerationConfig): headers: dict, # mutable-ok: provider JSON body and base-class dict signature ) -> dict: # mutable-ok: provider JSON body and base-class dict signature n: Final = optional_params.get("n") - requested_ratio = optional_params.get("aspect_ratio") - aspect = {"aspect_ratio": requested_ratio} if requested_ratio is not None else None # fmt: skip # mutable-ok: provider JSON body - count = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body + requested_ratio: Final = optional_params.get("aspect_ratio") + aspect: Final = {"aspect_ratio": requested_ratio} if requested_ratio is not None else None # fmt: skip # mutable-ok: provider JSON body + count: Final = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body return { # mutable-ok: provider JSON body and base-class dict signature "model": XAIModelInfo.get_base_model(model) or model, "prompt": prompt, diff --git a/litellm/llms/xai/videos/transformation.py b/litellm/llms/xai/videos/transformation.py index ef3786947f5..5b4cfb94fb4 100644 --- a/litellm/llms/xai/videos/transformation.py +++ b/litellm/llms/xai/videos/transformation.py @@ -27,7 +27,7 @@ from litellm.types.videos.utils import ( if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -_DROPPED = frozenset(("seconds", "size", "input_reference", "user", "extra_headers", "model")) +_DROPPED: Final = frozenset(("seconds", "size", "input_reference", "user", "extra_headers", "model")) _SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature "1024x1024": "1:1", "1792x1024": "16:9", @@ -78,19 +78,19 @@ class XAIVideoConfig(BaseVideoConfig): model: str, drop_params: bool, ) -> dict: # mutable-ok: provider JSON body and base-class dict signature - raw = video_create_optional_params + raw: Final = video_create_optional_params incoming: Final = dict(raw) # mutable-ok: provider JSON body and base-class dict signature size: Final = incoming.get("size") - seconds = incoming.get("seconds") - use_duration = "seconds" in incoming and "duration" not in incoming - duration_value = _duration_from_seconds(seconds) - duration = {"duration": duration_value} if use_duration else None # mutable-ok: provider JSON body - mapped_ratio = incoming.get("aspect_ratio") or _SIZE_TO_ASPECT_RATIO.get(str(size), "16:9") - use_ratio = bool(size) and "aspect_ratio" not in incoming - ratio = {"aspect_ratio": mapped_ratio} if use_ratio else None # mutable-ok: provider JSON body - image_ref = incoming.get("image") or incoming.get("input_reference") - use_image = bool(incoming.get("input_reference")) and "image" not in incoming - image = {"image": image_ref} if use_image else None # mutable-ok: provider JSON body + seconds: Final = incoming.get("seconds") + use_duration: Final = "seconds" in incoming and "duration" not in incoming + duration_value: Final = _duration_from_seconds(seconds) + duration: Final = {"duration": duration_value} if use_duration else None # mutable-ok: provider JSON body + mapped_ratio: Final = incoming.get("aspect_ratio") or _SIZE_TO_ASPECT_RATIO.get(str(size), "16:9") + use_ratio: Final = bool(size) and "aspect_ratio" not in incoming + ratio: Final = {"aspect_ratio": mapped_ratio} if use_ratio else None # mutable-ok: provider JSON body + image_ref: Final = incoming.get("image") or incoming.get("input_reference") + use_image: Final = bool(incoming.get("input_reference")) and "image" not in incoming + image: Final = {"image": image_ref} if use_image else None # mutable-ok: provider JSON body return { # mutable-ok: provider JSON body and base-class dict signature **{ # mutable-ok: provider JSON body and base-class dict signature key: value for key, value in incoming.items() if key not in _DROPPED @@ -104,9 +104,7 @@ class XAIVideoConfig(BaseVideoConfig): self, api_base: str | None, api_key: str | None, - litellm_params: GenericLiteLLMParams - | dict - | None, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: GenericLiteLLMParams | dict | None, # mutable-ok: get_complete_url receives the base-class dict ) -> str: from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth @@ -146,7 +144,7 @@ class XAIVideoConfig(BaseVideoConfig): should_use_xai_oauth, ) - dumped = litellm_params.model_dump() if litellm_params is not None else None + dumped: Final = litellm_params.model_dump() if litellm_params is not None else None params: Final = dumped or {} # mutable-ok: provider JSON body and base-class dict signature resolved_api_key: Final = api_key or (litellm_params.api_key if litellm_params else None) dynamic_api_key: Final = XAIModelInfo.get_api_key(resolved_api_key) @@ -211,8 +209,8 @@ class XAIVideoConfig(BaseVideoConfig): ) if video_create_optional_request_params.get(key) is not None } - prompt_body = {"prompt": prompt} if prompt else None # mutable-ok: provider JSON body - duration_body = {"duration": 6} if "duration" not in copied else None # mutable-ok: provider JSON body + prompt_body: Final = {"prompt": prompt} if prompt else None # mutable-ok: provider JSON body + duration_body: Final = {"duration": 6} if "duration" not in copied else None # mutable-ok: provider JSON body return ( { # mutable-ok: provider JSON body and base-class dict signature "model": XAIModelInfo.get_base_model(model) or model, @@ -248,7 +246,7 @@ class XAIVideoConfig(BaseVideoConfig): ) if custom_llm_provider: video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model) - usage_body = usage if isinstance(usage, dict) else None + usage_body: Final = usage if isinstance(usage, dict) else None video_obj.usage = usage_body or {} # mutable-ok: provider JSON body and base-class dict signature video_obj._hidden_params["video_url"] = None return video_obj @@ -280,7 +278,7 @@ class XAIVideoConfig(BaseVideoConfig): response_data: Final = raw_response.json() status_raw: Final = str(response_data.get("status") or "processing").lower() status: Final = _STATUS_MAP.get(status_raw, status_raw) - video_body = response_data.get("video") + video_body: Final = response_data.get("video") video_meta: Final = video_body or {} # mutable-ok: provider JSON body and base-class dict signature video_url: Final = video_meta.get("url") if isinstance(video_meta, dict) else None seconds: Final = ( From ab4de2f02a42e3cceb4240cf2e2c9c4dafdd5db6 Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Sun, 27 Sep 2026 18:29:06 +0800 Subject: [PATCH 5/6] fix(xai): stay within the basedpyright budget in imagine transforms Use the public module-level http client for CDN downloads, drop the dead create-time video_url hidden param, narrow file-like and reference image inputs without subscripting the FileTypes union, and restore the upstream signature of encode_character_id_in_response. Co-authored-by: Cursor --- litellm/llms/xai/image_edit/transformation.py | 20 ++++--- litellm/llms/xai/videos/transformation.py | 6 +-- litellm/proxy/video_endpoints/utils.py | 4 +- tests/unit/llms/xai/test_xai_image_edit.py | 19 +++++++ .../llms/xai/test_xai_video_generation.py | 53 ++++++++----------- 5 files changed, 58 insertions(+), 44 deletions(-) diff --git a/litellm/llms/xai/image_edit/transformation.py b/litellm/llms/xai/image_edit/transformation.py index b20382a7acf..bdb7ae35510 100644 --- a/litellm/llms/xai/image_edit/transformation.py +++ b/litellm/llms/xai/image_edit/transformation.py @@ -1,6 +1,6 @@ import base64 from io import BufferedReader, BytesIO -from typing import TYPE_CHECKING, Final +from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable import httpx from httpx._types import RequestFiles @@ -32,6 +32,11 @@ _SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-clas _XAI_NATIVE_PARAMS: Final = frozenset({"aspect_ratio", "n", "resolution"}) +@runtime_checkable +class _Readable(Protocol): + def read(self) -> bytes | str: ... + + def _read_seekable(image: BytesIO | BufferedReader) -> bytes: current_pos: Final = image.tell() image.seek(0) @@ -222,11 +227,12 @@ class XAIImageEditConfig(BaseImageEditConfig): if isinstance(image, str): return {"url": image} # mutable-ok: provider JSON body and base-class dict signature if isinstance(image, dict): - if image.get("url"): - return {"url": str(image["url"])} # mutable-ok: provider JSON body and base-class dict signature - if image.get("file_id"): - file_id: Final = str(image["file_id"]) - return {"file_id": file_id} # mutable-ok: provider JSON body and base-class dict signature + url: Final = image.get("url") + if url: + return {"url": str(url)} # mutable-ok: provider JSON body and base-class dict signature + file_id: Final = image.get("file_id") + if file_id: + return {"file_id": str(file_id)} # mutable-ok: provider JSON body and base-class dict signature mime: Final = ImageEditRequestUtils.get_image_content_type(image) encoded: Final = base64.b64encode(self._read_all_bytes(image)).decode("utf-8") @@ -239,7 +245,7 @@ class XAIImageEditConfig(BaseImageEditConfig): return bytes(image) if isinstance(image, (BytesIO, BufferedReader)): return _read_seekable(image) - if hasattr(image, "read"): + if isinstance(image, _Readable): raw: Final = image.read() if isinstance(raw, str): return raw.encode("utf-8") diff --git a/litellm/llms/xai/videos/transformation.py b/litellm/llms/xai/videos/transformation.py index 5b4cfb94fb4..4ad01b4ea2c 100644 --- a/litellm/llms/xai/videos/transformation.py +++ b/litellm/llms/xai/videos/transformation.py @@ -11,8 +11,6 @@ from litellm.litellm_core_utils.url_utils import async_safe_get, encode_url_path from litellm.llms.base_llm.videos.transformation import BaseVideoConfig from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, - HTTPHandler, - _get_httpx_client, get_async_httpx_client, ) from litellm.llms.xai.common_utils import XAIModelInfo @@ -248,7 +246,6 @@ class XAIVideoConfig(BaseVideoConfig): video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model) usage_body: Final = usage if isinstance(usage, dict) else None video_obj.usage = usage_body or {} # mutable-ok: provider JSON body and base-class dict signature - video_obj._hidden_params["video_url"] = None return video_obj def _video_resource_url(self, api_base: str, video_id: str) -> str: @@ -342,8 +339,7 @@ class XAIVideoConfig(BaseVideoConfig): url: Final = self._video_cdn_url(raw_response) if url is None: return raw_response.content - httpx_client: Final[HTTPHandler] = _get_httpx_client() - video_response: Final = safe_get(httpx_client, url) + video_response: Final = safe_get(litellm.module_level_client, url) video_response.raise_for_status() return video_response.content diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index 6e305bdff2a..b76a20d2642 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -1,5 +1,5 @@ from collections.abc import Mapping, Sequence -from typing import Final, Protocol +from typing import Any, Final, Protocol import orjson @@ -71,7 +71,7 @@ def get_custom_provider_from_data(data: Mapping[str, object]) -> str | None: return None -def encode_character_id_in_response(response: object, custom_llm_provider: str, model_id: str | None) -> object: +def encode_character_id_in_response(response: Any, custom_llm_provider: str, model_id: str | None) -> Any: if isinstance(response, dict) and response.get("id"): response["id"] = encode_character_id_with_provider( character_id=response["id"], diff --git a/tests/unit/llms/xai/test_xai_image_edit.py b/tests/unit/llms/xai/test_xai_image_edit.py index 575cbf459b6..3e996b94f4b 100644 --- a/tests/unit/llms/xai/test_xai_image_edit.py +++ b/tests/unit/llms/xai/test_xai_image_edit.py @@ -1,3 +1,4 @@ +import base64 from unittest.mock import MagicMock, patch import httpx @@ -134,6 +135,24 @@ def test_transform_bytes_to_data_uri_and_response(): assert response.data[0].url == "https://imgen.x.ai/edited.jpeg" +def test_transform_duck_typed_reader_to_data_uri(): + from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig + + class Reader: + def read(self) -> bytes: + return b"\x89PNG\r\n\x1a\nfakepng" + + data, _ = XAIImageEditConfig().transform_image_edit_request( + model="grok-imagine-image", + prompt="make it night", + image=Reader(), + image_edit_optional_request_params={}, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + assert data["image"]["url"] == "data:image/png;base64," + base64.b64encode(b"\x89PNG\r\n\x1a\nfakepng").decode() + + def test_transform_http_url_passthrough(): from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig diff --git a/tests/unit/llms/xai/test_xai_video_generation.py b/tests/unit/llms/xai/test_xai_video_generation.py index 261ce2e88c0..446b6fe47e2 100644 --- a/tests/unit/llms/xai/test_xai_video_generation.py +++ b/tests/unit/llms/xai/test_xai_video_generation.py @@ -196,13 +196,13 @@ 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 - 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: + with ( + patch("litellm.module_level_client", 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" safe_get.assert_called_once_with(cdn, "https://vidgen.x.ai/x.mp4") @@ -219,20 +219,20 @@ 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 - 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, 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() - ) + with ( + patch( + "litellm.llms.xai.videos.transformation.get_async_httpx_client", + return_value=async_client, + ), + patch("litellm.module_level_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() + assert sync_client.method_calls == [] async_safe_get.assert_awaited_once_with(async_client, "https://vidgen.x.ai/x.mp4") @@ -265,10 +265,7 @@ def test_content_response_rejects_internal_cdn_host(): 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 patch("litellm.module_level_client", MagicMock(get=boom)): with pytest.raises((SSRFError, ValueError)): XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock()) @@ -309,9 +306,7 @@ async def test_async_content_response_rejects_internal_cdn_host(): return_value=MagicMock(get=boom), ): with pytest.raises((SSRFError, ValueError)): - await XAIVideoConfig().async_transform_video_content_response( - status, logging_obj=MagicMock() - ) + await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock()) @pytest.mark.asyncio @@ -329,9 +324,7 @@ async def test_async_content_response_fetches_public_cdn_via_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() - ) + 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" From c510bd1aad2e17d52ee72d979aeaf24a6692910e Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Sun, 27 Sep 2026 18:50:13 +0800 Subject: [PATCH 6/6] fix(proxy): enforce per-caller ownership on proxy-created videos Record the creating caller's owner scope in LiteLLM_ManagedObjectTable (file_purpose=video, keyed by the provider-native video id) when videos are created, remixed, edited, or extended, and reject status, content, remix, edit, and extension calls on a video the caller does not own with 403. Proxy admins bypass the check, untracked videos are admin-only, and list results are filtered to the caller's videos. Without a database the checks are skipped, matching previous behavior. Co-authored-by: Cursor --- litellm/proxy/video_endpoints/endpoints.py | 19 +- litellm/proxy/video_endpoints/ownership.py | 167 ++++++++++++++ .../proxy/video_endpoints/test_ownership.py | 218 ++++++++++++++++++ 3 files changed, 403 insertions(+), 1 deletion(-) create mode 100644 litellm/proxy/video_endpoints/ownership.py create mode 100644 tests/test_litellm/proxy/video_endpoints/test_ownership.py diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 540860c1a5b..9c31876d298 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -16,6 +16,11 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_query, ) from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio +from litellm.proxy.video_endpoints.ownership import ( + assert_user_can_access_video, + filter_video_list_for_caller, + record_video_owner, +) from litellm.proxy.video_endpoints.utils import ( encode_character_id_in_response, extract_model_from_target_model_names, @@ -116,6 +121,7 @@ async def video_generation( version=version, ) else: + await record_video_owner(generated, user_api_key_dict) return generated @@ -203,7 +209,7 @@ async def video_list( version=version, ) else: - return listed + return await filter_video_list_for_caller(listed, user_api_key_dict) @router.get( @@ -250,6 +256,8 @@ async def video_status( version, ) + await assert_user_can_access_video(video_id, user_api_key_dict) + # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} @@ -353,6 +361,8 @@ async def video_content( version, ) + await assert_user_can_access_video(video_id, user_api_key_dict) + # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} @@ -462,6 +472,8 @@ async def video_remix( version, ) + await assert_user_can_access_video(video_id, user_api_key_dict) + data: Final = await _read_request_body(request=request) data["video_id"] = video_id @@ -514,6 +526,7 @@ async def video_remix( version=version, ) else: + await record_video_owner(remixed, user_api_key_dict) return remixed @@ -780,6 +793,7 @@ async def video_edit( data["video_id"] = "" else: data["video_id"] = video_reference_to_id(uploaded_video) + await assert_user_can_access_video(data["video_id"], user_api_key_dict) decoded: Final = decode_video_id_with_provider(data["video_id"]) provider_from_id: Final = decoded.get("custom_llm_provider") @@ -827,6 +841,7 @@ async def video_edit( version=version, ) else: + await record_video_owner(edited, user_api_key_dict) return edited @@ -877,6 +892,7 @@ async def video_extension( data: Final = await _read_request_body(request=request) data["video_id"] = video_reference_to_id(data.pop("video", None)) + await assert_user_can_access_video(data["video_id"], user_api_key_dict) decoded: Final = decode_video_id_with_provider(data["video_id"]) provider_from_id: Final = decoded.get("custom_llm_provider") @@ -924,4 +940,5 @@ async def video_extension( version=version, ) else: + await record_video_owner(extended, user_api_key_dict) return extended diff --git a/litellm/proxy/video_endpoints/ownership.py b/litellm/proxy/video_endpoints/ownership.py new file mode 100644 index 00000000000..49d8fe0eb13 --- /dev/null +++ b/litellm/proxy/video_endpoints/ownership.py @@ -0,0 +1,167 @@ +"""Per-caller ownership of videos created through the proxy. + +Ownership rows live in ``LiteLLM_ManagedObjectTable`` with ``file_purpose="video"`` +(the same table and owner scopes container ownership uses). A row is keyed by the +provider-native video id only, so re-wrapping an id with a different provider or +model_id cannot dodge the lookup. Proxy admins bypass the check; for everyone else, +a video without a row is treated as admin-only. + +Without a connected database there is nowhere to record ownership: recording and +checks are skipped and videos remain reachable by any caller allowed on the route. +""" + +import json +from collections.abc import Mapping, Sequence +from typing import TYPE_CHECKING, Final + +from fastapi import HTTPException +from typing_extensions import TypeIs # noqa: TID251 # narrows untyped wire payloads without a runtime conversion + +from litellm._logging import verbose_proxy_logger +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.resource_ownership import ( + get_primary_resource_owner_scope, + get_resource_owner_scopes, + is_proxy_admin, + user_can_access_resource_owner, +) +from litellm.repositories.table_repositories import ManagedObjectRepository +from litellm.types.videos.utils import extract_original_video_id + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +VIDEO_OBJECT_PURPOSE: Final = "video" + +# Status is polled; a short-lived cache of known owners keeps polling off the DB. +# Only positive answers are cached, so a missing row is always re-checked. +_VIDEO_OWNER_CACHE: Final = InMemoryCache(max_size_in_memory=10000, default_ttl=60) + + +def _video_model_object_id(video_id: str) -> str: + return f"{VIDEO_OBJECT_PURPOSE}:{extract_original_video_id(video_id)}" + + +def _video_id_of(item: object) -> str | None: + match item: + case {"id": str(video_id)} if video_id: + return video_id + case object(id=str(video_id)) if video_id: + return video_id + case _: + return None + + +def _is_object_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: provider JSON objects have str keys + return isinstance(value, Mapping) + + +def _is_object_sequence(value: object) -> TypeIs[Sequence[object]]: # guard-ok: a list is a Sequence of anything + return isinstance(value, list) + + +def _get_prisma_client() -> "PrismaClient | None": + from litellm.proxy.proxy_server import prisma_client + + return prisma_client + + +async def record_video_owner(response: object, user_api_key_dict: UserAPIKeyAuth) -> None: + """Stamp the caller as owner of a video the provider just created. + + Failures are logged rather than raised: the provider job already exists and is + billed, so the caller still gets its id; the untracked video is admin-only. + """ + prisma_client: Final = _get_prisma_client() + if prisma_client is None: + return + video_id: Final = _video_id_of(response) + owner: Final = get_primary_resource_owner_scope(user_api_key_dict) + if video_id is None or owner is None: + verbose_proxy_logger.warning( + "Skipping video ownership tracking: response has no id or caller has no identity scope" + ) + return + model_object_id: Final = _video_model_object_id(video_id) + row: Final = { # mutable-ok: prisma rows are plain dicts + "unified_object_id": video_id, + "model_object_id": model_object_id, + "file_object": json.dumps({"id": video_id, "object": "video"}), # mutable-ok: serialized immediately + "file_purpose": VIDEO_OBJECT_PURPOSE, + "created_by": owner, + "updated_by": owner, + } + try: + await ManagedObjectRepository(prisma_client).table.upsert( + where={"model_object_id": model_object_id}, # mutable-ok: prisma filters are plain dicts + data={"create": row, "update": {"updated_by": owner}}, # mutable-ok: prisma payloads are plain dicts + ) + except Exception as e: + verbose_proxy_logger.exception( + "Video ownership recording failed; video_id=%s is untracked and admin-only: %s", video_id, e + ) + + +async def _get_video_owner(prisma_client: "PrismaClient", model_object_id: str) -> str | None: + cached: Final = _VIDEO_OWNER_CACHE.get_cache(model_object_id) + if isinstance(cached, str): + return cached + row: Final = await ManagedObjectRepository(prisma_client).table.find_first( + where={ # mutable-ok: prisma filters are plain dicts + "model_object_id": model_object_id, + "file_purpose": VIDEO_OBJECT_PURPOSE, + } + ) + owner: Final = getattr(row, "created_by", None) if row is not None else None + if not isinstance(owner, str): + return None + _VIDEO_OWNER_CACHE.set_cache(model_object_id, owner) + return owner + + +async def assert_user_can_access_video(video_id: str, user_api_key_dict: UserAPIKeyAuth) -> None: + """Raise 403 unless the caller owns ``video_id`` (or is a proxy admin).""" + if not video_id or is_proxy_admin(user_api_key_dict): + return + prisma_client: Final = _get_prisma_client() + if prisma_client is None: + return + owner: Final = await _get_video_owner(prisma_client, _video_model_object_id(video_id)) + if not user_can_access_resource_owner(owner, user_api_key_dict): + raise HTTPException(status_code=403, detail="Forbidden") + + +async def filter_video_list_for_caller(listed: object, user_api_key_dict: UserAPIKeyAuth) -> object: + """Drop videos the caller does not own from a provider list page. + + Pagination cursors are left as the provider returned them so ``after`` keeps + walking the provider's list even when a page has no videos owned by the caller. + """ + if is_proxy_admin(user_api_key_dict) or not _is_object_mapping(listed): + return listed + prisma_client: Final = _get_prisma_client() + items: Final = listed.get("data") + if prisma_client is None or not _is_object_sequence(items): + return listed + ids_by_item: Final = tuple((item, _video_id_of(item)) for item in items) + candidate_ids: Final = tuple( + _video_model_object_id(video_id) for _, video_id in ids_by_item if video_id is not None + ) + owner_scopes: Final = get_resource_owner_scopes(user_api_key_dict) + rows: Final = ( + await ManagedObjectRepository(prisma_client).table.find_many( + where={ # mutable-ok: prisma filters are plain dicts + "model_object_id": {"in": list(candidate_ids)}, # mutable-ok: prisma filters are plain dicts + "file_purpose": VIDEO_OBJECT_PURPOSE, + "created_by": {"in": owner_scopes}, # mutable-ok: prisma filters are plain dicts + } + ) + if candidate_ids and owner_scopes + else () + ) + owned: Final = frozenset(row.model_object_id for row in rows) + kept: Final = tuple( + item for item, video_id in ids_by_item if video_id is not None and _video_model_object_id(video_id) in owned + ) + return {**listed, "data": kept} # mutable-ok: provider list JSON body diff --git a/tests/test_litellm/proxy/video_endpoints/test_ownership.py b/tests/test_litellm/proxy/video_endpoints/test_ownership.py new file mode 100644 index 00000000000..64dd83149d7 --- /dev/null +++ b/tests/test_litellm/proxy/video_endpoints/test_ownership.py @@ -0,0 +1,218 @@ +""" +Ownership tests for the proxy video endpoints. + +Requests go through the real FastAPI routes and the real ownership module. Only two +I/O boundaries are replaced: the provider call (``base_process_llm_request``) and the +Prisma ``litellm_managedobjecttable`` delegate, which is an in-memory table honoring +the subset of the query API the ownership module uses. +""" + +import uuid +from types import SimpleNamespace + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +import litellm.proxy.proxy_server as proxy_server +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, hash_token +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing +from litellm.proxy.video_endpoints import endpoints +from litellm.types.videos.utils import encode_video_id_with_provider + +ALICE = UserAPIKeyAuth(user_id="alice", api_key=hash_token("sk-alice")) +BOB = UserAPIKeyAuth(user_id="bob", api_key=hash_token("sk-bob")) +KEY_ONLY_A = UserAPIKeyAuth(api_key=hash_token("sk-service-a")) +KEY_ONLY_B = UserAPIKeyAuth(api_key=hash_token("sk-service-b")) +ADMIN = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + +def _matches(row: SimpleNamespace, where: dict) -> bool: + for field, condition in where.items(): + value = getattr(row, field, None) + if isinstance(condition, dict): + if value not in condition["in"]: + return False + elif value != condition: + return False + return True + + +class InMemoryManagedObjectTable: + def __init__(self, fail_writes: bool = False): + self.rows: dict[str, SimpleNamespace] = {} + self.fail_writes = fail_writes + + async def upsert(self, where: dict, data: dict) -> SimpleNamespace: + if self.fail_writes: + raise RuntimeError("database unavailable") + key = where["model_object_id"] + if key in self.rows: + self.rows[key] = SimpleNamespace(**{**vars(self.rows[key]), **data["update"]}) + else: + self.rows[key] = SimpleNamespace(**data["create"]) + return self.rows[key] + + async def find_first(self, where: dict) -> SimpleNamespace | None: + return next((row for row in self.rows.values() if _matches(row, where)), None) + + async def find_many(self, where: dict) -> list[SimpleNamespace]: + return [row for row in self.rows.values() if _matches(row, where)] + + +class Provider: + """The provider behind ``base_process_llm_request``: creates and serves videos.""" + + def __init__(self): + self.calls: list[str] = [] + + async def respond(self, processor, *, route_type: str, **kwargs): + self.calls.append(route_type) + if route_type in ("avideo_generation", "avideo_remix", "avideo_edit", "avideo_extension"): + return {"id": f"video_{uuid.uuid4().hex}", "object": "video", "status": "queued"} + if route_type == "avideo_content": + return b"mp4-bytes" + if route_type == "avideo_list": + return {"object": "list", "data": processor.data["listed"], "has_more": False} + return {"id": processor.data["video_id"], "object": "video", "status": "completed"} + + +@pytest.fixture +def provider(monkeypatch) -> Provider: + fake = Provider() + + async def base_process_llm_request(processor, **kwargs): + return await fake.respond(processor, **kwargs) + + monkeypatch.setattr(ProxyBaseLLMRequestProcessing, "base_process_llm_request", base_process_llm_request) + return fake + + +@pytest.fixture +def table(monkeypatch) -> InMemoryManagedObjectTable: + rows = InMemoryManagedObjectTable() + monkeypatch.setattr( + proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(litellm_managedobjecttable=rows)) + ) + return rows + + +def _client(auth: UserAPIKeyAuth) -> TestClient: + app = FastAPI() + app.include_router(endpoints.router) + app.dependency_overrides[user_api_key_auth] = lambda: auth + return TestClient(app) + + +def _create_video(auth: UserAPIKeyAuth) -> str: + response = _client(auth).post("/v1/videos", json={"model": "sora-2", "prompt": "a sunset"}) + assert response.status_code == 200 + return response.json()["id"] + + +def _access(auth: UserAPIKeyAuth, video_id: str) -> dict[str, int]: + client = _client(auth) + return { + "status": client.get(f"/v1/videos/{video_id}").status_code, + "content": client.get(f"/v1/videos/{video_id}/content").status_code, + "remix": client.post(f"/v1/videos/{video_id}/remix", json={"prompt": "again"}).status_code, + "edit": client.post("/v1/videos/edits", json={"prompt": "brighter", "video": {"id": video_id}}).status_code, + "extension": client.post( + "/v1/videos/extensions", json={"prompt": "continue", "video": {"id": video_id}} + ).status_code, + } + + +ALL_OK = {"status": 200, "content": 200, "remix": 200, "edit": 200, "extension": 200} +ALL_FORBIDDEN = {"status": 403, "content": 403, "remix": 403, "edit": 403, "extension": 403} + + +def test_owner_can_use_every_video_route(provider, table): + video_id = _create_video(ALICE) + + assert table.rows[f"video:{video_id}"].created_by == "alice" + assert _access(ALICE, video_id) == ALL_OK + + +def test_other_user_is_forbidden_and_the_provider_is_never_called(provider, table): + video_id = _create_video(ALICE) + provider.calls.clear() + + assert _access(BOB, video_id) == ALL_FORBIDDEN + assert provider.calls == [] + + +def test_key_scoped_owner_blocks_a_different_key(provider, table): + video_id = _create_video(KEY_ONLY_A) + + assert _access(KEY_ONLY_B, video_id) == ALL_FORBIDDEN + assert _access(KEY_ONLY_A, video_id) == ALL_OK + + +def test_proxy_admin_can_access_any_video(provider, table): + video_id = _create_video(ALICE) + + assert _access(ADMIN, video_id) == ALL_OK + + +def test_untracked_video_is_admin_only_when_a_database_is_connected(provider, table): + video_id = f"video_{uuid.uuid4().hex}" + + assert _access(ALICE, video_id) == ALL_FORBIDDEN + assert _access(ADMIN, video_id) == ALL_OK + + +def test_rewrapping_the_provider_id_does_not_bypass_the_owner_check(provider, table): + video_id = _create_video(ALICE) + rewrapped = encode_video_id_with_provider(video_id, "azure", "attacker-deployment") + + assert rewrapped != video_id + assert _access(BOB, rewrapped) == ALL_FORBIDDEN + + +def test_videos_derived_from_an_owned_video_belong_to_the_caller(provider, table): + video_id = _create_video(ALICE) + remixed = _client(ALICE).post(f"/v1/videos/{video_id}/remix", json={"prompt": "again"}).json()["id"] + + assert _access(ALICE, remixed)["status"] == 200 + assert _access(BOB, remixed)["status"] == 403 + + +def test_list_only_returns_the_callers_videos(provider, table): + alice_video = _create_video(ALICE) + bob_video = _create_video(BOB) + listed = [{"id": alice_video, "object": "video"}, {"id": bob_video, "object": "video"}] + provider_listing = {"listed": listed} + + def list_as(auth: UserAPIKeyAuth) -> list[str]: + async def with_listing(processor, **kwargs): + processor.data.update(provider_listing) + return await provider.respond(processor, **kwargs) + + with pytest.MonkeyPatch.context() as mp: + mp.setattr(ProxyBaseLLMRequestProcessing, "base_process_llm_request", with_listing) + response = _client(auth).get("/v1/videos") + assert response.status_code == 200 + return [item["id"] for item in response.json()["data"]] + + assert list_as(ALICE) == [alice_video] + assert list_as(BOB) == [bob_video] + assert list_as(ADMIN) == [alice_video, bob_video] + + +def test_without_a_database_ownership_is_not_enforced(provider, monkeypatch): + monkeypatch.setattr(proxy_server, "prisma_client", None) + video_id = _create_video(ALICE) + + assert _access(BOB, video_id) == ALL_OK + + +def test_failed_ownership_write_still_returns_the_created_video(provider, table): + table.fail_writes = True + + video_id = _create_video(ALICE) + + assert video_id.startswith("video_") + assert table.rows == {} + assert _access(ALICE, video_id)["status"] == 403