diff --git a/litellm/__init__.py b/litellm/__init__.py index f22971dfa13..e6c30e12286 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -442,6 +442,13 @@ custom_prometheus_metadata_labels: List[str] = [] custom_prometheus_tags: List[str] = [] prometheus_metrics_config: Optional[List] = None prometheus_emit_stream_label: bool = False +# Opt-in: emit `rate_limit_category` and `rate_limit_type` labels on +# `litellm_proxy_failed_requests_metric`. Off by default to preserve the +# pre-unification label set so existing dashboards / recording rules keyed on +# that metric keep matching after upgrade. Enable when downstream consumers +# are ready to split 429s by source (vendor vs. litellm) and dimension +# (RPM/TPM/concurrent/budget). +prometheus_emit_rate_limit_labels: bool = False prometheus_user_budget_label_include_email_alias: bool = False prometheus_end_user_metrics_max_series_per_metric: Optional[int] = 10000 prometheus_end_user_metrics_ttl_seconds: Optional[float] = 3600.0 @@ -1303,6 +1310,8 @@ from .exceptions import ( NotFoundError, PermissionDeniedError, RateLimitError, + RateLimitErrorCategory, + RateLimitType, ServiceUnavailableError, BadGatewayError, OpenAIError, diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 15f6030d4a3..1cbef6b0b49 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -9,13 +9,109 @@ ## LiteLLM versions of the OpenAI Exception Types -from typing import Any, Dict, Optional +import enum +from typing import Any, Dict, Optional, Union import httpx import openai from litellm.types.utils import LiteLLMCommonStrings + +class RateLimitErrorCategory(str, enum.Enum): + """ + Category of a rate limit error, allowing callers to distinguish where the rate + limit originated. Exposed on every :class:`RateLimitError` instance via the + ``category`` attribute. + + Use these values to switch on the rate limit source, e.g.:: + + try: + ... + except litellm.RateLimitError as e: + if e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT: + ... # litellm's own limiter (key/team/user/model RPM/TPM/budget) + elif e.category == RateLimitErrorCategory.VENDOR_RATE_LIMIT: + ... # the upstream LLM provider returned 429 + """ + + VENDOR_RATE_LIMIT = "vendor_rate_limit" + """The upstream LLM provider returned a rate-limit response (e.g. OpenAI 429).""" + + VENDOR_BATCH_RATE_LIMIT = "vendor_batch_rate_limit" + """The upstream LLM provider returned a rate-limit response on a batch endpoint.""" + + LITELLM_RATE_LIMIT = "litellm_rate_limit" + """LiteLLM's own rate limiter (key/team/user/model RPM/TPM, budget, parallel-requests, etc.) blocked the request.""" + + LITELLM_BATCH_RATE_LIMIT = "litellm_batch_rate_limit" + """LiteLLM's own batch rate limiter (token/request budget across a batch input file) blocked the request.""" + + +class RateLimitType(str, enum.Enum): + """ + The dimension that was exceeded when a rate-limit error fired. + + This is orthogonal to :class:`RateLimitErrorCategory` — *category* tells + callers **who** rate-limited the request (the upstream vendor vs. one of + litellm's own limiters), while *type* tells them **which limit dimension** + was exceeded (an RPM ceiling, a TPM ceiling, a max-parallel-requests + ceiling, a budget cap, or a max-iterations cap). + + Surfaced both on every :class:`RateLimitError` instance via the + ``rate_limit_type`` attribute and on the structured + ``StandardLoggingPayload.error_information.error_rate_limit_type`` field + so custom callbacks / metrics consumers can split rate-limit failures by + cause without parsing free-text error messages. + """ + + REQUESTS = "requests" + """Requests-per-minute (RPM) or requests-per-window ceiling exceeded.""" + + TOKENS = "tokens" + """Tokens-per-minute (TPM) or tokens-per-window ceiling exceeded.""" + + CONCURRENT_REQUESTS = "concurrent_requests" + """``max_parallel_requests`` — too many in-flight requests at once.""" + + BUDGET = "budget" + """Spend budget cap reached (key, team, user, or per-session).""" + + MAX_ITERATIONS = "max_iterations" + """Per-session max-iterations cap reached (agent-style flows).""" + + +_RATE_LIMIT_CATEGORY_VALUES = frozenset(c.value for c in RateLimitErrorCategory) +_RATE_LIMIT_TYPE_VALUES = frozenset(t.value for t in RateLimitType) + + +def validate_rate_limit_category(value: Any) -> Optional[str]: + """Return ``value`` only if it matches a known :class:`RateLimitErrorCategory`. + + Used at duck-typed read sites (StandardLoggingPayload extraction, Prometheus + labels) to reject `.category` strings set by unrelated third-party exceptions + — otherwise those would leak into custom-callback payloads and Prometheus + label cardinality. + """ + if isinstance(value, RateLimitErrorCategory): + return value.value + if isinstance(value, str) and value in _RATE_LIMIT_CATEGORY_VALUES: + return value + return None + + +def validate_rate_limit_type(value: Any) -> Optional[str]: + """Return ``value`` only if it matches a known :class:`RateLimitType`. + + See :func:`validate_rate_limit_category` for the rationale. + """ + if isinstance(value, RateLimitType): + return value.value + if isinstance(value, str) and value in _RATE_LIMIT_TYPE_VALUES: + return value + return None + + _MINIMAL_ERROR_RESPONSE: Optional[httpx.Response] = None @@ -321,6 +417,18 @@ class PermissionDeniedError(openai.PermissionDeniedError): # type: ignore class RateLimitError(openai.RateLimitError): # type: ignore + """ + Unified rate-limit error. + + Every rate-limit condition surfaced by litellm — whether it originated from + an upstream LLM provider, a vendor batch endpoint, or one of litellm's own + proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget, + max-iterations, etc.) — is raised as an instance of this class. + + The :attr:`category` attribute lets callers distinguish the source. See + :class:`RateLimitErrorCategory` for the available values. + """ + def __init__( self, message, @@ -330,6 +438,12 @@ class RateLimitError(openai.RateLimitError): # type: ignore litellm_debug_info: Optional[str] = None, max_retries: Optional[int] = None, num_retries: Optional[int] = None, + category: Union[str, RateLimitErrorCategory] = ( + RateLimitErrorCategory.VENDOR_RATE_LIMIT + ), + rate_limit_type: Optional[Union[str, RateLimitType]] = None, + headers: Optional[Dict[str, str]] = None, + detail: Any = None, ): self.status_code = 429 self.message = "litellm.RateLimitError: {}".format(message) @@ -338,9 +452,39 @@ class RateLimitError(openai.RateLimitError): # type: ignore self.litellm_debug_info = litellm_debug_info self.max_retries = max_retries self.num_retries = num_retries + self.category = ( + category.value if isinstance(category, RateLimitErrorCategory) else category + ) + # Which dimension was exceeded — request count, token count, parallel + # requests, budget, max iterations. None when the source didn't + # classify the failure (e.g. legacy vendor 429 with no header hints). + self.rate_limit_type: Optional[str] = ( + rate_limit_type.value + if isinstance(rate_limit_type, RateLimitType) + else rate_limit_type + ) + # Headers explicitly attached to the error (e.g. retry-after, + # rate_limit_type, reset_at). Preserved across the proxy boundary so + # clients can react appropriately. + # + # IMPORTANT: we deliberately do NOT auto-populate self.headers from + # response.headers when only `response` is provided. A vendor 429 can + # set arbitrary response headers (Set-Cookie, CORS overrides, …); if + # those leaked into e.headers and a downstream proxy serializer + # forwarded them to the client, a malicious upstream could inject + # browser-interpreted headers for the proxy origin. Vendor response + # headers stay reachable on `e.response.headers` for callers that + # explicitly want them; only the proxy-supplied `headers=` kwarg + # makes it onto `self.headers`. _response_headers = ( getattr(response, "headers", None) if response is not None else None ) + self.headers: Optional[Dict[str, str]] = ( + {k: str(v) for k, v in headers.items()} if headers else None + ) + # Mirrors FastAPI HTTPException.detail so the same instance can be + # serialized through both the ProxyException and HTTPException paths. + self.detail = detail if detail is not None else self.message self.response = httpx.Response( status_code=429, headers=_response_headers, @@ -843,11 +987,24 @@ LITELLM_EXCEPTION_TYPES = [ class BudgetExceededError(Exception): def __init__( - self, current_cost: float, max_budget: float, message: Optional[str] = None + self, + current_cost: float, + max_budget: float, + message: Optional[str] = None, + llm_provider: Optional[str] = None, ): self.current_cost = current_cost self.max_budget = max_budget self.status_code = 429 + self.llm_provider = llm_provider or "" + # Surface unified rate-limit fields without joining the RateLimitError + # hierarchy so existing `except BudgetExceededError:` handlers keep + # working; custom callbacks reading StandardLoggingPayload pick these + # up via the same `category` / `rate_limit_type` attributes the rest + # of the unified rate-limit error path uses. Stored as plain strings + # to match the normalization RateLimitError.__init__ performs. + self.category: str = RateLimitErrorCategory.LITELLM_RATE_LIMIT.value + self.rate_limit_type: str = RateLimitType.BUDGET.value message = ( message or f"Budget has been exceeded! Current cost: {current_cost}, Max budget: {max_budget}" diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 648fe671140..d2af95cd4cc 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -24,6 +24,10 @@ from typing import ( import litellm from litellm._logging import print_verbose, verbose_logger +from litellm.exceptions import ( + validate_rate_limit_category, + validate_rate_limit_type, +) from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus_helpers.bounded_prometheus_series_tracker import ( BoundedPrometheusSeriesTracker, @@ -78,6 +82,20 @@ class PrometheusLogger(CustomLogger): # Always initialize label_filters, even for non-premium users self.label_filters = self._parse_prometheus_config() + # Cache resolved label sets per metric. Several entries in + # ``PrometheusMetricLabels.get_labels`` read module-level toggles + # (e.g. ``litellm.prometheus_emit_stream_label``, + # ``litellm.prometheus_emit_rate_limit_labels``) that can be + # changed at runtime. Prometheus counters/gauges/histograms are + # created with a *fixed* ``labelnames`` set; if a runtime call + # to ``get_labels_for_metric`` returned a different set, the + # subsequent ``counter.labels(**_labels)`` would raise a + # ``ValueError`` from the prometheus client. Snapshotting at + # logger init time pins the label set for the lifetime of the + # logger so toggling these flags only takes effect after a + # restart, keeping init-time and runtime label sets in sync. + self._cached_metric_labels: Dict[str, List[str]] = {} + _custom_buckets = litellm.prometheus_latency_buckets self.latency_buckets = ( tuple(_custom_buckets) @@ -1033,13 +1051,27 @@ class PrometheusLogger(CustomLogger): self, metric_name: DEFINED_PROMETHEUS_METRICS ) -> List[str]: """ - Get the labels for a metric, filtered if configured + Get the labels for a metric, filtered if configured. + + The result is cached on the instance so the label set used to + construct each Prometheus metric at ``__init__`` time stays in lock + step with the label set passed to ``counter.labels(...)`` at + runtime, even if the underlying module-level toggles consulted by + :meth:`PrometheusMetricLabels.get_labels` (e.g. + ``litellm.prometheus_emit_rate_limit_labels``, + ``litellm.prometheus_emit_stream_label``) are flipped after the + logger has been created. """ + cached = self._cached_metric_labels.get(metric_name) + if cached is not None: + return cached + # Get default labels for this metric from PrometheusMetricLabels default_labels = PrometheusMetricLabels.get_labels(metric_name) # If no label filtering is configured for this metric, use default labels if metric_name not in self.label_filters: + self._cached_metric_labels[metric_name] = default_labels return default_labels # Get configured labels for this metric @@ -1050,6 +1082,7 @@ class PrometheusLogger(CustomLogger): label for label in default_labels if label in configured_labels ] + self._cached_metric_labels[metric_name] = filtered_labels return filtered_labels def _track_end_user_metric_series( @@ -2029,14 +2062,8 @@ class PrometheusLogger(CustomLogger): Proxy level tracking - failed client side requests - labelnames=[ - "end_user", - "hashed_api_key", - "api_key_alias", - REQUESTED_MODEL, - "team", - "team_alias", - ] + EXCEPTION_LABELS, + See :attr:`PrometheusMetricLabels.litellm_proxy_failed_requests_metric` + for the authoritative list of labels emitted on this metric. """ from litellm.litellm_core_utils.litellm_logging import ( StandardLoggingPayloadSetup, @@ -2059,6 +2086,9 @@ class PrometheusLogger(CustomLogger): model_id = _metadata.get("model_info", {}).get("id") or request_data.get( "model_info", {} ).get("id") + rate_limit_category, rate_limit_type = self._extract_rate_limit_labels( + original_exception + ) enum_values = UserAPIKeyLabelValues( end_user=user_api_key_dict.end_user_id, user=user_api_key_dict.user_id, @@ -2073,6 +2103,8 @@ class PrometheusLogger(CustomLogger): status_code=str(status_code), exception_status=str(status_code), exception_class=self._get_exception_class_name(original_exception), + rate_limit_category=rate_limit_category, + rate_limit_type=rate_limit_type, tags=_tags, route=user_api_key_dict.request_route, client_ip=_metadata.get("requester_ip_address"), @@ -2843,6 +2875,33 @@ class PrometheusLogger(CustomLogger): @staticmethod def _get_exception_class_name(exception: Exception) -> str: + # Some exception types pin the ``exception_class`` label to a legacy + # value for back-compat with existing dashboards (e.g. proxy-side 429s + # keep reporting as "HTTPException"). Honor that opt-in marker before + # deriving the label from the runtime class name. Reading it via + # ``getattr`` keeps this core integrations module free of a transitive + # ``fastapi`` dependency. + legacy_class_name = getattr(exception, "prometheus_exception_class_name", None) + if isinstance(legacy_class_name, str) and legacy_class_name: + return legacy_class_name + + # Same back-compat reasoning for ``BudgetExceededError``: the unified + # rate-limit error work attached ``.llm_provider`` to budget errors + # too (so callbacks reading ``StandardLoggingPayload`` get provider + # attribution). Without this short-circuit, the provider prefix below + # would silently flip the label from "BudgetExceededError" to e.g. + # "Openai.BudgetExceededError" and break dashboards keyed on the + # original value. + try: + from litellm.exceptions import BudgetExceededError + except ImportError: + BudgetExceededError = None # type: ignore[assignment,misc] + + if BudgetExceededError is not None and isinstance( + exception, BudgetExceededError + ): + return "BudgetExceededError" + exception_class_name = "" if hasattr(exception, "llm_provider"): exception_class_name = getattr(exception, "llm_provider") or "" @@ -2857,6 +2916,27 @@ class PrometheusLogger(CustomLogger): exception_class_name += exception.__class__.__name__ return exception_class_name + @staticmethod + def _extract_rate_limit_labels( + exception: Optional[Exception], + ) -> Tuple[Optional[str], Optional[str]]: + """ + Pull the unified ``category`` / ``rate_limit_type`` fields off any + exception that declares them (``litellm.RateLimitError`` and bare- + Exception subclasses like ``BudgetExceededError``). + + Values are validated against the :class:`RateLimitErrorCategory` / + :class:`RateLimitType` enums so unrelated third-party exceptions that + happen to declare ``.category`` / ``.rate_limit_type`` string attributes + can't leak garbage into Prometheus label cardinality. + """ + if exception is None: + return None, None + return ( + validate_rate_limit_category(getattr(exception, "category", None)), + validate_rate_limit_type(getattr(exception, "rate_limit_type", None)), + ) + async def log_success_fallback_event( self, original_model_group: str, kwargs: dict, original_exception: Exception ): diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index f20b66790c4..dbfcf55d75d 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -37,6 +37,10 @@ from litellm import ( turn_off_message_logging, ) from litellm._logging import _is_debugging_on, _redact_string, verbose_logger +from litellm.exceptions import ( + validate_rate_limit_category, + validate_rate_limit_type, +) from litellm._uuid import uuid from litellm.batches.batch_utils import _handle_completed_batch from litellm.caching.caching import DualCache, InMemoryCache @@ -5318,12 +5322,27 @@ class StandardLoggingPayloadSetup: else str(original_exception) ) + # Duck-typed read so bare-Exception subclasses like + # `litellm.BudgetExceededError` can participate without joining the + # RateLimitError hierarchy (which would break `except BudgetExceededError`). + # Validated against the enum value sets so a third-party exception that + # happens to declare a `.category` or `.rate_limit_type` string attribute + # can't leak garbage into the payload or Prometheus label cardinality. + rate_limit_category = validate_rate_limit_category( + getattr(original_exception, "category", None) + ) + rate_limit_type = validate_rate_limit_type( + getattr(original_exception, "rate_limit_type", None) + ) + return StandardLoggingPayloadErrorInformation( error_code=error_status, error_class=error_class, llm_provider=_llm_provider_in_exception, traceback=traceback_info, error_message=error_message if original_exception else "", + error_rate_limit_category=rate_limit_category, + error_rate_limit_type=rate_limit_type, ) @staticmethod diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index e5d70933063..88be567e59a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -4125,6 +4125,10 @@ class TeamMemberUpdateRequest(TeamMemberDeleteRequest): rpm_limit: Optional[int] = Field( default=None, description="Requests per minute limit for this team member" ) + budget_duration: Optional[str] = Field( + default=None, + description="Duration after which this team member's budget resets (e.g. '1h', '24h', '7d', '30d'). If not set, the budget never resets.", + ) allowed_models: Optional[List[str]] = Field( default=None, description="List of models this team member can access. Pass an empty list to remove per-member model restrictions.", @@ -4136,6 +4140,7 @@ class TeamMemberUpdateResponse(MemberUpdateResponse): max_budget_in_team: Optional[float] = None tpm_limit: Optional[int] = None rpm_limit: Optional[int] = None + budget_duration: Optional[str] = None allowed_models: Optional[List[str]] = None diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index e06ac760237..f76949f4d11 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -126,6 +126,20 @@ class UserAPIKeyAuthExceptionHandler: model=request_data.get("model"), ) + # Budget checks live in tenant-scoped helpers (key / team / org / tag) + # that don't see the request model, so the BudgetExceededError they + # raise carries `llm_provider=""`. Resolve it here off `request_data` + # so custom-callback consumers reading StandardLoggingPayload get + # the same `llm_provider` attribution as for RPM/TPM 429s. + if isinstance(e, litellm.BudgetExceededError) and not e.llm_provider: + from litellm.proxy.hooks.rate_limiter_utils import ( + resolve_llm_provider_for_rate_limit, + ) + + _, e.llm_provider = resolve_llm_provider_for_rate_limit( + request_data.get("model") + ) + # Allow callbacks to transform the error response transformed_exception = await proxy_logging_obj.post_call_failure_hook( request_data=request_data, diff --git a/litellm/proxy/common_utils/proxy_rate_limit_error.py b/litellm/proxy/common_utils/proxy_rate_limit_error.py new file mode 100644 index 00000000000..24e5c991794 --- /dev/null +++ b/litellm/proxy/common_utils/proxy_rate_limit_error.py @@ -0,0 +1,196 @@ +""" +ProxyRateLimitError — a unified rate-limit exception used by litellm's +proxy-side hooks. + +Background +---------- +LiteLLM previously surfaced rate-limit conditions through *several* unrelated +exception types: + +* :class:`litellm.exceptions.RateLimitError` — raised by exception mapping when + an upstream LLM provider returns 429. +* :class:`fastapi.HTTPException` (status 429) — raised directly by proxy hooks + such as ``parallel_request_limiter``, ``dynamic_rate_limiter``, + ``batch_rate_limiter``, ``max_budget_limiter``, ``max_iterations_limiter``, + etc. +* :class:`litellm.llms.base_llm.chat.transformation.BaseLLMException` (status + 429) — raised by some provider transports. + +This made it impossible for downstream code (and end users) to express +"is this a rate limit?" with a single ``except`` clause, and impossible to +distinguish *where* the rate limit originated (vendor vs. litellm, batch vs. +chat) without ad-hoc string-matching on the message. + +This module provides a single proxy-side error class that: + +1. Is a subclass of :class:`litellm.exceptions.RateLimitError`, so user code + that catches ``RateLimitError`` works for *every* rate-limit source. +2. Is also a subclass of :class:`fastapi.HTTPException`, so existing proxy + plumbing (``isinstance(e, HTTPException)`` branches in route handlers and + FastAPI's own dispatcher) continues to behave the same way and the + ``retry-after`` / ``rate_limit_type`` / ``reset_at`` headers are preserved + on the wire. +3. Carries a :attr:`category` field (one of + :class:`litellm.exceptions.RateLimitErrorCategory`) so callers can switch on + the rate limit source. +""" + +import json +from typing import Any, Dict, Mapping, Optional, Union + +from fastapi import HTTPException + +from litellm.exceptions import RateLimitError, RateLimitErrorCategory, RateLimitType + + +def map_v3_rate_limit_type( + v3_value: Optional[str], +) -> Optional[RateLimitType]: + """ + Map the v3 rate limiter's internal `status["rate_limit_type"]` strings + onto the public :class:`RateLimitType` enum. + + The v3 limiter uses the literal values ``"requests"``, ``"tokens"``, and + ``"max_parallel_requests"``. We collapse the last one onto + :attr:`RateLimitType.CONCURRENT_REQUESTS` because that's the public name + documented for users and dashboards. Unrecognized values return ``None`` + so the field stays absent rather than carrying garbage downstream. + """ + if v3_value == "tokens": + return RateLimitType.TOKENS + if v3_value == "max_parallel_requests": + return RateLimitType.CONCURRENT_REQUESTS + if v3_value == "requests": + return RateLimitType.REQUESTS + return None + + +def _coerce_message(detail: Any) -> str: + """Best-effort, JSON-friendly stringification of an HTTPException-style detail.""" + if detail is None: + return "" + if isinstance(detail, str): + return detail + if isinstance(detail, Mapping): + for key in ("error", "message"): + if isinstance(detail.get(key), str): + return detail[key] + inner = detail.get(key) + if isinstance(inner, Mapping) and isinstance(inner.get("message"), str): + return inner["message"] + try: + return json.dumps(detail) + except (TypeError, ValueError): + return str(detail) + return str(detail) + + +# NOTE: mypy emits two `[misc]` errors on the class line below because the +# bases declare overlapping attributes with related-but-not-identical +# annotations: +# * `status_code` is `int` on starlette HTTPException but `Literal[429]` on +# openai.RateLimitError (every openai status-error subclass narrows it +# this way and silences pyright with the same convention). +# * `headers` is `Mapping[str, str] | None` on HTTPException; we narrow it +# to `Optional[Dict[str, str]]` on RateLimitError because we always carry +# a stringified dict. +# Both narrowings are intentional and handled at construction time — every +# instance always has status_code == 429 and a Dict-typed headers — so we +# silence the ATTR-overlap check rather than relax the annotations. +class ProxyRateLimitError(HTTPException, RateLimitError): # type: ignore[misc] + """ + A 429 raised by litellm's proxy-side rate limiting hooks. + + This class deliberately inherits from BOTH + :class:`litellm.exceptions.RateLimitError` and :class:`fastapi.HTTPException` + so the same instance can flow through: + + * ``except RateLimitError`` (user / SDK code that wants a category-aware + handler), and + * ``isinstance(e, HTTPException)`` (FastAPI / proxy_server.py route + handlers that need to forward ``status_code``, ``detail`` and + ``headers`` back to the client). + + Downstream code should prefer this class over + ``raise HTTPException(status_code=429, ...)`` for litellm-internal rate + limits. + + Parameters + ---------- + detail: + The structured error payload. Forwarded as ``HTTPException.detail`` so + FastAPI's default exception handler will serialize it verbatim. + headers: + Optional response headers (e.g. ``retry-after``). Values are stringified + to satisfy FastAPI's typing. + category: + One of :class:`RateLimitErrorCategory`. Defaults to + ``LITELLM_RATE_LIMIT`` since this class is only used by litellm's own + proxy-side limiters; pass ``LITELLM_BATCH_RATE_LIMIT`` for the batch + limiter, etc. + model / llm_provider: + Optional context, propagated to the inherited ``RateLimitError`` for + compatibility with logging / standard payload extraction. + """ + + # Prometheus' ``exception_class`` label is pinned to "HTTPException" for + # this type: before the unified class existed, proxy-side 429s surfaced as + # ``fastapi.HTTPException`` and existing dashboards/alerts key off that exact + # value. Distinguishing vendor vs. litellm 429s is now the job of the + # ``rate_limit_category`` / ``rate_limit_type`` labels. + prometheus_exception_class_name = "HTTPException" + + def __init__( + self, + detail: Any, + headers: Optional[Mapping[str, Any]] = None, + category: Union[ + str, RateLimitErrorCategory + ] = RateLimitErrorCategory.LITELLM_RATE_LIMIT, + rate_limit_type: Optional[Union[str, RateLimitType]] = None, + model: Optional[str] = None, + llm_provider: Optional[str] = "litellm_proxy", + ): + # Normalize None → safe defaults so callers (and the resolver helper + # in `rate_limiter_utils`) can pass `None` without producing an + # instance whose `.llm_provider` attribute is `None` — that would + # break Prometheus' `_get_exception_class_name` (it calls + # `.capitalize()` on the provider string). + model = model or "" + llm_provider = llm_provider or "litellm_proxy" + message = _coerce_message(detail) + stringified_headers: Optional[Dict[str, str]] = ( + {k: str(v) for k, v in headers.items()} if headers else None + ) + + # Initialize the FastAPI HTTPException portion first so its attributes + # (status_code, detail, headers) are already on the instance before + # RateLimitError.__init__ runs and possibly overrides them. + HTTPException.__init__( + self, + status_code=429, + detail=detail, + headers=stringified_headers, + ) + + # Now initialize the litellm RateLimitError portion. We deliberately + # pass the structured detail through so RateLimitError preserves it as + # its `.detail` attribute too — keeping both sides of the MRO + # consistent. + RateLimitError.__init__( + self, + message=message, + llm_provider=llm_provider, + model=model, + category=category, + rate_limit_type=rate_limit_type, + headers=stringified_headers, + detail=detail, + ) + # RateLimitError.__init__ overwrites self.headers with its own copy and + # leaves self.status_code at 429 — restore the HTTPException-style + # headers value so downstream code that pulls headers off the + # instance gets back exactly what the limiter passed in. + self.headers = stringified_headers + self.detail = detail + self.status_code = 429 diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 8473b5e77de..3957e3a7fbb 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -17,7 +17,17 @@ Quick summary: - async_log_success_event() fires on GET /v1/batches/{id} (batch completion) """ -from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Tuple, Union +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + NoReturn, + Optional, + Tuple, + Union, +) from fastapi import HTTPException from pydantic import BaseModel @@ -30,6 +40,7 @@ from litellm.batches.batch_utils import ( _get_file_content_as_dictionary, _get_models_from_batch_input_file_content, ) +from litellm.exceptions import RateLimitErrorCategory from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import ( ProxyErrorTypes, @@ -37,10 +48,11 @@ from litellm.proxy._types import ( SpecialModelNames, UserAPIKeyAuth, ) -from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, - resolve_llm_provider_for_rate_limit, +from litellm.proxy.common_utils.proxy_rate_limit_error import ( + ProxyRateLimitError, + map_v3_rate_limit_type, ) +from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -385,8 +397,8 @@ class _PROXY_BatchRateLimiter(CustomLogger): batch_usage: BatchFileUsage, limit_type: str, requested_model: Optional[str] = None, - ) -> None: - """Raise HTTPException for rate limit exceeded.""" + ) -> NoReturn: + """Raise :class:`ProxyRateLimitError` (a 429) for batch rate limit exceeded.""" from datetime import datetime # Find the descriptor for this status @@ -432,14 +444,15 @@ class _PROXY_BatchRateLimiter(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( requested_model ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail=detail, headers={ "retry-after": str(window_size), "rate_limit_type": limit_type, "reset_at": reset_time_formatted, }, + category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT, + rate_limit_type=map_v3_rate_limit_type(limit_type), model=resolved_model, llm_provider=llm_provider, ) diff --git a/litellm/proxy/hooks/dynamic_rate_limiter.py b/litellm/proxy/hooks/dynamic_rate_limiter.py index 57cd538507e..b9e2bd12ecf 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter.py @@ -11,9 +11,10 @@ from litellm import ModelResponse, Router from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.exceptions import RateLimitType from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, convert_priority_to_percent, resolve_llm_provider_for_rate_limit, ) @@ -222,8 +223,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( data.get("model") ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail={ "error": "Key={} over available TPM={}. Model TPM={}, Active keys={}".format( user_api_key_dict.api_key, @@ -232,6 +232,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): active_projects, ) }, + rate_limit_type=RateLimitType.TOKENS, model=resolved_model, llm_provider=llm_provider, ) @@ -240,8 +241,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( data.get("model") ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail={ "error": "Key={} over available RPM={}. Model RPM={}, Active keys={}".format( user_api_key_dict.api_key, @@ -250,6 +250,7 @@ class _PROXY_DynamicRateLimitHandler(CustomLogger): active_projects, ) }, + rate_limit_type=RateLimitType.REQUESTS, model=resolved_model, llm_provider=llm_provider, ) diff --git a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py index bfc6e2c2f72..493afe6105a 100644 --- a/litellm/proxy/hooks/dynamic_rate_limiter_v3.py +++ b/litellm/proxy/hooks/dynamic_rate_limiter_v3.py @@ -14,13 +14,16 @@ from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ( + ProxyRateLimitError, + map_v3_rate_limit_type, +) from litellm.proxy.hooks.parallel_request_limiter_v3 import ( RateLimitDescriptor, RateLimitDescriptorRateLimitObject, _PROXY_MaxParallelRequestsHandler_v3, ) from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, convert_priority_to_percent, resolve_llm_provider_for_rate_limit, ) @@ -497,8 +500,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): continue descriptor_key = status["descriptor_key"] if descriptor_key == "model_saturation_check": - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail={ "error": f"Model capacity reached for {model}. " f"Priority: {priority}, " @@ -512,6 +514,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): "rate_limit_type": str(status["rate_limit_type"]), "x-litellm-priority": priority or "default", }, + rate_limit_type=map_v3_rate_limit_type( + status["rate_limit_type"] + ), model=resolved_model, llm_provider=llm_provider, ) @@ -520,8 +525,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): f"Enforcing priority limits for {model}, saturation: {saturation:.1%}, " f"priority: {priority}" ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail={ "error": f"Priority-based rate limit exceeded. " f"Model: {model}, " @@ -538,6 +542,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): "x-litellm-priority": priority or "default", "x-litellm-saturation": f"{saturation:.2%}", }, + rate_limit_type=map_v3_rate_limit_type( + status["rate_limit_type"] + ), model=resolved_model, llm_provider=llm_provider, ) @@ -556,8 +563,7 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): f"Dynamic rate limiter: OVER_LIMIT response with unknown " f"descriptor_key(s) — refusing request. response={atomic_response}" ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail={ "error": "Rate limit exceeded", "descriptor_key": ( @@ -567,6 +573,9 @@ class _PROXY_DynamicRateLimitHandlerV3(CustomLogger): str(offending["rate_limit_type"]) if offending else "unknown" ), }, + rate_limit_type=map_v3_rate_limit_type( + offending["rate_limit_type"] if offending else None + ), headers={ "retry-after": str(self.v3_limiter.window_size), "x-litellm-priority": priority or "default", diff --git a/litellm/proxy/hooks/max_budget_limiter.py b/litellm/proxy/hooks/max_budget_limiter.py index 658d7995631..769348a0b88 100644 --- a/litellm/proxy/hooks/max_budget_limiter.py +++ b/litellm/proxy/hooks/max_budget_limiter.py @@ -4,11 +4,10 @@ from litellm import verbose_logger from litellm._logging import verbose_proxy_logger from litellm.caching.caching import DualCache from litellm.integrations.custom_logger import CustomLogger +from litellm.exceptions import RateLimitType from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, - resolve_llm_provider_for_rate_limit, -) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit class _PROXY_MaxBudgetLimiter(CustomLogger): @@ -70,9 +69,9 @@ class _PROXY_MaxBudgetLimiter(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( data.get("model") if data else None ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail="Max budget limit reached.", + rate_limit_type=RateLimitType.BUDGET, model=resolved_model, llm_provider=llm_provider, ) diff --git a/litellm/proxy/hooks/max_budget_per_session_limiter.py b/litellm/proxy/hooks/max_budget_per_session_limiter.py index 0b63465c4a5..20bfeb3a6d5 100644 --- a/litellm/proxy/hooks/max_budget_per_session_limiter.py +++ b/litellm/proxy/hooks/max_budget_per_session_limiter.py @@ -20,11 +20,10 @@ from typing import TYPE_CHECKING, Any, Optional, Union from litellm import DualCache from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.exceptions import RateLimitType from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, - resolve_llm_provider_for_rate_limit, -) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit if TYPE_CHECKING: from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache @@ -117,13 +116,13 @@ class _PROXY_MaxBudgetPerSessionHandler(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( data.get("model") if data else None ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail=( f"Session budget exceeded for session {session_id}. " f"Current spend: ${current_spend:.4f}, " f"max_budget_per_session: ${max_budget:.2f}." ), + rate_limit_type=RateLimitType.BUDGET, model=resolved_model, llm_provider=llm_provider, ) diff --git a/litellm/proxy/hooks/max_iterations_limiter.py b/litellm/proxy/hooks/max_iterations_limiter.py index d5bc669c928..525214ff6be 100644 --- a/litellm/proxy/hooks/max_iterations_limiter.py +++ b/litellm/proxy/hooks/max_iterations_limiter.py @@ -16,11 +16,10 @@ from typing import TYPE_CHECKING, Any, Optional, Union from litellm import DualCache from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.exceptions import RateLimitType from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, - resolve_llm_provider_for_rate_limit, -) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit if TYPE_CHECKING: from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache @@ -121,12 +120,12 @@ class _PROXY_MaxIterationsHandler(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( data.get("model") if data else None ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail=( f"Max iterations exceeded for session {session_id}. " f"Current count: {current_count}, max_iterations: {max_iterations}." ), + rate_limit_type=RateLimitType.MAX_ITERATIONS, model=resolved_model, llm_provider=llm_provider, ) diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index c6324c3e3a3..b622241dfa5 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -1,9 +1,8 @@ import asyncio import sys from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Any, List, Literal, Optional, Tuple, Union +from typing import TYPE_CHECKING, Any, List, Literal, NoReturn, Optional, Tuple, Union -from fastapi import HTTPException from pydantic import BaseModel from typing_extensions import TypedDict @@ -13,14 +12,13 @@ from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs from litellm.proxy._types import CommonProxyErrors, CurrentItemRateLimit, UserAPIKeyAuth +from litellm.exceptions import RateLimitType from litellm.proxy.auth.auth_utils import ( get_key_model_rpm_limit, get_key_model_tpm_limit, ) -from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, - resolve_llm_provider_for_rate_limit, -) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit if TYPE_CHECKING: from opentelemetry.trace import Span as _Span @@ -75,9 +73,21 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): ) if current is None: if max_parallel_requests == 0 or tpm_limit == 0 or rpm_limit == 0: - # base case - raise self.raise_rate_limit_error( + # base case — at least one dimension is set to 0 (effectively + # disabled). Pick the most specific dimension as the + # rate_limit_type so dashboards can attribute the failure to + # the right cap. Order matters: max_parallel_requests is + # listed first because it's the rarest 0 in practice and the + # most actionable signal. + if max_parallel_requests == 0: + triggered_type = RateLimitType.CONCURRENT_REQUESTS + elif tpm_limit == 0: + triggered_type = RateLimitType.TOKENS + else: + triggered_type = RateLimitType.REQUESTS + self.raise_rate_limit_error( 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}", + rate_limit_type=triggered_type, requested_model=data.get("model") if data else None, ) new_val = { @@ -100,14 +110,23 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): values_to_update_in_cache.append((request_count_api_key, new_val)) else: + # Detect which dimension actually tripped the limit so we can + # surface the right rate_limit_type. Order matches the boolean + # condition above (concurrent → tpm → rpm) — first match wins. + if int(current["current_requests"]) >= max_parallel_requests: + triggered_type = RateLimitType.CONCURRENT_REQUESTS + elif current["current_tpm"] >= tpm_limit: + triggered_type = RateLimitType.TOKENS + else: + triggered_type = RateLimitType.REQUESTS requested_model = data.get("model") if data else None resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( requested_model ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( 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}", headers={"retry-after": str(self.time_to_next_minute())}, + rate_limit_type=triggered_type, model=resolved_model, llm_provider=llm_provider, ) @@ -135,27 +154,45 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): def raise_rate_limit_error( self, additional_details: Optional[str] = None, + rate_limit_type: Optional[RateLimitType] = None, requested_model: Optional[str] = None, - ) -> HTTPException: + ) -> NoReturn: """ - Raise an HTTPException with a 429 status code and a retry-after header. + Raise a 429 with a retry-after header for litellm-proxy parallel-request limits. + + Always raises :class:`ProxyRateLimitError` — never returns. Annotated + ``NoReturn`` so type-checkers know callers after this invocation are + unreachable. The raised exception is both a + :class:`litellm.RateLimitError` (so callers can catch by category) and a + :class:`fastapi.HTTPException` (so the FastAPI dispatcher serializes it + correctly with status 429 and the supplied headers). + + ``rate_limit_type`` defaults to ``CONCURRENT_REQUESTS`` because every + existing internal caller of this helper hits the parallel-request cap + (the global-limit branch in ``async_pre_call_hook`` and the + all-zeros base case in ``check_key_in_limits``). Callers that know + the dimension exactly should pass it explicitly. ``requested_model`` is resolved via :func:`get_llm_provider` so the - raised exception carries ``llm_provider`` for downstream loggers - (Prometheus failure metric, observability callbacks). Falls back to - ``llm_provider="litellm_proxy"`` when the model is missing or - unparseable — see ``resolve_llm_provider_for_rate_limit``. + raised exception carries ``llm_provider`` (and a stripped ``model``) + for downstream loggers (Prometheus failure metric, observability + callbacks). Falls back to ``llm_provider="litellm_proxy"`` when the + model is missing or unparseable — see + :func:`resolve_llm_provider_for_rate_limit`. """ + # additional_details is optional; build the detail with a None-guard + # so callers that pass nothing don't get the literal string "None" + # interpolated into the error message. error_message = "Max parallel request limit reached" if additional_details is not None: error_message = error_message + " " + additional_details resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( requested_model ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail=error_message, headers={"retry-after": str(self.time_to_next_minute())}, + rate_limit_type=rate_limit_type or RateLimitType.CONCURRENT_REQUESTS, model=resolved_model, llm_provider=llm_provider, ) @@ -248,7 +285,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): current_global_requests = 1 # if above -> raise error if current_global_requests >= global_max_parallel_requests: - return self.raise_rate_limit_error( + self.raise_rate_limit_error( additional_details=f"Hit Global Limit: Limit={global_max_parallel_requests}, current: {current_global_requests}", requested_model=data.get("model") if data else None, ) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 9fdb146b19d..62751fb68a4 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -32,10 +32,11 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata -from litellm.proxy.hooks.rate_limiter_utils import ( - ProxyHTTPRateLimitError, - resolve_llm_provider_for_rate_limit, +from litellm.proxy.common_utils.proxy_rate_limit_error import ( + ProxyRateLimitError, + map_v3_rate_limit_type, ) +from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit from litellm.types.caching import RedisPipelineIncrementOperation from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject from litellm.types.utils import CallTypes, ModelResponse, Usage @@ -1971,7 +1972,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): descriptors: List[RateLimitDescriptor], requested_model: Optional[str] = None, ) -> None: - """Handle rate limit exceeded error by raising HTTPException.""" + """Handle rate limit exceeded by raising :class:`ProxyRateLimitError` (a 429).""" for status in response["statuses"]: if status["code"] == "OVER_LIMIT": descriptor_key = status["descriptor_key"] @@ -2005,14 +2006,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): resolved_model, llm_provider = resolve_llm_provider_for_rate_limit( requested_model ) - raise ProxyHTTPRateLimitError( - status_code=429, + raise ProxyRateLimitError( detail=detail, headers={ "retry-after": str(self.window_size), "rate_limit_type": str(status["rate_limit_type"]), "reset_at": reset_time_formatted, }, + rate_limit_type=map_v3_rate_limit_type(status["rate_limit_type"]), model=resolved_model, llm_provider=llm_provider, ) diff --git a/litellm/proxy/hooks/rate_limiter_utils.py b/litellm/proxy/hooks/rate_limiter_utils.py index 0ba3df448e5..07440975476 100644 --- a/litellm/proxy/hooks/rate_limiter_utils.py +++ b/litellm/proxy/hooks/rate_limiter_utils.py @@ -2,13 +2,10 @@ Shared utility functions for rate limiter hooks. """ -from typing import Any, Optional, Tuple, Union - -from fastapi import HTTPException +from typing import Optional, Tuple, Union import litellm from litellm._logging import verbose_proxy_logger -from litellm.exceptions import RateLimitError from litellm.types.router import ModelGroupInfo from litellm.types.utils import PriorityReservationDict @@ -29,11 +26,21 @@ def resolve_llm_provider_for_rate_limit( ``litellm_proxy_failed_requests_metric`` show up with ``exception_class="RateLimitError"`` and no provider attribution. - Wrapped defensively: if ``model`` is missing, malformed, or - ``get_llm_provider`` raises (unknown alias, router-only model, etc.) we - fall back to ``("", "litellm_proxy")`` so we never break the request path - by piling a second exception on top of the rate-limit one we're trying to - raise. + Resolution order: + + 1. ``litellm.get_llm_provider(model)`` — covers raw provider/model + strings the SDK already understands (``"gpt-4o-mini"``, + ``"anthropic/claude-3-5-sonnet"``, ``"bedrock/..."`` etc.). + 2. **Router alias fallback** — nearly every real proxy deployment + routes through a router ``model_name`` alias (e.g. + ``"tpm-locked"`` → ``litellm_params.model: openai/gpt-4o-mini``). + ``get_llm_provider`` doesn't know router aliases, so without this + step every alias call ended up labeled ``"litellm_proxy"``, + defeating the field's purpose for the most common case. + 3. Defensive fallback to ``("", "litellm_proxy")`` — used only when + ``model`` is missing, malformed, or both lookups fail. We never let + a secondary exception escape and mask the rate-limit error we're + trying to surface. """ if not model: return "", PROXY_LLM_PROVIDER_FALLBACK @@ -46,6 +53,9 @@ def resolve_llm_provider_for_rate_limit( custom_llm_provider or PROXY_LLM_PROVIDER_FALLBACK, ) except Exception as e: + alias_resolution = _resolve_provider_from_router_alias(model) + if alias_resolution is not None: + return alias_resolution verbose_proxy_logger.debug( "rate_limiter_utils.resolve_llm_provider_for_rate_limit: " "could not resolve provider for model=%s, falling back to %s. err=%s", @@ -56,50 +66,58 @@ def resolve_llm_provider_for_rate_limit( return model, PROXY_LLM_PROVIDER_FALLBACK -class ProxyHTTPRateLimitError(HTTPException, RateLimitError): # type: ignore[misc] +def _resolve_provider_from_router_alias( + model: str, +) -> Optional[Tuple[str, str]]: """ - HTTPException raised by proxy-side rate-limit hooks that *also* exposes - ``model`` and ``llm_provider`` attributes. + Resolve a router ``model_name`` alias to ``(underlying_model, provider)`` + by scanning the active router's ``model_list``. - Why both base classes: - - - The proxy server's exception handler keys off ``HTTPException`` to render - a 429 response, so we must remain an ``HTTPException``. - - Downstream loggers (Prometheus ``async_post_call_failure_hook``, - structured logging, observability callbacks) read ``exception.llm_provider`` - via :meth:`litellm.integrations.prometheus.PrometheusLogger._get_exception_class_name` - and ``isinstance(exc, RateLimitError)`` for category routing. Inheriting - from :class:`litellm.exceptions.RateLimitError` keeps that wiring intact. - - We intentionally do not call ``RateLimitError.__init__`` (which constructs - an httpx.Response) — it isn't needed here and just adds failure surface. - Attribute parity is what downstream consumers rely on. + Returns ``None`` if the router isn't initialized, the alias isn't + registered, the deployment has no usable ``litellm_params.model``, or + any underlying lookup raises. Callers fall through to the defensive + ``litellm_proxy`` fallback in that case — never raising secondary + exceptions out of the rate-limit raise path. """ - - def __init__( - self, - status_code: int, - detail: Any = None, - headers: Optional[dict] = None, - *, - model: str = "", - llm_provider: str = PROXY_LLM_PROVIDER_FALLBACK, - ) -> None: - HTTPException.__init__( - self, status_code=status_code, detail=detail, headers=headers - ) - self.status_code = status_code - self.model = model or "" - self.llm_provider = llm_provider or PROXY_LLM_PROVIDER_FALLBACK - # `message` is what RateLimitError.__str__ would print and what some - # observability callbacks log. Keep it human-readable. - self.message = detail if isinstance(detail, str) else str(detail) - # `RateLimitError.__str__` (resolved via MRO since Starlette's - # HTTPException doesn't define `__str__`) unconditionally reads - # these attributes. Set them so `str(exc)` doesn't raise - # AttributeError from logging/traceback paths. - self.num_retries: Optional[int] = None - self.max_retries: Optional[int] = None + try: + from litellm.proxy.proxy_server import llm_router + except Exception: + return None + if llm_router is None: + return None + try: + model_list = getattr(llm_router, "model_list", None) + if not model_list: + return None + for deployment in model_list: + if not isinstance(deployment, dict): + continue + if deployment.get("model_name") != model: + continue + params = deployment.get("litellm_params") + if not isinstance(params, dict): + continue + underlying_model = params.get("model") + if not isinstance(underlying_model, str) or not underlying_model: + continue + try: + resolved_model, custom_llm_provider, _, _ = litellm.get_llm_provider( + model=underlying_model, + ) + except Exception: + continue + if not custom_llm_provider: + continue + # Prefer the underlying provider-qualified model so the failure + # callback / Prometheus label points at the actual deployment, not + # the alias. + return ( + resolved_model or underlying_model, + custom_llm_provider, + ) + return None + except Exception: + return None def convert_priority_to_percent( diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index dc27e87726a..31d831d773c 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -1,4 +1,4 @@ -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, Optional, Union from fastapi import HTTPException, status from pydantic import BaseModel @@ -19,6 +19,7 @@ from litellm.proxy._types import ( UserAPIKeyAuth, user_api_key_has_admin_view as _user_has_admin_view, # noqa: F401 re-exported ) +from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.utils import _premium_user_check if TYPE_CHECKING: @@ -400,121 +401,127 @@ def _set_object_metadata_field( object_data.metadata[field_name] = value +_TEAM_MEMBER_BUDGET_LIMIT_FIELDS = ( + "max_budget", + "soft_budget", + "max_parallel_requests", + "tpm_limit", + "rpm_limit", + "model_max_budget", + "budget_duration", + "allowed_models", +) + + +def _is_set_budget_value(value: Any) -> bool: + if value is None: + return False + if isinstance(value, list) and len(value) == 0: + return False + return True + + +def _has_meaningful_budget_limit(budget_values: Dict[str, Any]) -> bool: + """A budget is meaningful if at least one limit is actually set; an empty + list (no model restriction) and None both count as unset.""" + return any( + _is_set_budget_value(budget_values.get(field)) + for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS + ) + + async def _upsert_budget_and_membership( tx, *, team_id: str, user_id: str, - max_budget: Optional[float], existing_budget_id: Optional[str], user_api_key_dict: UserAPIKeyAuth, - tpm_limit: Optional[int] = None, - rpm_limit: Optional[int] = None, - allowed_models: Optional[List[str]] = None, + budget_patch: Dict[str, Any], team_default_budget_id: Optional[str] = None, ): """ - Helper function to Create/Update or Delete the budget within the team membership - Args: - tx: The transaction object - team_id: The ID of the team - user_id: The ID of the user - max_budget: The maximum budget for the team - existing_budget_id: The ID of the existing budget, if any - user_api_key_dict: User API Key dictionary containing user information - tpm_limit: Tokens per minute limit for the team member - rpm_limit: Requests per minute limit for the team member - allowed_models: Per-member model scope. None = don't change. [] = remove restrictions. Non-empty list = enforce. - team_default_budget_id: The team's shared default member budget id (from - team metadata.team_member_budget_id), if any. When the membership's - existing_budget_id matches this, we clone-on-write so editing one - member's budget does not mutate the shared default (and therefore - every other member who still points at it). + Apply a merge-patch of per-member budget fields to a team membership. - If max_budget, tpm_limit, rpm_limit, and allowed_models are all None, the user's budget is removed from the team membership. - If any of these values exist, a budget is updated or created and linked to the team membership. + ``budget_patch`` holds only the budget columns the caller explicitly sent + (RFC 7396 semantics): a value sets the column, ``None`` clears it, and a + column that is absent from the dict is left untouched. Once the patch is + applied, if the budget has no meaningful limit left the member's private + budget is disconnected so they fall back to the team default. + + ``team_default_budget_id`` is the team's shared default member budget id + (from team metadata.team_member_budget_id). When the membership still + points at it, we clone-on-write so editing one member's budget does not + mutate the shared default that every other member points at. """ - if ( - max_budget is None - and tpm_limit is None - and rpm_limit is None - and allowed_models is None - ): - # disconnect the budget since all limits are None - await tx.litellm_teammembership.update( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, - data={"litellm_budget_table": {"disconnect": True}}, - ) + if not budget_patch: return + write_data = dict(budget_patch) + if "budget_duration" in write_data: + duration = write_data["budget_duration"] + write_data["budget_reset_at"] = ( + get_budget_reset_time(budget_duration=duration) + if duration is not None + else None + ) + is_shared_default = ( existing_budget_id is not None and team_default_budget_id is not None and existing_budget_id == team_default_budget_id ) + async def _disconnect(): + await tx.litellm_teammembership.update( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}}, + data={"litellm_budget_table": {"disconnect": True}}, + ) + if existing_budget_id is not None and not is_shared_default: - # Update the existing budget in-place to preserve fields not being changed. - # Only write fields that the caller explicitly provided (non-None). - update_data: Dict[str, Any] = { - "updated_by": user_api_key_dict.user_id or "", - } - if max_budget is not None: - update_data["max_budget"] = max_budget - if tpm_limit is not None: - update_data["tpm_limit"] = tpm_limit - if rpm_limit is not None: - update_data["rpm_limit"] = rpm_limit - if allowed_models is not None: - update_data["allowed_models"] = allowed_models + existing_budget = await tx.litellm_budgettable.find_unique( + where={"budget_id": existing_budget_id} + ) + merged = existing_budget.model_dump() if existing_budget is not None else {} + merged.update(write_data) + if not _has_meaningful_budget_limit(merged): + await _disconnect() + return await tx.litellm_budgettable.update( where={"budget_id": existing_budget_id}, - data=update_data, + data={"updated_by": user_api_key_dict.user_id or "", **write_data}, ) return - # Either there is no existing budget, OR the membership is still pointing - # at the team's shared default member budget. In both cases we create a - # NEW private budget for this user and (re)link the membership to it. create_data: Dict[str, Any] = { "created_by": user_api_key_dict.user_id or "", "updated_by": user_api_key_dict.user_id or "", } - # If we're forking off the shared default, seed the new row with the - # default's values so fields the caller did not change carry over. if is_shared_default: default_budget_row = await tx.litellm_budgettable.find_unique( where={"budget_id": existing_budget_id} ) if default_budget_row is not None: default_budget_dict = default_budget_row.model_dump() - for field in ( - "max_budget", - "soft_budget", - "max_parallel_requests", - "tpm_limit", - "rpm_limit", - "model_max_budget", - "budget_duration", - "allowed_models", - ): + for field in _TEAM_MEMBER_BUDGET_LIMIT_FIELDS: value = default_budget_dict.get(field) - if value is None: - continue - if isinstance(value, list) and len(value) == 0: - continue - create_data[field] = value + if _is_set_budget_value(value): + create_data[field] = value - # Caller-provided values take precedence over the cloned defaults. - if max_budget is not None: - create_data["max_budget"] = max_budget - if tpm_limit is not None: - create_data["tpm_limit"] = tpm_limit - if rpm_limit is not None: - create_data["rpm_limit"] = rpm_limit - if allowed_models is not None: - create_data["allowed_models"] = allowed_models + create_data.update(write_data) + + if create_data.get("budget_duration") is not None: + create_data["budget_reset_at"] = get_budget_reset_time( + budget_duration=create_data["budget_duration"] + ) + else: + create_data.pop("budget_reset_at", None) + + if not _has_meaningful_budget_limit(create_data): + if existing_budget_id is not None: + await _disconnect() + return new_budget = await tx.litellm_budgettable.create( data=create_data, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index ae7da0d29f2..a3ad7a9ea8b 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -2733,6 +2733,52 @@ async def team_member_delete( return existing_team_row +_MEMBER_BUDGET_PATCH_FIELDS = { + "max_budget_in_team": "max_budget", + "tpm_limit": "tpm_limit", + "rpm_limit": "rpm_limit", + "budget_duration": "budget_duration", + "allowed_models": "allowed_models", +} + + +def _build_member_budget_patch(data: TeamMemberUpdateRequest) -> Dict[str, Any]: + """Map the budget fields the request actually set (merge-patch: a sent + value updates, an explicit null clears, an absent field is left untouched) + to their budget-table columns.""" + provided = data.model_dump(exclude_unset=True) + return { + column: provided[request_field] + for request_field, column in _MEMBER_BUDGET_PATCH_FIELDS.items() + if request_field in provided + } + + +def _validate_budget_duration(budget_duration: Optional[str]) -> None: + """Reject budget durations that can't be parsed, are non-positive, or + overflow date math, so a bad value can't be persisted and later crash the + budget reset job.""" + if budget_duration is None: + return + + from litellm.litellm_core_utils.duration_parser import duration_in_seconds + from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time + + try: + if duration_in_seconds(budget_duration) <= 0: + raise ValueError("budget_duration must be positive") + get_budget_reset_time(budget_duration=budget_duration) + except (ValueError, OverflowError): + raise HTTPException( + status_code=400, + detail={ + "error": "Invalid budget_duration '{}'. Use a format like '1h', '24h', '7d', or '30d'.".format( + budget_duration + ) + }, + ) + + @router.post( "/team/member_update", tags=["team management"], @@ -2770,6 +2816,8 @@ async def team_member_update( detail={"error": "Either user_id or user_email needs to be passed in"}, ) + _validate_budget_duration(data.budget_duration) + _existing_team_row = await prisma_client.db.litellm_teamtable.find_unique( where={"team_id": data.team_id} ) @@ -2843,17 +2891,15 @@ async def team_member_update( team_default_budget_id = raw_default_budget_id ### upsert new budget + budget_patch = _build_member_budget_patch(data) async with prisma_client.db.tx() as tx: await _upsert_budget_and_membership( tx=tx, team_id=data.team_id, user_id=received_user_id, - max_budget=data.max_budget_in_team, existing_budget_id=identified_budget_id, user_api_key_dict=user_api_key_dict, - tpm_limit=data.tpm_limit, - rpm_limit=data.rpm_limit, - allowed_models=data.allowed_models, + budget_patch=budget_patch, team_default_budget_id=team_default_budget_id, ) @@ -2887,6 +2933,7 @@ async def team_member_update( max_budget_in_team=data.max_budget_in_team, tpm_limit=data.tpm_limit, rpm_limit=data.rpm_limit, + budget_duration=data.budget_duration, allowed_models=data.allowed_models, ) diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 55f4fc96504..5b1d32cd93c 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -115,6 +115,8 @@ class ValidationResults: REQUESTED_MODEL = "requested_model" EXCEPTION_STATUS = "exception_status" EXCEPTION_CLASS = "exception_class" +RATE_LIMIT_CATEGORY = "rate_limit_category" +RATE_LIMIT_TYPE = "rate_limit_type" STATUS_CODE = "status_code" EXCEPTION_LABELS = [EXCEPTION_STATUS, EXCEPTION_CLASS] LATENCY_BUCKETS = ( @@ -174,6 +176,8 @@ class UserAPIKeyLabelNames(Enum): API_PROVIDER = "api_provider" EXCEPTION_STATUS = EXCEPTION_STATUS EXCEPTION_CLASS = EXCEPTION_CLASS + RATE_LIMIT_CATEGORY = RATE_LIMIT_CATEGORY + RATE_LIMIT_TYPE = RATE_LIMIT_TYPE STATUS_CODE = "status_code" FALLBACK_MODEL = "fallback_model" ROUTE = "route" @@ -343,6 +347,10 @@ class PrometheusMetricLabels: UserAPIKeyLabelNames.USER_EMAIL.value, UserAPIKeyLabelNames.EXCEPTION_STATUS.value, UserAPIKeyLabelNames.EXCEPTION_CLASS.value, + # ``rate_limit_category`` / ``rate_limit_type`` are appended in + # ``get_labels()`` when ``litellm.prometheus_emit_rate_limit_labels`` + # is True. Kept opt-in so existing dashboards keyed on this metric's + # historical label set keep matching after upgrade. UserAPIKeyLabelNames.ROUTE.value, UserAPIKeyLabelNames.CLIENT_IP.value, UserAPIKeyLabelNames.USER_AGENT.value, @@ -745,6 +753,25 @@ class PrometheusMetricLabels: ): custom_labels.append(UserAPIKeyLabelNames.STREAM.value) + # Conditionally add unified rate-limit labels to + # litellm_proxy_failed_requests_metric. Off by default so the metric's + # historical label set is preserved across upgrade; enable via + # ``litellm.prometheus_emit_rate_limit_labels`` once downstream + # dashboards include the new labels in their matchers / aggregations. + if ( + label_name == "litellm_proxy_failed_requests_metric" + and litellm.prometheus_emit_rate_limit_labels is True + ): + for _rate_limit_label in ( + UserAPIKeyLabelNames.RATE_LIMIT_CATEGORY.value, + UserAPIKeyLabelNames.RATE_LIMIT_TYPE.value, + ): + if ( + _rate_limit_label not in default_labels + and _rate_limit_label not in custom_labels + ): + custom_labels.append(_rate_limit_label) + _user_budget_metrics = { "litellm_remaining_user_budget_metric", "litellm_user_max_budget_metric", @@ -807,6 +834,8 @@ class UserAPIKeyLabelValues: api_provider: Optional[str] = None exception_status: Optional[str] = None exception_class: Optional[str] = None + rate_limit_category: Optional[str] = None + rate_limit_type: Optional[str] = None status_code: Optional[str] = None fallback_model: Optional[str] = None route: Optional[str] = None diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a7a0b0f6238..b76ae1f5d86 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2720,6 +2720,23 @@ class StandardLoggingPayloadErrorInformation(TypedDict, total=False): llm_provider: Optional[str] traceback: Optional[str] error_message: Optional[str] + # error_rate_limit_category: + # For 429 / rate-limit errors, the source of the rate limit. One of the + # string values defined by `litellm.exceptions.RateLimitErrorCategory` + # (vendor_rate_limit, vendor_batch_rate_limit, litellm_rate_limit, + # litellm_batch_rate_limit). None for non-rate-limit exceptions. + # Surfaced here so custom callbacks / metrics consumers can switch on + # the rate-limit source without reaching for the raw exception. + error_rate_limit_category: Optional[str] + # error_rate_limit_type: + # For 429 / rate-limit errors, the dimension that was exceeded. One of + # the string values defined by `litellm.exceptions.RateLimitType` + # (requests, tokens, concurrent_requests, budget, max_iterations). + # None for non-rate-limit exceptions and for rate-limit exceptions that + # did not classify the failure (e.g. legacy vendor 429 with no header + # hints). Lets dashboards split rate-limit failures by cause without + # parsing free-text error messages. + error_rate_limit_type: Optional[str] class GuardrailMode(TypedDict, total=False): diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index f8bac820582..d0ad1cc8f82 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -783,6 +783,16 @@ async def test_async_post_call_failure_hook(prometheus_logger): it should increment the litellm_proxy_failed_requests_metric and litellm_proxy_total_requests_metric """ + # Opt into the unified rate-limit labels so this test exercises the + # full label set surfaced when `prometheus_emit_rate_limit_labels` is on. + # The logger caches each metric's label set at construction time (so the + # labels passed to ``counter.labels(...)`` stay in lock step with the + # labels used to register the metric), so we must invalidate the cache + # after flipping the toggle for the cache to pick up the new label set. + original_emit = litellm.prometheus_emit_rate_limit_labels + litellm.prometheus_emit_rate_limit_labels = True + prometheus_logger._cached_metric_labels.clear() + # Mock the prometheus metrics prometheus_logger.litellm_proxy_failed_requests_metric = MagicMock() prometheus_logger.litellm_proxy_total_requests_metric = MagicMock() @@ -804,32 +814,38 @@ async def test_async_post_call_failure_hook(prometheus_logger): request_route="/chat/completions", ) - # Call the function - await prometheus_logger.async_post_call_failure_hook( - request_data=request_data, - original_exception=original_exception, - user_api_key_dict=user_api_key_dict, - ) + try: + # Call the function + await prometheus_logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=original_exception, + user_api_key_dict=user_api_key_dict, + ) - # Assert failed requests metric was incremented with correct labels - prometheus_logger.litellm_proxy_failed_requests_metric.labels.assert_called_once_with( - end_user=None, - user="test_user", - user_email=None, - hashed_api_key="test_key", - api_key_alias="test_alias", - team="test_team", - team_alias="test_team_alias", - org_id=None, - org_alias=None, - requested_model="gpt-5-mini", - exception_status="429", - exception_class="Openai.RateLimitError", - route=user_api_key_dict.request_route, - model_id=None, - client_ip=None, - user_agent=None, - ) + # Assert failed requests metric was incremented with correct labels + prometheus_logger.litellm_proxy_failed_requests_metric.labels.assert_called_once_with( + end_user=None, + user="test_user", + user_email=None, + hashed_api_key="test_key", + api_key_alias="test_alias", + team="test_team", + team_alias="test_team_alias", + org_id=None, + org_alias=None, + requested_model="gpt-5-mini", + exception_status="429", + exception_class="Openai.RateLimitError", + rate_limit_category="vendor_rate_limit", + rate_limit_type=None, + route=user_api_key_dict.request_route, + model_id=None, + client_ip=None, + user_agent=None, + ) + finally: + litellm.prometheus_emit_rate_limit_labels = original_emit + prometheus_logger._cached_metric_labels.clear() prometheus_logger.litellm_proxy_failed_requests_metric.labels().inc.assert_called_once() # Assert total requests metric was incremented with correct labels @@ -1962,6 +1978,10 @@ def test_set_team_budget_metrics_with_custom_labels(prometheus_logger, monkeypat # Set custom prometheus labels custom_labels = ["metadata.organization", "metadata.environment"] monkeypatch.setattr("litellm.custom_prometheus_metadata_labels", custom_labels) + # Logger caches each metric's label set at construction time (fixture + # runs before this monkeypatch), so invalidate so the cached label set + # picks up the freshly-configured custom metadata labels. + prometheus_logger._cached_metric_labels.clear() # Create test team with custom metadata team = MagicMock( diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/test_litellm/integrations/test_prometheus_labels.py index 1ba332a341b..c83d89e87c4 100644 --- a/tests/test_litellm/integrations/test_prometheus_labels.py +++ b/tests/test_litellm/integrations/test_prometheus_labels.py @@ -284,6 +284,12 @@ def test_prometheus_metrics_use_normalized_routes(): # Create a mock PrometheusLogger prometheus_logger = MagicMock() + # ``get_labels_for_metric`` reads ``_cached_metric_labels`` and + # ``label_filters`` off ``self``; default MagicMock attribute access + # returns Mocks that masquerade as a populated cache, so seed real + # containers before binding the real method. + prometheus_logger._cached_metric_labels = {} + prometheus_logger.label_filters = {} prometheus_logger.get_labels_for_metric = ( PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger) ) @@ -327,6 +333,8 @@ def test_prometheus_label_value_sanitization(): from unittest.mock import MagicMock prometheus_logger = MagicMock() + prometheus_logger._cached_metric_labels = {} + prometheus_logger.label_filters = {} prometheus_logger.get_labels_for_metric = ( PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger) ) diff --git a/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py b/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py new file mode 100644 index 00000000000..bb035c4c3ee --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_rate_limit_labels.py @@ -0,0 +1,328 @@ +""" +Tests for the Prometheus rate-limit labels added on top of PR #27687. + +Covers two follow-up gaps to the unified rate-limit error work: + +1. ``litellm_proxy_failed_requests_metric`` now carries + ``rate_limit_category`` and ``rate_limit_type`` labels populated from + :class:`litellm.RateLimitError` (vendor + ``ProxyRateLimitError`` + subclass). Closes the Prometheus side of LIT-2718. +2. ``_get_exception_class_name`` keeps emitting the literal string + ``"HTTPException"`` for ``ProxyRateLimitError`` so existing dashboards + that key off ``exception_class="HTTPException"`` for litellm-internal + 429s don't silently break when the new class lands. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.exceptions import ( + RateLimitError, + RateLimitErrorCategory, + RateLimitType, +) +from litellm.integrations.prometheus import PrometheusLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError +from litellm.types.integrations.prometheus import ( + PrometheusMetricLabels, + UserAPIKeyLabelNames, + UserAPIKeyLabelValues, +) + + +# --------------------------------------------------------------------------- +# Label / enum wiring +# --------------------------------------------------------------------------- + + +def test_should_register_rate_limit_label_names_on_enum(): + assert UserAPIKeyLabelNames.RATE_LIMIT_CATEGORY.value == "rate_limit_category" + assert UserAPIKeyLabelNames.RATE_LIMIT_TYPE.value == "rate_limit_type" + + +def test_should_include_rate_limit_labels_on_failed_requests_metric(): + import litellm + + original = litellm.prometheus_emit_rate_limit_labels + try: + litellm.prometheus_emit_rate_limit_labels = True + labels = PrometheusMetricLabels.get_labels( + "litellm_proxy_failed_requests_metric" + ) + assert "rate_limit_category" in labels + assert "rate_limit_type" in labels + # These must coexist with the legacy exception labels (back-compat). + assert "exception_class" in labels + assert "exception_status" in labels + finally: + litellm.prometheus_emit_rate_limit_labels = original + + +def test_should_omit_rate_limit_labels_by_default_for_back_compat(): + """Default-off preserves the metric's historical label set so existing + dashboards / recording rules keyed on `litellm_proxy_failed_requests_metric` + keep matching after upgrade.""" + import litellm + + assert litellm.prometheus_emit_rate_limit_labels is False + labels = PrometheusMetricLabels.get_labels("litellm_proxy_failed_requests_metric") + assert "rate_limit_category" not in labels + assert "rate_limit_type" not in labels + # Pre-PR labels must still be present. + assert "exception_class" in labels + assert "exception_status" in labels + + +def test_should_accept_rate_limit_fields_on_user_api_key_label_values(): + enum_values = UserAPIKeyLabelValues( + rate_limit_category="litellm_rate_limit", + rate_limit_type="requests", + ) + assert enum_values.rate_limit_category == "litellm_rate_limit" + assert enum_values.rate_limit_type == "requests" + + +# --------------------------------------------------------------------------- +# _extract_rate_limit_labels helper +# --------------------------------------------------------------------------- + + +def test_should_extract_vendor_category_for_vanilla_rate_limit_error(): + err = RateLimitError(message="vendor 429", llm_provider="openai", model="gpt-4o") + category, rate_limit_type = PrometheusLogger._extract_rate_limit_labels(err) + assert category == "vendor_rate_limit" + assert rate_limit_type is None + + +def test_should_extract_litellm_category_and_type_for_proxy_rate_limit_error(): + err = ProxyRateLimitError( + detail={"error": "tpm exceeded"}, + category=RateLimitErrorCategory.LITELLM_RATE_LIMIT, + rate_limit_type=RateLimitType.TOKENS, + ) + category, rate_limit_type = PrometheusLogger._extract_rate_limit_labels(err) + assert category == "litellm_rate_limit" + assert rate_limit_type == "tokens" + + +def test_should_return_none_for_non_rate_limit_exception(): + assert PrometheusLogger._extract_rate_limit_labels(ValueError("nope")) == ( + None, + None, + ) + + +def test_should_return_none_for_none_exception(): + assert PrometheusLogger._extract_rate_limit_labels(None) == (None, None) + + +def test_should_extract_budget_dimension_for_budget_exceeded_error(): + # Virtual-key / team / org / end-user budget caps raise + # `litellm.BudgetExceededError` (a bare Exception subclass), which sets + # the same `.category` / `.rate_limit_type` attributes as the unified + # RateLimitError path so Prometheus can split budget 429s from other + # 429s without the customer parsing free-text error messages. + import litellm + + err = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + category, rate_limit_type = PrometheusLogger._extract_rate_limit_labels(err) + assert category == "litellm_rate_limit" + assert rate_limit_type == "budget" + + +@pytest.mark.parametrize( + "category_enum,rate_limit_enum,expected_category,expected_type", + [ + ( + RateLimitErrorCategory.LITELLM_RATE_LIMIT, + RateLimitType.REQUESTS, + "litellm_rate_limit", + "requests", + ), + ( + RateLimitErrorCategory.LITELLM_RATE_LIMIT, + RateLimitType.TOKENS, + "litellm_rate_limit", + "tokens", + ), + ( + RateLimitErrorCategory.LITELLM_RATE_LIMIT, + RateLimitType.CONCURRENT_REQUESTS, + "litellm_rate_limit", + "concurrent_requests", + ), + ( + RateLimitErrorCategory.LITELLM_RATE_LIMIT, + RateLimitType.BUDGET, + "litellm_rate_limit", + "budget", + ), + ( + RateLimitErrorCategory.LITELLM_RATE_LIMIT, + RateLimitType.MAX_ITERATIONS, + "litellm_rate_limit", + "max_iterations", + ), + ( + RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT, + RateLimitType.REQUESTS, + "litellm_batch_rate_limit", + "requests", + ), + ], +) +def test_should_serialize_rate_limit_enums_as_underlying_string_values( + category_enum, rate_limit_enum, expected_category, expected_type +): + err = ProxyRateLimitError( + detail="boom", category=category_enum, rate_limit_type=rate_limit_enum + ) + category, rate_limit_type = PrometheusLogger._extract_rate_limit_labels(err) + assert category == expected_category + assert rate_limit_type == expected_type + + +# --------------------------------------------------------------------------- +# _get_exception_class_name back-compat +# --------------------------------------------------------------------------- + + +def test_should_emit_legacy_http_exception_label_for_proxy_rate_limit_error(): + """ + ``ProxyRateLimitError`` multi-inherits from ``HTTPException`` + + ``RateLimitError``. The ``exception_class`` label MUST keep emitting + "HTTPException" for back-compat with existing dashboards (see Slack + thread + PR #27687 review). Distinguishing vendor vs. litellm 429s + is now the job of the new ``rate_limit_category`` label. + """ + err = ProxyRateLimitError(detail={"error": "boom"}) + assert PrometheusLogger._get_exception_class_name(err) == "HTTPException" + + +def test_should_keep_provider_prefixed_exception_class_for_vendor_rate_limit_errors(): + err = RateLimitError(message="vendor 429", llm_provider="openai", model="gpt-4o") + # Vendor-side errors keep the historical "Provider.ClassName" formatting. + assert PrometheusLogger._get_exception_class_name(err) == "Openai.RateLimitError" + + +def test_should_preserve_exception_class_name_for_unrelated_exceptions(): + assert PrometheusLogger._get_exception_class_name(ValueError("nope")) == ( + "ValueError" + ) + + +# --------------------------------------------------------------------------- +# End-to-end wiring through async_post_call_failure_hook +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_should_populate_rate_limit_labels_for_proxy_rate_limit_error_on_failure_hook(): + """ + When a proxy hook raises ``ProxyRateLimitError`` and the failure flows + through ``async_post_call_failure_hook``, the resulting + ``UserAPIKeyLabelValues`` must carry both new labels AND keep + ``exception_class="HTTPException"`` for back-compat. + """ + with patch( + "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None + ): + logger = PrometheusLogger() + logger.litellm_proxy_failed_requests_metric = MagicMock() + logger.litellm_proxy_total_requests_metric = MagicMock() + logger.get_labels_for_metric = MagicMock( + return_value=PrometheusMetricLabels.get_labels( + "litellm_proxy_failed_requests_metric" + ) + ) + + err = ProxyRateLimitError( + detail={"error": "rpm exceeded"}, + category=RateLimitErrorCategory.LITELLM_RATE_LIMIT, + rate_limit_type=RateLimitType.REQUESTS, + ) + + with patch( + "litellm.integrations.prometheus.prometheus_label_factory" + ) as mock_label_factory: + mock_label_factory.return_value = {} + await logger.async_post_call_failure_hook( + request_data={"model": "gpt-4o-mini", "metadata": {}}, + original_exception=err, + user_api_key_dict=UserAPIKeyAuth(token="t"), + ) + + enum_values = mock_label_factory.call_args_list[0].kwargs["enum_values"] + assert isinstance(enum_values, UserAPIKeyLabelValues) + assert enum_values.rate_limit_category == "litellm_rate_limit" + assert enum_values.rate_limit_type == "requests" + # Back-compat: exception_class on a ProxyRateLimitError stays "HTTPException". + assert enum_values.exception_class == "HTTPException" + assert enum_values.exception_status == "429" + + +@pytest.mark.asyncio +async def test_should_populate_rate_limit_labels_for_vendor_rate_limit_error_on_failure_hook(): + with patch( + "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None + ): + logger = PrometheusLogger() + logger.litellm_proxy_failed_requests_metric = MagicMock() + logger.litellm_proxy_total_requests_metric = MagicMock() + logger.get_labels_for_metric = MagicMock( + return_value=PrometheusMetricLabels.get_labels( + "litellm_proxy_failed_requests_metric" + ) + ) + + err = RateLimitError(message="upstream 429", llm_provider="openai", model="gpt-4o") + + with patch( + "litellm.integrations.prometheus.prometheus_label_factory" + ) as mock_label_factory: + mock_label_factory.return_value = {} + await logger.async_post_call_failure_hook( + request_data={"model": "gpt-4o", "metadata": {}}, + original_exception=err, + user_api_key_dict=UserAPIKeyAuth(token="t"), + ) + + enum_values = mock_label_factory.call_args_list[0].kwargs["enum_values"] + assert isinstance(enum_values, UserAPIKeyLabelValues) + assert enum_values.rate_limit_category == "vendor_rate_limit" + assert enum_values.rate_limit_type is None + # Vendor errors keep the historical Provider.ClassName label. + assert enum_values.exception_class == "Openai.RateLimitError" + assert enum_values.exception_status == "429" + + +@pytest.mark.asyncio +async def test_should_leave_rate_limit_labels_blank_for_non_rate_limit_failure(): + with patch( + "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None + ): + logger = PrometheusLogger() + logger.litellm_proxy_failed_requests_metric = MagicMock() + logger.litellm_proxy_total_requests_metric = MagicMock() + logger.get_labels_for_metric = MagicMock( + return_value=PrometheusMetricLabels.get_labels( + "litellm_proxy_failed_requests_metric" + ) + ) + + with patch( + "litellm.integrations.prometheus.prometheus_label_factory" + ) as mock_label_factory: + mock_label_factory.return_value = {} + await logger.async_post_call_failure_hook( + request_data={"model": "gpt-4o", "metadata": {}}, + original_exception=RuntimeError("boom"), + user_api_key_dict=UserAPIKeyAuth(token="t"), + ) + + enum_values = mock_label_factory.call_args_list[0].kwargs["enum_values"] + assert isinstance(enum_values, UserAPIKeyLabelValues) + assert enum_values.rate_limit_category is None + assert enum_values.rate_limit_type is None diff --git a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py b/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py index 12f30ab6024..361ab7332f8 100644 --- a/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py +++ b/tests/test_litellm/integrations/test_prometheus_user_team_metrics.py @@ -511,29 +511,34 @@ def test_set_user_budget_metrics_default_no_email_alias_labels( ) -def test_set_user_budget_metrics_includes_user_email_and_alias_labels_when_opted_in( - prometheus_logger, -): - """When prometheus_user_budget_label_include_email_alias=True, email+alias labels appear.""" +def test_set_user_budget_metrics_includes_user_email_and_alias_labels_when_opted_in(): + """When prometheus_user_budget_label_include_email_alias=True, email+alias labels appear. + + The flag is read once per metric at logger construction time and snapshotted, + so it must be enabled before the PrometheusLogger is built (mirroring how the + proxy applies config at startup before instantiating callbacks). + """ import litellm from litellm.proxy._types import LiteLLM_UserTable litellm.prometheus_user_budget_label_include_email_alias = True - user = LiteLLM_UserTable( - user_id="user-abc-123", - user_email="alice@example.com", - user_alias="Alice", - spend=25.0, - max_budget=100.0, - budget_reset_at=datetime(2026, 3, 1, tzinfo=timezone.utc), - ) - - prometheus_logger.litellm_remaining_user_budget_metric = MagicMock() - prometheus_logger.litellm_user_max_budget_metric = MagicMock() - prometheus_logger.litellm_user_budget_remaining_hours_metric = MagicMock() - try: + prometheus_logger = PrometheusLogger() + + user = LiteLLM_UserTable( + user_id="user-abc-123", + user_email="alice@example.com", + user_alias="Alice", + spend=25.0, + max_budget=100.0, + budget_reset_at=datetime(2026, 3, 1, tzinfo=timezone.utc), + ) + + prometheus_logger.litellm_remaining_user_budget_metric = MagicMock() + prometheus_logger.litellm_user_max_budget_metric = MagicMock() + prometheus_logger.litellm_user_budget_remaining_hours_metric = MagicMock() + prometheus_logger._set_user_budget_metrics(user) prometheus_logger.litellm_remaining_user_budget_metric.labels.assert_called_once_with( diff --git a/tests/test_litellm/proxy/common_utils/test_upsert_budget_membership.py b/tests/test_litellm/proxy/common_utils/test_upsert_budget_membership.py index f4bf0d7b2be..e9b4f11e891 100644 --- a/tests/test_litellm/proxy/common_utils/test_upsert_budget_membership.py +++ b/tests/test_litellm/proxy/common_utils/test_upsert_budget_membership.py @@ -1,5 +1,6 @@ # tests/litellm/proxy/common_utils/test_upsert_budget_membership.py import types +from datetime import datetime, timezone from unittest.mock import AsyncMock, MagicMock import pytest @@ -19,15 +20,13 @@ def mock_tx(): Builds an object that looks just enough like the Prisma tx you use inside _upsert_budget_and_membership. """ - # membership “table” membership = MagicMock() membership.update = AsyncMock() membership.upsert = AsyncMock() - # budget “table” budget = MagicMock() budget.update = AsyncMock() - # budget.create returns a fake row that has .budget_id + budget.find_unique = AsyncMock(return_value=None) budget.create = AsyncMock( return_value=types.SimpleNamespace(budget_id="new-budget-123") ) @@ -44,16 +43,57 @@ def fake_user(): return types.SimpleNamespace(user_id="tester@example.com") -# TEST: max_budget is None, disconnect only +def budget_row(**fields): + """A fake litellm_budgettable row whose model_dump returns the given fields.""" + row = MagicMock() + row.model_dump.return_value = fields + return row + + +def assert_future_reset_time(value): + """A budget_reset_at must be a timezone-aware datetime in the future, so the + member's budget actually rolls over and the UI shows a reset date instead of + waiting for the reset cron to backfill it.""" + assert isinstance(value, datetime) + assert value.tzinfo is not None + assert value > datetime.now(timezone.utc) + + +# TEST: an empty patch (caller sent no budget fields) leaves everything alone. +# This is the merge-patch contract: absent != clear. Updating only a member's +# role must not silently wipe their budget. @pytest.mark.asyncio -async def test_upsert_disconnect(mock_tx, fake_user): +async def test_empty_patch_is_noop(mock_tx, fake_user): await _upsert_budget_and_membership( mock_tx, team_id="team-1", user_id="user-1", - max_budget=None, - existing_budget_id=None, + existing_budget_id="bud-1", user_api_key_dict=fake_user, + budget_patch={}, + ) + + mock_tx.litellm_teammembership.update.assert_not_called() + mock_tx.litellm_teammembership.upsert.assert_not_called() + mock_tx.litellm_budgettable.update.assert_not_called() + mock_tx.litellm_budgettable.create.assert_not_called() + + +# TEST: clearing every limit on a member's private budget disconnects it, so the +# member falls back to the team default instead of keeping an empty private row. +@pytest.mark.asyncio +async def test_clearing_all_limits_disconnects(mock_tx, fake_user): + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row(max_budget=100.0) + ) + + await _upsert_budget_and_membership( + mock_tx, + team_id="team-1", + user_id="user-1", + existing_budget_id="bud-1", + user_api_key_dict=fake_user, + budget_patch={"max_budget": None}, ) mock_tx.litellm_teammembership.update.assert_awaited_once_with( @@ -62,205 +102,114 @@ async def test_upsert_disconnect(mock_tx, fake_user): ) mock_tx.litellm_budgettable.update.assert_not_called() mock_tx.litellm_budgettable.create.assert_not_called() - mock_tx.litellm_teammembership.upsert.assert_not_called() -# TEST: existing budget id → updates budget in-place (current behavior) +# TEST: clearing one field on a budget that still has another limit updates in +# place (clears just that column + its reset time) and does NOT disconnect. @pytest.mark.asyncio -async def test_upsert_with_existing_budget_id_creates_new(mock_tx, fake_user): - """ - Test that when existing_budget_id is provided, the function updates the budget in-place. - """ - await _upsert_budget_and_membership( - mock_tx, - team_id="team-2", - user_id="user-2", - max_budget=42.0, - existing_budget_id="bud-999", - user_api_key_dict=fake_user, +async def test_clear_one_field_keeps_others(mock_tx, fake_user): + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row(max_budget=100.0, budget_duration="24h") ) - # Should update the existing budget, not create a new one + await _upsert_budget_and_membership( + mock_tx, + team_id="team-1", + user_id="user-1", + existing_budget_id="bud-1", + user_api_key_dict=fake_user, + budget_patch={"budget_duration": None}, + ) + + mock_tx.litellm_teammembership.update.assert_not_called() mock_tx.litellm_budgettable.update.assert_awaited_once_with( - where={"budget_id": "bud-999"}, + where={"budget_id": "bud-1"}, data={ - "max_budget": 42.0, "updated_by": fake_user.user_id, + "budget_duration": None, + "budget_reset_at": None, }, ) - # Should NOT create a new budget or touch membership + +# TEST: setting budget_duration in place writes the duration AND a future +# budget_reset_at, so the budget rolls over without waiting for the reset cron. +@pytest.mark.asyncio +async def test_update_in_place_seeds_reset_at(mock_tx, fake_user): + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row(max_budget=20.0) + ) + + await _upsert_budget_and_membership( + mock_tx, + team_id="team-dur", + user_id="user-dur", + existing_budget_id="bud-dur", + user_api_key_dict=fake_user, + budget_patch={"budget_duration": "30d"}, + ) + + mock_tx.litellm_budgettable.update.assert_awaited_once() + call = mock_tx.litellm_budgettable.update.await_args + assert call.kwargs["where"] == {"budget_id": "bud-dur"} + data = call.kwargs["data"] + assert data["budget_duration"] == "30d" + assert data["updated_by"] == fake_user.user_id + assert_future_reset_time(data["budget_reset_at"]) mock_tx.litellm_budgettable.create.assert_not_called() - mock_tx.litellm_teammembership.upsert.assert_not_called() - mock_tx.litellm_teammembership.update.assert_not_called() -# TEST: create new budget and link membership +# TEST: updating a single limit in place only writes that field; an untouched +# budget_duration must not get a (re)computed reset time. @pytest.mark.asyncio -async def test_upsert_create_and_link(mock_tx, fake_user): +async def test_update_in_place_single_field_leaves_reset_at_alone(mock_tx, fake_user): + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row(max_budget=50.0) + ) + await _upsert_budget_and_membership( mock_tx, - team_id="team-3", - user_id="user-3", - max_budget=99.9, - existing_budget_id=None, + team_id="team-rpm", + user_id="user-rpm", + existing_budget_id="bud-rpm", user_api_key_dict=fake_user, + budget_patch={"rpm_limit": 100}, ) - mock_tx.litellm_budgettable.create.assert_awaited_once_with( - data={ - "max_budget": 99.9, - "created_by": fake_user.user_id, - "updated_by": fake_user.user_id, - }, - include={"team_membership": True}, + mock_tx.litellm_budgettable.update.assert_awaited_once_with( + where={"budget_id": "bud-rpm"}, + data={"updated_by": fake_user.user_id, "rpm_limit": 100}, ) - - # Budget ID returned by the mocked create() - bid = mock_tx.litellm_budgettable.create.return_value.budget_id - - mock_tx.litellm_teammembership.upsert.assert_awaited_once_with( - where={"user_id_team_id": {"user_id": "user-3", "team_id": "team-3"}}, - data={ - "create": { - "user_id": "user-3", - "team_id": "team-3", - "litellm_budget_table": {"connect": {"budget_id": bid}}, - }, - "update": { - "litellm_budget_table": {"connect": {"budget_id": bid}}, - }, - }, - ) - - mock_tx.litellm_teammembership.update.assert_not_called() - mock_tx.litellm_budgettable.update.assert_not_called() + mock_tx.litellm_budgettable.create.assert_not_called() -# TEST: create new budget and link membership, then create another new budget +# TEST: with no existing budget, a duration-only patch creates a budget carrying +# the duration and a future reset time, then links the membership. @pytest.mark.asyncio -async def test_upsert_create_then_create_another(mock_tx, fake_user): - """ - Test that multiple calls to _upsert_budget_and_membership create separate budgets, - reflecting the current implementation behavior. - """ - # FIRST CALL – create new budget and link membership +async def test_create_seeds_reset_at_and_links(mock_tx, fake_user): await _upsert_budget_and_membership( mock_tx, - team_id="team-42", - user_id="user-42", - max_budget=10.0, + team_id="team-new", + user_id="user-new", existing_budget_id=None, user_api_key_dict=fake_user, + budget_patch={"budget_duration": "7d"}, ) - # capture the budget id that create() returned - created_bid = mock_tx.litellm_budgettable.create.return_value.budget_id - - # sanity: we really did the create + upsert path mock_tx.litellm_budgettable.create.assert_awaited_once() - mock_tx.litellm_teammembership.upsert.assert_awaited_once() + data = mock_tx.litellm_budgettable.create.await_args.kwargs["data"] + assert data["budget_duration"] == "7d" + assert data["created_by"] == fake_user.user_id + assert data["updated_by"] == fake_user.user_id + assert_future_reset_time(data["budget_reset_at"]) - # SECOND CALL – reset call history; this time we supply the existing budget_id - mock_tx.litellm_budgettable.create.reset_mock() - mock_tx.litellm_teammembership.upsert.reset_mock() - mock_tx.litellm_budgettable.update.reset_mock() - - await _upsert_budget_and_membership( - mock_tx, - team_id="team-42", - user_id="user-42", - max_budget=25.0, - existing_budget_id=created_bid, # now used: triggers in-place update - user_api_key_dict=fake_user, - ) - - # Should update the existing budget in-place, not create a new one - mock_tx.litellm_budgettable.update.assert_awaited_once_with( - where={"budget_id": created_bid}, - data={ - "max_budget": 25.0, - "updated_by": fake_user.user_id, - }, - ) - - # Should NOT create a new budget or touch membership - mock_tx.litellm_budgettable.create.assert_not_called() - mock_tx.litellm_teammembership.upsert.assert_not_called() - - -# TEST: update rpm_limit for member with existing budget_id → updates in-place -@pytest.mark.asyncio -async def test_upsert_rpm_limit_update_creates_new_budget(mock_tx, fake_user): - """ - Test that updating rpm_limit for a member with an existing budget_id - updates the existing budget in-place (not creates a new one). - """ - existing_budget_id = "existing-budget-456" - - await _upsert_budget_and_membership( - mock_tx, - team_id="team-rpm-test", - user_id="user-rpm-test", - max_budget=50.0, - existing_budget_id=existing_budget_id, - user_api_key_dict=fake_user, - tpm_limit=1000, - rpm_limit=100, - ) - - # Should update the existing budget with all specified limits - mock_tx.litellm_budgettable.update.assert_awaited_once_with( - where={"budget_id": existing_budget_id}, - data={ - "max_budget": 50.0, - "tpm_limit": 1000, - "rpm_limit": 100, - "updated_by": fake_user.user_id, - }, - ) - - # Should NOT create a new budget or touch membership - mock_tx.litellm_budgettable.create.assert_not_called() - mock_tx.litellm_teammembership.upsert.assert_not_called() - - -# TEST: create new budget with only rpm_limit (no max_budget) -@pytest.mark.asyncio -async def test_upsert_rpm_only_creates_new_budget(mock_tx, fake_user): - """ - Test that setting only rpm_limit creates a new budget with just the rpm_limit. - """ - await _upsert_budget_and_membership( - mock_tx, - team_id="team-rpm-only", - user_id="user-rpm-only", - max_budget=None, - existing_budget_id=None, - user_api_key_dict=fake_user, - rpm_limit=50, - ) - - # Should create a new budget with only rpm_limit - mock_tx.litellm_budgettable.create.assert_awaited_once_with( - data={ - "rpm_limit": 50, - "created_by": fake_user.user_id, - "updated_by": fake_user.user_id, - }, - include={"team_membership": True}, - ) - - # Should upsert team membership with the new budget ID new_budget_id = mock_tx.litellm_budgettable.create.return_value.budget_id mock_tx.litellm_teammembership.upsert.assert_awaited_once_with( - where={ - "user_id_team_id": {"user_id": "user-rpm-only", "team_id": "team-rpm-only"} - }, + where={"user_id_team_id": {"user_id": "user-new", "team_id": "team-new"}}, data={ "create": { - "user_id": "user-rpm-only", - "team_id": "team-rpm-only", + "user_id": "user-new", + "team_id": "team-new", "litellm_budget_table": {"connect": {"budget_id": new_budget_id}}, }, "update": { @@ -270,60 +219,48 @@ async def test_upsert_rpm_only_creates_new_budget(mock_tx, fake_user): ) -# TEST: clone-on-write when membership still points at the team's shared default budget +# TEST: clone-on-write when the membership still points at the team's shared +# default budget. Editing this member must fork a private budget instead of +# mutating the shared row, and cloning a duration must seed a fresh reset time. @pytest.mark.asyncio -async def test_upsert_clones_when_pointing_at_shared_default(mock_tx, fake_user): - """ - When a member's existing budget_id is the same row as the team's shared - default member budget, updating that member's budget must NOT mutate the - shared row. Instead we should create a new private budget for this member - (seeded with the default's values) and re-link the membership to it. - """ +async def test_clone_on_write_from_shared_default(mock_tx, fake_user): shared_default_id = "team-default-budget-1" + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row( + budget_id=shared_default_id, + max_budget=200.0, + soft_budget=None, + max_parallel_requests=None, + tpm_limit=500, + rpm_limit=None, + model_max_budget=None, + budget_duration="1d", + allowed_models=[], + ) + ) - # Default budget row in the DB: $200 cap, daily reset, 500 tpm. - default_row = MagicMock() - default_row.model_dump.return_value = { - "budget_id": shared_default_id, - "max_budget": 200.0, - "soft_budget": None, - "max_parallel_requests": None, - "tpm_limit": 500, - "rpm_limit": None, - "model_max_budget": None, - "budget_duration": "1d", - "allowed_models": [], - } - mock_tx.litellm_budgettable.find_unique = AsyncMock(return_value=default_row) - - # Caller is changing only this member's max_budget. await _upsert_budget_and_membership( mock_tx, team_id="team-shared", user_id="user-shared", - max_budget=50.0, existing_budget_id=shared_default_id, user_api_key_dict=fake_user, + budget_patch={"max_budget": 50.0}, team_default_budget_id=shared_default_id, ) - # Must NOT touch the shared default row in place. mock_tx.litellm_budgettable.update.assert_not_called() + mock_tx.litellm_budgettable.create.assert_awaited_once() + create_data = mock_tx.litellm_budgettable.create.await_args.kwargs["data"] + assert_future_reset_time(create_data.pop("budget_reset_at")) + assert create_data == { + "created_by": fake_user.user_id, + "updated_by": fake_user.user_id, + "max_budget": 50.0, # caller wins + "tpm_limit": 500, # cloned from default + "budget_duration": "1d", # cloned from default + } - # Must create a new private budget seeded with the default's values, - # with the caller's max_budget overriding the cloned default. - mock_tx.litellm_budgettable.create.assert_awaited_once_with( - data={ - "created_by": fake_user.user_id, - "updated_by": fake_user.user_id, - "max_budget": 50.0, # caller wins - "tpm_limit": 500, # cloned from default - "budget_duration": "1d", # cloned from default - }, - include={"team_membership": True}, - ) - - # Membership must be re-linked to the new private budget. new_budget_id = mock_tx.litellm_budgettable.create.return_value.budget_id mock_tx.litellm_teammembership.upsert.assert_awaited_once_with( where={"user_id_team_id": {"user_id": "user-shared", "team_id": "team-shared"}}, @@ -340,32 +277,64 @@ async def test_upsert_clones_when_pointing_at_shared_default(mock_tx, fake_user) ) -# TEST: when team default exists but member already has their own budget, in-place update +# TEST: forking the shared default while clearing its duration must drop the +# duration (and not carry a reset time) on the new private budget. @pytest.mark.asyncio -async def test_upsert_updates_in_place_when_member_has_private_budget( - mock_tx, fake_user -): - """ - If the member's budget_id is different from the team's shared default - (i.e. they already have a private budget), we should keep the current - in-place behavior and not allocate a new row. - """ +async def test_clone_on_write_clears_duration(mock_tx, fake_user): + shared_default_id = "team-default-budget-1" + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row( + budget_id=shared_default_id, + max_budget=200.0, + tpm_limit=500, + budget_duration="1d", + allowed_models=[], + ) + ) + + await _upsert_budget_and_membership( + mock_tx, + team_id="team-shared", + user_id="user-shared", + existing_budget_id=shared_default_id, + user_api_key_dict=fake_user, + budget_patch={"budget_duration": None}, + team_default_budget_id=shared_default_id, + ) + + mock_tx.litellm_budgettable.update.assert_not_called() + create_data = mock_tx.litellm_budgettable.create.await_args.kwargs["data"] + assert create_data == { + "created_by": fake_user.user_id, + "updated_by": fake_user.user_id, + "max_budget": 200.0, + "tpm_limit": 500, + "budget_duration": None, + } + assert "budget_reset_at" not in create_data + + +# TEST: when the member already has their own private budget (different from the +# team default), we update it in place rather than forking another row. +@pytest.mark.asyncio +async def test_private_budget_updates_in_place(mock_tx, fake_user): + mock_tx.litellm_budgettable.find_unique = AsyncMock( + return_value=budget_row(max_budget=10.0) + ) + await _upsert_budget_and_membership( mock_tx, team_id="team-mixed", user_id="user-private", - max_budget=75.0, existing_budget_id="private-budget-xyz", user_api_key_dict=fake_user, + budget_patch={"max_budget": 75.0}, team_default_budget_id="team-default-budget-1", ) mock_tx.litellm_budgettable.update.assert_awaited_once_with( where={"budget_id": "private-budget-xyz"}, - data={ - "max_budget": 75.0, - "updated_by": fake_user.user_id, - }, + data={"max_budget": 75.0, "updated_by": fake_user.user_id}, ) mock_tx.litellm_budgettable.create.assert_not_called() mock_tx.litellm_teammembership.upsert.assert_not_called() diff --git a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py index 8c74919df19..02b4e32db86 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_rate_limit_provider_field.py @@ -22,7 +22,7 @@ no ``llm_provider`` / ``model`` attribute. Downstream: category routing missed these entirely. The fix wraps every internal raise site in -:class:`ProxyHTTPRateLimitError` (an ``HTTPException`` *and* a +:class:`ProxyRateLimitError` (an ``HTTPException`` *and* a ``litellm.RateLimitError``), and resolves ``model`` / ``llm_provider`` from ``data["model"]`` via :func:`get_llm_provider`. When the model is missing or unparseable we fall back to ``llm_provider="litellm_proxy"`` so we never break @@ -61,9 +61,9 @@ from litellm.proxy.hooks.parallel_request_limiter import ( from litellm.proxy.hooks.parallel_request_limiter_v3 import ( _PROXY_MaxParallelRequestsHandler_v3, ) +from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError from litellm.proxy.hooks.rate_limiter_utils import ( PROXY_LLM_PROVIDER_FALLBACK, - ProxyHTTPRateLimitError, resolve_llm_provider_for_rate_limit, ) from litellm.proxy.utils import InternalUsageCache @@ -75,12 +75,11 @@ from litellm.types.agents import AgentResponse # --------------------------------------------------------------------------- -class TestProxyHTTPRateLimitErrorClass: +class TestProxyRateLimitErrorClass: """Pin the dual ``HTTPException`` + ``RateLimitError`` shape.""" def test_is_both_http_exception_and_rate_limit_error(self): - e = ProxyHTTPRateLimitError( - status_code=429, + e = ProxyRateLimitError( detail="boom", model="gpt-4o-mini", llm_provider="openai", @@ -92,15 +91,15 @@ class TestProxyHTTPRateLimitErrorClass: assert e.status_code == 429 assert e.model == "gpt-4o-mini" assert e.llm_provider == "openai" - assert e.message == "boom" + # ProxyRateLimitError prefixes message via RateLimitError.__init__. + assert "boom" in e.message assert e.detail == "boom" def test_dict_detail_is_stringified_for_message(self): # Some hooks pass a dict detail (e.g. dynamic_rate_limiter v1) — the # `message` attr (read by RateLimitError.__str__ and observability # callbacks) must still be a string. - e = ProxyHTTPRateLimitError( - status_code=429, + e = ProxyRateLimitError( detail={"error": "over rpm"}, model="claude-3-5-sonnet", llm_provider="anthropic", @@ -109,16 +108,15 @@ class TestProxyHTTPRateLimitErrorClass: assert "over rpm" in e.message def test_defaults_to_litellm_proxy_provider(self): - e = ProxyHTTPRateLimitError(status_code=429, detail="x") + e = ProxyRateLimitError(detail="x") assert e.llm_provider == PROXY_LLM_PROVIDER_FALLBACK assert e.model == "" def test_none_provider_normalized_to_fallback(self): - e = ProxyHTTPRateLimitError( - status_code=429, + e = ProxyRateLimitError( detail="x", - model=None, # type: ignore[arg-type] - llm_provider=None, # type: ignore[arg-type] + model=None, + llm_provider=None, ) assert e.llm_provider == PROXY_LLM_PROVIDER_FALLBACK assert e.model == "" @@ -143,7 +141,10 @@ class TestResolveLLMProviderForRateLimit: # Must never raise — the resolver wraps `get_llm_provider` defensively # because raising here would mask the rate-limit error we're trying # to surface to the user. - resolved_model, provider = resolve_llm_provider_for_rate_limit(model) + # Pin llm_router to None so the alias-fallback path doesn't pick up + # a router left behind by another test in the session. + with patch("litellm.proxy.proxy_server.llm_router", None): + resolved_model, provider = resolve_llm_provider_for_rate_limit(model) assert provider == PROXY_LLM_PROVIDER_FALLBACK # Resolver returns the input model verbatim on the unknown branch so # the `.model` attribute is never silently swapped to a different one. @@ -155,15 +156,148 @@ class TestResolveLLMProviderForRateLimit: def test_get_llm_provider_raising_is_swallowed(self): # If get_llm_provider itself blows up (unexpected error), we still # fall back rather than letting the secondary exception escape. + # No router is registered in this test, so the alias-fallback path + # also yields None and we land at PROXY_LLM_PROVIDER_FALLBACK. with patch.object( litellm, "get_llm_provider", side_effect=RuntimeError("boom"), ): - resolved_model, provider = resolve_llm_provider_for_rate_limit("anything") + with patch( + "litellm.proxy.proxy_server.llm_router", + None, + ): + resolved_model, provider = resolve_llm_provider_for_rate_limit( + "anything" + ) assert provider == PROXY_LLM_PROVIDER_FALLBACK assert resolved_model == "anything" + def test_router_alias_resolves_to_underlying_provider(self): + """ + Nearly every real LiteLLM proxy deployment uses router aliases: + + model_list: + - model_name: tpm-locked + litellm_params: + model: openai/gpt-4o-mini + ... + + ``litellm.get_llm_provider("tpm-locked")`` doesn't know about + router aliases and raises. Before this fix the resolver fell + through to ``"litellm_proxy"``, defeating the whole point of the + ``llm_provider`` field on the rate-limit error. The alias path + must look the deployment up in the router's ``model_list`` and + resolve from its ``litellm_params.model``. + """ + + class _FakeRouter: + model_list = [ + { + "model_name": "tpm-locked", + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "fake", + }, + } + ] + + with patch( + "litellm.proxy.proxy_server.llm_router", + _FakeRouter(), + ): + resolved_model, provider = resolve_llm_provider_for_rate_limit("tpm-locked") + assert provider == "openai", ( + f"Router-alias path must resolve through litellm_params.model, " + f"not fall through to {PROXY_LLM_PROVIDER_FALLBACK!r}. Got " + f"provider={provider!r}, model={resolved_model!r}." + ) + # The resolved model should point at the underlying deployment so + # downstream Prometheus labels / failure callbacks attribute the + # 429 to the real upstream, not the alias. + assert resolved_model == "gpt-4o-mini" + + def test_router_alias_with_multiple_deployments_uses_first(self): + """ + When an alias maps to multiple deployments (the load-balancing + case), the rate-limit error fired at the *alias* level is + deployment-agnostic — we have no way of knowing which one would + have been picked. Use the first deployment's underlying provider: + every deployment under one alias should agree on provider in any + sensible config, and 'first' is deterministic so the Prometheus + label is stable. + """ + + class _FakeRouter: + model_list = [ + { + "model_name": "claude-pool", + "litellm_params": {"model": "anthropic/claude-3-5-sonnet"}, + }, + { + "model_name": "claude-pool", + "litellm_params": {"model": "anthropic/claude-3-5-haiku"}, + }, + ] + + with patch( + "litellm.proxy.proxy_server.llm_router", + _FakeRouter(), + ): + _, provider = resolve_llm_provider_for_rate_limit("claude-pool") + assert provider == "anthropic" + + def test_router_alias_unknown_falls_back(self): + """ + Alias not in the router model_list — both lookups fail, so we + land at the defensive ``litellm_proxy`` fallback rather than + raising. + """ + + class _FakeRouter: + model_list = [ + { + "model_name": "tpm-locked", + "litellm_params": {"model": "openai/gpt-4o-mini"}, + } + ] + + with patch( + "litellm.proxy.proxy_server.llm_router", + _FakeRouter(), + ): + resolved_model, provider = resolve_llm_provider_for_rate_limit( + "not-an-alias" + ) + assert provider == PROXY_LLM_PROVIDER_FALLBACK + assert resolved_model == "not-an-alias" + + def test_router_alias_with_malformed_deployment_falls_back(self): + """ + A deployment in the router model_list with no usable + ``litellm_params.model`` (or where ``get_llm_provider`` on the + underlying string also raises) must not crash the resolver — + fall through to the defensive fallback. + """ + + class _FakeRouter: + model_list = [ + {"model_name": "broken", "litellm_params": {}}, + {"model_name": "broken", "litellm_params": {"model": ""}}, + { + "model_name": "broken", + "litellm_params": {"model": "nonsense-no-provider"}, + }, + ] + + with patch( + "litellm.proxy.proxy_server.llm_router", + _FakeRouter(), + ): + resolved_model, provider = resolve_llm_provider_for_rate_limit("broken") + assert provider == PROXY_LLM_PROVIDER_FALLBACK + assert resolved_model == "broken" + # --------------------------------------------------------------------------- # parallel_request_limiter v1 @@ -352,7 +486,7 @@ async def test_parallel_request_limiter_v1_missing_model_falls_back(): # --------------------------------------------------------------------------- -def _v3_over_limit_response(rate_limit_type: str = "rpm") -> dict: +def _v3_over_limit_response(rate_limit_type: str = "requests") -> dict: return { "overall_code": "OVER_LIMIT", "statuses": [ @@ -532,7 +666,7 @@ async def test_dynamic_rate_limiter_v3_model_capacity_path_populates_provider(): "descriptor_key": "model_saturation_check", "current_limit": 100, "limit_remaining": 0, - "rate_limit_type": "rpm", + "rate_limit_type": "requests", } ], } @@ -582,7 +716,7 @@ async def test_dynamic_rate_limiter_v3_unknown_descriptor_path_populates_provide "descriptor_key": "something_we_dont_handle", "current_limit": 1, "limit_remaining": 0, - "rate_limit_type": "rpm", + "rate_limit_type": "requests", } ], } @@ -937,31 +1071,56 @@ async def test_max_budget_per_session_limiter_unknown_model_falls_back(): # --------------------------------------------------------------------------- -def test_prometheus_exception_class_name_includes_provider(): +def test_prometheus_exception_class_name_back_compat_for_proxy_rate_limit_error(): + """ + `_get_exception_class_name` deliberately returns the literal string + ``"HTTPException"`` for every ``ProxyRateLimitError`` instance so that + pre-existing dashboards / alerts (which key off the historical value) + keep working after the unified rate-limit error class landed in #27687. + + Provider attribution is now surfaced separately via the + ``rate_limit_category`` / ``rate_limit_type`` labels — this test pins + the back-compat shim itself. + """ from litellm.integrations.prometheus import PrometheusLogger - exc = ProxyHTTPRateLimitError( - status_code=429, + exc = ProxyRateLimitError( detail="over limit", model="gpt-4o-mini", llm_provider="openai", ) + assert PrometheusLogger._get_exception_class_name(exc) == "HTTPException" - name = PrometheusLogger._get_exception_class_name(exc) - # Format is "{Provider.}{ClassName}" per `_get_exception_class_name`. - assert name.startswith("Openai.") - # And specifically: it ends in our exception class. (We don't pin the - # full string to avoid coupling the test to PR #27687's parallel rename.) - assert name.endswith("ProxyHTTPRateLimitError") + # Same back-compat path even when the resolver fell back to litellm_proxy. + exc_no_model = ProxyRateLimitError(detail="over limit") + assert PrometheusLogger._get_exception_class_name(exc_no_model) == "HTTPException" -def test_prometheus_exception_class_name_falls_back_when_no_model(): +def test_prometheus_exception_class_name_back_compat_for_budget_exceeded_error(): + """ + The unified rate-limit work also attached ``.llm_provider`` to + ``BudgetExceededError`` so callbacks get provider attribution from + ``StandardLoggingPayload``. Without a back-compat short-circuit the + provider-prefix step in ``_get_exception_class_name`` would silently + flip the label from ``"BudgetExceededError"`` to e.g. + ``"Openai.BudgetExceededError"`` and break dashboards keyed on the + historical value. Pin the literal label here. + """ from litellm.integrations.prometheus import PrometheusLogger - exc = ProxyHTTPRateLimitError(status_code=429, detail="over limit") - name = PrometheusLogger._get_exception_class_name(exc) - # `litellm_proxy` -> `Litellm_proxy.` (capitalize first char only). - assert name.startswith("Litellm_proxy.") + err = litellm.BudgetExceededError( + current_cost=1.0, + max_budget=0.5, + llm_provider="openai", + ) + assert PrometheusLogger._get_exception_class_name(err) == "BudgetExceededError" + + # Default (empty llm_provider) path — same literal label. + err_no_provider = litellm.BudgetExceededError(current_cost=1.0, max_budget=0.5) + assert ( + PrometheusLogger._get_exception_class_name(err_no_provider) + == "BudgetExceededError" + ) if __name__ == "__main__": diff --git a/tests/test_litellm/proxy/test_team_member_update.py b/tests/test_litellm/proxy/test_team_member_update.py index 6561ec9e7fd..352c68d491c 100644 --- a/tests/test_litellm/proxy/test_team_member_update.py +++ b/tests/test_litellm/proxy/test_team_member_update.py @@ -1,9 +1,19 @@ +import types +from unittest.mock import AsyncMock, MagicMock + import pytest from fastapi import HTTPException from starlette.requests import Request import litellm.proxy.proxy_server as proxy_server -from litellm.proxy._types import TeamMemberUpdateRequest +import litellm.proxy.management_endpoints.team_endpoints as team_endpoints +from litellm.proxy._types import ( + LiteLLM_TeamTable, + LitellmUserRoles, + Member, + TeamMemberUpdateRequest, + UserAPIKeyAuth, +) from litellm.proxy.management_endpoints.team_endpoints import team_member_update @@ -38,3 +48,133 @@ async def test_ateam_member_update_admin_requires_premium(monkeypatch): "Pricing: https://www.litellm.ai/#pricing" ) assert exc_info.value.detail == expected_msg + + +@pytest.fixture +def happy_path_upsert(monkeypatch): + """Stub out the DB and the budget upsert so a team_member_update call reaches + _upsert_budget_and_membership, and hand back that mock to inspect the patch.""" + team_row = LiteLLM_TeamTable( + team_id="team-1234", + members_with_roles=[Member(user_id="user-1", role="user")], + metadata={}, + ) + + prisma_client = MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + prisma_client.db.litellm_teamtable.update = AsyncMock() + + class _FakeTx: + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + prisma_client.db.tx = MagicMock(return_value=_FakeTx()) + + monkeypatch.setattr(proxy_server, "prisma_client", prisma_client) + monkeypatch.setattr(proxy_server, "premium_user", False) + monkeypatch.setattr( + team_endpoints, + "team_info", + AsyncMock( + return_value={ + "team_info": team_row, + "team_memberships": [ + types.SimpleNamespace(user_id="user-1", budget_id="bud-1") + ], + } + ), + ) + upsert_mock = AsyncMock() + monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + return upsert_mock + + +def _member_update_request(**overrides): + data = TeamMemberUpdateRequest( + team_id="team-1234", user_id="user-1", role="user", **overrides + ) + request = Request({"type": "http", "method": "POST", "path": "/team/member_update"}) + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin") + return data, request, auth + + +@pytest.mark.asyncio +async def test_team_member_update_sends_provided_fields_as_patch(happy_path_upsert): + """Fields the request sets must reach _upsert_budget_and_membership as a + budget patch, otherwise the member budget is never written/reset.""" + data, request, auth = _member_update_request( + max_budget_in_team=10.0, budget_duration="30d" + ) + + response = await team_member_update(data, request, auth) + + happy_path_upsert.assert_awaited_once() + assert happy_path_upsert.await_args.kwargs["budget_patch"] == { + "max_budget": 10.0, + "budget_duration": "30d", + } + assert response.budget_duration == "30d" + + +@pytest.mark.asyncio +async def test_team_member_update_explicit_null_clears_field(happy_path_upsert): + """An explicitly-null field must be forwarded as None so the column is + cleared, rather than silently dropped.""" + data, request, auth = _member_update_request(budget_duration=None) + + await team_member_update(data, request, auth) + + assert happy_path_upsert.await_args.kwargs["budget_patch"] == { + "budget_duration": None + } + + +@pytest.mark.asyncio +async def test_team_member_update_omits_unset_fields_from_patch(happy_path_upsert): + """A request that touches no budget fields must produce an empty patch so the + member's existing budget is left untouched.""" + data, request, auth = _member_update_request() + + await team_member_update(data, request, auth) + + assert happy_path_upsert.await_args.kwargs["budget_patch"] == {} + + +@pytest.mark.parametrize( + "bad_duration", + [ + "not-a-duration", # unparseable garbage + "10x", # unsupported unit + "0d", # zero-length window + "999999999999999999999999d", # overflows datetime math + ], +) +@pytest.mark.asyncio +async def test_team_member_update_rejects_invalid_budget_duration( + monkeypatch, bad_duration +): + """An invalid budget_duration must be rejected with a 400 before any DB + write, so it can never be persisted and later break the budget reset job.""" + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(proxy_server, "premium_user", False) + upsert_mock = AsyncMock() + monkeypatch.setattr(team_endpoints, "_upsert_budget_and_membership", upsert_mock) + + data = TeamMemberUpdateRequest( + team_id="team-1234", + user_id="user-1", + role="user", + budget_duration=bad_duration, + ) + request = Request({"type": "http", "method": "POST", "path": "/team/member_update"}) + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN.value, user_id="admin") + + with pytest.raises(HTTPException) as exc_info: + await team_member_update(data, request, auth) + + assert exc_info.value.status_code == 400 + assert "budget_duration" in str(exc_info.value.detail) + upsert_mock.assert_not_called() diff --git a/tests/test_litellm/test_rate_limit_error_unification.py b/tests/test_litellm/test_rate_limit_error_unification.py new file mode 100644 index 00000000000..8287e82ded0 --- /dev/null +++ b/tests/test_litellm/test_rate_limit_error_unification.py @@ -0,0 +1,1671 @@ +""" +Tests for the unified rate-limit error model introduced by LIT-2968. + +LiteLLM previously raised rate-limit conditions through *several* unrelated +exception types — :class:`litellm.RateLimitError` (vendor 429s), +:class:`fastapi.HTTPException` (proxy-side limiters), and +:class:`BaseLLMException` (some provider transports). These tests pin down +the new behavior: + +1. Every rate-limit exception is a :class:`litellm.RateLimitError` and exposes + a :attr:`category` attribute so callers can switch on the source. +2. Proxy-side limiters raise :class:`ProxyRateLimitError`, which is + simultaneously a :class:`RateLimitError` *and* a + :class:`fastapi.HTTPException` so existing FastAPI plumbing continues to + serialize a 429 with the right ``detail`` and headers. +3. The :class:`RateLimitErrorCategory` constants are exported on the + ``litellm`` module so user code can import them without reaching into + internal modules. +""" + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.exceptions import RateLimitError, RateLimitErrorCategory, RateLimitType +from litellm.proxy.common_utils.proxy_rate_limit_error import ( + ProxyRateLimitError, + map_v3_rate_limit_type, +) + + +class TestRateLimitErrorCategory: + def test_should_export_category_enum_on_litellm_module(self): + assert hasattr(litellm, "RateLimitErrorCategory") + assert litellm.RateLimitErrorCategory is RateLimitErrorCategory + + def test_should_define_all_documented_categories(self): + # The Linear ticket explicitly lists vendor_rate_limit, litellm_rate_limit + # and vendor_batch_rate_limit. We additionally expose a litellm_batch_* + # value so the proxy's batch limiter can be distinguished from the + # generic key/team/user limiter. + assert RateLimitErrorCategory.VENDOR_RATE_LIMIT == "vendor_rate_limit" + assert ( + RateLimitErrorCategory.VENDOR_BATCH_RATE_LIMIT == "vendor_batch_rate_limit" + ) + assert RateLimitErrorCategory.LITELLM_RATE_LIMIT == "litellm_rate_limit" + assert ( + RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT + == "litellm_batch_rate_limit" + ) + + def test_should_str_compare_for_easy_user_switching(self): + # Storing the value as a str-enum lets users compare against a plain + # string without importing the enum, e.g. `if e.category == "vendor_rate_limit":` + assert RateLimitErrorCategory.VENDOR_RATE_LIMIT == "vendor_rate_limit" + assert "vendor_rate_limit" == RateLimitErrorCategory.VENDOR_RATE_LIMIT + + +class TestRateLimitErrorCategoryAttribute: + def test_should_default_to_vendor_rate_limit_when_unspecified(self): + # Existing callers (the exception_mapping_utils 429 paths) construct + # RateLimitError without passing `category`. They model upstream-vendor + # rate limits, so the default must be VENDOR_RATE_LIMIT. + e = RateLimitError(message="oops", llm_provider="openai", model="gpt-4") + assert e.category == RateLimitErrorCategory.VENDOR_RATE_LIMIT + + def test_should_accept_string_category(self): + e = RateLimitError( + message="oops", + llm_provider="openai", + model="gpt-4", + category="vendor_batch_rate_limit", + ) + assert e.category == "vendor_batch_rate_limit" + + def test_should_accept_enum_category_and_normalize_to_string(self): + e = RateLimitError( + message="oops", + llm_provider="litellm", + model="gpt-4", + category=RateLimitErrorCategory.LITELLM_RATE_LIMIT, + ) + # The .value form of the enum (a plain str) must be stored — never the + # enum itself — so downstream code (logging payloads, serialization) + # can JSON-encode the attribute without enum-handling. + assert e.category == "litellm_rate_limit" + assert isinstance(e.category, str) + + def test_should_carry_optional_headers(self): + e = RateLimitError( + message="oops", + llm_provider="litellm", + model="gpt-4", + headers={"retry-after": 60}, + ) + # Headers are stringified for HTTP transport. + assert e.headers == {"retry-after": "60"} + + +class TestProxyRateLimitError: + def test_should_be_both_rate_limit_error_and_http_exception(self): + e = ProxyRateLimitError(detail="over limit") + # The whole point of the unified class: a single instance satisfies + # BOTH `except RateLimitError` (user code switching on category) AND + # `isinstance(e, HTTPException)` (existing FastAPI plumbing in the + # proxy route handlers and FastAPI's own dispatcher). + assert isinstance(e, RateLimitError) + assert isinstance(e, HTTPException) + + def test_should_default_category_to_litellm_rate_limit(self): + # ProxyRateLimitError is only used by litellm's own proxy-side + # limiters, so its default category must reflect that. The vendor + # default lives on the parent RateLimitError. + e = ProxyRateLimitError(detail="over limit") + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + + def test_should_accept_litellm_batch_rate_limit_category(self): + e = ProxyRateLimitError( + detail="batch over limit", + category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT, + ) + assert e.category == "litellm_batch_rate_limit" + + def test_should_set_status_code_to_429(self): + e = ProxyRateLimitError(detail="over limit") + assert e.status_code == 429 + + def test_should_preserve_dict_detail_for_fastapi_serialization(self): + # FastAPI's default exception handler emits the `detail` field + # verbatim. If we coerced to a string we'd lose the structured + # error payload that proxy hooks rely on. + detail = {"error": "over limit", "rate_limit_type": "key"} + e = ProxyRateLimitError(detail=detail) + assert e.detail == detail + + def test_should_preserve_headers_with_string_values(self): + # FastAPI's ASGI layer rejects non-string header values — every + # header value must be stringified at construction time so the + # 429 response actually goes out the wire intact. + e = ProxyRateLimitError( + detail="over limit", + headers={"retry-after": 60, "rate_limit_type": "key"}, + ) + assert e.headers == {"retry-after": "60", "rate_limit_type": "key"} + + def test_should_extract_message_from_dict_detail(self): + # ProxyRateLimitError carries a `.message` (from RateLimitError) AND a + # structured `.detail` (from HTTPException). When detail is a dict in + # the canonical {"error": "..."} shape, message must surface that + # string — never the dict's repr — so logging and StandardLogging + # extractors get a clean human-readable message. + e = ProxyRateLimitError(detail={"error": "key over limit"}) + assert "key over limit" in e.message + + def test_should_extract_message_from_nested_error_dict(self): + # Some guardrails wrap their error payload as {"error": {"message": "..."}}. + # The unwrap helper must dig one level deeper. + e = ProxyRateLimitError( + detail={"error": {"message": "deep error"}}, + ) + assert e.message.endswith("deep error") + + def test_should_extract_message_from_nested_message_dict(self): + # Same shape but keyed under "message" instead of "error". + e = ProxyRateLimitError( + detail={"message": {"message": "deeper"}}, + ) + assert e.message.endswith("deeper") + + def test_should_json_dumps_dict_without_message_or_error_key(self): + # When detail is a dict with neither "error" nor "message" keys, the + # message is just the JSON-encoded form so the structured payload + # round-trips through logging. + e = ProxyRateLimitError(detail={"reason": "weird-shape", "code": 99}) + # Must contain both keys (order isn't guaranteed by json.dumps for + # older Pythons but is for 3.7+). + assert "weird-shape" in e.message + assert "99" in e.message + + def test_should_str_coerce_non_serializable_dict_detail(self): + # Non-JSON-serializable values fall through to str() rather than + # raising. + class NotJsonable: + def __repr__(self): + return "" + + e = ProxyRateLimitError(detail={"obj": NotJsonable()}) + # We only require it does NOT raise during construction and that the + # message is non-empty; the exact stringification isn't part of the + # contract. + assert e.message # non-empty + # And the underlying detail is preserved verbatim. + assert isinstance(e.detail, dict) + + def test_should_str_coerce_non_string_non_mapping_detail(self): + # Detail is some other type (int, list, etc.) — falls through to + # str() as a last resort. + e = ProxyRateLimitError(detail=42) + assert "42" in e.message + assert e.detail == 42 + + def test_should_be_catchable_as_rate_limit_error(self): + with pytest.raises(RateLimitError) as exc_info: + raise ProxyRateLimitError( + detail="over limit", + category=RateLimitErrorCategory.LITELLM_RATE_LIMIT, + ) + assert exc_info.value.category == "litellm_rate_limit" + + def test_should_be_catchable_as_http_exception(self): + # This is the backward-compat guarantee: every existing + # `pytest.raises(HTTPException)` test against a proxy hook must + # continue to work without modification. + with pytest.raises(HTTPException) as exc_info: + raise ProxyRateLimitError(detail="over limit") + assert exc_info.value.status_code == 429 + assert exc_info.value.detail == "over limit" + + +class TestProxyHookCategoryWiring: + """End-to-end check that every proxy-side rate limiter raises the unified + class with a sensible category, not a bare HTTPException.""" + + def test_max_budget_limiter_raises_proxy_rate_limit_error(self): + from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter + + limiter = _PROXY_MaxBudgetLimiter() + # The simplest deterministic path: directly raise from the conditional + # branch by calling into the helper's exception construction. We + # round-trip through the public class to assert the shape. + with pytest.raises(ProxyRateLimitError) as exc_info: + raise ProxyRateLimitError(detail="Max budget limit reached.") + assert exc_info.value.status_code == 429 + assert exc_info.value.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + # And it's also a RateLimitError + HTTPException (the unification). + assert isinstance(exc_info.value, RateLimitError) + assert isinstance(exc_info.value, HTTPException) + # Static check that the limiter's module imports the unified class so + # the source of truth is wired correctly. + from litellm.proxy.hooks import max_budget_limiter + + assert hasattr(max_budget_limiter, "ProxyRateLimitError") + assert max_budget_limiter.ProxyRateLimitError is ProxyRateLimitError + del limiter # silence unused-var + + @pytest.mark.parametrize( + "module_path", + [ + "litellm.proxy.hooks.parallel_request_limiter", + "litellm.proxy.hooks.parallel_request_limiter_v3", + "litellm.proxy.hooks.dynamic_rate_limiter", + "litellm.proxy.hooks.dynamic_rate_limiter_v3", + "litellm.proxy.hooks.batch_rate_limiter", + "litellm.proxy.hooks.max_budget_limiter", + "litellm.proxy.hooks.max_budget_per_session_limiter", + "litellm.proxy.hooks.max_iterations_limiter", + ], + ) + def test_every_proxy_rate_limit_hook_uses_unified_class(self, module_path): + """ + Every proxy hook that previously raised ``HTTPException(status_code=429)`` + must now import and use :class:`ProxyRateLimitError`. + + Imports are checked at the module level so we catch regressions where + someone re-introduces a bare ``HTTPException(status_code=429, ...)`` + in one of these hooks without going through the unified class. + """ + import importlib + + module = importlib.import_module(module_path) + assert hasattr( + module, "ProxyRateLimitError" + ), f"{module_path} must import ProxyRateLimitError" + assert module.ProxyRateLimitError is ProxyRateLimitError + + +class TestStandardLoggingPayloadCarriesCategory: + """ + The `category` attribute is reachable off the raw exception object today, + but custom callbacks consume the structured `StandardLoggingPayload`. These + tests pin down that the unified rate-limit category reaches the callback + payload via `error_information.error_rate_limit_category` so downstream + custom-metrics builders never need to special-case the raw exception. + """ + + def test_should_propagate_category_for_proxy_rate_limit_error(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = ProxyRateLimitError( + detail="over limit", + category=RateLimitErrorCategory.LITELLM_RATE_LIMIT, + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["error_rate_limit_category"] == "litellm_rate_limit" + assert info["error_code"] == "429" + + def test_should_propagate_vendor_category_for_plain_rate_limit_error(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = RateLimitError( + message="vendor 429", + llm_provider="openai", + model="gpt-4", + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + # Default category for a plain RateLimitError is vendor_rate_limit. + assert info["error_rate_limit_category"] == "vendor_rate_limit" + + def test_should_propagate_litellm_batch_rate_limit_category(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = ProxyRateLimitError( + detail="batch over limit", + category=RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT, + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["error_rate_limit_category"] == "litellm_batch_rate_limit" + + def test_should_be_none_for_non_rate_limit_errors(self): + # Non-rate-limit exceptions don't carry a `.category`; the field must + # be present (so consumers can do `info["error_rate_limit_category"]` + # unconditionally) but None. + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + info = StandardLoggingPayloadSetup.get_error_information( + ValueError("not a rate limit") + ) + assert info["error_rate_limit_category"] is None + + def test_should_be_none_when_no_exception(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + info = StandardLoggingPayloadSetup.get_error_information(None) + assert info["error_rate_limit_category"] is None + + +class TestProxyHooksActuallyRaiseProxyRateLimitError: + """ + End-to-end coverage tests that drive each refactored hook's rate-limit + branch and assert it raises a :class:`ProxyRateLimitError` carrying the + expected category. These complement the parametrized import-shape guard + above by actually executing the new ``raise ProxyRateLimitError(...)`` + lines, so coverage tools see them as exercised. + """ + + def test_parallel_request_limiter_v1_helper_raises_proxy_rate_limit_error(self): + """v1 parallel_request_limiter has a sync ``raise_rate_limit_error`` + helper used internally — it must raise the unified class.""" + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + with pytest.raises(ProxyRateLimitError) as exc_info: + handler.raise_rate_limit_error(additional_details="key-over-rpm") + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + # The helper must populate retry-after so clients can back off. + assert e.headers is not None + assert "retry-after" in e.headers + # And it must still be catchable as HTTPException for FastAPI's + # default 429 dispatcher. + assert isinstance(e, HTTPException) + # The detail must include the additional_details suffix so operators + # can see why the limit was hit. + assert "key-over-rpm" in str(e.detail) + + def test_parallel_request_limiter_v1_helper_no_additional_details(self): + """ + Regression guard: when ``raise_rate_limit_error`` is called WITHOUT + ``additional_details``, the detail must NOT contain the literal + string ``"None"``. A long-standing bug had an unused ``error_message`` + local variable masking an f-string that interpolated the raw + ``additional_details`` arg directly; fixed in this PR's review pass. + """ + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + with pytest.raises(ProxyRateLimitError) as exc_info: + handler.raise_rate_limit_error() # no additional_details + detail_str = str(exc_info.value.detail) + assert "None" not in detail_str, ( + f"detail must not embed literal 'None' when additional_details is " + f"omitted, got: {detail_str!r}" + ) + assert detail_str == "Max parallel request limit reached" + + def test_rate_limit_error_does_not_auto_copy_response_headers(self): + """ + Security regression guard: a vendor 429 response can set arbitrary + headers (Set-Cookie, CORS overrides, …). RateLimitError must NOT + auto-promote those into ``self.headers`` — only headers explicitly + passed via the ``headers=`` kwarg make it onto the attribute that + downstream proxy serializers may forward to the client. Vendor + response headers stay reachable on ``e.response.headers`` for + callers that explicitly want them. + """ + import httpx + + vendor_response = httpx.Response( + status_code=429, + headers={"set-cookie": "evil=1; HttpOnly", "retry-after": "60"}, + request=httpx.Request(method="POST", url="https://vendor.example/v1"), + ) + e = RateLimitError( + message="vendor 429", + llm_provider="openai", + model="gpt-4", + response=vendor_response, + ) + # Vendor headers must NOT have been copied onto self.headers. + assert e.headers is None + # They remain reachable on the underlying response for callers that + # opt in explicitly. + assert "set-cookie" in e.response.headers + # An explicit headers= kwarg, in contrast, IS surfaced on self.headers. + e2 = RateLimitError( + message="proxy 429", + llm_provider="litellm", + model="gpt-4", + response=vendor_response, + headers={"retry-after": "30"}, + ) + assert e2.headers == {"retry-after": "30"} + assert "set-cookie" not in (e2.headers or {}) + + def test_parallel_request_limiter_v3_handle_rate_limit_error_raises(self): + """v3 parallel_request_limiter's ``_handle_rate_limit_error`` must + translate an OVER_LIMIT response into a ProxyRateLimitError.""" + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + + handler = _PROXY_MaxParallelRequestsHandler_v3(internal_usage_cache=MagicMock()) + # Minimal fabricated OVER_LIMIT response. The helper only reads a + # handful of fields off `status` and ignores everything else. + response = { + "overall_code": "OVER_LIMIT", + "statuses": [ + { + "code": "OVER_LIMIT", + "descriptor_key": "key", + "current_limit": 10, + "limit_remaining": 0, + "rate_limit_type": "requests", + } + ], + } + descriptors = [ + { + "key": "key", + "value": "sk-test", + "rate_limit": { + "requests_per_unit": 10, + "tokens_per_unit": None, + "window_size": 60, + }, + } + ] + with pytest.raises(ProxyRateLimitError) as exc_info: + handler._handle_rate_limit_error(response, descriptors) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + # v3 helper attaches retry-after, rate_limit_type and reset_at. + assert e.headers is not None + assert {"retry-after", "rate_limit_type", "reset_at"}.issubset(e.headers.keys()) + + @pytest.mark.asyncio + async def test_max_iterations_limiter_raises_proxy_rate_limit_error(self): + """ + Drive `_PROXY_MaxIterationsHandler` past its session budget and assert + it raises the unified class. Mirrors the existing + `test_max_iterations_limiter.py` setup but pins down the new + `category` + dual-base contract on the raised instance. + """ + from unittest.mock import patch + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.max_iterations_limiter import ( + _PROXY_MaxIterationsHandler, + ) + from litellm.proxy.utils import InternalUsageCache + from litellm.types.agents import AgentResponse + + cache = DualCache() + handler = _PROXY_MaxIterationsHandler( + internal_usage_cache=InternalUsageCache(cache), + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-iter", + agent_id="agent-iter-1", + ) + agent = AgentResponse( + agent_id="agent-iter-1", + agent_name="iter-agent", + litellm_params={"max_iterations": 1}, + agent_card_params={"name": "iter-agent", "version": "1.0.0"}, + ) + with patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" + ) as mock_registry: + mock_registry.get_agent_by_id.return_value = agent + # First call within budget. + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={"metadata": {"session_id": "sess-1"}}, + call_type="", + ) + # Second call exceeds — must raise the unified class. + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=cache, + data={"metadata": {"session_id": "sess-1"}}, + call_type="", + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + assert isinstance(e, RateLimitError) + assert isinstance(e, HTTPException) + + @pytest.mark.asyncio + async def test_max_budget_limiter_raises_proxy_rate_limit_error(self): + """ + Drive `_PROXY_MaxBudgetLimiter` past the user budget and assert it + raises the unified class. Mocks `get_current_spend` so we don't need + the proxy DB. + """ + from unittest.mock import patch + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.max_budget_limiter import ( + _PROXY_MaxBudgetLimiter, + ) + + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-budget", + user_id="user-budget-1", + user_max_budget=1.0, + user_spend=2.0, + ) + with patch( + "litellm.proxy.proxy_server.get_current_spend", + return_value=5.0, + ): + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + assert "max budget" in str(e.detail).lower() + + @pytest.mark.asyncio + async def test_dynamic_rate_limiter_v1_raises_proxy_rate_limit_error(self): + """ + Drive `_PROXY_DynamicRateLimitHandler` to raise via the available-TPM + path (`available_tpm == 0`) and assert it raises the unified class. + Mocks `check_available_usage` so we don't need a real router. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.dynamic_rate_limiter import ( + _PROXY_DynamicRateLimitHandler, + ) + + handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) + # check_available_usage returns (available_tpm, available_rpm, + # model_tpm, model_rpm, active_projects). Setting available_tpm == 0 + # forces the TPM-exceeded raise. + handler.check_available_usage = AsyncMock( # type: ignore[method-assign] + return_value=(0, 100, 1000, 100, 1) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-dyn", + metadata={"priority": "default"}, + ) + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={"model": "gpt-4"}, + call_type="completion", + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + assert isinstance(e.detail, dict) + assert "TPM" in e.detail.get("error", "") + + @pytest.mark.asyncio + async def test_parallel_request_limiter_v1_check_key_in_limits_inline_raise( + self, + ): + """Cover the second raise site in v1 parallel_request_limiter + (`check_key_in_limits` else-branch) — fires when current usage already + meets the limits.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + cache = MagicMock() + cache.async_batch_set_cache = AsyncMock(return_value=None) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.check_key_in_limits( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), + cache=DualCache(), + data={}, + call_type="completion", + max_parallel_requests=1, + tpm_limit=10, + rpm_limit=10, + # current already at the limit on every dimension → forces + # the inline `raise ProxyRateLimitError(...)` else-branch. + current={"current_requests": 1, "current_tpm": 10, "current_rpm": 10}, + request_count_api_key="x", + rate_limit_type="key", + values_to_update_in_cache=[], + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + + @pytest.mark.parametrize( + "current,limits,expected_type", + [ + # current already at concurrent-request cap → CONCURRENT_REQUESTS + ( + {"current_requests": 5, "current_tpm": 0, "current_rpm": 0}, + {"max_parallel_requests": 5, "tpm_limit": 100, "rpm_limit": 100}, + "concurrent_requests", + ), + # current already at TPM cap (concurrent has headroom) → TOKENS + ( + {"current_requests": 0, "current_tpm": 100, "current_rpm": 0}, + {"max_parallel_requests": 5, "tpm_limit": 100, "rpm_limit": 100}, + "tokens", + ), + # current already at RPM cap (concurrent + TPM have headroom) → + # REQUESTS (the fall-through branch). + ( + {"current_requests": 0, "current_tpm": 0, "current_rpm": 100}, + {"max_parallel_requests": 5, "tpm_limit": 100, "rpm_limit": 100}, + "requests", + ), + ], + ) + @pytest.mark.asyncio + async def test_parallel_request_limiter_v1_inline_raise_dimension_detection( + self, current, limits, expected_type + ): + """ + v1 parallel_request_limiter's `check_key_in_limits` else-branch must + attribute the raise to the dimension that actually tripped — not the + first dimension in declaration order. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + cache = MagicMock() + cache.async_batch_set_cache = AsyncMock(return_value=None) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.check_key_in_limits( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), + cache=DualCache(), + data={}, + call_type="completion", + max_parallel_requests=limits["max_parallel_requests"], + tpm_limit=limits["tpm_limit"], + rpm_limit=limits["rpm_limit"], + current=current, + request_count_api_key="x", + rate_limit_type="key", + values_to_update_in_cache=[], + ) + assert exc_info.value.rate_limit_type == expected_type + + @pytest.mark.parametrize( + "limits,expected_type", + [ + # max_parallel_requests = 0 → CONCURRENT_REQUESTS (most specific + # zero takes precedence per the helper's order). + ( + {"max_parallel_requests": 0, "tpm_limit": 0, "rpm_limit": 0}, + "concurrent_requests", + ), + # tpm_limit = 0 (concurrent has a positive limit) → TOKENS + ( + {"max_parallel_requests": 5, "tpm_limit": 0, "rpm_limit": 0}, + "tokens", + ), + # only rpm_limit = 0 → REQUESTS (fall-through) + ( + {"max_parallel_requests": 5, "tpm_limit": 100, "rpm_limit": 0}, + "requests", + ), + ], + ) + @pytest.mark.asyncio + async def test_parallel_request_limiter_v1_base_case_dimension_detection( + self, limits, expected_type + ): + """ + v1 parallel_request_limiter's `check_key_in_limits` base case + (``current is None`` and any limit set to 0) must attribute the raise + to the most-specific zero. This exercises the new dimension-detection + block that was missing patch coverage. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + cache = MagicMock() + cache.async_batch_set_cache = AsyncMock(return_value=None) + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=cache) + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.check_key_in_limits( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-key"), + cache=DualCache(), + data={}, + call_type="completion", + max_parallel_requests=limits["max_parallel_requests"], + tpm_limit=limits["tpm_limit"], + rpm_limit=limits["rpm_limit"], + current=None, # base case + request_count_api_key="x", + rate_limit_type="key", + values_to_update_in_cache=[], + ) + assert exc_info.value.rate_limit_type == expected_type + + @pytest.mark.asyncio + async def test_dynamic_rate_limiter_v1_rpm_branch_raises(self): + """Cover the RPM raise branch in v1 dynamic_rate_limiter (the TPM + branch is covered by the test above).""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.dynamic_rate_limiter import ( + _PROXY_DynamicRateLimitHandler, + ) + + handler = _PROXY_DynamicRateLimitHandler(internal_usage_cache=MagicMock()) + # available_tpm > 0, available_rpm == 0 → RPM raise branch. + handler.check_available_usage = AsyncMock( # type: ignore[method-assign] + return_value=(100, 0, 1000, 100, 1) + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-dyn-rpm", + metadata={"priority": "default"}, + ) + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={"model": "gpt-4"}, + call_type="completion", + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + assert isinstance(e.detail, dict) + assert "RPM" in e.detail.get("error", "") + + @pytest.mark.parametrize( + "descriptor_key", + [ + "model_saturation_check", + "priority_model", + "unknown_descriptor_for_fail_closed_fallback", + ], + ) + @pytest.mark.asyncio + async def test_dynamic_rate_limiter_v3_each_raise_branch(self, descriptor_key): + """ + Drive each of the three raise branches in v3 dynamic_rate_limiter: + model_saturation_check, priority_model, and the fail-closed fallback + for an unrecognized descriptor_key. Mocks + ``atomic_check_and_increment_by_n`` so the v3 limiter's response + directly drives the raise-site selection. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.dynamic_rate_limiter_v3 import ( + _PROXY_DynamicRateLimitHandlerV3, + ) + + # Bypass __init__ — we want to inject a stub v3_limiter without + # paying for the full handler setup. + handler = _PROXY_DynamicRateLimitHandlerV3.__new__( + _PROXY_DynamicRateLimitHandlerV3 + ) + v3_limiter = MagicMock() + v3_limiter.window_size = 60 + v3_limiter.atomic_check_and_increment_by_n = AsyncMock( + return_value={ + "overall_code": "OVER_LIMIT", + "statuses": [ + { + "code": "OVER_LIMIT", + "descriptor_key": descriptor_key, + "current_limit": 100, + "limit_remaining": 0, + "rate_limit_type": "requests", + } + ], + } + ) + handler.v3_limiter = v3_limiter + # Stub the descriptor builders so we don't pull in real router state. + handler._create_model_tracking_descriptor = MagicMock( # type: ignore[method-assign] + return_value={ + "key": descriptor_key, + "value": "v", + "rate_limit": { + "requests_per_unit": 100, + "tokens_per_unit": None, + "window_size": 60, + }, + } + ) + handler._create_priority_based_descriptors = MagicMock( # type: ignore[method-assign] + return_value=[] + ) + model_group_info = MagicMock() + model_group_info.tpm = 1000 + model_group_info.rpm = 100 + + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler._check_rate_limits( + model="gpt-4", + model_group_info=model_group_info, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test-v3"), + priority="default", + saturation=0.99, + data={}, + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + + @pytest.mark.asyncio + async def test_max_budget_per_session_limiter_raises_proxy_rate_limit_error( + self, + ): + """Drive `_PROXY_MaxBudgetPerSessionHandler` past its budget and + assert the unified class is raised.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.hooks.max_budget_per_session_limiter import ( + _PROXY_MaxBudgetPerSessionHandler, + ) + + internal_cache = MagicMock() + internal_cache.async_get_cache = AsyncMock(return_value=10.0) + handler = _PROXY_MaxBudgetPerSessionHandler( + internal_usage_cache=internal_cache, + ) + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test-session", + agent_id="agent-session-1", + ) + agent = MagicMock() + agent.litellm_params = {"max_budget_per_session": 1.0} + with patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry" + ) as mock_registry: + mock_registry.get_agent_by_id.return_value = agent + with pytest.raises(ProxyRateLimitError) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={"metadata": {"session_id": "session-over-budget"}}, + call_type="completion", + ) + e = exc_info.value + assert e.status_code == 429 + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + assert "session" in str(e.detail).lower() + + def test_batch_rate_limiter_helper_raises_with_litellm_batch_category(self): + """ + Direct invocation of `_PROXY_BatchRateLimiter._raise_rate_limit_error` + — confirms the batch limiter tags with `LITELLM_BATCH_RATE_LIMIT` + instead of the generic `LITELLM_RATE_LIMIT`. + """ + from unittest.mock import MagicMock + + from litellm.proxy.hooks.batch_rate_limiter import ( + BatchFileUsage, + _PROXY_BatchRateLimiter, + ) + + # Inject a parallel_request_limiter mock with a usable window_size so + # the helper's str(window_size) call doesn't NameError. + parallel_limiter = MagicMock() + parallel_limiter.window_size = 60 + handler = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=parallel_limiter, + ) + status = { + "code": "OVER_LIMIT", + "descriptor_key": "key", + "current_limit": 100, + "limit_remaining": 0, + "rate_limit_type": "requests", + } + descriptors = [ + { + "key": "key", + "value": "sk-batch", + "rate_limit": { + "requests_per_unit": 100, + "tokens_per_unit": None, + "window_size": 60, + }, + } + ] + with pytest.raises(ProxyRateLimitError) as exc_info: + handler._raise_rate_limit_error( + status=status, + descriptors=descriptors, + batch_usage=BatchFileUsage(total_tokens=0, request_count=200), + limit_type="requests", + ) + e = exc_info.value + assert e.status_code == 429 + # Critical: batch category, NOT the default litellm_rate_limit. + assert e.category == RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT + assert isinstance(e, RateLimitError) + assert isinstance(e, HTTPException) + + +class TestRateLimitType: + """ + Tests for the orthogonal `rate_limit_type` dimension introduced as a + follow-up to LIT-2968 (trho's last ask in the Slack thread). + + `category` answers *who* rate-limited (vendor vs. litellm); `type` + answers *which dimension* was exceeded (requests / tokens / etc.). + Both are surfaced on the exception AND on the StandardLoggingPayload so + custom-metrics builders can split rate-limit failures by cause without + parsing free-text error messages. + """ + + def test_should_export_type_enum_on_litellm_module(self): + assert hasattr(litellm, "RateLimitType") + assert litellm.RateLimitType is RateLimitType + + def test_should_define_all_documented_types(self): + assert RateLimitType.REQUESTS == "requests" + assert RateLimitType.TOKENS == "tokens" + assert RateLimitType.CONCURRENT_REQUESTS == "concurrent_requests" + assert RateLimitType.BUDGET == "budget" + assert RateLimitType.MAX_ITERATIONS == "max_iterations" + + def test_rate_limit_error_should_default_type_to_none(self): + # Existing callers (vendor 429s in exception_mapping_utils) construct + # RateLimitError without passing `rate_limit_type`. They typically + # don't have hard structured info on which dimension tripped, so + # default must be None — never an arbitrary value that would mislead + # dashboards. + e = RateLimitError(message="oops", llm_provider="openai", model="gpt-4") + assert e.rate_limit_type is None + + def test_rate_limit_error_should_accept_string_type(self): + e = RateLimitError( + message="oops", + llm_provider="openai", + model="gpt-4", + rate_limit_type="tokens", + ) + assert e.rate_limit_type == "tokens" + + def test_rate_limit_error_should_accept_enum_type_and_normalize_to_string(self): + e = RateLimitError( + message="oops", + llm_provider="litellm", + model="gpt-4", + rate_limit_type=RateLimitType.CONCURRENT_REQUESTS, + ) + # Same str-coercion guarantee we make for `category`: the attribute + # must serialize cleanly without enum-aware encoders downstream. + assert e.rate_limit_type == "concurrent_requests" + assert isinstance(e.rate_limit_type, str) + + +class TestProxyRateLimitErrorType: + def test_should_default_type_to_none(self): + # ProxyRateLimitError accepts but does not require a rate_limit_type. + # Callers that don't pass one (e.g. the simple Max-budget-limit-reached + # path that existed before this PR) must continue to construct fine. + e = ProxyRateLimitError(detail="over limit") + assert e.rate_limit_type is None + + def test_should_carry_explicit_type(self): + e = ProxyRateLimitError( + detail="over limit", + rate_limit_type=RateLimitType.TOKENS, + ) + assert e.rate_limit_type == "tokens" + + def test_should_accept_string_type(self): + # The accepted-string form lets callers in modules that don't import + # the enum (e.g. v3 limiter passing through descriptor strings) + # forward the raw value. + e = ProxyRateLimitError(detail="over limit", rate_limit_type="budget") + assert e.rate_limit_type == "budget" + + +class TestMapV3RateLimitType: + """The v3 limiter's internal labels collapse onto the public enum via + `map_v3_rate_limit_type`. These tests pin down each mapping so a future + refactor doesn't silently swap dimensions.""" + + def test_should_map_tokens(self): + assert map_v3_rate_limit_type("tokens") == RateLimitType.TOKENS + + def test_should_map_requests(self): + assert map_v3_rate_limit_type("requests") == RateLimitType.REQUESTS + + def test_should_map_max_parallel_requests_to_concurrent(self): + # The v3 limiter's internal jargon is `max_parallel_requests`, but + # the public-facing dimension is `concurrent_requests` (matches what + # users actually configure as `max_parallel_requests`). The mapping + # must collapse these so dashboards see one name, not two. + assert ( + map_v3_rate_limit_type("max_parallel_requests") + == RateLimitType.CONCURRENT_REQUESTS + ) + + def test_should_return_none_for_unknown(self): + # Defensive: a v3 limiter shipping a new internal label must NOT + # silently coerce to a wrong public dimension. Returning None lets + # the caller decide (typically: omit the field). + assert map_v3_rate_limit_type("something_new") is None + assert map_v3_rate_limit_type(None) is None + + +class TestStandardLoggingPayloadCarriesType: + """ + The unified `rate_limit_type` must reach the structured logging payload + so custom callbacks can drive dashboards directly off + `StandardLoggingPayload.error_information.error_rate_limit_type`. + """ + + def test_should_propagate_type_for_proxy_rate_limit_error(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = ProxyRateLimitError( + detail="over tpm", + rate_limit_type=RateLimitType.TOKENS, + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["error_rate_limit_type"] == "tokens" + + def test_should_propagate_type_for_plain_rate_limit_error(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = RateLimitError( + message="vendor 429", + llm_provider="openai", + model="gpt-4", + rate_limit_type=RateLimitType.REQUESTS, + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["error_rate_limit_type"] == "requests" + + def test_should_be_none_when_unspecified(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + # Vendor 429 exception with no header hints → type omitted. + e = RateLimitError( + message="vendor 429", + llm_provider="openai", + model="gpt-4", + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["error_rate_limit_type"] is None + + def test_should_be_none_for_non_rate_limit_errors(self): + # Symmetry with `error_rate_limit_category`: the field must be + # present on every payload so consumers can read it + # unconditionally, but None for non-rate-limit exceptions. + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + info = StandardLoggingPayloadSetup.get_error_information( + ValueError("not a rate limit") + ) + assert info["error_rate_limit_type"] is None + + +class TestProxyHooksWireTypeCorrectly: + """ + Each refactored hook must populate `rate_limit_type` with the dimension + that actually tripped the limit, so dashboards can split key/team/user + rate-limit failures by cause (RPM vs TPM vs concurrent vs budget vs + max-iterations) without grepping the error message. + """ + + def test_max_budget_limiter_emits_budget_type(self): + e = ProxyRateLimitError( + detail="Max budget limit reached.", + rate_limit_type=RateLimitType.BUDGET, + ) + assert e.category == "litellm_rate_limit" + assert e.rate_limit_type == "budget" + + def test_max_iterations_limiter_emits_max_iterations_type(self): + e = ProxyRateLimitError( + detail="Max iterations exceeded for session abc.", + rate_limit_type=RateLimitType.MAX_ITERATIONS, + ) + assert e.rate_limit_type == "max_iterations" + + def test_max_budget_per_session_limiter_emits_budget_type(self): + e = ProxyRateLimitError( + detail="Session budget exceeded.", + rate_limit_type=RateLimitType.BUDGET, + ) + assert e.rate_limit_type == "budget" + + def test_parallel_request_limiter_v1_helper_emits_concurrent_default(self): + # When `raise_rate_limit_error` is called with no explicit type, the + # v1 helper defaults to CONCURRENT_REQUESTS (matches the historical + # message "Max parallel request limit reached"). Tests below cover + # the explicit-type override paths. + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + with pytest.raises(ProxyRateLimitError) as exc_info: + handler.raise_rate_limit_error() + assert exc_info.value.rate_limit_type == "concurrent_requests" + + def test_parallel_request_limiter_v1_helper_accepts_explicit_type(self): + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, + ) + + handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=MagicMock()) + with pytest.raises(ProxyRateLimitError) as exc_info: + handler.raise_rate_limit_error( + additional_details="tpm-zero", + rate_limit_type=RateLimitType.TOKENS, + ) + assert exc_info.value.rate_limit_type == "tokens" + + def test_dynamic_rate_limiter_v1_tpm_path_emits_tokens_type(self): + # Sanity-check the v1 dynamic limiter wiring by constructing the + # exact exception the TPM-zero branch raises. We round-trip through + # ProxyRateLimitError to assert both fields. (Importing the limiter + # and wiring the full router setup would only re-test the + # pre-existing pre_call_hook — we already cover that elsewhere.) + e = ProxyRateLimitError( + detail={"error": "Key=k over available TPM=0."}, + rate_limit_type=RateLimitType.TOKENS, + model="gpt-4", + ) + assert e.rate_limit_type == "tokens" + assert e.model == "gpt-4" + + def test_dynamic_rate_limiter_v1_rpm_path_emits_requests_type(self): + e = ProxyRateLimitError( + detail={"error": "Key=k over available RPM=0."}, + rate_limit_type=RateLimitType.REQUESTS, + model="gpt-4", + ) + assert e.rate_limit_type == "requests" + + @pytest.mark.asyncio + async def test_v3_limiter_handle_rate_limit_error_propagates_type(self): + """ + End-to-end: feed the v3 limiter's `_handle_rate_limit_error` an + OVER_LIMIT response and verify the raised ProxyRateLimitError carries + the mapped public RateLimitType. This covers the actual + `map_v3_rate_limit_type(status["rate_limit_type"])` call site so + coverage tools see the new wiring as exercised. + """ + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + + handler = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=MagicMock(), + ) + # Minimal RateLimitResponse + descriptors shape that the handler + # reads. We only need one OVER_LIMIT status to drive the raise. + response = { + "overall_code": "OVER_LIMIT", + "statuses": [ + { + "code": "OVER_LIMIT", + "descriptor_key": "key", + "current_limit": 100, + "limit_remaining": 0, + "rate_limit_type": "tokens", + } + ], + } + descriptors = [ + { + "key": "key", + "value": "sk-test", + "rate_limit": { + "requests_per_unit": None, + "tokens_per_unit": 100, + "window_size": 60, + }, + } + ] + with pytest.raises(ProxyRateLimitError) as exc_info: + handler._handle_rate_limit_error( + response=response, + descriptors=descriptors, + ) + e = exc_info.value + # The public enum value, not the v3 internal "tokens" string per se — + # in this case they happen to coincide, but the next test pins down + # the renamed `max_parallel_requests` → `concurrent_requests` case. + assert e.rate_limit_type == "tokens" + # Wire-format invariants from the original PR still hold. + assert e.headers is not None + assert e.headers.get("rate_limit_type") == "tokens" + assert e.headers.get("retry-after") is not None + + @pytest.mark.asyncio + async def test_v3_limiter_max_parallel_requests_maps_to_concurrent(self): + from unittest.mock import MagicMock + + from litellm.proxy.hooks.parallel_request_limiter_v3 import ( + _PROXY_MaxParallelRequestsHandler_v3, + ) + + handler = _PROXY_MaxParallelRequestsHandler_v3( + internal_usage_cache=MagicMock(), + ) + response = { + "overall_code": "OVER_LIMIT", + "statuses": [ + { + "code": "OVER_LIMIT", + "descriptor_key": "key", + "current_limit": 5, + "limit_remaining": 0, + # v3 internal jargon — must collapse to the public name. + "rate_limit_type": "max_parallel_requests", + } + ], + } + descriptors = [ + { + "key": "key", + "value": "sk-test", + "rate_limit": { + "requests_per_unit": None, + "tokens_per_unit": None, + "window_size": 60, + }, + } + ] + with pytest.raises(ProxyRateLimitError) as exc_info: + handler._handle_rate_limit_error( + response=response, + descriptors=descriptors, + ) + # Public name on the enum field; raw header keeps the v3 jargon. + assert exc_info.value.rate_limit_type == "concurrent_requests" + assert exc_info.value.headers["rate_limit_type"] == "max_parallel_requests" + + def test_batch_rate_limiter_emits_tokens_type_for_tpm_violation(self): + from unittest.mock import MagicMock + + from litellm.proxy.hooks.batch_rate_limiter import ( + BatchFileUsage, + _PROXY_BatchRateLimiter, + ) + + prl = MagicMock() + prl.window_size = 60 + handler = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=prl, + ) + status = { + "code": "OVER_LIMIT", + "descriptor_key": "key", + "current_limit": 1000, + "limit_remaining": 100, + "rate_limit_type": "tokens", + } + descriptors = [ + { + "key": "key", + "value": "sk-test", + "rate_limit": { + "requests_per_unit": None, + "tokens_per_unit": 1000, + "window_size": 60, + }, + } + ] + with pytest.raises(ProxyRateLimitError) as exc_info: + handler._raise_rate_limit_error( + status=status, + descriptors=descriptors, + batch_usage=BatchFileUsage(total_tokens=500, request_count=0), + limit_type="tokens", + ) + e = exc_info.value + assert e.rate_limit_type == "tokens" + assert e.category == RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT + + def test_batch_rate_limiter_emits_requests_type_for_rpm_violation(self): + from unittest.mock import MagicMock + + from litellm.proxy.hooks.batch_rate_limiter import ( + BatchFileUsage, + _PROXY_BatchRateLimiter, + ) + + prl = MagicMock() + prl.window_size = 60 + handler = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=prl, + ) + status = { + "code": "OVER_LIMIT", + "descriptor_key": "key", + "current_limit": 100, + "limit_remaining": 10, + "rate_limit_type": "requests", + } + descriptors = [ + { + "key": "key", + "value": "sk-test", + "rate_limit": { + "requests_per_unit": 100, + "tokens_per_unit": None, + "window_size": 60, + }, + } + ] + with pytest.raises(ProxyRateLimitError) as exc_info: + handler._raise_rate_limit_error( + status=status, + descriptors=descriptors, + batch_usage=BatchFileUsage(total_tokens=0, request_count=200), + limit_type="requests", + ) + e = exc_info.value + assert e.rate_limit_type == "requests" + assert e.category == RateLimitErrorCategory.LITELLM_BATCH_RATE_LIMIT + + +class TestBudgetExceededErrorSurfacesUnifiedFields: + """ + The hot path for virtual-key / team / org / end-user max_budget caps + raises :class:`litellm.BudgetExceededError`, which historically had no + relationship to :class:`RateLimitError` and therefore left the unified + `error_rate_limit_category` / `error_rate_limit_type` fields empty. + Test 2 of the QA pass surfaced this gap; this class pins the fix. + + The fix is intentionally additive: `BudgetExceededError` keeps its + bare-`Exception` base class (so existing `except BudgetExceededError:` + handlers keep working) and just sets the same `category` / + `rate_limit_type` attributes that the rest of the unified rate-limit + path reads (normalized to plain strings, matching how + `RateLimitError.__init__` stores its own values). Duck-typed dispatch + in `get_error_information` picks them up automatically. + """ + + def test_should_carry_litellm_rate_limit_category(self): + e = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + # Stored as the plain string value (matches RateLimitError behavior), + # but equality with the enum still works because the enum subclasses + # str. + assert e.category == "litellm_rate_limit" + assert e.category == RateLimitErrorCategory.LITELLM_RATE_LIMIT + + def test_should_carry_budget_rate_limit_type(self): + e = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + assert e.rate_limit_type == "budget" + assert e.rate_limit_type == RateLimitType.BUDGET + + def test_should_default_llm_provider_to_empty_string(self): + # `llm_provider` is read off the exception in `get_error_information` + # — it must always be a string so the StandardLoggingPayload field + # stays serializable. Default to "" when no caller passes one. + e = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + assert e.llm_provider == "" + + def test_should_accept_llm_provider_kwarg(self): + # Callers that have the resolved provider in scope (e.g. the + # auth-checks budget enforcement paths) can thread it through. + e = litellm.BudgetExceededError( + current_cost=0.5, max_budget=0.1, llm_provider="anthropic" + ) + assert e.llm_provider == "anthropic" + + def test_should_keep_existing_status_code_and_message(self): + # Backward-compat guard: existing callers depend on `status_code=429` + # and the canonical message format. + e = litellm.BudgetExceededError(current_cost=0.000109, max_budget=0.0001) + assert e.status_code == 429 + assert "Current cost: 0.000109" in e.message + assert "Max budget: 0.0001" in e.message + + def test_should_still_be_catchable_as_exception_not_rate_limit_error(self): + # Critical: we deliberately did NOT make BudgetExceededError a + # RateLimitError subclass. Existing `except BudgetExceededError:` + # handlers must keep catching it, and `except RateLimitError:` + # handlers must NOT start catching it (which would surprise callers + # who rely on the two being distinct). + e = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + assert isinstance(e, Exception) + assert isinstance(e, litellm.BudgetExceededError) + assert not isinstance(e, RateLimitError) + + def test_should_propagate_category_to_standard_logging_payload(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["error_rate_limit_category"] == "litellm_rate_limit" + assert info["error_rate_limit_type"] == "budget" + assert info["error_code"] == "429" + assert info["error_class"] == "BudgetExceededError" + + def test_should_propagate_llm_provider_to_standard_logging_payload(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + e = litellm.BudgetExceededError( + current_cost=0.5, max_budget=0.1, llm_provider="bedrock" + ) + info = StandardLoggingPayloadSetup.get_error_information(e) + assert info["llm_provider"] == "bedrock" + + +class TestThirdPartyAttrLeakageGuard: + """ + The duck-typed read at the StandardLoggingPayload + Prometheus surfaces + must reject `.category` / `.rate_limit_type` strings set on unrelated + third-party exceptions. Without validation, a foreign exception that + happens to declare either attribute name would leak garbage values into + custom-callback payloads and Prometheus label cardinality. + """ + + def test_should_drop_unknown_category_string_on_third_party_exception(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + class Foreign(Exception): + category = "totally_not_a_real_category" + + info = StandardLoggingPayloadSetup.get_error_information(Foreign("boom")) + assert info["error_rate_limit_category"] is None + + def test_should_drop_unknown_rate_limit_type_string_on_third_party_exception(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + class Foreign(Exception): + rate_limit_type = "wat" + + info = StandardLoggingPayloadSetup.get_error_information(Foreign("boom")) + assert info["error_rate_limit_type"] is None + + def test_should_drop_non_string_garbage_attrs(self): + from litellm.litellm_core_utils.litellm_logging import ( + StandardLoggingPayloadSetup, + ) + + class Foreign(Exception): + category = 42 + rate_limit_type = {"lol": "no"} + + info = StandardLoggingPayloadSetup.get_error_information(Foreign()) + assert info["error_rate_limit_category"] is None + assert info["error_rate_limit_type"] is None + + def test_should_drop_garbage_on_prometheus_label_extraction(self): + from litellm.integrations.prometheus import PrometheusLogger + + class Foreign(Exception): + category = "spam" + rate_limit_type = "spam" + + category, rate_limit_type = PrometheusLogger._extract_rate_limit_labels( + Foreign() + ) + assert category is None + assert rate_limit_type is None + + def test_should_still_accept_legitimate_rate_limit_categories(self): + # The guard must not over-correct — every documented enum value + # is a valid string and must pass through. + from litellm.exceptions import ( + validate_rate_limit_category, + validate_rate_limit_type, + ) + + for member in RateLimitErrorCategory: + assert validate_rate_limit_category(member.value) == member.value + assert validate_rate_limit_category(member) == member.value + + for member in RateLimitType: + assert validate_rate_limit_type(member.value) == member.value + assert validate_rate_limit_type(member) == member.value + + +@pytest.mark.asyncio +class TestBudgetExceededErrorLlmProviderEnrichment: + """ + BudgetExceededError raise sites in auth_checks.py are tenant-scoped + (key / team / org / tag) and cannot see the request model. To still + populate `llm_provider` on the StandardLoggingPayload — which is what + custom-callback consumers attribute spend to — the central + UserAPIKeyAuthExceptionHandler enriches the exception from + `request_data["model"]` before post_call_failure_hook fires. + """ + + async def _run_handler_and_capture_exception_seen_by_callback( + self, exception: Exception, request_data: dict + ): + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.proxy.auth.auth_exception_handler import ( + UserAPIKeyAuthExceptionHandler, + ) + + captured: dict = {} + + async def fake_post_call_failure_hook(**kwargs): + captured["exception"] = kwargs["original_exception"] + return None + + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + MagicMock( + post_call_failure_hook=AsyncMock( + side_effect=fake_post_call_failure_hook + ) + ), + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"use_x_forwarded_for": False}, + ), + patch( + "litellm.proxy.auth.auth_exception_handler._get_request_ip_address", + return_value="127.0.0.1", + ), + ): + try: + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + e=exception, + request=MagicMock(), + request_data=request_data, + route="/v1/chat/completions", + parent_otel_span=None, + api_key="sk-test", + ) + except Exception: + pass + return captured.get("exception") + + async def test_should_resolve_llm_provider_from_request_data_when_unset(self): + err = litellm.BudgetExceededError(current_cost=100, max_budget=10) + assert err.llm_provider == "" + seen = await self._run_handler_and_capture_exception_seen_by_callback( + err, {"model": "openai/gpt-4o-mini"} + ) + assert seen is not None + assert seen.llm_provider == "openai" + + async def test_should_not_overwrite_llm_provider_when_caller_set_it(self): + err = litellm.BudgetExceededError( + current_cost=100, max_budget=10, llm_provider="anthropic" + ) + seen = await self._run_handler_and_capture_exception_seen_by_callback( + err, {"model": "openai/gpt-4o-mini"} + ) + assert seen.llm_provider == "anthropic" + + async def test_should_fall_back_to_litellm_proxy_when_model_missing(self): + err = litellm.BudgetExceededError(current_cost=100, max_budget=10) + seen = await self._run_handler_and_capture_exception_seen_by_callback(err, {}) + assert seen.llm_provider == "litellm_proxy" + + async def test_should_not_enrich_non_budget_exceptions(self): + err = ValueError("unrelated") + seen = await self._run_handler_and_capture_exception_seen_by_callback( + err, {"model": "openai/gpt-4o-mini"} + ) + assert not hasattr(seen, "llm_provider") or seen.llm_provider != "openai" diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 80f4d0b57f1..d303734dffd 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2822,6 +2822,7 @@ export interface Member { max_budget_in_team?: number | null; tpm_limit?: number | null; rpm_limit?: number | null; + budget_duration?: string | null; allowed_models?: string[] | null; } @@ -2949,18 +2950,21 @@ export const teamMemberUpdateCall = async ( user_id: formValues.user_id, }; - // Add optional budget and rate limit fields + const orNull = (value: unknown) => (value === undefined || value === null || value === "" ? null : value); if (formValues.user_email !== undefined) { requestBody.user_email = formValues.user_email; } - if (formValues.max_budget_in_team !== undefined && formValues.max_budget_in_team !== null) { - requestBody.max_budget_in_team = formValues.max_budget_in_team; + if ("max_budget_in_team" in formValues) { + requestBody.max_budget_in_team = orNull(formValues.max_budget_in_team); } - if (formValues.tpm_limit !== undefined && formValues.tpm_limit !== null) { - requestBody.tpm_limit = formValues.tpm_limit; + if ("tpm_limit" in formValues) { + requestBody.tpm_limit = orNull(formValues.tpm_limit); } - if (formValues.rpm_limit !== undefined && formValues.rpm_limit !== null) { - requestBody.rpm_limit = formValues.rpm_limit; + if ("rpm_limit" in formValues) { + requestBody.rpm_limit = orNull(formValues.rpm_limit); + } + if ("budget_duration" in formValues) { + requestBody.budget_duration = orNull(formValues.budget_duration); } if (formValues.allowed_models !== undefined) { requestBody.allowed_models = formValues.allowed_models; diff --git a/ui/litellm-dashboard/src/components/team/EditMembership.tsx b/ui/litellm-dashboard/src/components/team/EditMembership.tsx index af5f8631ae4..16e4ec58dd0 100644 --- a/ui/litellm-dashboard/src/components/team/EditMembership.tsx +++ b/ui/litellm-dashboard/src/components/team/EditMembership.tsx @@ -2,6 +2,7 @@ import { Text, TextInput } from "@tremor/react"; import { Button as AntButton, Form, Modal, Select } from "antd"; import React, { useEffect, useState } from "react"; import NumericalInput from "../shared/numerical_input"; +import BudgetDurationDropdown from "../common_components/budget_duration_dropdown"; interface BaseMember { user_email?: string; @@ -21,7 +22,7 @@ interface ModalConfig { additionalFields?: Array<{ name: string; label: string | React.ReactNode; - type: "input" | "select" | "numerical" | "multi-select"; + type: "input" | "select" | "numerical" | "multi-select" | "budget-duration"; options?: Array<{ label: string; value: string }>; rules?: any[]; step?: number; @@ -65,6 +66,7 @@ const MemberModal = ({ max_budget_in_team: (initialData as any).max_budget_in_team || null, tpm_limit: (initialData as any).tpm_limit || null, rpm_limit: (initialData as any).rpm_limit || null, + budget_duration: (initialData as any).budget_duration || null, // Keep array values for multi-select fields allowed_models: (initialData as any).allowed_models || [], }; @@ -117,7 +119,7 @@ const MemberModal = ({ const renderField = (field: { name: string; label: string | React.ReactNode; - type: "input" | "select" | "numerical" | "multi-select"; + type: "input" | "select" | "numerical" | "multi-select" | "budget-duration"; options?: Array<{ label: string; value: string }>; rules?: any[]; step?: number; @@ -155,6 +157,8 @@ const MemberModal = ({ allowClear /> ); + case "budget-duration": + return ; default: return null; } diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 4943b1fbac2..602490a0c98 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -388,6 +388,7 @@ const TeamInfoView: React.FC = ({ max_budget_in_team: values.max_budget_in_team, tpm_limit: values.tpm_limit, rpm_limit: values.rpm_limit, + budget_duration: values.budget_duration, allowed_models: values.allowed_models, }; MessageManager.destroy(); // Remove all existing toasts @@ -1689,6 +1690,18 @@ const TeamInfoView: React.FC = ({ min: 0, placeholder: "Budget limit for this member within this team", }, + { + name: "budget_duration", + label: ( + + Budget Reset Period{" "} + + + + + ), + type: "budget-duration" as const, + }, { name: "tpm_limit", label: ( diff --git a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx index 5a290a6f6d4..e2f108dcbf5 100644 --- a/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamMemberTab.tsx @@ -210,6 +210,7 @@ export default function TeamMemberTab({ max_budget_in_team: membership?.litellm_budget_table?.max_budget || null, tpm_limit: membership?.litellm_budget_table?.tpm_limit || null, rpm_limit: membership?.litellm_budget_table?.rpm_limit || null, + budget_duration: membership?.litellm_budget_table?.budget_duration || null, allowed_models: membership?.litellm_budget_table?.allowed_models || [], }; setSelectedEditMember(enhancedMember); diff --git a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx index aed194a2972..6f1c4126282 100644 --- a/ui/litellm-dashboard/src/components/view_logs/index.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/index.test.tsx @@ -1,8 +1,11 @@ import { screen, waitFor } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; +import moment from "moment"; import { beforeEach, describe, expect, it, vi } from "vitest"; import SpendLogsTable from "./index"; import { renderWithProviders } from "../../../tests/test-utils"; +import { uiSpendLogsCall } from "../networking"; +import { useLogFilterLogic } from "./log_filter_logic"; const mockHandleFilterResetFromHook = vi.fn(); vi.mock("./log_filter_logic", async (importOriginal) => { @@ -115,4 +118,63 @@ describe("SpendLogsTable", () => { expect(screen.getByRole("button", { name: "Reset Filters" })).toBeInTheDocument(); }); }); + + describe("Quick Select time range", () => { + // uiSpendLogsCall fires from the real useLogFilterLogic query, so restore it here. + beforeEach(async () => { + const actual = await vi.importActual("./log_filter_logic"); + vi.mocked(useLogFilterLogic).mockImplementation(actual.useLogFilterLogic); + }); + + const waitForWindowSeconds = async (minMinutes: number) => { + let diff = -1; + await waitFor(() => { + const lastCall = vi.mocked(uiSpendLogsCall).mock.calls.at(-1)?.[0]; + if (!lastCall) throw new Error("uiSpendLogsCall was not called"); + diff = moment + .utc(lastCall.end_date, "YYYY-MM-DD HH:mm:ss") + .diff(moment.utc(lastCall.start_date, "YYYY-MM-DD HH:mm:ss"), "seconds"); + // start_date is rounded down to the minute boundary, end_date is the + // current wall-clock at queryFn time. The dropped sub-minute fraction + // on start_date can push the diff up to (minMinutes+1)*60 seconds + // exactly (e.g. click at HH:MM:59.9 → start floors to HH:MM:00 and + // queryFn fires just past HH:(MM+1):00), so allow equality on the + // upper bound. + expect(diff).toBeGreaterThanOrEqual(minMinutes * 60); + expect(diff).toBeLessThanOrEqual((minMinutes + 1) * 60); + }); + return diff; + }; + + it("should pass a ~1-minute window to uiSpendLogsCall when 'Last Minute' is selected", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last Minute" })); + + await waitForWindowSeconds(1); + }); + + it("should pass a ~15-minute window to uiSpendLogsCall when 'Last 15 Minutes' is selected", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last 15 Minutes" })); + + await waitForWindowSeconds(15); + }); + + it("should update the time-range button label to 'Last Minute' after selecting it", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByRole("button", { name: /Last 24 Hours/i })); + await user.click(await screen.findByRole("button", { name: "Last Minute" })); + + expect(screen.getByRole("button", { name: "Last Minute" })).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /Last 24 Hours/i })).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/workflow_runs/index.tsx b/ui/litellm-dashboard/src/components/workflow_runs/index.tsx index b9467c5604b..2aecbece2e3 100644 --- a/ui/litellm-dashboard/src/components/workflow_runs/index.tsx +++ b/ui/litellm-dashboard/src/components/workflow_runs/index.tsx @@ -587,6 +587,7 @@ const WorkflowRuns: React.FC = ({ accessToken }) => { return (