diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 8208a269b84..89a0b6663ab 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -81,6 +81,37 @@ class RateLimitType(str, enum.Enum): """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 diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 4ce4014d077..3b2371490db 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, @@ -2686,23 +2690,18 @@ class PrometheusLogger(CustomLogger): """ Pull the unified ``category`` / ``rate_limit_type`` fields off any exception that declares them (``litellm.RateLimitError`` and bare- - Exception subclasses like ``BudgetExceededError`` that set these - attributes directly) so Prometheus can split 429s by source + - dimension without the consumer parsing free-text error messages. + Exception subclasses like ``BudgetExceededError``). - Returns ``(None, None)`` for exceptions that don't declare these - fields. Both classes normalize their values to plain ``str`` at - construction, so this helper only needs to coerce defensively. + 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 - - def _coerce(value: Any) -> Optional[str]: - return str(value) if value is not None else None - return ( - _coerce(getattr(exception, "category", None)), - _coerce(getattr(exception, "rate_limit_type", None)), + validate_rate_limit_category(getattr(exception, "category", None)), + validate_rate_limit_type(getattr(exception, "rate_limit_type", None)), ) async def log_success_fallback_event( diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e1585398831..7a5c26404f5 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 @@ -5159,18 +5163,17 @@ class StandardLoggingPayloadSetup: # Get additional error details error_message = str(original_exception) - # Surface the unified `category` and `rate_limit_type` fields off any - # exception that opts in by setting them. Duck-typed rather than - # isinstance-gated on RateLimitError so bare-Exception subclasses like + # Duck-typed read so bare-Exception subclasses like # `litellm.BudgetExceededError` can participate without joining the # RateLimitError hierarchy (which would break `except BudgetExceededError`). - # Both RateLimitError and BudgetExceededError normalize their values to - # plain strings at construction. - rate_limit_category: Optional[str] = getattr( - original_exception, "category", None + # 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: Optional[str] = getattr( - original_exception, "rate_limit_type", None + rate_limit_type = validate_rate_limit_type( + getattr(original_exception, "rate_limit_type", None) ) return StandardLoggingPayloadErrorInformation( diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index 431db4254eb..968380058d5 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -106,6 +106,21 @@ class UserAPIKeyAuthExceptionHandler: api_key=api_key, request_route=route, ) + + # 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/tests/test_litellm/test_rate_limit_error_unification.py b/tests/test_litellm/test_rate_limit_error_unification.py index 2913492ed63..8287e82ded0 100644 --- a/tests/test_litellm/test_rate_limit_error_unification.py +++ b/tests/test_litellm/test_rate_limit_error_unification.py @@ -1507,3 +1507,165 @@ class TestBudgetExceededErrorSurfacesUnifiedFields: ) 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"