mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): format xai transforms and bind video downloads to the creating key
This commit is contained in:
parent
55a34c313f
commit
0f8e5279a7
10 changed files with 211 additions and 59 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue