ai-gateway / tests /test_queue.py
basyx's picture
Upload 60 files
eb808a5 verified
Raw
History Blame Contribute Delete
2.23 kB
"""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"