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>
This commit is contained in:
Devin AI 2026-09-18 12:17:43 +00:00
parent 8fc9c46d1a
commit e45d65a0e9
2 changed files with 188 additions and 0 deletions

View file

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

View file

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