feat(router): generic tag-based rate limiting per model group

Adds token/request/dollar caps keyed by request tag, with per-entry tag_id and configurable day windows, enforced on every routing attempt
This commit is contained in:
Devin AI 2026-07-29 19:15:26 +00:00
parent 8b03315ac6
commit e579749d89
6 changed files with 568 additions and 58 deletions

View file

@ -130,9 +130,7 @@ def classify_value(value: object, key: str = "scan") -> ValueClass:
return "plaintext"
if value.startswith(_V2_GCM_PREFIX):
return "migrated"
decrypted = decrypt_value_helper(
value=value, key=key, exception_type="debug", return_original_value=False
)
decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False)
if decrypted is None:
# Did not decrypt under nacl and has no v2 marker: legacy plaintext.
return "plaintext"
@ -151,9 +149,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object:
return value
if value.startswith(_V2_GCM_PREFIX):
return value # idempotent: already migrated
decrypted = decrypt_value_helper(
value=value, key=key, exception_type="debug", return_original_value=False
)
decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False)
if decrypted is None:
# Either legacy plaintext (no ciphertext to migrate) or corrupt. Either
# way, do not overwrite — preserve the value as stored.
@ -161,9 +157,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object:
return encrypt_value_helper(decrypted)
def reencrypt_selective_dict(
data: dict[str, object], sensitive_keys: list[str]
) -> dict[str, object]:
def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) -> dict[str, object]:
"""Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted.
Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is.
@ -212,9 +206,7 @@ async def _migrate_config_settings_row(
dict with selected sensitive fields (vantage_settings / cloudzero_settings).
"""
report = LocationReport(location=param_name)
record = await prisma_client.db.litellm_config.find_unique(
where={"param_name": param_name}
)
record = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name})
if record is None or record.param_value is None:
return report
@ -266,9 +258,7 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR
every present string field.
"""
report = LocationReport(location="sso_config")
record = await prisma_client.db.litellm_ssoconfig.find_unique(
where={"id": "sso_config"}
)
record = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"})
if record is None or record.sso_settings is None:
return report
@ -344,9 +334,7 @@ async def _migrate_callback_vars_table(
rows = await table.find_many()
for row in rows or []:
metadata = getattr(row, "metadata", None)
if not isinstance(metadata, dict) or (
"logging" not in metadata and "callback_settings" not in metadata
):
if not isinstance(metadata, dict) or ("logging" not in metadata and "callback_settings" not in metadata):
continue
# Classify every callback-var value directly (strip the litellm_enc::
@ -534,9 +522,7 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport:
"""Scan the ``environment_variables`` config row (``param_value`` dict)."""
report = LocationReport(location="config_environment_variables")
try:
record = await prisma_client.db.litellm_config.find_unique(
where={"param_name": "environment_variables"}
)
record = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"})
except Exception as e: # pragma: no cover - defensive
verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e))
return report
@ -557,11 +543,7 @@ async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]:
"""Read-only classification of every rotation-covered table. No writes."""
reports: list[LocationReport] = []
for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS:
reports.append(
await _scan_one_table(
prisma_client, location, db_attr, json_cols, scalar_cols
)
)
reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols))
reports.append(await _scan_config_env_vars(prisma_client))
return reports
@ -575,9 +557,7 @@ _VANTAGE_SENSITIVE = ["api_key", "integration_token"]
_CLOUDZERO_SENSITIVE = ["api_key"]
async def _migrate_covered_tables(
prisma_client: object, user_api_key_dict: object
) -> list[LocationReport]:
async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: object) -> list[LocationReport]:
"""Re-encrypt the tables already covered by ``_rotate_master_key`` (model
table, credentials, MCP credential/env tables, config environment_variables)
by running that orchestrator in *same-key* mode. With the AES gate on, the
@ -597,8 +577,7 @@ async def _migrate_covered_tables(
current_key = _get_salt_key()
if current_key is None:
raise RuntimeError(
"Cannot migrate covered tables: no salt key / master key is set. "
"Set LITELLM_SALT_KEY before migrating."
"Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating."
)
await _rotate_master_key(
prisma_client=cast("PrismaClient", prisma_client),
@ -648,19 +627,9 @@ async def migrate_encryption(
# Net-new walkers (items 3, 4, 11, 12, 13).
report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run))
report.add(
await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run
)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run
)
)
report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run))
report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run))
report.add(await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run))
report.add(await _migrate_sso_config(prisma_client, dry_run))
return report
@ -683,20 +652,10 @@ async def check_encryption(prisma_client: object) -> MigrationReport:
# Net-new walker locations, in dry-run (read-only) mode.
report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True))
report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True))
report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True))
report.add(
await _migrate_callback_vars_table(
prisma_client, "verification_token", dry_run=True
)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True
)
)
report.add(
await _migrate_config_settings_row(
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True
)
await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True)
)
report.add(await _migrate_sso_config(prisma_client, dry_run=True))
return report

