From e579749d892a6f06ba0c9cb0b3d515f9abf2d466 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 29 Jul 2026 19:15:26 +0000 Subject: [PATCH] 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 --- .../credential_migration.py | 73 ++--- litellm/router.py | 11 +- litellm/router_strategy/tag_limits.py | 242 +++++++++++++++++ litellm/types/router.py | 1 + litellm/types/tag_limits.py | 47 ++++ .../test_router/test_tag_limits.py | 252 ++++++++++++++++++ 6 files changed, 568 insertions(+), 58 deletions(-) create mode 100644 litellm/router_strategy/tag_limits.py create mode 100644 litellm/types/tag_limits.py create mode 100644 tests/test_litellm/test_router/test_tag_limits.py diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 4d51295f8dc..6f79a39c883 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 6336b12258b..fc99c7d0a2f 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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 diff --git a/litellm/router_strategy/tag_limits.py b/litellm/router_strategy/tag_limits.py new file mode 100644 index 00000000000..78c2a56032a --- /dev/null +++ b/litellm/router_strategy/tag_limits.py @@ -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 `:` 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), + ) diff --git a/litellm/types/router.py b/litellm/types/router.py index 28e4a8272e8..3b4bf217500 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -790,6 +790,7 @@ OptionalPreCallChecks = List[ "forward_client_headers_by_model_group", "enforce_model_rate_limits", "encrypted_content_affinity", + "tag_limits", ] ] diff --git a/litellm/types/tag_limits.py b/litellm/types/tag_limits.py new file mode 100644 index 00000000000..cb48a600b3b --- /dev/null +++ b/litellm/types/tag_limits.py @@ -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) diff --git a/tests/test_litellm/test_router/test_tag_limits.py b/tests/test_litellm/test_router/test_tag_limits.py new file mode 100644 index 00000000000..bd2f50f30de --- /dev/null +++ b/tests/test_litellm/test_router/test_tag_limits.py @@ -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 + ]