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..bdb7ae35510 --- /dev/null +++ b/litellm/llms/xai/image_edit/transformation.py @@ -0,0 +1,253 @@ +import base64 +from io import BufferedReader, BytesIO +from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable + +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"}) + + +@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) + 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: 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: + 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: 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 ( + _SIZE_TO_ASPECT_RATIO.get(str(size), "1:1") if size else None + ) + n: Final = mapped.get("n") + resolution: Final = mapped.get("resolution") + 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 + + 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: 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 + 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], # 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,) + + 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): + 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") + 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 isinstance(image, _Readable): + 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..5db907fe809 --- /dev/null +++ b/litellm/llms/xai/image_generation/transformation.py @@ -0,0 +1,197 @@ +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: 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") + 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") + 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, + 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: 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, + **(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..4ad01b4ea2c --- /dev/null +++ b/litellm/llms/xai/videos/transformation.py @@ -0,0 +1,414 @@ +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, + 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: 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", + "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: 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: 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 + }, + **(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: get_complete_url receives the base-class dict + ) -> 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: 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) + 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: 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, + **(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: Final = usage if isinstance(usage, dict) else None + video_obj.usage = usage_body or {} # mutable-ok: provider JSON body and base-class dict signature + 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: 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 = ( + 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 + video_response: Final = safe_get(litellm.module_level_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 140e0cf0071..2f3e13fe87b 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -63428,7 +63428,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", @@ -63444,7 +63445,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", @@ -63461,7 +63463,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", @@ -63478,7 +63481,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", @@ -63495,7 +63499,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", @@ -63511,7 +63516,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", @@ -63527,7 +63533,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", @@ -63545,6 +63552,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", @@ -63563,6 +63573,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", @@ -63581,6 +63594,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", @@ -63599,6 +63615,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", @@ -63642,7 +63661,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 ad6e5857218..b825cee8160 100644 --- a/litellm/provider_endpoints_support_backup.json +++ b/litellm/provider_endpoints_support_backup.json @@ -2377,7 +2377,8 @@ "messages": true, "responses": true, "embeddings": false, - "image_generations": false, + "image_generations": true, + "image_edits": true, "audio_transcriptions": false, "audio_speech": false, "moderations": false, @@ -2385,7 +2386,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..9c31876d298 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -16,10 +16,16 @@ 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, get_custom_provider_from_data, + resolve_video_request_model, video_reference_to_id, ) from litellm.types.videos.utils import ( @@ -115,6 +121,7 @@ async def video_generation( version=version, ) else: + await record_video_owner(generated, user_api_key_dict) return generated @@ -202,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( @@ -249,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} @@ -256,23 +265,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 (None if resolved_model else "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 +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} @@ -366,12 +379,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: @@ -458,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 @@ -510,6 +526,7 @@ async def video_remix( version=version, ) else: + await record_video_owner(remixed, user_api_key_dict) return remixed @@ -776,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") @@ -823,6 +841,7 @@ async def video_edit( version=version, ) else: + await record_video_owner(edited, user_api_key_dict) return edited @@ -873,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") @@ -920,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/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index a38226cc253..b76a20d2642 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -1,16 +1,40 @@ -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Any, Final, Protocol import orjson from litellm.types.videos.utils import encode_character_id_with_provider -def extract_model_from_target_model_names(target_model_names: Any) -> str | None: +class VideoModelIdResolver(Protocol): + def resolve_model_name_from_model_id(self, model_id: str | None) -> str | 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): - target_model_names = [m.strip() for m in target_model_names.split(",") if m.strip()] - elif not isinstance(target_model_names, list): - return None - return target_model_names[0] if target_model_names else None + 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 +49,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") diff --git a/litellm/utils.py b/litellm/utils.py index 13a46840431..833aa60b42b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9603,6 +9603,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 @@ -9634,6 +9640,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 @@ -9772,6 +9782,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..32a56ad645a 100644 --- a/litellm/videos/main.py +++ b/litellm/videos/main.py @@ -30,6 +30,25 @@ from litellm.videos.utils import VideoGenerationRequestUtils llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler() +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 litellm.BadRequestError: + return None + return provider + + +def _provider_for_video_id(video_id: str, custom_llm_provider: str | None, model: object) -> str: + return ( + custom_llm_provider + or decode_video_id_with_provider(video_id).get("custom_llm_provider") + or _provider_from_model(model) + or "openai" + ) + + ##### Video Generation ####################### @client async def avideo_generation( @@ -317,10 +336,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 +428,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 +1032,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 140e0cf0071..2f3e13fe87b 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -63428,7 +63428,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", @@ -63444,7 +63445,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", @@ -63461,7 +63463,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", @@ -63478,7 +63481,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", @@ -63495,7 +63499,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", @@ -63511,7 +63516,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", @@ -63527,7 +63533,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", @@ -63545,6 +63552,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", @@ -63563,6 +63573,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", @@ -63581,6 +63594,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", @@ -63599,6 +63615,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", @@ -63642,7 +63661,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 44ef9363b64..528c2395112 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2723,7 +2723,8 @@ "messages": true, "responses": true, "embeddings": false, - "image_generations": false, + "image_generations": true, + "image_edits": true, "audio_transcriptions": false, "audio_speech": false, "moderations": false, @@ -2731,7 +2732,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..861b68e4953 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/video_endpoints/test_endpoints.py @@ -309,6 +309,40 @@ 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_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": "xai/grok-imagine-video-1.5"}, + ) + + harness.resolve_model.assert_not_called() + assert harness.processor_data() == { + "video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd", + "model": "xai/grok-imagine-video-1.5", + } + + # =========================================================================== # # GET /v1/videos/{video_id}/content - video_content # # =========================================================================== # @@ -365,6 +399,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_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 diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py index 9a2c208c075..5944fd7a394 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_utils.py +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -21,6 +21,7 @@ from litellm.proxy.video_endpoints.utils import ( encode_character_id_in_response, extract_model_from_target_model_names, get_custom_provider_from_data, + resolve_video_request_model, video_reference_to_id, ) from litellm.types.videos.utils import ( @@ -28,6 +29,52 @@ from litellm.types.videos.utils import ( encode_character_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" + ) + + # =========================================================================== # # extract_model_from_target_model_names # =========================================================================== # 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..3e996b94f4b --- /dev/null +++ b/tests/unit/llms/xai/test_xai_image_edit.py @@ -0,0 +1,186 @@ +import base64 +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_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 + + 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..446b6fe47e2 --- /dev/null +++ b/tests/unit/llms/xai/test_xai_video_generation.py @@ -0,0 +1,330 @@ +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.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") + + +@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.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" + assert sync_client.method_calls == [] + 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.module_level_client", 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..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,6 +179,47 @@ def test_video_content__plain_id_defaults_to_openai(seams): assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai" +@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="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") @@ -377,6 +419,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_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="xai/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