fix(proxy): format xai transforms and bind video downloads to the creating key

This commit is contained in:
hx 2026-09-23 09:44:16 +08:00
parent 55a34c313f
commit 0f8e5279a7
10 changed files with 211 additions and 59 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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