From 00d7706bfc99b6d008045c7806ba11f1d772233d Mon Sep 17 00:00:00 2001 From: Yash-Chindam Date: Sat, 3 Oct 2026 19:48:48 +0530 Subject: [PATCH] fix: hold streamed requests to the concurrency bound and define engine failure A streamed request returned its response from inside the admission slot, so the slot was freed before generation began and streaming escaped the concurrency bound that backpressure depends on. A stream now holds its slot until it ends, including when the client disconnects before the first chunk or the engine fails mid-stream. Section 13 also requires behaviour to be defined for GPU out-of-memory and node loss, and none was: every engine failure was one generic 502, each waiting out its own timeout. Out of memory now has its own error type and tells the caller what to change, since the same request at the same size cannot succeed. After consecutive failures a circuit opens: requests fail immediately with retry guidance instead of timing out against a node that is gone, and readiness fails so the orchestrator stops routing here. A healthy probe lets exactly one trial request through, and only its success closes the circuit. A cancelled or abandoned trial does not wedge it. An engine error body is checked for an out-of-memory report and then discarded rather than forwarded, since it can echo the prompt. Shutdown fails readiness first, then gives admitted requests a grace period before the engine client closes underneath them. Co-Authored-By: Claude Opus 5.5 --- README.md | 26 ++ src/llm_router/admission.py | 40 ++- src/llm_router/app.py | 137 ++++++++--- src/llm_router/backends.py | 32 ++- src/llm_router/config.py | 7 + src/llm_router/observability.py | 10 + src/llm_router/resilience.py | 158 ++++++++++++ tests/integration/test_resilience_api.py | 198 +++++++++++++++ tests/unit/test_resilience.py | 296 +++++++++++++++++++++++ 9 files changed, 863 insertions(+), 41 deletions(-) create mode 100644 src/llm_router/resilience.py create mode 100644 tests/integration/test_resilience_api.py create mode 100644 tests/unit/test_resilience.py 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