mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
8b03315ac6
commit
e579749d89
6 changed files with 568 additions and 58 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
242
litellm/router_strategy/tag_limits.py
Normal file
242
litellm/router_strategy/tag_limits.py
Normal 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),
|
||||
)
|
||||
|
|
@ -790,6 +790,7 @@ OptionalPreCallChecks = List[
|
|||
"forward_client_headers_by_model_group",
|
||||
"enforce_model_rate_limits",
|
||||
"encrypted_content_affinity",
|
||||
"tag_limits",
|
||||
]
|
||||
]
|
||||
|
||||
|
|
|
|||
47
litellm/types/tag_limits.py
Normal file
47
litellm/types/tag_limits.py
Normal 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)
|
||||
252
tests/test_litellm/test_router/test_tag_limits.py
Normal file
252
tests/test_litellm/test_router/test_tag_limits.py
Normal 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
|
||||
]
|
||||
Loading…
Add table
Reference in a new issue