From e45d65a0e912149ff1bafeab2d825ec5a1b981e6 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 18 Sep 2026 12:17:43 +0000 Subject: [PATCH] feat(proxy): enforce tpm_limit and rpm_limit set on tag objects Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../hooks/parallel_request_limiter_v3.py | 63 +++++++++ .../hooks/test_parallel_request_limiter_v3.py | 125 ++++++++++++++++++ 2 files changed, 188 insertions(+) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 89826694bb6..363e7121c8b 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -548,6 +548,38 @@ class RequestRateLimiterStash: batch_enqueued_reservation: BatchEnqueuedTokenReservation | None = None batch_tpd_refund_ops: tuple[ReservationAwareIncrementOperation, ...] = () reservation_released: bool = False + tpm_limited_tags: frozenset[str] = field(default_factory=frozenset) + + +@dataclass(frozen=True, slots=True) +class TagRateLimit: + rpm_limit: int | None + tpm_limit: int | None + + +TagRateLimitResolver: TypeAlias = Callable[[Sequence[str]], Awaitable[Mapping[str, TagRateLimit]]] + + +async def resolve_tag_rate_limits_from_db(tag_names: Sequence[str]) -> Mapping[str, TagRateLimit]: + """Read the rpm/tpm limits stored on each tag's budget row, served from the tag object cache.""" + from litellm.proxy.auth.auth_checks import get_tag_objects_batch + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + if prisma_client is None or not tag_names: + return MappingProxyType({}) + tag_objects: Final = await get_tag_objects_batch( + tag_names=tag_names, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + return MappingProxyType( + { + tag_name: TagRateLimit(rpm_limit=budget.rpm_limit, tpm_limit=budget.tpm_limit) + for tag_name, tag_object in tag_objects.items() + if (budget := tag_object.litellm_budget_table) is not None + and (budget.rpm_limit is not None or budget.tpm_limit is not None) + } + ) _request_stash: Final[ContextVar[RequestRateLimiterStash | None]] = ContextVar( @@ -613,9 +645,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self, internal_usage_cache: InternalUsageCache, time_provider: Callable[[], datetime] | None = None, + tag_rate_limit_resolver: TagRateLimitResolver = resolve_tag_rate_limits_from_db, ): self.internal_usage_cache = internal_usage_cache self._time_provider = time_provider or datetime.now + self._tag_rate_limit_resolver = tag_rate_limit_resolver if self.internal_usage_cache.dual_cache.redis_cache is not None: self.batch_rate_limiter_script = self.internal_usage_cache.dual_cache.redis_cache.async_register_script( BATCH_RATE_LIMITER_SCRIPT @@ -2686,6 +2720,26 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return descriptors + async def _create_tag_rate_limit_descriptors(self, data: Mapping[str, object]) -> tuple[RateLimitDescriptor, ...]: + """One ``tag`` descriptor per request tag whose tag object carries an rpm or tpm limit.""" + tags: Final = tuple(dict.fromkeys(get_tags_from_request_body(data))) + if not tags: + return () + tag_limits: Final = await self._tag_rate_limit_resolver(tags) + return tuple( + RateLimitDescriptor( + key="tag", + value=tag, + rate_limit={ + "requests_per_unit": limit.rpm_limit, + "tokens_per_unit": limit.tpm_limit, + "window_size": self.window_size, + }, + ) + for tag in tags + if (limit := tag_limits.get(tag)) is not None + ) + def _create_rate_limit_descriptors( self, user_api_key_dict: UserAPIKeyAuth, @@ -3480,6 +3534,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): return [ # mutable-ok: the shared generation reservation helpers require a list *descriptors, *self.create_organization_rate_limit_descriptor(user_api_key_dict, requested_model), + *await self._create_tag_rate_limit_descriptors(data), ] async def _release_request_capacity_when_admitted( @@ -3574,6 +3629,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): data=request_data, call_type=call_type, ) + stash.tpm_limited_tags = frozenset( + d["value"] + for d in descriptors + if d["key"] == "tag" and d["rate_limit"] is not None and d["rate_limit"].get("tokens_per_unit") is not None + ) # Only check rate limits if we have descriptors with actual limits if descriptors: @@ -4260,6 +4320,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): standard_logging_metadata: dict[str, Any], kwargs: object, model_group: str | None, + tpm_limited_tags: Set[str] = frozenset(), ) -> list[tuple[str, str]]: """ Enumerate every (scope_key, scope_value) pair that *might* carry a @@ -4315,6 +4376,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): targets.append(("agent", agent_id)) if session_id: targets.append(("agent_session", f"{agent_id}:{session_id}")) + targets.extend(("tag", tag) for tag in sorted(tpm_limited_tags)) return targets def _build_reservation_aware_tpm_ops( @@ -4488,6 +4550,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): standard_logging_metadata=standard_logging_metadata, kwargs=kwargs, model_group=reconcile_model, + tpm_limited_tags=stash.tpm_limited_tags if stash is not None else frozenset(), ) charged_targets: Final = ( [target for target in targets if target[0] != "model_per_team"] diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 85023b94207..f3335d4bce5 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -24,6 +24,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import ( PARALLEL_REQUEST_SLOT_TTL_SECONDS, ParallelSlotAcquisition, RequestRateLimiterStash, + TagRateLimit, _request_stash, get_or_create_request_stash, get_request_stash, @@ -5001,6 +5002,130 @@ async def test_per_tag_untagged_request_governed_by_key_limit_v3(monkeypatch): assert "tag_per_key" not in str(exc_info.value.detail) +def _static_tag_limits(limits: dict[str, TagRateLimit]): + calls: list[tuple[str, ...]] = [] + + async def resolver(tag_names: Sequence[str]): + calls.append(tuple(tag_names)) + return {name: limits[name] for name in tag_names if name in limits} + + return resolver, calls + + +@pytest.mark.asyncio +async def test_tag_object_rpm_limit_enforced_v3(monkeypatch): + """ + rpm_limit stored on the tag object (via /tag/new) is enforced for every key + sending that tag, independently of any key-level tag_rpm_limit metadata. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + _request_stash.set(None) + resolver, calls = _static_tag_limits({"cell-1": TagRateLimit(rpm_limit=2, tpm_limit=None)}) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + tag_rate_limit_resolver=resolver, + ) + + async def call(api_key: str, tags: list[str]) -> None: + await handler.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(api_key=hash_token(api_key)), + cache=local_cache, + data={"model": "gpt-3.5-turbo", "metadata": {"tags": tags}}, + call_type="", + ) + + await call("sk-a", ["cell-1"]) + await call("sk-b", ["cell-1", "cell-2"]) + with pytest.raises(HTTPException) as exc_info: + await call("sk-a", ["cell-1"]) + assert exc_info.value.status_code == 429 + assert "tag" in str(exc_info.value.detail) + + for _ in range(3): + await call("sk-a", ["cell-2"]) + await call("sk-a", []) + assert calls == [("cell-1",), ("cell-1", "cell-2"), ("cell-1",), ("cell-2",), ("cell-2",), ("cell-2",)] + + +@pytest.mark.asyncio +async def test_tag_object_tpm_limit_enforced_v3(monkeypatch): + """ + tpm_limit stored on the tag object is charged from actual usage on success + and blocks the tag once exhausted, while untagged traffic keeps flowing. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + monkeypatch.setenv("LITELLM_TPM_TOKEN_RESERVATION_ENABLED", "false") + _request_stash.set(None) + resolver, _ = _static_tag_limits({"cell-1": TagRateLimit(rpm_limit=None, tpm_limit=100)}) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache), + tag_rate_limit_resolver=resolver, + ) + monkeypatch.setattr(handler, "get_rate_limit_type", lambda: "total") + user_api_key_dict = UserAPIKeyAuth(api_key=hash_token("sk-tag-tpm")) + + async def call(tags: list[str]) -> None: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-3.5-turbo", "metadata": {"tags": tags}}, + call_type="", + ) + + await call(["cell-1"]) + tokens_before_success = await local_cache.async_get_cache("{tag:cell-1}:tokens") or 0 + await handler.async_log_success_event( + kwargs={ + "standard_logging_object": {"metadata": {"user_api_key_hash": user_api_key_dict.api_key}}, + "model": "gpt-3.5-turbo", + }, + response_obj=ModelResponse(usage=Usage(prompt_tokens=60, completion_tokens=60, total_tokens=120)), + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert await local_cache.async_get_cache("{tag:cell-1}:tokens") == tokens_before_success + 120 + + with pytest.raises(HTTPException) as exc_info: + await call(["cell-1"]) + assert exc_info.value.status_code == 429 + await call([]) + + +@pytest.mark.asyncio +async def test_resolve_tag_rate_limits_from_db_reads_budget_row(monkeypatch): + """Only tags whose budget row carries an rpm or tpm limit are returned.""" + from litellm.models.budget import LiteLLM_BudgetTable + from litellm.models.tag import LiteLLM_TagTable + from litellm.proxy import proxy_server + from litellm.proxy.auth import auth_checks + from litellm.proxy.hooks.parallel_request_limiter_v3 import resolve_tag_rate_limits_from_db + + async def fake_batch(tag_names, prisma_client, user_api_key_cache, **kwargs): + return { + "limited": LiteLLM_TagTable( + tag_name="limited", + litellm_budget_table=LiteLLM_BudgetTable(rpm_limit=3, tpm_limit=None), + ), + "spend-only": LiteLLM_TagTable( + tag_name="spend-only", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=5.0), + ), + "bare": LiteLLM_TagTable(tag_name="bare"), + } + + monkeypatch.setattr(proxy_server, "prisma_client", object()) + monkeypatch.setattr(auth_checks, "get_tag_objects_batch", fake_batch) + + assert dict(await resolve_tag_rate_limits_from_db(["limited", "spend-only", "bare"])) == { + "limited": TagRateLimit(rpm_limit=3, tpm_limit=None) + } + + monkeypatch.setattr(proxy_server, "prisma_client", None) + assert dict(await resolve_tag_rate_limits_from_db(["limited"])) == {} + + # -------------------------------------------------------------------------- # Streaming success logging mirrors x-ratelimit-* remaining values into # standard_logging_object.hidden_params.additional_headers so Prometheus /