Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/middleware/admission_control_middleware.py: 44%
177 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import asyncio
2import os
3from collections.abc import Callable, Mapping
4from dataclasses import dataclass
5from functools import lru_cache
6from typing import Annotated, Final, Protocol, TypeAlias, runtime_checkable
8from pydantic import Field, TypeAdapter, ValidationError
9from starlette.responses import JSONResponse
10from starlette.types import ASGIApp, Receive, Scope, Send
12from litellm._logging import verbose_proxy_logger
14_EXEMPT_PATHS: Final[frozenset[str]] = frozenset(
15 {
16 "/health/liveliness",
17 "/health/liveness",
18 "/health/readiness",
19 "/health/readiness/details",
20 "/health/backlog",
21 "/health/drain",
22 "/metrics",
23 "/metrics/",
24 }
25)
28@dataclass(frozen=True, slots=True)
29class AdmissionControlSettings:
30 max_in_flight_requests: int
31 max_queued_requests: int
32 queue_timeout_seconds: float
35@dataclass(frozen=True, slots=True)
36class AdmissionControlStats:
37 admitted: int
38 queued: int
39 rejected_total: int
42@runtime_checkable
43class _Gauge(Protocol):
44 def inc(self, amount: float = 1) -> None: ... 44 ↛ exitline 44 didn't return from function 'inc' because
46 def dec(self, amount: float = 1) -> None: ... 46 ↛ exitline 46 didn't return from function 'dec' because
49@runtime_checkable
50class _CounterChild(Protocol):
51 def inc(self, amount: float = 1) -> None: ... 51 ↛ exitline 51 didn't return from function 'inc' because
54@runtime_checkable
55class _Counter(Protocol):
56 def labels(self, reason: str) -> _CounterChild: ... 56 ↛ exitline 56 didn't return from function 'labels' because
59@dataclass(frozen=True, slots=True)
60class AdmissionControlMetrics:
61 admitted_gauge: _Gauge
62 queued_gauge: _Gauge
63 rejected_counter: _Counter
66class AdmissionControlState:
67 """Per-process admission counters and the in-flight semaphore shared by one worker's requests."""
69 def __init__(self, metrics_factory: Callable[[], AdmissionControlMetrics | None]) -> None:
70 self._metrics_factory = metrics_factory
71 self._metrics: AdmissionControlMetrics | None = None
72 self._metrics_init_attempted = False
73 self._admitted = 0
74 self._queued = 0
75 self._rejected_total = 0
76 self._semaphore: asyncio.Semaphore | None = None
77 self._semaphore_loop: asyncio.AbstractEventLoop | None = None
79 def get_stats(self) -> AdmissionControlStats:
80 return AdmissionControlStats(
81 admitted=self._admitted,
82 queued=self._queued,
83 rejected_total=self._rejected_total,
84 )
86 def get_semaphore(self, max_in_flight_requests: int) -> asyncio.Semaphore:
87 loop: Final = asyncio.get_running_loop()
88 if self._semaphore_loop is not loop:
89 self._semaphore = asyncio.Semaphore(max_in_flight_requests)
90 self._semaphore_loop = loop
91 semaphore: Final = self._semaphore
92 if semaphore is None:
93 raise RuntimeError("Admission control semaphore was not initialized")
94 return semaphore
96 def record_admission(self) -> None:
97 self._admitted += 1
98 metrics: Final = self._get_metrics()
99 if metrics is not None:
100 metrics.admitted_gauge.inc()
102 def record_release(self) -> None:
103 self._admitted -= 1
104 metrics: Final = self._get_metrics()
105 if metrics is not None:
106 metrics.admitted_gauge.dec()
108 def record_queue(self) -> None:
109 self._queued += 1
110 metrics: Final = self._get_metrics()
111 if metrics is not None:
112 metrics.queued_gauge.inc()
114 def record_dequeue(self) -> None:
115 self._queued -= 1
116 metrics: Final = self._get_metrics()
117 if metrics is not None:
118 metrics.queued_gauge.dec()
120 def record_rejection(self, reason: str) -> None:
121 self._rejected_total += 1
122 metrics: Final = self._get_metrics()
123 if metrics is not None:
124 metrics.rejected_counter.labels(reason=reason).inc()
126 def _get_metrics(self) -> AdmissionControlMetrics | None:
127 if not self._metrics_init_attempted:
128 self._metrics_init_attempted = True
129 self._metrics = self._metrics_factory()
130 return self._metrics
133class AdmissionControlMiddleware:
134 def __init__(
135 self,
136 app: ASGIApp,
137 get_settings: Callable[[], AdmissionControlSettings | None],
138 state: AdmissionControlState,
139 ) -> None:
140 self.app = app
141 self.get_settings = get_settings
142 self.state = state
144 async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
145 if scope["type"] != "http":
146 await self.app(scope, receive, send)
147 return
149 settings: Final = self.get_settings()
150 if settings is None or _get_route_path(scope) in _EXEMPT_PATHS: 150 ↛ 154line 150 didn't jump to line 154 because the condition on line 150 was always true
151 await self.app(scope, receive, send)
152 return
154 state: Final = self.state
155 semaphore: Final = state.get_semaphore(settings.max_in_flight_requests)
156 if not semaphore.locked():
157 await semaphore.acquire()
158 state.record_admission()
159 elif state.get_stats().queued >= settings.max_queued_requests:
160 state.record_rejection("queue_full")
161 await _overloaded_response(state)(scope, receive, send)
162 return
163 else:
164 state.record_queue()
165 try:
166 await asyncio.wait_for(
167 semaphore.acquire(),
168 timeout=settings.queue_timeout_seconds,
169 )
170 except asyncio.TimeoutError:
171 state.record_dequeue()
172 state.record_rejection("queue_timeout")
173 await _overloaded_response(state)(scope, receive, send)
174 return
175 except asyncio.CancelledError:
176 state.record_dequeue()
177 raise
178 state.record_dequeue()
179 state.record_admission()
181 try:
182 await self.app(scope, receive, send)
183 finally:
184 semaphore.release()
185 state.record_release()
188def _get_route_path(scope: Scope) -> str:
189 """Strip the ASGI root_path (SERVER_ROOT_PATH) the same way Starlette does before route matching."""
190 path: Final[str] = scope["path"]
191 root_path: Final[str] = scope.get("root_path", "")
192 if not root_path or not path.startswith(root_path):
193 return path
194 if path == root_path:
195 return ""
196 if path[len(root_path)] == "/":
197 return path[len(root_path) :]
198 return path
201def _create_gauge(gauge_type: Callable[..., object], name: str, description: str) -> _Gauge:
202 metric: Final = (
203 gauge_type(name, description, multiprocess_mode="livesum")
204 if "PROMETHEUS_MULTIPROC_DIR" in os.environ
205 else gauge_type(name, description)
206 )
207 if not isinstance(metric, _Gauge):
208 raise TypeError("Admission gauge has an unexpected type")
209 return metric
212def create_prometheus_admission_metrics() -> AdmissionControlMetrics | None:
213 try:
214 from prometheus_client import Counter, Gauge
216 return AdmissionControlMetrics(
217 admitted_gauge=_create_gauge(
218 Gauge,
219 "litellm_admission_admitted_requests",
220 "Number of requests admitted by this worker",
221 ),
222 queued_gauge=_create_gauge(
223 Gauge,
224 "litellm_admission_queued_requests",
225 "Number of requests queued by this worker",
226 ),
227 rejected_counter=Counter( # mutable-ok: Prometheus requires runtime Counter construction
228 "litellm_admission_rejected_requests_total",
229 "Number of requests rejected by this worker",
230 labelnames=("reason",),
231 ),
232 )
233 except (ImportError, ValueError):
234 return None
237admission_control_state: Final = AdmissionControlState(create_prometheus_admission_metrics)
240def get_admission_control_stats() -> AdmissionControlStats:
241 return admission_control_state.get_stats()
244_PositiveInt: TypeAlias = Annotated[int, Field(gt=0)]
245_NonNegativeInt: TypeAlias = Annotated[int, Field(ge=0)]
246_PositiveFloat: TypeAlias = Annotated[float, Field(gt=0)]
247_AdmissionControlRaw: TypeAlias = int | float | str | None
250def _hashable(value: object) -> _AdmissionControlRaw:
251 return value if value is None or isinstance(value, (int, float, str)) else repr(value)
254_POSITIVE_INT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(_PositiveInt)
255_NON_NEGATIVE_INT_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(_NonNegativeInt)
256_POSITIVE_FLOAT_ADAPTER: Final[TypeAdapter[float]] = TypeAdapter(_PositiveFloat)
259@lru_cache(maxsize=16)
260def _parse_admission_control_settings(
261 max_in_flight_raw: _AdmissionControlRaw,
262 max_queued_raw: _AdmissionControlRaw,
263 queue_timeout_raw: _AdmissionControlRaw,
264) -> AdmissionControlSettings | None:
265 try:
266 max_in_flight: Final = _POSITIVE_INT_ADAPTER.validate_python(max_in_flight_raw)
267 max_queued: Final = (
268 max_in_flight if max_queued_raw is None else _NON_NEGATIVE_INT_ADAPTER.validate_python(max_queued_raw)
269 )
270 queue_timeout: Final = _POSITIVE_FLOAT_ADAPTER.validate_python(queue_timeout_raw)
271 except ValidationError as exc:
272 verbose_proxy_logger.error(
273 "Ignoring invalid admission control settings, per-worker admission control is disabled: %s",
274 exc,
275 )
276 return None
277 return AdmissionControlSettings(
278 max_in_flight_requests=max_in_flight,
279 max_queued_requests=max_queued,
280 queue_timeout_seconds=queue_timeout,
281 )
284def get_admission_control_settings(settings: Mapping[str, object]) -> AdmissionControlSettings | None:
285 max_in_flight_raw: Final = settings.get("max_in_flight_requests_per_worker")
286 if max_in_flight_raw is None: 286 ↛ 288line 286 didn't jump to line 288 because the condition on line 286 was always true
287 return None
288 return _parse_admission_control_settings(
289 _hashable(max_in_flight_raw),
290 _hashable(settings.get("max_queued_requests_per_worker")),
291 _hashable(settings.get("admission_queue_timeout_seconds", 1.0)),
292 )
295def _overloaded_response(state: AdmissionControlState) -> JSONResponse:
296 stats: Final = state.get_stats()
297 return JSONResponse(
298 status_code=503,
299 headers={"retry-after": "1"}, # mutable-ok: Starlette expects a plain headers mapping
300 content={ # mutable-ok: Starlette serializes a plain response mapping
301 "error": { # mutable-ok: nested response mapping
302 "message": (
303 f"Worker at capacity: {stats.admitted} in-flight, {stats.queued} queued requests. Retry later."
304 ),
305 "type": "overloaded_error",
306 "code": "503",
307 }
308 },
309 )