diff --git a/server/api/routes.py b/server/api/routes.py index 168cf3d..d19ed9c 100644 --- a/server/api/routes.py +++ b/server/api/routes.py @@ -13,7 +13,7 @@ GenerationRequestState, TokenEvent, ) -from server.executor.worker import Worker +from server.executor.worker import Worker, WorkerShuttingDown from server.metrics.logging import log_event from server.metrics.timers import now_ns, ns_to_ms, timed from server.model.determinism import make_generator @@ -224,6 +224,11 @@ def _submit_or_fail(request: Request, req: GenerateRequest) -> GenerationRequest status_code=503, detail="Server at capacity. Please try again later.", ) + except WorkerShuttingDown: + raise HTTPException( + status_code=503, + detail="Worker is shutting down. Please try again later.", + ) return state diff --git a/server/executor/worker.py b/server/executor/worker.py index 790d3ec..ae87ba1 100644 --- a/server/executor/worker.py +++ b/server/executor/worker.py @@ -17,6 +17,10 @@ logger = logging.getLogger(__name__) +class WorkerShuttingDown(RuntimeError): + """Raised when a request is submitted while the worker is shutting down.""" + + class Worker: """Orchestrates generation requests by bridging a queue and an inference engine. @@ -150,11 +154,13 @@ def submit(self, request_state: GenerationRequestState) -> None: """Enqueue a request for the engine to process. Raises: - RuntimeError: If the worker is shutting down. + WorkerShuttingDown: If the worker is shutting down. queue.Full: If the inbound queue is at capacity. """ if self._shutdown_event.is_set(): - raise RuntimeError("Cannot submit new request, worker is shutting down") + raise WorkerShuttingDown( + "Cannot submit new request, worker is shutting down" + ) request_state.enqueued_ns = now_ns() self._inbound.put_nowait(request_state) diff --git a/tests/api/__init__.py b/tests/api/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/api/test_routes.py b/tests/api/test_routes.py new file mode 100644 index 0000000..dc78a29 --- /dev/null +++ b/tests/api/test_routes.py @@ -0,0 +1,110 @@ +from queue import Queue +from typing import Callable, cast + +import pytest +from fastapi import HTTPException, Request + +from server.api.routes import _submit_or_fail +from server.api.schema import GenerateRequest +from server.executor.engine import EngineCallbacks, EngineControl +from server.executor.types import GenerationRequestState +from server.executor.worker import Worker +from tests.executor.worker_helpers import make_req + + +class _NoopEngine: + """Minimal ``InferenceEngine`` stub. + + The engine is never run in these tests; only ``cancel_inflight`` is touched, + by ``Worker.stop()``'s drain path. + """ + + def run( + self, + inbound: Queue[GenerationRequestState], + control: EngineControl, + callbacks: EngineCallbacks, + ) -> None: + raise AssertionError("engine.run must not be called by these tests") + + def cancel_inflight( + self, + message: str, + cancel_request: Callable[[GenerationRequestState, str], None], + ) -> None: + return None + + +def make_worker(max_queue_size: int = 16) -> Worker: + """Lightweight local factory; the engine is never started here.""" + return Worker(_NoopEngine(), max_queue_size=max_queue_size) + + +class _FakeState: + def __init__(self, worker: Worker, device: str) -> None: + self.worker = worker + self.device = device + + +class _FakeApp: + def __init__(self, worker: Worker, device: str) -> None: + self.state = _FakeState(worker, device) + + +class _FakeRequest: + """Stand-in for ``fastapi.Request`` exposing only ``app.state``.""" + + def __init__(self, worker: Worker, device: str = "cpu") -> None: + self.app = _FakeApp(worker, device) + + +def make_request(worker: Worker, device: str = "cpu") -> Request: + """Build a minimal stand-in for ``fastapi.Request`` for unit testing.""" + return cast(Request, _FakeRequest(worker, device)) + + +def make_generate_request(prompt: str = "hello") -> GenerateRequest: + return GenerateRequest( + prompt=prompt, + max_new_tokens=1, + temperature=1.0, + top_p=0.95, + seed=None, + ) + + +def test_submit_or_fail_maps_worker_shutting_down_to_503() -> None: + # stop() sets the shutdown event without ever starting the worker thread, + # so the next submit() raises WorkerShuttingDown for real. + worker = make_worker() + worker.stop() + request = make_request(worker) + + with pytest.raises(HTTPException) as excinfo: + _submit_or_fail(request, make_generate_request()) + + assert excinfo.value.status_code == 503 + assert excinfo.value.detail == "Worker is shutting down. Please try again later." + + +def test_submit_or_fail_maps_queue_full_to_503() -> None: + worker = make_worker(max_queue_size=1) + worker.submit(make_req("filler")) # single-slot queue is now full + request = make_request(worker) + + with pytest.raises(HTTPException) as excinfo: + _submit_or_fail(request, make_generate_request()) + + assert excinfo.value.status_code == 503 + assert excinfo.value.detail == "Server at capacity. Please try again later." + + +def test_submit_or_fail_happy_path_returns_state() -> None: + worker = make_worker() + request = make_request(worker) + + state = _submit_or_fail(request, make_generate_request(prompt="hello")) + + assert isinstance(state, GenerationRequestState) + assert state.prompt == "hello" + assert worker._inbound.qsize() == 1 diff --git a/tests/executor/test_worker.py b/tests/executor/test_worker.py index de1eca6..e1fa806 100644 --- a/tests/executor/test_worker.py +++ b/tests/executor/test_worker.py @@ -18,7 +18,7 @@ RequestStatus, TokenEvent, ) -from server.executor.worker import Worker +from server.executor.worker import Worker, WorkerShuttingDown from server.metrics.timers import NS_PER_S, now_ns from .worker_helpers import ( @@ -166,7 +166,7 @@ def test_submit_after_stop_is_rejected() -> None: worker.start() worker.stop() - with pytest.raises(RuntimeError, match="shutting down"): + with pytest.raises(WorkerShuttingDown, match="shutting down"): worker.submit(make_req("r0"))