Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/parallel_request_limiter_v3.py: 30%

1646 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1""" 

2This is a rate limiter implementation based on a similar one by Envoy proxy. 

3 

4This is currently in development and not yet ready for production. 

5""" 

6 

7import asyncio 

8import binascii 

9import logging 

10import os 

11import uuid 

12from collections.abc import AsyncGenerator, Awaitable, Callable, Mapping, Sequence, Set 

13from contextlib import asynccontextmanager 

14from contextvars import ContextVar 

15from dataclasses import dataclass, field 

16from datetime import datetime, timezone 

17from types import MappingProxyType 

18from typing import ( 

19 TYPE_CHECKING, 

20 Any, 

21 Final, 

22 Literal, 

23 Protocol, 

24 TypeAlias, 

25 TypedDict, 

26) 

27 

28from pydantic import TypeAdapter 

29from typing_extensions import NotRequired, ReadOnly 

30 

31from litellm import DualCache 

32from litellm._logging import verbose_proxy_logger 

33from litellm.caching.redis_cache import log_redis_failure 

34from litellm.constants import DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE, INTERNAL_CALL_ORIGIN_METADATA_KEY 

35from litellm.integrations.custom_logger import CustomLogger 

36from litellm.litellm_core_utils.prompt_templates.common_utils import ( 

37 get_str_from_messages, 

38) 

39from litellm.litellm_core_utils.token_counter import offload_token_count 

40from litellm.proxy._types import UserAPIKeyAuth 

41from litellm.proxy.auth.auth_utils import ( 

42 ESTIMATED_OUTPUT_TOKENS_FIELD, 

43 get_estimated_output_tokens, 

44 get_key_own_model_rate_limit, 

45 get_key_tag_rpm_limit, 

46 get_model_rate_limit_from_metadata, 

47) 

48from litellm.proxy.auth.budget_throttle import throttled_limit 

49from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body 

50from litellm.proxy.common_utils.proxy_rate_limit_error import ( 

51 ProxyRateLimitError, 

52 map_v3_rate_limit_type, 

53) 

54from litellm.proxy.hooks.batch_enqueued_tokens import ( 

55 BATCH_ENQUEUED_REFUND_STATUSES, 

56 BatchEnqueuedTokenReservation, 

57 BatchEnqueuedTokenStore, 

58 batch_response_view, 

59 canonical_provider_batch_id, 

60) 

61from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit 

62from litellm.router_utils.add_retry_fallback_headers import ( 

63 ensure_response_additional_headers, 

64 response_has_hidden_params, 

65) 

66from litellm.router_utils.common_utils import resolve_model_group_alias 

67from litellm.types.caching import RedisPipelineIncrementOperation 

68from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject, ResponseAPIUsage 

69from litellm.types.utils import ( 

70 CallTypes, 

71 EmbeddingResponse, 

72 ModelResponse, 

73 RerankResponse, 

74 TextCompletionResponse, 

75 Usage, 

76) 

77 

78if TYPE_CHECKING: 78 ↛ 79line 78 didn't jump to line 79 because the condition on line 78 was never true

79 from opentelemetry.trace import Span as _Span 

80 

81 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache 

82 from litellm.types.agents import AgentResponse 

83 from litellm.types.caching import RedisPipelineIncrementOperation 

84 

85 Span = _Span | Any 

86 InternalUsageCache = _InternalUsageCache 

87else: 

88 Span = Any 

89 InternalUsageCache = Any 

90 

91 

92_REQUEST_RATE_LIMIT_DATA: Final = TypeAdapter(Mapping[str, object]) 

93 

94 

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

96class RateLimitedModel: 

97 requested: str 

98 group: str 

99 

100 def limit_in(self, limits: Mapping[str, int] | None) -> int | None: 

101 if limits is None: 101 ↛ 103line 101 didn't jump to line 103 because the condition on line 101 was always true

102 return None 

103 requested_limit: Final = limits.get(self.requested) 

104 return requested_limit if requested_limit is not None else limits.get(self.group) 

105 

106 

107def _resolve_model_group_alias_via_proxy_router(model: str) -> str | None: 

108 from litellm.proxy.proxy_server import llm_router 

109 

110 if llm_router is None: 110 ↛ 111line 110 didn't jump to line 111 because the condition on line 110 was never true

111 return None 

112 return resolve_model_group_alias(llm_router.model_group_alias, model) 

113 

114 

115def _sibling_counter_keys(window_key: str) -> tuple[str, str]: 

116 prefix: Final = window_key.removesuffix(":window") 

117 return f"{prefix}:requests", f"{prefix}:tokens" 

118 

119 

120BATCH_RATE_LIMITER_SCRIPT: Final = """ 

121local results = {} 

122local now = tonumber(ARGV[1]) 

123local window_size = tonumber(ARGV[2]) 

124 

125-- Process each window/counter pair 

126for i = 1, #KEYS, 2 do 

127 local window_key = KEYS[i] 

128 local counter_key = KEYS[i + 1] 

129 local increment_value = 1 

130 

131 -- Check if window exists and is valid 

132 local window_start = redis.call('GET', window_key) 

133 if not window_start or (now - tonumber(window_start)) >= window_size then 

134 -- Reset window and counter 

135 local prefix = string.sub(window_key, 1, -(#':window') - 1) 

136 redis.call('DEL', prefix .. ':requests', prefix .. ':tokens') 

137 redis.call('SET', window_key, tostring(now)) 

138 redis.call('SET', counter_key, increment_value) 

139 redis.call('EXPIRE', window_key, window_size) 

140 redis.call('EXPIRE', counter_key, window_size) 

141 table.insert(results, tostring(now)) -- window_start 

142 table.insert(results, increment_value) -- counter 

143 else 

144 local counter = redis.call('INCR', counter_key) 

145 -- This happens when window_key exists but counter_key doesn't (e.g., tokens key 

146 -- created after requests key when both share the same window_key) 

147 local current_ttl = redis.call('TTL', counter_key) 

148 if current_ttl == -1 then 

149 redis.call('EXPIRE', counter_key, window_size) 

150 end 

151 table.insert(results, window_start) -- window_start 

152 table.insert(results, counter) -- counter 

153 end 

154end 

155 

156return results 

157""" 

158 

159CHECK_AND_INCREMENT_BY_N_SCRIPT: Final = """ 

160-- Atomic check-and-increment-by-N across one or more descriptors. 

161-- All-or-nothing: if any descriptor would exceed its limit, no counter is 

162-- modified. 

163-- 

164-- Uses Redis server time (`redis.call('TIME')`) instead of a client-supplied 

165-- timestamp so that window resets are deterministic across replicas with 

166-- skewed wall-clocks. This prevents a clock-skew-induced reopening of the 

167-- TOCTOU window across multi-replica deployments. 

168-- 

169-- KEYS layout: pairs of (window_key, counter_key), one pair per descriptor. 

170-- ARGV layout: per-descriptor 4-tuple, starting at ARGV[1]: 

171-- ARGV[(i-1)*4 + 1] = limit 

172-- ARGV[(i-1)*4 + 2] = increment 

173-- ARGV[(i-1)*4 + 3] = ttl_seconds (counter TTL when window resets) 

174-- ARGV[(i-1)*4 + 4] = window_size_seconds (sliding-window length) 

175-- 

176-- Return on success: 

177-- { 0, new_counter_1, window_start_1, new_counter_2, window_start_2, ... } 

178-- Return on over-limit: { 1, descriptor_index, current_counter, limit } 

179local time_reply = redis.call('TIME') 

180local now = tonumber(time_reply[1]) 

181local descriptor_count = #KEYS / 2 

182local reset_windows = {} 

183 

184-- Pass 1: read state, validate. Abort without writing if any over limit. 

185local descriptor_state = {} 

186for i = 1, descriptor_count do 

187 local window_key = KEYS[(i - 1) * 2 + 1] 

188 local counter_key = KEYS[(i - 1) * 2 + 2] 

189 local arg_base = (i - 1) * 4 + 1 

190 local limit = tonumber(ARGV[arg_base]) 

191 local increment = tonumber(ARGV[arg_base + 1]) 

192 local window_size = tonumber(ARGV[arg_base + 3]) 

193 

194 local window_start = redis.call('GET', window_key) 

195 local window_expired = (not window_start) or 

196 ((now - tonumber(window_start)) >= window_size) 

197 

198 local current_counter 

199 if window_expired then 

200 current_counter = 0 

201 else 

202 current_counter = tonumber(redis.call('GET', counter_key) or 0) 

203 end 

204 

205 local blocked 

206 if increment > 0 then 

207 blocked = current_counter + increment > limit 

208 else 

209 blocked = current_counter >= limit 

210 end 

211 if blocked then 

212 return { 1, i, current_counter, limit } 

213 end 

214 

215 descriptor_state[i] = { window_expired, current_counter, window_start } 

216end 

217 

218-- Pass 2: all checks passed. Apply increments. 

219local results = { 0 } 

220for i = 1, descriptor_count do 

221 local window_key = KEYS[(i - 1) * 2 + 1] 

222 local counter_key = KEYS[(i - 1) * 2 + 2] 

223 local arg_base = (i - 1) * 4 + 1 

224 local increment = tonumber(ARGV[arg_base + 1]) 

225 local ttl = tonumber(ARGV[arg_base + 2]) 

226 local window_size = tonumber(ARGV[arg_base + 3]) 

227 

228 local window_expired = descriptor_state[i][1] 

229 local active_window_start 

230 

231 if window_expired then 

232 active_window_start = now 

233 if not reset_windows[window_key] then 

234 local prefix = string.sub(window_key, 1, -(#':window') - 1) 

235 redis.call('DEL', prefix .. ':requests', prefix .. ':tokens') 

236 reset_windows[window_key] = true 

237 end 

238 redis.call('SET', window_key, tostring(now)) 

239 redis.call('SET', counter_key, increment) 

240 redis.call('EXPIRE', window_key, window_size) 

241 if ttl > 0 then 

242 redis.call('EXPIRE', counter_key, ttl) 

243 end 

244 table.insert(results, increment) 

245 else 

246 active_window_start = tonumber(descriptor_state[i][3]) 

247 local new_counter = redis.call('INCRBY', counter_key, increment) 

248 local current_ttl = redis.call('TTL', counter_key) 

249 if current_ttl == -1 and ttl > 0 then 

250 redis.call('EXPIRE', counter_key, ttl) 

251 end 

252 table.insert(results, new_counter) 

253 end 

254 table.insert(results, active_window_start) 

255end 

256 

257return results 

258""" 

259 

260WINDOW_GUARDED_TOKEN_INCREMENT_SCRIPT: Final = """ 

261local results = {} 

262for i = 1, #KEYS, 2 do 

263 local window_key = KEYS[i] 

264 local counter_key = KEYS[i + 1] 

265 local arg_base = ((i - 1) / 2) * 3 + 1 

266 local expected_window_start = ARGV[arg_base] 

267 local increment = tonumber(ARGV[arg_base + 1]) 

268 local ttl = tonumber(ARGV[arg_base + 2]) 

269 local active_window_start = redis.call('GET', window_key) 

270 

271 if active_window_start and active_window_start == expected_window_start then 

272 local new_counter = redis.call('INCRBY', counter_key, increment) 

273 local current_ttl = redis.call('TTL', counter_key) 

274 if current_ttl == -1 and ttl > 0 then 

275 redis.call('EXPIRE', counter_key, ttl) 

276 end 

277 table.insert(results, 1) 

278 table.insert(results, new_counter) 

279 else 

280 table.insert(results, 0) 

281 table.insert(results, tonumber(redis.call('GET', counter_key) or 0)) 

282 end 

283end 

284return results 

285""" 

286 

287PARALLEL_ACQUIRE_SCRIPT: Final = """ 

288-- Atomic check-and-acquire for the max_parallel_requests concurrency gauge. 

289-- Each gauge key is a sorted set of per-request slot ids scored by acquire 

290-- time (Redis server clock). In-flight requests are counted by ZCARD after 

291-- pruning slots older than the slot TTL, so unlike the windowed RPM/TPM 

292-- counters the gauge is never reset while requests are in flight, a 

293-- rejected request never occupies a slot, and a slot leaked by a crashed 

294-- worker self-heals after the slot TTL even under continuous traffic. 

295-- 

296-- KEYS: one gauge zset key per descriptor. 

297-- ARGV: per-key triples (limit, slot_ttl_seconds, slot_id). 

298-- Success: { 0, in_flight_1, ... }. Over-limit: { 1, key_index, in_flight, limit }. 

299local time_reply = redis.call('TIME') 

300local now = tonumber(time_reply[1]) 

301for i = 1, #KEYS do 

302 local limit = tonumber(ARGV[(i - 1) * 3 + 1]) 

303 local slot_ttl = tonumber(ARGV[(i - 1) * 3 + 2]) 

304 redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', now - slot_ttl) 

305 local in_flight = redis.call('ZCARD', KEYS[i]) 

306 if in_flight + 1 > limit then 

307 return { 1, i, in_flight, limit } 

308 end 

309end 

310local results = { 0 } 

311for i = 1, #KEYS do 

312 local slot_ttl = tonumber(ARGV[(i - 1) * 3 + 2]) 

313 local slot_id = ARGV[(i - 1) * 3 + 3] 

314 redis.call('ZADD', KEYS[i], now, slot_id) 

315 redis.call('EXPIRE', KEYS[i], slot_ttl) 

316 table.insert(results, redis.call('ZCARD', KEYS[i])) 

317end 

318return results 

319""" 

320 

321PARALLEL_RELEASE_SCRIPT: Final = """ 

322-- Release one slot per gauge key by removing this request's slot id. 

323-- ZREM of an absent member (or key) is a no-op, so a release without a 

324-- matching acquire (proxy-side rejection, double-fired callback, slot 

325-- already expired) can never free a slot owned by another request. 

326-- KEYS: gauge zset keys. ARGV: per-key slot_id. 

327-- Returns the remaining in-flight count per key. 

328local results = {} 

329for i = 1, #KEYS do 

330 redis.call('ZREM', KEYS[i], ARGV[i]) 

331 table.insert(results, redis.call('ZCARD', KEYS[i])) 

332end 

333return results 

334""" 

335 

336PARALLEL_COUNT_SCRIPT: Final = """ 

337-- Read the current in-flight count per gauge key (prunes expired slots 

338-- first so leaked slots do not inflate the reading). 

339-- KEYS: gauge zset keys. ARGV: per-key slot_ttl_seconds. 

340local time_reply = redis.call('TIME') 

341local now = tonumber(time_reply[1]) 

342local results = {} 

343for i = 1, #KEYS do 

344 redis.call('ZREMRANGEBYSCORE', KEYS[i], '-inf', now - tonumber(ARGV[i])) 

345 table.insert(results, redis.call('ZCARD', KEYS[i])) 

346end 

347return results 

348""" 

349 

350TOKEN_INCREMENT_SCRIPT: Final = """ 

351local results = {} 

352 

353-- Process each key/increment_value/ttl triplet 

354for i = 1, #KEYS do 

355 local key = KEYS[i] 

356 local increment_value = tonumber(ARGV[i * 2 - 1]) 

357 local ttl_seconds = tonumber(ARGV[i * 2]) 

358 

359 -- Increment the value 

360 local new_value = redis.call('INCRBYFLOAT', key, increment_value) 

361 

362 -- Handle TTL: only set expire if ttl_seconds > 0 and key has no current TTL 

363 -- ttl_seconds can be 0 (no TTL) or positive (set TTL) 

364 if ttl_seconds and ttl_seconds > 0 then 

365 local current_ttl = redis.call('TTL', key) 

366 if current_ttl == -1 then 

367 redis.call('EXPIRE', key, ttl_seconds) 

368 end 

369 end 

370 

371 table.insert(results, new_value) 

372end 

373 

374return results 

375""" 

376 

377# Redis cluster slot count 

378REDIS_CLUSTER_SLOTS: Final = 16384 

379REDIS_NODE_HASHTAG_NAME: Final = "all_keys" 

380 

381# TPM token reservation tuning constants. 

382# When max_tokens is not specified in the request we still need to reserve 

383# *some* output budget; these define that fallback estimate. 

384DEFAULT_MAX_TOKENS_ESTIMATE: Final = 4096 

385DEFAULT_CHARS_PER_TOKEN: Final = 4 

386# Fraction of the available output budget reserved as the upfront floor when 

387# the request omits max_tokens. Applied to both DEFAULT_MAX_TOKENS_ESTIMATE 

388# (baseline floor) and to the smallest configured TPM limit (capped floor for 

389# small per-tenant TPM caps). 

390_TPM_FLOOR_FRACTION: Final = 4 

391# Both embeddings and the Responses API put their prompt in data["input"], 

392# but only embeddings have no output tokens. Every "is this an embedding" 

393# check on data["input"] must exclude these call types, or a Responses call 

394# gets misclassified as an embedding and skips output-token reservation/caps. 

395RESPONSES_API_CALL_TYPES: Final = ("aresponses", "responses") 

396EMBEDDING_API_CALL_TYPES: Final = ("aembedding", "embedding") 

397TEXT_COMPLETION_API_CALL_TYPES: Final = ("atext_completion", "text_completion") 

398RERANK_API_CALL_TYPES: Final = (CallTypes.rerank.value, CallTypes.arerank.value) 

399GOOGLE_GENAI_NATIVE_CALL_TYPES: Final = ( 

400 CallTypes.generate_content.value, 

401 CallTypes.agenerate_content.value, 

402 CallTypes.generate_content_stream.value, 

403 CallTypes.agenerate_content_stream.value, 

404) 

405RESPONSES_API_MIN_OUTPUT_TOKENS: Final = 16 

406# litellm.token_counter has no per-type handling for "input_audio" content 

407# blocks (unlike images, which use use_default_image_token_count) -- it 

408# silently contributes 0 tokens for them. When the block carries a base64 

409# payload, the estimate is derived from the decoded byte count; when the 

410# block is a reference without a payload (or the payload is missing), this 

411# flat per-block floor is used instead. 

412DEFAULT_AUDIO_TOKEN_ESTIMATE: Final = 300 

413# Conservative bytes-per-token assumption for size-based audio estimation: 

414# equivalent to 8 kHz mono PCM-16 (16 000 bytes/s) at 10 tokens/s. Choosing 

415# the lowest reasonable bitrate means we never under-reserve for higher- 

416# quality audio recorded at the same wall-clock duration. 

417_AUDIO_BYTES_PER_TOKEN: Final = 1600 

418# Descriptor "key" values for project-scoped ITPM/OTPM. Distinct from 

419# "model_per_project" (the combined-TPM descriptor) so both can be enforced 

420# on the same project+model simultaneously without colliding on cache keys. 

421PROJECT_ITPM_DESCRIPTOR_KEY: Final = "model_per_project_itpm" 

422PROJECT_OTPM_DESCRIPTOR_KEY: Final = "model_per_project_otpm" 

423# How long an acquired slot counts toward the in-flight total before it is 

424# considered leaked (worker crashed without any release callback firing) and 

425# pruned. Also the longest request duration the gauge can track: a request 

426# running longer than this stops occupying its slot. 

427PARALLEL_REQUEST_SLOT_TTL_SECONDS: Final = 3600 

428 

429 

430CacheCounterValue: TypeAlias = int | float | str | bytes 

431 

432CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None] 

433 

434ReservationWindowIdentity: TypeAlias = tuple[str, str, Literal["redis", "local"]] 

435 

436ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes 

437 

438 

439class _AsyncLuaScript(Protocol): 

440 """A Lua script registered against the async Redis client, called with KEYS and ARGV.""" 

441 

442 def __call__(self, *, keys: Sequence[str], args: Sequence[object]) -> Awaitable[list[CacheCounterValue]]: ... 442 ↛ exitline 442 didn't return from function '__call__' because

443 

444 

445class RateLimitDescriptorRateLimitObject(TypedDict, total=False): 

446 requests_per_unit: int | None 

447 tokens_per_unit: int | None 

448 max_parallel_requests: int | None 

449 window_size: int | None 

450 

451 

452class RateLimitDescriptor(TypedDict): 

453 key: str 

454 value: str 

455 rate_limit: RateLimitDescriptorRateLimitObject | None 

456 

457 

458class ParallelRequestGauge(TypedDict): 

459 counter_key: str 

460 limit: int 

461 descriptor_key: str 

462 

463 

464class ParallelSlotAcquisition(TypedDict): 

465 slot_id: str 

466 counter_keys: list[str] 

467 

468 

469class RateLimitStatus(TypedDict): 

470 code: str 

471 current_limit: int 

472 limit_remaining: int 

473 rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"] 

474 descriptor_key: str 

475 # Populated by the atomic_check_and_increment_by_n and windowed 

476 # sliding-window paths. A caller matching a status back to its 

477 # descriptor must key on (descriptor_key, descriptor_value) when this 

478 # is present, not descriptor_key alone -- e.g. a batch charging several 

479 # models' project ITPM/OTPM in one call, or a request carrying multiple 

480 # rate-limited tags, produces statuses sharing the same descriptor_key. 

481 descriptor_value: NotRequired[ReadOnly[str]] 

482 

483 

484class RateLimitResponse(TypedDict): 

485 overall_code: str 

486 statuses: list[RateLimitStatus] 

487 reservation_windows: NotRequired[ReadOnly[frozenset[tuple[str, str, Literal["redis", "local"]]]]] 

488 

489 

490class ReservationAwareIncrementOperation(RedisPipelineIncrementOperation): 

491 window_key: NotRequired[str] 

492 expected_window_start: NotRequired[str] 

493 reservation_backend: NotRequired[Literal["redis", "local"]] 

494 

495 

496class RateLimitResponseWithDescriptors(TypedDict): 

497 descriptors: list[RateLimitDescriptor] 

498 response: RateLimitResponse 

499 

500 

501class _RateLimitDescriptorSink(Protocol): 

502 def append(self, descriptor: RateLimitDescriptor, /) -> None: ... 502 ↛ exitline 502 didn't return from function 'append' because

503 

504 

505class WindowKeyMetadata(TypedDict): 

506 requests_limit: int | None 

507 tokens_limit: int | None 

508 window_size: int 

509 descriptor_key: str 

510 descriptor_value: ReadOnly[str] 

511 

512 

513class AtomicCounterMeta(TypedDict): 

514 descriptor_key: str 

515 descriptor_value: ReadOnly[str] 

516 current_limit: int 

517 rate_limit_type: Literal["requests", "tokens"] 

518 window_key: str 

519 counter_key: str 

520 increment: int 

521 ttl: int 

522 window_size: int 

523 

524 

525class AtomicCounterState(TypedDict): 

526 window_expired: bool 

527 current: int 

528 window_start: ReadOnly[str] 

529 

530 

531DescriptorAtomicGroup: TypeAlias = tuple[list[str], list[int], list[AtomicCounterMeta]] 

532 

533 

534class CallTypeRateLimiter(Protocol): 

535 async def async_pre_call_hook( 535 ↛ exitline 535 didn't return from function 'async_pre_call_hook' because

536 self, 

537 user_api_key_dict: UserAPIKeyAuth, 

538 cache: DualCache, 

539 data: dict[str, object], 

540 call_type: str, 

541 ) -> Exception | str | dict[str, object] | None: ... 

542 

543 

544@dataclass(slots=True) 

545class RequestRateLimiterStash: 

546 """ 

547 Per-request bookkeeping the pre-call hook hands to the success/failure/ 

548 disconnect callbacks. Lives on a ContextVar instead of the request body so 

549 it never reaches provider-facing ``metadata`` channels. 

550 

551 A single mutable instance is shared by every context forked from the 

552 request task (the SDK call, streaming generators, and the logging worker's 

553 captured context all see the same object), which is what makes the 

554 ``reservation_released`` flag and ``parallel_slot`` clearing effective 

555 across sibling callbacks: the first release wins, later callbacks observe 

556 the cleared state. 

557 

558 Because the stash is context-inherited, nested LiteLLM calls made inside 

559 the request (LLM-judge guardrails, silent experiments) would also see it 

560 from their own logging callbacks. ``owner_litellm_call_id`` pins the stash 

561 to the proxy request's ``litellm_call_id`` so those callbacks can tell the 

562 owning request's events apart from a nested call's: router retries and 

563 fallbacks reuse the request's call id and keep access, while nested calls 

564 mint fresh ids and are ignored. 

565 """ 

566 

567 owner_litellm_call_id: str | None = None 

568 rate_limit_response: RateLimitResponse | None = None 

569 parallel_slot: ParallelSlotAcquisition | None = None 

570 parallel_slot_release_lock: asyncio.Lock = field(default_factory=asyncio.Lock, repr=False, compare=False) 

571 reserved_tokens: int = 0 

572 reserved_model: RateLimitedModel | None = None 

573 reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) 

574 itpm_reserved_tokens: int = 0 

575 itpm_reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) 

576 itpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field( 

577 default_factory=frozenset 

578 ) 

579 otpm_reserved_tokens: int = 0 

580 otpm_reserved_scopes: frozenset[tuple[str, str]] = field(default_factory=frozenset) 

581 otpm_reserved_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]] = field( 

582 default_factory=frozenset 

583 ) 

584 batch_enqueued_reservation: BatchEnqueuedTokenReservation | None = None 

585 batch_tpd_refund_ops: tuple[ReservationAwareIncrementOperation, ...] = () 

586 reservation_released: bool = False 

587 tpm_limited_tags: frozenset[str] = field(default_factory=frozenset) 

588 

589 

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

591class TagRateLimit: 

592 rpm_limit: int | None 

593 tpm_limit: int | None 

594 

595 

596class TagRateLimitResolver(Protocol): 

597 def __call__(self, tag_names: Sequence[str], /) -> Awaitable[Mapping[str, TagRateLimit]]: ... 597 ↛ exitline 597 didn't return from function '__call__' because

598 

599 

600def _tag_rate_limit_descriptor(tag: str, limit: TagRateLimit, window_size: int) -> RateLimitDescriptor: 

601 rate_limit: Final[RateLimitDescriptorRateLimitObject] = { 

602 "requests_per_unit": limit.rpm_limit, 

603 "tokens_per_unit": limit.tpm_limit, 

604 "window_size": window_size, 

605 } 

606 return RateLimitDescriptor(key="tag", value=tag, rate_limit=rate_limit) 

607 

608 

609async def resolve_tag_rate_limits_from_db(tag_names: Sequence[str]) -> Mapping[str, TagRateLimit]: 

610 from litellm.proxy.auth.auth_checks import get_tag_objects_batch 

611 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache 

612 

613 if prisma_client is None or not tag_names: 

614 return MappingProxyType({}) 

615 tag_objects: Final = await get_tag_objects_batch( 

616 tag_names=tag_names, 

617 prisma_client=prisma_client, 

618 user_api_key_cache=user_api_key_cache, 

619 ) 

620 return MappingProxyType( 

621 { 

622 tag_name: TagRateLimit(rpm_limit=budget.rpm_limit, tpm_limit=budget.tpm_limit) 

623 for tag_name, tag_object in tag_objects.items() 

624 if (budget := tag_object.litellm_budget_table) is not None 

625 and (budget.rpm_limit is not None or budget.tpm_limit is not None) 

626 } 

627 ) 

628 

629 

630_request_stash: Final[ContextVar[RequestRateLimiterStash | None]] = ContextVar( 

631 "litellm_v3_rate_limiter_request_stash", default=None 

632) 

633 

634 

635def get_request_stash() -> RequestRateLimiterStash | None: 

636 return _request_stash.get() 

637 

638 

639def get_or_create_request_stash() -> RequestRateLimiterStash: 

640 stash = _request_stash.get() 

641 if stash is None: 641 ↛ 644line 641 didn't jump to line 644 because the condition on line 641 was always true

642 stash = RequestRateLimiterStash() 

643 _request_stash.set(stash) 

644 return stash 

645 

646 

647def claim_request_stash_for_data(data: dict) -> RequestRateLimiterStash: 

648 stash: Final = get_or_create_request_stash() 

649 owner_call_id: Final = data.get("litellm_call_id") 

650 if isinstance(owner_call_id, str): 

651 stash.owner_litellm_call_id = owner_call_id 

652 return stash 

653 

654 

655def get_request_stash_for_call(litellm_call_id: str | None) -> RequestRateLimiterStash | None: 

656 stash: Final = _request_stash.get() 

657 if stash is None: 

658 return None 

659 if stash.owner_litellm_call_id is None or litellm_call_id is None: 

660 return stash 

661 return stash if litellm_call_id == stash.owner_litellm_call_id else None 

662 

663 

664def _call_id_from_callback_kwargs(kwargs: object) -> str | None: 

665 if not isinstance(kwargs, dict): 665 ↛ 666line 665 didn't jump to line 666 because the condition on line 665 was never true

666 return None 

667 call_id: Final = kwargs.get("litellm_call_id") 

