diff --git a/README.md b/README.md index f83b0d6..e79ccc9 100644 --- a/README.md +++ b/README.md @@ -215,6 +215,29 @@ be served. A request can never introduce a model path, revision, or adapter. Send `routing.domain` to request a domain adapter; the router applies the promoted adapter with the largest measured quality gain for that base revision and task, or none at all. +## Failure behaviour + +| Condition | What the gateway does | +|---|---| +| Capacity saturated | `503` `overloaded` after the bounded admission wait, with `Retry-After: 1`. A streamed response holds its slot until the stream ends, so streaming cannot escape the concurrency bound. | +| GPU out of memory | `503` `engine_out_of_memory`, `Retry-After: 10`. The same request at the same size will fail again, so the message says to shorten the prompt or lower `max_tokens`. Counts toward the circuit. | +| Node lost or engine unreachable | `502` `backend_unavailable` for each of the first `ROUTER_ENGINE_FAILURE_THRESHOLD` consecutive failures; then the circuit opens. | +| Circuit open | `503` `engine_unavailable` immediately, without contacting the engine, with `Retry-After` set to the remaining cooldown. `/readyz` fails, so the orchestrator stops routing here. | +| Engine recovers | A healthy probe ends the cooldown early and lets exactly one trial request through. Only that request succeeding closes the circuit; if it fails, the cooldown restarts. | +| Shutdown | `/readyz` fails first, then admitted requests get `ROUTER_SHUTDOWN_GRACE_SECONDS` to finish before the engine client closes. Keep it under the pod's `terminationGracePeriodSeconds` (60). | + +An engine error body is inspected for an out-of-memory report and then discarded, never forwarded: +it can echo the prompt it rejected. `router_engine_circuit_open` reports 0 closed, 0.5 half-open, +1 open. + +A request that fails is not retried on another model. Every local model shares the one engine a +gateway faces, so a retry would meet the same failure; the caller is told when to come back. + +Cold start is measured, not assumed (see `router_model_load_seconds` below). Tiers that keep a +warm replica never pay it on the request path. The high-capability tier scales to zero, so its +first request after an idle period waits for a full model load: budget for it, or raise its +`min_replicas`. + ## Tenants What a caller may use is governance and lives in the catalog; the credential that proves which @@ -380,6 +403,9 @@ All settings use the `ROUTER_` prefix. | `ROUTER_BACKEND` | `mock` | `mock` or `vllm`. | | `ROUTER_VLLM_BASE_URL` | `http://127.0.0.1:8001` | vLLM OpenAI-compatible endpoint. | | `ROUTER_BACKEND_TIMEOUT_SECONDS` | `60` | Per-request engine timeout. | +| `ROUTER_ENGINE_FAILURE_THRESHOLD` | `5` | Consecutive engine failures before the circuit opens. | +| `ROUTER_ENGINE_COOLDOWN_SECONDS` | `30` | How long the circuit stays open before a trial request. | +| `ROUTER_SHUTDOWN_GRACE_SECONDS` | `20` | How long shutdown waits for admitted requests. | | `ROUTER_REGISTRY_PATH` | `config/registry.yaml` | Governed model catalog; built-in profiles are used if absent. | | `ROUTER_ROUTING_POLICY_VERSION` | `v1` | Invalidates router and response caches when changed. | | `ROUTER_CACHE_ENABLED` | `true` | Master switch for all cache tiers. | diff --git a/src/llm_router/admission.py b/src/llm_router/admission.py index e8ee735..f9f19e3 100644 --- a/src/llm_router/admission.py +++ b/src/llm_router/admission.py @@ -17,17 +17,51 @@ class AdmissionController: def __init__(self, max_concurrency: int, timeout_seconds: float) -> None: self._semaphore = asyncio.Semaphore(max_concurrency) self._timeout_seconds = timeout_seconds + self._in_flight = 0 + self._idle = asyncio.Event() + self._idle.set() + + @property + def in_flight(self) -> int: + return self._in_flight + + async def acquire(self) -> None: + """Take a slot, or refuse once the bounded wait has run out. + + A streamed response outlives the handler that admitted it, so it + holds its slot through `acquire` and `release` rather than a block + that would end, and free the slot, before generation began. + """ - @asynccontextmanager - async def slot(self) -> AsyncIterator[None]: try: await asyncio.wait_for(self._semaphore.acquire(), timeout=self._timeout_seconds) except TimeoutError as error: raise AdmissionRejectedError("inference capacity is saturated") from error + self._in_flight += 1 + self._idle.clear() + + def release(self) -> None: + self._in_flight -= 1 + if self._in_flight == 0: + self._idle.set() + self._semaphore.release() + + @asynccontextmanager + async def slot(self) -> AsyncIterator[None]: + await self.acquire() try: yield finally: - self._semaphore.release() + self.release() + + async def drain(self, timeout_seconds: float) -> bool: + """Wait for admitted work to finish; False if the grace period ran out.""" + + try: + await asyncio.wait_for(self._idle.wait(), timeout=timeout_seconds) + except TimeoutError: + return False + return True class SlidingWindowQuota: diff --git a/src/llm_router/app.py b/src/llm_router/app.py index 8c2cb59..bf8a93a 100644 --- a/src/llm_router/app.py +++ b/src/llm_router/app.py @@ -7,7 +7,7 @@ from contextlib import asynccontextmanager from dataclasses import dataclass from pathlib import Path -from typing import Annotated +from typing import Annotated, Any import httpx from fastapi import Depends, FastAPI, Header, HTTPException, Request, Response, status @@ -21,6 +21,7 @@ SlidingWindowQuota, ) from llm_router.backends import ( + BackendOutOfMemoryError, BackendResult, BackendUnavailableError, InferenceBackend, @@ -62,6 +63,7 @@ load_registry, strictest_privacy, ) +from llm_router.resilience import CircuitBreaker, EngineCircuitOpenError, ResilientBackend from llm_router.routing import NoEligibleModelError, Router, default_model_profiles from llm_router.tracing import RequestSpan, Tracing, build_tracer_provider @@ -144,7 +146,7 @@ def create_app( engine_client = ( httpx.AsyncClient() if backend is None and runtime_settings.backend == "vllm" else None ) - inference_backend: InferenceBackend = backend or ( + raw_backend: InferenceBackend = backend or ( VLLMBackend( base_url=runtime_settings.vllm_base_url, client=engine_client, @@ -153,6 +155,11 @@ def create_app( if engine_client is not None else MockInferenceBackend() ) + circuit = CircuitBreaker( + failure_threshold=runtime_settings.engine_failure_threshold, + cooldown_seconds=runtime_settings.engine_cooldown_seconds, + ) + inference_backend: InferenceBackend = ResilientBackend(raw_backend, circuit) telemetry = metrics if metrics is not None else Metrics() engine_telemetry = engine_stats or ( EngineStatsCollector(base_url=runtime_settings.vllm_base_url, client=engine_client) @@ -193,7 +200,11 @@ def create_app( async def lifespan(app: FastAPI) -> AsyncIterator[None]: app.state.ready = True yield + # Readiness fails first so no new traffic arrives, then admitted + # requests are given the grace period to finish before the engine + # client they depend on is closed underneath them. app.state.ready = False + await admission.drain(runtime_settings.shutdown_grace_seconds) if engine_client is not None: await engine_client.aclose() @@ -246,6 +257,26 @@ async def admission_handler(_: Request, error: AdmissionRejectedError) -> JSONRe content={"error": {"message": str(error), "type": "overloaded"}}, ) + @app.exception_handler(BackendOutOfMemoryError) + async def out_of_memory_handler(_: Request, error: BackendOutOfMemoryError) -> JSONResponse: + # Not a retry-as-is condition: the same request at the same size will + # fail again, so the message says what to change. + telemetry.record_rejection("engine_out_of_memory") + return JSONResponse( + status_code=503, + headers={"Retry-After": "10"}, + content={"error": {"message": str(error), "type": "engine_out_of_memory"}}, + ) + + @app.exception_handler(EngineCircuitOpenError) + async def circuit_handler(_: Request, error: EngineCircuitOpenError) -> JSONResponse: + telemetry.record_rejection("engine_unavailable") + return JSONResponse( + status_code=503, + headers={"Retry-After": str(max(1, round(error.retry_after_seconds)))}, + content={"error": {"message": str(error), "type": "engine_unavailable"}}, + ) + @app.exception_handler(BackendUnavailableError) async def backend_handler(_: Request, error: BackendUnavailableError) -> JSONResponse: telemetry.record_rejection("backend_unavailable") @@ -287,6 +318,7 @@ async def prometheus_metrics() -> Response: # Engine and accelerator state is pulled through on scrape so the # gateway stays the single scrape target for the whole serving path and # no background poller runs when nobody is collecting. + telemetry.record_circuit_state(circuit.state, engine=engine_label) if engine_telemetry is not None: stats = await engine_telemetry.sample() if stats is not None: @@ -513,6 +545,39 @@ def _chunk(completion_id: str, created: int, model_id: str, **choice: object) -> } return f"data: {json.dumps(document)}\n\n" + class _StreamLease: + """What a live stream holds until it ends: a slot, a gauge, a span.""" + + def __init__(self, span: RequestSpan) -> None: + self._span = span + self._closed = False + + def close(self) -> None: + if self._closed: + return + self._closed = True + telemetry.inflight_requests.dec() + admission.release() + self._span.end() + + class _LeasedStreamingResponse(StreamingResponse): + """Releases the lease even if the body iterator is never started. + + A generator that is never iterated never runs its own cleanup, which + happens when the client disconnects before the first chunk. The + response is always invoked, so it closes the lease as well. + """ + + def __init__(self, *args: Any, lease: _StreamLease, **kwargs: Any) -> None: + super().__init__(*args, **kwargs) + self._lease = lease + + async def __call__(self, scope: Any, receive: Any, send: Any) -> None: + try: + await super().__call__(scope, receive, send) + finally: + self._lease.close() + async def _stream_completion( payload: ChatCompletionRequest, prompt: str, @@ -522,6 +587,7 @@ async def _stream_completion( started: float, queue_seconds: float, span: RequestSpan, + lease: "_StreamLease", ) -> AsyncIterator[str]: completion_id = f"chatcmpl-{uuid.uuid4().hex}" created = int(time.time()) @@ -565,7 +631,7 @@ async def _stream_completion( span.fail(error) raise finally: - span.end() + lease.close() @app.post("/v1/chat/completions", response_model=None) async def chat_completions( @@ -679,36 +745,43 @@ async def _complete( telemetry.queued_requests.inc() try: - async with admission.slot(): - telemetry.queued_requests.dec() - queue_seconds = time.perf_counter() - started - # What this request actually waited becomes the next - # request's estimate for the same model. - load.observe_queue(decision.profile.id, queue_seconds * 1000) - telemetry.inflight_requests.inc() - try: - if payload.stream: - span.handed_off = True - return StreamingResponse( - _stream_completion( - payload, - prompt, - cache_key, - subject, - decision, - started, - queue_seconds, - span, - ), - media_type="text/event-stream", - headers=_route_headers(decision, "miss"), - ) - result = await inference_backend.generate(payload, decision) - finally: - telemetry.inflight_requests.dec() - except AdmissionRejectedError: + await admission.acquire() + finally: telemetry.queued_requests.dec() - raise + queue_seconds = time.perf_counter() - started + # What this request actually waited becomes the next request's + # estimate for the same model. + load.observe_queue(decision.profile.id, queue_seconds * 1000) + telemetry.inflight_requests.inc() + + if payload.stream: + # The stream holds its admission slot until it ends. Returning it + # from inside a slot block would free the slot before generation + # began, and streamed work would escape the concurrency bound. + lease = _StreamLease(span) + span.handed_off = True + return _LeasedStreamingResponse( + _stream_completion( + payload, + prompt, + cache_key, + subject, + decision, + started, + queue_seconds, + span, + lease, + ), + lease=lease, + media_type="text/event-stream", + headers=_route_headers(decision, "miss"), + ) + + try: + result = await inference_backend.generate(payload, decision) + finally: + telemetry.inflight_requests.dec() + admission.release() span.set_usage( prompt_tokens=result.prompt_tokens, completion_tokens=result.completion_tokens diff --git a/src/llm_router/backends.py b/src/llm_router/backends.py index 5a4a59e..48f7641 100644 --- a/src/llm_router/backends.py +++ b/src/llm_router/backends.py @@ -19,6 +19,28 @@ class BackendUnavailableError(RuntimeError): """Raised when an inference engine is unreachable or returns an error.""" +class BackendOutOfMemoryError(BackendUnavailableError): + """Raised when the engine reports it ran out of GPU memory for a request.""" + + +OUT_OF_MEMORY_MARKERS = ("out of memory", "outofmemoryerror", "cuda oom") + + +def engine_failure(status_code: int, body: str, context: str) -> BackendUnavailableError: + """Classify an engine error response without repeating what it said. + + The body is inspected for an out-of-memory report and then discarded: an + engine error can echo the prompt it rejected, so it is never forwarded. + """ + + if any(marker in body.lower() for marker in OUT_OF_MEMORY_MARKERS): + return BackendOutOfMemoryError( + f"inference engine ran out of GPU memory {context}; " + "retry with a shorter prompt or a smaller max_tokens" + ) + return BackendUnavailableError(f"inference engine returned {status_code} {context}") + + @dataclass(frozen=True) class BackendResult: text: str @@ -103,9 +125,8 @@ async def generate( raise BackendUnavailableError(f"inference engine unreachable: {error}") from error if response.status_code >= 400: - raise BackendUnavailableError( - f"inference engine returned {response.status_code} for " - f"{served_model_name(decision)}" + raise engine_failure( + response.status_code, response.text, f"for {served_model_name(decision)}" ) body = response.json() @@ -133,9 +154,8 @@ async def stream( timeout=self.request_timeout_seconds, ) as response: if response.status_code >= 400: - raise BackendUnavailableError( - f"inference engine returned {response.status_code} while streaming" - ) + body = (await response.aread()).decode(errors="replace") + raise engine_failure(response.status_code, body, "while streaming") async for line in response.aiter_lines(): delta = _parse_stream_line(line) if delta: diff --git a/src/llm_router/config.py b/src/llm_router/config.py index 00d49c4..6dd34e3 100644 --- a/src/llm_router/config.py +++ b/src/llm_router/config.py @@ -24,6 +24,13 @@ class Settings(BaseSettings): backend: Literal["mock", "vllm"] = "mock" vllm_base_url: str = "http://127.0.0.1:8001" backend_timeout_seconds: float = Field(default=60.0, gt=0) + # Consecutive engine failures before requests fail fast, and how long + # the gateway waits before letting a trial request through again. + engine_failure_threshold: int = Field(default=5, ge=1) + engine_cooldown_seconds: float = Field(default=30.0, gt=0) + # How long shutdown waits for admitted requests before closing the engine + # client; keep it under the pod's termination grace period. + shutdown_grace_seconds: float = Field(default=20.0, ge=0) registry_path: str = "config/registry.yaml" routing_policy_version: str = "v1" # Labelled prompts the task and complexity classifier is trained on at diff --git a/src/llm_router/observability.py b/src/llm_router/observability.py index 0457683..48904fb 100644 --- a/src/llm_router/observability.py +++ b/src/llm_router/observability.py @@ -129,6 +129,12 @@ def __init__(self, registry: CollectorRegistry | None = None) -> None: buckets=LATENCY_BUCKETS, registry=self.registry, ) + self.engine_circuit_open = Gauge( + "router_engine_circuit_open", + "Engine circuit state: 0 closed, 0.5 half-open, 1 open.", + ["engine"], + registry=self.registry, + ) self.engine_running_requests = Gauge( "router_engine_running_requests", "Requests the engine is currently decoding, which is its live batch size.", @@ -237,6 +243,10 @@ def record_structured_output(self, decision: RouteDecision, *, valid: bool) -> N abs(decision.profile.quality - observed) ) + def record_circuit_state(self, state: str, *, engine: str) -> None: + value = {"closed": 0.0, "half-open": 0.5, "open": 1.0}[state] + self.engine_circuit_open.labels(engine=engine).set(value) + def record_model_load(self, model: str, seconds: float) -> None: """Record a measured cold start, from first unready observation to ready.""" diff --git a/src/llm_router/resilience.py b/src/llm_router/resilience.py new file mode 100644 index 0000000..6bb4e37 --- /dev/null +++ b/src/llm_router/resilience.py @@ -0,0 +1,158 @@ +"""Defined behaviour when the engine fails (section 13). + +Section 13 requires the platform to stop routing to unhealthy replicas and to +define what happens on GPU out-of-memory and node loss, rather than leave it to +whatever the HTTP client does. + +**Node loss** shows up as connection failures. After a run of them the circuit +opens: requests fail immediately with retry guidance instead of each waiting +out a timeout against a node that is gone, and readiness fails so the +orchestrator stops sending traffic here. + +**Out of memory** is reported by the engine. The request is refused with its +own error type, because retrying the same request at the same size cannot +succeed, and the failure counts toward the circuit because an engine that has +exhausted its memory usually needs to reload. + +Recovery is driven by the readiness probe: once the engine answers its health +check, a single trial request is let through, and only its success closes the +circuit again. +""" + +import time +from collections.abc import AsyncIterator, Callable +from dataclasses import dataclass, field +from typing import Literal + +from llm_router.backends import BackendResult, BackendUnavailableError, InferenceBackend +from llm_router.models import ChatCompletionRequest, RouteDecision + +CircuitState = Literal["closed", "open", "half-open"] + + +class EngineCircuitOpenError(BackendUnavailableError): + """Raised without contacting the engine while its circuit is open.""" + + def __init__(self, retry_after_seconds: float) -> None: + super().__init__("inference engine is unavailable; the circuit is open") + self.retry_after_seconds = retry_after_seconds + + +@dataclass +class CircuitBreaker: + """Opens after consecutive failures and closes only on a successful trial.""" + + failure_threshold: int = 5 + cooldown_seconds: float = 30.0 + clock: Callable[[], float] = time.monotonic + _failures: int = 0 + _opened_at: float | None = None + _trial_in_flight: bool = False + _half_open: bool = field(default=False) + + @property + def state(self) -> CircuitState: + if self._opened_at is None: + return "closed" + if self._half_open or self.clock() - self._opened_at >= self.cooldown_seconds: + return "half-open" + return "open" + + @property + def retry_after_seconds(self) -> float: + if self._opened_at is None: + return 0.0 + return max(0.0, self.cooldown_seconds - (self.clock() - self._opened_at)) + + def allow(self) -> bool: + """Admit a request, letting exactly one trial through while half-open.""" + + state = self.state + if state == "closed": + return True + if state == "open" or self._trial_in_flight: + return False + self._trial_in_flight = True + return True + + def record_success(self) -> None: + self._failures = 0 + self._opened_at = None + self._half_open = False + self._trial_in_flight = False + + def record_failure(self) -> None: + self._failures += 1 + failed_trial = self._trial_in_flight + self._trial_in_flight = False + if failed_trial or self._failures >= self.failure_threshold: + # A failed trial restarts the cooldown rather than letting traffic + # straight back onto an engine that just proved it is still down. + self._opened_at = self.clock() + self._half_open = False + + def abandon_trial(self) -> None: + """Free the trial slot when a request ends without a verdict. + + A cancelled request or a client that leaves mid-stream says nothing + about the engine, and must not leave the circuit waiting on a trial + that will never report. + """ + + self._trial_in_flight = False + + def probe_succeeded(self) -> None: + """A healthy probe ends the cooldown early; it does not close the circuit.""" + + if self._opened_at is not None: + self._half_open = True + + +@dataclass +class ResilientBackend: + """Wraps an engine so its failures open a circuit instead of piling up.""" + + inner: InferenceBackend + breaker: CircuitBreaker + + def _admit(self) -> None: + if not self.breaker.allow(): + raise EngineCircuitOpenError(self.breaker.retry_after_seconds) + + async def generate( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> BackendResult: + self._admit() + try: + result = await self.inner.generate(request, decision) + except BackendUnavailableError: + self.breaker.record_failure() + raise + except BaseException: + self.breaker.abandon_trial() + raise + self.breaker.record_success() + return result + + async def stream( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> AsyncIterator[str]: + self._admit() + try: + async for delta in self.inner.stream(request, decision): + yield delta + except BackendUnavailableError: + self.breaker.record_failure() + raise + except BaseException: + self.breaker.abandon_trial() + raise + self.breaker.record_success() + + async def healthy(self) -> bool: + """Report ready only when the engine answers and the circuit is not open.""" + + if not await self.inner.healthy(): + return False + self.breaker.probe_succeeded() + return self.breaker.state != "open" diff --git a/tests/integration/test_resilience_api.py b/tests/integration/test_resilience_api.py new file mode 100644 index 0000000..1d8bf75 --- /dev/null +++ b/tests/integration/test_resilience_api.py @@ -0,0 +1,198 @@ +import asyncio +from collections.abc import AsyncIterator + +import httpx +import pytest +from fastapi.testclient import TestClient + +from llm_router.app import create_app +from llm_router.backends import ( + BackendOutOfMemoryError, + BackendResult, + BackendUnavailableError, + MockInferenceBackend, +) +from llm_router.config import Settings +from llm_router.models import ChatCompletionRequest, RouteDecision + +HEADERS = {"Authorization": "Bearer resilience-key"} +BODY = {"model": "auto", "messages": [{"role": "user", "content": "Classify this ticket"}]} + + +class LostNode(MockInferenceBackend): + def __init__(self) -> None: + self.calls = 0 + self.down = True + + async def generate( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> BackendResult: + self.calls += 1 + if self.down: + raise BackendUnavailableError("inference engine unreachable: connection refused") + return await super().generate(request, decision) + + async def healthy(self) -> bool: + return not self.down + + +class OutOfMemory(MockInferenceBackend): + async def generate( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> BackendResult: + raise BackendOutOfMemoryError("inference engine ran out of GPU memory for m") + + +class GatedStream(MockInferenceBackend): + """Holds its stream open until released, so an in-flight stream can be observed.""" + + def __init__(self) -> None: + self.started = asyncio.Event() + self.release = asyncio.Event() + + async def stream( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> AsyncIterator[str]: + self.started.set() + await self.release.wait() + yield "done " + + +def settings(**overrides: object) -> Settings: + return Settings(api_keys="resilience-key", **overrides) # type: ignore[arg-type] + + +def test_a_lost_node_opens_the_circuit_and_requests_then_fail_fast() -> None: + backend = LostNode() + app = create_app(settings(engine_failure_threshold=2), backend=backend) + with TestClient(app) as client: + first = client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + rejected = client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + ready = client.get("/readyz") + metrics = client.get("/metrics").text + + assert first.status_code == 502 + assert rejected.status_code == 503 + assert rejected.json()["error"]["type"] == "engine_unavailable" + assert int(rejected.headers["Retry-After"]) >= 1 + # The third request never reached the engine. + assert backend.calls == 2 + assert ready.status_code == 503 + assert 'router_engine_circuit_open{engine="mock"} 1.0' in metrics + assert 'router_rejections_total{type="engine_unavailable"} 1.0' in metrics + + +def test_the_gateway_recovers_once_the_engine_answers_its_probe() -> None: + backend = LostNode() + app = create_app(settings(engine_failure_threshold=1), backend=backend) + with TestClient(app) as client: + client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + assert client.get("/readyz").status_code == 503 + + backend.down = False + # The probe half-opens the circuit, and the trial request closes it. + assert client.get("/readyz").status_code == 200 + trial = client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + after = client.post( + "/v1/chat/completions", + headers=HEADERS, + json={**BODY, "messages": [{"role": "user", "content": "Classify this other one"}]}, + ) + metrics = client.get("/metrics").text + + assert trial.status_code == 200 and after.status_code == 200 + assert 'router_engine_circuit_open{engine="mock"} 0.0' in metrics + + +def test_out_of_memory_has_its_own_error_and_says_what_to_change() -> None: + with TestClient(create_app(settings(), backend=OutOfMemory())) as client: + response = client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + metrics = client.get("/metrics").text + + assert response.status_code == 503 + assert response.json()["error"]["type"] == "engine_out_of_memory" + assert response.headers["Retry-After"] == "10" + assert 'router_rejections_total{type="engine_out_of_memory"} 1.0' in metrics + + +def test_a_buffered_request_returns_its_slot_even_when_the_engine_fails() -> None: + app = create_app( + settings(max_concurrency=1, engine_failure_threshold=50), backend=OutOfMemory() + ) + with TestClient(app) as client: + statuses = [ + client.post("/v1/chat/completions", headers=HEADERS, json=BODY).status_code + for _ in range(3) + ] + + # Each failure released the only slot; none was rejected as overloaded. + assert statuses == [503, 503, 503] + + +@pytest.mark.asyncio +async def test_a_live_stream_holds_its_admission_slot_until_it_ends() -> None: + backend = GatedStream() + app = create_app(settings(max_concurrency=1, admission_timeout_seconds=0.05), backend=backend) + transport = httpx.ASGITransport(app=app) + async with httpx.AsyncClient(transport=transport, base_url="http://gateway") as client: + streaming = asyncio.create_task( + client.post("/v1/chat/completions", headers=HEADERS, json={**BODY, "stream": True}) + ) + await asyncio.wait_for(backend.started.wait(), timeout=5) + + # The only slot is held by a stream that is still generating. + blocked = await client.post( + "/v1/chat/completions", + headers=HEADERS, + json={**BODY, "messages": [{"role": "user", "content": "Classify another"}]}, + ) + assert blocked.status_code == 503 + assert blocked.json()["error"]["type"] == "overloaded" + + backend.release.set() + finished = await asyncio.wait_for(streaming, timeout=5) + assert finished.status_code == 200 + assert "data: [DONE]" in finished.text + + # Ending the stream returned the slot. + admitted = await client.post( + "/v1/chat/completions", + headers=HEADERS, + json={**BODY, "messages": [{"role": "user", "content": "Classify a third"}]}, + ) + assert admitted.status_code == 200 + + +def test_a_failed_stream_returns_its_slot() -> None: + class BrokenStream(MockInferenceBackend): + async def stream( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> AsyncIterator[str]: + raise BackendUnavailableError("inference engine unreachable while streaming") + yield "" # pragma: no cover - unreachable, keeps this an async generator + + app = create_app( + settings(max_concurrency=1, engine_failure_threshold=50), backend=BrokenStream() + ) + with TestClient(app, raise_server_exceptions=False) as client: + for _ in range(2): + with client.stream( + "POST", "/v1/chat/completions", headers=HEADERS, json={**BODY, "stream": True} + ) as response: + try: + list(response.iter_lines()) + except Exception: # the stream aborts mid-flight; that is the point + pass + healthy = client.post("/v1/chat/completions", headers=HEADERS, json=BODY) + + # Had either broken stream kept the only slot, this would be 503 overloaded. + assert healthy.status_code == 200 + + +def test_shutdown_stops_readiness_before_the_engine_client_closes() -> None: + app = create_app(settings(shutdown_grace_seconds=0.01)) + with TestClient(app) as client: + assert client.get("/readyz").status_code == 200 + + assert app.state.ready is False diff --git a/tests/unit/test_resilience.py b/tests/unit/test_resilience.py new file mode 100644 index 0000000..6806e33 --- /dev/null +++ b/tests/unit/test_resilience.py @@ -0,0 +1,296 @@ +import asyncio +from collections.abc import AsyncIterator + +import httpx +import pytest + +from llm_router.admission import AdmissionController, AdmissionRejectedError +from llm_router.backends import ( + BackendOutOfMemoryError, + BackendResult, + BackendUnavailableError, + MockInferenceBackend, + VLLMBackend, + engine_failure, +) +from llm_router.models import ChatCompletionRequest, ModelProfile, RouteDecision, TaskClass +from llm_router.resilience import CircuitBreaker, EngineCircuitOpenError, ResilientBackend + +REQUEST = ChatCompletionRequest.model_validate( + {"messages": [{"role": "user", "content": "patient record 4471"}]} +) +DECISION = RouteDecision( + profile=ModelProfile( + id="general-local", + revision="rev", + local=True, + context_limit=8192, + supported_tasks=frozenset(TaskClass), + quality=0.9, + ), + task=TaskClass.GENERAL, + reason="test", + score=1.0, + candidate_count=1, +) + + +class Clock: + def __init__(self) -> None: + self.now = 0.0 + + def __call__(self) -> float: + return self.now + + +class ScriptedBackend(MockInferenceBackend): + """Fails or succeeds on command, and counts how often it was reached.""" + + def __init__(self) -> None: + self.failing = True + self.engine_up = True + self.calls = 0 + + async def generate( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> BackendResult: + self.calls += 1 + if self.failing: + raise BackendUnavailableError("inference engine unreachable") + return BackendResult(text="ok", prompt_tokens=1, completion_tokens=1) + + async def stream( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> AsyncIterator[str]: + self.calls += 1 + if self.failing: + raise BackendUnavailableError("inference engine unreachable") + yield "one " + yield "two " + + async def healthy(self) -> bool: + return self.engine_up + + +def build(threshold: int = 3) -> tuple[ResilientBackend, ScriptedBackend, CircuitBreaker, Clock]: + clock = Clock() + breaker = CircuitBreaker(failure_threshold=threshold, cooldown_seconds=30.0, clock=clock) + inner = ScriptedBackend() + return ResilientBackend(inner, breaker), inner, breaker, clock + + +async def fail(backend: ResilientBackend, times: int) -> None: + for _ in range(times): + with pytest.raises(BackendUnavailableError): + await backend.generate(REQUEST, DECISION) + + +@pytest.mark.asyncio +async def test_the_circuit_opens_after_consecutive_failures_and_fails_fast() -> None: + backend, inner, breaker, _ = build(threshold=3) + + await fail(backend, 3) + + assert breaker.state == "open" + with pytest.raises(EngineCircuitOpenError) as raised: + await backend.generate(REQUEST, DECISION) + # The engine was not contacted again: a lost node costs no further timeouts. + assert inner.calls == 3 + assert raised.value.retry_after_seconds == pytest.approx(30.0) + + +@pytest.mark.asyncio +async def test_a_success_resets_the_failure_count() -> None: + backend, inner, breaker, _ = build(threshold=3) + + await fail(backend, 2) + inner.failing = False + await backend.generate(REQUEST, DECISION) + inner.failing = True + await fail(backend, 2) + + assert breaker.state == "closed" + + +@pytest.mark.asyncio +async def test_after_the_cooldown_exactly_one_trial_is_admitted() -> None: + backend, _, breaker, clock = build(threshold=1) + await fail(backend, 1) + + clock.now = 31.0 + + assert breaker.state == "half-open" + assert breaker.allow() is True + assert breaker.allow() is False + + +@pytest.mark.asyncio +async def test_a_successful_trial_closes_the_circuit() -> None: + backend, inner, breaker, clock = build(threshold=1) + await fail(backend, 1) + clock.now = 31.0 + inner.failing = False + + await backend.generate(REQUEST, DECISION) + + assert breaker.state == "closed" + assert breaker.retry_after_seconds == 0.0 + + +@pytest.mark.asyncio +async def test_a_failed_trial_reopens_the_circuit_and_restarts_the_cooldown() -> None: + backend, _, breaker, clock = build(threshold=3) + await fail(backend, 3) + clock.now = 31.0 + + await fail(backend, 1) + + assert breaker.state == "open" + assert breaker.retry_after_seconds == pytest.approx(30.0) + + +@pytest.mark.asyncio +async def test_readiness_fails_while_open_and_a_healthy_probe_only_half_opens() -> None: + backend, inner, breaker, _ = build(threshold=1) + await fail(backend, 1) + + inner.engine_up = False + assert await backend.healthy() is False + assert breaker.state == "open" + + # The engine answers its health check again, well inside the cooldown. + inner.engine_up = True + assert await backend.healthy() is True + assert breaker.state == "half-open" + + # Only a real request succeeding closes it. + inner.failing = False + await backend.generate(REQUEST, DECISION) + assert breaker.state == "closed" + + +@pytest.mark.asyncio +async def test_a_healthy_engine_with_a_closed_circuit_is_simply_ready() -> None: + backend, _, breaker, _ = build() + + assert await backend.healthy() is True + assert breaker.state == "closed" + + +@pytest.mark.asyncio +async def test_stream_failures_count_and_stream_successes_close() -> None: + backend, inner, breaker, clock = build(threshold=1) + + with pytest.raises(BackendUnavailableError): + async for _ in backend.stream(REQUEST, DECISION): + pass + assert breaker.state == "open" + with pytest.raises(EngineCircuitOpenError): + async for _ in backend.stream(REQUEST, DECISION): + pass + + clock.now = 31.0 + inner.failing = False + assert [delta async for delta in backend.stream(REQUEST, DECISION)] == ["one ", "two "] + assert breaker.state == "closed" + + +@pytest.mark.asyncio +async def test_a_trial_abandoned_mid_stream_does_not_wedge_the_circuit() -> None: + backend, inner, breaker, clock = build(threshold=1) + await fail(backend, 1) + clock.now = 31.0 + inner.failing = False + + # The client leaves after the first chunk: no verdict on the engine. + stream = backend.stream(REQUEST, DECISION) + assert await anext(stream) == "one " + await stream.aclose() + + assert breaker.state == "half-open" + assert breaker.allow() is True + + +@pytest.mark.asyncio +async def test_a_cancelled_trial_does_not_wedge_the_circuit() -> None: + class Hanging(ScriptedBackend): + async def generate( + self, request: ChatCompletionRequest, decision: RouteDecision + ) -> BackendResult: + await asyncio.sleep(60) + raise AssertionError("unreachable") + + clock = Clock() + breaker = CircuitBreaker(failure_threshold=1, cooldown_seconds=30.0, clock=clock) + breaker.record_failure() + clock.now = 31.0 + backend = ResilientBackend(Hanging(), breaker) + + task = asyncio.create_task(backend.generate(REQUEST, DECISION)) + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + assert breaker.allow() is True + + +def test_engine_failures_are_classified_without_repeating_the_engine() -> None: + oom = engine_failure( + 500, "torch.OutOfMemoryError: CUDA out of memory while serving 'patient 4471'", "for m" + ) + plain = engine_failure(500, "internal error for prompt 'patient 4471'", "for m") + + assert isinstance(oom, BackendOutOfMemoryError) + assert "smaller max_tokens" in str(oom) + assert not isinstance(plain, BackendOutOfMemoryError) + # Neither message carries what the engine said back. + assert "4471" not in str(oom) and "4471" not in str(plain) + + +@pytest.mark.asyncio +async def test_the_vllm_backend_reports_out_of_memory_on_both_paths() -> None: + def handler(_: httpx.Request) -> httpx.Response: + return httpx.Response(500, text="CUDA out of memory. Tried to allocate 2.00 GiB") + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client: + backend = VLLMBackend(base_url="http://engine", client=client) + + with pytest.raises(BackendOutOfMemoryError): + await backend.generate(REQUEST, DECISION) + with pytest.raises(BackendOutOfMemoryError): + async for _ in backend.stream(REQUEST, DECISION): + pass + + +@pytest.mark.asyncio +async def test_admission_tracks_in_flight_work_and_drains() -> None: + admission = AdmissionController(max_concurrency=2, timeout_seconds=0.01) + + assert await admission.drain(0.01) is True + + await admission.acquire() + await admission.acquire() + assert admission.in_flight == 2 + with pytest.raises(AdmissionRejectedError): + await admission.acquire() + assert await admission.drain(0.01) is False + + admission.release() + admission.release() + assert admission.in_flight == 0 + assert await admission.drain(0.01) is True + + +@pytest.mark.asyncio +async def test_drain_returns_as_soon_as_the_last_request_finishes() -> None: + admission = AdmissionController(max_concurrency=1, timeout_seconds=0.01) + await admission.acquire() + + async def finish() -> None: + await asyncio.sleep(0.01) + admission.release() + + task = asyncio.create_task(finish()) + assert await admission.drain(5.0) is True + await task