47 lines
1.7 KiB
Python
47 lines
1.7 KiB
Python
import asyncio
|
|
from collections.abc import Awaitable, Callable
|
|
from typing import TypeVar
|
|
|
|
T = TypeVar("T")
|
|
|
|
|
|
class QueueFullError(Exception):
|
|
pass
|
|
|
|
|
|
class LLMQueue:
|
|
"""Bounds concurrent Ollama calls and rejects work once too much is queued."""
|
|
|
|
def __init__(self, max_concurrency: int, max_queue_depth: int):
|
|
self._semaphore = asyncio.Semaphore(max_concurrency)
|
|
self._max_queue_depth = max_queue_depth
|
|
self._waiting = 0
|
|
self._lock = asyncio.Lock()
|
|
|
|
async def submit(self, coro_factory: Callable[[], Awaitable[T]]) -> T:
|
|
async with self._lock:
|
|
if self._waiting >= self._max_queue_depth:
|
|
raise QueueFullError("too many seekers right now")
|
|
self._waiting += 1
|
|
|
|
acquired = False
|
|
try:
|
|
# _waiting counts requests queued behind the concurrency limit,
|
|
# not requests currently running. It must drop the instant a slot
|
|
# is won (before coro_factory() runs), otherwise a long-running
|
|
# call keeps counting against the queue depth for its whole
|
|
# duration and wrongly rejects callers that should have been
|
|
# queued. The decrement below has no `await` before it, so it
|
|
# can't be split by cancellation from the acquire() that preceded
|
|
# it - either we own a slot and will decrement, or we don't and
|
|
# the `finally` below decrements instead.
|
|
await self._semaphore.acquire()
|
|
acquired = True
|
|
self._waiting -= 1
|
|
return await coro_factory()
|
|
finally:
|
|
if acquired:
|
|
self._semaphore.release()
|
|
else:
|
|
self._waiting -= 1
|