668 return call_id if isinstance(call_id, str) else None 

669 

670 

671def _parse_output_cap_value(raw_value: object) -> int | None: 

672 if isinstance(raw_value, bool) or not isinstance(raw_value, (int, float, str)): 

673 return None 

674 try: 

675 return int(float(raw_value)) 

676 except (ValueError, OverflowError): 

677 return None 

678 

679 

680class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): 

681 batch_rate_limiter_script: _AsyncLuaScript | None 

682 token_increment_script: _AsyncLuaScript | None 

683 check_and_increment_by_n_script: _AsyncLuaScript | None 

684 window_guarded_token_increment_script: _AsyncLuaScript | None 

685 parallel_acquire_script: _AsyncLuaScript | None 

686 parallel_release_script: _AsyncLuaScript | None 

687 parallel_count_script: _AsyncLuaScript | None 

688 

689 def __init__( 

690 self, 

691 internal_usage_cache: InternalUsageCache, 

692 time_provider: Callable[[], datetime] | None = None, 

693 tag_rate_limit_resolver: TagRateLimitResolver = resolve_tag_rate_limits_from_db, 

694 model_group_resolver: Callable[[str], str | None] = _resolve_model_group_alias_via_proxy_router, 

695 ): 

696 self.internal_usage_cache = internal_usage_cache 

697 self._time_provider = time_provider or datetime.now 

698 self._tag_rate_limit_resolver = tag_rate_limit_resolver 

699 self._model_group_resolver = model_group_resolver 

700 if self.internal_usage_cache.dual_cache.redis_cache is not None: 700 ↛ 701line 700 didn't jump to line 701 because the condition on line 700 was never true

701 self.batch_rate_limiter_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

702 BATCH_RATE_LIMITER_SCRIPT 

703 ) 

704 self.token_increment_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

705 TOKEN_INCREMENT_SCRIPT 

706 ) 

707 self.check_and_increment_by_n_script = ( 

708 self.internal_usage_cache.dual_cache.redis_cache.async_register_script(CHECK_AND_INCREMENT_BY_N_SCRIPT) 

709 ) 

710 self.window_guarded_token_increment_script = ( 

711 self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

712 WINDOW_GUARDED_TOKEN_INCREMENT_SCRIPT 

713 ) 

714 ) 

715 self.parallel_acquire_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

716 PARALLEL_ACQUIRE_SCRIPT 

717 ) 

718 self.parallel_release_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

719 PARALLEL_RELEASE_SCRIPT 

720 ) 

721 self.parallel_count_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( 

722 PARALLEL_COUNT_SCRIPT 

723 ) 

724 else: 

725 self.batch_rate_limiter_script = None 

726 self.token_increment_script = None 

727 self.check_and_increment_by_n_script = None 

728 self.window_guarded_token_increment_script = None 

729 self.parallel_acquire_script = None 

730 self.parallel_release_script = None 

731 self.parallel_count_script = None 

732 

733 self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) 

734 

735 # When disabled, TPM is enforced post-call from actual usage (pre-v1.82 

736 # behavior) instead of reserving an estimated budget upfront, shedding 

737 # the extra per-request Redis Lua round-trip and the global-lock 

738 # in-memory fallback that the reservation path incurs. 

739 self.tpm_reservation_enabled = os.getenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "true").lower() == "true" 

740 

741 # Batch rate limiter (lazy loaded) 

742 self._batch_rate_limiter: CallTypeRateLimiter | None = None 

743 self.batch_enqueued_token_store = BatchEnqueuedTokenStore(internal_usage_cache=internal_usage_cache) 

744 

745 # Serializes multi-phase check+increment sequences (batch + dynamic 

746 # limiters) within this process to close the TOCTOU window between 

747 # read-only check and counter increment. Multi-replica deployments 

748 # additionally rely on Redis Lua atomicity for cross-process safety. 

749 # 

750 # Coarse granularity: this single lock serializes ALL atomic check+ 

751 # increment operations across batch and dynamic limiters on this 

752 # instance. A slow batch input-file fetch (which happens upstream of 

753 # the lock) does not block here, but Redis Lua latency does. If 

754 # contention shows up under load (visible as p99 latency spikes 

755 # correlated with batch traffic), shard to a per-descriptor-key lock 

756 # via a `weakref.WeakValueDictionary[str, asyncio.Lock]`. Punted as a 

757 # follow-up because Lua dominates wall-time and the lock is held for 

758 # one round-trip. 

759 self._check_and_increment_lock = asyncio.Lock() 

760 

761 def _get_batch_rate_limiter(self) -> CallTypeRateLimiter | None: 

762 """Get or lazy-load the batch rate limiter.""" 

763 if self._batch_rate_limiter is None: 

764 try: 

765 from litellm.proxy.hooks.batch_rate_limiter import ( 

766 _PROXY_BatchRateLimiter, 

767 ) 

768 

769 self._batch_rate_limiter = _PROXY_BatchRateLimiter( 

770 internal_usage_cache=self.internal_usage_cache, 

771 parallel_request_limiter=self, 

772 time_provider=self._time_provider, 

773 ) 

774 except Exception as e: 

775 verbose_proxy_logger.debug("Could not load batch rate limiter: %s", e) 

776 return self._batch_rate_limiter 

777 

778 def _get_current_time(self) -> datetime: 

779 """Return the current time for rate limiting calculations.""" 

780 return self._time_provider() 

781 

782 @staticmethod 

783 def no_max_tokens_output_floor( 

784 min_configured_tpm_limit: int | None, 

785 ) -> int: 

786 """Output-budget floor used when the request omits max_tokens. 

787 

788 Capped at a fraction of the smallest configured TPM limit so a small 

789 per-tenant cap can't be tripped by the floor alone. Returns the 

790 baseline floor when no limit is provided. 

791 """ 

792 baseline: Final = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION 

793 if min_configured_tpm_limit is None: 

794 return baseline 

795 return min(baseline, max(1, min_configured_tpm_limit // _TPM_FLOOR_FRACTION)) 

796 

797 @staticmethod 

798 def _is_embedding_request(data: object, call_type: str | None) -> bool: 

799 if call_type in EMBEDDING_API_CALL_TYPES: 

800 return True 

801 if call_type in RESPONSES_API_CALL_TYPES: 

802 return False 

803 if call_type: 

804 return False 

805 if not isinstance(data, dict): 

806 return False 

807 return data.get("input") is not None 

808 

809 @staticmethod 

810 def _translate_google_genai_native_request( 

811 data: object, 

812 call_type: str | None, 

813 ) -> Mapping[str, object] | None: 

814 contents: Final = data.get("contents") if isinstance(data, dict) else None 

815 if ( 

816 not isinstance(data, dict) 

817 or call_type not in GOOGLE_GENAI_NATIVE_CALL_TYPES 

818 or not isinstance(contents, (dict, list)) 

819 ): 

820 return None 

821 from litellm.google_genai.adapters.transformation import GoogleGenAIAdapter 

822 

823 config: Final = data.get("config") if "config" in data else data.get("generationConfig") 

824 return GoogleGenAIAdapter().translate_generate_content_to_completion( 

825 model=data.get("model") if isinstance(data.get("model"), str) else "", 

826 contents=contents, 

827 config=config if isinstance(config, dict) else None, 

828 systemInstruction=data.get("systemInstruction"), 

829 system_instruction=data.get("system_instruction"), 

830 tools=data.get("tools"), 

831 toolConfig=data.get("toolConfig"), 

832 tool_config=data.get("tool_config"), 

833 ) 

834 

835 @staticmethod 

836 def _get_explicit_output_cap(data: object, call_type: str | None) -> int | None: 

837 if not isinstance(data, dict): 

838 return None 

839 if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES: 

840 config: Final = data.get("config") if "config" in data else data.get("generationConfig") 

841 google_cap_values: Final = tuple( 

842 parsed 

843 for field in ("maxOutputTokens", "max_output_tokens") 

844 if isinstance(config, dict) 

845 for parsed in (_parse_output_cap_value(config.get(field)),) 

846 if parsed is not None 

847 ) 

848 return max(google_cap_values, default=None) 

849 if call_type in RESPONSES_API_CALL_TYPES: 

850 responses_cap: Final = _parse_output_cap_value(data.get("max_output_tokens")) 

851 if responses_cap is None: 

852 return None 

853 return max(RESPONSES_API_MIN_OUTPUT_TOKENS, responses_cap) 

854 if call_type in EMBEDDING_API_CALL_TYPES: 

855 return None 

856 fields: Final = ( 

857 ("max_tokens", "max_completion_tokens") 

858 if call_type 

859 else ("max_tokens", "max_completion_tokens", "max_output_tokens") 

860 ) 

861 output_cap_values: Final = tuple( 

862 parsed for field in fields for parsed in (_parse_output_cap_value(data.get(field)),) if parsed is not None 

863 ) 

864 return max(output_cap_values, default=None) 

865 

866 @classmethod 

867 def _has_explicit_output_cap(cls, data: object, call_type: str | None) -> bool: 

868 """Whether the caller explicitly set an output-token cap. 

869 

870 Checked via ``is not None`` (not truthiness) so an explicit 0 -- 

871 a legitimate zero-output request -- counts as explicit. 

872 """ 

873 return cls._get_explicit_output_cap(data, call_type) is not None 

874 

875 @staticmethod 

876 def get_output_candidate_count(data: object, call_type: str | None = None) -> int: 

877 if not isinstance(data, Mapping): 

878 return 1 

879 config: Final = ( 

880 (data.get("config") if "config" in data else data.get("generationConfig")) 

881 if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES 

882 else None 

883 ) 

884 candidate_values: Final = ( 

885 data.get("n"), 

886 data.get("best_of"), 

887 config.get("candidateCount") if isinstance(config, dict) else None, 

888 config.get("candidate_count") if isinstance(config, dict) else None, 

889 ) 

890 candidate_count = 1 # rebind-ok: running maximum across candidate-count aliases 

891 for value in candidate_values: 

892 try: 

893 candidate_count = max(candidate_count, int(value or 1)) 

894 except (TypeError, ValueError, OverflowError): 

895 continue 

896 return candidate_count 

897 

898 @staticmethod 

899 def _apply_implicit_output_cap( 

900 data: object, 

901 min_configured_limit: int | None, 

902 call_type: str | None, 

903 configured_output_tokens: int | None = None, 

904 ) -> None: 

905 """Hard-cap generation length when the request has no explicit cap. 

906 

907 Guards against an unbounded response overshooting a small TPM/OTPM 

908 budget before post-call reconciliation runs. Skips requests that 

909 already set an explicit cap and embeddings, which have no generation 

910 budget. The Responses API only honors ``max_output_tokens`` (its 

911 underlying chat-completion transformation ignores ``max_tokens``), so 

912 the cap must be written to that field for Responses call types. 

913 

914 ``configured_output_tokens`` is the operator-declared per-tenant 

915 estimate; when it exceeds the safety floor, the cap is raised to that 

916 value instead of clamping every tenant to the same floor. 

917 """ 

918 if not isinstance(data, dict): 

919 return 

920 base_capped_floor: Final = _PROXY_MaxParallelRequestsHandler_v3.no_max_tokens_output_floor(min_configured_limit) 

921 capped_floor: Final = ( 

922 max(base_capped_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) 

923 if call_type in RESPONSES_API_CALL_TYPES 

924 else base_capped_floor 

925 ) 

926 baseline_floor: Final = DEFAULT_MAX_TOKENS_ESTIMATE // _TPM_FLOOR_FRACTION 

927 is_embedding: Final = _PROXY_MaxParallelRequestsHandler_v3._is_embedding_request(data, call_type) 

928 if ( 

929 capped_floor >= baseline_floor 

930 or _PROXY_MaxParallelRequestsHandler_v3._has_explicit_output_cap(data, call_type) 

931 or is_embedding 

932 ): 

933 return 

934 effective_cap: Final = max(capped_floor, configured_output_tokens or 0) 

935 if call_type in GOOGLE_GENAI_NATIVE_CALL_TYPES: 

936 config_field: Final = "config" if "config" in data or "generationConfig" not in data else "generationConfig" 

937 config: Final = data.get(config_field) 

938 if config is None or isinstance(config, dict): 

939 data[config_field] = { # rebind-ok: routed request needs cap # mutable-ok: downstream needs dict 

940 **(config or {}), # mutable-ok: downstream native routing requires a mutable request config 

941 "maxOutputTokens": effective_cap, 

942 } 

943 return 

944 cap_field: Final = "max_output_tokens" if call_type in RESPONSES_API_CALL_TYPES else "max_tokens" 

945 existing_cap: Final = data.get(cap_field) 

946 if existing_cap is None or effective_cap < existing_cap: 

947 data[cap_field] = effective_cap # rebind-ok: downstream routing requires the bounded output cap 

948 

949 def _estimate_tokens_for_request( 

950 self, 

951 data: dict, 

952 model: str | None = None, 

953 min_configured_tpm_limit: int | None = None, 

954 call_type: str | None = None, 

955 configured_output_tokens: int | None = None, 

956 ) -> int: 

957 """ 

958 Estimate total tokens this request will consume so we can reserve them 

959 upfront (input + output budget): 

960 estimated = input_tokens + max_tokens. 

961 

962 Supports chat (messages), completions (prompt), embeddings (input), 

963 and the Responses API (also `input`, disambiguated from embeddings 

964 via ``call_type``). 

965 

966 ``min_configured_tpm_limit`` is the smallest ``tokens_per_unit`` among 

967 the TPM-bearing descriptors this request will be charged against. When 

968 provided, the no-``max_tokens`` output-budget floor is capped at a 

969 fraction of that limit so small TPM caps remain usable. Omit to 

970 preserve the unconstrained floor. 

971 

972 ``configured_output_tokens`` is the operator-declared estimate resolved 

973 from key or team metadata. When provided it replaces the heuristic 

974 floor entirely, so the reservation reflects what this tenant's model 

975 actually emits rather than one constant shared by every tenant. 

976 """ 

977 estimated_input_tokens, max_tokens_estimate = self._estimate_input_and_output_tokens( 

978 data=data, 

979 min_configured_tpm_limit=min_configured_tpm_limit, 

980 call_type=call_type, 

981 configured_output_tokens=configured_output_tokens, 

982 ) 

983 total_estimated: Final = estimated_input_tokens + max_tokens_estimate 

984 

985 verbose_proxy_logger.debug( 

986 "TPM reservation estimate: input=%s, max_tokens=%s, total=%s", 

987 estimated_input_tokens, 

988 max_tokens_estimate, 

989 total_estimated, 

990 ) 

991 

992 return total_estimated 

993 

994 def _estimate_input_and_output_tokens( 

995 self, 

996 data: object, 

997 min_configured_tpm_limit: int | None = None, 

998 call_type: str | None = None, 

999 configured_output_tokens: int | None = None, 

1000 ) -> tuple[int, int]: 

1001 """ 

1002 Estimate input tokens and output (max_tokens) budget separately, so 

1003 callers needing independent ITPM/OTPM reservations (rather than one 

1004 combined TPM reservation) can use each half on its own. 

1005 

1006 ``min_configured_tpm_limit`` is the smallest ``tokens_per_unit`` among 

1007 the TPM-bearing descriptors this request will be charged against. When 

1008 provided, the no-``max_tokens`` output-budget floor is capped at a 

1009 fraction of that limit so small TPM caps remain usable. Omit to 

1010 preserve the unconstrained floor. 

1011 

1012 ``call_type`` disambiguates embeddings from the Responses API: both 

1013 put their prompt in ``data["input"]``, but only embeddings have no 

1014 output tokens. Unset (the default) preserves the historical 

1015 "any `input` means zero output" behavior for callers that don't have 

1016 a call type to pass. 

1017 

1018 ``configured_output_tokens`` is the operator-declared estimate resolved 

1019 from key or team metadata. When provided it replaces the heuristic 

1020 floor entirely, so the reservation reflects what this tenant's model 

1021 actually emits rather than one constant shared by every tenant. 

1022 """ 

1023 if not isinstance(data, dict): 

1024 return 0, 0 

1025 translated_data: Final = self._translate_google_genai_native_request(data, call_type) 

1026 estimable_data: Final = translated_data if translated_data is not None else data 

1027 selected_fields: Final[tuple[object | None, object | None, object | None]] = ( 

1028 (None, None, estimable_data.get("input")) 

1029 if call_type in RESPONSES_API_CALL_TYPES or call_type in EMBEDDING_API_CALL_TYPES 

1030 else (None, estimable_data.get("prompt"), None) 

1031 if call_type in TEXT_COMPLETION_API_CALL_TYPES 

1032 else (estimable_data.get("messages"), None, None) 

1033 if call_type 

1034 else ( 

1035 estimable_data.get("messages"), 

1036 estimable_data.get("prompt"), 

1037 estimable_data.get("input"), 

1038 ) 

1039 ) 

1040 messages, prompt, input_text = selected_fields 

1041 

1042 total_chars: Final = ( 

1043 len(get_str_from_messages(messages)) 

1044 if isinstance(messages, list) and messages 

1045 else len(prompt) 

1046 if isinstance(prompt, str) 

1047 else sum(len(str(item)) for item in prompt) 

1048 if isinstance(prompt, list) 

1049 else len(input_text) 

1050 if isinstance(input_text, str) 

1051 else sum(len(str(item)) for item in input_text) 

1052 if isinstance(input_text, list) 

1053 else 0 

1054 ) 

1055 

1056 estimated_input_tokens: Final = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) if total_chars > 0 else 0 

1057 

1058 explicit_max_tokens: Final = self._get_explicit_output_cap(data, call_type) 

1059 is_embedding: Final = self._is_embedding_request(data, call_type) 

1060 

1061 base_output_floor: Final = self.no_max_tokens_output_floor(min_configured_tpm_limit) 

1062 output_floor: Final = ( 

1063 max(base_output_floor, RESPONSES_API_MIN_OUTPUT_TOKENS) 

1064 if call_type in RESPONSES_API_CALL_TYPES 

1065 else base_output_floor 

1066 ) 

1067 max_tokens_estimate: Final = ( 

1068 0 

1069 if is_embedding or (explicit_max_tokens is None and total_chars == 0 and configured_output_tokens is None) 

1070 else explicit_max_tokens 

1071 if explicit_max_tokens is not None 

1072 else configured_output_tokens 

1073 if configured_output_tokens is not None 

1074 else max(estimated_input_tokens, output_floor) 

1075 ) 

1076 

1077 return estimated_input_tokens, max_tokens_estimate * self.get_output_candidate_count(data, call_type) 

1078 

1079 def _is_redis_cluster(self) -> bool: 

1080 """ 

1081 Check if the dual cache is using Redis cluster. 

1082 

1083 Returns: 

1084 bool: True if using Redis cluster, False otherwise. 

1085 """ 

1086 from litellm.caching.redis_cluster_cache import RedisClusterCache 

1087 

1088 return self.internal_usage_cache.dual_cache.redis_cache is not None and isinstance( 

1089 self.internal_usage_cache.dual_cache.redis_cache, RedisClusterCache 

1090 ) 

1091 

1092 async def in_memory_cache_sliding_window( 

1093 self, 

1094 keys: list[str], 

1095 now_int: int, 

1096 window_size: int, 

1097 ) -> CacheCounterValues: 

1098 """ 

1099 Implement sliding window rate limiting logic using in-memory cache operations. 

1100 This follows the same logic as the Redis Lua script but uses async cache operations. 

1101 """ 

1102 async with self._check_and_increment_lock: 

1103 return await self._in_memory_cache_sliding_window(keys=keys, now_int=now_int, window_size=window_size) 

1104 

1105 async def _in_memory_cache_sliding_window( 

1106 self, 

1107 keys: list[str], 

1108 now_int: int, 

1109 window_size: int, 

1110 ) -> CacheCounterValues: 

1111 results: Final[list[CacheCounterValue | None]] = [] 

1112 

1113 # Process each window/counter pair 

1114 for i in range(0, len(keys), 2): 

1115 window_key = keys[i] 

1116 counter_key = keys[i + 1] 

1117 increment_value = 1 

1118 

1119 # Get the window start time 

1120 window_start: CacheCounterValue | None = await self.internal_usage_cache.async_get_cache( 

1121 key=window_key, 

1122 litellm_parent_otel_span=None, 

1123 local_only=True, 

1124 ) 

1125 

1126 # Check if window exists and is valid 

1127 if window_start is None or (now_int - int(window_start)) >= window_size: 

1128 # Reset window and counter 

1129 for sibling_counter_key in _sibling_counter_keys(window_key): 

1130 await self.internal_usage_cache.async_set_cache( 

1131 key=sibling_counter_key, 

1132 value=0, 

1133 ttl=window_size, 

1134 litellm_parent_otel_span=None, 

1135 local_only=True, 

1136 ) 

1137 await self.internal_usage_cache.async_set_cache( 

1138 key=window_key, 

1139 value=str(now_int), 

1140 ttl=window_size, 

1141 litellm_parent_otel_span=None, 

1142 local_only=True, 

1143 ) 

1144 await self.internal_usage_cache.async_set_cache( 

1145 key=counter_key, 

1146 value=increment_value, 

1147 ttl=window_size, 

1148 litellm_parent_otel_span=None, 

1149 local_only=True, 

1150 ) 

1151 results.append(str(now_int)) # window_start 

1152 results.append(increment_value) # counter 

1153 else: 

1154 # Increment the counter 

1155 current_counter: CacheCounterValue | None = await self.internal_usage_cache.async_get_cache( 

1156 key=counter_key, 

1157 litellm_parent_otel_span=None, 

1158 local_only=True, 

1159 ) 

1160 new_counter_value = (int(current_counter) if current_counter is not None else 0) + increment_value 

1161 await self.internal_usage_cache.async_set_cache( 

1162 key=counter_key, 

1163 value=new_counter_value, 

1164 ttl=window_size, 

1165 litellm_parent_otel_span=None, 

1166 local_only=True, 

1167 ) 

1168 results.append(window_start) # window_start 

1169 results.append(new_counter_value) # counter 

1170 

1171 return results 

1172 

1173 def create_rate_limit_keys( 

1174 self, 

1175 key: str, 

1176 value: str, 

1177 rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"], 

1178 ) -> str: 

1179 """ 

1180 Create the rate limit keys for the given key and value. 

1181 """ 

1182 counter_key: Final = f"{{{key}:{value}}}:{rate_limit_type}" 

1183 

1184 return counter_key 

1185 

1186 def is_cache_list_over_limit( 

1187 self, 

1188 keys_to_fetch: list[str], 

1189 cache_values: CacheCounterValues, 

1190 key_metadata: dict[str, WindowKeyMetadata], 

1191 ) -> RateLimitResponse: 

1192 """ 

1193 Check if the cache values are over the limit. 

1194 """ 

1195 statuses: Final[list[RateLimitStatus]] = [] 

1196 overall_code = "OK" 

1197 

1198 for i in range(0, len(cache_values), 2): 

1199 item_code = "OK" 

1200 window_key = keys_to_fetch[i] 

1201 counter_key = keys_to_fetch[i + 1] 

1202 counter_value = cache_values[i + 1] 

1203 requests_limit = key_metadata[window_key]["requests_limit"] 

1204 tokens_limit = key_metadata[window_key]["tokens_limit"] 

1205 

1206 # Determine which limit to use for current_limit and limit_remaining 

1207 current_limit: int | None = None 

1208 rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"] | None = None 

1209 if counter_key.endswith(":requests"): 

1210 current_limit = requests_limit 

1211 rate_limit_type = "requests" 

1212 elif counter_key.endswith(":tokens"): 

1213 current_limit = tokens_limit 

1214 rate_limit_type = "tokens" 

1215 

1216 if current_limit is None or rate_limit_type is None: 

1217 continue 

1218 

1219 if counter_value is not None and int(counter_value) > current_limit: 

1220 overall_code = "OVER_LIMIT" 

1221 item_code = "OVER_LIMIT" 

1222 

1223 # Only compute limit_remaining if current_limit is not None 

1224 limit_remaining = current_limit - int(counter_value) if counter_value is not None else current_limit 

1225 

1226 statuses.append( 

1227 { 

1228 "code": item_code, 

1229 "current_limit": current_limit, 

1230 "limit_remaining": limit_remaining, 

1231 "rate_limit_type": rate_limit_type, 

1232 "descriptor_key": key_metadata[window_key]["descriptor_key"], 

1233 "descriptor_value": key_metadata[window_key]["descriptor_value"], 

1234 } 

1235 ) 

1236 

1237 return RateLimitResponse(overall_code=overall_code, statuses=statuses) 

1238 

1239 def keyslot_for_redis_cluster(self, key: str) -> int: 

1240 """ 

1241 Compute the Redis Cluster slot for a given key. 

1242 

1243 Simple implementation of `HASH_SLOT = CRC16(key) mod 16384` 

1244 

1245 Read more about hash slots here: https://medium.com/@linz07m/how-hash-slots-power-data-distribution-in-redis-cluster-bc5b7e74ca7d 

1246 

1247 Args: 

1248 key (str): The Redis key. 

1249 

1250 Returns: 

1251 int: The slot number (0-16383). 

1252 

1253 

1254 """ 

1255 # Handle hash tags: use substring between { and } 

1256 start: Final = key.find("{") 

1257 if start != -1: 

1258 end: Final = key.find("}", start + 1) 

1259 if end != -1 and end != start + 1: 

1260 key = key[start + 1 : end] 

1261 

1262 # Compute CRC16 and mod 16384 

1263 crc: Final = binascii.crc_hqx(key.encode("utf-8"), 0) 

1264 return crc % REDIS_CLUSTER_SLOTS 

1265 

1266 def _group_keys_by_hash_tag(self, keys: list[str]) -> dict[str, list[str]]: 

1267 """ 

1268 Group keys by their Redis hash tag to ensure cluster compatibility. 

1269 

1270 For Redis clusters, uses slot calculation to group keys that belong to the same slot. 

1271 For regular Redis, no grouping is needed - all keys can be processed together. 

1272 """ 

1273 groups: Final[dict[str, list[str]]] = {} 

1274 

1275 # Use slot calculation for Redis clusters only 

1276 if self._is_redis_cluster(): 

1277 for key in keys: 

1278 slot = self.keyslot_for_redis_cluster(key) 

1279 slot_key = f"slot_{slot}" 

1280 

1281 if slot_key not in groups: 

1282 groups[slot_key] = [] 

1283 groups[slot_key].append(key) 

1284 else: 

1285 # For regular Redis, no grouping needed - process all keys together 

1286 groups[REDIS_NODE_HASHTAG_NAME] = keys 

1287 

1288 return groups 

1289 

1290 async def _batch_get_counter_values( 

1291 self, 

1292 keys: list[str], 

1293 parent_otel_span: Span | None, 

1294 local_only: bool, 

1295 ) -> CacheCounterValues | None: 

1296 """Typed view over the DualCache batch read of window/counter keys.""" 

1297 return await self.internal_usage_cache.async_batch_get_cache( 

1298 keys=keys, 

1299 parent_otel_span=parent_otel_span, 

1300 local_only=local_only, 

1301 ) 

1302 

1303 async def _batch_get_gauge_values( 

1304 self, 

1305 keys: list[str], 

1306 parent_otel_span: Span | None, 

1307 ) -> Sequence[ParallelGaugeCacheValue | None] | None: 

1308 """Typed view over the DualCache batch read of parallel-request gauges.""" 

1309 return await self.internal_usage_cache.async_batch_get_cache( 

1310 keys=keys, 

1311 parent_otel_span=parent_otel_span, 

1312 local_only=True, 

1313 ) 

1314 

1315 async def _execute_redis_batch_rate_limiter_script( 

1316 self, 

1317 keys_to_fetch: list[str], 

1318 now_int: int, 

1319 ) -> CacheCounterValues: 

1320 """ 

1321 Execute Redis operations grouped by hash tag for cluster compatibility. 

1322 

1323 Args: 

1324 keys_to_fetch: List[str] - List of keys to fetch 

1325 now_int: int - Current timestamp 

1326 

1327 Returns: 

1328 List of cache values 

1329 """ 

1330 if self.batch_rate_limiter_script is None: 

1331 return [] 

1332 

1333 key_groups: Final = self._group_keys_by_hash_tag(keys_to_fetch) 

1334 all_cache_values: Final[list[CacheCounterValue | None]] = [] 

1335 

1336 for hash_tag, group_keys in key_groups.items(): 

1337 try: 

1338 group_cache_values: CacheCounterValues = await self.batch_rate_limiter_script( 

1339 keys=group_keys, 

1340 args=[now_int, self.window_size], # Use integer timestamp 

1341 ) 

1342 all_cache_values.extend(group_cache_values) 

1343 except Exception as e: 

1344 log_redis_failure( 

1345 verbose_proxy_logger, logging.WARNING, f"Redis Lua script failed for hash tag {hash_tag}", e 

1346 ) 

1347 # Fallback to in-memory cache for this group 

1348 group_cache_values = await self.in_memory_cache_sliding_window( 

1349 keys=group_keys, 

1350 now_int=now_int, 

1351 window_size=self.window_size, 

1352 ) 

1353 all_cache_values.extend(group_cache_values) 

1354 

1355 return all_cache_values 

1356 

1357 async def should_rate_limit( 

1358 self, 

1359 descriptors: Sequence[RateLimitDescriptor], 

1360 parent_otel_span: Span | None = None, 

1361 read_only: bool = False, 

1362 skip_tpm_check: bool = False, 

1363 parallel_slot_id: str | None = None, 

1364 ) -> RateLimitResponse: 

1365 """ 

1366 Check if any of the rate limit descriptors should be rate limited. 

1367 Returns a RateLimitResponse with the overall code and status for each descriptor. 

1368 Uses batch operations for Redis to improve performance. 

1369 

1370 Args: 

1371 descriptors: List of rate limit descriptors to check 

1372 parent_otel_span: Optional OpenTelemetry span for tracing 

1373 read_only: If True, only check limits without incrementing counters 

1374 skip_tpm_check: If True, ignore each descriptor's ``tokens_per_unit`` 

1375 — the :tokens counter is neither read nor incremented by this 

1376 pass. Callers that handle TPM via the atomic 

1377 ``reserve_tpm_tokens`` reservation path should set this to 

1378 avoid the +1-per-key Lua / in-memory increment double-charging 

1379 the tokens counter. 

1380 

1381 ``max_parallel_requests`` descriptors are enforced by the dedicated 

1382 concurrency-gauge path (``_check_parallel_request_gauges``), never by 

1383 the windowed counters. The gauge phase must stay AFTER the windowed 

1384 check so a windowed rejection never strands an acquired slot; the 

1385 reverse order would leak one gauge slot per RPM/TPM rejection. 

1386 ``parallel_slot_id`` names the slot an admission registers; callers 

1387 that enforce (not read_only) should pass the id they will later 

1388 release with — when omitted, a generated slot id is used and the slot 

1389 can only be reclaimed by TTL expiry. 

1390 """ 

1391 

1392 current_time: Final = self._get_current_time() 

1393 now: Final = current_time.timestamp() 

1394 now_int: Final = int(now) # Convert to integer for Redis Lua script 

1395 

1396 keys_to_fetch, key_metadata, gauges = self._collect_windowed_keys_and_gauges( 

1397 descriptors=descriptors, 

1398 skip_tpm_check=skip_tpm_check, 

1399 ) 

1400 

1401 windowed_response = RateLimitResponse(overall_code="OK", statuses=[]) 

1402 if keys_to_fetch: 

1403 ## CHECK IN-MEMORY CACHE 

1404 cache_values = await self._batch_get_counter_values( # rebind-ok: refreshed by the Redis read below when the in-memory pass is under limit 

1405 keys=keys_to_fetch, 

1406 parent_otel_span=parent_otel_span, 

1407 local_only=True, 

1408 ) 

1409 

1410 if cache_values is not None: 

1411 rate_limit_response: Final = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata) 

