Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 23 additions & 3 deletions decart/queue/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,19 @@
INITIAL_DELAY = 0.5 # seconds


async def _wait_or_cancel(delay: float, cancel_token: Optional[asyncio.Event]) -> None:
if cancel_token is None:
await asyncio.sleep(delay)
return
if cancel_token.is_set():
raise asyncio.CancelledError("Queue polling cancelled by user")
try:
await asyncio.wait_for(cancel_token.wait(), timeout=delay)
except asyncio.TimeoutError:
return
raise asyncio.CancelledError("Queue polling cancelled by user")


class QueueClient:
"""
Queue client for async jobs.
Expand Down Expand Up @@ -186,7 +199,12 @@ async def submit_and_poll(
QueueStatusError: If status check fails
QueueResultError: If result retrieval fails
"""
options = options.copy()
on_status_change: Optional[OnStatusChangeCallback] = options.pop("on_status_change", None)
cancel_token: Optional[asyncio.Event] = options.get("cancel_token")

if cancel_token and cancel_token.is_set():
raise asyncio.CancelledError("Queue polling cancelled by user")

# Submit the job
job = await self.submit(options)
Expand All @@ -196,13 +214,14 @@ async def submit_and_poll(
on_status_change(JobStatusResponse(job_id=job.job_id, status=job.status))

# Initial delay before polling
await asyncio.sleep(INITIAL_DELAY)
await _wait_or_cancel(INITIAL_DELAY, cancel_token)

last_status = job.status
# Poll until complete
while True:
status = await self.status(job.job_id)

if on_status_change:
if on_status_change and status.status != last_status:
on_status_change(status)

if status.status == "completed":
Expand All @@ -212,5 +231,6 @@ async def submit_and_poll(
if status.status == "failed":
return QueueJobResultFailed(status="failed", error="Job failed")

last_status = status.status
# Still pending or processing
await asyncio.sleep(POLLING_INTERVAL)
await _wait_or_cancel(POLLING_INTERVAL, cancel_token)
37 changes: 26 additions & 11 deletions tests/test_queue.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""Tests for the queue API."""

import asyncio
import pytest
from unittest.mock import AsyncMock, patch, MagicMock
from decart import (
Expand Down Expand Up @@ -203,23 +204,37 @@ def on_status_change(job):

mock_submit.return_value = MagicMock(job_id="job-123", status="pending")
mock_status.side_effect = [
MagicMock(job_id="job-123", status="pending"),
MagicMock(job_id="job-123", status="processing"),
MagicMock(job_id="job-123", status="completed"),
]
mock_content.return_value = b"fake video data"

await client.queue.submit_and_poll(
{
"model": models.video("lucy-clip"),
"prompt": "Add anime shading and crisp outlines",
"data": b"fake video data",
"on_status_change": on_status_change,
}
)
options = {
"model": models.video("lucy-clip"),
"prompt": "Add anime shading and crisp outlines",
"data": b"fake video data",
"on_status_change": on_status_change,
}
await client.queue.submit_and_poll(options)

assert status_changes == ["pending", "processing", "completed"]
assert "on_status_change" in options


@pytest.mark.asyncio
async def test_queue_submit_and_poll_can_be_cancelled() -> None:
client = DecartClient(api_key="test-key")
cancel_token = asyncio.Event()
cancel_token.set()

with patch("decart.queue.client.submit_job") as mock_submit:
with pytest.raises(asyncio.CancelledError):
await client.queue.submit_and_poll(
{"model": models.video("lucy-clip"), "cancel_token": cancel_token}
)

assert "pending" in status_changes
assert "processing" in status_changes
assert "completed" in status_changes
mock_submit.assert_not_called()


@pytest.mark.asyncio
Expand Down