fix: async video polling for ModelsLab — use asyncio.sleep, not time.sleep

- Add BaseVideoConfig.async_transform_video_create_response() (optional
  override; defaults to sync version) so providers can implement non-blocking
  polling for long-running generation jobs.
- Update async_video_generation_handler to await the new async method instead
  of calling the sync transform_video_create_response directly.
- Add ModelsLabVideoConfig._poll_async() that uses asyncio.sleep(interval) so
  it never blocks the event loop during up-to-300s video jobs.
- Add ModelsLabVideoConfig.async_transform_video_create_response() that drives
  _poll_async; the sync path is unchanged.
- Fix _poll_sync: wrap httpx.HTTPStatusError in BaseLLMException so poll
  failures surface a structured error (status code + message) instead of a
  raw httpx exception.
- 4 new tests: async success, async poll-until-done, async poll-error message,
  sync _poll_sync HTTP error wrapping. All 17 tests pass.
This commit is contained in:
Adhik Joshi 2026-03-06 08:43:05 +05:30
parent 91ef877761
commit 10a46739c3
4 changed files with 241 additions and 6 deletions

View file

@ -111,6 +111,28 @@ class BaseVideoConfig(ABC):
) -> VideoObject:
pass
async def async_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:
"""
Async transform video creation response to VideoObject.
Optional override for providers that do async polling (e.g. ModelsLab).
Default falls back to the synchronous transform_video_create_response.
"""
return self.transform_video_create_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
request_data=request_data,
)
@abstractmethod
def transform_video_content_request(
self,

View file

@ -5402,7 +5402,7 @@ class BaseLLMHTTPHandler:
provider_config=video_generation_provider_config,
)
return video_generation_provider_config.transform_video_create_response(
return await video_generation_provider_config.async_transform_video_create_response(
model=model,
raw_response=response,
logging_obj=logging_obj,

View file

@ -5,6 +5,7 @@ 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 asyncio
import time
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
@ -14,7 +15,12 @@ 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.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
HTTPHandler,
_get_httpx_client,
get_async_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
@ -218,7 +224,7 @@ class ModelsLabVideoConfig(BaseVideoConfig):
timeout: int = MODELSLAB_POLL_TIMEOUT_SECONDS,
interval: int = MODELSLAB_POLL_INTERVAL_SECONDS,
) -> Dict:
"""Poll the ModelsLab fetch endpoint until status is success or error."""
"""Poll the ModelsLab fetch endpoint until status is success or error (sync path)."""
fetch_url = f"{MODELSLAB_VIDEO_BASE_URL}/fetch/{request_id}"
body = {"key": self._api_key}
client: HTTPHandler = _get_httpx_client()
@ -226,8 +232,15 @@ class ModelsLabVideoConfig(BaseVideoConfig):
while time.time() < deadline:
time.sleep(interval)
resp = client.post(fetch_url, json=body)
resp.raise_for_status()
try:
resp = client.post(fetch_url, json=body)
resp.raise_for_status()
except httpx.HTTPStatusError as e:
raise BaseLLMException(
status_code=e.response.status_code,
message=f"ModelsLab poll request failed: {e.response.text}",
headers=dict(e.response.headers),
) from e
data = resp.json()
status = data.get("status", "")
if status in ("success", "error"):
@ -240,6 +253,101 @@ class ModelsLabVideoConfig(BaseVideoConfig):
headers={},
)
async def _poll_async(
self,
request_id: str,
timeout: int = MODELSLAB_POLL_TIMEOUT_SECONDS,
interval: int = MODELSLAB_POLL_INTERVAL_SECONDS,
) -> Dict:
"""Poll the ModelsLab fetch endpoint using asyncio.sleep (async path)."""
fetch_url = f"{MODELSLAB_VIDEO_BASE_URL}/fetch/{request_id}"
body = {"key": self._api_key}
async_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.MODELSLAB,
)
deadline = time.time() + timeout
while time.time() < deadline:
await asyncio.sleep(interval)
try:
resp = await async_client.post(fetch_url, json=body)
resp.raise_for_status()
except httpx.HTTPStatusError as e:
raise BaseLLMException(
status_code=e.response.status_code,
message=f"ModelsLab poll request failed: {e.response.text}",
headers=dict(e.response.headers),
) from e
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={},
)
async def async_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:
"""
Async version of transform_video_create_response.
Uses asyncio.sleep in the poll loop instead of time.sleep, so it does
not block the event loop during long-running video generation jobs.
"""
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":
response_data = await self._poll_async(request_id)
status = response_data.get("status", "")
if status == "error":
raise BaseLLMException(
status_code=500,
message=response_data.get("message", "ModelsLab video generation failed"),
headers={},
)
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 transform_video_status_retrieve_request(
self,
video_id: str,

View file

@ -2,7 +2,7 @@
Tests for ModelsLab video generation transformation.
All tests are mocked — no real network calls.
"""
from unittest.mock import MagicMock, Mock, patch
from unittest.mock import AsyncMock, MagicMock, Mock, patch
import httpx
import pytest
@ -250,3 +250,108 @@ class TestModelsLabVideoTransformation:
)
assert "Insufficient credits" in str(exc_info.value)
def test_poll_sync_wraps_http_error_as_base_llm_exception(self):
"""_poll_sync wraps httpx.HTTPStatusError in BaseLLMException."""
from litellm.llms.base_llm.chat.transformation import BaseLLMException
mock_http_resp = Mock(spec=httpx.Response)
mock_http_resp.status_code = 403
mock_http_resp.text = "Forbidden"
mock_http_resp.headers = {}
mock_client = Mock()
mock_client.post.side_effect = httpx.HTTPStatusError(
"403 Forbidden",
request=Mock(),
response=mock_http_resp,
)
with patch(
"litellm.llms.modelslab.videos.transformation._get_httpx_client",
return_value=mock_client,
):
with pytest.raises(BaseLLMException) as exc_info:
self.config._poll_sync("req_403", timeout=10, interval=0)
assert exc_info.value.status_code == 403
assert "ModelsLab poll request failed" in str(exc_info.value)
@pytest.mark.asyncio
async def test_async_transform_video_create_response_success(self):
"""async_transform_video_create_response returns completed VideoObject on immediate success."""
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = {
"status": "success",
"request_id": "req_async_ok",
"output": ["https://cdn.modelslab.com/output/async_video.mp4"],
}
result = await self.config.async_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/async_video.mp4"
@pytest.mark.asyncio
async def test_async_transform_video_create_response_polls_with_asyncio_sleep(self):
"""async_transform_video_create_response uses _poll_async (asyncio.sleep, not time.sleep)."""
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = {
"status": "processing",
"request_id": "req_async_poll",
}
poll_result = {
"status": "success",
"request_id": "req_async_poll",
"output": ["https://cdn.modelslab.com/output/polled.mp4"],
}
with patch.object(self.config, "_poll_async", new=AsyncMock(return_value=poll_result)) as mock_poll:
result = await self.config.async_transform_video_create_response(
model="i2vgen-xl",
raw_response=mock_response,
logging_obj=self.mock_logging_obj,
)
mock_poll.assert_awaited_once_with("req_async_poll")
assert result.status == "completed"
assert result._hidden_params.get("output_url") == "https://cdn.modelslab.com/output/polled.mp4"
@pytest.mark.asyncio
async def test_async_transform_video_create_response_poll_error_surfaces_message(self):
"""async path: when _poll_async returns status=error, error message is raised."""
from litellm.llms.base_llm.chat.transformation import BaseLLMException
mock_response = Mock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.headers = {}
mock_response.json.return_value = {
"status": "processing",
"request_id": "req_async_err",
}
poll_result = {
"status": "error",
"message": "Model overloaded",
}
with patch.object(self.config, "_poll_async", new=AsyncMock(return_value=poll_result)):
with pytest.raises(BaseLLMException) as exc_info:
await self.config.async_transform_video_create_response(
model="i2vgen-xl",
raw_response=mock_response,
logging_obj=self.mock_logging_obj,
)
assert "Model overloaded" in str(exc_info.value)