import asyncio from collections.abc import AsyncIterator, 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 async def submit_stream( self, gen_factory: Callable[[], AsyncIterator[T]] ) -> AsyncIterator[T]: """Like submit(), but for async generators (token streams). The concurrency slot is held for the stream's whole lifetime, since the Ollama box stays busy until the last token.""" async with self._lock: if self._waiting >= self._max_queue_depth: raise QueueFullError("too many seekers right now") self._waiting += 1 acquired = False try: await self._semaphore.acquire() acquired = True self._waiting -= 1 async for item in gen_factory(): yield item finally: if acquired: self._semaphore.release() else: self._waiting -= 1