From 25733020ce4b9117d7ebe712bf197d8f59c3548f Mon Sep 17 00:00:00 2001 From: shivam Date: Fri, 4 Sep 2026 20:42:55 +0000 Subject: [PATCH] fix(rate_limiter): resolve model aliases and access groups for per-model limits Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/proxy/auth/auth_utils.py | 18 ++ .../hooks/parallel_request_limiter_v3.py | 153 +++++++++-------- .../proxy/auth/test_auth_utils.py | 12 ++ .../hooks/test_parallel_request_limiter_v3.py | 162 ++++++++++++++++++ 4 files changed, 276 insertions(+), 69 deletions(-) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index d6007a2d56e..ecdfaae27b0 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -1271,6 +1271,24 @@ def get_model_rate_limit_from_metadata( return None +def resolve_rate_limited_model_name( + requested_model: str, + configured_models: Collection[str], + llm_router: Router | None, +) -> str | None: + """The configured name that governs `requested_model`: itself, its model_group_alias target, + or a model access group serving it.""" + if requested_model in configured_models: + return requested_model + if llm_router is None: + return None + candidates: Final = ( + llm_router._get_model_from_alias(requested_model), + *llm_router.get_model_access_groups(model_name=requested_model), + ) + return next((c for c in candidates if c is not None and c in configured_models), None) + + def get_team_model_rpm_limit( user_api_key_dict: UserAPIKeyAuth, ) -> dict[str, int] | None: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 31437af7770..8231e722c50 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -37,6 +37,7 @@ from litellm.proxy.auth.auth_utils import ( get_estimated_output_tokens, get_key_tag_rpm_limit, get_model_rate_limit_from_metadata, + resolve_rate_limited_model_name, ) from litellm.proxy.auth.budget_throttle import throttled_limit from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body @@ -67,6 +68,7 @@ if TYPE_CHECKING: from opentelemetry.trace import Span as _Span from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache + from litellm.router import Router from litellm.types.agents import AgentResponse from litellm.types.caching import RedisPipelineIncrementOperation @@ -379,6 +381,12 @@ PROJECT_OTPM_DESCRIPTOR_KEY: Final = "model_per_project_otpm" PARALLEL_REQUEST_SLOT_TTL_SECONDS: Final = 3600 +def _proxy_llm_router() -> "Router | None": + from litellm.proxy.proxy_server import llm_router + + return llm_router + + CacheCounterValue: TypeAlias = int | float | str | bytes CacheCounterValues: TypeAlias = Sequence[CacheCounterValue | None] @@ -2277,35 +2285,31 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): or get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_tpm_limit") is not None ): - _tpm_limit_for_team_model: Final = ( + _tpm_limit_for_org_model: Final = ( get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_tpm_limit") or {} ) - _rpm_limit_for_team_model: Final = ( + _rpm_limit_for_org_model: Final = ( get_model_rate_limit_from_metadata(user_api_key_dict, "organization_metadata", "model_rpm_limit") or {} ) - - should_check_rate_limit = False - if requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model: - should_check_rate_limit = True - - if should_check_rate_limit: - model_specific_tpm_limit = None - model_specific_rpm_limit = None - if requested_model in _tpm_limit_for_team_model: - model_specific_tpm_limit = _tpm_limit_for_team_model[requested_model] - if requested_model in _rpm_limit_for_team_model: - model_specific_rpm_limit = _rpm_limit_for_team_model[requested_model] - descriptors.append( - RateLimitDescriptor( - key="model_per_organization", - value=f"{user_api_key_dict.org_id}:{requested_model}", - rate_limit={ - "requests_per_unit": model_specific_rpm_limit, - "tokens_per_unit": model_specific_tpm_limit, - "window_size": self.window_size, - }, - ) + if requested_model is not None: + configured: Final = _tpm_limit_for_org_model.keys() | _rpm_limit_for_org_model.keys() + limited_model: Final = resolve_rate_limited_model_name( + requested_model=requested_model, + configured_models=configured, + llm_router=_proxy_llm_router(), ) + if limited_model is not None: + descriptors.append( + RateLimitDescriptor( + key="model_per_organization", + value=f"{user_api_key_dict.org_id}:{limited_model}", + rate_limit={ + "requests_per_unit": _rpm_limit_for_org_model.get(limited_model), + "tokens_per_unit": _tpm_limit_for_org_model.get(limited_model), + "window_size": self.window_size, + }, + ) + ) return descriptors @@ -2340,22 +2344,23 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): _tpm_limit_for_key_model = _tpm_limit_for_key_model or {} _rpm_limit_for_key_model = _rpm_limit_for_key_model or {} - # Check if model has any rate limits configured - should_check_rate_limit: Final = ( - requested_model in _tpm_limit_for_key_model or requested_model in _rpm_limit_for_key_model + configured: Final = _tpm_limit_for_key_model.keys() | _rpm_limit_for_key_model.keys() + limited_model: Final = resolve_rate_limited_model_name( + requested_model=requested_model, + configured_models=configured, + llm_router=_proxy_llm_router(), ) - - if not should_check_rate_limit: + if limited_model is None: return # Get model-specific limits - model_specific_tpm_limit: Final[int | None] = _tpm_limit_for_key_model.get(requested_model) - model_specific_rpm_limit: Final[int | None] = _rpm_limit_for_key_model.get(requested_model) + model_specific_tpm_limit: Final[int | None] = _tpm_limit_for_key_model.get(limited_model) + model_specific_rpm_limit: Final[int | None] = _rpm_limit_for_key_model.get(limited_model) descriptors.append( RateLimitDescriptor( key="model_per_key", - value=f"{user_api_key_dict.api_key}:{requested_model}", + value=f"{user_api_key_dict.api_key}:{limited_model}", rate_limit={ "requests_per_unit": model_specific_rpm_limit, "tokens_per_unit": model_specific_tpm_limit, @@ -2900,24 +2905,25 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): _rpm_limit_for_team_model: Final = ( get_model_rate_limit_from_metadata(user_api_key_dict, "team_metadata", "model_rpm_limit") or {} ) - should_check_rate_limit: Final = ( - requested_model in _tpm_limit_for_team_model or requested_model in _rpm_limit_for_team_model - ) - - if should_check_rate_limit and requested_model is not None: - model_specific_tpm_limit: Final = _tpm_limit_for_team_model.get(requested_model) - model_specific_rpm_limit: Final = _rpm_limit_for_team_model.get(requested_model) - descriptors.append( - RateLimitDescriptor( - key="model_per_team", - value=f"{user_api_key_dict.team_id}:{requested_model}", - rate_limit={ - "requests_per_unit": model_specific_rpm_limit, - "tokens_per_unit": model_specific_tpm_limit, - "window_size": self.window_size, - }, - ) + if requested_model is not None: + configured: Final = _tpm_limit_for_team_model.keys() | _rpm_limit_for_team_model.keys() + limited_model: Final = resolve_rate_limited_model_name( + requested_model=requested_model, + configured_models=configured, + llm_router=_proxy_llm_router(), ) + if limited_model is not None: + descriptors.append( + RateLimitDescriptor( + key="model_per_team", + value=f"{user_api_key_dict.team_id}:{limited_model}", + rate_limit={ + "requests_per_unit": _rpm_limit_for_team_model.get(limited_model), + "tokens_per_unit": _tpm_limit_for_team_model.get(limited_model), + "window_size": self.window_size, + }, + ) + ) def _add_project_model_rate_limit_descriptor_from_metadata( self, @@ -2936,24 +2942,25 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): _rpm_limit_for_project_model: Final = ( get_model_rate_limit_from_metadata(user_api_key_dict, "project_metadata", "model_rpm_limit") or {} ) - should_check_rate_limit: Final = ( - requested_model in _tpm_limit_for_project_model or requested_model in _rpm_limit_for_project_model - ) - - if should_check_rate_limit and requested_model is not None: - model_specific_tpm_limit: Final = _tpm_limit_for_project_model.get(requested_model) - model_specific_rpm_limit: Final = _rpm_limit_for_project_model.get(requested_model) - descriptors.append( - RateLimitDescriptor( - key="model_per_project", - value=f"{user_api_key_dict.project_id}:{requested_model}", - rate_limit={ - "requests_per_unit": model_specific_rpm_limit, - "tokens_per_unit": model_specific_tpm_limit, - "window_size": self.window_size, - }, - ) + if requested_model is not None: + configured: Final = _tpm_limit_for_project_model.keys() | _rpm_limit_for_project_model.keys() + limited_model: Final = resolve_rate_limited_model_name( + requested_model=requested_model, + configured_models=configured, + llm_router=_proxy_llm_router(), ) + if limited_model is not None: + descriptors.append( + RateLimitDescriptor( + key="model_per_project", + value=f"{user_api_key_dict.project_id}:{limited_model}", + rate_limit={ + "requests_per_unit": _rpm_limit_for_project_model.get(limited_model), + "tokens_per_unit": _tpm_limit_for_project_model.get(limited_model), + "window_size": self.window_size, + }, + ) + ) def add_project_io_token_rate_limit_descriptors_from_metadata( self, @@ -2979,13 +2986,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): or {} # mutable-ok: metadata helper returns an optional mapping ) - model_itpm_limit: Final = itpm_limit_for_project_model.get(requested_model) - model_otpm_limit: Final = otpm_limit_for_project_model.get(requested_model) + configured: Final = itpm_limit_for_project_model.keys() | otpm_limit_for_project_model.keys() + limited_model: Final = resolve_rate_limited_model_name( + requested_model=requested_model, + configured_models=configured, + llm_router=_proxy_llm_router(), + ) + if limited_model is None: + return + model_itpm_limit: Final = itpm_limit_for_project_model.get(limited_model) + model_otpm_limit: Final = otpm_limit_for_project_model.get(limited_model) if model_itpm_limit is None and model_otpm_limit is None: return - descriptor_value: Final = f"{user_api_key_dict.project_id}:{requested_model}" + descriptor_value: Final = f"{user_api_key_dict.project_id}:{limited_model}" if model_itpm_limit is not None: descriptors.append( RateLimitDescriptor( diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index f513f397b64..4b0401a06f4 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -26,9 +26,21 @@ from litellm.proxy.auth.auth_utils import ( get_project_model_tpm_limit, get_request_route_template, is_request_body_safe, + resolve_rate_limited_model_name, ) +@pytest.mark.parametrize( + ("requested_model", "configured_models", "expected"), + [ + ("gpt-4", {"gpt-4"}, "gpt-4"), + ("gpt4", {"gpt-4"}, None), + ], +) +def test_resolve_rate_limited_model_name_without_router(requested_model, configured_models, expected): + assert resolve_rate_limited_model_name(requested_model, configured_models, None) == expected + + class TestCustomAuthCommonChecksWarning: """custom_auth_common_checks_warning only warns when custom auth is configured and the common-checks opt-in is off, since that is the only state where diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 4003286d887..12b4cdf543e 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -51,6 +51,27 @@ class TimeController: self._current += timedelta(seconds=seconds) +@pytest.fixture +def model_rate_limit_router() -> Router: + return Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": {"model": "openai/gpt-4"}, + "model_info": {"access_groups": ["gpt4-family"]}, + }, + { + "model_name": "gpt-4-turbo", + "litellm_params": {"model": "openai/gpt-4-turbo"}, + "model_info": {"access_groups": ["gpt4-family"]}, + }, + ], + model_group_alias={"gpt4": "gpt-4"}, + set_verbose=False, + num_retries=0, + ) + + @pytest.fixture def time_controller(monkeypatch): controller = TimeController() @@ -65,6 +86,147 @@ def _isolated_request_stash(): _request_stash.reset(token) +@pytest.mark.parametrize( + ("metadata", "requested_model", "expected_model", "expected_tokens"), + [ + ({"model_tpm_limit": {"gpt-4": 100}}, "gpt4", "gpt-4", 100), + ({"model_tpm_limit": {"gpt4-family": 100}}, "gpt-4-turbo", "gpt4-family", 100), + ( + {"model_tpm_limit": {"gpt-4": 100, "gpt4": 50}}, + "gpt4", + "gpt4", + 50, + ), + ({"model_tpm_limit": {"other-model": 100}}, "gpt4", None, None), + ], +) +def test_model_per_key_rate_limit_resolves_alias_and_access_group( + monkeypatch, + model_rate_limit_router, + metadata, + requested_model, + expected_model, + expected_tokens, +): + monkeypatch.setattr( + "litellm.proxy.hooks.parallel_request_limiter_v3._proxy_llm_router", + lambda: model_rate_limit_router, + ) + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + descriptors = handler._create_rate_limit_descriptors( + user_api_key_dict=UserAPIKeyAuth(api_key="sk-model-key", metadata=metadata), + data={"model": requested_model}, + rpm_limit_type=None, + tpm_limit_type=None, + model_has_failures=False, + ) + model_descriptor = next( + (descriptor for descriptor in descriptors if descriptor["key"] == "model_per_key"), + None, + ) + if expected_model is None: + assert model_descriptor is None + return + assert model_descriptor is not None + assert model_descriptor["value"] == f"{hash_token('sk-model-key')}:{expected_model}" + assert model_descriptor["rate_limit"]["tokens_per_unit"] == expected_tokens + + +def test_model_per_team_rate_limit_resolves_alias(monkeypatch, model_rate_limit_router): + monkeypatch.setattr( + "litellm.proxy.hooks.parallel_request_limiter_v3._proxy_llm_router", + lambda: model_rate_limit_router, + ) + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(DualCache()) + ) + descriptors = [] + handler._add_team_model_rate_limit_descriptor_from_metadata( + user_api_key_dict=UserAPIKeyAuth( + team_id="team-model-key", + team_metadata={"model_tpm_limit": {"gpt-4": 100}}, + ), + requested_model="gpt4", + descriptors=descriptors, + ) + assert descriptors == [ + { + "key": "model_per_team", + "value": "team-model-key:gpt-4", + "rate_limit": { + "requests_per_unit": None, + "tokens_per_unit": 100, + "window_size": handler.window_size, + }, + } + ] + + +@pytest.mark.asyncio +async def test_model_per_key_rate_limit_alias_shares_tpm_counter( + monkeypatch, model_rate_limit_router +): + monkeypatch.setattr( + "litellm.proxy.hooks.parallel_request_limiter_v3._proxy_llm_router", + lambda: model_rate_limit_router, + ) + api_key = hash_token("sk-model-counter") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + metadata={"model_tpm_limit": {"gpt-4": 10}}, + ) + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + first_request = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 5, + } + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data=first_request, + call_type="", + ) + await handler.async_log_success_event( + kwargs={ + "standard_logging_object": { + "metadata": {"user_api_key_hash": api_key}, + }, + "model": "gpt-4", + }, + response_obj=ModelResponse( + id="model-counter", + object="chat.completion", + created=int(datetime.now().timestamp()), + model="gpt-4", + usage=Usage(prompt_tokens=5, completion_tokens=5, total_tokens=10), + choices=[], + ), + start_time=datetime.now(), + end_time=datetime.now(), + ) + counter_key = handler.create_rate_limit_keys( + "model_per_key", + f"{api_key}:gpt-4", + "tokens", + ) + assert await local_cache.async_get_cache(key=counter_key) == 10 + + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt4"}, + call_type="", + ) + assert exc_info.value.status_code == 429 + + @pytest.mark.parametrize( "throttle_pct, expected_rpm, expected_tpm", [