View file

@ -31,7 +31,6 @@ from typing import (
Generator,
List,
Literal,
Mapping,
Optional,
Set,
Tuple,
@ -84,6 +83,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import (
)
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.router_strategy.tag_limits import TagLimitCheck, has_tag_limits
from litellm.router_strategy.least_busy import LeastBusyLoggingHandler
from litellm.router_strategy.lowest_cost import LowestCostLoggingHandler
from litellm.router_strategy.lowest_latency import LowestLatencyLoggingHandler
@ -720,6 +720,11 @@ class Router:
optional_pre_call_checks.append("router_budget_limiting")
else:
optional_pre_call_checks = ["router_budget_limiting"]
if has_tag_limits(model_list):
if optional_pre_call_checks is not None:
optional_pre_call_checks.append("tag_limits")
else:
optional_pre_call_checks = ["tag_limits"]
self.retry_policy: Optional[RetryPolicy] = None
if retry_policy is not None:
if isinstance(retry_policy, dict):
@ -1708,6 +1713,10 @@ class Router:
self.router_budget_logger = _callback
elif pre_call_check == "enforce_model_rate_limits":
_callback = ModelRateLimitingCheck(dual_cache=self.cache)
elif pre_call_check == "tag_limits":
if any(isinstance(_cb, TagLimitCheck) for _cb in (self.optional_callbacks or [])):
continue
_callback = TagLimitCheck(dual_cache=self.cache)
if _callback is None:
continue

View file

@ -0,0 +1,242 @@
"""
Tag-based rate limiting.
A model group declares one grouping per unit on `model_info`, and each limit entry
carries its own `tag_id`, so limits of the same unit can key off different tags:
```yaml
model_info:
token_limits:
limits:
- name: daily
tag_id: end_user_id
limit: 500000
period_days: 1
request_limits:
limits:
- name: daily
tag_id: end_user_id
limit: 1000
period_days: 1
dollar_limits:
limits:
- name: monthly
tag_id: team_id
limit: 50.0
period_days: 30
```
Callers identify themselves with the existing tag mechanism, e.g.
`X-Litellm-Tags: end_user_id:user-123,team_id:team-a`. Usage is counted in fixed
windows of `period_days` and checked on every routing attempt, so a fallback model
group's own limits are enforced on its own hop. Requests carrying no matching
`tag_id` are unaffected.
"""
import json
import time
from collections.abc import Callable, Mapping, Sequence
from functools import lru_cache
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_router_logger
from litellm.caching.dual_cache import DualCache
from litellm.exceptions import RateLimitErrorCategory, RateLimitType
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs
from litellm.types.llms.openai import AllMessageValues
from litellm.types.tag_limits import (
TAG_LIMIT_FIELD_BY_UNIT,
DeploymentTagLimits,
TagLimit,
TagLimitUnit,
)
SECONDS_PER_DAY = 86400
_LIMIT_FIELDS: tuple[str, ...] = tuple(TAG_LIMIT_FIELD_BY_UNIT.values())
_RATE_LIMIT_TYPE_BY_UNIT: Mapping[TagLimitUnit, RateLimitType] = {
"tokens": RateLimitType.TOKENS,
"requests": RateLimitType.REQUESTS,
"dollars": RateLimitType.BUDGET,
}
_MODEL_INFO_ADAPTER: TypeAdapter[Mapping[str, object]] = TypeAdapter(Mapping[str, object])
@lru_cache(maxsize=1024)
def _parse_tag_limits(serialized_limits: str) -> DeploymentTagLimits | None:
try:
limits = DeploymentTagLimits.model_validate_json(serialized_limits)
except ValidationError as e:
verbose_router_logger.warning("tag_limits: ignoring invalid limit config %s: %s", serialized_limits, e)
return None
return limits if limits.limits_by_unit() else None
def get_tag_limits(model_info: object) -> DeploymentTagLimits | None:
"""Parse (and memoize) the tag limits declared on a deployment's `model_info`."""
if not isinstance(model_info, dict) or all(field not in model_info for field in _LIMIT_FIELDS):
return None
info = _MODEL_INFO_ADAPTER.validate_python(model_info)
declared = {field: info[field] for field in _LIMIT_FIELDS if field in info}
return _parse_tag_limits(json.dumps(declared, sort_keys=True, default=str))
def has_tag_limits(model_list: Sequence[Mapping[str, object]] | None) -> bool:
if model_list is None:
return False
return any(_deployment_limits(deployment).limits_by_unit() for deployment in model_list)
def _tag_values_from_request(kwargs: Mapping[str, object] | None) -> Mapping[str, str]:
"""
Turn `metadata.tags` entries of the form `<tag_id>:<value>` into a {tag_id: value} map.
Tags without a `:` separator carry no identity and are ignored. On duplicate tag ids
the last one wins.
"""
request_kwargs = dict(kwargs or {})
tags = _get_tags_from_request_kwargs(
request_kwargs=request_kwargs,
metadata_variable_name=get_metadata_variable_name_from_kwargs(request_kwargs),
)
return {tag_id: tag_value for tag_id, _, tag_value in (tag.partition(":") for tag in tags) if tag_id and tag_value}
def _window_seconds(limit: TagLimit) -> int:
return max(int(limit.period_days * SECONDS_PER_DAY), 1)
def _usage_key(unit: TagLimitUnit, limit: TagLimit, tag_value: str, now: float) -> str:
window = _window_seconds(limit)
bucket = int(now // window)
return f"tag_limit:{unit}:{limit.tag_id}:{tag_value}:{limit.name}:{limit.period_days}:{bucket}"
class _RequestUsage(BaseModel):
"""The `StandardLoggingPayload` fields tag counters are incremented by."""
model_config = ConfigDict(extra="ignore")
total_tokens: float = 0.0
response_cost: float = 0.0
def _nested(container: object, key: str) -> object:
return container.get(key) if isinstance(container, dict) else None
def _deployment_limits(deployment: Mapping[str, object]) -> DeploymentTagLimits:
return get_tag_limits(deployment.get("model_info")) or DeploymentTagLimits()
class TagLimitCheck(CustomLogger):
"""Enforces `model_info` tag limits on every routing attempt, and counts usage on success."""
def __init__(self, dual_cache: DualCache, now: Callable[[], float] = time.time):
self.dual_cache = dual_cache
self.now = now
async def async_filter_deployments(
self,
model: str,
healthy_deployments: list[dict[str, object]] # mutable-ok: router hands the live deployment dicts around
| dict[str, object],
messages: Sequence[AllMessageValues] | None,
request_kwargs: Mapping[str, object] | None = None,
parent_otel_span: Span | None = None,
) -> list[dict[str, object]]: # mutable-ok: CustomLogger.async_filter_deployments returns a list of deployments
deployments: list[dict[str, object]] = ( # mutable-ok: mirrors the list contract above
[healthy_deployments] if isinstance(healthy_deployments, dict) else list(healthy_deployments)
)
if not deployments:
return deployments
tag_values = _tag_values_from_request(request_kwargs)
if not tag_values:
return deployments
now = self.now()
checks: tuple[tuple[int, TagLimitUnit, TagLimit, str, str], ...] = tuple(
(idx, unit, limit, tag_values[limit.tag_id], _usage_key(unit, limit, tag_values[limit.tag_id], now))
for idx, deployment in enumerate(deployments)
for unit, limit in _deployment_limits(deployment).limits_by_unit()
if limit.tag_id in tag_values
)
if not checks:
return deployments
keys = tuple(dict.fromkeys(check[4] for check in checks))
values = await self.dual_cache.async_batch_get_cache(keys=list(keys), parent_otel_span=parent_otel_span)
usage: Mapping[str, float] = {key: float(value or 0.0) for key, value in zip(keys, values or [0.0] * len(keys))}
violations = tuple(check for check in checks if usage[check[4]] >= check[2].limit)
blocked_indexes = frozenset(violation[0] for violation in violations)
allowed = [deployment for idx, deployment in enumerate(deployments) if idx not in blocked_indexes]
if allowed:
return allowed
_, unit, limit, tag_value, key = violations[0]
error = {
"message": "tag_rate_limit_exceeded",
"type": unit,
"tag_id": limit.tag_id,
"tag_value": tag_value,
"limit_name": limit.name,
"limit": limit.limit,
"period_days": limit.period_days,
}
verbose_router_logger.info("tag_limits: blocking request, usage=%s, %s", usage[key], error)
raise litellm.RateLimitError(
message=json.dumps({"error": error}),
llm_provider="litellm",
model=model,
category=RateLimitErrorCategory.LITELLM_RATE_LIMIT,
rate_limit_type=_RATE_LIMIT_TYPE_BY_UNIT[unit],
detail={"error": error},
)
async def async_log_success_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: object,
end_time: object,
) -> None:
tag_limits = get_tag_limits(_nested(kwargs.get("litellm_params"), "model_info"))
if tag_limits is None:
return
tag_values = _tag_values_from_request(kwargs)
if not tag_values:
return
standard_logging_payload = kwargs.get("standard_logging_object")
if not isinstance(standard_logging_payload, dict):
return
usage = _RequestUsage.model_validate(standard_logging_payload)
usage_by_unit: Mapping[TagLimitUnit, float] = {
"tokens": usage.total_tokens,
"requests": 1.0,
"dollars": usage.response_cost,
}
now = self.now()
for unit, limit in tag_limits.limits_by_unit():
tag_value = tag_values.get(limit.tag_id)
if tag_value is None or usage_by_unit[unit] == 0:
continue
await self.dual_cache.async_increment_cache(
key=_usage_key(unit, limit, tag_value, now),
value=usage_by_unit[unit],
ttl=2 * _window_seconds(limit),
)

View file

@ -790,6 +790,7 @@ OptionalPreCallChecks = List[
"forward_client_headers_by_model_group",
"enforce_model_rate_limits",
"encrypted_content_affinity",
"tag_limits",
]
]