1412 if rate_limit_response["overall_code"] == "OVER_LIMIT": 

1413 return rate_limit_response 

1414 

1415 ## IF under limit in-memory, check Redis 

1416 if read_only: 

1417 # READ-ONLY MODE: Just read current values without incrementing 

1418 cache_values = await self._batch_get_counter_values( # rebind-ok: read-only mode replaces the in-memory snapshot with Redis values 

1419 keys=keys_to_fetch, 

1420 parent_otel_span=parent_otel_span, 

1421 local_only=False, # Check Redis too 

1422 ) 

1423 

1424 # For keys that don't exist yet, set them to 0 

1425 if cache_values is None: 

1426 cache_values = [ # rebind-ok: missing keys default to a zeroed window snapshot 

1427 str(now_int) if key.endswith(":window") else 0 for key in keys_to_fetch 

1428 ] 

1429 elif self.batch_rate_limiter_script is not None: 

1430 # NORMAL MODE: Increment counters in Redis 

1431 # Group keys by hash tag for Redis cluster compatibility 

1432 cache_values = await self._execute_redis_batch_rate_limiter_script( 

1433 keys_to_fetch=keys_to_fetch, 

1434 now_int=now_int, 

1435 ) 

1436 

1437 # update in-memory cache with new values 

1438 for i in range(0, len(cache_values), 2): 

1439 window_key = keys_to_fetch[i] 

1440 counter_key = keys_to_fetch[i + 1] 

1441 window_value = cache_values[i] 

1442 counter_value = cache_values[i + 1] 

1443 await self.internal_usage_cache.async_set_cache( 

1444 key=counter_key, 

1445 value=counter_value, 

1446 ttl=self.window_size, 

1447 litellm_parent_otel_span=parent_otel_span, 

1448 local_only=True, 

1449 ) 

1450 await self.internal_usage_cache.async_set_cache( 

1451 key=window_key, 

1452 value=window_value, 

1453 ttl=self.window_size, 

1454 litellm_parent_otel_span=parent_otel_span, 

1455 local_only=True, 

1456 ) 

1457 else: 

1458 # NORMAL MODE: In-memory sliding window (no Redis) 

1459 cache_values = await self.in_memory_cache_sliding_window( 

1460 keys=keys_to_fetch, 

1461 now_int=now_int, 

1462 window_size=self.window_size, 

1463 ) 

1464 

1465 windowed_response = self.is_cache_list_over_limit(keys_to_fetch, cache_values, key_metadata) 

1466 if windowed_response["overall_code"] == "OVER_LIMIT": 

1467 return windowed_response 

1468 

1469 if not gauges: 

1470 return windowed_response 

1471 

1472 gauge_response: Final = await self._check_parallel_request_gauges( 

1473 gauges=gauges, 

1474 slot_id=parallel_slot_id or uuid.uuid4().hex, 

1475 parent_otel_span=parent_otel_span, 

1476 read_only=read_only, 

1477 ) 

1478 return RateLimitResponse( 

1479 overall_code=gauge_response["overall_code"], 

1480 statuses=[*windowed_response["statuses"], *gauge_response["statuses"]], 

1481 ) 

1482 

1483 def _collect_windowed_keys_and_gauges( 

1484 self, 

1485 descriptors: Sequence[RateLimitDescriptor], 

1486 skip_tpm_check: bool, 

1487 ) -> tuple[list[str], dict[str, WindowKeyMetadata], list[ParallelRequestGauge]]: 

1488 """ 

1489 Split descriptors into the windowed (window_key, counter_key) fetch 

1490 list with its per-window metadata, and the concurrency gauges for 

1491 descriptors carrying a max_parallel_requests limit. 

1492 """ 

1493 keys_to_fetch: Final[list[str]] = [] 

1494 key_metadata: Final[dict[str, WindowKeyMetadata]] = {} 

1495 gauges: Final[list[ParallelRequestGauge]] = [] 

1496 for descriptor in descriptors: 

1497 descriptor_key = descriptor["key"] 

1498 descriptor_value = descriptor["value"] 

1499 rate_limit: RateLimitDescriptorRateLimitObject = ( 

1500 descriptor.get("rate_limit") or RateLimitDescriptorRateLimitObject() 

1501 ) 

1502 requests_limit = rate_limit.get("requests_per_unit") 

1503 tokens_limit = None if skip_tpm_check else rate_limit.get("tokens_per_unit") 

1504 max_parallel_requests_limit = rate_limit.get("max_parallel_requests") 

1505 window_size = rate_limit.get("window_size") or self.window_size 

1506 

1507 window_key = f"{{{descriptor_key}:{descriptor_value}}}:window" 

1508 

1509 if max_parallel_requests_limit is not None: 

1510 gauges.append( 

1511 ParallelRequestGauge( 

1512 counter_key=self.create_rate_limit_keys( 

1513 descriptor_key, descriptor_value, "max_parallel_requests" 

1514 ), 

1515 limit=int(max_parallel_requests_limit), 

1516 descriptor_key=descriptor_key, 

1517 ) 

1518 ) 

1519 

1520 rate_limit_set = False 

1521 if requests_limit is not None: 

1522 rpm_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, "requests") 

1523 keys_to_fetch.extend([window_key, rpm_key]) 

1524 rate_limit_set = True 

1525 if tokens_limit is not None: 

1526 tpm_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, "tokens") 

1527 keys_to_fetch.extend([window_key, tpm_key]) 

1528 rate_limit_set = True 

1529 

1530 if not rate_limit_set: 

1531 continue 

1532 

1533 key_metadata[window_key] = { 

1534 "requests_limit": (int(requests_limit) if requests_limit is not None else None), 

1535 "tokens_limit": int(tokens_limit) if tokens_limit is not None else None, 

1536 "window_size": int(window_size), 

1537 "descriptor_key": descriptor_key, 

1538 "descriptor_value": descriptor_value, 

1539 } 

1540 return keys_to_fetch, key_metadata, gauges 

1541 

1542 def _gauge_status(self, gauge: ParallelRequestGauge, in_flight: int, code: str) -> RateLimitStatus: 

1543 return RateLimitStatus( 

1544 code=code, 

1545 current_limit=gauge["limit"], 

1546 limit_remaining=max(0, gauge["limit"] - in_flight), 

1547 rate_limit_type="max_parallel_requests", 

1548 descriptor_key=gauge["descriptor_key"], 

1549 ) 

1550 

1551 def _gauge_in_flight_from_cache_value(self, raw_value: ParallelGaugeCacheValue | None) -> int: 

1552 """ 

1553 In-flight count from a cached gauge value: a dict of slot_id -> 

1554 acquire timestamp when the in-memory registry is authoritative, or 

1555 the mirrored integer count from the last Redis script result. 

1556 """ 

1557 if raw_value is None: 

1558 return 0 

1559 if isinstance(raw_value, dict): 

1560 cutoff: Final = self._get_current_time().timestamp() - PARALLEL_REQUEST_SLOT_TTL_SECONDS 

1561 return sum(1 for ts in raw_value.values() if isinstance(ts, (int, float)) and ts >= cutoff) 

1562 return max(0, int(raw_value)) 

1563 

1564 async def _check_parallel_request_gauges( 

1565 self, 

1566 gauges: list[ParallelRequestGauge], 

1567 slot_id: str, 

1568 parent_otel_span: Span | None = None, 

1569 read_only: bool = False, 

1570 ) -> RateLimitResponse: 

1571 """ 

1572 Enforce max_parallel_requests as a concurrency gauge over a per-slot 

1573 registry: each admitted request registers ``slot_id`` with its 

1574 acquire time, and admission requires in_flight + 1 <= limit over the 

1575 unexpired slots. Unlike the windowed RPM/TPM counters, the gauge is 

1576 never reset while requests are in flight, a rejected request never 

1577 occupies a slot, and a slot leaked by a crashed worker is pruned 

1578 after PARALLEL_REQUEST_SLOT_TTL_SECONDS even under continuous 

1579 traffic. Releases remove exactly this request's slot id, so a 

1580 double-fired or unmatched release can never free another request's 

1581 slot. 

1582 """ 

1583 gauge_keys: Final = [gauge["counter_key"] for gauge in gauges] 

1584 

1585 if read_only: 

1586 if self.parallel_count_script is not None: 

1587 try: 

1588 raw_counts: Final[list[CacheCounterValue]] = await self.parallel_count_script( 

1589 keys=gauge_keys, 

1590 args=[PARALLEL_REQUEST_SLOT_TTL_SECONDS for _ in gauges], 

1591 ) 

1592 counts = [max(0, int(value)) for value in raw_counts] 

1593 except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the local mirror, never a 500 

1594 log_redis_failure( 

1595 verbose_proxy_logger, logging.WARNING, "parallel_count_script failed, using local mirror", e 

1596 ) 

1597 counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span) 

1598 else: 

1599 counts = await self._read_local_gauge_counts(gauge_keys, parent_otel_span) 

1600 statuses = [] 

1601 overall_code = "OK" 

1602 for gauge, in_flight in zip(gauges, counts): 

1603 code = "OVER_LIMIT" if in_flight >= gauge["limit"] else "OK" 

1604 if code == "OVER_LIMIT": 

1605 overall_code = "OVER_LIMIT" 

1606 statuses.append(self._gauge_status(gauge, in_flight, code)) 

1607 return RateLimitResponse(overall_code=overall_code, statuses=statuses) 

1608 

1609 local_counts: Final = await self._read_local_gauge_counts(gauge_keys, parent_otel_span) 

1610 for gauge, in_flight in zip(gauges, local_counts): 

1611 if in_flight >= gauge["limit"]: 

1612 return RateLimitResponse( 

1613 overall_code="OVER_LIMIT", 

1614 statuses=[self._gauge_status(gauge, in_flight, "OVER_LIMIT")], 

1615 ) 

1616 

1617 if self.parallel_acquire_script is not None: 

1618 try: 

1619 raw: Final[list[CacheCounterValue]] = await self.parallel_acquire_script( 

1620 keys=gauge_keys, 

1621 args=[ 

1622 arg for gauge in gauges for arg in (gauge["limit"], PARALLEL_REQUEST_SLOT_TTL_SECONDS, slot_id) 

1623 ], 

1624 ) 

1625 except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to in-memory enforcement, never a 500 

1626 log_redis_failure( 

1627 verbose_proxy_logger, 

1628 logging.WARNING, 

1629 "parallel_acquire_script failed, falling back to in-memory gauge", 

1630 e, 

1631 ) 

1632 async with self._check_and_increment_lock: 

1633 return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span) 

1634 if int(raw[0]) == 1: 

1635 gauge = gauges[int(raw[1]) - 1] 

1636 return RateLimitResponse( 

1637 overall_code="OVER_LIMIT", 

1638 statuses=[self._gauge_status(gauge, int(raw[2]), "OVER_LIMIT")], 

1639 ) 

1640 statuses = [] 

1641 for gauge, in_flight in zip(gauges, raw[1:]): 

1642 await self.internal_usage_cache.async_set_cache( 

1643 key=gauge["counter_key"], 

1644 value=int(in_flight), 

1645 ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, 

1646 litellm_parent_otel_span=parent_otel_span, 

1647 local_only=True, 

1648 ) 

1649 statuses.append(self._gauge_status(gauge, int(in_flight), "OK")) 

1650 return RateLimitResponse(overall_code="OK", statuses=statuses) 

1651 

1652 async with self._check_and_increment_lock: 

1653 return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span) 

1654 

1655 async def _read_local_gauge_counts( 

1656 self, 

1657 gauge_keys: list[str], 

1658 parent_otel_span: Span | None = None, 

1659 ) -> list[int]: 

1660 values: Final = await self._batch_get_gauge_values( 

1661 keys=gauge_keys, 

1662 parent_otel_span=parent_otel_span, 

1663 ) 

1664 if values is None: 

1665 return [0 for _ in gauge_keys] 

1666 return [self._gauge_in_flight_from_cache_value(value) for value in values] 

1667 

1668 async def _acquire_parallel_slots_in_memory( 

1669 self, 

1670 gauges: list[ParallelRequestGauge], 

1671 slot_id: str, 

1672 parent_otel_span: Span | None = None, 

1673 ) -> RateLimitResponse: 

1674 """ 

1675 All-or-nothing in-memory slot-registry acquire. Caller holds the lock. 

1676 

1677 A cached dict is the authoritative in-memory registry. A cached 

1678 integer is the count mirrored from the last successful Redis script 

1679 call: when Redis fails over to this path, that mirror still counts 

1680 the slots in flight on the Redis side, so it is carried forward as 

1681 an integer counter (not discarded as an empty registry, which would 

1682 briefly double the admitted concurrency during a Redis outage). 

1683 """ 

1684 now: Final = self._get_current_time().timestamp() 

1685 cutoff: Final = now - PARALLEL_REQUEST_SLOT_TTL_SECONDS 

1686 states: Final[list[tuple[dict[str, float] | None, int]]] = [] 

1687 for gauge in gauges: 

1688 raw_value: ParallelGaugeCacheValue | None = await self.internal_usage_cache.async_get_cache( 

1689 key=gauge["counter_key"], 

1690 litellm_parent_otel_span=parent_otel_span, 

1691 local_only=True, 

1692 ) 

1693 if isinstance(raw_value, dict): 

1694 registry: dict[str, float] | None = { 

1695 key: float(ts) for key, ts in raw_value.items() if isinstance(ts, (int, float)) and ts >= cutoff 

1696 } 

1697 in_flight = len(registry or {}) 

1698 elif raw_value is None: 

1699 registry = {} 

1700 in_flight = 0 

1701 else: 

1702 registry = None 

1703 in_flight = max(0, int(raw_value)) 

1704 if in_flight + 1 > gauge["limit"]: 

1705 return RateLimitResponse( 

1706 overall_code="OVER_LIMIT", 

1707 statuses=[self._gauge_status(gauge, in_flight, "OVER_LIMIT")], 

1708 ) 

1709 states.append((registry, in_flight)) 

1710 

1711 statuses: Final = [] 

1712 for gauge, (registry, in_flight) in zip(gauges, states): 

1713 new_value: dict[str, float] | int = {**registry, slot_id: now} if registry is not None else in_flight + 1 

1714 await self.internal_usage_cache.async_set_cache( 

1715 key=gauge["counter_key"], 

1716 value=new_value, 

1717 ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, 

1718 litellm_parent_otel_span=parent_otel_span, 

1719 local_only=True, 

1720 ) 

1721 statuses.append(self._gauge_status(gauge, in_flight + 1, "OK")) 

1722 return RateLimitResponse(overall_code="OK", statuses=statuses) 

1723 

1724 async def _release_stashed_parallel_slot( 

1725 self, 

1726 stash: RequestRateLimiterStash | None, 

1727 parent_otel_span: Span | None, 

1728 ) -> None: 

1729 if stash is None: 

1730 return 

1731 async with stash.parallel_slot_release_lock: 

1732 acquisition: Final = stash.parallel_slot 

1733 if acquisition is None: 1733 ↛ 1735line 1733 didn't jump to line 1735 because the condition on line 1733 was always true

1734 return 

1735 await self._release_parallel_request_slots(acquisition, parent_otel_span) 

1736 stash.parallel_slot = None # rebind-ok: marks this request's slot as released 

1737 

1738 async def _release_parallel_request_slots( 

1739 self, 

1740 acquisition: ParallelSlotAcquisition, 

1741 parent_otel_span: Span | None = None, 

1742 ) -> None: 

1743 """ 

1744 Release the max_parallel_requests slots acquired at pre-call by 

1745 removing this request's slot id from every gauge it was registered 

1746 under. Removing an absent slot id is a no-op, so a release without a 

1747 matching acquire or a double-fired release can never free another 

1748 request's slot. The in-memory fallback decrements integer mirror 

1749 values (floored at 0) because the mirror carries no per-slot ids. 

1750 """ 

1751 counter_keys: Final = acquisition["counter_keys"] 

1752 slot_id: Final = acquisition["slot_id"] 

1753 if not counter_keys or not slot_id: 

1754 return 

1755 if self.parallel_release_script is not None: 

1756 try: 

1757 raw: Final[list[CacheCounterValue]] = await self.parallel_release_script( 

1758 keys=counter_keys, 

1759 args=[slot_id for _ in counter_keys], 

1760 ) 

1761 for counter_key, remaining in zip(counter_keys, raw): 

1762 await self.internal_usage_cache.async_set_cache( 

1763 key=counter_key, 

1764 value=max(0, int(remaining)), 

1765 ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, 

1766 litellm_parent_otel_span=parent_otel_span, 

1767 local_only=True, 

1768 ) 

1769 return 

1770 except Exception as e: # noqa: BLE001 - any Redis/Lua failure degrades to the in-memory release, never a 500 

1771 log_redis_failure( 

1772 verbose_proxy_logger, 

1773 logging.WARNING, 

1774 "parallel_release_script failed, falling back to in-memory release", 

1775 e, 

1776 ) 

1777 

1778 async with self._check_and_increment_lock: 

1779 for counter_key in counter_keys: 

1780 raw_value: ParallelGaugeCacheValue | None = await self.internal_usage_cache.async_get_cache( 

1781 key=counter_key, 

1782 litellm_parent_otel_span=parent_otel_span, 

1783 local_only=True, 

1784 ) 

1785 if isinstance(raw_value, dict): 

1786 if slot_id not in raw_value: 

1787 continue 

1788 new_value: dict[str, object] | int = {key: ts for key, ts in raw_value.items() if key != slot_id} 

1789 elif raw_value is None: 

1790 continue 

1791 else: 

1792 new_value = max(0, int(raw_value) - 1) 

1793 await self.internal_usage_cache.async_set_cache( 

1794 key=counter_key, 

1795 value=new_value, 

1796 ttl=PARALLEL_REQUEST_SLOT_TTL_SECONDS, 

1797 litellm_parent_otel_span=parent_otel_span, 

1798 local_only=True, 

1799 ) 

1800 

1801 async def atomic_check_and_increment_by_n( 

1802 self, 

1803 descriptors: list[RateLimitDescriptor], 

1804 increments: list[dict[Literal["requests", "tokens"], int]], 

1805 parent_otel_span: Span | None = None, 

1806 ) -> RateLimitResponse: 

1807 """ 

1808 Atomic check-and-increment-by-N across one or more descriptors. 

1809 

1810 All-or-nothing: if any descriptor would exceed its limit, no counter is 

1811 modified and the response carries `overall_code = "OVER_LIMIT"` with 

1812 the offending descriptor's status. Closes the TOCTOU window between 

1813 read and increment in both single-process and multi-process (Redis) 

1814 deployments. 

1815 

1816 Cluster-safety: each descriptor's keys all share a `{key:value}` hash 

1817 tag, so the Redis Lua path issues one Lua call per descriptor — every 

1818 call's keys co-locate on a single Redis Cluster slot, avoiding 

1819 CROSSSLOT errors. Cross-descriptor atomicity is preserved via 

1820 refund-on-rollback: if descriptor i is OVER_LIMIT, descriptors 0..i-1 

1821 get a direct INCRBY refund (refunds need no atomicity guarantee). 

1822 

1823 Args: 

1824 descriptors: rate-limit descriptors to check 

1825 increments: per-descriptor increment amounts, indexed parallel to 

1826 `descriptors`. Each entry is `{"requests": int, "tokens": int}` 

1827 — values default to 0 when a descriptor has no matching limit. 

1828 

1829 Returns: 

1830 RateLimitResponse with one status per (descriptor, rate_limit_type) 

1831 counter, mirroring `should_rate_limit`'s shape. 

1832 """ 

1833 if len(descriptors) != len(increments): 

1834 raise ValueError("atomic_check_and_increment_by_n: descriptors and increments must have the same length") 

1835 

1836 # Build per-descriptor (keys, args, meta) groups. All keys within a 

1837 # group share the descriptor's {key:value} hash tag, so a single Lua 

1838 # call per group never triggers CROSSSLOT on Redis Cluster. 

1839 descriptor_groups: Final[list[DescriptorAtomicGroup]] = [] 

1840 for descriptor, increment_amounts in zip(descriptors, increments): 

1841 keys, args, meta = self._build_descriptor_atomic_payload( 

1842 descriptor=descriptor, 

1843 increment_amounts=increment_amounts, 

1844 ) 

1845 if keys: 

1846 descriptor_groups.append((keys, args, meta)) 

1847 

1848 if not descriptor_groups: 

1849 return RateLimitResponse(overall_code="OK", statuses=[]) 

1850 

1851 # Multi-process atomicity via Redis Lua, per descriptor for slot 

1852 # co-location. Single-process atomicity falls back to the 

1853 # asyncio.Lock + in-memory sliding window below — there are no 

1854 # cluster slot concerns locally, so we keep the batched 2-phase 

1855 # critical section for true cross-descriptor atomicity. 

1856 if self.check_and_increment_by_n_script is not None: 

1857 return await self._atomic_lua_per_descriptor( 

1858 descriptor_groups=descriptor_groups, 

1859 parent_otel_span=parent_otel_span, 

1860 ) 

1861 

1862 flat_meta: Final[list[AtomicCounterMeta]] = [ 

1863 m for _keys, _args, group_meta in descriptor_groups for m in group_meta 

1864 ] 

1865 async with self._check_and_increment_lock: 

1866 return await self._atomic_check_and_increment_in_memory( 

1867 per_counter_meta=flat_meta, 

1868 parent_otel_span=parent_otel_span, 

1869 ) 

1870 

1871 def _build_descriptor_atomic_payload( 

1872 self, 

1873 descriptor: RateLimitDescriptor, 

1874 increment_amounts: dict[Literal["requests", "tokens"], int], 

1875 ) -> DescriptorAtomicGroup: 

1876 """ 

1877 Build (KEYS, ARGV, per-counter meta) for a single descriptor's Lua 

1878 call. All keys returned share the descriptor's {key:value} hash tag. 

1879 """ 

1880 descriptor_key: Final = descriptor["key"] 

1881 descriptor_value: Final = descriptor["value"] 

1882 rate_limit: Final[RateLimitDescriptorRateLimitObject] = ( 

1883 descriptor.get("rate_limit") or RateLimitDescriptorRateLimitObject() 

1884 ) 

1885 window_size: Final = rate_limit.get("window_size") or self.window_size 

1886 window_key: Final = f"{{{descriptor_key}:{descriptor_value}}}:window" 

1887 

1888 keys: Final[list[str]] = [] 

1889 args: Final[list[int]] = [] 

1890 meta: Final[list[AtomicCounterMeta]] = [] 

1891 

1892 rate_limit_types: Final[tuple[Literal["requests", "tokens"], ...]] = ("requests", "tokens") 

1893 for rlt in rate_limit_types: 

1894 if rlt == "requests": 

1895 limit_value = rate_limit.get("requests_per_unit") 

1896 inc_amount = int(increment_amounts.get("requests", 0) or 0) 

1897 else: 

1898 limit_value = rate_limit.get("tokens_per_unit") 

1899 inc_amount = int(increment_amounts.get("tokens", 0) or 0) 

1900 if limit_value is None or inc_amount < 0: 

1901 continue 

1902 counter_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, rlt) 

1903 # Counter-key TTL and window_size are conceptually distinct 

1904 # ("how long the counter Redis key lives" vs "how long the 

1905 # sliding window is"). Kept as separate values so a future 

1906 # custom-TTL descriptor doesn't reintroduce a silent expiry bug. 

1907 ttl_seconds = int(window_size) 

1908 window_size_seconds = int(window_size) 

1909 keys.extend([window_key, counter_key]) 

1910 # 4-tuple matches the Lua ARGV layout: 

1911 # [limit, increment, ttl_seconds, window_size_seconds]. 

1912 args.extend([int(limit_value), inc_amount, ttl_seconds, window_size_seconds]) 

1913 meta.append( 

1914 { 

1915 "descriptor_key": descriptor_key, 

1916 "descriptor_value": descriptor_value, 

1917 "current_limit": int(limit_value), 

1918 "rate_limit_type": rlt, 

1919 "window_key": window_key, 

1920 "counter_key": counter_key, 

1921 "increment": inc_amount, 

1922 "ttl": ttl_seconds, 

1923 "window_size": window_size_seconds, 

1924 } 

1925 ) 

1926 return keys, args, meta 

1927 

1928 async def _atomic_lua_per_descriptor( 

1929 self, 

1930 descriptor_groups: list[DescriptorAtomicGroup], 

1931 parent_otel_span: Span | None = None, 

1932 ) -> RateLimitResponse: 

1933 """ 

1934 Run Lua check-and-increment one descriptor at a time so each call's 

1935 keys co-locate on a single Redis Cluster slot. On OVER_LIMIT for 

1936 descriptor i, refund descriptors 0..i-1's increments. On Lua failure 

1937 mid-loop, refund applied increments and fall back to in-memory. 

1938 """ 

1939 if not descriptor_groups: 

1940 return RateLimitResponse( 

1941 overall_code="OK", 

1942 statuses=[], # mutable-ok: response contract requires a status list 

1943 ) 

1944 applied: Final[list[list[AtomicCounterMeta]]] = [] 

1945 statuses: Final[list[RateLimitStatus]] = [] 

1946 reservation_windows: Final[set[ReservationWindowIdentity]] = set() # mutable-ok: filled by the group loop 

1947 raw: list[CacheCounterValue] 

1948 

1949 for _idx, (keys, args, meta) in enumerate(descriptor_groups): 

1950 try: 

1951 raw = await self.check_and_increment_by_n_script( # pyright: ignore[reportOptionalCall] # sole caller guards it is not None 

1952 keys=keys, 

1953 args=args, 

1954 ) 

1955 except Exception as e: 

1956 # Lua failure (timeout, OOM, network partition) leaves Redis 

