feat: Add ModelsLab video generation provider

- Implements ModelsLabVideoConfig following the RunwayML BaseVideoConfig pattern
- Supports text-to-video and image-to-video generation
- Key-in-body authentication (MODELSLAB_API_KEY env var)
- Async polling: processing status → poll /fetch/{id} until success
- Models: i2vgen-xl, stable-video-diffusion, wan-i2v-480p, animate-diff
- Adds MODELSLAB to LlmProviders enum and get_provider_video_config()
- 10 unit tests (all mocked, no real network calls)
This commit is contained in:
adhikjoshi 2026-03-01 05:43:33 +05:30
parent 98974771fd
commit 89b88ee831
8 changed files with 563 additions and 0 deletions

View file

View file

@ -0,0 +1,370 @@
"""
ModelsLab video generation transformation for LiteLLM.
NOTE: ModelsLab uses key-in-body authentication. The MODELSLAB_API_KEY
will appear in the request body (not headers). LiteLLM's logging pipeline
may log this — treat the key accordingly.
"""
import time
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from httpx._types import RequestFiles
import litellm
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.custom_httpx.http_handler import HTTPHandler, _get_httpx_client
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
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
MODELSLAB_VIDEO_BASE_URL = "https://modelslab.com/api/v6/video"
MODELSLAB_POLL_INTERVAL_SECONDS = 5
MODELSLAB_POLL_TIMEOUT_SECONDS = 300
class ModelsLabVideoConfig(BaseVideoConfig):
"""
Configuration class for ModelsLab video generation.
ModelsLab uses an async pattern:
1. POST /api/v6/video/text2video (or img2video) creates a job
2. Response is either {status: success, output: [...]} or {status: processing, request_id: ...}
3. When processing, poll POST /api/v6/video/fetch/{request_id} with {key} body until done
"""
def __init__(self):
super().__init__()
self._api_key: Optional[str] = None
def get_supported_openai_params(self, model: str) -> list:
return [
"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:
mapped_params: Dict[str, Any] = {}
# Parse size "WxH" → width, height
if "size" in video_create_optional_params:
size = video_create_optional_params["size"]
if isinstance(size, str) and "x" in size:
try:
w, h = size.split("x", 1)
mapped_params["width"] = int(w)
mapped_params["height"] = int(h)
except (ValueError, TypeError):
pass
# input_reference → init_image (for img2video)
if "input_reference" in video_create_optional_params:
mapped_params["init_image"] = video_create_optional_params["input_reference"]
# seconds → num_frames (approximate at ~8fps default)
if "seconds" in video_create_optional_params:
seconds = video_create_optional_params["seconds"]
try:
mapped_params["num_frames"] = max(8, int(float(str(seconds))) * 8)
except (ValueError, TypeError):
pass
return mapped_params
def validate_environment(
self,
headers: dict,
model: str,
api_key: Optional[str] = None,
litellm_params: Optional[GenericLiteLLMParams] = None,
) -> dict:
"""
Validate environment. ModelsLab uses key-in-body — only Content-Type goes in headers.
"""
if litellm_params and litellm_params.api_key:
api_key = api_key or litellm_params.api_key
api_key = (
api_key
or litellm.api_key
or get_secret_str("MODELSLAB_API_KEY")
)
if not api_key:
raise ValueError(
"ModelsLab API key is required. Set MODELSLAB_API_KEY environment variable "
"or pass api_key parameter."
)
self._api_key = api_key
# Key-in-body: DO NOT set Authorization header
headers["Content-Type"] = "application/json"
return headers
def get_complete_url(
self,
model: str,
api_base: Optional[str],
litellm_params: dict,
) -> str:
if api_base:
return api_base.rstrip("/")
# Use img2video if init_image is present in litellm_params
if litellm_params.get("init_image") or litellm_params.get("input_reference"):
return f"{MODELSLAB_VIDEO_BASE_URL}/img2video"
return f"{MODELSLAB_VIDEO_BASE_URL}/text2video"
def transform_video_create_request(
self,
model: str,
prompt: str,
api_base: str,
video_create_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[Dict, RequestFiles, str]:
request_data: Dict[str, Any] = {
"key": self._api_key,
"model_id": model,
"prompt": prompt,
}
request_data.update(video_create_optional_request_params)
files_list: List[Tuple[str, Any]] = []
return request_data, files_list, api_base
def transform_video_create_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str] = None,
request_data: Optional[Dict] = None,
) -> VideoObject:
response_data = raw_response.json()
status = response_data.get("status", "")
request_id = str(response_data.get("request_id", response_data.get("id", "")))
if status == "error":
raise BaseLLMException(
status_code=raw_response.status_code,
message=response_data.get("message", "ModelsLab video generation failed"),
headers=dict(raw_response.headers),
)
if status == "processing":
# Poll until done
response_data = self._poll_sync(request_id)
status = response_data.get("status", "")
if status == "success":
output = response_data.get("output", [])
output_url = output[0] if output else None
video_obj = VideoObject(
id=request_id,
object="video",
status="completed",
created_at=int(time.time()),
) # type: ignore[arg-type]
if output_url:
video_obj._hidden_params["output_url"] = output_url
if custom_llm_provider and video_obj.id:
video_obj.id = encode_video_id_with_provider(
video_obj.id, custom_llm_provider, model
)
return video_obj
raise BaseLLMException(
status_code=raw_response.status_code,
message=f"Unexpected ModelsLab video status: {status}",
headers=dict(raw_response.headers),
)
def _poll_sync(
self,
request_id: str,
timeout: int = MODELSLAB_POLL_TIMEOUT_SECONDS,
interval: int = MODELSLAB_POLL_INTERVAL_SECONDS,
) -> Dict:
"""Poll the ModelsLab fetch endpoint until status is success or error."""
fetch_url = f"{MODELSLAB_VIDEO_BASE_URL}/fetch/{request_id}"
body = {"key": self._api_key}
client: HTTPHandler = _get_httpx_client()
deadline = time.time() + timeout
while time.time() < deadline:
time.sleep(interval)
resp = client.post(fetch_url, json=body)
resp.raise_for_status()
data = resp.json()
status = data.get("status", "")
if status in ("success", "error"):
return data
# still processing — keep polling
raise BaseLLMException(
status_code=408,
message=f"ModelsLab video generation timed out after {timeout}s (request_id={request_id})",
headers={},
)
def transform_video_status_retrieve_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
original_id = extract_original_video_id(video_id)
url = f"{MODELSLAB_VIDEO_BASE_URL}/fetch/{original_id}"
body = {"key": self._api_key}
return url, body
def transform_video_status_retrieve_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str] = None,
) -> VideoObject:
response_data = raw_response.json()
status = response_data.get("status", "")
request_id = str(response_data.get("request_id", ""))
status_map = {
"processing": "in_progress",
"success": "completed",
"error": "failed",
}
mapped_status = status_map.get(status, "queued")
output = response_data.get("output", [])
output_url = output[0] if output else None
video_obj = VideoObject(
id=request_id,
object="video",
status=mapped_status,
created_at=int(time.time()),
) # type: ignore[arg-type]
if output_url:
video_obj._hidden_params["output_url"] = output_url
if custom_llm_provider and video_obj.id:
video_obj.id = encode_video_id_with_provider(
video_obj.id, custom_llm_provider, None
)
return video_obj
def transform_video_content_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
variant: Optional[str] = None,
) -> Tuple[str, Dict]:
raise NotImplementedError("Video content download not supported by ModelsLab API")
def transform_video_content_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> bytes:
raise NotImplementedError("Video content download not supported by ModelsLab API")
async def async_transform_video_content_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> bytes:
raise NotImplementedError("Video content download not supported by ModelsLab API")
def transform_video_remix_request(
self,
video_id: str,
prompt: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
raise NotImplementedError("Video remix not supported by ModelsLab API")
def transform_video_remix_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str] = None,
) -> VideoObject:
raise NotImplementedError("Video remix not supported by ModelsLab API")
def transform_video_list_request(
self,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
after: Optional[str] = None,
limit: Optional[int] = None,
order: Optional[str] = None,
extra_query: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
raise NotImplementedError("Video listing not supported by ModelsLab API")
def transform_video_list_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str] = None,
) -> Dict[str, str]:
raise NotImplementedError("Video listing not supported by ModelsLab API")
def transform_video_delete_request(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
raise NotImplementedError("Video deletion not supported by ModelsLab API")
def transform_video_delete_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> VideoObject:
raise NotImplementedError("Video deletion not supported by ModelsLab API")
def get_error_class(
self,
error_message: str,
status_code: int,
headers: Union[dict, httpx.Headers],
) -> BaseLLMException:
raise BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -3099,6 +3099,7 @@ class LlmProviders(str, Enum):
DEEPINFRA = "deepinfra"
PERPLEXITY = "perplexity"
MISTRAL = "mistral"
MODELSLAB = "modelslab"
MILVUS = "milvus"
GROQ = "groq"
A2A = "a2a"

View file

@ -8697,6 +8697,10 @@ class ProviderConfigManager:
from litellm.llms.runwayml.videos.transformation import RunwayMLVideoConfig
return RunwayMLVideoConfig()
elif LlmProviders.MODELSLAB == provider:
from litellm.llms.modelslab.videos.transformation import ModelsLabVideoConfig
return ModelsLabVideoConfig()
return None
@staticmethod

View file

@ -0,0 +1,188 @@
"""
Tests for ModelsLab video generation transformation.
All tests are mocked — no real network calls.
"""
from unittest.mock import MagicMock, Mock, patch
import httpx
import pytest
from litellm.llms.modelslab.videos.transformation import ModelsLabVideoConfig
from litellm.types.router import GenericLiteLLMParams
from litellm.types.videos.main import VideoObject
class TestModelsLabVideoTransformation:
"""Test ModelsLabVideoConfig transformation class."""
def setup_method(self):
self.config = ModelsLabVideoConfig()
self.config._api_key = "test-api-key"
self.mock_logging_obj = Mock()
# -------------------------------------------------------------------------
# validate_environment
# -------------------------------------------------------------------------
def test_validate_environment_no_auth_header(self):
"""Key-in-body auth: only Content-Type in headers, no Authorization."""
with patch(
"litellm.llms.modelslab.videos.transformation.get_secret_str",
return_value="test-key",
):
headers = self.config.validate_environment(
headers={}, model="i2vgen-xl"
)
assert "Content-Type" in headers
assert headers["Content-Type"] == "application/json"
assert "Authorization" not in headers
def test_validate_environment_raises_without_key(self):
"""Raises ValueError when no API key is available."""
with patch(
"litellm.llms.modelslab.videos.transformation.get_secret_str",
return_value=None,
):
with pytest.raises(ValueError, match="MODELSLAB_API_KEY"):
self.config.validate_environment(headers={}, model="i2vgen-xl")
# -------------------------------------------------------------------------
# get_supported_openai_params / map_openai_params
# -------------------------------------------------------------------------
def test_get_supported_openai_params(self):
params = self.config.get_supported_openai_params("i2vgen-xl")
assert "prompt" in params
assert "input_reference" in params
assert "size" in params
assert "seconds" in params
def test_map_openai_params_size_parsing(self):
"""'512x768' → width=512, height=768."""
result = self.config.map_openai_params(
video_create_optional_params={"size": "512x768"},
model="i2vgen-xl",
drop_params=False,
)
assert result["width"] == 512
assert result["height"] == 768
def test_map_openai_params_input_reference(self):
"""input_reference maps to init_image."""
result = self.config.map_openai_params(
video_create_optional_params={"input_reference": "https://example.com/img.jpg"},
model="stable-video-diffusion",
drop_params=False,
)
assert result["init_image"] == "https://example.com/img.jpg"
# -------------------------------------------------------------------------
# transform_video_create_request
# -------------------------------------------------------------------------
def test_transform_video_create_request_text2video(self):
"""text2video: key in body, model_id, prompt; URL is text2video endpoint."""
data, files, url = self.config.transform_video_create_request(
model="i2vgen-xl",
prompt="A cat playing with a ball",
api_base="https://modelslab.com/api/v6/video/text2video",
video_create_optional_request_params={"width": 512, "height": 512},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data["key"] == "test-api-key"
assert data["model_id"] == "i2vgen-xl"
assert data["prompt"] == "A cat playing with a ball"
assert data["width"] == 512
assert data["height"] == 512
assert files == []
assert "text2video" in url
def test_transform_video_create_request_img2video(self):
"""img2video: init_image in body; URL is img2video endpoint."""
data, files, url = self.config.transform_video_create_request(
model="stable-video-diffusion",
prompt="Camera panning right",
api_base="https://modelslab.com/api/v6/video/img2video",
video_create_optional_request_params={"init_image": "https://example.com/frame.jpg"},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert data["key"] == "test-api-key"
assert data["init_image"] == "https://example.com/frame.jpg"
assert "img2video" in url
# -------------------------------------------------------------------------
# transform_video_create_response
# -------------------------------------------------------------------------
def test_transform_video_create_response_success(self):
"""Immediate success response → completed VideoObject with output_url."""
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = {
"status": "success",
"request_id": "req_123",
"output": ["https://cdn.modelslab.com/output/video.mp4"],
}
result = self.config.transform_video_create_response(
model="i2vgen-xl",
raw_response=mock_response,
logging_obj=self.mock_logging_obj,
custom_llm_provider="modelslab",
)
assert isinstance(result, VideoObject)
assert result.status == "completed"
assert result._hidden_params.get("output_url") == "https://cdn.modelslab.com/output/video.mp4"
def test_transform_video_create_response_processing_polls(self):
"""Processing status → _poll_sync() is called → returns completed VideoObject."""
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = {
"status": "processing",
"request_id": "req_456",
"eta": 10,
}
poll_result = {
"status": "success",
"request_id": "req_456",
"output": ["https://cdn.modelslab.com/output/video2.mp4"],
}
with patch.object(self.config, "_poll_sync", return_value=poll_result) as mock_poll:
result = self.config.transform_video_create_response(
model="i2vgen-xl",
raw_response=mock_response,
logging_obj=self.mock_logging_obj,
)
mock_poll.assert_called_once_with("req_456")
assert result.status == "completed"
assert result._hidden_params.get("output_url") == "https://cdn.modelslab.com/output/video2.mp4"
# -------------------------------------------------------------------------
# transform_video_status_retrieve_request
# -------------------------------------------------------------------------
def test_transform_video_status_retrieve_request(self):
"""Fetch URL includes request_id; body has 'key'."""
from litellm.types.videos.utils import encode_video_id_with_provider
video_id = encode_video_id_with_provider("req_789", "modelslab", "i2vgen-xl")
url, body = self.config.transform_video_status_retrieve_request(
video_id=video_id,
api_base="https://modelslab.com/api/v6/video",
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "req_789" in url
assert "fetch" in url
assert body["key"] == "test-api-key"