fix(rate-limit): validate enum membership at duck-typed read sites + enrich BudgetExceededError llm_provider
Some checks are pending
Unit Tests: Caching (Redis) / caching-redis (push) Waiting to run
Unit Tests: Proxy DB Operations / assert-shard-coverage (push) Waiting to run
Unit Tests: Proxy DB Operations / auth-checks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / budgets (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / custom-logging (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / db-and-spend (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-runtime (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / endpoints-and-responses (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / guardrails-hooks (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / jwt-and-keys (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / key-generation (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / logging-misc (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-server-core (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / schema-migration (push) Blocked by required conditions
Unit Tests: Proxy DB Operations / proxy-utils (push) Blocked by required conditions
Unit Tests: Security / security (push) Waiting to run

Two follow-ups uncovered during the second QA pass on PR #27687:

1. Guard third-party `.category` / `.rate_limit_type` attribute leakage.
   The duck-typed read in `get_error_information` and
   `_extract_rate_limit_labels` would forward any string attribute named
   `category` / `rate_limit_type` on an unrelated third-party exception
   into the StandardLoggingPayload and Prometheus labels — silently
   mislabeling custom-callback payloads and blowing out Prometheus label
   cardinality. Add `validate_rate_limit_category` /
   `validate_rate_limit_type` helpers that gate on the documented enum
   value sets; non-matching values are dropped to None.

2. Enrich BudgetExceededError.llm_provider from request_data.
   Budget checks live in tenant-scoped helpers (key / team / org / tag /
   end-user / project) that don't see the request model, so the
   BudgetExceededError they raise carried llm_provider="" — leaving
   custom-metrics consumers without provider attribution for the most
   common 429 case. Resolve it once at the central
   UserAPIKeyAuthExceptionHandler seam, before post_call_failure_hook
   fires, so the StandardLoggingPayload the callback sees has the same
   provider attribution as RPM/TPM 429s.

Regression tests pin both: 4 leakage tests + 4 enrichment tests. The
leakage tests would fail under the pre-validation version of either read
site; the enrichment tests would fail if the handler skipped the
resolver call.
This commit is contained in:
mateo-berri 2026-05-13 22:32:09 -07:00
parent 807d1b29d6
commit f2d324310a
5 changed files with 231 additions and 21 deletions

View file

@ -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

View file

@ -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(

View file

@ -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(

View file

@ -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,

View file

@ -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"