1957 # state ambiguous. Refund any prior groups so Redis returns 

1958 # to its pre-call state, then fall back to in-memory for the 

1959 # whole call (counters there are independent of Redis). 

1960 log_redis_failure( 

1961 verbose_proxy_logger, 

1962 logging.ERROR, 

1963 f"atomic_check_and_increment_by_n: Redis Lua execution failed ({type(e).__name__}). Refunding " 

1964 f"{len(applied)} prior descriptors and falling back to in-memory enforcement, counters will " 

1965 f"diverge from Redis until window expires (window_size={self.window_size}s)", 

1966 e, 

1967 ) 

1968 await self._refund_applied_descriptor_groups(applied) 

1969 flat_meta: list[AtomicCounterMeta] = [m for _k, _a, group_meta in descriptor_groups for m in group_meta] 

1970 async with self._check_and_increment_lock: 

1971 return await self._atomic_check_and_increment_in_memory( 

1972 per_counter_meta=flat_meta, 

1973 parent_otel_span=parent_otel_span, 

1974 ) 

1975 

1976 response = self._build_atomic_response(raw, meta) 

1977 if response["overall_code"] == "OVER_LIMIT": 

1978 await self._refund_applied_descriptor_groups(applied) 

1979 return response 

1980 if len(descriptor_groups) == 1: 

1981 return response 

1982 applied.append(meta) 

1983 statuses.extend(response["statuses"]) 

1984 reservation_windows.update(response.get("reservation_windows", frozenset())) 

1985 

1986 return RateLimitResponse( 

1987 overall_code="OK", 

1988 statuses=statuses, 

1989 reservation_windows=frozenset(reservation_windows), 

1990 ) 

1991 

1992 async def _refund_applied_descriptor_groups( 

1993 self, 

1994 applied: list[list[AtomicCounterMeta]], 

1995 ) -> None: 

1996 """ 

1997 Decrement counters for descriptor groups already applied via Lua. 

1998 Best-effort: refund failures are logged but not raised — the original 

1999 OVER_LIMIT / fallback decision is what matters to the caller. 

2000 """ 

2001 if not applied: 

2002 return 

2003 redis_cache: Final = self.internal_usage_cache.dual_cache.redis_cache 

2004 if redis_cache is None: 

2005 return 

2006 for group_meta in applied: 

2007 for entry in group_meta: 

2008 try: 

2009 await redis_cache.async_increment( 

2010 key=entry["counter_key"], 

2011 value=-entry["increment"], 

2012 ) 

2013 except Exception as e: 

2014 log_redis_failure( 

2015 verbose_proxy_logger, 

2016 logging.WARNING, 

2017 f"Failed to refund {entry['counter_key']} on cross-descriptor rollback", 

2018 e, 

2019 ) 

2020 

2021 def _build_atomic_response( 

2022 self, 

2023 raw: list[CacheCounterValue], 

2024 per_counter_meta: list[AtomicCounterMeta], 

2025 ) -> RateLimitResponse: 

2026 """Convert Lua script return value to RateLimitResponse. 

2027 

2028 Indexing invariant: `per_counter_meta` and `KEYS` are parallel-indexed 

2029 at the COUNTER level, not the descriptor level. A descriptor with both 

2030 RPM and TPM limits emits two `(window_key, counter_key)` pairs and 

2031 two meta entries — one per counter. The Lua script's loop variable 

2032 `i` therefore enumerates counters, and the over-limit return tuple 

2033 `{1, i, ...}` carries a counter index that maps directly to 

2034 `per_counter_meta[i - 1]`. Keep these arrays parallel at the counter 

2035 level when modifying this code. 

2036 """ 

2037 if not raw: 

2038 return RateLimitResponse(overall_code="OK", statuses=[]) 

2039 

2040 status_code: Final = int(raw[0]) 

2041 if status_code == 1: 

2042 # Over limit: { 1, counter_index (1-based), current_counter, limit } 

2043 descriptor_index: Final = int(raw[1]) - 1 

2044 current_counter: Final = int(raw[2]) 

2045 limit: Final = int(raw[3]) 

2046 meta = per_counter_meta[descriptor_index] 

2047 return RateLimitResponse( 

2048 overall_code="OVER_LIMIT", 

2049 statuses=[ 

2050 RateLimitStatus( 

2051 code="OVER_LIMIT", 

2052 current_limit=limit, 

2053 limit_remaining=max(0, limit - current_counter), 

2054 rate_limit_type=meta["rate_limit_type"], 

2055 descriptor_key=meta["descriptor_key"], 

2056 descriptor_value=meta["descriptor_value"], 

2057 ) 

2058 ], 

2059 ) 

2060 

2061 statuses: Final[list[RateLimitStatus]] = [] 

2062 for index, meta in enumerate(per_counter_meta): 

2063 new_counter = raw[1 + index * 2] 

2064 statuses.append( 

2065 RateLimitStatus( 

2066 code="OK", 

2067 current_limit=meta["current_limit"], 

2068 limit_remaining=max(0, meta["current_limit"] - int(new_counter)), 

2069 rate_limit_type=meta["rate_limit_type"], 

2070 descriptor_key=meta["descriptor_key"], 

2071 descriptor_value=meta["descriptor_value"], 

2072 ) 

2073 ) 

2074 return RateLimitResponse( 

2075 overall_code="OK", 

2076 statuses=statuses, 

2077 reservation_windows=frozenset( 

2078 ( 

2079 meta["counter_key"], 

2080 str(int(raw[2 + index * 2])), 

2081 "redis", 

2082 ) 

2083 for index, meta in enumerate(per_counter_meta) 

2084 ), 

2085 ) 

2086 

2087 async def _atomic_check_and_increment_in_memory( 

2088 self, 

2089 per_counter_meta: list[AtomicCounterMeta], 

2090 parent_otel_span: Span | None = None, 

2091 ) -> RateLimitResponse: 

2092 """In-memory all-or-nothing check-and-increment. Caller holds lock. 

2093 

2094 Reads/writes the LOCAL DualCache (`local_only=True`) — note this is 

2095 a different store from Redis. When this fallback fires after a Lua 

2096 failure, in-memory counters will diverge from Redis until each key's 

2097 window expires (TTL bounds divergence). 

2098 """ 

2099 # Use a single 'now' for the duration of this critical section so all 

2100 # descriptors evaluate window expiry consistently. 

2101 now_int: Final = int(self._get_current_time().timestamp()) 

2102 

2103 # Pass 1: read state, validate. 

2104 descriptor_state: Final[list[AtomicCounterState]] = [] 

2105 for meta in per_counter_meta: 

2106 window_size = meta["window_size"] 

2107 window_start: CacheCounterValue | None = await self.internal_usage_cache.async_get_cache( 

2108 key=meta["window_key"], 

2109 litellm_parent_otel_span=parent_otel_span, 

2110 local_only=True, 

2111 ) 

2112 window_expired = window_start is None or (now_int - int(window_start)) >= window_size 

2113 raw_counter: CacheCounterValue | None = ( 

2114 None 

2115 if window_expired 

2116 else await self.internal_usage_cache.async_get_cache( 

2117 key=meta["counter_key"], 

2118 litellm_parent_otel_span=parent_otel_span, 

2119 local_only=True, 

2120 ) 

2121 ) 

2122 current_counter = 0 if window_expired else int(raw_counter or 0) 

2123 over_limit = ( 

2124 current_counter + meta["increment"] > meta["current_limit"] 

2125 if meta["increment"] > 0 

2126 else current_counter >= meta["current_limit"] 

2127 ) 

2128 if over_limit: 

2129 return RateLimitResponse( 

2130 overall_code="OVER_LIMIT", 

2131 statuses=[ 

2132 RateLimitStatus( 

2133 code="OVER_LIMIT", 

2134 current_limit=meta["current_limit"], 

2135 limit_remaining=max(0, meta["current_limit"] - current_counter), 

2136 rate_limit_type=meta["rate_limit_type"], 

2137 descriptor_key=meta["descriptor_key"], 

2138 descriptor_value=meta["descriptor_value"], 

2139 ) 

2140 ], 

2141 ) 

2142 descriptor_state.append( 

2143 { # mutable-ok: local atomic-counter state is updated during pass two 

2144 "window_expired": window_expired, 

2145 "current": current_counter, 

2146 "window_start": str(now_int if window_expired else int(window_start)), 

2147 } 

2148 ) 

2149 

2150 # Pass 2: apply increments. 

2151 expired_windows: Final[Mapping[str, int]] = { 

2152 meta["window_key"]: meta["window_size"] 

2153 for meta, state in zip(per_counter_meta, descriptor_state) 

2154 if state["window_expired"] 

2155 } 

2156 for window_key, window_size in expired_windows.items(): 

2157 for sibling_counter_key in _sibling_counter_keys(window_key): 

2158 await self.internal_usage_cache.async_set_cache( 

2159 key=sibling_counter_key, 

2160 value=0, 

2161 ttl=window_size, 

2162 litellm_parent_otel_span=parent_otel_span, 

2163 local_only=True, 

2164 ) 

2165 statuses: Final[list[RateLimitStatus]] = [] 

2166 for meta, state in zip(per_counter_meta, descriptor_state): 

2167 new_counter = meta["increment"] if state["window_expired"] else state["current"] + meta["increment"] 

2168 if state["window_expired"]: 

2169 await self.internal_usage_cache.async_set_cache( 

2170 key=meta["window_key"], 

2171 value=str(now_int), 

2172 ttl=meta["window_size"], 

2173 litellm_parent_otel_span=parent_otel_span, 

2174 local_only=True, 

2175 ) 

2176 await self.internal_usage_cache.async_set_cache( 

2177 key=meta["counter_key"], 

2178 value=new_counter, 

2179 ttl=meta["ttl"], 

2180 litellm_parent_otel_span=parent_otel_span, 

2181 local_only=True, 

2182 ) 

2183 statuses.append( 

2184 RateLimitStatus( 

2185 code="OK", 

2186 current_limit=meta["current_limit"], 

2187 limit_remaining=max(0, meta["current_limit"] - new_counter), 

2188 rate_limit_type=meta["rate_limit_type"], 

2189 descriptor_key=meta["descriptor_key"], 

2190 descriptor_value=meta["descriptor_value"], 

2191 ) 

2192 ) 

2193 return RateLimitResponse( 

2194 overall_code="OK", 

2195 statuses=statuses, 

2196 reservation_windows=frozenset( 

2197 (meta["counter_key"], state["window_start"], "local") 

2198 for meta, state in zip(per_counter_meta, descriptor_state) 

2199 ), 

2200 ) 

2201 

2202 async def reserve_tpm_tokens( 

2203 self, 

2204 descriptors: list[RateLimitDescriptor], 

2205 estimated_tokens: int, 

2206 parent_otel_span: Span | None = None, 

2207 ) -> RateLimitResponse: 

2208 """ 

2209 Reserve ``estimated_tokens`` against every TPM-bearing descriptor 

2210 BEFORE the upstream call, so concurrent requests cannot all observe 

2211 "under limit" before any of them increments the counter. 

2212 

2213 Thin wrapper around ``atomic_check_and_increment_by_n``: builds a 

2214 TPM-only descriptor/increment list and delegates the all-or-nothing 

2215 atomicity (Lua on Redis, asyncio-locked DualCache otherwise) to the 

2216 shared primitive. 

2217 

2218 Excludes project ITPM/OTPM descriptors -- those are reserved 

2219 separately (different estimate per bucket) via ``reserve_io_tokens``. 

2220 """ 

2221 tpm_descriptors: Final[list[RateLimitDescriptor]] = [ 

2222 RateLimitDescriptor( 

2223 key=d["key"], 

2224 value=d["value"], 

2225 rate_limit=RateLimitDescriptorRateLimitObject( 

2226 tokens_per_unit=(d.get("rate_limit") or {}).get("tokens_per_unit"), 

2227 window_size=(d.get("rate_limit") or {}).get("window_size"), 

2228 ), 

2229 ) 

2230 for d in descriptors 

2231 if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) 

2232 and (d.get("rate_limit") or {}).get("tokens_per_unit") is not None # mutable-ok: optional descriptor 

2233 ] 

2234 if not tpm_descriptors: 

2235 return RateLimitResponse(overall_code="OK", statuses=[]) 

2236 

2237 increments: Final[list[dict[Literal["requests", "tokens"], int]]] = [ 

2238 {"tokens": estimated_tokens} for _ in tpm_descriptors 

2239 ] 

2240 return await self.atomic_check_and_increment_by_n( 

2241 descriptors=tpm_descriptors, 

2242 increments=increments, 

2243 parent_otel_span=parent_otel_span, 

2244 ) 

2245 

2246 async def _refund_reserved_tokens( 

2247 self, 

2248 scopes: Sequence[tuple[str, str]], 

2249 amount: int, 

2250 reservation_windows: frozenset[tuple[str, str, Literal["redis", "local"]]] = frozenset(), 

2251 parent_otel_span: Span | None = None, 

2252 ) -> None: 

2253 """ 

2254 Directly decrement previously-reserved token counters for ``scopes`` 

2255 by ``amount``. Used to roll back a reservation that already 

2256 succeeded once a *different* bucket in the same request turns out to 

2257 be over its limit (e.g. ITPM reserved fine, OTPM then hits its 

2258 limit -- the ITPM reservation must not be left inflated). 

2259 """ 

2260 if amount <= 0 or not scopes: 

2261 return 

2262 if not reservation_windows: 

2263 await self.async_increment_tokens_with_ttl_preservation( 

2264 pipeline_operations=self._build_reservation_aware_tpm_ops( 

2265 targets=scopes, 

2266 reserved_scopes=frozenset(scopes), 

2267 actual_tokens=0, 

2268 reserved_tokens=amount, 

2269 ), 

2270 parent_otel_span=parent_otel_span, 

2271 ) 

2272 return 

2273 pipeline_operations: Final = self._build_project_reservation_ops( 

2274 targets=scopes, 

2275 reserved_scopes=frozenset(scopes), 

2276 actual_tokens=0, 

2277 reserved_tokens=amount, 

2278 reservation_window_identities=reservation_windows, 

2279 ) 

2280 await self.async_increment_reservation_aware_tokens( 

2281 pipeline_operations=pipeline_operations, 

2282 parent_otel_span=parent_otel_span, 

2283 ) 

2284 

2285 async def reserve_io_tokens( 

2286 self, 

2287 descriptors: Sequence[RateLimitDescriptor], 

2288 estimated_input_tokens: int, 

2289 estimated_output_tokens: int, 

2290 parent_otel_span: Span | None = None, 

2291 ) -> tuple[RateLimitResponse, int, int]: 

2292 """ 

2293 Reserve ``estimated_input_tokens`` against project ITPM descriptors 

2294 and ``estimated_output_tokens`` against project OTPM descriptors. 

2295 

2296 ITPM and OTPM are reserved from different-sized estimates, so unlike 

2297 same-size TPM descriptors they can't share a single 

2298 ``atomic_check_and_increment_by_n`` call -- each bucket gets its own 

2299 all-or-nothing atomic call. If the OTPM reservation is over limit 

2300 after ITPM already succeeded, the ITPM reservation this call made is 

2301 rolled back before returning, so a partial reservation never leaks. 

2302 

2303 Returns ``(response, itpm_reserved, otpm_reserved)`` -- the latter two 

2304 are the amounts actually reserved (0 if that bucket wasn't 

2305 configured, or if the reservation failed), for the caller to stash 

2306 for post-call reconciliation. 

2307 """ 

2308 itpm_descriptors: Final = [ # mutable-ok: atomic limiter API requires lists 

2309 d for d in descriptors if d["key"] == PROJECT_ITPM_DESCRIPTOR_KEY 

2310 ] 

2311 otpm_descriptors: Final = [ # mutable-ok: atomic limiter API requires lists 

2312 d for d in descriptors if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY 

2313 ] 

2314 

2315 if not itpm_descriptors and not otpm_descriptors: 

2316 return RateLimitResponse(overall_code="OK", statuses=[]), 0, 0 # mutable-ok: response contract uses a list 

2317 

2318 itpm_response: Final = ( 

2319 await self.atomic_check_and_increment_by_n( 

2320 descriptors=itpm_descriptors, 

2321 increments=[ # mutable-ok: atomic limiter API requires mutable increment records 

2322 {"tokens": estimated_input_tokens} # mutable-ok: atomic limiter increment record 

2323 for _ in itpm_descriptors 

2324 ], 

2325 parent_otel_span=parent_otel_span, 

2326 ) 

2327 if itpm_descriptors 

2328 else None 

2329 ) 

2330 if itpm_response is not None and itpm_response["overall_code"] == "OVER_LIMIT": 

2331 return itpm_response, 0, 0 

2332 itpm_reserved: Final = estimated_input_tokens if itpm_response is not None else 0 

2333 

2334 if otpm_descriptors: 

2335 otpm_response: Final = await self.atomic_check_and_increment_by_n( 

2336 descriptors=otpm_descriptors, 

2337 increments=[ # mutable-ok: atomic limiter API requires mutable increment records 

2338 {"tokens": estimated_output_tokens} # mutable-ok: atomic limiter increment record 

2339 for _ in otpm_descriptors 

2340 ], 

2341 parent_otel_span=parent_otel_span, 

2342 ) 

2343 if otpm_response["overall_code"] == "OVER_LIMIT": 

2344 if itpm_reserved > 0: 

2345 await self._refund_reserved_tokens( 

2346 scopes=[ # mutable-ok: reservation rollback accepts collected scopes 

2347 (d["key"], d["value"]) for d in itpm_descriptors 

2348 ], 

2349 amount=itpm_reserved, 

2350 reservation_windows=itpm_response.get("reservation_windows", frozenset()), 

2351 parent_otel_span=parent_otel_span, 

2352 ) 

2353 return otpm_response, 0, 0 

2354 statuses: Final = ( 

2355 [ # mutable-ok: response contract uses a list 

2356 *itpm_response["statuses"], 

2357 *otpm_response["statuses"], 

2358 ] 

2359 if itpm_response is not None 

2360 else otpm_response["statuses"] 

2361 ) 

2362 return ( 

2363 RateLimitResponse( 

2364 overall_code="OK", 

2365 statuses=statuses, 

2366 reservation_windows=( 

2367 ( 

2368 itpm_response.get("reservation_windows", frozenset()) 

2369 if itpm_response is not None 

2370 else frozenset() 

2371 ) 

2372 | otpm_response.get("reservation_windows", frozenset()) 

2373 ), 

2374 ), 

2375 itpm_reserved, 

2376 estimated_output_tokens, 

2377 ) 

2378 

2379 assert itpm_response is not None 

2380 return itpm_response, itpm_reserved, 0 

2381 

2382 async def enforce_project_io_token_quota_for_frame( 

2383 self, 

2384 user_api_key_dict: UserAPIKeyAuth | None, 

2385 requested_model: str | None, 

2386 estimated_input_tokens: int, 

2387 estimated_output_tokens: int, 

2388 ) -> None: 

2389 """Reserve one WebSocket ``response.create`` frame's tokens against 

2390 the caller's project ITPM/OTPM quota. 

2391 

2392 The Responses WebSocket connection-level pre-call hook only runs once 

2393 per connection, but a connection accepts many ``response.create`` 

2394 frames over its lifetime. Without this, a project caller could send 

2395 unlimited high-token generations after a single minimal reservation. 

2396 There is no per-frame post-call hook to reconcile against, so -- 

2397 like the batch rate limiter -- this charges the estimate immediately 

2398 and never refunds it. 

2399 """ 

2400 if user_api_key_dict is None: 

2401 return 

2402 descriptors: Final[list[RateLimitDescriptor]] = [] # mutable-ok: descriptor helper appends in place 

2403 self.add_project_io_token_rate_limit_descriptors_from_metadata( 

2404 user_api_key_dict=user_api_key_dict, 

2405 requested_model=requested_model, 

2406 descriptors=descriptors, 

2407 ) 

2408 if not descriptors: 

2409 return 

2410 response, _itpm_reserved, _otpm_reserved = await self.reserve_io_tokens( 

2411 descriptors=descriptors, 

2412 estimated_input_tokens=estimated_input_tokens, 

2413 estimated_output_tokens=estimated_output_tokens, 

2414 parent_otel_span=user_api_key_dict.parent_otel_span, 

2415 ) 

2416 if response["overall_code"] == "OVER_LIMIT": 

2417 self._handle_rate_limit_error(response, descriptors, requested_model) 

2418 

2419 def _rate_limited_model(self, requested_model: str | None) -> RateLimitedModel | None: 

2420 if not requested_model: 

2421 return None 

2422 return RateLimitedModel( 

2423 requested=requested_model, 

2424 group=self._model_group_resolver(requested_model) or requested_model, 

2425 ) 

2426 

2427 def create_organization_rate_limit_descriptor( 

2428 self, user_api_key_dict: UserAPIKeyAuth, requested_model: str | None = None 

2429 ) -> list[RateLimitDescriptor]: 

2430 descriptors: Final[list[RateLimitDescriptor]] = [] 

2431 

2432 # Global org rate limits 

2433 if user_api_key_dict.org_id is not None and ( 2433 ↛ 2436line 2433 didn't jump to line 2436 because the condition on line 2433 was never true

2434 user_api_key_dict.organization_rpm_limit is not None or user_api_key_dict.organization_tpm_limit is not None 

2435 ): 

2436 descriptors.append( 

2437 RateLimitDescriptor( 

2438 key="organization", 

2439 value=user_api_key_dict.org_id, 

2440 rate_limit={ 

2441 "requests_per_unit": user_api_key_dict.organization_rpm_limit, 

2442 "tokens_per_unit": user_api_key_dict.organization_tpm_limit, 

2443 "window_size": self.window_size, 

2444 }, 

2445 ) 

2446 ) 

2447 

2448 model: Final = self._rate_limited_model(requested_model) 

2449 if model is None: 

2450 return descriptors 

2451 model_specific_tpm_limit: Final = model.limit_in( 

2452 get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_tpm_limit") 

2453 ) 

2454 model_specific_rpm_limit: Final = model.limit_in( 

2455 get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_rpm_limit") 

2456 ) 

2457 if model_specific_tpm_limit is None and model_specific_rpm_limit is None: 2457 ↛ 2459line 2457 didn't jump to line 2459 because the condition on line 2457 was always true

2458 return descriptors 

2459 descriptors.append( 

2460 RateLimitDescriptor( 

2461 key="model_per_organization", 

2462 value=f"{user_api_key_dict.org_id}:{model.group}", 

2463 rate_limit={ 

2464 "requests_per_unit": model_specific_rpm_limit, 

2465 "tokens_per_unit": model_specific_tpm_limit, 

2466 "window_size": self.window_size, 

2467 }, 

2468 ) 

2469 ) 

2470 return descriptors 

2471 

2472 def _add_model_per_key_rate_limit_descriptor( 

2473 self, 

2474 user_api_key_dict: UserAPIKeyAuth, 

2475 requested_model: str | None, 

2476 descriptors: list[RateLimitDescriptor], 

2477 ) -> None: 

2478 """ 

2479 Add model-specific rate limit descriptor for API key if applicable. 

2480 

2481 Args: 

2482 user_api_key_dict: User API key authentication dictionary 

2483 requested_model: The model being requested 

2484 descriptors: List of rate limit descriptors to append to 

2485 """ 

2486 from litellm.proxy.auth.auth_utils import ( 

2487 get_key_model_rpm_limit, 

2488 get_key_model_tpm_limit, 

2489 ) 

2490 

2491 model: Final = self._rate_limited_model(requested_model) 

2492 if model is None: 

2493 return 

2494 model_specific_tpm_limit: Final = model.limit_in( 

2495 get_key_model_tpm_limit(user_api_key_dict, model_name=model.group) 

2496 ) 

2497 model_specific_rpm_limit: Final = model.limit_in( 

2498 get_key_model_rpm_limit(user_api_key_dict, model_name=model.group) 

2499 ) 

2500 if model_specific_tpm_limit is None and model_specific_rpm_limit is None: 2500 ↛ 2503line 2500 didn't jump to line 2503 because the condition on line 2500 was always true

2501 return 

2502 

2503 descriptors.append( 

2504 RateLimitDescriptor( 

2505 key="model_per_key", 

2506 value=f"{user_api_key_dict.api_key}:{model.group}", 

2507 rate_limit={ 

2508 "requests_per_unit": model_specific_rpm_limit, 

2509 "tokens_per_unit": model_specific_tpm_limit, 

2510 "window_size": self.window_size, 

2511 }, 

2512 ) 

2513 ) 

2514 

2515 def _add_tag_per_key_rate_limit_descriptor( 

2516 self, 

2517 user_api_key_dict: UserAPIKeyAuth, 

2518 data: dict, 

2519 descriptors: list[RateLimitDescriptor], 

2520 ) -> None: 

2521 """ 

2522 Add per-request-tag rpm limit descriptors for the API key. 

2523 

2524 Each tag carried on the request that has a configured limit gets its own 

2525 ``{api_key}:{tag}`` counter, so a burst on one tag/group never consumes 

2526 another's budget. Tags without a configured limit fall through to the 

2527 key-level descriptor. 

2528 """ 

2529 if not user_api_key_dict.api_key: 2529 ↛ 2530line 2529 didn't jump to line 2530 because the condition on line 2529 was never true

2530 return 

2531 

2532 tag_rpm_limit: Final = get_key_tag_rpm_limit(user_api_key_dict) or {} 

2533 if not tag_rpm_limit: 2533 ↛ 2536line 2533 didn't jump to line 2536 because the condition on line 2533 was always true

2534 return 

2535 

2536 for tag in dict.fromkeys(get_tags_from_request_body(data)): 

2537 rpm_limit = tag_rpm_limit.get(tag) 

2538 if rpm_limit is None: 

2539 continue 

2540 descriptors.append( 

2541 RateLimitDescriptor( 

2542 key="tag_per_key", 

2543 value=f"{user_api_key_dict.api_key}:{tag}", 

2544 rate_limit={ 

2545 "requests_per_unit": rpm_limit, 

2546 "tokens_per_unit": None, 

2547 "window_size": self.window_size, 

2548 }, 

2549 ) 

2550 ) 

2551 

2552 def _add_mcp_per_key_rate_limit_descriptor( 

2553 self, 

2554 user_api_key_dict: UserAPIKeyAuth, 

2555 mcp_server_name: str | None, 

2556 descriptors: list[RateLimitDescriptor], 

2557 ) -> None: 

2558 """ 

2559 Add a per-MCP-server rpm descriptor for the API key, if a limit is 

2560 configured for the server being called. 

2561 

2562 MCP tool calls have no token usage, so only requests_per_unit is set; 

2563 tokens_per_unit stays None so the TPM reservation path is never engaged. 

2564 """ 

2565 from litellm.proxy.auth.auth_utils import get_key_mcp_rpm_limit 

2566 

2567 if not mcp_server_name or not user_api_key_dict.api_key: 

2568 return 

2569 

2570 mcp_rpm_limit: Final = get_key_mcp_rpm_limit(user_api_key_dict) 

2571 if not mcp_rpm_limit: 

2572 return 

2573 

2574 server_rpm_limit: Final = mcp_rpm_limit.get(mcp_server_name) 

2575 if server_rpm_limit is None: 

2576 return 

2577 

2578 descriptors.append( 

2579 RateLimitDescriptor( 

2580 key="mcp_per_key", 

2581 value=f"{user_api_key_dict.api_key}:{mcp_server_name}", 

2582 rate_limit={ 

2583 "requests_per_unit": server_rpm_limit, 

2584 "tokens_per_unit": None, 

2585 "window_size": self.window_size, 

2586 }, 

2587 ) 

2588 ) 

2589 

2590 def _add_mcp_per_team_rate_limit_descriptor( 

2591 self, 

2592 user_api_key_dict: UserAPIKeyAuth, 

2593 mcp_server_name: str | None, 

2594 descriptors: list[RateLimitDescriptor], 

2595 ) -> None: 

2596 """ 

2597 Add a per-MCP-server rpm descriptor for the team, if a limit is 

2598 configured for the server being called. 

2599 """ 

2600 from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit 

2601 

2602 if not mcp_server_name: 

2603 return 

2604 

2605 # Which teams' buckets does this call charge? A key is pinned to exactly one team. A keyless 

2606 # MCP-admitted subject reaches servers through SEVERAL teams at once and has no team_id, so 

2607 # without the second source below its calls charged no team bucket at all and it outran every 

2608 # team's mcp_rpm_limit. Every applicable team is charged rather than one being picked: the 

2609 # limiter enforces all descriptors, so each team's own ceiling binds on a call made through 

2610 # its grant, and there is no arbitrary attribution when several teams grant the same server. 

2611 team_limits: Final[list[tuple[str | None, dict[str, int] | None]]] = [] 

