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

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 

7 

8from pydantic import Field, TypeAdapter, ValidationError 

9from starlette.responses import JSONResponse 

10from starlette.types import ASGIApp, Receive, Scope, Send 

11 

12from litellm._logging import verbose_proxy_logger 

13 

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) 

26 

27 

28@dataclass(frozen=True, slots=True) 

29class AdmissionControlSettings: 

30 max_in_flight_requests: int 

31 max_queued_requests: int 

32 queue_timeout_seconds: float 

33 

34 

35@dataclass(frozen=True, slots=True) 

36class AdmissionControlStats: 

37 admitted: int 

38 queued: int 

39 rejected_total: int 

40 

41 

42@runtime_checkable 

43class _Gauge(Protocol): 

44 def inc(self, amount: float = 1) -> None: ... 44 ↛ exitline 44 didn't return from function 'inc' because

45 

46 def dec(self, amount: float = 1) -> None: ... 46 ↛ exitline 46 didn't return from function 'dec' because

47 

48 

49@runtime_checkable 

50class _CounterChild(Protocol): 

51 def inc(self, amount: float = 1) -> None: ... 51 ↛ exitline 51 didn't return from function 'inc' because

52 

53 

54@runtime_checkable 

55class _Counter(Protocol): 

56 def labels(self, reason: str) -> _CounterChild: ... 56 ↛ exitline 56 didn't return from function 'labels' because

57 

58 

59@dataclass(frozen=True, slots=True) 

60class AdmissionControlMetrics: 

61 admitted_gauge: _Gauge 

62 queued_gauge: _Gauge 

63 rejected_counter: _Counter 

64 

65 

66class AdmissionControlState: 

67 """Per-process admission counters and the in-flight semaphore shared by one worker's requests.""" 

68 

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 

78 

79 def get_stats(self) -> AdmissionControlStats: 

80 return AdmissionControlStats( 

81 admitted=self._admitted, 

82 queued=self._queued, 

83 rejected_total=self._rejected_total, 

84 ) 

85 

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 

95 

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() 

101 

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() 

107 

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() 

113 

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() 

119 

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() 

125 

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 

131 

132 

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 

143 

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 

148 

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 

153 

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() 

180 

181 try: 

182 await self.app(scope, receive, send) 

183 finally: 

184 semaphore.release() 

185 state.record_release() 

186 

187 

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 

199 

200 

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 

210 

211 

212def create_prometheus_admission_metrics() -> AdmissionControlMetrics | None: 

213 try: 

214 from prometheus_client import Counter, Gauge 

215 

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 

235 

236 

237admission_control_state: Final = AdmissionControlState(create_prometheus_admission_metrics) 

238 

239 

240def get_admission_control_stats() -> AdmissionControlStats: 

241 return admission_control_state.get_stats() 

242 

243 

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 

248 

249 

250def _hashable(value: object) -> _AdmissionControlRaw: 

251 return value if value is None or isinstance(value, (int, float, str)) else repr(value) 

252 

253 

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) 

257 

258 

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 ) 

282 

283 

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 ) 

293 

294 

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 )