View file

@ -0,0 +1,47 @@
"""
Types for tag-based rate limiting (`model_info.token_limits` / `request_limits` / `dollar_limits`).
"""
from collections.abc import Mapping
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
TagLimitUnit = Literal["tokens", "requests", "dollars"]
TAG_LIMIT_FIELD_BY_UNIT: Mapping[TagLimitUnit, str] = {
"tokens": "token_limits",
"requests": "request_limits",
"dollars": "dollar_limits",
}
class TagLimit(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
name: str
tag_id: str
limit: float = Field(gt=0)
period_days: float = Field(gt=0)
class TagLimitGroup(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid")
limits: tuple[TagLimit, ...] = ()
class DeploymentTagLimits(BaseModel):
model_config = ConfigDict(frozen=True, extra="ignore")
token_limits: TagLimitGroup | None = None
request_limits: TagLimitGroup | None = None
dollar_limits: TagLimitGroup | None = None
def limits_by_unit(self) -> tuple[tuple[TagLimitUnit, TagLimit], ...]:
groups: tuple[tuple[TagLimitUnit, TagLimitGroup | None], ...] = (
("tokens", self.token_limits),
("requests", self.request_limits),
("dollars", self.dollar_limits),
)
return tuple((unit, limit) for unit, group in groups if group is not None for limit in group.limits)

View file

@ -0,0 +1,252 @@
"""
Tests for tag-based rate limiting (`model_info.token_limits` / `request_limits` / `dollar_limits`).
"""
import asyncio
import json
from typing import Any, Dict, List, Optional
import pytest
import litellm
from litellm.caching.dual_cache import DualCache
from litellm.router_strategy.tag_limits import (
TagLimitCheck,
get_tag_limits,
has_tag_limits,
)
def _deployment(model_info: Optional[Dict[str, Any]] = None, model_id: str = "id-1") -> Dict[str, Any]:
return {
"model_name": "gpt-4o",
"litellm_params": {"model": "openai/gpt-4o"},
"model_info": {"id": model_id, **(model_info or {})},
}
def _limits(
unit: str = "request_limits",
name: str = "daily",
tag_id: str = "end_user_id",
limit: float = 2,
period_days: float = 1,
) -> Dict[str, Any]:
return {unit: {"limits": [{"name": name, "tag_id": tag_id, "limit": limit, "period_days": period_days}]}}
def _request_kwargs(tags: List[str]) -> Dict[str, Any]:
return {"metadata": {"tags": tags}}
def _success_kwargs(
model_info: Dict[str, Any],
tags: List[str],
total_tokens: int = 0,
response_cost: float = 0.0,
) -> Dict[str, Any]:
return {
"litellm_params": {"model_info": {"id": "id-1", **model_info}, "metadata": {"tags": tags}},
"standard_logging_object": {"total_tokens": total_tokens, "response_cost": response_cost},
}
async def _filter(check: TagLimitCheck, deployments: List[Dict[str, Any]], tags: List[str]) -> List[dict]:
return await check.async_filter_deployments(
model="gpt-4o",
healthy_deployments=deployments,
messages=None,
request_kwargs=_request_kwargs(tags),
)
def test_get_tag_limits_parses_per_entry_tag_ids():
model_info = {
"token_limits": {
"limits": [
{"name": "daily", "tag_id": "end_user_id", "limit": 500000, "period_days": 1},
{"name": "weekly", "tag_id": "team_id", "limit": 2000000, "period_days": 7},
]
},
"dollar_limits": {"limits": [{"name": "monthly", "tag_id": "team_id", "limit": 50.0, "period_days": 30}]},
}
limits = get_tag_limits(model_info)
assert limits is not None
assert [(unit, limit.name, limit.tag_id, limit.limit) for unit, limit in limits.limits_by_unit()] == [
("tokens", "daily", "end_user_id", 500000),
("tokens", "weekly", "team_id", 2000000),
("dollars", "monthly", "team_id", 50.0),
]
def test_get_tag_limits_ignores_deployments_without_limits_and_invalid_config():
assert get_tag_limits({"id": "id-1"}) is None
assert get_tag_limits(None) is None
assert get_tag_limits({"request_limits": {"limits": [{"name": "daily", "limit": 5}]}}) is None
assert has_tag_limits([_deployment()]) is False
assert has_tag_limits([_deployment(_limits())]) is True
@pytest.mark.asyncio
async def test_request_without_matching_tag_is_unaffected():
check = TagLimitCheck(dual_cache=DualCache())
deployments = [_deployment(_limits(limit=1))]
assert await _filter(check, deployments, []) == deployments
assert await _filter(check, deployments, ["team_id:team-a"]) == deployments
assert await _filter(check, deployments, ["end_user_id"]) == deployments
# untagged usage never lands in a counter, so the deployment stays available
await check.async_log_success_event(_success_kwargs(_limits(limit=1), ["team_id:team-a"]), None, None, None)
assert await _filter(check, deployments, ["team_id:team-a"]) == deployments
@pytest.mark.asyncio
@pytest.mark.parametrize(
"unit, field, usage_kwargs",
[
("requests", "request_limits", {}),
("tokens", "token_limits", {"total_tokens": 1}),
("dollars", "dollar_limits", {"response_cost": 1.0}),
],
)
async def test_limit_blocks_with_429_once_exceeded(unit: str, field: str, usage_kwargs: Dict[str, Any]):
model_info = _limits(unit=field, limit=2)
check = TagLimitCheck(dual_cache=DualCache())
deployments = [_deployment(model_info)]
tags = ["end_user_id:user-123"]
assert await _filter(check, deployments, tags) == deployments
for _ in range(2):
await check.async_log_success_event(_success_kwargs(model_info, tags, **usage_kwargs), None, None, None)
with pytest.raises(litellm.RateLimitError) as exc_info:
await _filter(check, deployments, tags)
assert exc_info.value.status_code == 429
assert exc_info.value.detail == {
"error": {
"message": "tag_rate_limit_exceeded",
"type": unit,
"tag_id": "end_user_id",
"tag_value": "user-123",
"limit_name": "daily",
"limit": 2.0,
"period_days": 1.0,
}
}
assert json.loads(exc_info.value.message.split("litellm.RateLimitError: ")[1]) == exc_info.value.detail
# a different tag value has its own counter
assert await _filter(check, deployments, ["end_user_id:user-456"]) == deployments
@pytest.mark.asyncio
async def test_usage_is_counted_per_limit_entry_tag_id():
model_info = {
"token_limits": {
"limits": [
{"name": "daily", "tag_id": "end_user_id", "limit": 10, "period_days": 1},
{"name": "weekly", "tag_id": "team_id", "limit": 100, "period_days": 7},
]
}
}
check = TagLimitCheck(dual_cache=DualCache())
deployments = [_deployment(model_info)]
tags = ["end_user_id:user-123", "team_id:team-a"]
for _ in range(2):
await check.async_log_success_event(_success_kwargs(model_info, tags, total_tokens=6), None, None, None)
# end_user_id daily limit (10) is exceeded by 12 tokens; team_id weekly (100) is not
with pytest.raises(litellm.RateLimitError) as exc_info:
await _filter(check, deployments, tags)
assert exc_info.value.detail["error"]["tag_id"] == "end_user_id"
# another end user on the same team is still under both limits
assert await _filter(check, deployments, ["end_user_id:user-999", "team_id:team-a"]) == deployments
@pytest.mark.asyncio
async def test_usage_from_a_previous_window_is_not_counted():
model_info = _limits(limit=1, period_days=1)
clock = {"now": 1_700_000_000.0}
check = TagLimitCheck(dual_cache=DualCache(), now=lambda: clock["now"])
deployments = [_deployment(model_info)]
tags = ["end_user_id:user-123"]
await check.async_log_success_event(_success_kwargs(model_info, tags), None, None, None)
with pytest.raises(litellm.RateLimitError):
await _filter(check, deployments, tags)
clock["now"] += 86400
assert await _filter(check, deployments, tags) == deployments
@pytest.mark.asyncio
async def test_only_deployments_over_their_own_limits_are_filtered_out():
limited = _deployment(_limits(limit=1), model_id="limited")
unlimited = _deployment(model_id="unlimited")
check = TagLimitCheck(dual_cache=DualCache())
tags = ["end_user_id:user-123"]
await check.async_log_success_event(_success_kwargs(_limits(limit=1), tags), None, None, None)
assert await _filter(check, [limited, unlimited], tags) == [unlimited]
@pytest.mark.asyncio
async def test_router_registers_and_enforces_tag_limits():
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4o",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-fake", "mock_response": "hi"},
"model_info": _limits(limit=1),
}
],
num_retries=0,
)
checks = [cb for cb in (router.optional_callbacks or []) if isinstance(cb, TagLimitCheck)]
assert len(checks) == 1
request = {
"model": "gpt-4o",
"messages": [{"role": "user", "content": "hello"}],
"metadata": {"tags": ["end_user_id:user-123"]},
}
assert (await router.acompletion(**request)).choices[0].message.content == "hi"
for _ in range(50):
if await checks[0].dual_cache.async_get_cache(
key=next(iter(await _counter_keys(checks[0], request["metadata"]["tags"])))
):
break
await asyncio.sleep(0.05)
with pytest.raises(litellm.RateLimitError) as exc_info:
await router.acompletion(**request)
assert exc_info.value.detail["error"]["message"] == "tag_rate_limit_exceeded"
# a request without the configured tag id is unaffected
assert (
await router.acompletion(
model="gpt-4o",
messages=[{"role": "user", "content": "hello"}],
metadata={"tags": ["team_id:team-a"]},
)
).choices[0].message.content == "hi"
async def _counter_keys(check: TagLimitCheck, tags: List[str]) -> List[str]:
from litellm.router_strategy.tag_limits import _tag_values_from_request, _usage_key
tag_values = _tag_values_from_request(_request_kwargs(tags))
limits = get_tag_limits(_limits(limit=1))
assert limits is not None
return [
_usage_key(unit, limit, tag_values[limit.tag_id], check.now())
for unit, limit in limits.limits_by_unit()
if limit.tag_id in tag_values
]