2612 if user_api_key_dict.team_id: 

2613 team_limits.append((user_api_key_dict.team_id, get_team_mcp_rpm_limit(user_api_key_dict))) 

2614 for source_team_id, source_limit in (user_api_key_dict.mcp_source_team_rpm_limits or {}).items(): 

2615 team_limits.append((source_team_id, source_limit)) 

2616 

2617 for team_id, mcp_rpm_limit in team_limits: 

2618 if not team_id or not mcp_rpm_limit: 

2619 continue 

2620 server_rpm_limit = mcp_rpm_limit.get(mcp_server_name) 

2621 if server_rpm_limit is None: 

2622 continue 

2623 descriptors.append( 

2624 RateLimitDescriptor( 

2625 key="mcp_per_team", 

2626 value=f"{team_id}:{mcp_server_name}", 

2627 rate_limit={ 

2628 "requests_per_unit": server_rpm_limit, 

2629 "tokens_per_unit": None, 

2630 "window_size": self.window_size, 

2631 }, 

2632 ) 

2633 ) 

2634 

2635 def _should_enforce_rate_limit( 

2636 self, 

2637 limit_type: str | None, 

2638 model_has_failures: bool, 

2639 ) -> bool: 

2640 """ 

2641 Determine if rate limit should be enforced based on limit type and model health. 

2642 

2643 Args: 

2644 limit_type: Type of rate limit ("dynamic", "guaranteed_throughput", "best_effort_throughput", or None) 

2645 model_has_failures: Whether the model has recent failures 

2646 

2647 Returns: 

2648 True if rate limit should be enforced, False otherwise 

2649 """ 

2650 if limit_type == "dynamic": 

2651 # Dynamic mode: only enforce if model has failures 

2652 return model_has_failures 

2653 # All other modes (including None): always enforce 

2654 return True 

2655 

2656 def _get_enforced_limit( 

2657 self, 

2658 limit_value: int | None, 

2659 limit_type: str | None, 

2660 model_has_failures: bool, 

2661 ) -> int | None: 

2662 """ 

2663 Get the rate limit value to enforce based on limit type and model health. 

2664 

2665 Args: 

2666 limit_value: The configured limit value 

2667 limit_type: Type of rate limit ("dynamic", "guaranteed_throughput", "best_effort_throughput", or None) 

2668 model_has_failures: Whether the model has recent failures 

2669 

2670 Returns: 

2671 The limit value if it should be enforced, None otherwise 

2672 """ 

2673 if limit_value is None: 

2674 return None 

2675 

2676 if self._should_enforce_rate_limit( 

2677 limit_type=limit_type, 

2678 model_has_failures=model_has_failures, 

2679 ): 

2680 return limit_value 

2681 

2682 return None 

2683 

2684 def _is_dynamic_rate_limiting_enabled( 

2685 self, 

2686 rpm_limit_type: str | None, 

2687 tpm_limit_type: str | None, 

2688 ) -> bool: 

2689 """ 

2690 Check if dynamic rate limiting is enabled for either RPM or TPM. 

2691 

2692 Args: 

2693 rpm_limit_type: RPM rate limit type 

2694 tpm_limit_type: TPM rate limit type 

2695 

2696 Returns: 

2697 True if dynamic mode is enabled for either limit type 

2698 """ 

2699 return rpm_limit_type == "dynamic" or tpm_limit_type == "dynamic" 

2700 

2701 def _get_agent_from_registry(self, agent_id: str) -> "AgentResponse | None": 

2702 """Look up an agent from the in-memory registry by ID.""" 

2703 from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry 

2704 

2705 return global_agent_registry.get_agent_by_id(agent_id=agent_id) 

2706 

2707 def _get_resolved_agent_id(self, user_api_key_dict: UserAPIKeyAuth, data: dict) -> str | None: 

2708 """ 

2709 Resolve the agent_id from either the API key or request metadata. 

2710 Key-level agent_id takes precedence over metadata/header-supplied agent_id. 

2711 """ 

2712 key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None) 

2713 if key_agent_id: 2713 ↛ 2714line 2713 didn't jump to line 2714 because the condition on line 2713 was never true

2714 return key_agent_id 

2715 metadata: Final = data.get("metadata") or {} 

2716 return metadata.get("agent_id") 

2717 

2718 def _get_session_id_from_data(self, data: dict) -> str | None: 

2719 """Extract session_id from request metadata or litellm_session_id.""" 

2720 session_id = data.get("litellm_session_id") 

2721 if session_id: 

2722 return str(session_id) 

2723 metadata: Final = data.get("metadata") or {} 

2724 session_id = metadata.get("session_id") 

2725 if session_id: 

2726 return str(session_id) 

2727 litellm_metadata: Final = data.get("litellm_metadata") or {} 

2728 session_id = litellm_metadata.get("session_id") 

2729 if session_id: 

2730 return str(session_id) 

2731 return None 

2732 

2733 def _create_agent_rate_limit_descriptors( 

2734 self, 

2735 agent_id: str, 

2736 data: dict, 

2737 ) -> list[RateLimitDescriptor]: 

2738 """ 

2739 Create rate limit descriptors for agent-level and session-level limits. 

2740 

2741 Agent-level: caps total RPM/TPM across all sessions for a given agent. 

2742 Session-level: caps RPM/TPM within a single session (identified by session_id). 

2743 """ 

2744 descriptors: Final[list[RateLimitDescriptor]] = [] 

2745 

2746 agent: Final = self._get_agent_from_registry(agent_id) 

2747 if agent is None: 

2748 return descriptors 

2749 

2750 agent_rpm: Final = getattr(agent, "rpm_limit", None) 

2751 agent_tpm: Final = getattr(agent, "tpm_limit", None) 

2752 if agent_rpm is not None or agent_tpm is not None: 

2753 descriptors.append( 

2754 RateLimitDescriptor( 

2755 key="agent", 

2756 value=agent_id, 

2757 rate_limit={ 

2758 "requests_per_unit": agent_rpm, 

2759 "tokens_per_unit": agent_tpm, 

2760 "window_size": self.window_size, 

2761 }, 

2762 ) 

2763 ) 

2764 

2765 session_rpm: Final = getattr(agent, "session_rpm_limit", None) 

2766 session_tpm: Final = getattr(agent, "session_tpm_limit", None) 

2767 if session_rpm is not None or session_tpm is not None: 

2768 session_id: Final = self._get_session_id_from_data(data) 

2769 if session_id is not None: 

2770 descriptors.append( 

2771 RateLimitDescriptor( 

2772 key="agent_session", 

2773 value=f"{agent_id}:{session_id}", 

2774 rate_limit={ 

2775 "requests_per_unit": session_rpm, 

2776 "tokens_per_unit": session_tpm, 

2777 "window_size": self.window_size, 

2778 }, 

2779 ) 

2780 ) 

2781 

2782 return descriptors 

2783 

2784 async def _create_tag_rate_limit_descriptors(self, data: Mapping[str, object]) -> tuple[RateLimitDescriptor, ...]: 

2785 tags: Final = tuple(dict.fromkeys(get_tags_from_request_body(data))) 

2786 if not tags: 2786 ↛ 2788line 2786 didn't jump to line 2788 because the condition on line 2786 was always true

2787 return () 

2788 tag_limits: Final = await self._tag_rate_limit_resolver(tags) 

2789 return tuple( 

2790 _tag_rate_limit_descriptor(tag, limit, self.window_size) 

2791 for tag in tags 

2792 if (limit := tag_limits.get(tag)) is not None 

2793 ) 

2794 

2795 def _create_rate_limit_descriptors( 

2796 self, 

2797 user_api_key_dict: UserAPIKeyAuth, 

2798 data: dict, 

2799 rpm_limit_type: str | None, 

2800 tpm_limit_type: str | None, 

2801 model_has_failures: bool, 

2802 call_type: str | None = None, 

2803 ) -> list[RateLimitDescriptor]: 

2804 """ 

2805 Create all rate limit descriptors for the request. 

2806 

2807 Returns list of descriptors for API key, user, team, team member, end user, 

2808 model-specific, agent, and agent-session limits. 

2809 """ 

2810 descriptors: Final[list[RateLimitDescriptor]] = [] # mutable-ok: existing descriptor helpers append in place 

2811 

2812 # API Key rate limits 

2813 if user_api_key_dict.api_key and ( 2813 ↛ 2818line 2813 didn't jump to line 2818 because the condition on line 2813 was never true

2814 user_api_key_dict.rpm_limit is not None 

2815 or user_api_key_dict.tpm_limit is not None 

2816 or user_api_key_dict.max_parallel_requests is not None 

2817 ): 

2818 throttle_pct: Final = user_api_key_dict.budget_throttle_pct 

2819 descriptors.append( 

2820 RateLimitDescriptor( 

2821 key="api_key", 

2822 value=user_api_key_dict.api_key, 

2823 rate_limit={ 

2824 "requests_per_unit": self._get_enforced_limit( 

2825 limit_value=throttled_limit(user_api_key_dict.rpm_limit, throttle_pct), 

2826 limit_type=rpm_limit_type, 

2827 model_has_failures=model_has_failures, 

2828 ), 

2829 "tokens_per_unit": self._get_enforced_limit( 

2830 limit_value=throttled_limit(user_api_key_dict.tpm_limit, throttle_pct), 

2831 limit_type=tpm_limit_type, 

2832 model_has_failures=model_has_failures, 

2833 ), 

2834 "max_parallel_requests": user_api_key_dict.max_parallel_requests, 

2835 "window_size": self.window_size, 

2836 }, 

2837 ) 

2838 ) 

2839 

2840 # User rate limits 

2841 if user_api_key_dict.user_id and ( 2841 ↛ 2844line 2841 didn't jump to line 2844 because the condition on line 2841 was never true

2842 user_api_key_dict.user_rpm_limit is not None or user_api_key_dict.user_tpm_limit is not None 

2843 ): 

2844 descriptors.append( 

2845 RateLimitDescriptor( 

2846 key="user", 

2847 value=user_api_key_dict.user_id, 

2848 rate_limit={ 

2849 "requests_per_unit": user_api_key_dict.user_rpm_limit, 

2850 "tokens_per_unit": user_api_key_dict.user_tpm_limit, 

2851 "window_size": self.window_size, 

2852 }, 

2853 ) 

2854 ) 

2855 

2856 # Team rate limits 

2857 if user_api_key_dict.team_id and ( 2857 ↛ 2860line 2857 didn't jump to line 2860 because the condition on line 2857 was never true

2858 user_api_key_dict.team_rpm_limit is not None or user_api_key_dict.team_tpm_limit is not None 

2859 ): 

2860 descriptors.append( 

2861 RateLimitDescriptor( 

2862 key="team", 

2863 value=user_api_key_dict.team_id, 

2864 rate_limit={ 

2865 "requests_per_unit": user_api_key_dict.team_rpm_limit, 

2866 "tokens_per_unit": user_api_key_dict.team_tpm_limit, 

2867 "window_size": self.window_size, 

2868 }, 

2869 ) 

2870 ) 

2871 

2872 # Team Member rate limits 

2873 if user_api_key_dict.user_id and ( 2873 ↛ 2876line 2873 didn't jump to line 2876 because the condition on line 2873 was never true

2874 user_api_key_dict.team_member_rpm_limit is not None or user_api_key_dict.team_member_tpm_limit is not None 

2875 ): 

2876 team_member_value: Final = f"{user_api_key_dict.team_id}:{user_api_key_dict.user_id}" 

2877 descriptors.append( 

2878 RateLimitDescriptor( 

2879 key="team_member", 

2880 value=team_member_value, 

2881 rate_limit={ 

2882 "requests_per_unit": user_api_key_dict.team_member_rpm_limit, 

2883 "tokens_per_unit": user_api_key_dict.team_member_tpm_limit, 

2884 "window_size": self.window_size, 

2885 }, 

2886 ) 

2887 ) 

2888 

2889 # End user rate limits 

2890 if user_api_key_dict.end_user_id and ( 2890 ↛ 2893line 2890 didn't jump to line 2893 because the condition on line 2890 was never true

2891 user_api_key_dict.end_user_rpm_limit is not None or user_api_key_dict.end_user_tpm_limit is not None 

2892 ): 

2893 descriptors.append( 

2894 RateLimitDescriptor( 

2895 key="end_user", 

2896 value=user_api_key_dict.end_user_id, 

2897 rate_limit={ 

2898 "requests_per_unit": user_api_key_dict.end_user_rpm_limit, 

2899 "tokens_per_unit": user_api_key_dict.end_user_tpm_limit, 

2900 "window_size": self.window_size, 

2901 }, 

2902 ) 

2903 ) 

2904 

2905 # Model rate limits 

2906 requested_model: Final = data.get("model", None) 

2907 self._add_model_per_key_rate_limit_descriptor( 

2908 user_api_key_dict=user_api_key_dict, 

2909 requested_model=requested_model, 

2910 descriptors=descriptors, 

2911 ) 

2912 

2913 # Per-request-tag rate limits scoped to this key 

2914 self._add_tag_per_key_rate_limit_descriptor( 

2915 user_api_key_dict=user_api_key_dict, 

2916 data=data, 

2917 descriptors=descriptors, 

2918 ) 

2919 

2920 # REST MCP calls pass the raw body through this hook before server 

2921 # resolution; only the later synthetic hook payload may carry this key. 

2922 if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data: 2922 ↛ 2923line 2922 didn't jump to line 2923 because the condition on line 2922 was never true

2923 mcp_server_name: Final = data.get("mcp_server_name", None) 

2924 self._add_mcp_per_key_rate_limit_descriptor( 

2925 user_api_key_dict=user_api_key_dict, 

2926 mcp_server_name=mcp_server_name, 

2927 descriptors=descriptors, 

2928 ) 

2929 self._add_mcp_per_team_rate_limit_descriptor( 

2930 user_api_key_dict=user_api_key_dict, 

2931 mcp_server_name=mcp_server_name, 

2932 descriptors=descriptors, 

2933 ) 

2934 

2935 self._add_team_model_rate_limit_descriptor_from_metadata( 

2936 user_api_key_dict=user_api_key_dict, 

2937 requested_model=requested_model if isinstance(requested_model, str) else None, 

2938 descriptors=descriptors, 

2939 ) 

2940 

2941 # Agent-level and session-level rate limits 

2942 resolved_agent_id: Final = self._get_resolved_agent_id(user_api_key_dict, data) 

2943 

2944 if resolved_agent_id: 2944 ↛ 2945line 2944 didn't jump to line 2945 because the condition on line 2944 was never true

2945 descriptors.extend( 

2946 self._create_agent_rate_limit_descriptors( 

2947 agent_id=resolved_agent_id, 

2948 data=data, 

2949 ) 

2950 ) 

2951 

2952 return descriptors 

2953 

2954 async def _check_model_has_recent_failures( 

2955 self, 

2956 model: str, 

2957 parent_otel_span: Span | None = None, 

2958 ) -> bool: 

2959 """ 

2960 Check if any deployment for this model has recent failures by using 

2961 the router's existing failure tracking. 

2962 

2963 Returns True if any deployment has failures in the current minute. 

2964 """ 

2965 from litellm.proxy.proxy_server import llm_router 

2966 from litellm.router_utils.router_callbacks.track_deployment_metrics import ( 

2967 get_deployment_failures_for_current_minute, 

2968 ) 

2969 

2970 if llm_router is None: 

2971 return False 

2972 

2973 try: 

2974 # Get all deployments for this model 

2975 model_list: Final = llm_router.get_model_list(model_name=model) 

2976 if not model_list: 

2977 return False 

2978 

2979 # Check each deployment's failure count 

2980 for deployment in model_list: 

2981 deployment_id = deployment.get("model_info", {}).get("id") 

2982 if not deployment_id: 

2983 continue 

2984 

2985 # Use router's existing failure tracking 

2986 failure_count = get_deployment_failures_for_current_minute( 

2987 litellm_router_instance=llm_router, 

2988 deployment_id=deployment_id, 

2989 ) 

2990 

2991 if failure_count > DYNAMIC_RATE_LIMIT_ERROR_THRESHOLD_PER_MINUTE: 

2992 verbose_proxy_logger.debug( 

2993 "[Dynamic Rate Limit] Deployment %s has %s failures in current minute - enforcing rate limits for model %s", 

2994 deployment_id, 

2995 failure_count, 

2996 model, 

2997 ) 

2998 return True 

2999 

3000 verbose_proxy_logger.debug( 

3001 "[Dynamic Rate Limit] No failures detected for model %s - allowing dynamic exceeding", model 

3002 ) 

3003 return False 

3004 

3005 except Exception as e: 

3006 verbose_proxy_logger.debug("Error checking model failure status: %s, defaulting to enforce limits", e) 

3007 # Fail safe: enforce limits if we can't check 

3008 return True 

3009 

3010 def get_rate_limiter_for_call_type(self, call_type: str) -> CallTypeRateLimiter | None: 

3011 """Get the rate limiter for the call type.""" 

3012 if call_type == "acreate_batch": 

3013 batch_limiter: Final = self._get_batch_rate_limiter() 

3014 return batch_limiter 

3015 return None 

3016 

3017 def _key_owns_model_limit( 

3018 self, 

3019 user_api_key_dict: UserAPIKeyAuth, 

3020 model: RateLimitedModel, 

3021 rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"], 

3022 ) -> bool: 

3023 return model.limit_in(get_key_own_model_rate_limit(user_api_key_dict, rate_limit_key)) is not None 

3024 

3025 def _inherited_team_model_limit( 

3026 self, 

3027 user_api_key_dict: UserAPIKeyAuth, 

3028 model: RateLimitedModel, 

3029 rate_limit_key: Literal["model_rpm_limit", "model_tpm_limit"], 

3030 ) -> int | None: 

3031 team_limit: Final = model.limit_in( 

3032 get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", rate_limit_key) 

3033 ) 

3034 if team_limit is None or self._key_owns_model_limit(user_api_key_dict, model, rate_limit_key): 3034 ↛ 3036line 3034 didn't jump to line 3036 because the condition on line 3034 was always true

3035 return None 

3036 return team_limit 

3037 

3038 def _key_owns_model_tpm_limit_from_request_metadata( 

3039 self, 

3040 request_metadata: Mapping[str, object], 

3041 model: RateLimitedModel | None, 

3042 ) -> bool: 

3043 if model is None: 3043 ↛ 3045line 3043 didn't jump to line 3045 because the condition on line 3043 was always true

3044 return False 

3045 key_view: Final = UserAPIKeyAuth.model_validate( 

3046 { 

3047 "metadata": request_metadata.get("user_api_key_metadata") or {}, 

3048 "model_max_budget": request_metadata.get("user_api_key_model_max_budget") or {}, 

3049 } 

3050 ) 

3051 return self._key_owns_model_limit(key_view, model, "model_tpm_limit") 

3052 

3053 def _add_team_model_rate_limit_descriptor_from_metadata( 

3054 self, 

3055 user_api_key_dict: UserAPIKeyAuth, 

3056 requested_model: str | None, 

3057 descriptors: list[RateLimitDescriptor], 

3058 ) -> None: 

3059 model: Final = self._rate_limited_model(requested_model) 

3060 if model is None: 

3061 return 

3062 team_rpm_limit: Final = self._inherited_team_model_limit(user_api_key_dict, model, "model_rpm_limit") 

3063 team_tpm_limit: Final = self._inherited_team_model_limit(user_api_key_dict, model, "model_tpm_limit") 

3064 if team_rpm_limit is None and team_tpm_limit is None: 3064 ↛ 3066line 3064 didn't jump to line 3066 because the condition on line 3064 was always true

3065 return 

3066 descriptors.append( 

3067 RateLimitDescriptor( 

3068 key="model_per_team", 

3069 value=f"{user_api_key_dict.team_id}:{model.group}", 

3070 rate_limit={ 

3071 "requests_per_unit": team_rpm_limit, 

3072 "tokens_per_unit": team_tpm_limit, 

3073 "window_size": self.window_size, 

3074 }, 

3075 ) 

3076 ) 

3077 

3078 def _add_project_model_rate_limit_descriptor_from_metadata( 

3079 self, 

3080 user_api_key_dict: UserAPIKeyAuth, 

3081 requested_model: str | None, 

3082 descriptors: list[RateLimitDescriptor], 

3083 ) -> None: 

3084 """Add project model rate limit descriptor from project_metadata if applicable.""" 

3085 model: Final = self._rate_limited_model(requested_model) 

3086 if model is None: 

3087 return 

3088 model_specific_tpm_limit: Final = model.limit_in( 

3089 get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_tpm_limit") 

3090 ) 

3091 model_specific_rpm_limit: Final = model.limit_in( 

3092 get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_rpm_limit") 

3093 ) 

3094 if model_specific_tpm_limit is None and model_specific_rpm_limit is None: 3094 ↛ 3096line 3094 didn't jump to line 3096 because the condition on line 3094 was always true

3095 return 

3096 descriptors.append( 

3097 RateLimitDescriptor( 

3098 key="model_per_project", 

3099 value=f"{user_api_key_dict.project_id}:{model.group}", 

3100 rate_limit={ 

3101 "requests_per_unit": model_specific_rpm_limit, 

3102 "tokens_per_unit": model_specific_tpm_limit, 

3103 "window_size": self.window_size, 

3104 }, 

3105 ) 

3106 ) 

3107 

3108 def add_project_io_token_rate_limit_descriptors_from_metadata( 

3109 self, 

3110 user_api_key_dict: UserAPIKeyAuth, 

3111 requested_model: str | None, 

3112 descriptors: _RateLimitDescriptorSink, 

3113 ) -> None: 

3114 """Add project-scoped ITPM/OTPM descriptors from project_metadata. 

3115 

3116 Enforced independently of, and alongside, the combined ``model_per_project`` 

3117 TPM descriptor above -- these give Bedrock Mantle-style separate input/output 

3118 token quotas at the project level. 

3119 """ 

3120 model: Final = self._rate_limited_model(requested_model) 

3121 if model is None or user_api_key_dict.project_id is None: 3121 ↛ 3124line 3121 didn't jump to line 3124 because the condition on line 3121 was always true

3122 return 

3123 

3124 model_itpm_limit: Final = model.limit_in( 

3125 get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_itpm_limit") 

3126 ) 

3127 model_otpm_limit: Final = model.limit_in( 

3128 get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_otpm_limit") 

3129 ) 

3130 

3131 if model_itpm_limit is None and model_otpm_limit is None: 

3132 return 

3133 

3134 descriptor_value: Final = f"{user_api_key_dict.project_id}:{model.group}" 

3135 if model_itpm_limit is not None: 

3136 descriptors.append( 

3137 RateLimitDescriptor( 

3138 key=PROJECT_ITPM_DESCRIPTOR_KEY, 

3139 value=descriptor_value, 

3140 rate_limit={ # mutable-ok: descriptor TypedDict requires a runtime dict 

3141 "requests_per_unit": None, 

3142 "tokens_per_unit": model_itpm_limit, 

3143 "window_size": self.window_size, 

3144 }, 

3145 ) 

3146 ) 

3147 if model_otpm_limit is not None: 

3148 descriptors.append( 

3149 RateLimitDescriptor( 

3150 key=PROJECT_OTPM_DESCRIPTOR_KEY, 

3151 value=descriptor_value, 

3152 rate_limit={ # mutable-ok: descriptor TypedDict requires a runtime dict 

3153 "requests_per_unit": None, 

3154 "tokens_per_unit": model_otpm_limit, 

3155 "window_size": self.window_size, 

3156 }, 

3157 ) 

3158 ) 

3159 

3160 def _handle_rate_limit_error( 

3161 self, 

3162 response: RateLimitResponse, 

3163 descriptors: list[RateLimitDescriptor], 

3164 requested_model: str | None = None, 

3165 ) -> None: 

3166 """Handle rate limit exceeded by raising :class:`ProxyRateLimitError` (a 429).""" 

3167 for status in response["statuses"]: 

3168 if status["code"] == "OVER_LIMIT": 

3169 descriptor_key = status["descriptor_key"] 

3170 matching_descriptor = next( 

3171 ( 

3172 desc 

3173 for desc in descriptors 

3174 if desc["key"] == descriptor_key 

3175 and ((status_value := status.get("descriptor_value")) is None or desc["value"] == status_value) 

3176 ), 

3177 None, 

3178 ) 

3179 descriptor_value = matching_descriptor["value"] if matching_descriptor is not None else "unknown" 

3180 

3181 now = self._get_current_time().timestamp() 

3182 reset_time = now + self.window_size 

3183 reset_time_formatted = datetime.fromtimestamp(reset_time, tz=timezone.utc).strftime( 

3184 "%Y-%m-%d %H:%M:%S UTC" 

3185 ) 

3186 

3187 remaining_display = max(0, status["limit_remaining"]) 

3188 rate_limit_type = status["rate_limit_type"] 

3189 current_limit = status["current_limit"] 

3190 

3191 detail = ( 

3192 f"Rate limit exceeded for {descriptor_key}: {descriptor_value}. " 

3193 f"Limit type: {rate_limit_type}. " 

3194 f"Current limit: {current_limit}, Remaining: {remaining_display}. " 

3195 f"Limit resets at: {reset_time_formatted}" 

3196 ) 

3197 

3198 resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(requested_model) 

3199 raise ProxyRateLimitError( 

3200 detail=detail, 

3201 headers={ 

3202 "retry-after": str(self.window_size), 

3203 "rate_limit_type": str(status["rate_limit_type"]), 

3204 "reset_at": reset_time_formatted, 

3205 }, 

3206 rate_limit_type=map_v3_rate_limit_type(status["rate_limit_type"]), 

3207 model=resolved_model, 

3208 llm_provider=llm_provider, 

3209 ) 

3210 

3211 @staticmethod 

3212 def _estimate_audio_block_tokens(block: object) -> int: 

3213 """ 

3214 Token estimate for one ``input_audio`` content block. 

3215 

3216 When the block carries a base64 ``data`` payload, the estimate comes 

3217 from the decoded byte count (``len(b64) * 3 // 4 // _AUDIO_BYTES_PER_TOKEN``), 

3218 assuming the lowest reasonable audio bitrate so we never under-reserve 

3219 for higher-quality recordings of the same duration. 

3220 

3221 When no payload is present (reference-only block or missing ``data``), 

3222 falls back to ``DEFAULT_AUDIO_TOKEN_ESTIMATE``. 

3223 """ 

3224 if not isinstance(block, dict): 

3225 return DEFAULT_AUDIO_TOKEN_ESTIMATE 

3226 input_audio: Final = block.get("input_audio") 

3227 b64_data: Final = input_audio.get("data") if isinstance(input_audio, dict) else None 

3228 if b64_data and isinstance(b64_data, str): 

3229 decoded_bytes: Final = len(b64_data) * 3 // 4 

