Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/hooks/parallel_request_limiter.py: 10%
316 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1import asyncio
2import sys
3from datetime import datetime, timedelta
4from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn
6from pydantic import BaseModel
7from typing_extensions import TypedDict
9import litellm
10from litellm import DualCache, EmbeddingResponse, ModelResponse, TextCompletionResponse
11from litellm._logging import verbose_proxy_logger
12from litellm.exceptions import RateLimitType
13from litellm.integrations.custom_logger import CustomLogger
14from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
15from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth
16from litellm.proxy.auth.auth_utils import (
17 get_key_model_rpm_limit,
18 get_key_model_tpm_limit,
19)
20from litellm.proxy.auth.budget_throttle import throttled_limit
21from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
22from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
23from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
24from litellm.types.utils import Usage
26if TYPE_CHECKING: 26 ↛ 27line 26 didn't jump to line 27 because the condition on line 26 was never true
27 from opentelemetry.trace import Span as _Span
29 from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
31 Span = _Span
32 InternalUsageCache = _InternalUsageCache
33else:
34 Span = Any
35 InternalUsageCache = Any
38def _response_total_tokens(response_obj: object) -> int:
39 if not isinstance(response_obj, (ModelResponse, EmbeddingResponse, TextCompletionResponse)):
40 return 0
41 response_usage: Final = getattr(response_obj, "usage", None)
42 return response_usage.total_tokens if isinstance(response_usage, Usage) else 0
45class CacheObject(TypedDict):
46 current_global_requests: dict | None
47 request_count_api_key: dict | None
48 request_count_api_key_model: dict | None
49 request_count_user_id: dict | None
50 request_count_team_id: dict | None
51 request_count_end_user_id: dict | None
54class _PROXY_MaxParallelRequestsHandler(CustomLogger):
55 # Class variables or attributes
56 def __init__(self, internal_usage_cache: InternalUsageCache):
57 self.internal_usage_cache = internal_usage_cache
59 def print_verbose(self, print_statement):
60 try:
61 verbose_proxy_logger.debug(print_statement)
62 if litellm.set_verbose:
63 print(print_statement) # noqa: T201
64 except Exception:
65 pass
67 async def check_key_in_limits(
68 self,
69 user_api_key_dict: UserAPIKeyAuth,
70 cache: DualCache,
71 data: dict,
72 call_type: str,
73 max_parallel_requests: int,
74 tpm_limit: int,
75 rpm_limit: int,
76 current: dict | None,
77 request_count_api_key: str,
78 rate_limit_type: Literal["key", "model_per_key", "user", "customer", "team"],
79 values_to_update_in_cache: list[tuple[str, object]],
80 ) -> dict:
81 verbose_proxy_logger.info("Current Usage of %s in this minute: %s", rate_limit_type, current)
82 if current is None:
83 if max_parallel_requests == 0 or tpm_limit == 0 or rpm_limit == 0:
84 # base case — at least one dimension is set to 0 (effectively
85 # disabled). Pick the most specific dimension as the
86 # rate_limit_type so dashboards can attribute the failure to
87 # the right cap. Order matters: max_parallel_requests is
88 # listed first because it's the rarest 0 in practice and the
89 # most actionable signal.
90 if max_parallel_requests == 0:
91 triggered_type = RateLimitType.CONCURRENT_REQUESTS
92 elif tpm_limit == 0:
93 triggered_type = RateLimitType.TOKENS
94 else:
95 triggered_type = RateLimitType.REQUESTS
96 self.raise_rate_limit_error(
97 additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}",
98 rate_limit_type=triggered_type,
99 requested_model=data.get("model") if data else None,
100 )
101 new_val = {
102 "current_requests": 1,
103 "current_tpm": 0,
104 "current_rpm": 1,
105 }
106 values_to_update_in_cache.append((request_count_api_key, new_val))
107 elif (
108 int(current["current_requests"]) < max_parallel_requests
109 and current["current_tpm"] < tpm_limit
110 and current["current_rpm"] < rpm_limit
111 ):
112 # Increase count for this token
113 new_val = {
114 "current_requests": current["current_requests"] + 1,
115 "current_tpm": current["current_tpm"],
116 "current_rpm": current["current_rpm"] + 1,
117 }
118 values_to_update_in_cache.append((request_count_api_key, new_val))
120 else:
121 # Detect which dimension actually tripped the limit so we can
122 # surface the right rate_limit_type. Order matches the boolean
123 # condition above (concurrent → tpm → rpm) — first match wins.
124 if int(current["current_requests"]) >= max_parallel_requests:
125 triggered_type = RateLimitType.CONCURRENT_REQUESTS
126 elif current["current_tpm"] >= tpm_limit:
127 triggered_type = RateLimitType.TOKENS
128 else:
129 triggered_type = RateLimitType.REQUESTS
130 requested_model: Final = data.get("model") if data else None
131 resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(requested_model)
132 raise ProxyRateLimitError(
133 detail=f"LiteLLM Rate Limit Handler for rate limit type = {rate_limit_type}. {CommonProxyErrors.max_parallel_request_limit_reached.value}. current rpm: {current['current_rpm']}, rpm limit: {rpm_limit}, current tpm: {current['current_tpm']}, tpm limit: {tpm_limit}, current max_parallel_requests: {current['current_requests']}, max_parallel_requests: {max_parallel_requests}",
134 headers={"retry-after": str(self.time_to_next_minute())},
135 rate_limit_type=triggered_type,
136 model=resolved_model,
137 llm_provider=llm_provider,
138 )
140 await self.internal_usage_cache.async_batch_set_cache(
141 cache_list=values_to_update_in_cache,
142 ttl=60,
143 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
144 local_only=True,
145 )
146 return new_val
148 def time_to_next_minute(self) -> float:
149 # Get the current time
150 now: Final = datetime.now()
152 # Calculate the next minute
153 next_minute: Final = (now + timedelta(minutes=1)).replace(second=0, microsecond=0)
155 # Calculate the difference in seconds
156 seconds_to_next_minute: Final = (next_minute - now).total_seconds()
158 return seconds_to_next_minute
160 def raise_rate_limit_error(
161 self,
162 additional_details: str | None = None,
163 rate_limit_type: RateLimitType | None = None,
164 requested_model: str | None = None,
165 ) -> NoReturn:
166 """
167 Raise a 429 with a retry-after header for litellm-proxy parallel-request limits.
169 Always raises :class:`ProxyRateLimitError` — never returns. Annotated
170 ``NoReturn`` so type-checkers know callers after this invocation are
171 unreachable. The raised exception is both a
172 :class:`litellm.RateLimitError` (so callers can catch by category) and a
173 :class:`fastapi.HTTPException` (so the FastAPI dispatcher serializes it
174 correctly with status 429 and the supplied headers).
176 ``rate_limit_type`` defaults to ``CONCURRENT_REQUESTS`` because every
177 existing internal caller of this helper hits the parallel-request cap
178 (the global-limit branch in ``async_pre_call_hook`` and the
179 all-zeros base case in ``check_key_in_limits``). Callers that know
180 the dimension exactly should pass it explicitly.
182 ``requested_model`` is resolved via :func:`get_llm_provider` so the
183 raised exception carries ``llm_provider`` (and a stripped ``model``)
184 for downstream loggers (Prometheus failure metric, observability
185 callbacks). Falls back to ``llm_provider="litellm_proxy"`` when the
186 model is missing or unparseable — see
187 :func:`resolve_llm_provider_for_rate_limit`.
188 """
189 # additional_details is optional; build the detail with a None-guard
190 # so callers that pass nothing don't get the literal string "None"
191 # interpolated into the error message.
192 error_message = "Max parallel request limit reached"
193 if additional_details is not None:
194 error_message = error_message + " " + additional_details
195 resolved_model, llm_provider = resolve_llm_provider_for_rate_limit(requested_model)
196 raise ProxyRateLimitError(
197 detail=error_message,
198 headers={"retry-after": str(self.time_to_next_minute())},
199 rate_limit_type=rate_limit_type or RateLimitType.CONCURRENT_REQUESTS,
200 model=resolved_model,
201 llm_provider=llm_provider,
202 )
204 async def get_all_cache_objects(
205 self,
206 current_global_requests: str | None,
207 request_count_api_key: str | None,
208 request_count_api_key_model: str | None,
209 request_count_user_id: str | None,
210 request_count_team_id: str | None,
211 request_count_end_user_id: str | None,
212 parent_otel_span: Span | None = None,
213 ) -> CacheObject:
214 keys: Final = [
215 current_global_requests,
216 request_count_api_key,
217 request_count_api_key_model,
218 request_count_user_id,
219 request_count_team_id,
220 request_count_end_user_id,
221 ]
222 results: Final = await self.internal_usage_cache.async_batch_get_cache(
223 keys=keys,
224 parent_otel_span=parent_otel_span,
225 )
227 if results is None:
228 return CacheObject(
229 current_global_requests=None,
230 request_count_api_key=None,
231 request_count_api_key_model=None,
232 request_count_user_id=None,
233 request_count_team_id=None,
234 request_count_end_user_id=None,
235 )
237 return CacheObject(
238 current_global_requests=results[0],
239 request_count_api_key=results[1],
240 request_count_api_key_model=results[2],
241 request_count_user_id=results[3],
242 request_count_team_id=results[4],
243 request_count_end_user_id=results[5],
244 )
246 async def async_pre_call_hook(
247 self,
248 user_api_key_dict: UserAPIKeyAuth,
249 cache: DualCache,
250 data: dict,
251 call_type: str,
252 ):
253 self.print_verbose("Inside Max Parallel Request Pre-Call Hook")
254 api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
255 max_parallel_requests = user_api_key_dict.max_parallel_requests
256 if max_parallel_requests is None:
257 max_parallel_requests = sys.maxsize
258 if data is None:
259 data = {}
260 global_max_parallel_requests: Final = data.get("metadata", {}).get("global_max_parallel_requests", None)
261 throttle_pct: Final = getattr(user_api_key_dict, "budget_throttle_pct", None)
262 tpm_limit = throttled_limit(getattr(user_api_key_dict, "tpm_limit", sys.maxsize), throttle_pct)
263 if tpm_limit is None:
264 tpm_limit = sys.maxsize
265 rpm_limit = throttled_limit(getattr(user_api_key_dict, "rpm_limit", sys.maxsize), throttle_pct)
266 if rpm_limit is None:
267 rpm_limit = sys.maxsize
269 values_to_update_in_cache: list[
270 tuple[str, object]
271 ] = [] # values that need to get updated in cache, will run a batch_set_cache after this function
273 # ------------
274 # Setup values
275 # ------------
276 new_val: dict | None = None
278 if global_max_parallel_requests is not None:
279 # get value from cache
280 _key: Final = "global_max_parallel_requests"
281 current_global_requests = await self.internal_usage_cache.async_get_cache(
282 key=_key,
283 local_only=True,
284 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
285 )
286 # check if below limit
287 if current_global_requests is None:
288 current_global_requests = 1
289 # if above -> raise error
290 if current_global_requests >= global_max_parallel_requests:
291 self.raise_rate_limit_error(
292 additional_details=f"Hit Global Limit: Limit={global_max_parallel_requests}, current: {current_global_requests}",
293 requested_model=data.get("model") if data else None,
294 )
295 # if below -> increment
296 else:
297 await self.internal_usage_cache.async_increment_cache(
298 key=_key,
299 value=1,
300 local_only=True,
301 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
302 )
303 _model = data.get("model", None)
305 current_date: Final = datetime.now().strftime("%Y-%m-%d")
306 current_hour: Final = datetime.now().strftime("%H")
307 current_minute: Final = datetime.now().strftime("%M")
308 precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
310 cache_objects: Final[CacheObject] = await self.get_all_cache_objects(
311 current_global_requests=(
312 "global_max_parallel_requests" if global_max_parallel_requests is not None else None
313 ),
314 request_count_api_key=(f"{api_key}::{precise_minute}::request_count" if api_key is not None else None),
315 request_count_api_key_model=(
316 f"{api_key}::{_model}::{precise_minute}::request_count"
317 if api_key is not None and _model is not None
318 else None
319 ),
320 request_count_user_id=(
321 f"{user_api_key_dict.user_id}::{precise_minute}::request_count"
322 if user_api_key_dict.user_id is not None
323 else None
324 ),
325 request_count_team_id=(
326 f"{user_api_key_dict.team_id}::{precise_minute}::request_count"
327 if user_api_key_dict.team_id is not None
328 else None
329 ),
330 request_count_end_user_id=(
331 f"{user_api_key_dict.end_user_id}::{precise_minute}::request_count"
332 if user_api_key_dict.end_user_id is not None
333 else None
334 ),
335 parent_otel_span=user_api_key_dict.parent_otel_span,
336 )
337 if api_key is not None:
338 request_count_api_key = f"{api_key}::{precise_minute}::request_count"
339 # CHECK IF REQUEST ALLOWED for key
340 await self.check_key_in_limits(
341 user_api_key_dict=user_api_key_dict,
342 cache=cache,
343 data=data,
344 call_type=call_type,
345 max_parallel_requests=max_parallel_requests,
346 current=cache_objects["request_count_api_key"],
347 request_count_api_key=request_count_api_key,
348 tpm_limit=tpm_limit,
349 rpm_limit=rpm_limit,
350 rate_limit_type="key",
351 values_to_update_in_cache=values_to_update_in_cache,
352 )
354 # Check if request under RPM/TPM per model for a given API Key
355 _model = data.get("model", None)
356 _tpm_limit_for_key_model: Final = get_key_model_tpm_limit(user_api_key_dict, model_name=_model)
357 _rpm_limit_for_key_model: Final = get_key_model_rpm_limit(user_api_key_dict, model_name=_model)
358 if _tpm_limit_for_key_model is not None or _rpm_limit_for_key_model is not None:
359 request_count_api_key = f"{api_key}::{_model}::{precise_minute}::request_count"
360 tpm_limit_for_model = None
361 rpm_limit_for_model = None
363 if _model is not None:
364 if _tpm_limit_for_key_model:
365 tpm_limit_for_model = _tpm_limit_for_key_model.get(_model)
367 if _rpm_limit_for_key_model:
368 rpm_limit_for_model = _rpm_limit_for_key_model.get(_model)
370 new_val = await self.check_key_in_limits(
371 user_api_key_dict=user_api_key_dict,
372 cache=cache,
373 data=data,
374 call_type=call_type,
375 max_parallel_requests=sys.maxsize, # TODO: Support max parallel requests for a model
376 current=cache_objects["request_count_api_key_model"],
377 request_count_api_key=request_count_api_key,
378 tpm_limit=tpm_limit_for_model or sys.maxsize,
379 rpm_limit=rpm_limit_for_model or sys.maxsize,
380 rate_limit_type="model_per_key",
381 values_to_update_in_cache=values_to_update_in_cache,
382 )
383 _remaining_tokens = None
384 _remaining_requests = None
385 # Add remaining tokens, requests to metadata
386 if new_val:
387 if tpm_limit_for_model is not None:
388 _remaining_tokens = tpm_limit_for_model - new_val["current_tpm"]
389 if rpm_limit_for_model is not None:
390 _remaining_requests = rpm_limit_for_model - new_val["current_rpm"]
392 _remaining_limits_data: Final = {
393 f"litellm-key-remaining-tokens-{_model}": _remaining_tokens,
394 f"litellm-key-remaining-requests-{_model}": _remaining_requests,
395 }
397 if "metadata" not in data:
398 data["metadata"] = {}
399 data["metadata"].update(_remaining_limits_data)
401 # check if REQUEST ALLOWED for user_id
402 user_id: Final = user_api_key_dict.user_id
403 if user_id is not None:
404 user_tpm_limit = user_api_key_dict.user_tpm_limit
405 user_rpm_limit = user_api_key_dict.user_rpm_limit
406 if user_tpm_limit is None:
407 user_tpm_limit = sys.maxsize
408 if user_rpm_limit is None:
409 user_rpm_limit = sys.maxsize
411 request_count_api_key = f"{user_id}::{precise_minute}::request_count"
412 # print(f"Checking if {request_count_api_key} is allowed to make request for minute {precise_minute}")
413 await self.check_key_in_limits(
414 user_api_key_dict=user_api_key_dict,
415 cache=cache,
416 data=data,
417 call_type=call_type,
418 max_parallel_requests=sys.maxsize, # TODO: Support max parallel requests for a user
419 current=cache_objects["request_count_user_id"],
420 request_count_api_key=request_count_api_key,
421 tpm_limit=user_tpm_limit,
422 rpm_limit=user_rpm_limit,
423 rate_limit_type="user",
424 values_to_update_in_cache=values_to_update_in_cache,
425 )
427 # TEAM RATE LIMITS
428 ## get team tpm/rpm limits
429 team_id: Final = user_api_key_dict.team_id
430 if team_id is not None:
431 team_tpm_limit = user_api_key_dict.team_tpm_limit
432 team_rpm_limit = user_api_key_dict.team_rpm_limit
434 if team_tpm_limit is None:
435 team_tpm_limit = sys.maxsize
436 if team_rpm_limit is None:
437 team_rpm_limit = sys.maxsize
439 request_count_api_key = f"{team_id}::{precise_minute}::request_count"
440 # print(f"Checking if {request_count_api_key} is allowed to make request for minute {precise_minute}")
441 await self.check_key_in_limits(
442 user_api_key_dict=user_api_key_dict,
443 cache=cache,
444 data=data,
445 call_type=call_type,
446 max_parallel_requests=sys.maxsize, # TODO: Support max parallel requests for a team
447 current=cache_objects["request_count_team_id"],
448 request_count_api_key=request_count_api_key,
449 tpm_limit=team_tpm_limit,
450 rpm_limit=team_rpm_limit,
451 rate_limit_type="team",
452 values_to_update_in_cache=values_to_update_in_cache,
453 )
455 # End-User Rate Limits
456 # Only enforce if user passed `user` to /chat, /completions, /embeddings
457 if user_api_key_dict.end_user_id:
458 end_user_tpm_limit = getattr(user_api_key_dict, "end_user_tpm_limit", sys.maxsize)
459 end_user_rpm_limit = getattr(user_api_key_dict, "end_user_rpm_limit", sys.maxsize)
461 if end_user_tpm_limit is None:
462 end_user_tpm_limit = sys.maxsize
463 if end_user_rpm_limit is None:
464 end_user_rpm_limit = sys.maxsize
466 # now do the same tpm/rpm checks
467 request_count_api_key = f"{user_api_key_dict.end_user_id}::{precise_minute}::request_count"
469 # print(f"Checking if {request_count_api_key} is allowed to make request for minute {precise_minute}")
470 await self.check_key_in_limits(
471 user_api_key_dict=user_api_key_dict,
472 cache=cache,
473 data=data,
474 call_type=call_type,
475 max_parallel_requests=sys.maxsize, # TODO: Support max parallel requests for an End-User
476 request_count_api_key=request_count_api_key,
477 current=cache_objects["request_count_end_user_id"],
478 tpm_limit=end_user_tpm_limit,
479 rpm_limit=end_user_rpm_limit,
480 rate_limit_type="customer",
481 values_to_update_in_cache=values_to_update_in_cache,
482 )
484 asyncio.create_task(
485 self.internal_usage_cache.async_batch_set_cache(
486 cache_list=values_to_update_in_cache,
487 ttl=60,
488 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
489 ) # don't block execution for cache updates
490 )
492 async def async_log_success_event(self, kwargs, response_obj: object, start_time, end_time):
493 from litellm.proxy.common_utils.callback_utils import (
494 get_model_group_from_litellm_kwargs,
495 )
497 litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
498 try:
499 self.print_verbose("INSIDE parallel request limiter ASYNC SUCCESS LOGGING")
501 global_max_parallel_requests: Final = kwargs["litellm_params"]["metadata"].get(
502 "global_max_parallel_requests", None
503 )
504 user_api_key: Final = kwargs["litellm_params"]["metadata"]["user_api_key"]
505 user_api_key_user_id: Final = kwargs["litellm_params"]["metadata"].get("user_api_key_user_id", None)
506 user_api_key_team_id: Final = kwargs["litellm_params"]["metadata"].get("user_api_key_team_id", None)
507 user_api_key_model_max_budget: Final = kwargs["litellm_params"]["metadata"].get(
508 "user_api_key_model_max_budget", None
509 )
510 user_api_key_end_user_id: Final = kwargs.get("user")
512 user_api_key_metadata: Final = kwargs["litellm_params"]["metadata"].get("user_api_key_metadata", {}) or {}
513 user_api_key_team_metadata = kwargs["litellm_params"]["metadata"].get("user_api_key_team_metadata", None)
514 user_api_key_dict: Final = UserAPIKeyAuth(
515 api_key=user_api_key,
516 metadata=user_api_key_metadata,
517 model_max_budget=user_api_key_model_max_budget,
518 team_metadata=user_api_key_team_metadata,
519 )
521 # ------------
522 # Setup values
523 # ------------
525 if global_max_parallel_requests is not None:
526 # get value from cache
527 _key: Final = "global_max_parallel_requests"
528 # decrement
529 await self.internal_usage_cache.async_increment_cache(
530 key=_key,
531 value=-1,
532 local_only=True,
533 litellm_parent_otel_span=litellm_parent_otel_span,
534 )
536 current_date: Final = datetime.now().strftime("%Y-%m-%d")
537 current_hour: Final = datetime.now().strftime("%H")
538 current_minute: Final = datetime.now().strftime("%M")
539 precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
541 total_tokens: int = _response_total_tokens(response_obj)
543 # ------------
544 # Update usage - API Key
545 # ------------
547 values_to_update_in_cache: Final[list[tuple[str, object]]] = []
549 if user_api_key is not None:
550 request_count_api_key = f"{user_api_key}::{precise_minute}::request_count"
552 current: dict[str, int] = await self.internal_usage_cache.async_get_cache(
553 key=request_count_api_key,
554 litellm_parent_otel_span=litellm_parent_otel_span,
555 ) or {
556 "current_requests": 1,
557 "current_tpm": 0,
558 "current_rpm": 0,
559 }
561 new_val = {
562 "current_requests": max(current["current_requests"] - 1, 0),
563 "current_tpm": current["current_tpm"] + total_tokens,
564 "current_rpm": current["current_rpm"],
565 }
567 self.print_verbose(f"updated_value in success call: {new_val}, precise_minute: {precise_minute}")
568 values_to_update_in_cache.append((request_count_api_key, new_val))
570 # ------------
571 # Update usage - model group + API Key
572 # ------------
573 model_group: Final = get_model_group_from_litellm_kwargs(kwargs)
574 _success_tpm_limit: Final = (
575 get_key_model_tpm_limit(user_api_key_dict, model_name=model_group) if model_group is not None else None
576 )
577 _success_rpm_limit: Final = (
578 get_key_model_rpm_limit(user_api_key_dict, model_name=model_group) if model_group is not None else None
579 )
580 if (
581 user_api_key is not None
582 and model_group is not None
583 and (
584 "model_rpm_limit" in user_api_key_metadata
585 or "model_tpm_limit" in user_api_key_metadata
586 or user_api_key_model_max_budget is not None
587 or _success_tpm_limit is not None
588 or _success_rpm_limit is not None
589 )
590 ):
591 request_count_api_key = f"{user_api_key}::{model_group}::{precise_minute}::request_count"
593 current = await self.internal_usage_cache.async_get_cache(
594 key=request_count_api_key,
595 litellm_parent_otel_span=litellm_parent_otel_span,
596 ) or {
597 "current_requests": 1,
598 "current_tpm": 0,
599 "current_rpm": 0,
600 }
602 new_val = {
603 "current_requests": max(current["current_requests"] - 1, 0),
604 "current_tpm": current["current_tpm"] + total_tokens,
605 "current_rpm": current["current_rpm"],
606 }
608 self.print_verbose(f"updated_value in success call: {new_val}, precise_minute: {precise_minute}")
609 values_to_update_in_cache.append((request_count_api_key, new_val))
611 # ------------
612 # Update usage - User
613 # ------------
614 if user_api_key_user_id is not None:
615 total_tokens = _response_total_tokens(response_obj)
617 request_count_api_key = f"{user_api_key_user_id}::{precise_minute}::request_count"
619 current = await self.internal_usage_cache.async_get_cache(
620 key=request_count_api_key,
621 litellm_parent_otel_span=litellm_parent_otel_span,
622 ) or {
623 "current_requests": 1,
624 "current_tpm": total_tokens,
625 "current_rpm": 1,
626 }
628 new_val = {
629 "current_requests": max(current["current_requests"] - 1, 0),
630 "current_tpm": current["current_tpm"] + total_tokens,
631 "current_rpm": current["current_rpm"],
632 }
634 self.print_verbose(f"updated_value in success call: {new_val}, precise_minute: {precise_minute}")
635 values_to_update_in_cache.append((request_count_api_key, new_val))
637 # ------------
638 # Update usage - Team
639 # ------------
640 if user_api_key_team_id is not None:
641 total_tokens = _response_total_tokens(response_obj)
643 request_count_api_key = f"{user_api_key_team_id}::{precise_minute}::request_count"
645 current = await self.internal_usage_cache.async_get_cache(
646 key=request_count_api_key,
647 litellm_parent_otel_span=litellm_parent_otel_span,
648 ) or {
649 "current_requests": 1,
650 "current_tpm": total_tokens,
651 "current_rpm": 1,
652 }
654 new_val = {
655 "current_requests": max(current["current_requests"] - 1, 0),
656 "current_tpm": current["current_tpm"] + total_tokens,
657 "current_rpm": current["current_rpm"],
658 }
660 self.print_verbose(f"updated_value in success call: {new_val}, precise_minute: {precise_minute}")
661 values_to_update_in_cache.append((request_count_api_key, new_val))
663 # ------------
664 # Update usage - End User
665 # ------------
666 if user_api_key_end_user_id is not None:
667 total_tokens = _response_total_tokens(response_obj)
669 request_count_api_key = f"{user_api_key_end_user_id}::{precise_minute}::request_count"
671 current = await self.internal_usage_cache.async_get_cache(
672 key=request_count_api_key,
673 litellm_parent_otel_span=litellm_parent_otel_span,
674 ) or {
675 "current_requests": 1,
676 "current_tpm": total_tokens,
677 "current_rpm": 1,
678 }
680 new_val = {
681 "current_requests": max(current["current_requests"] - 1, 0),
682 "current_tpm": current["current_tpm"] + total_tokens,
683 "current_rpm": current["current_rpm"],
684 }
686 self.print_verbose(f"updated_value in success call: {new_val}, precise_minute: {precise_minute}")
687 values_to_update_in_cache.append((request_count_api_key, new_val))
689 await self.internal_usage_cache.async_batch_set_cache(
690 cache_list=values_to_update_in_cache,
691 ttl=60,
692 litellm_parent_otel_span=litellm_parent_otel_span,
693 )
694 except Exception as e:
695 self.print_verbose(e)
697 async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
698 try:
699 self.print_verbose("Inside Max Parallel Request Failure Hook")
700 litellm_parent_otel_span: Final[Span | None] = _get_parent_otel_span_from_kwargs(kwargs=kwargs)
701 _metadata: Final = kwargs["litellm_params"].get("metadata", {}) or {}
702 global_max_parallel_requests: Final = _metadata.get("global_max_parallel_requests", None)
703 user_api_key: Final = _metadata.get("user_api_key", None)
704 self.print_verbose(f"user_api_key: [set={user_api_key is not None}]")
705 if user_api_key is None:
706 return
708 ## decrement call count if call failed
709 if CommonProxyErrors.max_parallel_request_limit_reached.value in str(kwargs["exception"]):
710 pass # ignore failed calls due to max limit being reached
711 else:
712 # ------------
713 # Setup values
714 # ------------
716 if global_max_parallel_requests is not None:
717 # get value from cache
718 _key: Final = "global_max_parallel_requests"
719 (
720 await self.internal_usage_cache.async_get_cache(
721 key=_key,
722 local_only=True,
723 litellm_parent_otel_span=litellm_parent_otel_span,
724 )
725 )
726 # decrement
727 await self.internal_usage_cache.async_increment_cache(
728 key=_key,
729 value=-1,
730 local_only=True,
731 litellm_parent_otel_span=litellm_parent_otel_span,
732 )
734 current_date: Final = datetime.now().strftime("%Y-%m-%d")
735 current_hour: Final = datetime.now().strftime("%H")
736 current_minute: Final = datetime.now().strftime("%M")
737 precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
739 request_count_api_key: Final = f"{user_api_key}::{precise_minute}::request_count"
741 # ------------
742 # Update usage
743 # ------------
744 current: Final = await self.internal_usage_cache.async_get_cache(
745 key=request_count_api_key,
746 litellm_parent_otel_span=litellm_parent_otel_span,
747 ) or {
748 "current_requests": 1,
749 "current_tpm": 0,
750 "current_rpm": 0,
751 }
753 new_val: Final = {
754 "current_requests": max(current["current_requests"] - 1, 0),
755 "current_tpm": current["current_tpm"],
756 "current_rpm": current["current_rpm"],
757 }
759 self.print_verbose(f"updated_value in failure call: {new_val}")
760 await self.internal_usage_cache.async_set_cache(
761 request_count_api_key,
762 new_val,
763 ttl=60,
764 litellm_parent_otel_span=litellm_parent_otel_span,
765 ) # save in cache for up to 1 min.
766 except Exception as e:
767 verbose_proxy_logger.exception("Inside Parallel Request Limiter: An exception occurred - %s", e)
769 async def get_internal_user_object(
770 self,
771 user_id: str,
772 user_api_key_dict: UserAPIKeyAuth,
773 ) -> dict | None:
774 """
775 Helper to get the 'Internal User Object'
777 It uses the `get_user_object` function from `litellm.proxy.auth.auth_checks`
779 We need this because the UserApiKeyAuth object does not contain the rpm/tpm limits for a User AND there could be a perf impact by additionally reading the UserTable.
780 """
781 from litellm._logging import verbose_proxy_logger
782 from litellm.proxy.auth.auth_checks import get_user_object
783 from litellm.proxy.proxy_server import prisma_client
785 try:
786 _user_id_rate_limits: Final = await get_user_object(
787 user_id=user_id,
788 prisma_client=prisma_client,
789 user_api_key_cache=self.internal_usage_cache.dual_cache,
790 user_id_upsert=False,
791 parent_otel_span=user_api_key_dict.parent_otel_span,
792 proxy_logging_obj=None,
793 )
795 if _user_id_rate_limits is None:
796 return None
798 return _user_id_rate_limits.model_dump()
799 except Exception as e:
800 verbose_proxy_logger.debug("Parallel Request Limiter: Error getting user object", str(e))
801 return None
803 async def async_post_call_success_hook(self, data: dict, user_api_key_dict: UserAPIKeyAuth, response):
804 """
805 Retrieve the key's remaining rate limits.
806 """
807 api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
808 current_date: Final = datetime.now().strftime("%Y-%m-%d")
809 current_hour: Final = datetime.now().strftime("%H")
810 current_minute: Final = datetime.now().strftime("%M")
811 precise_minute: Final = f"{current_date}-{current_hour}-{current_minute}"
812 request_count_api_key: Final = f"{api_key}::{precise_minute}::request_count"
813 current: Final[CurrentItemRateLimit | None] = await self.internal_usage_cache.async_get_cache(
814 key=request_count_api_key,
815 litellm_parent_otel_span=user_api_key_dict.parent_otel_span,
816 )
818 key_remaining_rpm_limit: int | None = None
819 key_rpm_limit: int | None = None
820 key_remaining_tpm_limit: int | None = None
821 key_tpm_limit: int | None = None
822 if current is not None:
823 if user_api_key_dict.rpm_limit is not None:
824 key_remaining_rpm_limit = user_api_key_dict.rpm_limit - current["current_rpm"]
825 key_rpm_limit = user_api_key_dict.rpm_limit
826 if user_api_key_dict.tpm_limit is not None:
827 key_remaining_tpm_limit = user_api_key_dict.tpm_limit - current["current_tpm"]
828 key_tpm_limit = user_api_key_dict.tpm_limit
830 if hasattr(response, "_hidden_params"):
831 _hidden_params = getattr(response, "_hidden_params")
832 else:
833 _hidden_params = None
834 if _hidden_params is not None and (isinstance(_hidden_params, BaseModel) or isinstance(_hidden_params, dict)):
835 if isinstance(_hidden_params, BaseModel):
836 _hidden_params = _hidden_params.model_dump()
838 _additional_headers: Final = _hidden_params.get("additional_headers", {}) or {}
840 if key_remaining_rpm_limit is not None:
841 _additional_headers["x-ratelimit-remaining-requests"] = key_remaining_rpm_limit
842 if key_rpm_limit is not None:
843 _additional_headers["x-ratelimit-limit-requests"] = key_rpm_limit
844 if key_remaining_tpm_limit is not None:
845 _additional_headers["x-ratelimit-remaining-tokens"] = key_remaining_tpm_limit
846 if key_tpm_limit is not None:
847 _additional_headers["x-ratelimit-limit-tokens"] = key_tpm_limit
849 setattr(
850 response,
851 "_hidden_params",
852 {**_hidden_params, "additional_headers": _additional_headers},
853 )
855 return await super().async_post_call_success_hook(data, user_api_key_dict, response)