diff --git a/decart/queue/client.py b/decart/queue/client.py index ef14a77..035a6e4 100644 --- a/decart/queue/client.py +++ b/decart/queue/client.py @@ -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. @@ -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) @@ -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": @@ -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) diff --git a/tests/test_queue.py b/tests/test_queue.py index b45ed0c..13d8ee6 100644 --- a/tests/test_queue.py +++ b/tests/test_queue.py @@ -1,5 +1,6 @@ """Tests for the queue API.""" +import asyncio import pytest from unittest.mock import AsyncMock, patch, MagicMock from decart import ( @@ -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