3230 return max(decoded_bytes // _AUDIO_BYTES_PER_TOKEN, DEFAULT_AUDIO_TOKEN_ESTIMATE) 

3231 return DEFAULT_AUDIO_TOKEN_ESTIMATE 

3232 

3233 @classmethod 

3234 def _estimate_audio_content_tokens(cls, messages: object) -> int: 

3235 """ 

3236 Sum of per-block audio token estimates across all ``messages``. 

3237 Returns 0 when there are no ``input_audio`` blocks, which the caller 

3238 uses to skip the (relatively expensive) strip pass. 

3239 """ 

3240 if not isinstance(messages, list): 

3241 return 0 

3242 return sum( 

3243 cls._estimate_audio_block_tokens(block) 

3244 for message in messages 

3245 if isinstance(message, dict) 

3246 for content in (message.get("content"),) 

3247 if isinstance(content, list) 

3248 for block in content 

3249 if isinstance(block, dict) and block.get("type") == "input_audio" 

3250 ) 

3251 

3252 @staticmethod 

3253 def _strip_audio_content_blocks(messages: object) -> object: 

3254 """ 

3255 Drop ``input_audio`` content blocks before passing ``messages`` to 

3256 ``token_counter``, which raises ``ValueError`` on them (no per-type 

3257 handling, unlike images). The audio contribution is added back 

3258 separately via ``DEFAULT_AUDIO_TOKEN_ESTIMATE`` so the rest of the 

3259 message (text/images/tools) still gets counted accurately instead of 

3260 the whole call falling back to the cheap char-count estimate. 

3261 """ 

3262 if not isinstance(messages, list): 

3263 return messages 

3264 sanitized: Final[list[object]] = [] # mutable-ok: token_counter requires a list of message dicts 

3265 for message in messages: 

3266 if not isinstance(message, dict): 

3267 sanitized.append(message) 

3268 continue 

3269 content = message.get("content") 

3270 if not isinstance(content, list): 

3271 sanitized.append(message) 

3272 continue 

3273 filtered_content = [ # mutable-ok: token_counter requires list content blocks 

3274 block for block in content if not (isinstance(block, dict) and block.get("type") == "input_audio") 

3275 ] 

3276 sanitized.append( 

3277 {**message, "content": filtered_content} # mutable-ok: token_counter requires message dicts 

3278 ) 

3279 return sanitized 

3280 

3281 @staticmethod 

3282 def _responses_input_to_chat_messages(data: object) -> Sequence[object]: 

3283 """ 

3284 Convert a Responses API ``input`` (string or list of input items) into 

3285 chat-completion-style messages via the standard LiteLLM transformation 

3286 (the same one guardrails use, e.g. ``purview_dlp.py``), so multimodal 

3287 ``input_image``/``input_text`` content blocks get counted by 

3288 ``token_counter``'s ``messages`` path instead of silently contributing 

3289 zero tokens via its ``text`` path, which only joins plain strings. 

3290 """ 

3291 from litellm.responses.litellm_completion_transformation.transformation import ( 

3292 LiteLLMCompletionResponsesConfig, 

3293 ) 

3294 

3295 if not isinstance(data, dict): 

3296 return () 

3297 return LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages( 

3298 input=data.get("input") or "", 

3299 responses_api_request=data, 

3300 ) 

3301 

3302 @staticmethod 

3303 def _count_pretokenized_embedding_input(value: object) -> int | None: 

3304 if not isinstance(value, list): 

3305 return None 

3306 if all(isinstance(token, int) for token in value): 

3307 return len(value) 

3308 if all( 

3309 isinstance(token_ids, list) and all(isinstance(token, int) for token in token_ids) for token_ids in value 

3310 ): 

3311 return sum(len(token_ids) for token_ids in value) 

3312 return None 

3313 

3314 @staticmethod 

3315 def _rerank_input_to_text(data: Mapping[str, object]) -> str: 

3316 documents: Final = data.get("documents") 

3317 document_items: Final[Sequence[object]] = documents if isinstance(documents, list) else () # pyright: ignore[reportUnknownVariableType] # rerank documents are validated runtime JSON 

3318 input_parts: Final[tuple[object, ...]] = ( # pyright: ignore[reportUnknownVariableType] # list narrowing preserves unknown JSON element types 

3319 data.get("query"), 

3320 *document_items, 

3321 ) 

3322 return "\n".join( 

3323 str(part) # pyright: ignore[reportUnknownArgumentType] # accepted document dicts have provider-defined fields 

3324 for part in input_parts # pyright: ignore[reportUnknownVariableType] # runtime JSON list elements remain unknown after list narrowing 

3325 if isinstance(part, (str, dict)) 

3326 ) 

3327 

3328 def _estimate_precise_input_tokens(self, data: object, model: str | None, call_type: str | None = None) -> int: 

3329 """ 

3330 Model-aware input token estimate for the project ITPM reservation, 

3331 using ``litellm.token_counter`` -- the same approach the 

3332 deployment-level itpm/otpm check uses in 

3333 ``io_token_rate_limit_check.py``. Unlike the cheap char-count 

3334 estimate the combined-TPM path uses, this accounts for image/tool 

3335 content and derives per-``input_audio``-block estimates from the 

3336 base64 payload size (assuming the lowest reasonable bitrate so 

3337 longer recordings always reserve proportionally more), so a burst 

3338 of multimodal, tool-heavy, or audio-heavy requests can't each 

3339 reserve only the one-token floor and blow past ITPM before 

3340 post-call reconciliation catches up. 

3341 

3342 For the Responses API, ``input`` is converted to chat messages first 

3343 (via ``_responses_input_to_chat_messages``) so its own multimodal 

3344 content blocks are counted the same way; ``token_counter``'s ``text`` 

3345 argument can only see plain strings in a list, not content blocks. 

3346 

3347 Falls back to the cheap char-count estimate if ``token_counter`` 

3348 can't resolve a tokenizer for this model (e.g. an unrecognized 

3349 custom model name) or otherwise raises -- the audio add-on still 

3350 applies on top of the fallback. 

3351 """ 

3352 from litellm import token_counter 

3353 

3354 if not isinstance(data, dict): 

3355 return 0 

3356 is_responses_request: Final = call_type in RESPONSES_API_CALL_TYPES 

3357 translated_request: Final = ( 

3358 None if is_responses_request else self._translate_google_genai_native_request(data, call_type) 

3359 ) 

3360 is_embedding_request: Final = self._is_embedding_request(data, call_type) 

3361 embedding_text: Final = data.get("input") if is_embedding_request else None 

3362 pretokenized_input_tokens: Final = ( 

3363 self._count_pretokenized_embedding_input(embedding_text) if is_embedding_request else None 

3364 ) 

3365 if pretokenized_input_tokens is not None: 

3366 return pretokenized_input_tokens 

3367 

3368 prompt: Final = data.get("prompt") 

3369 fallback_text: Final = prompt if prompt is not None else data.get("input") 

3370 selected_inputs: Final[tuple[object | None, object | None, object | None, object | None]] = ( 

3371 (self._responses_input_to_chat_messages(data), None, data.get("tools"), data.get("tool_choice")) 

3372 if is_responses_request 

3373 else ( 

3374 translated_request.get("messages"), 

3375 None, 

3376 translated_request.get("tools"), 

3377 translated_request.get("tool_choice"), 

3378 ) 

3379 if translated_request is not None 

3380 else (None, embedding_text, data.get("tools"), data.get("tool_choice")) 

3381 if is_embedding_request 

3382 else (None, self._rerank_input_to_text(data), data.get("tools"), data.get("tool_choice")) 

3383 if call_type in RERANK_API_CALL_TYPES 

3384 else (None, prompt, data.get("tools"), data.get("tool_choice")) 

3385 if call_type in TEXT_COMPLETION_API_CALL_TYPES 

3386 else (data.get("messages"), fallback_text, data.get("tools"), data.get("tool_choice")) 

3387 ) 

3388 messages, selected_text, countable_tools, countable_tool_choice = selected_inputs 

3389 

3390 audio_token_estimate: Final = self._estimate_audio_content_tokens(messages) 

3391 countable_messages: Final = self._strip_audio_content_blocks(messages) if audio_token_estimate > 0 else messages 

3392 

3393 try: 

3394 estimate: Final = max( 

3395 0, 

3396 int( 

3397 token_counter( 

3398 model=model or "", 

3399 messages=countable_messages, 

3400 text=selected_text, 

3401 tools=countable_tools, 

3402 tool_choice=countable_tool_choice, 

3403 use_default_image_token_count=True, 

3404 ) 

3405 ), 

3406 ) 

3407 return estimate + audio_token_estimate 

3408 except Exception: # noqa: BLE001 # tokenizer failures degrade to the cheap estimate 

3409 if call_type in RERANK_API_CALL_TYPES and isinstance(selected_text, str): 

3410 return max(0, len(selected_text) // DEFAULT_CHARS_PER_TOKEN) 

3411 estimated_input_tokens, _ = self._estimate_input_and_output_tokens(data=data, call_type=call_type) 

3412 return estimated_input_tokens + audio_token_estimate 

3413 

3414 async def _reserve_project_io_tokens_or_raise( 

3415 self, 

3416 descriptors: Sequence[RateLimitDescriptor], 

3417 data: object, 

3418 requested_model: str | None, 

3419 user_api_key_dict: UserAPIKeyAuth, 

3420 tpm_reservation_scopes: Sequence[tuple[str, str]], 

3421 tpm_reservation_amount: int, 

3422 call_type: str | None = None, 

3423 ) -> None: 

3424 """ 

3425 Reserve project-scoped ITPM/OTPM tokens (Bedrock Mantle-style 

3426 separate input/output token buckets), independently of -- and, when 

3427 both are configured, in addition to -- the combined-TPM reservation 

3428 the caller already made. Raises (via ``_handle_rate_limit_error``) on 

3429 an over-limit reservation, first rolling back the combined-TPM 

3430 reservation named by ``tpm_reservation_scopes``/``tpm_reservation_amount`` 

3431 if one was made, so a partial reservation never leaks. 

3432 """ 

3433 if not isinstance(data, dict): 

3434 return 

3435 stash: Final = claim_request_stash_for_data(data) 

3436 io_token_descriptors: Final = [ # mutable-ok: reservation API requires descriptor lists 

3437 d for d in descriptors if d["key"] in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) 

3438 ] 

3439 if not io_token_descriptors: 

3440 return 

3441 

3442 configured_otpm_limits: Final = [ # mutable-ok: min calculation materializes validated limits 

3443 int(v) 

3444 for d in io_token_descriptors 

3445 if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY 

3446 for v in [ # mutable-ok: comprehension binds the optional descriptor value 

3447 (d.get("rate_limit") or {}).get( # mutable-ok: optional descriptor fallback 

3448 "tokens_per_unit" 

3449 ) 

3450 ] 

3451 if v is not None 

3452 ] 

3453 min_configured_otpm_limit: Final = min(configured_otpm_limits) if configured_otpm_limits else None 

3454 _, raw_estimated_output_tokens = self._estimate_input_and_output_tokens( 

3455 data=data, 

3456 min_configured_tpm_limit=min_configured_otpm_limit, 

3457 call_type=call_type, 

3458 ) 

3459 raw_estimated_input_tokens: Final = await offload_token_count(self._estimate_precise_input_tokens)( 

3460 data=data, model=requested_model, call_type=call_type 

3461 ) 

3462 estimated_input_tokens: Final = max(raw_estimated_input_tokens, 1) 

3463 estimated_output_tokens: Final = ( 

3464 raw_estimated_output_tokens 

3465 if self._has_explicit_output_cap(data, call_type) 

3466 else max(raw_estimated_output_tokens, 1) 

3467 ) 

3468 

3469 # Hard-cap generation length so an unbounded response can't overshoot 

3470 # the OTPM budget before post-call reconciliation runs, mirroring the 

3471 # combined-TPM floor cap in the caller. 

3472 self._apply_implicit_output_cap( 

3473 data=data, 

3474 min_configured_limit=min_configured_otpm_limit, 

3475 call_type=call_type, 

3476 ) 

3477 

3478 io_response, itpm_reserved, otpm_reserved = await self.reserve_io_tokens( 

3479 descriptors=io_token_descriptors, 

3480 estimated_input_tokens=estimated_input_tokens, 

3481 estimated_output_tokens=estimated_output_tokens, 

3482 parent_otel_span=user_api_key_dict.parent_otel_span, 

3483 ) 

3484 

3485 if io_response["overall_code"] == "OVER_LIMIT": 

3486 # A combined-TPM reservation may have already succeeded above for 

3487 # this same request; refund it too, or its counter stays inflated 

3488 # until the window's TTL expires. Mark it released so the 

3489 # ProxyRateLimitError we're about to raise doesn't get refunded 

3490 # a second time when async_post_call_failure_hook sees the same 

3491 # (still-stashed) reservation and refunds it again. 

3492 if tpm_reservation_amount > 0: 

3493 await self._refund_reserved_tokens( 

3494 scopes=tpm_reservation_scopes, 

3495 amount=tpm_reservation_amount, 

3496 parent_otel_span=user_api_key_dict.parent_otel_span, 

3497 ) 

3498 stash.reservation_released = True 

3499 await self._release_stashed_parallel_slot(stash, user_api_key_dict.parent_otel_span) 

3500 self._handle_rate_limit_error( 

3501 response=io_response, 

3502 descriptors=descriptors, 

3503 requested_model=requested_model, 

3504 ) 

3505 

3506 if itpm_reserved > 0: 

3507 itpm_scopes: Final = tuple( 

3508 (d["key"], d["value"]) for d in io_token_descriptors if d["key"] == PROJECT_ITPM_DESCRIPTOR_KEY 

3509 ) 

3510 stash.itpm_reserved_tokens = itpm_reserved 

3511 stash.itpm_reserved_scopes = frozenset(itpm_scopes) 

3512 stash.itpm_reserved_window_identities = frozenset( 

3513 (counter_key, window_start, backend) 

3514 for counter_key, window_start, backend in io_response.get("reservation_windows", frozenset()) 

3515 if "model_per_project_itpm" in counter_key 

3516 ) 

3517 if otpm_reserved > 0: 

3518 otpm_scopes: Final = tuple( 

3519 (d["key"], d["value"]) for d in io_token_descriptors if d["key"] == PROJECT_OTPM_DESCRIPTOR_KEY 

3520 ) 

3521 stash.otpm_reserved_tokens = otpm_reserved 

3522 stash.otpm_reserved_scopes = frozenset(otpm_scopes) 

3523 stash.otpm_reserved_window_identities = frozenset( 

3524 (counter_key, window_start, backend) 

3525 for counter_key, window_start, backend in io_response.get("reservation_windows", frozenset()) 

3526 if "model_per_project_otpm" in counter_key 

3527 ) 

3528 

3529 if stash.rate_limit_response is not None: 

3530 stash.rate_limit_response["statuses"].extend(io_response["statuses"]) 

3531 elif io_response["statuses"]: 

3532 stash.rate_limit_response = io_response 

3533 

3534 verbose_proxy_logger.debug( 

3535 "ITPM/OTPM tokens reserved: itpm=%s, otpm=%s for model %s", 

3536 itpm_reserved, 

3537 otpm_reserved, 

3538 requested_model, 

3539 ) 

3540 

3541 async def _build_request_rate_limit_descriptors( 

3542 self, 

3543 user_api_key_dict: UserAPIKeyAuth, 

3544 data: Mapping[str, object], 

3545 call_type: str | None, 

3546 ) -> list[RateLimitDescriptor]: # mutable-ok: the shared generation reservation helpers require a list 

3547 metadata: Final = _REQUEST_RATE_LIMIT_DATA.validate_python( 

3548 user_api_key_dict.metadata or MappingProxyType({}) # pyright: ignore[reportUnknownMemberType] # validates the legacy auth metadata boundary 

3549 ) 

3550 rpm_value: Final = metadata.get("rpm_limit_type") 

3551 tpm_value: Final = metadata.get("tpm_limit_type") 

3552 rpm_limit_type: Final = rpm_value if isinstance(rpm_value, str) else None 

3553 tpm_limit_type: Final = tpm_value if isinstance(tpm_value, str) else None 

3554 model_value: Final = data.get("model") 

3555 requested_model: Final = model_value if isinstance(model_value, str) else None 

3556 model_has_failures: Final = ( 

3557 await self._check_model_has_recent_failures( 

3558 model=requested_model, 

3559 parent_otel_span=user_api_key_dict.parent_otel_span, 

3560 ) 

3561 if requested_model and self._is_dynamic_rate_limiting_enabled(rpm_limit_type, tpm_limit_type) 

3562 else False 

3563 ) 

3564 descriptors: Final = self._create_rate_limit_descriptors( # pyright: ignore[reportUnknownMemberType] # legacy helper reads a dictionary with validated keys 

3565 user_api_key_dict=user_api_key_dict, 

3566 data=dict(data), # mutable-ok: legacy descriptor helpers accept a request dictionary 

3567 rpm_limit_type=rpm_limit_type, 

3568 tpm_limit_type=tpm_limit_type, 

3569 model_has_failures=model_has_failures, 

3570 call_type=call_type, 

3571 ) 

3572 self._add_project_model_rate_limit_descriptor_from_metadata( 

3573 user_api_key_dict=user_api_key_dict, 

3574 requested_model=requested_model, 

3575 descriptors=descriptors, 

3576 ) 

3577 self.add_project_io_token_rate_limit_descriptors_from_metadata( 

3578 user_api_key_dict=user_api_key_dict, 

3579 requested_model=requested_model, 

3580 descriptors=descriptors, 

3581 ) 

3582 return [ # mutable-ok: the shared generation reservation helpers require a list 

3583 *descriptors, 

3584 *self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model), 

3585 *await self._create_tag_rate_limit_descriptors(data), 

3586 ] 

3587 

3588 async def _release_request_capacity_when_admitted( 

3589 self, 

3590 admission: asyncio.Task[RateLimitResponse], 

3591 acquisition: ParallelSlotAcquisition, 

3592 user_api_key_dict: UserAPIKeyAuth, 

3593 ) -> None: 

3594 response: Final = await admission 

3595 if response["overall_code"] == "OK": 

3596 await self._release_parallel_request_slots(acquisition, user_api_key_dict.parent_otel_span) 

3597 

3598 @asynccontextmanager 

3599 async def request_capacity( 

3600 self, 

3601 user_api_key_dict: UserAPIKeyAuth, 

3602 model: str, 

3603 *, 

3604 request_data: Mapping[str, object] | None = None, 

3605 ) -> AsyncGenerator[None, None]: 

3606 """Charge one non-generation provider request to RPM and hold its concurrency slot.""" 

3607 data: Final = MappingProxyType({**(request_data or MappingProxyType({})), "model": model}) 

3608 descriptors: Final = await self._build_request_rate_limit_descriptors(user_api_key_dict, data, None) 

3609 acquisition: Final = ParallelSlotAcquisition( 

3610 slot_id=uuid.uuid4().hex, 

3611 counter_keys=[ # mutable-ok: the shared slot-release contract requires a list 

3612 self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") 

3613 for d in descriptors 

3614 if d["rate_limit"] is not None and d["rate_limit"].get("max_parallel_requests") is not None 

3615 ], 

3616 ) 

3617 admission: Final = asyncio.create_task( 

3618 self.should_rate_limit( 

3619 descriptors=descriptors, 

3620 parent_otel_span=user_api_key_dict.parent_otel_span, 

3621 skip_tpm_check=True, 

3622 parallel_slot_id=acquisition["slot_id"], 

3623 ) 

3624 ) 

3625 try: 

3626 response: Final = await asyncio.shield(admission) 

3627 if response["overall_code"] == "OVER_LIMIT": 

3628 self._handle_rate_limit_error(response, descriptors, model) 

3629 yield 

3630 finally: 

3631 cleanup: Final = asyncio.create_task( 

3632 self._release_request_capacity_when_admitted(admission, acquisition, user_api_key_dict) 

3633 ) 

3634 cancellation: asyncio.CancelledError | None = None # rebind-ok: retain cancellation until cleanup finishes 

3635 while not cleanup.done(): 

3636 try: 

3637 await asyncio.shield(cleanup) 

3638 except asyncio.CancelledError as exc: 

3639 cancellation = exc 

3640 cleanup.result() 

3641 if cancellation is not None: 

3642 raise cancellation 

3643 

3644 async def async_pre_call_hook( 

3645 self, 

3646 user_api_key_dict: UserAPIKeyAuth, 

3647 cache: DualCache, 

3648 data: dict, 

3649 call_type: str, 

3650 ): 

3651 """ 

3652 Pre-call hook to check rate limits before making the API call. 

3653 Supports dynamic rate limiting based on deployment health. 

3654 """ 

3655 verbose_proxy_logger.debug("Inside Rate Limit Pre-Call Hook") 

3656 

3657 stash: Final = claim_request_stash_for_data(data) 

3658 

3659 ######################################################### 

3660 # Check if the call type has a specific rate limiter 

3661 # eg. for Batch APIs we need to use the batch rate limiter to read the input file and count the tokens and requests 

3662 ######################################################### 

3663 call_type_specific_rate_limiter: Final = self.get_rate_limiter_for_call_type(call_type=call_type) 

3664 if call_type_specific_rate_limiter: 

3665 return await call_type_specific_rate_limiter.async_pre_call_hook( 

3666 user_api_key_dict=user_api_key_dict, 

3667 cache=cache, 

3668 data=data, 

3669 call_type=call_type, 

3670 ) 

3671 

3672 request_data: Final = _REQUEST_RATE_LIMIT_DATA.validate_python(data) 

3673 model_value: Final = request_data.get("model") 

3674 requested_model: Final = model_value if isinstance(model_value, str) else None 

3675 descriptors: Final = await self._build_request_rate_limit_descriptors( 

3676 user_api_key_dict=user_api_key_dict, 

3677 data=request_data, 

3678 call_type=call_type, 

3679 ) 

3680 stash.tpm_limited_tags = frozenset( 

3681 d["value"] 

3682 for d in descriptors 

3683 if d["key"] == "tag" and d["rate_limit"] is not None and d["rate_limit"].get("tokens_per_unit") is not None 

3684 ) 

3685 

3686 # Only check rate limits if we have descriptors with actual limits 

3687 if descriptors: 3687 ↛ 3702line 3687 didn't jump to line 3702 because the condition on line 3687 was never true

3688 # First pass: RPM and max_parallel_requests sliding-window check. 

3689 # When reservation is enabled, `skip_tpm_check=True` tells 

3690 # should_rate_limit to ignore each descriptor's tokens_per_unit so 

3691 # its +1-per-key Lua / in-memory increment never touches the 

3692 # :tokens counters — those are owned exclusively by the atomic 

3693 # reserve_tpm_tokens path below. Without this, every concurrent 

3694 # in-flight request would pre-inflate the :tokens counter by 1, 

3695 # shrinking the effective TPM budget by N and causing 

3696 # false-positive 429s under bursts. When reservation is disabled, 

3697 # this pass enforces TPM directly from the post-call counters -- 

3698 # except for project ITPM/OTPM descriptors, which are excluded 

3699 # then because _reserve_project_io_tokens_or_raise below charges 

3700 # them unconditionally and counting them here too would 

3701 # double-charge every request. 

3702 parallel_counter_keys: Final = [ 

3703 self.create_rate_limit_keys(d["key"], d["value"], "max_parallel_requests") 

3704 for d in descriptors 

3705 if (d.get("rate_limit") or {}).get("max_parallel_requests") is not None 

3706 ] 

3707 parallel_slot_id: Final = uuid.uuid4().hex if parallel_counter_keys else None 

3708 

3709 first_pass_descriptors: Final = ( 

3710 descriptors 

3711 if self.tpm_reservation_enabled 

3712 else tuple( 

3713 d for d in descriptors if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) 

3714 ) 

3715 ) 

3716 response: Final = await self.should_rate_limit( 

3717 descriptors=first_pass_descriptors, 

3718 parent_otel_span=user_api_key_dict.parent_otel_span, 

3719 skip_tpm_check=self.tpm_reservation_enabled, 

3720 parallel_slot_id=parallel_slot_id, 

3721 ) 

3722 

3723 if response["overall_code"] == "OVER_LIMIT": 

3724 self._handle_rate_limit_error( 

3725 response=response, 

3726 descriptors=descriptors, 

3727 requested_model=requested_model, 

3728 ) 

3729 else: 

3730 stash.rate_limit_response = response 

3731 if parallel_slot_id is not None: 

3732 stash.parallel_slot = ParallelSlotAcquisition( 

3733 slot_id=parallel_slot_id, 

3734 counter_keys=parallel_counter_keys, 

3735 ) 

3736 

3737 # ---------------------------------------------------------------- 

3738 # TPM token reservation 

3739 # Atomically reserve estimated tokens upfront so concurrent 

3740 # requests cannot all observe "under limit" before any of them 

3741 # has incremented the counter. atomic_check_and_increment_by_n 

3742 # uses Redis Lua when available and falls back to an asyncio-locked 

3743 # in-memory check otherwise — single-worker protection still holds 

3744 # even without Redis. 

3745 # ---------------------------------------------------------------- 

3746 configured_tpm_limits: Final = [ 

3747 int(v) 

3748 for d in descriptors 

3749 if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) 

3750 for v in [(d.get("rate_limit") or {}).get("tokens_per_unit")] 

3751 if v is not None 

3752 ] 

3753 has_tpm_limits: Final = bool(configured_tpm_limits) 

3754 

3755 # Populated on a successful combined-TPM reservation below, so the 

3756 # project ITPM/OTPM block further down can roll it back if a 

3757 # different bucket in the same request subsequently hits its 

3758 # limit. Stays empty/0 whenever no combined-TPM reservation was 

3759 # made (or it was over limit, in which case execution never 

3760 # reaches the ITPM/OTPM block -- `_handle_rate_limit_error` raises). 

3761 tpm_reservation_scopes: Sequence[tuple[str, str]] = () # rebind-ok: set after successful reservation 

3762 tpm_reservation_amount = 0 # rebind-ok: set after successful reservation 

3763 

3764 if has_tpm_limits and self.tpm_reservation_enabled: 

3765 min_configured_tpm_limit: Final = min(configured_tpm_limits) 

3766 

3767 configured_output_tokens: Final = get_estimated_output_tokens( 

3768 user_api_key_dict=user_api_key_dict, 

3769 model_name=requested_model, 

3770 ) 

3771 

3772 # When the configured TPM cap is small enough to constrain the 

3773 # no-max_tokens floor, also hard-cap the model output so 

3774 # concurrent unbounded generations can't spend past the limit 

3775 # before post-call reconciliation runs. 

3776 self._apply_implicit_output_cap( 

3777 data=data, 

3778 min_configured_limit=min_configured_tpm_limit, 

3779 call_type=call_type, 

3780 configured_output_tokens=configured_output_tokens, 

3781 ) 

3782 

3783 # Floor at 1 token so contentless requests (/responses, 

3784 # tool-call continuations, empty messages) still flow 

3785 # through the atomic counter and get backpressure when at 

3786 # limit. Without this floor, N concurrent contentless 

3787 # requests would all pass pre-call with no enforcement. 

3788 # Post-call reconciliation refunds the over-reservation 

3789 # delta when actual usage comes in below the floor. 

3790 estimated_tokens: Final = max( 

3791 self._estimate_tokens_for_request( 

3792 data=data, 

3793 model=requested_model, 

3794 min_configured_tpm_limit=min_configured_tpm_limit, 

3795 call_type=call_type, 

3796 configured_output_tokens=configured_output_tokens, 

3797 ), 

3798 1, 

3799 ) 

3800 

3801 if configured_output_tokens is not None and estimated_tokens > min_configured_tpm_limit: 

3802 verbose_proxy_logger.debug( 

3803 "Reserving %s tokens for model %s (declared %s=%s plus the input estimate) exceeds the " 

3804 "smallest TPM limit this request is charged against (%s), so it cannot be admitted even " 

3805 "against an empty window. Lower the declared estimate or raise the TPM limit.", 

3806 estimated_tokens, 

3807 requested_model, 

3808 ESTIMATED_OUTPUT_TOKENS_FIELD, 

3809 configured_output_tokens, 

3810 min_configured_tpm_limit, 

3811 ) 

3812 

3813 tpm_response: Final = await self.reserve_tpm_tokens( 

3814 descriptors=descriptors, 

3815 estimated_tokens=estimated_tokens, 

3816 parent_otel_span=user_api_key_dict.parent_otel_span, 

3817 ) 

3818 

3819 if tpm_response["overall_code"] == "OVER_LIMIT": 

3820 await self._release_stashed_parallel_slot(stash, user_api_key_dict.parent_otel_span) 

3821 self._handle_rate_limit_error( 

3822 response=tpm_response, 

3823 descriptors=descriptors, 

3824 requested_model=requested_model, 

3825 ) 

3826 else: 

3827 # Capture the exact (key, value) scopes the reservation 

3828 # incremented so post-call reconciliation only applies 

3829 # the (actual - reserved) delta to those — unreserved 

3830 # scopes get charged the full actual usage instead. 

3831 stash.reserved_tokens = estimated_tokens 

3832 stash.reserved_model = self._rate_limited_model(requested_model) 

3833 stash.reserved_scopes = frozenset( 

3834 (d["key"], d["value"]) 

3835 for d in descriptors 

3836 if d["key"] not in (PROJECT_ITPM_DESCRIPTOR_KEY, PROJECT_OTPM_DESCRIPTOR_KEY) 

3837 and (d.get("rate_limit") or {}).get( # mutable-ok: optional descriptor fallback 

3838 "tokens_per_unit" 

3839 ) 

3840 is not None 

3841 ) 

3842 tpm_reservation_scopes = tuple( # rebind-ok: record successful reservation scopes 

3843 stash.reserved_scopes 

3844 ) 

3845 tpm_reservation_amount = estimated_tokens # rebind-ok: record successful reservation amount 

3846 

3847 # Merge TPM statuses into the stored rate-limit response 

3848 # so x-ratelimit-{key}-remaining-tokens / -limit-tokens 

3849 # headers reach the client. Without this, the RPM-only 

3850 # response from should_rate_limit (skip_tpm_check=True) 

3851 # silently drops all token headers. 

3852 stored_response: Final = stash.rate_limit_response 

3853 if stored_response is not None: 

3854 stored_response["statuses"].extend(tpm_response["statuses"]) 

3855 

3856 verbose_proxy_logger.debug( 

3857 "TPM tokens reserved: %s for model %s", estimated_tokens, requested_model 

3858 ) 

3859 await self._reserve_project_io_tokens_or_raise( 

3860 descriptors=descriptors, 

3861 data=data, 

3862 requested_model=requested_model, 

3863 user_api_key_dict=user_api_key_dict, 

3864 tpm_reservation_scopes=tpm_reservation_scopes, 

3865 tpm_reservation_amount=tpm_reservation_amount, 

3866 call_type=call_type, 

3867 ) 

3868 

3869 def _create_pipeline_operations( 

3870 self, 

3871 key: str, 

3872 value: str, 

3873 rate_limit_type: Literal["requests", "tokens", "max_parallel_requests"], 

3874 total_tokens: int, 

3875 ) -> list["RedisPipelineIncrementOperation"]: 

3876 """ 

3877 Create pipeline operations for TPM increments 

3878 """ 

3879 pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [] 

3880 counter_key: Final = self.create_rate_limit_keys( 

3881 key=key, 

3882 value=value, 

3883 rate_limit_type="tokens", 

3884 ) 

3885 pipeline_operations.append( 

3886 RedisPipelineIncrementOperation( 

3887 key=counter_key, 

3888 increment_value=total_tokens, 

3889 ttl=self.window_size, 

3890 ) 

3891 ) 

3892 

3893 return pipeline_operations 

3894 

3895 def _get_total_tokens_from_usage( 

3896 self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"] 

3897 ) -> int: 

3898 """ 

3899 Get total tokens from response usage for rate limiting. 

3900 

3901 For 'input' and 'total' rate limit types, cached tokens are excluded 

3902 because providers like AWS Bedrock don't count cached tokens toward 

3903 rate limits. This aligns LiteLLM's TPM calculation with provider behavior. 

3904 """ 

3905 total_tokens = 0 

3906 cached_tokens = 0 

3907 

3908 if usage: 3908 ↛ 3909line 3908 didn't jump to line 3909 because the condition on line 3908 was never true

3909 if isinstance(usage, Usage): 

3910 if rate_limit_type == "output": 

3911 total_tokens = usage.completion_tokens or 0 

3912 elif rate_limit_type == "input": 

3913 total_tokens = usage.prompt_tokens or 0 

3914 elif rate_limit_type == "total": 

3915 total_tokens = usage.total_tokens or 0 

3916 

3917 # Get cached tokens to exclude from input/total 

3918 if rate_limit_type in ("input", "total"): 

3919 if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details is not None: 

3920 cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0 

3921 

3922 elif isinstance(usage, dict): 

3923 # Responses API usage comes as a dict 

3924 if rate_limit_type == "output": 

3925 total_tokens = usage.get("completion_tokens", 0) or 0 

3926 elif rate_limit_type == "input": 

3927 total_tokens = usage.get("prompt_tokens", 0) or 0 

3928 elif rate_limit_type == "total": 

3929 total_tokens = usage.get("total_tokens", 0) or 0 

3930 

3931 # Get cached tokens from dict 

3932 if rate_limit_type in ("input", "total"): 

3933 prompt_details: Final = usage.get("prompt_tokens_details") or {} 

3934 if isinstance(prompt_details, dict): 

3935 cached_tokens = prompt_details.get("cached_tokens", 0) or 0 

3936 

3937 # Subtract cached tokens for input/total (providers don't count them) 

3938 if cached_tokens > 0: 3938 ↛ 3939line 3938 didn't jump to line 3939 because the condition on line 3938 was never true

3939 total_tokens = max(0, total_tokens - cached_tokens) 

3940 

3941 return total_tokens 

3942 

3943 @staticmethod 

3944 def _aggregate_only_total_tokens(usage: Usage | ResponseAPIUsage | Mapping[str, object] | None) -> int: 

3945 """Total for usage that carries no input/output split, else 0. 

3946 

3947 A source that can only report one number for the whole request (a 

3948 pass-through target pricing its own multi-model call) charges that 

3949 number under every ``token_rate_limit_type``. Splitting it is 

3950 impossible, and reading 0 out of it would leave the window 

3951 uncharged, which is how pass-through traffic slips past a TPM limit 

3952 it is supposed to share. 

3953 """ 

3954 if usage is None: 3954 ↛ 3956line 3954 didn't jump to line 3956 because the condition on line 3954 was always true

3955 return 0 

3956 token_counts: Final = ( 

3957 (usage.prompt_tokens or 0, usage.completion_tokens or 0, usage.total_tokens or 0) 

3958 if isinstance(usage, Usage) 

3959 else (usage.input_tokens or 0, usage.output_tokens or 0, usage.total_tokens or 0) 

3960 if isinstance(usage, ResponseAPIUsage) 

3961 else ( 

3962 usage.get("prompt_tokens") or usage.get("input_tokens") or 0, 

3963 usage.get("completion_tokens") or usage.get("output_tokens") or 0, 

3964 usage.get("total_tokens") or 0, 

3965 ) 

3966 ) 

3967 prompt_tokens, completion_tokens, total_tokens = token_counts 

3968 if prompt_tokens or completion_tokens or not isinstance(total_tokens, int): 

3969 return 0 

3970 return total_tokens 

3971 

3972 @staticmethod 

3973 def _response_usage( 

3974 response_obj: object, 

3975 ) -> Usage | ResponseAPIUsage | Mapping[str, object] | None: 

3976 if isinstance(response_obj, (Usage, ResponseAPIUsage)): 

3977 return response_obj 

3978 if isinstance( 

3979 response_obj, 

3980 (ModelResponse, EmbeddingResponse, TextCompletionResponse, BaseLiteLLMOpenAIResponseObject), 

3981 ): 

3982 usage: Final = getattr(response_obj, "usage", None) 

3983 return usage if isinstance(usage, (Usage, ResponseAPIUsage, dict)) else None 

3984 if isinstance(response_obj, dict): 

3985 nested_usage: Final = response_obj.get("usage") 

3986 if isinstance(nested_usage, (Usage, ResponseAPIUsage, dict)): 

3987 return nested_usage 

3988 return response_obj 

3989 return None 

3990 

3991 async def _execute_token_increment_script( 

3992 self, 

3993 pipeline_operations: list["RedisPipelineIncrementOperation"], 

3994 ) -> None: 

3995 """ 

3996 Execute token increment script grouped by hash tag for cluster compatibility. 

3997 """ 

3998 if self.token_increment_script is None: 

3999 return 

4000 

4001 # Group operations by hash tag for Redis cluster compatibility 

4002 operation_keys: Final = [op["key"] for op in pipeline_operations] 

4003 key_groups: Final = self._group_keys_by_hash_tag(operation_keys) 

4004 

4005 for _hash_tag, group_keys in key_groups.items(): 

4006 # Get operations for this hash tag group 

4007 group_operations = [op for op in pipeline_operations if op["key"] in group_keys] 

4008 

4009 keys = [] 

4010 args = [] 

4011 

4012 for op in group_operations: 

4013 # Convert None TTL to 0 for Lua script 

4014 ttl_value = op["ttl"] if op["ttl"] is not None else 0 

4015 

4016 verbose_proxy_logger.debug( 

4017 "Executing TTL-preserving increment for key=%s, increment=%s, ttl=%s", 

4018 op["key"], 

4019 op["increment_value"], 

4020 ttl_value, 

4021 ) 

4022 keys.append(op["key"]) 

4023 args.extend([op["increment_value"], ttl_value]) 

4024 

4025 await self.token_increment_script( 

4026 keys=keys, 

4027 args=args, 

4028 ) 

4029 

4030 async def async_increment_tokens_with_ttl_preservation( 

4031 self, 

4032 pipeline_operations: list["RedisPipelineIncrementOperation"], 

4033 parent_otel_span: Span | None = None, 

4034 ) -> None: 

4035 """ 

4036 Increment token counters using Lua script to preserve existing TTL. 

4037 This prevents TTL reset on every token increment. 

4038 """ 

4039 if not pipeline_operations: 

4040 return 

4041 

4042 # Check if script is available 

4043 if self.token_increment_script is None: 

4044 verbose_proxy_logger.debug("TTL preservation script not available, using regular pipeline") 

4045 await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( 

4046 increment_list=pipeline_operations, 

4047 litellm_parent_otel_span=parent_otel_span, 

4048 ) 

4049 return 

4050 

4051 try: 

4052 await self._execute_token_increment_script(pipeline_operations) 

4053 

4054 verbose_proxy_logger.debug( 

4055 "Successfully executed TTL-preserving increment for %s keys", len(pipeline_operations) 

4056 ) 

4057 

4058 except Exception as e: 

4059 log_redis_failure( 

4060 verbose_proxy_logger, logging.WARNING, "TTL preservation failed, falling back to regular pipeline", e 

4061 ) 

4062 # Fallback to regular pipeline on error 

4063 await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( 

4064 increment_list=pipeline_operations, 

4065 litellm_parent_otel_span=parent_otel_span, 

4066 ) 

4067 

4068 async def _apply_local_window_guarded_token_increments( 

4069 self, 

4070 operations: Sequence[ReservationAwareIncrementOperation], 

4071 parent_otel_span: Span | None = None, 

4072 ) -> None: 

4073 async with self._check_and_increment_lock: 

4074 for operation in operations: 

4075 window_key = operation.get("window_key") 

4076 expected_window_start = operation.get("expected_window_start") 

4077 if window_key is None or expected_window_start is None: 

4078 continue 

4079 active_window_start: CacheCounterValue | None = await self.internal_usage_cache.async_get_cache( 

4080 key=window_key, 

4081 litellm_parent_otel_span=parent_otel_span, 

4082 local_only=True, 

4083 ) 

4084 if active_window_start is None or str(active_window_start) != expected_window_start: 

4085 continue 

4086 current_counter = ( 

4087 await self.internal_usage_cache.async_get_cache( 

4088 key=operation["key"], 

4089 litellm_parent_otel_span=parent_otel_span, 

4090 local_only=True, 

4091 ) 

4092 or 0 

4093 ) 

4094 await self.internal_usage_cache.async_set_cache( 

4095 key=operation["key"], 

4096 value=float(current_counter) + operation["increment_value"], 

4097 ttl=operation["ttl"], 

4098 litellm_parent_otel_span=parent_otel_span, 

4099 local_only=True, 

4100 ) 

4101 

4102 async def _apply_redis_window_guarded_token_increments( 

4103 self, 

4104 operations: Sequence[ReservationAwareIncrementOperation], 

4105 parent_otel_span: Span | None = None, 

4106 ) -> None: 

4107 for operation in operations: 

4108 window_key = operation.get("window_key") 

4109 expected_window_start = operation.get("expected_window_start") 

4110 if window_key is None or expected_window_start is None: 

4111 continue 

4112 if self.window_guarded_token_increment_script is not None: 

4113 try: 

4114 await self.window_guarded_token_increment_script( 

4115 keys=[ # mutable-ok: Redis script interface requires a key list 

4116 window_key, 

4117 operation["key"], 

4118 ], 

4119 args=[ # mutable-ok: Redis script interface requires an argument list 

4120 expected_window_start, 

4121 operation["increment_value"], 

4122 operation["ttl"] or 0, 

4123 ], 

4124 ) 

4125 continue 

4126 except Exception as e: # noqa: BLE001 # Redis failures use the plain increment fallback 

4127 log_redis_failure( 

4128 verbose_proxy_logger, 

4129 logging.WARNING, 

4130 f"Window-guarded token adjustment failed for {operation['key']}", 

4131 e, 

4132 ) 

4133 if operation["increment_value"] > 0: 

4134 await self.internal_usage_cache.async_increment_cache( 

4135 key=operation["key"], 

4136 value=operation["increment_value"], 

4137 litellm_parent_otel_span=parent_otel_span, 

4138 ttl=operation["ttl"], 

4139 ) 

4140 

4141 async def async_increment_reservation_aware_tokens( 

4142 self, 

4143 pipeline_operations: Sequence[ReservationAwareIncrementOperation], 

4144 parent_otel_span: Span | None = None, 

4145 ) -> None: 

4146 for operation in pipeline_operations: 

4147 if operation.get("window_key") is None or operation.get("expected_window_start") is None: 

4148 await self.internal_usage_cache.async_increment_cache( 

4149 key=operation["key"], 

4150 value=operation["increment_value"], 

4151 litellm_parent_otel_span=parent_otel_span, 

4152 ttl=operation["ttl"], 

4153 ) 

4154 local_guarded_operations: Final = tuple( 

4155 operation 

4156 for operation in pipeline_operations 

4157 if operation.get("window_key") is not None 

4158 and operation.get("expected_window_start") is not None 

4159 and operation.get("reservation_backend") == "local" 

4160 ) 

4161 redis_guarded_operations: Final = tuple( 

4162 operation 

4163 for operation in pipeline_operations 

4164 if operation.get("window_key") is not None 

4165 and operation.get("expected_window_start") is not None 

4166 and operation.get("reservation_backend") != "local" 

4167 ) 

4168 if local_guarded_operations: 

4169 await self._apply_local_window_guarded_token_increments( 

4170 operations=local_guarded_operations, 

4171 parent_otel_span=parent_otel_span, 

4172 ) 

4173 if redis_guarded_operations: 

4174 await self._apply_redis_window_guarded_token_increments( 

4175 operations=redis_guarded_operations, 

4176 parent_otel_span=parent_otel_span, 

4177 ) 

4178 

4179 def get_rate_limit_type(self) -> Literal["output", "input", "total"]: 

4180 from litellm.proxy.proxy_server import general_settings 

4181 

4182 specified_rate_limit_type: Final = general_settings.get("token_rate_limit_type", "total") 

4183 if specified_rate_limit_type not in [ 4183 ↛ 4188line 4183 didn't jump to line 4188 because the condition on line 4183 was never true

4184 "output", 

4185 "input", 

4186 "total", 

4187 ]: 

4188 return "total" # default to total 

4189 return specified_rate_limit_type 

4190 

4191 @staticmethod 

4192 def _merge_ratelimit_statuses_into_additional_headers( 

4193 additional_headers: dict[str, object], 

4194 statuses: list[RateLimitStatus], 

4195 ) -> dict[str, object]: 

4196 """ 

4197 Return ``additional_headers`` extended with 

4198 ``x-ratelimit-{descriptor_key}-{remaining|limit}-{rate_limit_type}`` 

4199 entries. Non-mutating so callers pick their own target dict. 

4200 """ 

4201 merged: Final[dict[str, object]] = dict(additional_headers) 

4202 for status in statuses: 

4203 prefix = f"x-ratelimit-{status['descriptor_key']}" 

4204 merged[f"{prefix}-remaining-{status['rate_limit_type']}"] = status["limit_remaining"] 

4205 merged[f"{prefix}-limit-{status['rate_limit_type']}"] = status["current_limit"] 

4206 return merged 

4207 

4208 @staticmethod 

4209 def _resolve_rerank_token_usage(response_obj: object) -> tuple[int, int, bool] | None: 

4210 if not isinstance(response_obj, RerankResponse) or response_obj.meta is None: 

4211 return None 

4212 

4213 rerank_tokens: Final = response_obj.meta.get("tokens") # pyright: ignore[reportUnknownMemberType] # TypedDict's optional generic metadata widens get overloads 

4214 if rerank_tokens is not None: 

4215 input_tokens: Final = rerank_tokens.get("input_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # token fields are typed integers despite the generic get overload 

4216 output_tokens: Final = rerank_tokens.get("output_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # token fields are typed integers despite the generic get overload 

4217 if input_tokens or output_tokens: 

4218 return max(0, input_tokens), max(0, output_tokens), True 

4219 

4220 billed_units: Final = response_obj.meta.get("billed_units") # pyright: ignore[reportUnknownMemberType] # TypedDict's optional generic metadata widens get overloads 

4221 if billed_units is not None: 

4222 total_tokens: Final = billed_units.get("total_tokens") or 0 # pyright: ignore[reportUnknownMemberType] # billed total is a typed integer despite the generic get overload 

4223 if total_tokens: 

4224 return max(0, total_tokens), 0, True 

4225 return None 

4226 

4227 def _resolve_io_token_reconcile_usage( 

4228 self, 

4229 response_obj: object, 

4230 ) -> tuple[int, int, bool]: 

4231 """ 

4232 Resolve ``(billable_input_tokens, completion_tokens, usage_resolved)`` 

4233 for ITPM/OTPM reconciliation. Cache-read tokens are excluded from 

4234 billable input -- Bedrock Mantle doesn't count them toward ITPM -- 

4235 but they're untouched everywhere else (cost/usage logging still sees 

4236 the full prompt token count). 

4237 """ 

4238 rerank_usage: Final = self._resolve_rerank_token_usage(response_obj) 

4239 if rerank_usage is not None: 

4240 return rerank_usage 

4241 

4242 usage: Final = self._response_usage(response_obj) 

4243 

4244 if isinstance(usage, Usage): 

4245 prompt_tokens: Final = usage.prompt_tokens or 0 

4246 completion_tokens: Final = usage.completion_tokens or 0 

4247 cached_tokens: Final = ( 

4248 getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0 

4249 if usage.prompt_tokens_details is not None 

4250 else 0 

4251 ) 

4252 if prompt_tokens == 0 and completion_tokens == 0: 

4253 return 0, 0, False 

4254 return max(0, prompt_tokens - cached_tokens), completion_tokens, True 

4255 

4256 if isinstance(usage, ResponseAPIUsage): 

4257 response_input_tokens: Final = usage.input_tokens or 0 

4258 response_output_tokens: Final = usage.output_tokens or 0 

4259 response_cached_tokens: Final = ( 

4260 usage.input_tokens_details.cached_tokens or 0 if usage.input_tokens_details is not None else 0 

4261 ) 

4262 if response_input_tokens == 0 and response_output_tokens == 0: 

4263 return 0, 0, False 

4264 return max(0, response_input_tokens - response_cached_tokens), response_output_tokens, True 

4265 

4266 if isinstance(usage, Mapping): 

4267 raw_prompt_tokens: Final = usage.get("prompt_tokens") or usage.get("input_tokens") or 0 

4268 raw_completion_tokens: Final = usage.get("completion_tokens") or usage.get("output_tokens") or 0 

4269 mapped_prompt_tokens: Final = raw_prompt_tokens if isinstance(raw_prompt_tokens, int) else 0 

4270 mapped_completion_tokens: Final = raw_completion_tokens if isinstance(raw_completion_tokens, int) else 0 

4271 prompt_details: Final = usage.get("prompt_tokens_details") or usage.get("input_tokens_details") 

4272 raw_cached_tokens: Final = ( 

4273 (prompt_details.get("cached_tokens", 0) if isinstance(prompt_details, dict) else 0) 

4274 or usage.get("cache_read_input_tokens") 

4275 or 0 

4276 ) 

4277 mapped_cached_tokens: Final = raw_cached_tokens if isinstance(raw_cached_tokens, int) else 0 

4278 if mapped_prompt_tokens == 0 and mapped_completion_tokens == 0: 

4279 return 0, 0, False 

4280 return max(0, mapped_prompt_tokens - mapped_cached_tokens), mapped_completion_tokens, True 

4281 

4282 return 0, 0, False 

4283 

4284 def _build_io_token_reservation_ops( 

4285 self, 

4286 kwargs: object, 

4287 response_obj: object, 

4288 ) -> Sequence[RedisPipelineIncrementOperation]: 

4289 """ 

4290 Reconcile project ITPM/OTPM reservations to actual usage on success: 

4291 ITPM to billable input tokens, OTPM to actual completion tokens. 

4292 Reuses ``_build_reservation_aware_tpm_ops``'s delta pattern -- ITPM/OTPM 

4293 are stored in the same ":tokens" cache bucket as combined TPM, just 

4294 under distinct scope keys, so the reservation-aware increment math is 

4295 identical; only the usage fields being reconciled against differ. 

4296 """ 

4297 if not isinstance(kwargs, dict): 4297 ↛ 4298line 4297 didn't jump to line 4298 because the condition on line 4297 was never true

4298 return () 

4299 stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) 

4300 if stash is None: 

4301 return () 

4302 

4303 itpm_reserved: Final = stash.itpm_reserved_tokens 

4304 otpm_reserved: Final = stash.otpm_reserved_tokens 

4305 if itpm_reserved <= 0 and otpm_reserved <= 0: 4305 ↛ 4308line 4305 didn't jump to line 4308 because the condition on line 4305 was always true

4306 return () 

4307 

4308 response_usage: Final = self._resolve_io_token_reconcile_usage(response_obj) 

4309 combined_usage: Final = self._resolve_io_token_reconcile_usage(kwargs.get("combined_usage_object")) 

4310 aggregate_total: Final = self._aggregate_only_total_tokens( 

4311 self._response_usage(response_obj) 

4312 ) or self._aggregate_only_total_tokens(self._response_usage(kwargs.get("combined_usage_object"))) 

4313 

4314 if not response_usage[2] and not combined_usage[2] and aggregate_total <= 0 and not stash.reservation_released: 

4315 return () 

4316 resolved_usage: Final = ( 

4317 response_usage 

4318 if response_usage[2] 

4319 else combined_usage 

4320 if combined_usage[2] 

4321 else (aggregate_total, aggregate_total, True) 

4322 if aggregate_total > 0 

4323 else (itpm_reserved, otpm_reserved, False) 

4324 ) 

4325 billable_input, completion_tokens, _ = resolved_usage 

4326 

4327 if stash.reservation_released or ( 

4328 not stash.itpm_reserved_window_identities and not stash.otpm_reserved_window_identities 

4329 ): 

4330 return self._build_reservation_aware_tpm_ops( 

4331 targets=tuple(stash.itpm_reserved_scopes), 

4332 reserved_scopes=frozenset() if stash.reservation_released else stash.itpm_reserved_scopes, 

4333 actual_tokens=billable_input, 

4334 reserved_tokens=0 if stash.reservation_released else itpm_reserved, 

4335 ) + self._build_reservation_aware_tpm_ops( 

4336 targets=tuple(stash.otpm_reserved_scopes), 

4337 reserved_scopes=frozenset() if stash.reservation_released else stash.otpm_reserved_scopes, 

4338 actual_tokens=completion_tokens, 

4339 reserved_tokens=0 if stash.reservation_released else otpm_reserved, 

4340 ) 

4341 

4342 itpm_ops: Final[Sequence[ReservationAwareIncrementOperation]] = ( 

4343 self._build_project_reservation_ops( 

4344 targets=tuple(stash.itpm_reserved_scopes), 

4345 reserved_scopes=frozenset() if stash.reservation_released else stash.itpm_reserved_scopes, 

4346 actual_tokens=billable_input, 

4347 reserved_tokens=itpm_reserved, 

4348 reservation_window_identities=stash.itpm_reserved_window_identities, 

4349 ) 

4350 if itpm_reserved > 0 

4351 else () 

4352 ) 

4353 otpm_ops: Final[Sequence[ReservationAwareIncrementOperation]] = ( 

4354 self._build_project_reservation_ops( 

4355 targets=tuple(stash.otpm_reserved_scopes), 

4356 reserved_scopes=frozenset() if stash.reservation_released else stash.otpm_reserved_scopes, 

4357 actual_tokens=completion_tokens, 

4358 reserved_tokens=otpm_reserved, 

4359 reservation_window_identities=stash.otpm_reserved_window_identities, 

4360 ) 

4361 if otpm_reserved > 0 

4362 else () 

4363 ) 

4364 return tuple((*itpm_ops, *otpm_ops)) 

4365 

4366 def _collect_tpm_scope_targets( 

4367 self, 

4368 standard_logging_metadata: dict[str, Any], 

4369 kwargs: object, 

4370 model_group: str | None, 

4371 tpm_limited_tags: Set[str] = frozenset(), 

4372 ) -> list[tuple[str, str]]: 

4373 """ 

4374 Enumerate every (scope_key, scope_value) pair that *might* carry a 

4375 TPM counter for this request — independent of whether each scope had 

4376 a configured TPM limit at pre-call. Reservation awareness happens at 

4377 the emitter; this helper just lists the candidate scopes so callers 

4378 can split reserved-vs-unreserved. 

4379 """ 

4380 user_api_key: Final = standard_logging_metadata.get("user_api_key_hash") 

4381 user_api_key_user_id: Final = standard_logging_metadata.get("user_api_key_user_id") 

4382 user_api_key_team_id: Final = standard_logging_metadata.get("user_api_key_team_id") 

4383 user_api_key_organization_id: Final = standard_logging_metadata.get("user_api_key_org_id") 

4384 user_api_key_project_id: Final = standard_logging_metadata.get("user_api_key_project_id") 

4385 user_api_key_end_user_id: Final = ( 

4386 kwargs.get("user") if isinstance(kwargs, dict) else None 

4387 ) or standard_logging_metadata.get("user_api_key_end_user_id") 

4388 agent_id: Final = standard_logging_metadata.get("agent_id") 

4389 session_id: Final = standard_logging_metadata.get("session_id") or standard_logging_metadata.get("trace_id") 

4390 

4391 targets: Final[list[tuple[str, str]]] = [] 

4392 if user_api_key: 

4393 targets.append(("api_key", user_api_key)) 

4394 if user_api_key_user_id: 

4395 targets.append(("user", user_api_key_user_id)) 

4396 if user_api_key_team_id: 4396 ↛ 4397line 4396 didn't jump to line 4397 because the condition on line 4396 was never true

4397 targets.append(("team", user_api_key_team_id)) 

4398 if user_api_key_team_id and user_api_key_user_id: 4398 ↛ 4399line 4398 didn't jump to line 4399 because the condition on line 4398 was never true

4399 targets.append(("team_member", f"{user_api_key_team_id}:{user_api_key_user_id}")) 

4400 if user_api_key_end_user_id: 4400 ↛ 4401line 4400 didn't jump to line 4401 because the condition on line 4400 was never true

4401 targets.append(("end_user", user_api_key_end_user_id)) 

4402 if user_api_key_organization_id: 4402 ↛ 4403line 4402 didn't jump to line 4403 because the condition on line 4402 was never true

4403 targets.append(("organization", user_api_key_organization_id)) 

4404 if model_group: 4404 ↛ 4405line 4404 didn't jump to line 4405 because the condition on line 4404 was never true

4405 if user_api_key: 

4406 targets.append(("model_per_key", f"{user_api_key}:{model_group}")) 

4407 if user_api_key_team_id: 

4408 targets.append(("model_per_team", f"{user_api_key_team_id}:{model_group}")) 

4409 if user_api_key_organization_id: 

4410 targets.append( 

4411 ( 

4412 "model_per_organization", 

4413 f"{user_api_key_organization_id}:{model_group}", 

4414 ) 

4415 ) 

4416 if user_api_key_project_id: 

4417 targets.append( 

4418 ( 

4419 "model_per_project", 

4420 f"{user_api_key_project_id}:{model_group}", 

4421 ) 

4422 ) 

4423 if agent_id: 4423 ↛ 4424line 4423 didn't jump to line 4424 because the condition on line 4423 was never true

4424 targets.append(("agent", agent_id)) 

4425 if session_id: 

4426 targets.append(("agent_session", f"{agent_id}:{session_id}")) 

4427 targets.extend(("tag", tag) for tag in sorted(tpm_limited_tags)) 

4428 return targets 

4429 

4430 def _build_reservation_aware_tpm_ops( 

4431 self, 

4432 targets: Sequence[tuple[str, str]], 

4433 reserved_scopes: Set[tuple[str, str]], 

4434 actual_tokens: int, 

4435 reserved_tokens: int, 

4436 ) -> list[RedisPipelineIncrementOperation]: 

4437 """ 

4438 Emit per-scope TPM increment ops with reservation awareness. 

4439 

4440 - Reserved scope (counter already at +reserved from pre-call): 

4441 reconcile to actual via ``actual - reserved``. 

4442 - Unreserved scope (counter never touched at pre-call): 

4443 charge the full ``actual``. 

4444 

4445 Same primitive serves success reconciliation, over-reservation 

4446 release, and failure refund — pass ``actual_tokens=0`` for the pure 

4447 refund case (reserved scopes get -reserved, unreserved get 0/skip). 

4448 """ 

4449 ops: Final[list[RedisPipelineIncrementOperation]] = [] 

4450 for scope_key, scope_value in targets: 

4451 if (scope_key, scope_value) in reserved_scopes: 4451 ↛ 4452line 4451 didn't jump to line 4452 because the condition on line 4451 was never true

4452 increment = actual_tokens - reserved_tokens 

4453 else: 

4454 increment = actual_tokens 

4455 if increment == 0: 4455 ↛ 4457line 4455 didn't jump to line 4457 because the condition on line 4455 was always true

4456 continue 

4457 ops.append( 

4458 RedisPipelineIncrementOperation( 

4459 key=self.create_rate_limit_keys(scope_key, scope_value, "tokens"), 

4460 increment_value=increment, 

4461 ttl=self.window_size, 

4462 ) 

4463 ) 

4464 return ops 

4465 

4466 def _build_project_reservation_op( 

4467 self, 

4468 scope: tuple[str, str], 

4469 reserved_scopes: Set[tuple[str, str]], 

4470 actual_tokens: int, 

4471 reserved_tokens: int, 

4472 reservation_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]], 

