diff --git a/backend/app/llm/__init__.py b/backend/app/llm/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/app/llm/client.py b/backend/app/llm/client.py new file mode 100644 index 0000000..f75f2fd --- /dev/null +++ b/backend/app/llm/client.py @@ -0,0 +1,18 @@ +import httpx + +from app.config import settings + + +class OllamaClient: + def __init__(self, base_url: str | None = None): + self._base_url = base_url or settings.ollama_base_url + + async def generate(self, model: str, prompt: str, system: str | None = None) -> str: + payload = {"model": model, "prompt": prompt, "stream": False} + if system is not None: + payload["system"] = system + + async with httpx.AsyncClient(base_url=self._base_url, timeout=120.0) as http_client: + response = await http_client.post("/api/generate", json=payload) + response.raise_for_status() + return response.json()["response"] diff --git a/backend/app/llm/queue.py b/backend/app/llm/queue.py new file mode 100644 index 0000000..ae51a97 --- /dev/null +++ b/backend/app/llm/queue.py @@ -0,0 +1,46 @@ +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 diff --git a/backend/tests/test_llm_queue.py b/backend/tests/test_llm_queue.py new file mode 100644 index 0000000..6c45767 --- /dev/null +++ b/backend/tests/test_llm_queue.py @@ -0,0 +1,43 @@ +import asyncio + +import pytest + +from app.llm.queue import LLMQueue, QueueFullError + + +@pytest.mark.asyncio +async def test_queue_runs_calls_up_to_concurrency_limit(): + queue = LLMQueue(max_concurrency=2, max_queue_depth=10) + concurrent_count = 0 + max_observed = 0 + + async def slow_call(): + nonlocal concurrent_count, max_observed + concurrent_count += 1 + max_observed = max(max_observed, concurrent_count) + await asyncio.sleep(0.05) + concurrent_count -= 1 + return "done" + + results = await asyncio.gather(*(queue.submit(slow_call) for _ in range(5))) + + assert results == ["done"] * 5 + assert max_observed == 2 + + +@pytest.mark.asyncio +async def test_queue_raises_when_depth_exceeded(): + queue = LLMQueue(max_concurrency=1, max_queue_depth=1) + + async def slow_call(): + await asyncio.sleep(0.1) + return "done" + + task1 = asyncio.create_task(queue.submit(slow_call)) + task2 = asyncio.create_task(queue.submit(slow_call)) + await asyncio.sleep(0.01) # let task1 start running, task2 start waiting + + with pytest.raises(QueueFullError): + await queue.submit(slow_call) + + await asyncio.gather(task1, task2)