mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge c510bd1aad into f4a217d005
This commit is contained in:
commit
b6db2d034c
27 changed files with 2389 additions and 96 deletions
|
|
@ -384,6 +384,7 @@ def image_generation(
|
|||
litellm.LlmProviders.QWENCLOUD,
|
||||
litellm.LlmProviders.QWEN_AI_PLATFORM,
|
||||
litellm.LlmProviders.EDENAI,
|
||||
litellm.LlmProviders.XAI,
|
||||
):
|
||||
if image_generation_config is None:
|
||||
raise ValueError(f"image generation config is not supported for {custom_llm_provider}")
|
||||
|
|
@ -393,7 +394,7 @@ def image_generation(
|
|||
litellm_params_dict["api_base"] = _api_base
|
||||
|
||||
return llm_http_handler.image_generation_handler(
|
||||
api_key=api_key,
|
||||
api_key=api_key or dynamic_api_key,
|
||||
model=model,
|
||||
prompt=prompt,
|
||||
image_generation_provider_config=image_generation_config,
|
||||
|
|
|
|||
3
litellm/llms/xai/image_edit/__init__.py
Normal file
3
litellm/llms/xai/image_edit/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import XAIImageEditConfig
|
||||
|
||||
__all__ = ["XAIImageEditConfig"] # mutable-ok: provider JSON body and base-class dict signature
|
||||
253
litellm/llms/xai/image_edit/transformation.py
Normal file
253
litellm/llms/xai/image_edit/transformation.py
Normal file
|
|
@ -0,0 +1,253 @@
|
|||
import base64
|
||||
from io import BufferedReader, BytesIO
|
||||
from typing import TYPE_CHECKING, Final, Protocol, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
from litellm.constants import XAI_API_BASE
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.images.utils import ImageEditRequestUtils
|
||||
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
|
||||
from litellm.llms.xai.common_utils import XAIModelInfo
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.images.main import ImageEditOptionalRequestParams
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import FileTypes, ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
_SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature
|
||||
"1024x1024": "1:1",
|
||||
"1792x1024": "16:9",
|
||||
"1024x1792": "9:16",
|
||||
"1536x1024": "3:2",
|
||||
"1024x1536": "2:3",
|
||||
"1280x720": "16:9",
|
||||
"720x1280": "9:16",
|
||||
"1920x1080": "16:9",
|
||||
"1080x1920": "9:16",
|
||||
}
|
||||
_XAI_NATIVE_PARAMS: Final = frozenset({"aspect_ratio", "n", "resolution"})
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class _Readable(Protocol):
|
||||
def read(self) -> bytes | str: ...
|
||||
|
||||
|
||||
def _read_seekable(image: BytesIO | BufferedReader) -> bytes:
|
||||
current_pos: Final = image.tell()
|
||||
image.seek(0)
|
||||
data: Final = image.read()
|
||||
image.seek(current_pos)
|
||||
return data
|
||||
|
||||
|
||||
class XAIImageEditConfig(BaseImageEditConfig):
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> list: # mutable-ok: provider JSON body and base-class dict signature
|
||||
return ["n", "response_format", "size", "user"] # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
image_edit_optional_params: ImageEditOptionalRequestParams,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
supported: Final = frozenset(self.get_supported_openai_params(model))
|
||||
allowed: Final = supported | _XAI_NATIVE_PARAMS
|
||||
raw: Final = image_edit_optional_params
|
||||
incoming: Final = dict(raw) # mutable-ok: provider JSON body and base-class dict signature
|
||||
unknown: Final = tuple(key for key in incoming if key not in allowed)
|
||||
if unknown and not drop_params:
|
||||
raise ValueError(
|
||||
f"Parameter {unknown[0]} is not supported for model {model}. "
|
||||
f"Supported parameters are {sorted(allowed)}. "
|
||||
"Set drop_params=True to drop unsupported parameters."
|
||||
)
|
||||
|
||||
pairs: Final = ((key, value) for key, value in incoming.items() if key in allowed)
|
||||
mapped: Final = dict(pairs) # mutable-ok: provider JSON body and base-class dict signature
|
||||
size: Final = mapped.get("size")
|
||||
aspect_ratio: Final = mapped.get("aspect_ratio") or (
|
||||
_SIZE_TO_ASPECT_RATIO.get(str(size), "1:1") if size else None
|
||||
)
|
||||
n: Final = mapped.get("n")
|
||||
resolution: Final = mapped.get("resolution")
|
||||
fields: Final = (
|
||||
("aspect_ratio", aspect_ratio),
|
||||
("n", int(n) if n is not None else None),
|
||||
("resolution", resolution),
|
||||
)
|
||||
return {key: value for key, value in fields if value is not None} # mutable-ok: base class returns a dict
|
||||
|
||||
def use_multipart_form_data(self) -> bool:
|
||||
return False
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
model: str,
|
||||
api_base: str | None,
|
||||
litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> str:
|
||||
from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth
|
||||
|
||||
api_key: Final = litellm_params.get("api_key") if isinstance(litellm_params, dict) else None
|
||||
resolved_base: Final = (
|
||||
XAIOAuthAuthenticator().get_api_base()
|
||||
if should_use_xai_oauth(litellm_params) and not XAIModelInfo.get_api_key(api_key)
|
||||
else (api_base or get_secret_str("XAI_API_BASE") or get_secret_str("XAI_OAUTH_API_BASE") or XAI_API_BASE)
|
||||
)
|
||||
base: Final = (resolved_base or XAI_API_BASE).rstrip("/")
|
||||
if base.endswith("/v1"):
|
||||
return f"{base}/images/edits"
|
||||
return f"{base}/v1/images/edits"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
litellm_params: dict | None = None, # mutable-ok: provider JSON body and base-class dict signature
|
||||
api_base: str | None = None,
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
from litellm.llms.xai.oauth import (
|
||||
XAIOAuthAuthenticator,
|
||||
XAIOAuthError,
|
||||
should_use_xai_oauth,
|
||||
)
|
||||
|
||||
params: Final = litellm_params or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key)
|
||||
if should_use_xai_oauth(params) and not dynamic_api_key:
|
||||
try:
|
||||
headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().get_access_token()}"
|
||||
except XAIOAuthError as exc:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
llm_provider="xai",
|
||||
message=str(exc),
|
||||
) from exc
|
||||
else:
|
||||
if not dynamic_api_key:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
llm_provider="xai",
|
||||
message=(
|
||||
"Missing xAI credentials for image edit. Pass api_key / XAI_API_KEY, or set use_xai_oauth=True."
|
||||
),
|
||||
)
|
||||
headers["Authorization"] = f"Bearer {dynamic_api_key}"
|
||||
|
||||
if "content-type" not in headers and "Content-Type" not in headers:
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def transform_image_edit_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str | None,
|
||||
image: FileTypes | None,
|
||||
image_edit_optional_request_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> tuple[dict, RequestFiles]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
if image is None:
|
||||
raise ValueError("xAI image edit requires at least one reference image.")
|
||||
|
||||
image_payloads: Final = tuple(self._to_image_url(item) for item in self._as_image_list(image))
|
||||
if not image_payloads:
|
||||
raise ValueError("xAI image edit requires at least one reference image.")
|
||||
|
||||
n: Final = image_edit_optional_request_params.get("n")
|
||||
prompt_body: Final = {"prompt": prompt} if prompt is not None else None # mutable-ok: provider JSON body
|
||||
many: Final = {"images": list(image_payloads)} # mutable-ok: provider JSON body
|
||||
one: Final = image_payloads[0]
|
||||
image_body: Final = {"image": one} if len(image_payloads) == 1 else many # mutable-ok: provider JSON body
|
||||
prompt_field: Final = prompt_body or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
image_field: Final = image_body
|
||||
request: Final[dict[str, object]] = { # mutable-ok: provider JSON body and base-class dict signature
|
||||
"model": XAIModelInfo.get_base_model(model) or model,
|
||||
**prompt_field,
|
||||
**image_field,
|
||||
**{ # mutable-ok: provider JSON body and base-class dict signature
|
||||
key: image_edit_optional_request_params[key]
|
||||
for key in ("aspect_ratio", "resolution")
|
||||
if image_edit_optional_request_params.get(key) is not None
|
||||
},
|
||||
**({"n": int(n)} if n is not None else {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
}
|
||||
return request, [] # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def transform_image_edit_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
) -> ImageResponse:
|
||||
try:
|
||||
response_data: Final = raw_response.json()
|
||||
except Exception:
|
||||
raise self.get_error_class(
|
||||
error_message=raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
images: Final = tuple(
|
||||
ImageObject(
|
||||
url=item.get("url"),
|
||||
b64_json=item.get("b64_json") or item.get("b64"),
|
||||
)
|
||||
for item in response_data.get("data") or ()
|
||||
if isinstance(item, dict)
|
||||
)
|
||||
if not images:
|
||||
raise self.get_error_class(
|
||||
error_message=f"xAI image edit returned no image data: {response_data}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
return ImageResponse(data=list(images)) # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def _as_image_list(
|
||||
self,
|
||||
image: FileTypes | list[FileTypes], # mutable-ok: the proxy passes multi-image edits as a list of uploads
|
||||
) -> tuple[FileTypes, ...]:
|
||||
if isinstance(image, list):
|
||||
return tuple(item for item in image if item is not None)
|
||||
return (image,)
|
||||
|
||||
def _to_image_url(
|
||||
self, image: FileTypes
|
||||
) -> dict[str, str]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
if isinstance(image, str):
|
||||
return {"url": image} # mutable-ok: provider JSON body and base-class dict signature
|
||||
if isinstance(image, dict):
|
||||
url: Final = image.get("url")
|
||||
if url:
|
||||
return {"url": str(url)} # mutable-ok: provider JSON body and base-class dict signature
|
||||
file_id: Final = image.get("file_id")
|
||||
if file_id:
|
||||
return {"file_id": str(file_id)} # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
mime: Final = ImageEditRequestUtils.get_image_content_type(image)
|
||||
encoded: Final = base64.b64encode(self._read_all_bytes(image)).decode("utf-8")
|
||||
return {"url": f"data:{mime};base64,{encoded}"} # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def _read_all_bytes(self, image: FileTypes) -> bytes:
|
||||
if isinstance(image, bytes):
|
||||
return image
|
||||
if isinstance(image, bytearray):
|
||||
return bytes(image)
|
||||
if isinstance(image, (BytesIO, BufferedReader)):
|
||||
return _read_seekable(image)
|
||||
if isinstance(image, _Readable):
|
||||
raw: Final = image.read()
|
||||
if isinstance(raw, str):
|
||||
return raw.encode("utf-8")
|
||||
return bytes(raw)
|
||||
raise ValueError(f"Unsupported image input type for xAI image edit: {type(image)}")
|
||||
11
litellm/llms/xai/image_generation/__init__.py
Normal file
11
litellm/llms/xai/image_generation/__init__.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
|
||||
from .transformation import XAIImageGenerationConfig
|
||||
|
||||
__all__ = ("XAIImageGenerationConfig", "get_xai_image_generation_config")
|
||||
|
||||
|
||||
def get_xai_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
||||
return XAIImageGenerationConfig()
|
||||
197
litellm/llms/xai/image_generation/transformation.py
Normal file
197
litellm/llms/xai/image_generation/transformation.py
Normal file
|
|
@ -0,0 +1,197 @@
|
|||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.constants import XAI_API_BASE
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.llms.base_llm.image_generation.transformation import (
|
||||
BaseImageGenerationConfig,
|
||||
)
|
||||
from litellm.llms.xai.common_utils import XAIModelInfo
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
OpenAIImageGenerationOptionalParams,
|
||||
)
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import tiktoken
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
_SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature
|
||||
"1024x1024": "1:1",
|
||||
"1792x1024": "16:9",
|
||||
"1024x1792": "9:16",
|
||||
"1536x1024": "3:2",
|
||||
"1024x1536": "2:3",
|
||||
"1280x720": "16:9",
|
||||
"720x1280": "9:16",
|
||||
"1920x1080": "16:9",
|
||||
"1080x1920": "9:16",
|
||||
}
|
||||
_XAI_NATIVE_PARAMS: Final = frozenset({"aspect_ratio", "n"})
|
||||
|
||||
|
||||
class XAIImageGenerationConfig(BaseImageGenerationConfig):
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> list[OpenAIImageGenerationOptionalParams]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
return ["n", "response_format", "size", "user"] # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
supported_params: Final = frozenset(self.get_supported_openai_params(model))
|
||||
allowed: Final = supported_params | _XAI_NATIVE_PARAMS
|
||||
unknown: Final = tuple(key for key in non_default_params if key not in optional_params and key not in allowed)
|
||||
if unknown and not drop_params:
|
||||
raise ValueError(
|
||||
f"Parameter {unknown[0]} is not supported for model {model}. "
|
||||
f"Supported parameters are {sorted(allowed)}. "
|
||||
"Set drop_params=True to drop unsupported parameters."
|
||||
)
|
||||
|
||||
pairs: Final = ((k, v) for k, v in non_default_params.items() if k in allowed)
|
||||
native: Final = dict(pairs) # mutable-ok: provider JSON body and base-class dict signature
|
||||
merged: Final = {**optional_params, **native} # mutable-ok: provider JSON body and base-class dict signature
|
||||
size: Final = merged.get("size")
|
||||
aspect_ratio: Final = merged.get("aspect_ratio") or (
|
||||
_SIZE_TO_ASPECT_RATIO.get(str(size), "1:1") if size else None
|
||||
)
|
||||
n: Final = merged.get("n")
|
||||
fields: Final = (("aspect_ratio", aspect_ratio), ("n", int(n) if n is not None else None))
|
||||
return {key: value for key, value in fields if value is not None} # mutable-ok: base class returns a dict
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth
|
||||
|
||||
resolved_base: Final = (
|
||||
XAIOAuthAuthenticator().get_api_base()
|
||||
if should_use_xai_oauth(litellm_params) and not XAIModelInfo.get_api_key(api_key)
|
||||
else (api_base or get_secret_str("XAI_API_BASE") or get_secret_str("XAI_OAUTH_API_BASE") or XAI_API_BASE)
|
||||
)
|
||||
base: Final = (resolved_base or XAI_API_BASE).rstrip("/")
|
||||
if base.endswith("/v1"):
|
||||
return f"{base}/images/generations"
|
||||
return f"{base}/v1/images/generations"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
model: str,
|
||||
messages: list[AllMessageValues], # mutable-ok: provider JSON body and base-class dict signature
|
||||
optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
from litellm.llms.xai.oauth import (
|
||||
XAIOAuthAuthenticator,
|
||||
XAIOAuthError,
|
||||
should_use_xai_oauth,
|
||||
)
|
||||
|
||||
dynamic_api_key: Final = XAIModelInfo.get_api_key(api_key)
|
||||
if should_use_xai_oauth(litellm_params) and not dynamic_api_key:
|
||||
try:
|
||||
headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().get_access_token()}"
|
||||
except XAIOAuthError as exc:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
llm_provider="xai",
|
||||
message=str(exc),
|
||||
) from exc
|
||||
else:
|
||||
if not dynamic_api_key:
|
||||
raise AuthenticationError(
|
||||
model=model,
|
||||
llm_provider="xai",
|
||||
message=(
|
||||
"Missing xAI credentials for image generation. "
|
||||
"Pass api_key / XAI_API_KEY, or set use_xai_oauth=True."
|
||||
),
|
||||
)
|
||||
headers["Authorization"] = f"Bearer {dynamic_api_key}"
|
||||
|
||||
if "content-type" not in headers and "Content-Type" not in headers:
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def transform_image_generation_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
n: Final = optional_params.get("n")
|
||||
requested_ratio: Final = optional_params.get("aspect_ratio")
|
||||
aspect: Final = {"aspect_ratio": requested_ratio} if requested_ratio is not None else None # fmt: skip # mutable-ok: provider JSON body
|
||||
count: Final = {"n": int(n)} if n is not None else None # mutable-ok: provider JSON body
|
||||
return { # mutable-ok: provider JSON body and base-class dict signature
|
||||
"model": XAIModelInfo.get_base_model(model) or model,
|
||||
"prompt": prompt,
|
||||
**(aspect or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
**(count or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
}
|
||||
|
||||
def transform_image_generation_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ImageResponse,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
request_data: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
optional_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ImageResponse:
|
||||
try:
|
||||
response_data: Final = raw_response.json()
|
||||
except Exception:
|
||||
raise self.get_error_class(
|
||||
error_message=raw_response.text,
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
logging_obj.post_call(
|
||||
input=request_data.get("prompt", ""),
|
||||
api_key=api_key,
|
||||
additional_args={"complete_input_dict": request_data}, # mutable-ok: provider JSON body
|
||||
original_response=response_data,
|
||||
)
|
||||
|
||||
images: Final = tuple(
|
||||
ImageObject(
|
||||
url=item.get("url"),
|
||||
b64_json=item.get("b64_json") or item.get("b64"),
|
||||
)
|
||||
for item in response_data.get("data") or ()
|
||||
if isinstance(item, dict)
|
||||
)
|
||||
if not images:
|
||||
raise self.get_error_class(
|
||||
error_message=f"xAI image generation returned no image data: {response_data}",
|
||||
status_code=raw_response.status_code,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
model_response.data = list(images) # mutable-ok: provider JSON body and base-class dict signature
|
||||
return model_response
|
||||
3
litellm/llms/xai/videos/__init__.py
Normal file
3
litellm/llms/xai/videos/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import XAIVideoConfig
|
||||
|
||||
__all__ = ["XAIVideoConfig"] # mutable-ok: provider JSON body and base-class dict signature
|
||||
414
litellm/llms/xai/videos/transformation.py
Normal file
414
litellm/llms/xai/videos/transformation.py
Normal file
|
|
@ -0,0 +1,414 @@
|
|||
import time
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
import litellm
|
||||
from litellm.constants import XAI_API_BASE
|
||||
from litellm.exceptions import AuthenticationError
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get, encode_url_path_segment, safe_get
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.llms.xai.common_utils import XAIModelInfo
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
|
||||
from litellm.types.videos.utils import (
|
||||
encode_video_id_with_provider,
|
||||
extract_original_video_id,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
_DROPPED: Final = frozenset(("seconds", "size", "input_reference", "user", "extra_headers", "model"))
|
||||
_SIZE_TO_ASPECT_RATIO: Final = { # mutable-ok: provider JSON body and base-class dict signature
|
||||
"1024x1024": "1:1",
|
||||
"1792x1024": "16:9",
|
||||
"1024x1792": "9:16",
|
||||
"1280x720": "16:9",
|
||||
"720x1280": "9:16",
|
||||
"1920x1080": "16:9",
|
||||
"1080x1920": "9:16",
|
||||
}
|
||||
|
||||
|
||||
def _duration_from_seconds(seconds: object) -> int:
|
||||
try:
|
||||
return int(seconds) if seconds is not None else 6
|
||||
except (TypeError, ValueError):
|
||||
return 6
|
||||
|
||||
|
||||
_STATUS_MAP: Final = { # mutable-ok: provider JSON body and base-class dict signature
|
||||
"done": "completed",
|
||||
"completed": "completed",
|
||||
"succeeded": "completed",
|
||||
"failed": "failed",
|
||||
"expired": "failed",
|
||||
"pending": "processing",
|
||||
"processing": "processing",
|
||||
"in_progress": "processing",
|
||||
}
|
||||
|
||||
|
||||
class XAIVideoConfig(BaseVideoConfig):
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> list: # mutable-ok: provider JSON body and base-class dict signature
|
||||
return [ # mutable-ok: provider JSON body and base-class dict signature
|
||||
"model",
|
||||
"prompt",
|
||||
"input_reference",
|
||||
"seconds",
|
||||
"size",
|
||||
"user",
|
||||
"extra_headers",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
video_create_optional_params: VideoCreateOptionalRequestParams,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
raw: Final = video_create_optional_params
|
||||
incoming: Final = dict(raw) # mutable-ok: provider JSON body and base-class dict signature
|
||||
size: Final = incoming.get("size")
|
||||
seconds: Final = incoming.get("seconds")
|
||||
use_duration: Final = "seconds" in incoming and "duration" not in incoming
|
||||
duration_value: Final = _duration_from_seconds(seconds)
|
||||
duration: Final = {"duration": duration_value} if use_duration else None # mutable-ok: provider JSON body
|
||||
mapped_ratio: Final = incoming.get("aspect_ratio") or _SIZE_TO_ASPECT_RATIO.get(str(size), "16:9")
|
||||
use_ratio: Final = bool(size) and "aspect_ratio" not in incoming
|
||||
ratio: Final = {"aspect_ratio": mapped_ratio} if use_ratio else None # mutable-ok: provider JSON body
|
||||
image_ref: Final = incoming.get("image") or incoming.get("input_reference")
|
||||
use_image: Final = bool(incoming.get("input_reference")) and "image" not in incoming
|
||||
image: Final = {"image": image_ref} if use_image else None # mutable-ok: provider JSON body
|
||||
return { # mutable-ok: provider JSON body and base-class dict signature
|
||||
**{ # mutable-ok: provider JSON body and base-class dict signature
|
||||
key: value for key, value in incoming.items() if key not in _DROPPED
|
||||
},
|
||||
**(duration or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
**(ratio or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
**(image or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
}
|
||||
|
||||
def _resolve_api_base(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
litellm_params: GenericLiteLLMParams | dict | None, # mutable-ok: get_complete_url receives the base-class dict
|
||||
) -> str:
|
||||
from litellm.llms.xai.oauth import XAIOAuthAuthenticator, should_use_xai_oauth
|
||||
|
||||
params: Final = (
|
||||
litellm_params.model_dump()
|
||||
if isinstance(litellm_params, GenericLiteLLMParams)
|
||||
else (litellm_params or {}) # mutable-ok: provider JSON body and base-class dict signature
|
||||
)
|
||||
if should_use_xai_oauth(params) and not XAIModelInfo.get_api_key(api_key):
|
||||
return XAIOAuthAuthenticator().get_api_base().rstrip("/")
|
||||
|
||||
resolved: Final = (
|
||||
api_base
|
||||
or (params.get("api_base") if isinstance(params, dict) else None)
|
||||
or get_secret_str("XAI_API_BASE")
|
||||
or get_secret_str("XAI_OAUTH_API_BASE")
|
||||
or XAI_API_BASE
|
||||
)
|
||||
return str(resolved).rstrip("/")
|
||||
|
||||
def _v1_root(self, api_base: str) -> str:
|
||||
base: Final = api_base.rstrip("/")
|
||||
if base.endswith("/v1"):
|
||||
return base
|
||||
return f"{base}/v1"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
litellm_params: GenericLiteLLMParams | None = None,
|
||||
) -> dict: # mutable-ok: provider JSON body and base-class dict signature
|
||||
from litellm.llms.xai.oauth import (
|
||||
XAIOAuthAuthenticator,
|
||||
XAIOAuthError,
|
||||
should_use_xai_oauth,
|
||||
)
|
||||
|
||||
dumped: Final = litellm_params.model_dump() if litellm_params is not None else None
|
||||
params: Final = dumped or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
resolved_api_key: Final = api_key or (litellm_params.api_key if litellm_params else None)
|
||||
dynamic_api_key: Final = XAIModelInfo.get_api_key(resolved_api_key)
|
||||
if should_use_xai_oauth(params) and not dynamic_api_key:
|
||||
try:
|
||||
headers["Authorization"] = f"Bearer {XAIOAuthAuthenticator().get_access_token()}"
|
||||
except XAIOAuthError as exc:
|
||||
raise AuthenticationError(
|
||||
model=model or "xai-video",
|
||||
llm_provider="xai",
|
||||
message=str(exc),
|
||||
) from exc
|
||||
else:
|
||||
if not dynamic_api_key:
|
||||
raise AuthenticationError(
|
||||
model=model or "xai-video",
|
||||
llm_provider="xai",
|
||||
message=(
|
||||
"Missing xAI credentials for video generation. "
|
||||
"Pass api_key / XAI_API_KEY, or set use_xai_oauth=True."
|
||||
),
|
||||
)
|
||||
headers["Authorization"] = f"Bearer {dynamic_api_key}"
|
||||
|
||||
if "content-type" not in headers and "Content-Type" not in headers:
|
||||
headers["Content-Type"] = "application/json"
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
model: str,
|
||||
api_base: str | None,
|
||||
litellm_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> str:
|
||||
resolved: Final = self._resolve_api_base(
|
||||
api_base=api_base,
|
||||
api_key=litellm_params.get("api_key") if litellm_params else None,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
if not model:
|
||||
return self._v1_root(resolved)
|
||||
return f"{self._v1_root(resolved)}/videos/generations"
|
||||
|
||||
def transform_video_create_request(
|
||||
self,
|
||||
model: str,
|
||||
prompt: str,
|
||||
api_base: str,
|
||||
video_create_optional_request_params: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> tuple[dict, RequestFiles, str]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
copied: Final = { # mutable-ok: provider JSON body and base-class dict signature
|
||||
key: video_create_optional_request_params[key]
|
||||
for key in (
|
||||
"image",
|
||||
"images",
|
||||
"duration",
|
||||
"resolution_name",
|
||||
"aspect_ratio",
|
||||
"size",
|
||||
)
|
||||
if video_create_optional_request_params.get(key) is not None
|
||||
}
|
||||
prompt_body: Final = {"prompt": prompt} if prompt else None # mutable-ok: provider JSON body
|
||||
duration_body: Final = {"duration": 6} if "duration" not in copied else None # mutable-ok: provider JSON body
|
||||
return (
|
||||
{ # mutable-ok: provider JSON body and base-class dict signature
|
||||
"model": XAIModelInfo.get_base_model(model) or model,
|
||||
**(prompt_body or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
**copied,
|
||||
**(duration_body or {}), # mutable-ok: provider JSON body and base-class dict signature
|
||||
},
|
||||
[], # mutable-ok: provider JSON body and base-class dict signature
|
||||
api_base,
|
||||
)
|
||||
|
||||
def transform_video_create_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str | None = None,
|
||||
request_data: dict | None = None, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> VideoObject:
|
||||
response_data: Final = raw_response.json()
|
||||
request_id: Final = response_data.get("request_id") or response_data.get("id")
|
||||
if not request_id:
|
||||
raise ValueError(f"xAI video generation response missing request_id: {response_data}")
|
||||
|
||||
usage: Final = response_data.get("usage") or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
video_obj: Final = VideoObject(
|
||||
id=str(request_id),
|
||||
object="video",
|
||||
status="processing",
|
||||
created_at=int(time.time()),
|
||||
model=XAIModelInfo.get_base_model(model) or model,
|
||||
progress=0,
|
||||
)
|
||||
if custom_llm_provider:
|
||||
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, model)
|
||||
usage_body: Final = usage if isinstance(usage, dict) else None
|
||||
video_obj.usage = usage_body or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
return video_obj
|
||||
|
||||
def _video_resource_url(self, api_base: str, video_id: str) -> str:
|
||||
encoded_video_id: Final = encode_url_path_segment(
|
||||
extract_original_video_id(video_id),
|
||||
field_name="video_id",
|
||||
)
|
||||
return f"{self._v1_root(api_base)}/videos/{encoded_video_id}"
|
||||
|
||||
def transform_video_status_retrieve_request(
|
||||
self,
|
||||
video_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
return self._video_resource_url(
|
||||
api_base, video_id
|
||||
), {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def transform_video_status_retrieve_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> VideoObject:
|
||||
response_data: Final = raw_response.json()
|
||||
status_raw: Final = str(response_data.get("status") or "processing").lower()
|
||||
status: Final = _STATUS_MAP.get(status_raw, status_raw)
|
||||
video_body: Final = response_data.get("video")
|
||||
video_meta: Final = video_body or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
video_url: Final = video_meta.get("url") if isinstance(video_meta, dict) else None
|
||||
seconds: Final = (
|
||||
str(video_meta.get("duration"))
|
||||
if isinstance(video_meta, dict) and video_meta.get("duration") is not None
|
||||
else None
|
||||
)
|
||||
request_id: Final = (
|
||||
response_data.get("request_id")
|
||||
or response_data.get("id")
|
||||
or (video_url.split("/")[-1].replace(".mp4", "") if video_url else "unknown")
|
||||
)
|
||||
video_obj: Final = VideoObject(
|
||||
id=str(request_id),
|
||||
object="video",
|
||||
status=status,
|
||||
created_at=response_data.get("created_at") or int(time.time()),
|
||||
completed_at=int(time.time()) if status == "completed" else None,
|
||||
model=response_data.get("model"),
|
||||
progress=response_data.get("progress"),
|
||||
seconds=seconds,
|
||||
usage=response_data.get("usage")
|
||||
if isinstance(response_data.get("usage"), dict)
|
||||
else {}, # mutable-ok: provider JSON body and base-class dict signature
|
||||
)
|
||||
video_obj._hidden_params["video_url"] = video_url
|
||||
if custom_llm_provider and video_obj.id and video_obj.id != "unknown":
|
||||
video_obj.id = encode_video_id_with_provider(video_obj.id, custom_llm_provider, response_data.get("model"))
|
||||
return video_obj
|
||||
|
||||
def transform_video_content_request(
|
||||
self,
|
||||
video_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
variant: str | None = None,
|
||||
) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
return self._video_resource_url(
|
||||
api_base, video_id
|
||||
), {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
|
||||
def _video_cdn_url(self, raw_response: httpx.Response) -> str | None:
|
||||
content_type: Final = (raw_response.headers.get("content-type") or "").lower()
|
||||
if "application/json" not in content_type and raw_response.content[:1] != b"{":
|
||||
return None
|
||||
payload: Final = raw_response.json()
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
video_meta: Final = payload.get("video") or {} # mutable-ok: provider JSON body and base-class dict signature
|
||||
url: Final = video_meta.get("url") if isinstance(video_meta, dict) else None
|
||||
if isinstance(url, str) and url:
|
||||
return url
|
||||
raise ValueError(f"xAI video not ready for download (status={payload.get('status')}): {payload}")
|
||||
|
||||
def transform_video_content_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
) -> bytes:
|
||||
url: Final = self._video_cdn_url(raw_response)
|
||||
if url is None:
|
||||
return raw_response.content
|
||||
video_response: Final = safe_get(litellm.module_level_client, url)
|
||||
video_response.raise_for_status()
|
||||
return video_response.content
|
||||
|
||||
async def async_transform_video_content_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
) -> bytes:
|
||||
url: Final = self._video_cdn_url(raw_response)
|
||||
if url is None:
|
||||
return raw_response.content
|
||||
async_httpx_client: Final[AsyncHTTPHandler] = get_async_httpx_client(
|
||||
llm_provider=litellm.LlmProviders.XAI,
|
||||
)
|
||||
video_response: Final = await async_safe_get(async_httpx_client, url)
|
||||
video_response.raise_for_status()
|
||||
return video_response.content
|
||||
|
||||
def transform_video_remix_request(
|
||||
self,
|
||||
video_id: str,
|
||||
prompt: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
extra_body: dict[str, object] | None = None, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
raise NotImplementedError("Video remix is not supported by xAI Imagine API")
|
||||
|
||||
def transform_video_remix_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> VideoObject:
|
||||
raise NotImplementedError("Video remix is not supported by xAI Imagine API")
|
||||
|
||||
def transform_video_list_request(
|
||||
self,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
order: str | None = None,
|
||||
extra_query: dict[str, object] | None = None, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
raise NotImplementedError("Video listing is not supported by xAI Imagine API")
|
||||
|
||||
def transform_video_list_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
custom_llm_provider: str | None = None,
|
||||
) -> dict[str, str]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
raise NotImplementedError("Video listing is not supported by xAI Imagine API")
|
||||
|
||||
def transform_video_delete_request(
|
||||
self,
|
||||
video_id: str,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
headers: dict, # mutable-ok: provider JSON body and base-class dict signature
|
||||
) -> tuple[str, dict]: # mutable-ok: provider JSON body and base-class dict signature
|
||||
raise NotImplementedError("Video delete is not supported by xAI Imagine API")
|
||||
|
||||
def transform_video_delete_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
) -> VideoObject:
|
||||
raise NotImplementedError("Video delete is not supported by xAI Imagine API")
|
||||
|
|
@ -63428,7 +63428,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63444,7 +63445,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63461,7 +63463,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63478,7 +63481,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63495,7 +63499,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63511,7 +63516,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63527,7 +63533,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63545,6 +63552,9 @@
|
|||
"output_cost_per_second_480p": 0.05,
|
||||
"output_cost_per_second_720p": 0.07,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63563,6 +63573,9 @@
|
|||
"output_cost_per_second_480p": 0.08,
|
||||
"output_cost_per_second_720p": 0.14,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63581,6 +63594,9 @@
|
|||
"output_cost_per_second_480p": 0.08,
|
||||
"output_cost_per_second_720p": 0.14,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63599,6 +63615,9 @@
|
|||
"output_cost_per_second_480p": 0.08,
|
||||
"output_cost_per_second_720p": 0.14,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63642,7 +63661,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
|
|||
|
|
@ -2377,7 +2377,8 @@
|
|||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"image_generations": true,
|
||||
"image_edits": true,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
|
|
@ -2385,7 +2386,8 @@
|
|||
"rerank": false,
|
||||
"a2a": true,
|
||||
"interactions": true,
|
||||
"realtime": true
|
||||
"realtime": true,
|
||||
"video_generations": true
|
||||
}
|
||||
},
|
||||
"xinference": {
|
||||
|
|
|
|||
|
|
@ -1691,6 +1691,7 @@ _MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS: Final = (
|
|||
"/batches",
|
||||
"/skills",
|
||||
"/evals",
|
||||
"/videos",
|
||||
)
|
||||
_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS: Final = (
|
||||
"/files",
|
||||
|
|
|
|||
|
|
@ -1,11 +1,13 @@
|
|||
import asyncio
|
||||
import io
|
||||
from collections.abc import Sequence
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import Final, get_type_hints
|
||||
|
||||
import orjson
|
||||
from fastapi import APIRouter, Depends, File, HTTPException, Request, Response, UploadFile, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||
from fastapi.responses import ORJSONResponse
|
||||
from starlette.datastructures import FormData, UploadFile
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
|
|
@ -19,6 +21,7 @@ from litellm.proxy.common_request_processing import (
|
|||
resolve_litellm_call_id,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_is_form_content_type,
|
||||
coerce_numeric_form_fields,
|
||||
numeric_form_fields,
|
||||
)
|
||||
|
|
@ -39,6 +42,8 @@ IMAGE_EDIT_NUMERIC_FORM_FIELDS: Final = numeric_form_fields(get_type_hints(Image
|
|||
IMAGE_ARRAY_FIELD: Final = "image[]"
|
||||
MASK_ARRAY_FIELD: Final = "mask[]"
|
||||
BRACKETED_FILE_FIELDS: Final = frozenset({IMAGE_ARRAY_FIELD, MASK_ARRAY_FIELD})
|
||||
IMAGE_EDIT_FILE_FIELDS: Final = MappingProxyType({"image": IMAGE_ARRAY_FIELD, "mask": MASK_ARRAY_FIELD})
|
||||
IMAGE_REFERENCE_PREFIXES: Final = ("http://", "https://", "data:image/")
|
||||
|
||||
|
||||
async def uploadfile_to_bytesio(upload: UploadFile) -> io.BytesIO:
|
||||
|
|
@ -63,6 +68,58 @@ async def batch_to_bytesio(
|
|||
return [await uploadfile_to_bytesio(u) for u in uploads]
|
||||
|
||||
|
||||
async def _image_edit_part(value: object, field: str) -> io.BytesIO | str:
|
||||
if isinstance(value, UploadFile):
|
||||
return await uploadfile_to_bytesio(value)
|
||||
if isinstance(value, str) and value.startswith(IMAGE_REFERENCE_PREFIXES):
|
||||
return value
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail=f"'{field}' must be a multipart file upload, an http(s) URL, or a data:image URI.",
|
||||
)
|
||||
|
||||
|
||||
def _image_edit_values(form: FormData | None, body: Mapping[str, object], name: str) -> tuple[object, ...]:
|
||||
if form is not None:
|
||||
return tuple(form.getlist(name))
|
||||
raw: Final = body.get(name)
|
||||
if raw is None:
|
||||
return ()
|
||||
return tuple(raw) if isinstance(raw, list) else (raw,)
|
||||
|
||||
|
||||
async def _image_edit_field(
|
||||
form: FormData | None, body: Mapping[str, object], field: str, alias: str
|
||||
) -> list[io.BytesIO | str] | str | None:
|
||||
values: Final = _image_edit_values(form, body, field)
|
||||
alias_values: Final = _image_edit_values(form, body, alias)
|
||||
if values and alias_values:
|
||||
raise HTTPException(status_code=422, detail=f"Cannot specify both '{field}' and '{alias}'")
|
||||
parts: Final = tuple([await _image_edit_part(value, field) for value in values or alias_values])
|
||||
if not parts:
|
||||
return None
|
||||
if len(parts) == 1 and isinstance(parts[0], str):
|
||||
return parts[0]
|
||||
return list(parts) # mutable-ok: provider image edit handlers take a list of image parts
|
||||
|
||||
|
||||
async def image_edit_assets(request: Request, body: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""
|
||||
Collect ``image`` / ``mask`` (or their ``[]`` aliases) from a multipart form or JSON body.
|
||||
|
||||
Each part is an uploaded file, an http(s) URL, or a ``data:image/`` URI; any other string is rejected so it can
|
||||
never be treated as a filesystem path downstream.
|
||||
"""
|
||||
form: Final = await request.form() if _is_form_content_type(request.headers.get("content-type", "")) else None
|
||||
return MappingProxyType(
|
||||
{
|
||||
field: parts
|
||||
for field, alias in IMAGE_EDIT_FILE_FIELDS.items()
|
||||
if (parts := await _image_edit_field(form, body, field, alias)) is not None
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/images/generations",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -247,10 +304,6 @@ async def image_edit_api(
|
|||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
image: list[UploadFile] | None = File(None),
|
||||
image_array: list[UploadFile] | None = File(None, alias=IMAGE_ARRAY_FIELD),
|
||||
mask: list[UploadFile] | None = File(None),
|
||||
mask_array: list[UploadFile] | None = File(None, alias=MASK_ARRAY_FIELD),
|
||||
model: str | None = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -266,20 +319,6 @@ async def image_edit_api(
|
|||
-F 'prompt=Create a studio ghibli image of this'
|
||||
```
|
||||
"""
|
||||
if image is not None and image_array is not None:
|
||||
raise HTTPException(status_code=422, detail="Cannot specify both 'image' and 'image[]'")
|
||||
if mask is not None and mask_array is not None:
|
||||
raise HTTPException(status_code=422, detail="Cannot specify both 'mask' and 'mask[]'")
|
||||
if image is None and image_array is not None:
|
||||
image = image_array
|
||||
if mask is None and mask_array is not None:
|
||||
mask = mask_array
|
||||
|
||||
# if image is None:
|
||||
# raise HTTPException(status_code=422, detail="Field required: image")
|
||||
# Note: Image is optional for some models (e.g., Bedrock Stability style-transfer)
|
||||
# The validation will be done at the model level if image is truly required
|
||||
|
||||
from litellm.proxy.proxy_server import (
|
||||
_read_request_body,
|
||||
general_settings,
|
||||
|
|
@ -298,27 +337,14 @@ async def image_edit_api(
|
|||
#########################################################
|
||||
# Read request body and convert UploadFiles to BytesIO
|
||||
#########################################################
|
||||
parsed_body: Final = coerce_numeric_form_fields(
|
||||
parsed_body=await _read_request_body(request=request),
|
||||
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
|
||||
)
|
||||
data: Final = {
|
||||
key: value
|
||||
for key, value in coerce_numeric_form_fields(
|
||||
parsed_body=await _read_request_body(request=request),
|
||||
numeric_fields=IMAGE_EDIT_NUMERIC_FORM_FIELDS,
|
||||
).items()
|
||||
if key not in BRACKETED_FILE_FIELDS
|
||||
**{key: value for key, value in parsed_body.items() if key not in BRACKETED_FILE_FIELDS},
|
||||
**await image_edit_assets(request, parsed_body),
|
||||
}
|
||||
image_files: Final = await batch_to_bytesio(image)
|
||||
mask_files: Final = await batch_to_bytesio(mask)
|
||||
if image_files:
|
||||
data["image"] = image_files
|
||||
if mask_files:
|
||||
data["mask"] = mask_files
|
||||
|
||||
for _field in ("image", "mask"):
|
||||
if _field in data and isinstance(data[_field], str):
|
||||
raise HTTPException(
|
||||
status_code=422,
|
||||
detail=f"'{_field}' must be provided as a multipart file upload, not a string.",
|
||||
)
|
||||
|
||||
# Ensure prompt exists in data (default to None for models that don't require it)
|
||||
if "prompt" not in data:
|
||||
|
|
|
|||
|
|
@ -16,10 +16,16 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.proxy.image_endpoints.endpoints import batch_to_bytesio
|
||||
from litellm.proxy.video_endpoints.ownership import (
|
||||
assert_user_can_access_video,
|
||||
filter_video_list_for_caller,
|
||||
record_video_owner,
|
||||
)
|
||||
from litellm.proxy.video_endpoints.utils import (
|
||||
encode_character_id_in_response,
|
||||
extract_model_from_target_model_names,
|
||||
get_custom_provider_from_data,
|
||||
resolve_video_request_model,
|
||||
video_reference_to_id,
|
||||
)
|
||||
from litellm.types.videos.utils import (
|
||||
|
|
@ -115,6 +121,7 @@ async def video_generation(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(generated, user_api_key_dict)
|
||||
return generated
|
||||
|
||||
|
||||
|
|
@ -202,7 +209,7 @@ async def video_list(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
return listed
|
||||
return await filter_video_list_for_caller(listed, user_api_key_dict)
|
||||
|
||||
|
||||
@router.get(
|
||||
|
|
@ -249,6 +256,8 @@ async def video_status(
|
|||
version,
|
||||
)
|
||||
|
||||
await assert_user_can_access_video(video_id, user_api_key_dict)
|
||||
|
||||
# Create data with video_id
|
||||
data: Final[dict[str, object]] = {"video_id": video_id}
|
||||
|
||||
|
|
@ -256,23 +265,25 @@ async def video_status(
|
|||
provider_from_id: Final = decoded.get("custom_llm_provider")
|
||||
model_id_from_decoded: Final = decoded.get("model_id")
|
||||
|
||||
custom_llm_provider: Final = (
|
||||
explicit_provider: Final = (
|
||||
get_custom_llm_provider_from_request_headers(request=request)
|
||||
or get_custom_llm_provider_from_request_query(request=request)
|
||||
or await get_custom_llm_provider_from_request_body(request=request)
|
||||
or provider_from_id
|
||||
or "openai"
|
||||
)
|
||||
|
||||
resolved_model: Final = resolve_video_request_model(
|
||||
model_id_from_decoded=model_id_from_decoded,
|
||||
query_model=request.query_params.get("model"),
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if resolved_model:
|
||||
data["model"] = resolved_model
|
||||
|
||||
custom_llm_provider: Final = explicit_provider or (None if resolved_model else "openai")
|
||||
if custom_llm_provider:
|
||||
data["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
# Resolve model_name from model_id if available
|
||||
# This allows the router to automatically inject litellm_params from the model config
|
||||
if model_id_from_decoded and llm_router:
|
||||
resolved_model: Final = llm_router.resolve_model_name_from_model_id(model_id_from_decoded)
|
||||
if resolved_model:
|
||||
data["model"] = resolved_model
|
||||
|
||||
# Process request using ProxyBaseLLMRequestProcessing
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
|
|
@ -350,6 +361,8 @@ async def video_content(
|
|||
version,
|
||||
)
|
||||
|
||||
await assert_user_can_access_video(video_id, user_api_key_dict)
|
||||
|
||||
# Create data with video_id
|
||||
data: Final[dict[str, object]] = {"video_id": video_id}
|
||||
|
||||
|
|
@ -366,12 +379,13 @@ async def video_content(
|
|||
if custom_llm_provider:
|
||||
data["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
# Resolve model_name from model_id if available
|
||||
# This allows the router to automatically inject litellm_params from the model config
|
||||
if model_id_from_decoded and llm_router:
|
||||
resolved_model: Final = llm_router.resolve_model_name_from_model_id(model_id_from_decoded)
|
||||
if resolved_model:
|
||||
data["model"] = resolved_model
|
||||
resolved_content_model: Final = resolve_video_request_model(
|
||||
model_id_from_decoded=model_id_from_decoded,
|
||||
query_model=request.query_params.get("model"),
|
||||
llm_router=llm_router,
|
||||
)
|
||||
if resolved_content_model:
|
||||
data["model"] = resolved_content_model
|
||||
# Process request using ProxyBaseLLMRequestProcessing
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
|
|
@ -458,6 +472,8 @@ async def video_remix(
|
|||
version,
|
||||
)
|
||||
|
||||
await assert_user_can_access_video(video_id, user_api_key_dict)
|
||||
|
||||
data: Final = await _read_request_body(request=request)
|
||||
data["video_id"] = video_id
|
||||
|
||||
|
|
@ -510,6 +526,7 @@ async def video_remix(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(remixed, user_api_key_dict)
|
||||
return remixed
|
||||
|
||||
|
||||
|
|
@ -776,6 +793,7 @@ async def video_edit(
|
|||
data["video_id"] = ""
|
||||
else:
|
||||
data["video_id"] = video_reference_to_id(uploaded_video)
|
||||
await assert_user_can_access_video(data["video_id"], user_api_key_dict)
|
||||
|
||||
decoded: Final = decode_video_id_with_provider(data["video_id"])
|
||||
provider_from_id: Final = decoded.get("custom_llm_provider")
|
||||
|
|
@ -823,6 +841,7 @@ async def video_edit(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(edited, user_api_key_dict)
|
||||
return edited
|
||||
|
||||
|
||||
|
|
@ -873,6 +892,7 @@ async def video_extension(
|
|||
|
||||
data: Final = await _read_request_body(request=request)
|
||||
data["video_id"] = video_reference_to_id(data.pop("video", None))
|
||||
await assert_user_can_access_video(data["video_id"], user_api_key_dict)
|
||||
|
||||
decoded: Final = decode_video_id_with_provider(data["video_id"])
|
||||
provider_from_id: Final = decoded.get("custom_llm_provider")
|
||||
|
|
@ -920,4 +940,5 @@ async def video_extension(
|
|||
version=version,
|
||||
)
|
||||
else:
|
||||
await record_video_owner(extended, user_api_key_dict)
|
||||
return extended
|
||||
|
|
|
|||
167
litellm/proxy/video_endpoints/ownership.py
Normal file
167
litellm/proxy/video_endpoints/ownership.py
Normal file
|
|
@ -0,0 +1,167 @@
|
|||
"""Per-caller ownership of videos created through the proxy.
|
||||
|
||||
Ownership rows live in ``LiteLLM_ManagedObjectTable`` with ``file_purpose="video"``
|
||||
(the same table and owner scopes container ownership uses). A row is keyed by the
|
||||
provider-native video id only, so re-wrapping an id with a different provider or
|
||||
model_id cannot dodge the lookup. Proxy admins bypass the check; for everyone else,
|
||||
a video without a row is treated as admin-only.
|
||||
|
||||
Without a connected database there is nowhere to record ownership: recording and
|
||||
checks are skipped and videos remain reachable by any caller allowed on the route.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from fastapi import HTTPException
|
||||
from typing_extensions import TypeIs # noqa: TID251 # narrows untyped wire payloads without a runtime conversion
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.resource_ownership import (
|
||||
get_primary_resource_owner_scope,
|
||||
get_resource_owner_scopes,
|
||||
is_proxy_admin,
|
||||
user_can_access_resource_owner,
|
||||
)
|
||||
from litellm.repositories.table_repositories import ManagedObjectRepository
|
||||
from litellm.types.videos.utils import extract_original_video_id
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.proxy.utils import PrismaClient
|
||||
|
||||
VIDEO_OBJECT_PURPOSE: Final = "video"
|
||||
|
||||
# Status is polled; a short-lived cache of known owners keeps polling off the DB.
|
||||
# Only positive answers are cached, so a missing row is always re-checked.
|
||||
_VIDEO_OWNER_CACHE: Final = InMemoryCache(max_size_in_memory=10000, default_ttl=60)
|
||||
|
||||
|
||||
def _video_model_object_id(video_id: str) -> str:
|
||||
return f"{VIDEO_OBJECT_PURPOSE}:{extract_original_video_id(video_id)}"
|
||||
|
||||
|
||||
def _video_id_of(item: object) -> str | None:
|
||||
match item:
|
||||
case {"id": str(video_id)} if video_id:
|
||||
return video_id
|
||||
case object(id=str(video_id)) if video_id:
|
||||
return video_id
|
||||
case _:
|
||||
return None
|
||||
|
||||
|
||||
def _is_object_mapping(value: object) -> TypeIs[Mapping[str, object]]: # guard-ok: provider JSON objects have str keys
|
||||
return isinstance(value, Mapping)
|
||||
|
||||
|
||||
def _is_object_sequence(value: object) -> TypeIs[Sequence[object]]: # guard-ok: a list is a Sequence of anything
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
def _get_prisma_client() -> "PrismaClient | None":
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
return prisma_client
|
||||
|
||||
|
||||
async def record_video_owner(response: object, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
"""Stamp the caller as owner of a video the provider just created.
|
||||
|
||||
Failures are logged rather than raised: the provider job already exists and is
|
||||
billed, so the caller still gets its id; the untracked video is admin-only.
|
||||
"""
|
||||
prisma_client: Final = _get_prisma_client()
|
||||
if prisma_client is None:
|
||||
return
|
||||
video_id: Final = _video_id_of(response)
|
||||
owner: Final = get_primary_resource_owner_scope(user_api_key_dict)
|
||||
if video_id is None or owner is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Skipping video ownership tracking: response has no id or caller has no identity scope"
|
||||
)
|
||||
return
|
||||
model_object_id: Final = _video_model_object_id(video_id)
|
||||
row: Final = { # mutable-ok: prisma rows are plain dicts
|
||||
"unified_object_id": video_id,
|
||||
"model_object_id": model_object_id,
|
||||
"file_object": json.dumps({"id": video_id, "object": "video"}), # mutable-ok: serialized immediately
|
||||
"file_purpose": VIDEO_OBJECT_PURPOSE,
|
||||
"created_by": owner,
|
||||
"updated_by": owner,
|
||||
}
|
||||
try:
|
||||
await ManagedObjectRepository(prisma_client).table.upsert(
|
||||
where={"model_object_id": model_object_id}, # mutable-ok: prisma filters are plain dicts
|
||||
data={"create": row, "update": {"updated_by": owner}}, # mutable-ok: prisma payloads are plain dicts
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Video ownership recording failed; video_id=%s is untracked and admin-only: %s", video_id, e
|
||||
)
|
||||
|
||||
|
||||
async def _get_video_owner(prisma_client: "PrismaClient", model_object_id: str) -> str | None:
|
||||
cached: Final = _VIDEO_OWNER_CACHE.get_cache(model_object_id)
|
||||
if isinstance(cached, str):
|
||||
return cached
|
||||
row: Final = await ManagedObjectRepository(prisma_client).table.find_first(
|
||||
where={ # mutable-ok: prisma filters are plain dicts
|
||||
"model_object_id": model_object_id,
|
||||
"file_purpose": VIDEO_OBJECT_PURPOSE,
|
||||
}
|
||||
)
|
||||
owner: Final = getattr(row, "created_by", None) if row is not None else None
|
||||
if not isinstance(owner, str):
|
||||
return None
|
||||
_VIDEO_OWNER_CACHE.set_cache(model_object_id, owner)
|
||||
return owner
|
||||
|
||||
|
||||
async def assert_user_can_access_video(video_id: str, user_api_key_dict: UserAPIKeyAuth) -> None:
|
||||
"""Raise 403 unless the caller owns ``video_id`` (or is a proxy admin)."""
|
||||
if not video_id or is_proxy_admin(user_api_key_dict):
|
||||
return
|
||||
prisma_client: Final = _get_prisma_client()
|
||||
if prisma_client is None:
|
||||
return
|
||||
owner: Final = await _get_video_owner(prisma_client, _video_model_object_id(video_id))
|
||||
if not user_can_access_resource_owner(owner, user_api_key_dict):
|
||||
raise HTTPException(status_code=403, detail="Forbidden")
|
||||
|
||||
|
||||
async def filter_video_list_for_caller(listed: object, user_api_key_dict: UserAPIKeyAuth) -> object:
|
||||
"""Drop videos the caller does not own from a provider list page.
|
||||
|
||||
Pagination cursors are left as the provider returned them so ``after`` keeps
|
||||
walking the provider's list even when a page has no videos owned by the caller.
|
||||
"""
|
||||
if is_proxy_admin(user_api_key_dict) or not _is_object_mapping(listed):
|
||||
return listed
|
||||
prisma_client: Final = _get_prisma_client()
|
||||
items: Final = listed.get("data")
|
||||
if prisma_client is None or not _is_object_sequence(items):
|
||||
return listed
|
||||
ids_by_item: Final = tuple((item, _video_id_of(item)) for item in items)
|
||||
candidate_ids: Final = tuple(
|
||||
_video_model_object_id(video_id) for _, video_id in ids_by_item if video_id is not None
|
||||
)
|
||||
owner_scopes: Final = get_resource_owner_scopes(user_api_key_dict)
|
||||
rows: Final = (
|
||||
await ManagedObjectRepository(prisma_client).table.find_many(
|
||||
where={ # mutable-ok: prisma filters are plain dicts
|
||||
"model_object_id": {"in": list(candidate_ids)}, # mutable-ok: prisma filters are plain dicts
|
||||
"file_purpose": VIDEO_OBJECT_PURPOSE,
|
||||
"created_by": {"in": owner_scopes}, # mutable-ok: prisma filters are plain dicts
|
||||
}
|
||||
)
|
||||
if candidate_ids and owner_scopes
|
||||
else ()
|
||||
)
|
||||
owned: Final = frozenset(row.model_object_id for row in rows)
|
||||
kept: Final = tuple(
|
||||
item for item, video_id in ids_by_item if video_id is not None and _video_model_object_id(video_id) in owned
|
||||
)
|
||||
return {**listed, "data": kept} # mutable-ok: provider list JSON body
|
||||
|
|
@ -1,16 +1,40 @@
|
|||
from typing import Any, Final
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Protocol
|
||||
|
||||
import orjson
|
||||
|
||||
from litellm.types.videos.utils import encode_character_id_with_provider
|
||||
|
||||
|
||||
def extract_model_from_target_model_names(target_model_names: Any) -> str | None:
|
||||
class VideoModelIdResolver(Protocol):
|
||||
def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None: ...
|
||||
|
||||
|
||||
def resolve_video_request_model(
|
||||
*,
|
||||
model_id_from_decoded: str | None,
|
||||
query_model: str | None,
|
||||
llm_router: VideoModelIdResolver | None,
|
||||
) -> str | None:
|
||||
if model_id_from_decoded:
|
||||
if llm_router is not None:
|
||||
resolved: Final = llm_router.resolve_model_name_from_model_id(model_id_from_decoded)
|
||||
if isinstance(resolved, str) and resolved:
|
||||
return resolved
|
||||
return model_id_from_decoded
|
||||
if isinstance(query_model, str) and query_model:
|
||||
return query_model
|
||||
return None
|
||||
|
||||
|
||||
def extract_model_from_target_model_names(target_model_names: object) -> str | None:
|
||||
if isinstance(target_model_names, str):
|
||||
target_model_names = [m.strip() for m in target_model_names.split(",") if m.strip()]
|
||||
elif not isinstance(target_model_names, list):
|
||||
return None
|
||||
return target_model_names[0] if target_model_names else None
|
||||
names: Final = tuple(m.strip() for m in target_model_names.split(",") if m.strip())
|
||||
return names[0] if names else None
|
||||
if isinstance(target_model_names, Sequence) and not isinstance(target_model_names, (str, bytes)):
|
||||
first: Final = target_model_names[0] if target_model_names else None
|
||||
return first if isinstance(first, str) else None
|
||||
return None
|
||||
|
||||
|
||||
def video_reference_to_id(video_ref: object) -> str:
|
||||
|
|
@ -25,9 +49,9 @@ def video_reference_to_id(video_ref: object) -> str:
|
|||
return parsed_ref.get("id", "") if isinstance(parsed_ref, dict) else video_ref
|
||||
|
||||
|
||||
def get_custom_provider_from_data(data: dict[str, Any]) -> str | None:
|
||||
def get_custom_provider_from_data(data: Mapping[str, object]) -> str | None:
|
||||
custom_llm_provider: Final = data.get("custom_llm_provider")
|
||||
if custom_llm_provider:
|
||||
if isinstance(custom_llm_provider, str) and custom_llm_provider:
|
||||
return custom_llm_provider
|
||||
|
||||
extra_body = data.get("extra_body")
|
||||
|
|
|
|||
|
|
@ -9603,6 +9603,12 @@ class ProviderConfigManager:
|
|||
return get_modelscope_image_generation_config(model)
|
||||
elif LlmProviders.EDENAI == provider:
|
||||
return litellm.EdenAIImageGenerationConfig()
|
||||
elif LlmProviders.XAI == provider:
|
||||
from litellm.llms.xai.image_generation import (
|
||||
get_xai_image_generation_config,
|
||||
)
|
||||
|
||||
return get_xai_image_generation_config(model)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -9634,6 +9640,10 @@ class ProviderConfigManager:
|
|||
from litellm.llms.fal_ai.videos.transformation import FalAIVideoConfig
|
||||
|
||||
return FalAIVideoConfig()
|
||||
elif LlmProviders.XAI == provider:
|
||||
from litellm.llms.xai.videos.transformation import XAIVideoConfig
|
||||
|
||||
return XAIVideoConfig()
|
||||
elif LlmProviders.HOSTED_VLLM == provider:
|
||||
from litellm.llms.hosted_vllm.videos import get_hosted_vllm_video_config
|
||||
|
||||
|
|
@ -9772,6 +9782,10 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return get_openrouter_image_edit_config(model)
|
||||
elif LlmProviders.XAI == provider:
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
return XAIImageEditConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -30,6 +30,25 @@ from litellm.videos.utils import VideoGenerationRequestUtils
|
|||
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
|
||||
|
||||
|
||||
def _provider_from_model(model: object) -> str | None:
|
||||
if not isinstance(model, str) or not model:
|
||||
return None
|
||||
try:
|
||||
_, provider, _, _ = get_llm_provider(model=model)
|
||||
except litellm.BadRequestError:
|
||||
return None
|
||||
return provider
|
||||
|
||||
|
||||
def _provider_for_video_id(video_id: str, custom_llm_provider: str | None, model: object) -> str:
|
||||
return (
|
||||
custom_llm_provider
|
||||
or decode_video_id_with_provider(video_id).get("custom_llm_provider")
|
||||
or _provider_from_model(model)
|
||||
or "openai"
|
||||
)
|
||||
|
||||
|
||||
##### Video Generation #######################
|
||||
@client
|
||||
async def avideo_generation(
|
||||
|
|
@ -317,10 +336,7 @@ def video_content(
|
|||
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
|
||||
_is_async: Final = kwargs.pop("async_call", False) is True
|
||||
|
||||
# Try to decode provider from video_id if not explicitly provided
|
||||
if custom_llm_provider is None:
|
||||
decoded: Final = decode_video_id_with_provider(video_id)
|
||||
custom_llm_provider = decoded.get("custom_llm_provider") or "openai"
|
||||
custom_llm_provider = _provider_for_video_id(video_id, custom_llm_provider, kwargs.get("model"))
|
||||
|
||||
# get llm provider logic
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
|
@ -412,10 +428,7 @@ async def avideo_content(
|
|||
loop: Final = asyncio.get_event_loop()
|
||||
kwargs["async_call"] = True
|
||||
|
||||
# Try to decode provider from video_id if not explicitly provided
|
||||
if custom_llm_provider is None:
|
||||
decoded: Final = decode_video_id_with_provider(video_id)
|
||||
custom_llm_provider = decoded.get("custom_llm_provider") or "openai"
|
||||
custom_llm_provider = _provider_for_video_id(video_id, custom_llm_provider, kwargs.get("model"))
|
||||
|
||||
func: Final = partial(
|
||||
video_content,
|
||||
|
|
@ -1019,10 +1032,7 @@ def video_status(
|
|||
response: Final = VideoObject(**mock_response)
|
||||
return response
|
||||
|
||||
# Try to decode provider from video_id if not explicitly provided
|
||||
if custom_llm_provider is None:
|
||||
decoded: Final = decode_video_id_with_provider(video_id)
|
||||
custom_llm_provider = decoded.get("custom_llm_provider") or "openai"
|
||||
custom_llm_provider = _provider_for_video_id(video_id, custom_llm_provider, kwargs.get("model"))
|
||||
|
||||
# get llm provider logic
|
||||
litellm_params: Final = GenericLiteLLMParams(**kwargs)
|
||||
|
|
|
|||
|
|
@ -63428,7 +63428,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63444,7 +63445,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63461,7 +63463,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63478,7 +63481,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63495,7 +63499,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63511,7 +63516,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63527,7 +63533,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
@ -63545,6 +63552,9 @@
|
|||
"output_cost_per_second_480p": 0.05,
|
||||
"output_cost_per_second_720p": 0.07,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63563,6 +63573,9 @@
|
|||
"output_cost_per_second_480p": 0.08,
|
||||
"output_cost_per_second_720p": 0.14,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63581,6 +63594,9 @@
|
|||
"output_cost_per_second_480p": 0.08,
|
||||
"output_cost_per_second_720p": 0.14,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63599,6 +63615,9 @@
|
|||
"output_cost_per_second_480p": 0.08,
|
||||
"output_cost_per_second_720p": 0.14,
|
||||
"source": "https://docs.x.ai/docs/models/grok-imagine-video-1.5",
|
||||
"supported_endpoints": [
|
||||
"/v1/videos"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
|
|
@ -63642,7 +63661,8 @@
|
|||
"mode": "image_generation",
|
||||
"source": "https://docs.x.ai/docs/models",
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations"
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
|
|
|
|||
|
|
@ -2723,7 +2723,8 @@
|
|||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"image_generations": true,
|
||||
"image_edits": true,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
|
|
@ -2731,7 +2732,8 @@
|
|||
"rerank": false,
|
||||
"a2a": true,
|
||||
"interactions": true,
|
||||
"realtime": true
|
||||
"realtime": true,
|
||||
"video_generations": true
|
||||
}
|
||||
},
|
||||
"xinference": {
|
||||
|
|
|
|||
|
|
@ -1022,6 +1022,25 @@ def test_get_model_from_request_includes_fine_tuning_target_model_query():
|
|||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_includes_video_query_model_for_plain_id():
|
||||
result = get_model_from_request(
|
||||
request_data={"video_id": "plain-xai-id"},
|
||||
route="/v1/videos/{video_id}",
|
||||
request_query_params={"model": "grok-imagine-video-1.5"},
|
||||
)
|
||||
assert result == "grok-imagine-video-1.5"
|
||||
|
||||
|
||||
def test_get_model_from_request_includes_video_query_model_on_content_route():
|
||||
result = get_model_from_request(
|
||||
request_data={"video_id": "plain-xai-id"},
|
||||
route="/v1/videos/{video_id}/content",
|
||||
request_query_params={"model": "restricted-xai-model"},
|
||||
request_headers={"x-litellm-model": "also-restricted"},
|
||||
)
|
||||
assert result == ["restricted-xai-model", "also-restricted"]
|
||||
|
||||
|
||||
def test_get_model_from_request_extracts_video_id_model():
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
|
||||
|
|
|
|||
|
|
@ -222,6 +222,54 @@ def test_image_edit_multipart_n_that_is_not_a_number_is_left_alone(monkeypatch):
|
|||
assert captured["n"] == "two"
|
||||
|
||||
|
||||
def test_image_edit_http_url_is_accepted(monkeypatch):
|
||||
captured: Dict[str, Any] = {}
|
||||
|
||||
response = _image_edit_client(monkeypatch, captured).post(
|
||||
"/v1/images/edits",
|
||||
data={
|
||||
"model": "grok-imagine-image",
|
||||
"prompt": "make it night",
|
||||
"image": "https://imgen.x.ai/source.jpeg",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["image"] == "https://imgen.x.ai/source.jpeg"
|
||||
|
||||
|
||||
def test_image_edit_data_uri_is_accepted(monkeypatch):
|
||||
captured: Dict[str, Any] = {}
|
||||
|
||||
response = _image_edit_client(monkeypatch, captured).post(
|
||||
"/v1/images/edits",
|
||||
data={
|
||||
"model": "grok-imagine-image",
|
||||
"prompt": "make it red",
|
||||
"image": "data:image/jpeg;base64,abc",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert captured["image"] == "data:image/jpeg;base64,abc"
|
||||
|
||||
|
||||
def test_image_edit_plain_string_image_is_rejected(monkeypatch):
|
||||
captured: Dict[str, Any] = {}
|
||||
|
||||
response = _image_edit_client(monkeypatch, captured).post(
|
||||
"/v1/images/edits",
|
||||
data={
|
||||
"model": "grok-imagine-image",
|
||||
"prompt": "make it red",
|
||||
"image": "not-a-url-or-file",
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
assert "multipart file" in response.json()["detail"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_model_the_router_cannot_serve_answers_an_openai_typed_error(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A bare HTTPException carries no type or param, so the tail used to ship the
|
||||
|
|
|
|||
|
|
@ -309,6 +309,40 @@ async def test_status__header_provider_beats_decoded_id(harness):
|
|||
assert data["model"] == "azure-sora"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_status__resolve_fail_keeps_decoded_model_id(harness):
|
||||
encoded = encode_video_id_with_provider(
|
||||
"9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
"xai",
|
||||
"grok-imagine-video-1.5",
|
||||
)
|
||||
|
||||
await call_status(harness, encoded)
|
||||
|
||||
harness.resolve_model.assert_called_once_with("grok-imagine-video-1.5")
|
||||
assert harness.processor_data() == {
|
||||
"video_id": encoded,
|
||||
"custom_llm_provider": "xai",
|
||||
"model": "grok-imagine-video-1.5",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_status__query_model_on_plain_id_leaves_provider_to_the_model(harness):
|
||||
# an "openai" default here would override the provider the router / SDK derive from the model
|
||||
await call_status(
|
||||
harness,
|
||||
"9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
query={"model": "xai/grok-imagine-video-1.5"},
|
||||
)
|
||||
|
||||
harness.resolve_model.assert_not_called()
|
||||
assert harness.processor_data() == {
|
||||
"video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
"model": "xai/grok-imagine-video-1.5",
|
||||
}
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# GET /v1/videos/{video_id}/content - video_content #
|
||||
# =========================================================================== #
|
||||
|
|
@ -365,6 +399,42 @@ async def test_content__model_encoded_id(harness):
|
|||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_content__query_model_on_plain_id(harness):
|
||||
harness.base_process.return_value = b"x"
|
||||
|
||||
await call_content(
|
||||
harness,
|
||||
"9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
query={"model": "grok-imagine-video-1.5"},
|
||||
)
|
||||
|
||||
harness.resolve_model.assert_not_called()
|
||||
assert harness.processor_data() == {
|
||||
"video_id": "9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
"model": "grok-imagine-video-1.5",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_content__resolve_fail_keeps_decoded_model_id(harness):
|
||||
harness.base_process.return_value = b"x"
|
||||
encoded = encode_video_id_with_provider(
|
||||
"9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
"xai",
|
||||
"grok-imagine-video-1.5",
|
||||
)
|
||||
|
||||
await call_content(harness, encoded)
|
||||
|
||||
harness.resolve_model.assert_called_once_with("grok-imagine-video-1.5")
|
||||
assert harness.processor_data() == {
|
||||
"video_id": encoded,
|
||||
"custom_llm_provider": "xai",
|
||||
"model": "grok-imagine-video-1.5",
|
||||
}
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# POST /v1/videos/edits - video_edit #
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
218
tests/test_litellm/proxy/video_endpoints/test_ownership.py
Normal file
218
tests/test_litellm/proxy/video_endpoints/test_ownership.py
Normal file
|
|
@ -0,0 +1,218 @@
|
|||
"""
|
||||
Ownership tests for the proxy video endpoints.
|
||||
|
||||
Requests go through the real FastAPI routes and the real ownership module. Only two
|
||||
I/O boundaries are replaced: the provider call (``base_process_llm_request``) and the
|
||||
Prisma ``litellm_managedobjecttable`` delegate, which is an in-memory table honoring
|
||||
the subset of the query API the ownership module uses.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth, hash_token
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.video_endpoints import endpoints
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
|
||||
ALICE = UserAPIKeyAuth(user_id="alice", api_key=hash_token("sk-alice"))
|
||||
BOB = UserAPIKeyAuth(user_id="bob", api_key=hash_token("sk-bob"))
|
||||
KEY_ONLY_A = UserAPIKeyAuth(api_key=hash_token("sk-service-a"))
|
||||
KEY_ONLY_B = UserAPIKeyAuth(api_key=hash_token("sk-service-b"))
|
||||
ADMIN = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
|
||||
def _matches(row: SimpleNamespace, where: dict) -> bool:
|
||||
for field, condition in where.items():
|
||||
value = getattr(row, field, None)
|
||||
if isinstance(condition, dict):
|
||||
if value not in condition["in"]:
|
||||
return False
|
||||
elif value != condition:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class InMemoryManagedObjectTable:
|
||||
def __init__(self, fail_writes: bool = False):
|
||||
self.rows: dict[str, SimpleNamespace] = {}
|
||||
self.fail_writes = fail_writes
|
||||
|
||||
async def upsert(self, where: dict, data: dict) -> SimpleNamespace:
|
||||
if self.fail_writes:
|
||||
raise RuntimeError("database unavailable")
|
||||
key = where["model_object_id"]
|
||||
if key in self.rows:
|
||||
self.rows[key] = SimpleNamespace(**{**vars(self.rows[key]), **data["update"]})
|
||||
else:
|
||||
self.rows[key] = SimpleNamespace(**data["create"])
|
||||
return self.rows[key]
|
||||
|
||||
async def find_first(self, where: dict) -> SimpleNamespace | None:
|
||||
return next((row for row in self.rows.values() if _matches(row, where)), None)
|
||||
|
||||
async def find_many(self, where: dict) -> list[SimpleNamespace]:
|
||||
return [row for row in self.rows.values() if _matches(row, where)]
|
||||
|
||||
|
||||
class Provider:
|
||||
"""The provider behind ``base_process_llm_request``: creates and serves videos."""
|
||||
|
||||
def __init__(self):
|
||||
self.calls: list[str] = []
|
||||
|
||||
async def respond(self, processor, *, route_type: str, **kwargs):
|
||||
self.calls.append(route_type)
|
||||
if route_type in ("avideo_generation", "avideo_remix", "avideo_edit", "avideo_extension"):
|
||||
return {"id": f"video_{uuid.uuid4().hex}", "object": "video", "status": "queued"}
|
||||
if route_type == "avideo_content":
|
||||
return b"mp4-bytes"
|
||||
if route_type == "avideo_list":
|
||||
return {"object": "list", "data": processor.data["listed"], "has_more": False}
|
||||
return {"id": processor.data["video_id"], "object": "video", "status": "completed"}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def provider(monkeypatch) -> Provider:
|
||||
fake = Provider()
|
||||
|
||||
async def base_process_llm_request(processor, **kwargs):
|
||||
return await fake.respond(processor, **kwargs)
|
||||
|
||||
monkeypatch.setattr(ProxyBaseLLMRequestProcessing, "base_process_llm_request", base_process_llm_request)
|
||||
return fake
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def table(monkeypatch) -> InMemoryManagedObjectTable:
|
||||
rows = InMemoryManagedObjectTable()
|
||||
monkeypatch.setattr(
|
||||
proxy_server, "prisma_client", SimpleNamespace(db=SimpleNamespace(litellm_managedobjecttable=rows))
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def _client(auth: UserAPIKeyAuth) -> TestClient:
|
||||
app = FastAPI()
|
||||
app.include_router(endpoints.router)
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: auth
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def _create_video(auth: UserAPIKeyAuth) -> str:
|
||||
response = _client(auth).post("/v1/videos", json={"model": "sora-2", "prompt": "a sunset"})
|
||||
assert response.status_code == 200
|
||||
return response.json()["id"]
|
||||
|
||||
|
||||
def _access(auth: UserAPIKeyAuth, video_id: str) -> dict[str, int]:
|
||||
client = _client(auth)
|
||||
return {
|
||||
"status": client.get(f"/v1/videos/{video_id}").status_code,
|
||||
"content": client.get(f"/v1/videos/{video_id}/content").status_code,
|
||||
"remix": client.post(f"/v1/videos/{video_id}/remix", json={"prompt": "again"}).status_code,
|
||||
"edit": client.post("/v1/videos/edits", json={"prompt": "brighter", "video": {"id": video_id}}).status_code,
|
||||
"extension": client.post(
|
||||
"/v1/videos/extensions", json={"prompt": "continue", "video": {"id": video_id}}
|
||||
).status_code,
|
||||
}
|
||||
|
||||
|
||||
ALL_OK = {"status": 200, "content": 200, "remix": 200, "edit": 200, "extension": 200}
|
||||
ALL_FORBIDDEN = {"status": 403, "content": 403, "remix": 403, "edit": 403, "extension": 403}
|
||||
|
||||
|
||||
def test_owner_can_use_every_video_route(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert table.rows[f"video:{video_id}"].created_by == "alice"
|
||||
assert _access(ALICE, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_other_user_is_forbidden_and_the_provider_is_never_called(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
provider.calls.clear()
|
||||
|
||||
assert _access(BOB, video_id) == ALL_FORBIDDEN
|
||||
assert provider.calls == []
|
||||
|
||||
|
||||
def test_key_scoped_owner_blocks_a_different_key(provider, table):
|
||||
video_id = _create_video(KEY_ONLY_A)
|
||||
|
||||
assert _access(KEY_ONLY_B, video_id) == ALL_FORBIDDEN
|
||||
assert _access(KEY_ONLY_A, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_proxy_admin_can_access_any_video(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert _access(ADMIN, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_untracked_video_is_admin_only_when_a_database_is_connected(provider, table):
|
||||
video_id = f"video_{uuid.uuid4().hex}"
|
||||
|
||||
assert _access(ALICE, video_id) == ALL_FORBIDDEN
|
||||
assert _access(ADMIN, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_rewrapping_the_provider_id_does_not_bypass_the_owner_check(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
rewrapped = encode_video_id_with_provider(video_id, "azure", "attacker-deployment")
|
||||
|
||||
assert rewrapped != video_id
|
||||
assert _access(BOB, rewrapped) == ALL_FORBIDDEN
|
||||
|
||||
|
||||
def test_videos_derived_from_an_owned_video_belong_to_the_caller(provider, table):
|
||||
video_id = _create_video(ALICE)
|
||||
remixed = _client(ALICE).post(f"/v1/videos/{video_id}/remix", json={"prompt": "again"}).json()["id"]
|
||||
|
||||
assert _access(ALICE, remixed)["status"] == 200
|
||||
assert _access(BOB, remixed)["status"] == 403
|
||||
|
||||
|
||||
def test_list_only_returns_the_callers_videos(provider, table):
|
||||
alice_video = _create_video(ALICE)
|
||||
bob_video = _create_video(BOB)
|
||||
listed = [{"id": alice_video, "object": "video"}, {"id": bob_video, "object": "video"}]
|
||||
provider_listing = {"listed": listed}
|
||||
|
||||
def list_as(auth: UserAPIKeyAuth) -> list[str]:
|
||||
async def with_listing(processor, **kwargs):
|
||||
processor.data.update(provider_listing)
|
||||
return await provider.respond(processor, **kwargs)
|
||||
|
||||
with pytest.MonkeyPatch.context() as mp:
|
||||
mp.setattr(ProxyBaseLLMRequestProcessing, "base_process_llm_request", with_listing)
|
||||
response = _client(auth).get("/v1/videos")
|
||||
assert response.status_code == 200
|
||||
return [item["id"] for item in response.json()["data"]]
|
||||
|
||||
assert list_as(ALICE) == [alice_video]
|
||||
assert list_as(BOB) == [bob_video]
|
||||
assert list_as(ADMIN) == [alice_video, bob_video]
|
||||
|
||||
|
||||
def test_without_a_database_ownership_is_not_enforced(provider, monkeypatch):
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert _access(BOB, video_id) == ALL_OK
|
||||
|
||||
|
||||
def test_failed_ownership_write_still_returns_the_created_video(provider, table):
|
||||
table.fail_writes = True
|
||||
|
||||
video_id = _create_video(ALICE)
|
||||
|
||||
assert video_id.startswith("video_")
|
||||
assert table.rows == {}
|
||||
assert _access(ALICE, video_id)["status"] == 403
|
||||
|
|
@ -21,6 +21,7 @@ from litellm.proxy.video_endpoints.utils import (
|
|||
encode_character_id_in_response,
|
||||
extract_model_from_target_model_names,
|
||||
get_custom_provider_from_data,
|
||||
resolve_video_request_model,
|
||||
video_reference_to_id,
|
||||
)
|
||||
from litellm.types.videos.utils import (
|
||||
|
|
@ -28,6 +29,52 @@ from litellm.types.videos.utils import (
|
|||
encode_character_id_with_provider,
|
||||
)
|
||||
|
||||
# =========================================================================== #
|
||||
# resolve_video_request_model
|
||||
# =========================================================================== #
|
||||
|
||||
|
||||
class _Resolver:
|
||||
def __init__(self, mapping: dict[str, str | None]):
|
||||
self.mapping = mapping
|
||||
|
||||
def resolve_model_name_from_model_id(self, model_id: str | None) -> str | None:
|
||||
return self.mapping.get(model_id) if model_id else None
|
||||
|
||||
|
||||
def test_resolve_video_request_model__router_hit():
|
||||
assert (
|
||||
resolve_video_request_model(
|
||||
model_id_from_decoded="deployment-123",
|
||||
query_model="ignored",
|
||||
llm_router=_Resolver({"deployment-123": "azure-sora"}),
|
||||
)
|
||||
== "azure-sora"
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_video_request_model__keeps_decoded_id_when_router_misses():
|
||||
assert (
|
||||
resolve_video_request_model(
|
||||
model_id_from_decoded="grok-imagine-video-1.5",
|
||||
query_model=None,
|
||||
llm_router=_Resolver({}),
|
||||
)
|
||||
== "grok-imagine-video-1.5"
|
||||
)
|
||||
|
||||
|
||||
def test_resolve_video_request_model__query_model_on_plain_id():
|
||||
assert (
|
||||
resolve_video_request_model(
|
||||
model_id_from_decoded=None,
|
||||
query_model="grok-imagine-video-1.5",
|
||||
llm_router=None,
|
||||
)
|
||||
== "grok-imagine-video-1.5"
|
||||
)
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# extract_model_from_target_model_names
|
||||
# =========================================================================== #
|
||||
|
|
|
|||
186
tests/unit/llms/xai/test_xai_image_edit.py
Normal file
186
tests/unit/llms/xai/test_xai_image_edit.py
Normal file
|
|
@ -0,0 +1,186 @@
|
|||
import base64
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_xai_image_edit_config():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
config = ProviderConfigManager.get_provider_image_edit_config(
|
||||
model="grok-imagine-image",
|
||||
provider=LlmProviders.XAI,
|
||||
)
|
||||
assert isinstance(config, XAIImageEditConfig)
|
||||
|
||||
|
||||
def test_get_complete_url_default():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
url = XAIImageEditConfig().get_complete_url(
|
||||
model="grok-imagine-image",
|
||||
api_base=None,
|
||||
litellm_params={},
|
||||
)
|
||||
assert url.endswith("/v1/images/edits")
|
||||
|
||||
|
||||
def test_uses_json_not_multipart():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
assert XAIImageEditConfig().use_multipart_form_data() is False
|
||||
|
||||
|
||||
def test_map_size_to_aspect_ratio():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
mapped = XAIImageEditConfig().map_openai_params(
|
||||
image_edit_optional_params={"size": "1024x1792", "n": 1},
|
||||
model="grok-imagine-image",
|
||||
drop_params=True,
|
||||
)
|
||||
assert mapped["aspect_ratio"] == "9:16"
|
||||
assert mapped["n"] == 1
|
||||
assert "size" not in mapped
|
||||
|
||||
|
||||
def test_validate_environment_requires_credentials():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
with pytest.raises(Exception, match="Missing xAI credentials"):
|
||||
XAIImageEditConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-imagine-image",
|
||||
api_key=None,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
|
||||
def test_validate_environment_oauth_injects_bearer():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
with patch(
|
||||
"litellm.llms.xai.oauth.XAIOAuthAuthenticator.get_access_token",
|
||||
return_value="oauth-token",
|
||||
):
|
||||
headers = XAIImageEditConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-imagine-image",
|
||||
api_key=None,
|
||||
litellm_params={"use_xai_oauth": True},
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer oauth-token"
|
||||
|
||||
|
||||
def test_map_string_n_to_int():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
mapped = XAIImageEditConfig().map_openai_params(
|
||||
image_edit_optional_params={"n": "1"},
|
||||
model="grok-imagine-image",
|
||||
drop_params=True,
|
||||
)
|
||||
assert mapped["n"] == 1
|
||||
assert isinstance(mapped["n"], int)
|
||||
|
||||
|
||||
def test_transform_string_n_to_int():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
data, _ = XAIImageEditConfig().transform_image_edit_request(
|
||||
model="xai/grok-imagine-image",
|
||||
prompt="make the cube red",
|
||||
image="https://imgen.x.ai/source.jpeg",
|
||||
image_edit_optional_request_params={"n": "1"},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert data["n"] == 1
|
||||
assert isinstance(data["n"], int)
|
||||
|
||||
|
||||
def test_transform_bytes_to_data_uri_and_response():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
config = XAIImageEditConfig()
|
||||
data, files = config.transform_image_edit_request(
|
||||
model="xai/grok-imagine-image",
|
||||
prompt="make the cube red",
|
||||
image=b"\xff\xd8\xfffakejpeg",
|
||||
image_edit_optional_request_params={"aspect_ratio": "1:1"},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert files == []
|
||||
assert data["model"] == "grok-imagine-image"
|
||||
assert data["prompt"] == "make the cube red"
|
||||
assert data["aspect_ratio"] == "1:1"
|
||||
assert data["image"]["url"].startswith("data:image/jpeg;base64,")
|
||||
|
||||
raw = httpx.Response(
|
||||
200,
|
||||
json={"data": [{"url": "https://imgen.x.ai/edited.jpeg", "mime_type": "image/jpeg"}]},
|
||||
)
|
||||
response = config.transform_image_edit_response(
|
||||
model="grok-imagine-image",
|
||||
raw_response=raw,
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
assert response.data is not None
|
||||
assert response.data[0].url == "https://imgen.x.ai/edited.jpeg"
|
||||
|
||||
|
||||
def test_transform_duck_typed_reader_to_data_uri():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
class Reader:
|
||||
def read(self) -> bytes:
|
||||
return b"\x89PNG\r\n\x1a\nfakepng"
|
||||
|
||||
data, _ = XAIImageEditConfig().transform_image_edit_request(
|
||||
model="grok-imagine-image",
|
||||
prompt="make it night",
|
||||
image=Reader(),
|
||||
image_edit_optional_request_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert data["image"]["url"] == "data:image/png;base64," + base64.b64encode(b"\x89PNG\r\n\x1a\nfakepng").decode()
|
||||
|
||||
|
||||
def test_transform_http_url_passthrough():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
data, files = XAIImageEditConfig().transform_image_edit_request(
|
||||
model="grok-imagine-image",
|
||||
prompt="make it night",
|
||||
image="https://imgen.x.ai/source.jpeg",
|
||||
image_edit_optional_request_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert files == []
|
||||
assert data["image"] == {"url": "https://imgen.x.ai/source.jpeg"}
|
||||
|
||||
|
||||
def test_transform_multiple_images_uses_images_array():
|
||||
from litellm.llms.xai.image_edit.transformation import XAIImageEditConfig
|
||||
|
||||
data, _ = XAIImageEditConfig().transform_image_edit_request(
|
||||
model="grok-imagine-image",
|
||||
prompt="combine styles",
|
||||
image=["https://imgen.x.ai/a.jpeg", "https://imgen.x.ai/b.jpeg"],
|
||||
image_edit_optional_request_params={},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert "image" not in data
|
||||
assert data["images"] == [
|
||||
{"url": "https://imgen.x.ai/a.jpeg"},
|
||||
{"url": "https://imgen.x.ai/b.jpeg"},
|
||||
]
|
||||
129
tests/unit/llms/xai/test_xai_image_generation.py
Normal file
129
tests/unit/llms/xai/test_xai_image_generation.py
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.xai.image_generation.transformation import XAIImageGenerationConfig
|
||||
from litellm.types.utils import ImageResponse, LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_xai_image_config():
|
||||
config = ProviderConfigManager.get_provider_image_generation_config(
|
||||
model="grok-imagine-image",
|
||||
provider=LlmProviders.XAI,
|
||||
)
|
||||
assert isinstance(config, XAIImageGenerationConfig)
|
||||
|
||||
|
||||
def test_map_size_to_aspect_ratio():
|
||||
mapped = XAIImageGenerationConfig().map_openai_params(
|
||||
non_default_params={"size": "1024x1792"},
|
||||
optional_params={},
|
||||
model="grok-imagine-image",
|
||||
drop_params=True,
|
||||
)
|
||||
assert mapped["aspect_ratio"] == "9:16"
|
||||
assert "size" not in mapped
|
||||
|
||||
|
||||
def test_get_complete_url_default():
|
||||
url = XAIImageGenerationConfig().get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="grok-imagine-image",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
assert url.endswith("/v1/images/generations")
|
||||
|
||||
|
||||
def test_validate_environment_requires_credentials():
|
||||
with pytest.raises(Exception, match="Missing xAI credentials"):
|
||||
XAIImageGenerationConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-imagine-image",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key=None,
|
||||
)
|
||||
|
||||
|
||||
def test_validate_environment_oauth_injects_bearer():
|
||||
with patch(
|
||||
"litellm.llms.xai.oauth.XAIOAuthAuthenticator.get_access_token",
|
||||
return_value="oauth-token",
|
||||
):
|
||||
headers = XAIImageGenerationConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-imagine-image",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={"use_xai_oauth": True},
|
||||
api_key=None,
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer oauth-token"
|
||||
|
||||
|
||||
def test_map_string_n_to_int():
|
||||
mapped = XAIImageGenerationConfig().map_openai_params(
|
||||
non_default_params={"n": "1"},
|
||||
optional_params={},
|
||||
model="grok-imagine-image",
|
||||
drop_params=True,
|
||||
)
|
||||
assert mapped["n"] == 1
|
||||
assert isinstance(mapped["n"], int)
|
||||
|
||||
|
||||
def test_transform_string_n_to_int():
|
||||
request = XAIImageGenerationConfig().transform_image_generation_request(
|
||||
model="xai/grok-imagine-image",
|
||||
prompt="a red apple",
|
||||
optional_params={"n": "1"},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert request["n"] == 1
|
||||
assert isinstance(request["n"], int)
|
||||
|
||||
|
||||
def test_transform_request_and_response():
|
||||
config = XAIImageGenerationConfig()
|
||||
request = config.transform_image_generation_request(
|
||||
model="xai/grok-imagine-image",
|
||||
prompt="a red apple",
|
||||
optional_params={"aspect_ratio": "1:1"},
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
assert request == {
|
||||
"model": "grok-imagine-image",
|
||||
"prompt": "a red apple",
|
||||
"aspect_ratio": "1:1",
|
||||
}
|
||||
|
||||
raw = httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"data": [
|
||||
{
|
||||
"url": "https://imgen.x.ai/example.jpeg",
|
||||
"mime_type": "image/jpeg",
|
||||
}
|
||||
]
|
||||
},
|
||||
)
|
||||
response = config.transform_image_generation_response(
|
||||
model="grok-imagine-image",
|
||||
raw_response=raw,
|
||||
model_response=ImageResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data=request,
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
assert response.data is not None
|
||||
assert response.data[0].url == "https://imgen.x.ai/example.jpeg"
|
||||
330
tests/unit/llms/xai/test_xai_video_generation.py
Normal file
330
tests/unit/llms/xai/test_xai_video_generation.py
Normal file
|
|
@ -0,0 +1,330 @@
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError
|
||||
from litellm.llms.xai.videos.transformation import XAIVideoConfig
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
def test_provider_config_manager_returns_xai_video_config():
|
||||
config = ProviderConfigManager.get_provider_video_config(
|
||||
model="grok-imagine-video",
|
||||
provider=LlmProviders.XAI,
|
||||
)
|
||||
assert isinstance(config, XAIVideoConfig)
|
||||
|
||||
|
||||
def test_map_seconds_and_size():
|
||||
mapped = XAIVideoConfig().map_openai_params(
|
||||
video_create_optional_params={"seconds": "10", "size": "1280x720"},
|
||||
model="grok-imagine-video",
|
||||
drop_params=True,
|
||||
)
|
||||
assert mapped["duration"] == 10
|
||||
assert mapped["aspect_ratio"] == "16:9"
|
||||
|
||||
|
||||
def test_map_seconds_nine_is_not_clamped_to_six():
|
||||
mapped = XAIVideoConfig().map_openai_params(
|
||||
video_create_optional_params={"seconds": "9"},
|
||||
model="grok-imagine-video-1.5",
|
||||
drop_params=True,
|
||||
)
|
||||
assert mapped["duration"] == 9
|
||||
assert "seconds" not in mapped
|
||||
|
||||
|
||||
def test_get_complete_url_create_and_status_root():
|
||||
config = XAIVideoConfig()
|
||||
create_url = config.get_complete_url(
|
||||
model="grok-imagine-video",
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params={},
|
||||
)
|
||||
assert create_url == "https://api.x.ai/v1/videos/generations"
|
||||
|
||||
status_root = config.get_complete_url(
|
||||
model="",
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params={},
|
||||
)
|
||||
assert status_root == "https://api.x.ai/v1"
|
||||
|
||||
|
||||
def test_validate_environment_oauth_injects_bearer():
|
||||
with patch(
|
||||
"litellm.llms.xai.oauth.XAIOAuthAuthenticator.get_access_token",
|
||||
return_value="oauth-token",
|
||||
):
|
||||
headers = XAIVideoConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-imagine-video",
|
||||
api_key=None,
|
||||
litellm_params=GenericLiteLLMParams(use_xai_oauth=True),
|
||||
)
|
||||
assert headers["Authorization"] == "Bearer oauth-token"
|
||||
|
||||
|
||||
def test_transform_create_and_status_response():
|
||||
config = XAIVideoConfig()
|
||||
data, files, api_base = config.transform_video_create_request(
|
||||
model="xai/grok-imagine-video",
|
||||
prompt="a cat walking",
|
||||
api_base="https://api.x.ai/v1/videos/generations",
|
||||
video_create_optional_request_params={"duration": 6},
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert data["model"] == "grok-imagine-video"
|
||||
assert data["prompt"] == "a cat walking"
|
||||
assert data["duration"] == 6
|
||||
assert files == []
|
||||
|
||||
created = config.transform_video_create_response(
|
||||
model="grok-imagine-video",
|
||||
raw_response=httpx.Response(200, json={"request_id": "req-123"}),
|
||||
logging_obj=MagicMock(),
|
||||
custom_llm_provider="xai",
|
||||
)
|
||||
assert created.status == "processing"
|
||||
assert created.id
|
||||
|
||||
status_url, params = config.transform_video_status_retrieve_request(
|
||||
video_id=created.id,
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert status_url.endswith("/videos/req-123")
|
||||
assert params == {}
|
||||
|
||||
status = config.transform_video_status_retrieve_response(
|
||||
raw_response=httpx.Response(
|
||||
200,
|
||||
json={
|
||||
"status": "done",
|
||||
"request_id": "req-123",
|
||||
"video": {"url": "https://vidgen.x.ai/x.mp4", "duration": 6},
|
||||
"progress": 100,
|
||||
"model": "grok-imagine-video",
|
||||
},
|
||||
),
|
||||
logging_obj=MagicMock(),
|
||||
custom_llm_provider="xai",
|
||||
)
|
||||
assert status.status == "completed"
|
||||
assert status.seconds == "6"
|
||||
assert status._hidden_params.get("video_url") == "https://vidgen.x.ai/x.mp4"
|
||||
|
||||
|
||||
def test_validate_environment_requires_credentials():
|
||||
with pytest.raises(Exception, match="Missing xAI credentials"):
|
||||
XAIVideoConfig().validate_environment(
|
||||
headers={},
|
||||
model="grok-imagine-video",
|
||||
api_key=None,
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
)
|
||||
|
||||
|
||||
def test_content_request_is_get_status():
|
||||
url, params = XAIVideoConfig().transform_video_content_request(
|
||||
video_id="req-123",
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert url == "https://api.x.ai/v1/videos/req-123"
|
||||
assert params == {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("video_id", ["..", ""])
|
||||
@pytest.mark.parametrize(
|
||||
"transform_name",
|
||||
["transform_video_status_retrieve_request", "transform_video_content_request"],
|
||||
)
|
||||
def test_video_id_rejects_empty_and_dot_path_segments(video_id, transform_name):
|
||||
transform = getattr(XAIVideoConfig(), transform_name)
|
||||
with pytest.raises(ValueError, match="video_id"):
|
||||
transform(
|
||||
video_id=video_id,
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"transform_name",
|
||||
["transform_video_status_retrieve_request", "transform_video_content_request"],
|
||||
)
|
||||
def test_video_id_parent_path_is_one_encoded_segment(transform_name):
|
||||
transform = getattr(XAIVideoConfig(), transform_name)
|
||||
url, params = transform(
|
||||
video_id="../models",
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert url == "https://api.x.ai/v1/videos/..%2Fmodels"
|
||||
assert "/videos/../" not in url
|
||||
assert params == {}
|
||||
|
||||
|
||||
def test_video_id_is_percent_encoded_as_one_segment():
|
||||
url, params = XAIVideoConfig().transform_video_status_retrieve_request(
|
||||
video_id="req-123?x=1#frag",
|
||||
api_base="https://api.x.ai/v1",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
headers={},
|
||||
)
|
||||
assert url == "https://api.x.ai/v1/videos/req-123%3Fx%3D1%23frag"
|
||||
assert params == {}
|
||||
|
||||
|
||||
def test_content_response_fetches_cdn_via_shared_client():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}},
|
||||
)
|
||||
cdn = MagicMock()
|
||||
video_resp = MagicMock()
|
||||
video_resp.content = b"mp4-bytes"
|
||||
video_resp.raise_for_status.return_value = None
|
||||
with (
|
||||
patch("litellm.module_level_client", cdn),
|
||||
patch(
|
||||
"litellm.llms.xai.videos.transformation.safe_get",
|
||||
return_value=video_resp,
|
||||
) as safe_get,
|
||||
):
|
||||
body = XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock())
|
||||
assert body == b"mp4-bytes"
|
||||
safe_get.assert_called_once_with(cdn, "https://vidgen.x.ai/x.mp4")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_content_response_does_not_use_sync_client():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}},
|
||||
)
|
||||
async_client = MagicMock()
|
||||
video_resp = MagicMock()
|
||||
video_resp.content = b"async-mp4"
|
||||
video_resp.raise_for_status.return_value = None
|
||||
with (
|
||||
patch(
|
||||
"litellm.llms.xai.videos.transformation.get_async_httpx_client",
|
||||
return_value=async_client,
|
||||
),
|
||||
patch("litellm.module_level_client") as sync_client,
|
||||
patch(
|
||||
"litellm.llms.xai.videos.transformation.async_safe_get",
|
||||
new=AsyncMock(return_value=video_resp),
|
||||
) as async_safe_get,
|
||||
):
|
||||
body = await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock())
|
||||
assert body == b"async-mp4"
|
||||
assert sync_client.method_calls == []
|
||||
async_safe_get.assert_awaited_once_with(async_client, "https://vidgen.x.ai/x.mp4")
|
||||
|
||||
|
||||
def test_content_response_raises_when_status_has_no_url():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "pending"},
|
||||
)
|
||||
with pytest.raises(ValueError, match="not ready"):
|
||||
XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock())
|
||||
|
||||
|
||||
def test_content_response_returns_raw_bytes_when_not_json():
|
||||
raw = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "video/mp4"},
|
||||
content=b"already-mp4",
|
||||
)
|
||||
assert XAIVideoConfig().transform_video_content_response(raw, logging_obj=MagicMock()) == b"already-mp4"
|
||||
|
||||
|
||||
def test_content_response_rejects_internal_cdn_host():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "done", "video": {"url": "http://127.0.0.1/secret.mp4"}},
|
||||
)
|
||||
|
||||
def boom(*args, **kwargs):
|
||||
raise AssertionError("unsafe CDN fetch must not run")
|
||||
|
||||
with patch("litellm.module_level_client", MagicMock(get=boom)):
|
||||
with pytest.raises((SSRFError, ValueError)):
|
||||
XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock())
|
||||
|
||||
|
||||
def test_content_response_fetches_public_cdn_via_safe_get():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}},
|
||||
)
|
||||
video_resp = MagicMock()
|
||||
video_resp.content = b"safe-mp4"
|
||||
video_resp.raise_for_status.return_value = None
|
||||
with patch(
|
||||
"litellm.llms.xai.videos.transformation.safe_get",
|
||||
return_value=video_resp,
|
||||
create=True,
|
||||
) as safe_get:
|
||||
body = XAIVideoConfig().transform_video_content_response(status, logging_obj=MagicMock())
|
||||
assert body == b"safe-mp4"
|
||||
assert safe_get.call_count == 1
|
||||
assert safe_get.call_args.args[1] == "https://vidgen.x.ai/x.mp4"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_content_response_rejects_internal_cdn_host():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "done", "video": {"url": "http://127.0.0.1/secret.mp4"}},
|
||||
)
|
||||
|
||||
async def boom(*args, **kwargs):
|
||||
raise AssertionError("unsafe async CDN fetch must not run")
|
||||
|
||||
with patch(
|
||||
"litellm.llms.xai.videos.transformation.get_async_httpx_client",
|
||||
return_value=MagicMock(get=boom),
|
||||
):
|
||||
with pytest.raises((SSRFError, ValueError)):
|
||||
await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_content_response_fetches_public_cdn_via_async_safe_get():
|
||||
status = httpx.Response(
|
||||
200,
|
||||
headers={"content-type": "application/json"},
|
||||
json={"status": "done", "video": {"url": "https://vidgen.x.ai/x.mp4"}},
|
||||
)
|
||||
video_resp = MagicMock()
|
||||
video_resp.content = b"async-safe-mp4"
|
||||
video_resp.raise_for_status.return_value = None
|
||||
with patch(
|
||||
"litellm.llms.xai.videos.transformation.async_safe_get",
|
||||
new=AsyncMock(return_value=video_resp),
|
||||
create=True,
|
||||
) as async_safe_get:
|
||||
body = await XAIVideoConfig().async_transform_video_content_response(status, logging_obj=MagicMock())
|
||||
assert body == b"async-safe-mp4"
|
||||
async_safe_get.assert_awaited()
|
||||
assert async_safe_get.call_args.args[1] == "https://vidgen.x.ai/x.mp4"
|
||||
|
|
@ -39,6 +39,7 @@ import pytest
|
|||
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.videos.main import CharacterObject, VideoObject
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider
|
||||
|
|
@ -178,6 +179,47 @@ def test_video_content__plain_id_defaults_to_openai(seams):
|
|||
assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "openai"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def real_provider_resolution(seams):
|
||||
with patch.object(videos_main, "get_llm_provider", get_llm_provider):
|
||||
yield seams
|
||||
|
||||
|
||||
def test_video_content__plain_id_with_xai_model_uses_xai(real_provider_resolution):
|
||||
seams = real_provider_resolution
|
||||
videos_main.video_content(
|
||||
video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
model="xai/grok-imagine-video-1.5",
|
||||
)
|
||||
|
||||
assert seams.kwargs_of("video_content_handler")["custom_llm_provider"] == "xai"
|
||||
|
||||
|
||||
def test_video_status__plain_id_with_xai_model_uses_xai(real_provider_resolution):
|
||||
seams = real_provider_resolution
|
||||
videos_main.video_status(
|
||||
video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
model="xai/grok-imagine-video-1.5",
|
||||
)
|
||||
|
||||
assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "xai"
|
||||
assert seams.get_config.call_args.kwargs["provider"] == litellm.LlmProviders.XAI
|
||||
|
||||
|
||||
def test_video_status__plain_id_with_unresolvable_model_defaults_to_openai(real_provider_resolution):
|
||||
seams = real_provider_resolution
|
||||
videos_main.video_status(video_id="video_plain", model="not-a-known-model")
|
||||
|
||||
assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "openai"
|
||||
|
||||
|
||||
def test_video_status__decoded_provider_beats_model(real_provider_resolution):
|
||||
seams = real_provider_resolution
|
||||
videos_main.video_status(video_id=AZURE_VIDEO_ID, model="xai/grok-imagine-video-1.5")
|
||||
|
||||
assert seams.kwargs_of("video_status_handler")["custom_llm_provider"] == "azure"
|
||||
|
||||
|
||||
def test_video_remix__dispatch_and_provider_from_id(seams):
|
||||
result = videos_main.video_remix(video_id=AZURE_VIDEO_ID, prompt="new colors")
|
||||
|
||||
|
|
@ -377,6 +419,21 @@ async def test_avideo_content__pre_decodes_provider_before_delegating():
|
|||
assert sync.call_args.kwargs["custom_llm_provider"] == "azure"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_avideo_content__plain_id_with_xai_model_uses_xai():
|
||||
sentinel = b"mp4-bytes"
|
||||
with patch.object(
|
||||
videos_main, "video_content", MagicMock(return_value=sentinel)
|
||||
) as sync:
|
||||
result = await videos_main.avideo_content(
|
||||
video_id="9b444cea-aaaa-bbbb-cccc-dddddddddddd",
|
||||
model="xai/grok-imagine-video-1.5",
|
||||
)
|
||||
|
||||
assert result is sentinel
|
||||
assert sync.call_args.kwargs["custom_llm_provider"] == "xai"
|
||||
|
||||
|
||||
# =========================================================================== #
|
||||
# Credential passthrough - DB/YAML model-config credentials the router injects
|
||||
# via kwargs must reach the provider call for EVERY video handler, carried in
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue