Spaces:
Running on Zero
Running on Zero
| """Concurrency guarantees for the inference queue.""" | |
| from __future__ import annotations | |
| import asyncio | |
| import threading | |
| import time | |
| from contextvars import ContextVar | |
| import pytest | |
| from core.errors import GatewayError | |
| from core.queue import InferenceQueue | |
| async def test_queue_executes_only_one_operation_at_a_time() -> None: | |
| queue = InferenceQueue(capacity=4, timeout_seconds=5) | |
| active = 0 | |
| maximum_active = 0 | |
| lock = threading.Lock() | |
| def operation(value: int) -> int: | |
| nonlocal active, maximum_active | |
| with lock: | |
| active += 1 | |
| maximum_active = max(maximum_active, active) | |
| time.sleep(0.02) | |
| with lock: | |
| active -= 1 | |
| return value | |
| results = await asyncio.gather( | |
| *( | |
| queue.submit(str(index), "test", lambda index=index: operation(index)) | |
| for index in range(3) | |
| ) | |
| ) | |
| await queue.stop() | |
| assert results == [0, 1, 2] | |
| assert maximum_active == 1 | |
| async def test_queue_timeout_cancels_job_before_inference() -> None: | |
| queue = InferenceQueue(capacity=2, timeout_seconds=0.01) | |
| queued_job_ran = False | |
| async def submit_first() -> None: | |
| await queue.submit("first", "slow", lambda: time.sleep(0.04)) | |
| def queued_operation() -> None: | |
| nonlocal queued_job_ran | |
| queued_job_ran = True | |
| first = asyncio.create_task(submit_first()) | |
| await asyncio.sleep(0.005) | |
| with pytest.raises(GatewayError, match="timed out") as captured: | |
| await queue.submit("second", "queued", queued_operation) | |
| assert captured.value.status_code == 504 | |
| await first | |
| await asyncio.sleep(0) | |
| assert queued_job_ran is False | |
| assert queue.active is False | |
| await queue.stop() | |
| async def test_queue_preserves_request_context_for_worker() -> None: | |
| queue = InferenceQueue(capacity=1, timeout_seconds=5) | |
| request_context: ContextVar[str] = ContextVar("test_request", default="missing") | |
| token = request_context.set("request-identity") | |
| try: | |
| result = await queue.submit("request", "context", request_context.get) | |
| finally: | |
| request_context.reset(token) | |
| await queue.stop() | |
| assert result == "request-identity" | |