mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
91ef877761
commit
10a46739c3
4 changed files with 241 additions and 6 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue