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
« 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.
4This is currently in development and not yet ready for production.
5"""
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)
28from pydantic import TypeAdapter
29from typing_extensions import NotRequired, ReadOnly
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)
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
81 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
82 from litellm.types.agents import AgentResponse
83 from litellm.types.caching import RedisPipelineIncrementOperation
85 Span = _Span | Any
86 InternalUsageCache = _InternalUsageCache
87else:
88 Span = Any
89 InternalUsageCache = Any
92_REQUEST_RATE_LIMIT_DATA: Final = TypeAdapter(Mapping[str, object])
95@dataclass(frozen=True, slots=True)
96class RateLimitedModel:
97 requested: str
98 group: str
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)
107def _resolve_model_group_alias_via_proxy_router(model: str) -> str | None:
108 from litellm.proxy.proxy_server import llm_router
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)
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"
120BATCH_RATE_LIMITER_SCRIPT: Final = """
121local results = {}
122local now = tonumber(ARGV[1])
123local window_size = tonumber(ARGV[2])
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
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
156return results
157"""
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 = {}
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])
194 local window_start = redis.call('GET', window_key)
195 local window_expired = (not window_start) or
196 ((now - tonumber(window_start)) >= window_size)
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
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
215 descriptor_state[i] = { window_expired, current_counter, window_start }
216end
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])
228 local window_expired = descriptor_state[i][1]
229 local active_window_start
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
257return results
258"""
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)
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"""
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"""
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"""
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"""
350TOKEN_INCREMENT_SCRIPT: Final = """
351local results = {}
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])
359 -- Increment the value
360 local new_value = redis.call('INCRBYFLOAT', key, increment_value)
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
371 table.insert(results, new_value)
372end
374return results
375"""
377# Redis cluster slot count
378REDIS_CLUSTER_SLOTS: Final = 16384
379REDIS_NODE_HASHTAG_NAME: Final = "all_keys"
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
430CacheCounterValue: TypeAlias = int | float | str | bytes
432CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None]
434ReservationWindowIdentity: TypeAlias = tuple[str, str, Literal["redis", "local"]]
436ParallelGaugeCacheValue: TypeAlias = dict[str, object] | int | float | str | bytes
439class _AsyncLuaScript(Protocol):
440 """A Lua script registered against the async Redis client, called with KEYS and ARGV."""
442 def __call__(self, *, keys: Sequence[str], args: Sequence[object]) -> Awaitable[list[CacheCounterValue]]: ... 442 ↛ exitline 442 didn't return from function '__call__' because
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
452class RateLimitDescriptor(TypedDict):
453 key: str
454 value: str
455 rate_limit: RateLimitDescriptorRateLimitObject | None
458class ParallelRequestGauge(TypedDict):
459 counter_key: str
460 limit: int
461 descriptor_key: str
464class ParallelSlotAcquisition(TypedDict):
465 slot_id: str
466 counter_keys: list[str]
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]]
484class RateLimitResponse(TypedDict):
485 overall_code: str
486 statuses: list[RateLimitStatus]
487 reservation_windows: NotRequired[ReadOnly[frozenset[tuple[str, str, Literal["redis", "local"]]]]]
490class ReservationAwareIncrementOperation(RedisPipelineIncrementOperation):
491 window_key: NotRequired[str]
492 expected_window_start: NotRequired[str]
493 reservation_backend: NotRequired[Literal["redis", "local"]]
496class RateLimitResponseWithDescriptors(TypedDict):
497 descriptors: list[RateLimitDescriptor]
498 response: RateLimitResponse
501class _RateLimitDescriptorSink(Protocol):
502 def append(self, descriptor: RateLimitDescriptor, /) -> None: ... 502 ↛ exitline 502 didn't return from function 'append' because
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]
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
525class AtomicCounterState(TypedDict):
526 window_expired: bool
527 current: int
528 window_start: ReadOnly[str]
531DescriptorAtomicGroup: TypeAlias = tuple[list[str], list[int], list[AtomicCounterMeta]]
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: ...
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.
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.
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 """
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)
590@dataclass(frozen=True, slots=True)
591class TagRateLimit:
592 rpm_limit: int | None
593 tpm_limit: int | None
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
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)
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
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 )
630_request_stash: Final[ContextVar[RequestRateLimiterStash | None]] = ContextVar(
631 "litellm_v3_rate_limiter_request_stash", default=None
632)
635def get_request_stash() -> RequestRateLimiterStash | None:
636 return _request_stash.get()
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
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
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
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
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
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
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
733 self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
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"
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)
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()
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 )
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
778 def _get_current_time(self) -> datetime:
779 """Return the current time for rate limiting calculations."""
780 return self._time_provider()
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.
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))
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
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
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 )
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)
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.
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
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
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.
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.
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
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.
962 Supports chat (messages), completions (prompt), embeddings (input),
963 and the Responses API (also `input`, disambiguated from embeddings
964 via ``call_type``).
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.
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
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 )
992 return total_estimated
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.
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.
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.
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
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 )
1056 estimated_input_tokens: Final = max(1, total_chars // DEFAULT_CHARS_PER_TOKEN) if total_chars > 0 else 0
1058 explicit_max_tokens: Final = self._get_explicit_output_cap(data, call_type)
1059 is_embedding: Final = self._is_embedding_request(data, call_type)
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 )
1077 return estimated_input_tokens, max_tokens_estimate * self.get_output_candidate_count(data, call_type)
1079 def _is_redis_cluster(self) -> bool:
1080 """
1081 Check if the dual cache is using Redis cluster.
1083 Returns:
1084 bool: True if using Redis cluster, False otherwise.
1085 """
1086 from litellm.caching.redis_cluster_cache import RedisClusterCache
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 )
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)
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]] = []
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
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 )
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
1171 return results
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}"
1184 return counter_key
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"
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"]
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"
1216 if current_limit is None or rate_limit_type is None:
1217 continue
1219 if counter_value is not None and int(counter_value) > current_limit:
1220 overall_code = "OVER_LIMIT"
1221 item_code = "OVER_LIMIT"
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
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 )
1237 return RateLimitResponse(overall_code=overall_code, statuses=statuses)
1239 def keyslot_for_redis_cluster(self, key: str) -> int:
1240 """
1241 Compute the Redis Cluster slot for a given key.
1243 Simple implementation of `HASH_SLOT = CRC16(key) mod 16384`
1245 Read more about hash slots here: https://medium.com/@linz07m/how-hash-slots-power-data-distribution-in-redis-cluster-bc5b7e74ca7d
1247 Args:
1248 key (str): The Redis key.
1250 Returns:
1251 int: The slot number (0-16383).
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]
1262 # Compute CRC16 and mod 16384
1263 crc: Final = binascii.crc_hqx(key.encode("utf-8"), 0)
1264 return crc % REDIS_CLUSTER_SLOTS
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.
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]]] = {}
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}"
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
1288 return groups
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 )
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 )
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.
1323 Args:
1324 keys_to_fetch: List[str] - List of keys to fetch
1325 now_int: int - Current timestamp
1327 Returns:
1328 List of cache values
1329 """
1330 if self.batch_rate_limiter_script is None:
1331 return []
1333 key_groups: Final = self._group_keys_by_hash_tag(keys_to_fetch)
1334 all_cache_values: Final[list[CacheCounterValue | None]] = []
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)
1355 return all_cache_values
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.
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.
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 """
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
1396 keys_to_fetch, key_metadata, gauges = self._collect_windowed_keys_and_gauges(
1397 descriptors=descriptors,
1398 skip_tpm_check=skip_tpm_check,
1399 )
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 )
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
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 )
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 )
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 )
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
1469 if not gauges:
1470 return windowed_response
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 )
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
1507 window_key = f"{{{descriptor_key}:{descriptor_value}}}:window"
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 )
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
1530 if not rate_limit_set:
1531 continue
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
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 )
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))
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]
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)
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 )
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)
1652 async with self._check_and_increment_lock:
1653 return await self._acquire_parallel_slots_in_memory(gauges, slot_id, parent_otel_span)
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]
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.
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))
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)
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
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 )
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 )
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.
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.
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).
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.
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")
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))
1848 if not descriptor_groups:
1849 return RateLimitResponse(overall_code="OK", statuses=[])
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 )
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 )
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"
1888 keys: Final[list[str]] = []
1889 args: Final[list[int]] = []
1890 meta: Final[list[AtomicCounterMeta]] = []
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
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]
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 )
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()))
1986 return RateLimitResponse(
1987 overall_code="OK",
1988 statuses=statuses,
1989 reservation_windows=frozenset(reservation_windows),
1990 )
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 )
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.
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=[])
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 )
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 )
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.
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())
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 )
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 )
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.
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.
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=[])
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 )
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 )
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.
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.
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 ]
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
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
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 )
2379 assert itpm_response is not None
2380 return itpm_response, itpm_reserved, 0
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.
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)
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 )
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]] = []
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 )
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
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.
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 )
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
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 )
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.
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
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
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 )
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.
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
2567 if not mcp_server_name or not user_api_key_dict.api_key:
2568 return
2570 mcp_rpm_limit: Final = get_key_mcp_rpm_limit(user_api_key_dict)
2571 if not mcp_rpm_limit:
2572 return
2574 server_rpm_limit: Final = mcp_rpm_limit.get(mcp_server_name)
2575 if server_rpm_limit is None:
2576 return
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 )
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
2602 if not mcp_server_name:
2603 return
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))
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 )
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.
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
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
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.
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
2670 Returns:
2671 The limit value if it should be enforced, None otherwise
2672 """
2673 if limit_value is None:
2674 return None
2676 if self._should_enforce_rate_limit(
2677 limit_type=limit_type,
2678 model_has_failures=model_has_failures,
2679 ):
2680 return limit_value
2682 return None
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.
2692 Args:
2693 rpm_limit_type: RPM rate limit type
2694 tpm_limit_type: TPM rate limit type
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"
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
2705 return global_agent_registry.get_agent_by_id(agent_id=agent_id)
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")
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
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.
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]] = []
2746 agent: Final = self._get_agent_from_registry(agent_id)
2747 if agent is None:
2748 return descriptors
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 )
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 )
2782 return descriptors
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 )
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.
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
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 )
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 )
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 )
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 )
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 )
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 )
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 )
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 )
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 )
2941 # Agent-level and session-level rate limits
2942 resolved_agent_id: Final = self._get_resolved_agent_id(user_api_key_dict, data)
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 )
2952 return descriptors
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.
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 )
2970 if llm_router is None:
2971 return False
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
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
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 )
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
3000 verbose_proxy_logger.debug(
3001 "[Dynamic Rate Limit] No failures detected for model %s - allowing dynamic exceeding", model
3002 )
3003 return False
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
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
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
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
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")
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 )
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 )
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.
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
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 )
3131 if model_itpm_limit is None and model_otpm_limit is None:
3132 return
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 )
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"
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 )
3187 remaining_display = max(0, status["limit_remaining"])
3188 rate_limit_type = status["rate_limit_type"]
3189 current_limit = status["current_limit"]
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 )
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 )
3211 @staticmethod
3212 def _estimate_audio_block_tokens(block: object) -> int:
3213 """
3214 Token estimate for one ``input_audio`` content block.
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.
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
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 )
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
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 )
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 )
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
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 )
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.
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.
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
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
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
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
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
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
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 )
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 )
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 )
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 )
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 )
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
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 )
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 ]
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)
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
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")
3657 stash: Final = claim_request_stash_for_data(data)
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 )
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 )
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
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 )
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 )
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)
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
3764 if has_tpm_limits and self.tpm_reservation_enabled:
3765 min_configured_tpm_limit: Final = min(configured_tpm_limits)
3767 configured_output_tokens: Final = get_estimated_output_tokens(
3768 user_api_key_dict=user_api_key_dict,
3769 model_name=requested_model,
3770 )
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 )
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 )
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 )
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 )
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
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"])
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 )
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 )
3893 return pipeline_operations
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.
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
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
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
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
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
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)
3941 return total_tokens
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.
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
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
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
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)
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]
4009 keys = []
4010 args = []
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
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])
4025 await self.token_increment_script(
4026 keys=keys,
4027 args=args,
4028 )
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
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
4051 try:
4052 await self._execute_token_increment_script(pipeline_operations)
4054 verbose_proxy_logger.debug(
4055 "Successfully executed TTL-preserving increment for %s keys", len(pipeline_operations)
4056 )
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 )
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 )
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 )
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 )
4179 def get_rate_limit_type(self) -> Literal["output", "input", "total"]:
4180 from litellm.proxy.proxy_server import general_settings
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
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
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
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
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
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
4242 usage: Final = self._response_usage(response_obj)
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
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
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
4282 return 0, 0, False
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 ()
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 ()
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")))
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
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 )
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))
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")
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
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.
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``.
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
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 )
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 )
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 )
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 {}
4549 model_group: Final = get_model_group_from_litellm_kwargs(kwargs)
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)
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)
4586 pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = []
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 )
4626 return pipeline_operations
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 )
4636 rate_limit_type: Final = self.get_rate_limit_type()
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")
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)
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 )
4671 except Exception as e:
4672 verbose_proxy_logger.exception("Error in rate limit success event: %s", e)
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
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
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
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
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 )
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 )
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 )
4757 try:
4758 litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs)
4760 pipeline_operations: Final[list[RedisPipelineIncrementOperation]] = []
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)
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)
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 )
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 )
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 )
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)
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.
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)
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)
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 )
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 )
4898 except Exception as e:
4899 verbose_proxy_logger.exception("Error in rate limit post-call hook: %s", e)
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)
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)
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.
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)
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
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 = ()
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)
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