From 0f8e5279a7b910bc1b8c25162363c0ffc077f553 Mon Sep 17 00:00:00 2001 From: hx <1367557521@qq.com> Date: Wed, 23 Sep 2026 09:44:16 +0800 Subject: [PATCH] fix(proxy): format xai transforms and bind video downloads to the creating key --- litellm/llms/xai/image_edit/transformation.py | 40 ++++++++--- litellm/llms/xai/image_generation/__init__.py | 5 +- .../xai/image_generation/transformation.py | 21 ++++-- litellm/llms/xai/videos/transformation.py | 70 ++++++++++++++----- litellm/proxy/image_endpoints/endpoints.py | 36 +++++++--- litellm/proxy/video_endpoints/endpoints.py | 22 +++++- litellm/proxy/video_endpoints/utils.py | 33 ++++++++- litellm/types/videos/main.py | 1 + litellm/types/videos/utils.py | 10 ++- .../proxy/video_endpoints/test_utils.py | 32 +++++---- 10 files changed, 211 insertions(+), 59 deletions(-) diff --git a/litellm/llms/xai/image_edit/transformation.py b/litellm/llms/xai/image_edit/transformation.py index a109f1051b2..79b8b953d94 100644 --- a/litellm/llms/xai/image_edit/transformation.py +++ b/litellm/llms/xai/image_edit/transformation.py @@ -41,7 +41,9 @@ def _read_seekable(image: BytesIO | BufferedReader) -> bytes: class XAIImageEditConfig(BaseImageEditConfig): - def get_supported_openai_params(self, model: str) -> list: # mutable-ok: provider JSON body and base-class dict signature + 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( @@ -52,7 +54,9 @@ class XAIImageEditConfig(BaseImageEditConfig): ) -> dict: # mutable-ok: provider JSON body and base-class dict signature supported: Final = frozenset(self.get_supported_openai_params(model)) allowed: Final = supported | _XAI_NATIVE_PARAMS - incoming: Final = dict(image_edit_optional_params) # mutable-ok: provider JSON body and base-class dict signature + incoming: Final = dict( + image_edit_optional_params + ) # 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( @@ -61,7 +65,9 @@ class XAIImageEditConfig(BaseImageEditConfig): "Set drop_params=True to drop unsupported parameters." ) - mapped: Final = {key: value for key, value in incoming.items() if key in allowed} # mutable-ok: provider JSON body and base-class dict signature + mapped: Final = { + key: value for key, value in incoming.items() if key in allowed + } # 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 @@ -69,9 +75,13 @@ class XAIImageEditConfig(BaseImageEditConfig): n: Final = mapped.get("n") resolution: Final = mapped.get("resolution") return { # mutable-ok: provider JSON body and base-class dict signature - **({"aspect_ratio": aspect_ratio} if aspect_ratio is not None else {}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"aspect_ratio": aspect_ratio} if aspect_ratio is not None else {} + ), # mutable-ok: provider JSON body and base-class dict signature **({"n": int(n)} if n is not None else {}), # mutable-ok: provider JSON body and base-class dict signature - **({"resolution": resolution} if resolution is not None else {}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"resolution": resolution} if resolution is not None else {} + ), # mutable-ok: provider JSON body and base-class dict signature } def use_multipart_form_data(self) -> bool: @@ -155,8 +165,12 @@ class XAIImageEditConfig(BaseImageEditConfig): n: Final = image_edit_optional_request_params.get("n") request: Final[dict[str, object]] = { # mutable-ok: provider JSON body and base-class dict signature "model": XAIModelInfo.get_base_model(model) or model, - **({"prompt": prompt} if prompt is not None else {}), # mutable-ok: provider JSON body and base-class dict signature - **({"image": image_payloads[0]} if len(image_payloads) == 1 else {"images": list(image_payloads)}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"prompt": prompt} if prompt is not None else {} + ), # mutable-ok: provider JSON body and base-class dict signature + **( + {"image": image_payloads[0]} if len(image_payloads) == 1 else {"images": list(image_payloads)} + ), # mutable-ok: provider JSON body and base-class dict signature **{ # mutable-ok: provider JSON body and base-class dict signature key: image_edit_optional_request_params[key] for key in ("aspect_ratio", "resolution") @@ -197,19 +211,25 @@ class XAIImageEditConfig(BaseImageEditConfig): ) return ImageResponse(data=list(images)) # mutable-ok: provider JSON body and base-class dict signature - def _as_image_list(self, image: FileTypes | list[FileTypes]) -> tuple[FileTypes, ...]: # mutable-ok: provider JSON body and base-class dict signature + def _as_image_list( + self, image: FileTypes | list[FileTypes] + ) -> tuple[FileTypes, ...]: # mutable-ok: provider JSON body and base-class dict signature if isinstance(image, list): return tuple(item for item in image if item is not None) return (image,) - def _to_image_url(self, image: FileTypes) -> dict[str, str]: # mutable-ok: provider JSON body and base-class dict signature + def _to_image_url( + self, image: FileTypes + ) -> dict[str, str]: # mutable-ok: provider JSON body and base-class dict signature if isinstance(image, str): return {"url": image} # mutable-ok: provider JSON body and base-class dict signature if isinstance(image, dict): if image.get("url"): return {"url": str(image["url"])} # mutable-ok: provider JSON body and base-class dict signature if image.get("file_id"): - return {"file_id": str(image["file_id"])} # mutable-ok: provider JSON body and base-class dict signature + return { + "file_id": str(image["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") diff --git a/litellm/llms/xai/image_generation/__init__.py b/litellm/llms/xai/image_generation/__init__.py index c1b24b609c3..a9207e9a404 100644 --- a/litellm/llms/xai/image_generation/__init__.py +++ b/litellm/llms/xai/image_generation/__init__.py @@ -4,7 +4,10 @@ from litellm.llms.base_llm.image_generation.transformation import ( from .transformation import XAIImageGenerationConfig -__all__ = ["XAIImageGenerationConfig", "get_xai_image_generation_config"] # mutable-ok: provider JSON body and base-class dict signature +__all__ = [ + "XAIImageGenerationConfig", + "get_xai_image_generation_config", +] # mutable-ok: provider JSON body and base-class dict signature def get_xai_image_generation_config(model: str) -> BaseImageGenerationConfig: diff --git a/litellm/llms/xai/image_generation/transformation.py b/litellm/llms/xai/image_generation/transformation.py index 2db407eb450..e5a35733dd2 100644 --- a/litellm/llms/xai/image_generation/transformation.py +++ b/litellm/llms/xai/image_generation/transformation.py @@ -35,7 +35,9 @@ _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 + 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( @@ -55,14 +57,19 @@ class XAIImageGenerationConfig(BaseImageGenerationConfig): "Set drop_params=True to drop unsupported parameters." ) - merged: Final = {**optional_params, **{k: v for k, v in non_default_params.items() if k in allowed}} # mutable-ok: provider JSON body and base-class dict signature + merged: Final = { + **optional_params, + **{k: v for k, v in non_default_params.items() if k in allowed}, + } # 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") return { # mutable-ok: provider JSON body and base-class dict signature - **({"aspect_ratio": aspect_ratio} if aspect_ratio is not None else {}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"aspect_ratio": aspect_ratio} if aspect_ratio is not None else {} + ), # mutable-ok: provider JSON body and base-class dict signature **({"n": int(n)} if n is not None else {}), # mutable-ok: provider JSON body and base-class dict signature } @@ -142,7 +149,9 @@ class XAIImageGenerationConfig(BaseImageGenerationConfig): "model": XAIModelInfo.get_base_model(model) or model, "prompt": prompt, **( - {"aspect_ratio": optional_params["aspect_ratio"]} # mutable-ok: provider JSON body and base-class dict signature + { + "aspect_ratio": optional_params["aspect_ratio"] + } # mutable-ok: provider JSON body and base-class dict signature if optional_params.get("aspect_ratio") is not None else {} # mutable-ok: provider JSON body and base-class dict signature ), @@ -174,7 +183,9 @@ class XAIImageGenerationConfig(BaseImageGenerationConfig): 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 and base-class dict signature + additional_args={ + "complete_input_dict": request_data + }, # mutable-ok: provider JSON body and base-class dict signature original_response=response_data, ) diff --git a/litellm/llms/xai/videos/transformation.py b/litellm/llms/xai/videos/transformation.py index 039dce3a9fa..64e83c8c673 100644 --- a/litellm/llms/xai/videos/transformation.py +++ b/litellm/llms/xai/videos/transformation.py @@ -58,7 +58,9 @@ _STATUS_MAP: Final = { # mutable-ok: provider JSON body and base-class dict sig class XAIVideoConfig(BaseVideoConfig): - def get_supported_openai_params(self, model: str) -> list: # mutable-ok: provider JSON body and base-class dict signature + 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", @@ -75,26 +77,42 @@ class XAIVideoConfig(BaseVideoConfig): model: str, drop_params: bool, ) -> dict: # mutable-ok: provider JSON body and base-class dict signature - incoming: Final = dict(video_create_optional_params) # mutable-ok: provider JSON body and base-class dict signature + incoming: Final = dict( + video_create_optional_params + ) # mutable-ok: provider JSON body and base-class dict signature size: Final = incoming.get("size") 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 {"seconds", "size", "input_reference", "user", "extra_headers", "model"} # mutable-ok: provider JSON body and base-class dict signature + if key + not in { + "seconds", + "size", + "input_reference", + "user", + "extra_headers", + "model", + } # mutable-ok: provider JSON body and base-class dict signature }, **( - {"duration": _duration_from_seconds(incoming.get("seconds"))} # mutable-ok: provider JSON body and base-class dict signature + { + "duration": _duration_from_seconds(incoming.get("seconds")) + } # mutable-ok: provider JSON body and base-class dict signature if "seconds" in incoming and "duration" not in incoming else {} # mutable-ok: provider JSON body and base-class dict signature ), **( - {"aspect_ratio": incoming.get("aspect_ratio") or _SIZE_TO_ASPECT_RATIO.get(str(size), "16:9")} # mutable-ok: provider JSON body and base-class dict signature + { + "aspect_ratio": incoming.get("aspect_ratio") or _SIZE_TO_ASPECT_RATIO.get(str(size), "16:9") + } # mutable-ok: provider JSON body and base-class dict signature if size and "aspect_ratio" not in incoming else {} # mutable-ok: provider JSON body and base-class dict signature ), **( - {"image": incoming.get("image") or incoming.get("input_reference")} # mutable-ok: provider JSON body and base-class dict signature + { + "image": incoming.get("image") or incoming.get("input_reference") + } # mutable-ok: provider JSON body and base-class dict signature if incoming.get("input_reference") and "image" not in incoming else {} # mutable-ok: provider JSON body and base-class dict signature ), @@ -104,12 +122,16 @@ class XAIVideoConfig(BaseVideoConfig): self, api_base: str | None, api_key: str | None, - litellm_params: GenericLiteLLMParams | dict | None, # mutable-ok: provider JSON body and base-class dict signature + litellm_params: GenericLiteLLMParams + | dict + | None, # mutable-ok: provider JSON body and base-class dict signature ) -> str: from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth params: Final = ( - litellm_params.model_dump() if isinstance(litellm_params, GenericLiteLLMParams) else (litellm_params or {}) # mutable-ok: provider JSON body and base-class dict signature + 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("/") @@ -142,7 +164,9 @@ class XAIVideoConfig(BaseVideoConfig): should_use_xai_oauth, ) - params: Final = litellm_params.model_dump() if litellm_params is not None else {} # mutable-ok: provider JSON body and base-class dict signature + params: Final = ( + litellm_params.model_dump() if litellm_params is not None else {} + ) # 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: @@ -209,9 +233,13 @@ class XAIVideoConfig(BaseVideoConfig): return ( { # mutable-ok: provider JSON body and base-class dict signature "model": XAIModelInfo.get_base_model(model) or model, - **({"prompt": prompt} if prompt else {}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"prompt": prompt} if prompt else {} + ), # mutable-ok: provider JSON body and base-class dict signature **copied, - **({"duration": 6} if "duration" not in copied else {}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"duration": 6} if "duration" not in copied else {} + ), # mutable-ok: provider JSON body and base-class dict signature }, [], # mutable-ok: provider JSON body and base-class dict signature api_base, @@ -241,7 +269,9 @@ class XAIVideoConfig(BaseVideoConfig): ) if custom_llm_provider: video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model) - video_obj.usage = usage if isinstance(usage, dict) else {} # mutable-ok: provider JSON body and base-class dict signature + video_obj.usage = ( + usage if isinstance(usage, dict) else {} + ) # mutable-ok: provider JSON body and base-class dict signature video_obj._hidden_params["video_url"] = None return video_obj @@ -259,7 +289,9 @@ class XAIVideoConfig(BaseVideoConfig): 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 + 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, @@ -270,7 +302,9 @@ class XAIVideoConfig(BaseVideoConfig): response_data: Final = raw_response.json() status_raw: Final = str(response_data.get("status") or "processing").lower() status: Final = _STATUS_MAP.get(status_raw, status_raw) - video_meta: Final = response_data.get("video") or {} # mutable-ok: provider JSON body and base-class dict signature + video_meta: Final = ( + response_data.get("video") 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")) @@ -291,7 +325,9 @@ class XAIVideoConfig(BaseVideoConfig): 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 + 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": @@ -306,7 +342,9 @@ class XAIVideoConfig(BaseVideoConfig): 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 + 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() diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index 615e821270d..6d65748387f 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -60,7 +60,9 @@ async def batch_to_bytesio( """ if not uploads: return None - return [await uploadfile_to_bytesio(u) for u in uploads] # mutable-ok: provider JSON body and base-class dict signature + return [ + await uploadfile_to_bytesio(u) for u in uploads + ] # mutable-ok: provider JSON body and base-class dict signature def _is_image_reference_string(value: str) -> bool: @@ -98,7 +100,9 @@ async def _normalize_image_values(values: tuple[object, ...], field: str) -> obj return list(coerced) # mutable-ok: provider JSON body and base-class dict signature -def _json_image_values(data: dict[str, object], field: str) -> tuple[object, ...]: # mutable-ok: provider JSON body and base-class dict signature +def _json_image_values( + data: dict[str, object], field: str +) -> tuple[object, ...]: # mutable-ok: provider JSON body and base-class dict signature if field not in data: return () raw: Final = data[field] @@ -113,7 +117,9 @@ async def _normalized_image_edit_fields( image: Final = await _normalize_image_values(values_by_field["image"], "image") mask: Final = await _normalize_image_values(values_by_field["mask"], "mask") return { # mutable-ok: provider JSON body and base-class dict signature - **({"image": image} if image is not None else {}), # mutable-ok: provider JSON body and base-class dict signature + **( + {"image": image} if image is not None else {} + ), # mutable-ok: provider JSON body and base-class dict signature **({"mask": mask} if mask is not None else {}), # mutable-ok: provider JSON body and base-class dict signature } @@ -125,9 +131,13 @@ async def _image_edit_assets_from_request( form: Final = await request.form() if _is_form_content_type(request.headers.get("content-type", "")) else None if form is None: return { # mutable-ok: provider JSON body and base-class dict signature - **{key: value for key, value in data.items() if key not in {"image[]", "mask[]"}}, # mutable-ok: provider JSON body and base-class dict signature + **{ + key: value for key, value in data.items() if key not in {"image[]", "mask[]"} + }, # mutable-ok: provider JSON body and base-class dict signature **await _normalized_image_edit_fields( - {field: _json_image_values(data, field) for field, _alias in _IMAGE_EDIT_FILE_FIELDS} # mutable-ok: provider JSON body and base-class dict signature + { + field: _json_image_values(data, field) for field, _alias in _IMAGE_EDIT_FILE_FIELDS + } # mutable-ok: provider JSON body and base-class dict signature ), } @@ -143,9 +153,13 @@ async def _image_edit_assets_from_request( detail=f"Cannot specify both '{conflicts[0]}' and '{conflicts[0]}[]'", ) return { # mutable-ok: provider JSON body and base-class dict signature - **{key: value for key, value in data.items() if key not in {"image[]", "mask[]"}}, # mutable-ok: provider JSON body and base-class dict signature + **{ + key: value for key, value in data.items() if key not in {"image[]", "mask[]"} + }, # mutable-ok: provider JSON body and base-class dict signature **await _normalized_image_edit_fields( - {field: form_values[field] or form_values[alias] for field, alias in _IMAGE_EDIT_FILE_FIELDS} # mutable-ok: provider JSON body and base-class dict signature + { + field: form_values[field] or form_values[alias] for field, alias in _IMAGE_EDIT_FILE_FIELDS + } # mutable-ok: provider JSON body and base-class dict signature ), } @@ -257,7 +271,9 @@ async def image_generation( ) ### RESPONSE HEADERS ### - hidden_params: Final = getattr(response, "_hidden_params", {}) or {} # mutable-ok: provider JSON body and base-class dict signature + hidden_params: Final = ( + getattr(response, "_hidden_params", {}) or {} + ) # mutable-ok: provider JSON body and base-class dict signature model_id: Final = hidden_params.get("model_id", None) or "" cache_key: Final = hidden_params.get("cache_key", None) or "" api_base: Final = hidden_params.get("api_base", None) or "" @@ -373,7 +389,9 @@ async def image_edit_api( with_assets: Final = await _image_edit_assets_from_request(request, parsed_body) data: Final = { # mutable-ok: provider JSON body and base-class dict signature **with_assets, - **({} if "prompt" in with_assets else {"prompt": None}), # mutable-ok: provider JSON body and base-class dict signature + **( + {} if "prompt" in with_assets else {"prompt": None} + ), # mutable-ok: provider JSON body and base-class dict signature "model": ( model or general_settings.get("image_generation_model", None) or user_model or with_assets.get("model") ), diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 6da6301aac5..84fec0e33ee 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -17,11 +17,14 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( ) from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio from litellm.proxy.video_endpoints.utils import ( + assert_video_owner, encode_character_id_in_response, extract_model_from_target_model_names, get_custom_provider_from_data, infer_video_provider_from_model, resolve_video_request_model, + stamp_video_owner, + video_owner_from_key, video_reference_to_id, ) from litellm.types.videos.utils import ( @@ -91,7 +94,7 @@ async def video_generation( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + result = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -109,6 +112,9 @@ async def video_generation( user_api_base=user_api_base, version=version, ) + return _stamp_generated_video_owner( + result, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key) + ) except Exception as e: raise await processor._handle_llm_api_exception( e=e, @@ -118,6 +124,16 @@ async def video_generation( ) +def _stamp_generated_video_owner(result: object, owner: str | None) -> object: + video_id: Final = getattr(result, "id", None) + if isinstance(video_id, str): + result.id = stamp_video_owner(video_id, owner) # mutable-ok: response object id is rewritten before return + return result + if isinstance(result, dict) and isinstance(result.get("id"), str): + result["id"] = stamp_video_owner(result["id"], owner) # mutable-ok: JSON response body + return result + + @router.get( "/v1/videos", dependencies=[Depends(user_api_key_auth)], @@ -247,6 +263,8 @@ async def video_status( version, ) + assert_video_owner(video_id, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key)) + # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} @@ -348,6 +366,8 @@ async def video_content( version, ) + assert_video_owner(video_id, video_owner_from_key(user_api_key_dict.token, user_api_key_dict.api_key)) + # Create data with video_id data: Final[dict[str, object]] = {"video_id": video_id} diff --git a/litellm/proxy/video_endpoints/utils.py b/litellm/proxy/video_endpoints/utils.py index 9795e925863..baa2aea81c3 100644 --- a/litellm/proxy/video_endpoints/utils.py +++ b/litellm/proxy/video_endpoints/utils.py @@ -3,13 +3,44 @@ from typing import Final, Protocol import orjson -from litellm.types.videos.utils import encode_character_id_with_provider +from litellm.proxy._types import ProxyException +from litellm.types.videos.utils import ( + decode_video_id_with_provider, + encode_character_id_with_provider, + encode_video_id_with_provider, +) class VideoModelIdResolver(Protocol): def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: ... +def video_owner_from_key(token: str | None, api_key: str | None) -> str | None: + return token or api_key + + +def assert_video_owner(video_id: str, owner: str | None) -> None: + recorded: Final = decode_video_id_with_provider(video_id).get("owner") + if recorded and recorded != owner: + raise ProxyException( + message="Video does not belong to this API key", + type="permission_error", + param="video_id", + code=403, + ) + + +def stamp_video_owner(video_id: str, owner: str | None) -> str: + if not owner: + return video_id + decoded: Final = decode_video_id_with_provider(video_id) + provider: Final = decoded.get("custom_llm_provider") + raw_id: Final = decoded.get("video_id") + if not provider or not raw_id or decoded.get("owner"): + return video_id + return encode_video_id_with_provider(raw_id, provider, decoded.get("model_id"), owner) + + def infer_video_provider_from_model(model: str | None) -> str | None: if not isinstance(model, str) or not model: return None diff --git a/litellm/types/videos/main.py b/litellm/types/videos/main.py index f4369fd95af..010bcfd9ddd 100644 --- a/litellm/types/videos/main.py +++ b/litellm/types/videos/main.py @@ -102,6 +102,7 @@ class DecodedVideoId(TypedDict, total=False): custom_llm_provider: str | None model_id: str | None video_id: str + owner: str | None class CharacterObject(BaseModel): diff --git a/litellm/types/videos/utils.py b/litellm/types/videos/utils.py index b23b2269543..e07f71b3240 100644 --- a/litellm/types/videos/utils.py +++ b/litellm/types/videos/utils.py @@ -35,7 +35,9 @@ def _add_base64_padding(value: str) -> str: return value -def encode_video_id_with_provider(video_id: str, provider: str, model_id: str | None = None) -> str: +def encode_video_id_with_provider( + video_id: str, provider: str, model_id: str | None = None, owner: str | None = None +) -> str: """Encode provider and model_id into video_id using base64.""" if not provider or not video_id: return video_id @@ -50,6 +52,8 @@ def encode_video_id_with_provider(video_id: str, provider: str, model_id: str | # ID is not encoded (even if it starts with video_), so encode it assembled_id = str(SpecialEnums.LITELLM_MANAGED_VIDEO_COMPLETE_STR.value).format(provider, model_id or "", video_id) + if owner: + assembled_id = f"{assembled_id};owner:{owner}" base64_encoded_id: Final[str] = base64.b64encode(assembled_id.encode("utf-8")).decode("utf-8") @@ -89,20 +93,24 @@ def decode_video_id_with_provider(encoded_video_id: str) -> DecodedVideoId: custom_llm_provider = None model_id = None decoded_video_id = encoded_video_id + owner = None if len(parts) >= 3: custom_llm_provider_part: Final = parts[0] model_id_part: Final = parts[1] video_id_part: Final = parts[2] + owner_part: Final = next((part for part in parts[3:] if part.startswith("owner:")), None) custom_llm_provider = custom_llm_provider_part.replace("litellm:custom_llm_provider:", "") model_id = model_id_part.replace("model_id:", "") decoded_video_id = video_id_part.replace("video_id:", "") + owner = owner_part.removeprefix("owner:") if owner_part else None return DecodedVideoId( custom_llm_provider=custom_llm_provider, model_id=model_id, video_id=decoded_video_id, + owner=owner, ) except Exception as e: verbose_logger.debug("Error decoding video_id '%s': %s", encoded_video_id, e) diff --git a/tests/test_litellm/proxy/video_endpoints/test_utils.py b/tests/test_litellm/proxy/video_endpoints/test_utils.py index af8af6cdcf2..3fb523944aa 100644 --- a/tests/test_litellm/proxy/video_endpoints/test_utils.py +++ b/tests/test_litellm/proxy/video_endpoints/test_utils.py @@ -13,21 +13,23 @@ is encode_character_id_with_provider, which runs for real; encoding assertions are checked by the genuine decode round-trip. """ - import pytest - +from litellm.proxy._types import ProxyException from litellm.proxy.video_endpoints.utils import ( + assert_video_owner, encode_character_id_in_response, extract_model_from_target_model_names, get_custom_provider_from_data, infer_video_provider_from_model, resolve_video_request_model, + stamp_video_owner, video_reference_to_id, ) from litellm.types.videos.utils import ( decode_character_id_with_provider, encode_character_id_with_provider, + encode_video_id_with_provider, ) # =========================================================================== # @@ -166,12 +168,7 @@ def test_provider__falsy_top_level_falls_through_to_extra_body(falsy): def test_provider__from_extra_body_dict(): - assert ( - get_custom_provider_from_data( - {"extra_body": {"custom_llm_provider": "bedrock"}} - ) - == "bedrock" - ) + assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": "bedrock"}}) == "bedrock" def test_provider__from_extra_body_json_string(): @@ -189,10 +186,7 @@ def test_provider__json_string_parsing_to_non_dict_is_none(): def test_provider__extra_body_provider_not_a_string_is_none(): - assert ( - get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) - is None - ) + assert get_custom_provider_from_data({"extra_body": {"custom_llm_provider": 123}}) is None @pytest.mark.parametrize( @@ -224,9 +218,7 @@ def test_encode__dict_with_id_mutates_in_place_and_preserves_other_keys(): assert out is response # same dict, mutated in place assert out["object"] == "character" and out["name"] == "hero" - assert out["id"] == encode_character_id_with_provider( - "char_raw", "azure", "model-1" - ) + assert out["id"] == encode_character_id_with_provider("char_raw", "azure", "model-1") decoded = decode_character_id_with_provider(out["id"]) assert decoded["custom_llm_provider"] == "azure" assert decoded["model_id"] == "model-1" @@ -264,6 +256,16 @@ def test_encode__object_non_str_or_empty_id_unchanged(bad_id): assert resp.id == bad_id # untouched +def test_stamp_and_assert_video_owner_round_trip(): + encoded = encode_video_id_with_provider("req_1", "xai", "grok-imagine-video") + stamped = stamp_video_owner(encoded, "key-hash") + assert stamped != encoded + assert_video_owner(stamped, "key-hash") + with pytest.raises(ProxyException): + assert_video_owner(stamped, "other-key") + assert_video_owner(encoded, "other-key") + + def test_encode__object_without_id_attr_returned_unchanged(): resp = _Resp() out = encode_character_id_in_response(resp, "azure", "model-1")