4473 ) -> ReservationAwareIncrementOperation | None: 

4474 scope_key, scope_value = scope 

4475 is_reserved_scope: Final = scope in reserved_scopes 

4476 increment: Final = actual_tokens - reserved_tokens if is_reserved_scope else actual_tokens 

4477 if increment == 0: 

4478 return None 

4479 counter_key: Final = self.create_rate_limit_keys(scope_key, scope_value, "tokens") 

4480 window_identity: Final = next( 

4481 ( 

4482 (window_start, backend) 

4483 for identity_counter_key, window_start, backend in reservation_window_identities 

4484 if identity_counter_key == counter_key 

4485 ), 

4486 None, 

4487 ) 

4488 if not is_reserved_scope or window_identity is None: 

4489 return ReservationAwareIncrementOperation( 

4490 key=counter_key, 

4491 increment_value=increment, 

4492 ttl=self.window_size, 

4493 ) 

4494 return ReservationAwareIncrementOperation( 

4495 key=counter_key, 

4496 increment_value=increment, 

4497 ttl=self.window_size, 

4498 window_key=f"{{{scope_key}:{scope_value}}}:window", 

4499 expected_window_start=window_identity[0], 

4500 reservation_backend=window_identity[1], 

4501 ) 

4502 

4503 def _build_project_reservation_ops( 

4504 self, 

4505 targets: Sequence[tuple[str, str]], 

4506 reserved_scopes: Set[tuple[str, str]], 

4507 actual_tokens: int, 

4508 reserved_tokens: int, 

4509 reservation_window_identities: frozenset[tuple[str, str, Literal["redis", "local"]]], 

4510 ) -> tuple[ReservationAwareIncrementOperation, ...]: 

4511 return tuple( 

4512 operation 

4513 for scope in targets 

4514 if ( 

4515 operation := self._build_project_reservation_op( 

4516 scope=scope, 

4517 reserved_scopes=reserved_scopes, 

4518 actual_tokens=actual_tokens, 

4519 reserved_tokens=reserved_tokens, 

4520 reservation_window_identities=reservation_window_identities, 

4521 ) 

4522 ) 

4523 is not None 

4524 ) 

4525 

4526 def _build_success_event_pipeline_operations( 

4527 self, 

4528 kwargs: dict[str, Any], 

4529 response_obj: object, 

4530 rate_limit_type: Literal["output", "input", "total"], 

4531 ) -> list[RedisPipelineIncrementOperation]: 

4532 """Build Redis pipeline increment ops for TPM / parallel-request counters.""" 

4533 from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs 

4534 from litellm.proxy.common_utils.callback_utils import ( 

4535 get_model_group_from_litellm_kwargs, 

4536 ) 

4537 

4538 # Get metadata from standard_logging_object - this correctly handles both 

4539 # 'metadata' and 'litellm_metadata' fields from litellm_params 

4540 standard_logging_object: Final = kwargs.get("standard_logging_object") or {} 

4541 request_metadata: Final = get_litellm_metadata_from_kwargs(kwargs) 

4542 origin: Final = request_metadata.get(INTERNAL_CALL_ORIGIN_METADATA_KEY) 

4543 if origin and origin != "autorouter_compaction": 4543 ↛ 4546line 4543 didn't jump to line 4546 because the condition on line 4543 was never true

4544 # Background evaluations keep their exemption; foreground compaction 

4545 # is necessary caller traffic and consumes the caller's token limits. 

4546 return [] 

4547 standard_logging_metadata: Final = standard_logging_object.get("metadata") or {} 

4548 

4549 model_group: Final = get_model_group_from_litellm_kwargs(kwargs) 

4550 

4551 # Get total tokens from response. Responses LiteLLM does not model 

4552 # (e.g. pass-through, whose usage is reported by the upstream rather 

4553 # than parsed out of the body) carry their usage in 

4554 # ``combined_usage_object`` instead, and would otherwise never charge 

4555 # the TPM window. 

4556 _usage: Usage | dict | None = None 

4557 if isinstance( 4557 ↛ 4566line 4557 didn't jump to line 4566 because the condition on line 4557 was never true

4558 response_obj, 

4559 ( 

4560 ModelResponse, 

4561 EmbeddingResponse, 

4562 TextCompletionResponse, 

4563 BaseLiteLLMOpenAIResponseObject, 

4564 ), 

4565 ): 

4566 _usage = getattr(response_obj, "usage", None) 

4567 else: 

4568 _combined_usage: Final = kwargs.get("combined_usage_object") 

4569 if isinstance(_combined_usage, Usage): 4569 ↛ 4570line 4569 didn't jump to line 4570 because the condition on line 4569 was never true

4570 _usage = _combined_usage 

4571 total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type) 

4572 if total_tokens == 0: 4572 ↛ 4575line 4572 didn't jump to line 4575 because the condition on line 4572 was always true

4573 total_tokens = self._aggregate_only_total_tokens(usage=_usage) 

4574 

4575 stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) 

4576 reserved_tokens: Final = stash.reserved_tokens if stash is not None else 0 

4577 reserved_model: Final = stash.reserved_model if stash is not None else None 

4578 reserved_scopes: Final[frozenset[tuple[str, str]]] = stash.reserved_scopes if stash is not None else frozenset() 

4579 # Reconciliation must target the same model-scoped counter that the 

4580 # pre-call reservation incremented. If a reservation was made, 

4581 # ``reserved_model`` (resolved at admission, so an alias map reload 

4582 # mid-flight cannot move the charge) is authoritative; otherwise fall 

4583 # back to the router's ``model_group`` (the no-reservation charge path). 

4584 reconcile_model: Final = reserved_model if reserved_model is not None else self._rate_limited_model(model_group) 

4585 

4586 pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [] 

4587 

4588 # ---------------------------------------------------------------- 

4589 # TPM reconciliation 

4590 # Per-scope behavior: 

4591 # reserved scope -> apply (actual - reserved) delta to settle 

4592 # the counter at +actual. 

4593 # unreserved scope -> charge the full actual usage (the 

4594 # reservation never incremented this scope). 

4595 # When no reservation was made, reserved_tokens=0 and reserved_scopes 

4596 # is empty, so every scope falls through the unreserved branch and 

4597 # gets the full actual charge — matching pre-PR behavior. 

4598 # ---------------------------------------------------------------- 

4599 targets: Final = self._collect_tpm_scope_targets( 

4600 standard_logging_metadata=standard_logging_metadata, 

4601 kwargs=kwargs, 

4602 tpm_limited_tags=stash.tpm_limited_tags if stash is not None else frozenset(), 

4603 model_group=reconcile_model.group if reconcile_model is not None else None, 

4604 ) 

4605 charged_targets: Final = ( 

4606 [target for target in targets if target[0] != "model_per_team"] 

4607 if self._key_owns_model_tpm_limit_from_request_metadata(request_metadata, reconcile_model) 

4608 else targets 

4609 ) 

4610 if reserved_tokens > 0 and total_tokens < reserved_tokens: 4610 ↛ 4611line 4610 didn't jump to line 4611 because the condition on line 4610 was never true

4611 verbose_proxy_logger.debug( 

4612 "Releasing unused TPM budget on success: reserved=%s, actual=%s, release=%s", 

4613 reserved_tokens, 

4614 total_tokens, 

4615 reserved_tokens - total_tokens, 

4616 ) 

4617 pipeline_operations.extend( 

4618 self._build_reservation_aware_tpm_ops( 

4619 targets=charged_targets, 

4620 reserved_scopes=reserved_scopes, 

4621 actual_tokens=total_tokens, 

4622 reserved_tokens=reserved_tokens, 

4623 ) 

4624 ) 

4625 

4626 return pipeline_operations 

4627 

4628 async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): 

4629 """ 

4630 Update TPM usage on successful API calls by incrementing counters using pipeline 

4631 """ 

4632 from litellm.litellm_core_utils.core_helpers import ( 

4633 _get_parent_otel_span_from_kwargs, 

4634 ) 

4635 

4636 rate_limit_type: Final = self.get_rate_limit_type() 

4637 

4638 litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs) 

4639 try: 

4640 verbose_proxy_logger.debug("INSIDE parallel request limiter ASYNC SUCCESS LOGGING") 

4641 

4642 stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) 

4643 await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span) 

4644 

4645 pipeline_operations: Final = self._build_success_event_pipeline_operations( 

4646 kwargs=kwargs, 

4647 response_obj=response_obj, 

4648 rate_limit_type=rate_limit_type, 

4649 ) 

4650 if pipeline_operations: 4650 ↛ 4651line 4650 didn't jump to line 4651 because the condition on line 4650 was never true

4651 await self.async_increment_tokens_with_ttl_preservation( 

4652 pipeline_operations=pipeline_operations, 

4653 parent_otel_span=litellm_parent_otel_span, 

4654 ) 

4655 io_token_operations: Final = self._build_io_token_reservation_ops( 

4656 kwargs=kwargs, 

4657 response_obj=response_obj, 

4658 ) 

4659 if io_token_operations: 4659 ↛ 4660line 4659 didn't jump to line 4660 because the condition on line 4659 was never true

4660 if isinstance(io_token_operations, list): 

4661 await self.async_increment_tokens_with_ttl_preservation( 

4662 pipeline_operations=io_token_operations, 

4663 parent_otel_span=litellm_parent_otel_span, 

4664 ) 

4665 else: 

4666 await self.async_increment_reservation_aware_tokens( 

4667 pipeline_operations=io_token_operations, 

4668 parent_otel_span=litellm_parent_otel_span, 

4669 ) 

4670 

4671 except Exception as e: 

4672 verbose_proxy_logger.exception("Error in rate limit success event: %s", e) 

4673 

4674 async def async_logging_hook( 

4675 self, 

4676 kwargs: dict, 

4677 result: object, 

4678 call_type: str, 

4679 ) -> tuple[dict, object]: 

4680 """ 

4681 Mirror the pre-call rate-limit snapshot into the SLP so streaming 

4682 success callbacks see the same ``x-ratelimit-*`` headers the 

4683 non-streaming path writes via ``async_post_call_success_hook``. 

4684 Runs in the earlier of the two callback loops inside 

4685 ``async_success_handler`` so downstream callbacks see the values 

4686 regardless of registration order. Idempotent for non-streaming. 

4687 """ 

4688 self._mirror_ratelimit_response_into_logging_payload( 

4689 kwargs=kwargs, 

4690 response_obj=result, 

4691 ) 

4692 return kwargs, result 

4693 

4694 def _mirror_ratelimit_response_into_logging_payload( 

4695 self, 

4696 kwargs: object, 

4697 response_obj: object, 

4698 ) -> None: 

4699 """ 

4700 Copy the stashed ``RateLimitResponse`` into the SLP's 

4701 ``hidden_params.additional_headers`` and the response object's 

4702 ``_hidden_params.additional_headers`` (when the latter is a dict). 

4703 """ 

4704 if not isinstance(kwargs, dict): 4704 ↛ 4705line 4704 didn't jump to line 4705 because the condition on line 4704 was never true

4705 return 

4706 

4707 stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) 

4708 rate_limit_response: Final = stash.rate_limit_response if stash is not None else None 

4709 statuses: Final = rate_limit_response["statuses"] if rate_limit_response is not None else [] 

4710 if not statuses: 4710 ↛ 4713line 4710 didn't jump to line 4713 because the condition on line 4710 was always true

4711 return 

4712 

4713 standard_logging_object: Final = kwargs.get("standard_logging_object") 

4714 if isinstance(standard_logging_object, dict): 

4715 hidden_params = standard_logging_object.get("hidden_params") 

4716 if not isinstance(hidden_params, dict): 

4717 hidden_params = {} 

4718 existing = hidden_params.get("additional_headers") 

4719 hidden_params["additional_headers"] = self._merge_ratelimit_statuses_into_additional_headers( 

4720 additional_headers=existing if isinstance(existing, dict) else {}, 

4721 statuses=statuses, 

4722 ) 

4723 standard_logging_object["hidden_params"] = hidden_params 

4724 

4725 response_hidden: Final = getattr(response_obj, "_hidden_params", None) 

4726 if isinstance(response_hidden, dict): 

4727 existing = response_hidden.get("additional_headers") 

4728 response_hidden["additional_headers"] = self._merge_ratelimit_statuses_into_additional_headers( 

4729 additional_headers=existing if isinstance(existing, dict) else {}, 

4730 statuses=statuses, 

4731 ) 

4732 

4733 def _recovered_partial_usage_tokens(self, source: Mapping[str, object]) -> tuple[int, int, int]: 

4734 usage: Final = source.get("combined_usage_object") 

4735 if not isinstance(usage, Usage) or (usage.completion_tokens or 0) <= 0: 4735 ↛ 4737line 4735 didn't jump to line 4737 because the condition on line 4735 was always true

4736 return 0, 0, 0 

4737 billable_input, completion_tokens, _ = self._resolve_io_token_reconcile_usage(usage) 

4738 return ( 

4739 self._get_total_tokens_from_usage(usage=usage, rate_limit_type=self.get_rate_limit_type()), 

4740 billable_input, 

4741 completion_tokens, 

4742 ) 

4743 

4744 async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): 

4745 """ 

4746 On failure: decrement max_parallel_requests and refund the upfront 

4747 TPM reservation only against the scopes the reservation actually 

4748 charged. Unreserved scopes were never incremented at pre-call, so 

4749 refunding them would drive their counter negative. A failed stream 

4750 whose partial usage was recovered settles the reservation at that 

4751 usage instead of refunding it. 

4752 """ 

4753 from litellm.litellm_core_utils.core_helpers import ( 

4754 _get_parent_otel_span_from_kwargs, 

4755 ) 

4756 

4757 try: 

4758 litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs) 

4759 

4760 pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = [] 

4761 

4762 stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(kwargs)) 

4763 await self._release_stashed_parallel_slot(stash, litellm_parent_otel_span) 

4764 

4765 # Skip the reservation refund if async_post_call_failure_hook 

4766 # already released it (proxy-level rejection that also bubbles up 

4767 # here as an LLM-error callback). max_parallel_requests is its 

4768 # own counter and is always decremented per call. 

4769 reserved_tokens, itpm_reserved, otpm_reserved = ( 

4770 (0, 0, 0) 

4771 if stash is None or stash.reservation_released 

4772 else (stash.reserved_tokens, stash.itpm_reserved_tokens, stash.otpm_reserved_tokens) 

4773 ) 

4774 tpm_actual, itpm_actual, otpm_actual = self._recovered_partial_usage_tokens(kwargs) 

4775 

4776 if stash is not None and reserved_tokens > 0: 4776 ↛ 4777line 4776 didn't jump to line 4777 because the condition on line 4776 was never true

4777 verbose_proxy_logger.debug( 

4778 "Settling reserved TPM tokens on failure: reserved=%s actual=%s", reserved_tokens, tpm_actual 

4779 ) 

4780 # Settle only against the scopes the reservation actually 

4781 # charged: unreserved scopes were never incremented, so a 

4782 # refund there would drive their counter negative. 

4783 pipeline_operations.extend( 

4784 self._build_reservation_aware_tpm_ops( 

4785 targets=list(stash.reserved_scopes), 

4786 reserved_scopes=stash.reserved_scopes, 

4787 actual_tokens=tpm_actual, 

4788 reserved_tokens=reserved_tokens, 

4789 ) 

4790 ) 

4791 

4792 # Settle project ITPM/OTPM reservations the same way: at the 

4793 # recovered partial usage, or a full refund when there is none. 

4794 itpm_operations: Final = ( 

4795 self._build_project_reservation_ops( 

4796 targets=tuple(stash.itpm_reserved_scopes), 

4797 reserved_scopes=stash.itpm_reserved_scopes, 

4798 actual_tokens=itpm_actual, 

4799 reserved_tokens=itpm_reserved, 

4800 reservation_window_identities=stash.itpm_reserved_window_identities, 

4801 ) 

4802 if stash is not None and itpm_reserved > 0 and stash.itpm_reserved_window_identities 

4803 else self._build_reservation_aware_tpm_ops( 

4804 targets=tuple(stash.itpm_reserved_scopes), 

4805 reserved_scopes=stash.itpm_reserved_scopes, 

4806 actual_tokens=itpm_actual, 

4807 reserved_tokens=itpm_reserved, 

4808 ) 

4809 if stash is not None and itpm_reserved > 0 

4810 else () 

4811 ) 

4812 

4813 otpm_operations: Final = ( 

4814 self._build_project_reservation_ops( 

4815 targets=tuple(stash.otpm_reserved_scopes), 

4816 reserved_scopes=stash.otpm_reserved_scopes, 

4817 actual_tokens=otpm_actual, 

4818 reserved_tokens=otpm_reserved, 

4819 reservation_window_identities=stash.otpm_reserved_window_identities, 

4820 ) 

4821 if stash is not None and otpm_reserved > 0 and stash.otpm_reserved_window_identities 

4822 else self._build_reservation_aware_tpm_ops( 

4823 targets=tuple(stash.otpm_reserved_scopes), 

4824 reserved_scopes=stash.otpm_reserved_scopes, 

4825 actual_tokens=otpm_actual, 

4826 reserved_tokens=otpm_reserved, 

4827 ) 

4828 if stash is not None and otpm_reserved > 0 

4829 else () 

4830 ) 

4831 

4832 if pipeline_operations: 4832 ↛ 4833line 4832 didn't jump to line 4833 because the condition on line 4832 was never true

4833 await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( 

4834 increment_list=pipeline_operations, 

4835 litellm_parent_otel_span=litellm_parent_otel_span, 

4836 ) 

4837 for project_operations in (itpm_operations, otpm_operations): 

4838 if isinstance(project_operations, list): 4838 ↛ 4839line 4838 didn't jump to line 4839 because the condition on line 4838 was never true

4839 await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( 

4840 increment_list=project_operations, 

4841 litellm_parent_otel_span=litellm_parent_otel_span, 

4842 ) 

4843 elif project_operations: 4843 ↛ 4844line 4843 didn't jump to line 4844 because the condition on line 4843 was never true

4844 await self.async_increment_reservation_aware_tokens( 

4845 pipeline_operations=project_operations, 

4846 parent_otel_span=litellm_parent_otel_span, 

4847 ) 

4848 if stash is not None and (reserved_tokens > 0 or itpm_reserved > 0 or otpm_reserved > 0): 4848 ↛ 4849line 4848 didn't jump to line 4849 because the condition on line 4848 was never true

4849 stash.reservation_released = True 

4850 except Exception as e: 

4851 verbose_proxy_logger.exception("Error in rate limit failure event: %s", e) 

4852 

4853 async def async_release_max_parallel_requests_on_disconnect( 

4854 self, 

4855 user_api_key_dict: UserAPIKeyAuth, 

4856 ) -> None: 

4857 """ 

4858 Release the api-key ``max_parallel_requests`` slot that 

4859 ``async_pre_call_hook`` acquired, for a request that ended without 

4860 either logging callback firing. 

4861 

4862 The slot is normally released by ``async_log_success_event`` (natural 

4863 stream completion) or ``async_log_failure_event`` (LLM error). When a 

4864 client cancels a stream mid-flight, the cancellation surfaces as 

4865 ``asyncio.CancelledError`` / ``GeneratorExit`` and neither callback 

4866 runs, so without this the slot leaks per cancelled stream until its 

4867 TTL prunes it. The stashed acquisition's presence (not the key 

4868 object's current max_parallel_requests configuration, which can 

4869 change mid-request) decides whether there is anything to release. 

4870 """ 

4871 await self._release_stashed_parallel_slot(get_request_stash(), None) 

4872 

4873 async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response): 

4874 """ 

4875 Release completed-request slots and update rate limit headers in the response. 

4876 """ 

4877 try: 

4878 slot_stash: Final = get_request_stash_for_call(_call_id_from_callback_kwargs(data)) 

4879 await self._release_stashed_parallel_slot(slot_stash, user_api_key_dict.parent_otel_span) 

4880 except Exception as e: 

4881 verbose_proxy_logger.exception("Error releasing parallel request slot in post-call hook: %s", e) 

4882 

4883 try: 

4884 header_stash: Final = get_request_stash() 

4885 litellm_proxy_rate_limit_response: Final = ( 

4886 header_stash.rate_limit_response if header_stash is not None else None 

4887 ) 

4888 

4889 if litellm_proxy_rate_limit_response is not None and response_has_hidden_params(response): 4889 ↛ 4890line 4889 didn't jump to line 4890 because the condition on line 4889 was never true

4890 additional_headers: Final = ensure_response_additional_headers(response) 

4891 additional_headers.update( 

4892 self._merge_ratelimit_statuses_into_additional_headers( 

4893 additional_headers={}, 

4894 statuses=litellm_proxy_rate_limit_response["statuses"], 

4895 ) 

4896 ) 

4897 

4898 except Exception as e: 

4899 verbose_proxy_logger.exception("Error in rate limit post-call hook: %s", e) 

4900 

4901 try: 

4902 await self._handle_batch_enqueued_post_call(user_api_key_dict=user_api_key_dict, response=response) 

4903 except Exception as e: # noqa: BLE001 # post-call batch accounting must never fail the response 

4904 verbose_proxy_logger.exception("Error in batch enqueued-token post-call hook: %s", e) 

4905 

4906 async def _handle_batch_enqueued_post_call(self, user_api_key_dict: UserAPIKeyAuth, response: object) -> None: 

4907 view: Final = batch_response_view(response) 

4908 if view is None: 4908 ↛ 4910line 4908 didn't jump to line 4910 because the condition on line 4908 was always true

4909 return 

4910 span: Final = user_api_key_dict.parent_otel_span 

4911 stash: Final = get_request_stash() 

4912 if stash is not None and stash.batch_enqueued_reservation is not None: 

4913 await self.batch_enqueued_token_store.save_reservation( 

4914 batch_id=canonical_provider_batch_id(view.id), 

4915 reservation=stash.batch_enqueued_reservation, 

4916 litellm_parent_otel_span=span, 

4917 ) 

4918 stash.batch_enqueued_reservation = None 

4919 if view.status.lower() in BATCH_ENQUEUED_REFUND_STATUSES: 

4920 popped: Final = await self.batch_enqueued_token_store.pop_reservation( 

4921 batch_id=canonical_provider_batch_id(view.id), 

4922 litellm_parent_otel_span=span, 

4923 ) 

4924 if popped is not None: 

4925 await self.batch_enqueued_token_store.refund(reservation=popped, litellm_parent_otel_span=span) 

4926 

4927 async def async_post_call_failure_hook( 

4928 self, 

4929 request_data: dict, 

4930 original_exception: Exception, 

4931 user_api_key_dict: UserAPIKeyAuth, 

4932 traceback_str: str | None = None, 

4933 ) -> None: 

4934 """ 

4935 Release the parallel-request slot and any TPM/ITPM/OTPM reservation 

4936 when the request is rejected after the pre-call hook acquired them 

4937 but before the LLM call ran (e.g. a downstream guardrail/auth hook 

4938 raised). Without this, those resources are stranded — 

4939 async_log_failure_event is a litellm completion-level callback and 

4940 never fires for proxy-side rejections, so a leaked slot would occupy 

4941 the gauge for the full PARALLEL_REQUEST_SLOT_TTL_SECONDS. 

4942 

4943 Idempotent: the slot release clears the stashed acquisition (and slot 

4944 removal is a no-op ZREM on a second run), and the TPM/ITPM/OTPM 

4945 refund is guarded by the stash's ``reservation_released`` flag — if 

4946 both this hook and async_log_failure_event end up running in the same 

4947 flow, only the first release/refund applies. A mid-stream failure 

4948 relayed here with recovered partial usage settles the reservation at 

4949 that usage instead of refunding it. 

4950 """ 

4951 try: 

4952 stash: Final = get_request_stash() 

4953 if stash is None: 

4954 return 

4955 await self._release_stashed_parallel_slot(stash, user_api_key_dict.parent_otel_span) 

4956 

4957 if stash.batch_enqueued_reservation is not None: 4957 ↛ 4958line 4957 didn't jump to line 4958 because the condition on line 4957 was never true

4958 await self.batch_enqueued_token_store.refund( 

4959 reservation=stash.batch_enqueued_reservation, 

4960 litellm_parent_otel_span=user_api_key_dict.parent_otel_span, 

4961 ) 

4962 stash.batch_enqueued_reservation = None 

4963 

4964 if stash.batch_tpd_refund_ops: 4964 ↛ 4965line 4964 didn't jump to line 4965 because the condition on line 4964 was never true

4965 await self.async_increment_reservation_aware_tokens( 

4966 pipeline_operations=stash.batch_tpd_refund_ops, 

4967 parent_otel_span=user_api_key_dict.parent_otel_span, 

4968 ) 

4969 stash.batch_tpd_refund_ops = () 

4970 

4971 if stash.reservation_released: 4971 ↛ 4972line 4971 didn't jump to line 4972 because the condition on line 4971 was never true

4972 return 

4973 reserved_tokens: Final = stash.reserved_tokens 

4974 itpm_reserved: Final = stash.itpm_reserved_tokens 

4975 otpm_reserved: Final = stash.otpm_reserved_tokens 

4976 if reserved_tokens <= 0 and itpm_reserved <= 0 and otpm_reserved <= 0: 4976 ↛ 4978line 4976 didn't jump to line 4978 because the condition on line 4976 was always true

4977 return 

4978 tpm_actual, itpm_actual, otpm_actual = self._recovered_partial_usage_tokens(request_data) 

4979 

4980 combined_ops: Final = ( 

4981 self._build_reservation_aware_tpm_ops( 

4982 targets=tuple(stash.reserved_scopes), 

4983 reserved_scopes=stash.reserved_scopes, 

4984 actual_tokens=tpm_actual, 

4985 reserved_tokens=reserved_tokens, 

4986 ) 

4987 if reserved_tokens > 0 

4988 else () 

4989 ) 

4990 itpm_ops: Final = ( 

4991 self._build_project_reservation_ops( 

4992 targets=tuple(stash.itpm_reserved_scopes), 

4993 reserved_scopes=stash.itpm_reserved_scopes, 

4994 actual_tokens=itpm_actual, 

4995 reserved_tokens=itpm_reserved, 

4996 reservation_window_identities=stash.itpm_reserved_window_identities, 

4997 ) 

4998 if itpm_reserved > 0 and stash.itpm_reserved_window_identities 

4999 else self._build_reservation_aware_tpm_ops( 

5000 targets=tuple(stash.itpm_reserved_scopes), 

5001 reserved_scopes=stash.itpm_reserved_scopes, 

5002 actual_tokens=itpm_actual, 

5003 reserved_tokens=itpm_reserved, 

5004 ) 

5005 if itpm_reserved > 0 

5006 else () 

5007 ) 

5008 otpm_ops: Final = ( 

5009 self._build_project_reservation_ops( 

5010 targets=tuple(stash.otpm_reserved_scopes), 

5011 reserved_scopes=stash.otpm_reserved_scopes, 

5012 actual_tokens=otpm_actual, 

5013 reserved_tokens=otpm_reserved, 

5014 reservation_window_identities=stash.otpm_reserved_window_identities, 

5015 ) 

5016 if otpm_reserved > 0 and stash.otpm_reserved_window_identities 

5017 else self._build_reservation_aware_tpm_ops( 

5018 targets=tuple(stash.otpm_reserved_scopes), 

5019 reserved_scopes=stash.otpm_reserved_scopes, 

5020 actual_tokens=otpm_actual, 

5021 reserved_tokens=otpm_reserved, 

5022 ) 

5023 if otpm_reserved > 0 

5024 else () 

5025 ) 

5026 if combined_ops or itpm_ops or otpm_ops: 

5027 verbose_proxy_logger.debug( 

5028 "Releasing reserved tokens on proxy-level rejection: tpm=%s, itpm=%s, otpm=%s", 

5029 reserved_tokens, 

5030 itpm_reserved, 

5031 otpm_reserved, 

5032 ) 

5033 if combined_ops: 

5034 await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( 

5035 increment_list=combined_ops, 

5036 litellm_parent_otel_span=user_api_key_dict.parent_otel_span, 

5037 ) 

5038 for project_ops in (itpm_ops, otpm_ops): 

5039 if isinstance(project_ops, list): 

5040 await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline( 

5041 increment_list=project_ops, 

5042 litellm_parent_otel_span=user_api_key_dict.parent_otel_span, 

5043 ) 

5044 elif project_ops: 

5045 await self.async_increment_reservation_aware_tokens( 

5046 pipeline_operations=project_ops, 

5047 parent_otel_span=user_api_key_dict.parent_otel_span, 

5048 ) 

5049 stash.reservation_released = True 

5050 except Exception as e: 

5051 verbose_proxy_logger.exception("Error releasing TPM reservation on post-call failure: %s", e) 

5052 return