From 132bdb13804fa212f8429db51b36ebf3032801ba Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Thu, 1 May 2025 22:09:26 -0700 Subject: [PATCH] Add user + team based multi-instance rate limiting (#10497) * fix(parallel_request_limiter_v2.py): add user multi-instance rate limiting * fix(parallel_request_limiter_v2.py): add user multi-instance rpm limiting * fix(parallel_request_limiter_v2.py): add team based multi-instance rate limiting --- .../parallel_request_limiter_v2.py | 92 +++++++++---- .../test_parallel_request_limiter_v2.py | 124 ++++++++++++++---- 2 files changed, 165 insertions(+), 51 deletions(-) diff --git a/enterprise/enterprise_hooks/parallel_request_limiter_v2.py b/enterprise/enterprise_hooks/parallel_request_limiter_v2.py index c3d3171307..0ca5d03f4d 100644 --- a/enterprise/enterprise_hooks/parallel_request_limiter_v2.py +++ b/enterprise/enterprise_hooks/parallel_request_limiter_v2.py @@ -3,6 +3,7 @@ V2 Implementation of Parallel Requests, TPM, RPM Limiting on the proxy Designed to work on a multi-instance setup, where multiple instances are writing to redis simultaneously """ +import asyncio import sys from datetime import datetime, timedelta from typing import ( @@ -51,6 +52,7 @@ class CacheObject(TypedDict): RateLimitGroups = Literal["request_count", "tpm", "rpm"] +RateLimitTypes = Literal["key", "model_per_key", "user", "customer", "team"] class _PROXY_MaxParallelRequestsHandler(BaseRoutingStrategy, CustomLogger): @@ -110,10 +112,10 @@ class _PROXY_MaxParallelRequestsHandler(BaseRoutingStrategy, CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, data: dict, - max_parallel_requests: int, + max_parallel_requests: Optional[int], precise_minute: str, - tpm_limit: int, - rpm_limit: int, + tpm_limit: Optional[int], + rpm_limit: Optional[int], rate_limit_type: Literal["key", "model_per_key", "user", "customer", "team"], ): ## INCREMENT CURRENT USAGE @@ -133,16 +135,23 @@ class _PROXY_MaxParallelRequestsHandler(BaseRoutingStrategy, CustomLogger): ) increment_list.append((key, increment_value_by_group[group])) + if ( + not max_parallel_requests and not rpm_limit and not tpm_limit + ): # no rate limits + return + results = await self._increment_value_list_in_current_window( increment_list=increment_list, ttl=60, ) - - if ( - results[0] > max_parallel_requests - or results[1] > rpm_limit - or results[2] > tpm_limit - ): + should_raise_error = False + if max_parallel_requests is not None: + should_raise_error = results[0] > max_parallel_requests + if rpm_limit is not None: + should_raise_error = should_raise_error or results[1] > rpm_limit + if tpm_limit is not None: + should_raise_error = should_raise_error or results[2] > tpm_limit + if should_raise_error: raise self.raise_rate_limit_error( additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current limits: max_parallel_requests: {max_parallel_requests}, tpm_limit: {tpm_limit}, rpm_limit: {rpm_limit}" ) @@ -231,17 +240,46 @@ class _PROXY_MaxParallelRequestsHandler(BaseRoutingStrategy, CustomLogger): current_minute = datetime.now().strftime("%M") precise_minute = f"{current_date}-{current_hour}-{current_minute}" + tasks = [] if api_key is not None: # CHECK IF REQUEST ALLOWED for key - await self.check_key_in_limits_v2( - user_api_key_dict=user_api_key_dict, - data=data, - max_parallel_requests=max_parallel_requests, - precise_minute=precise_minute, - tpm_limit=tpm_limit, - rpm_limit=rpm_limit, - rate_limit_type="key", + tasks.append( + self.check_key_in_limits_v2( + user_api_key_dict=user_api_key_dict, + data=data, + max_parallel_requests=max_parallel_requests, + precise_minute=precise_minute, + tpm_limit=tpm_limit, + rpm_limit=rpm_limit, + rate_limit_type="key", + ) ) + elif user_api_key_dict.user_id is not None: + # CHECK IF REQUEST ALLOWED for key + tasks.append( + self.check_key_in_limits_v2( + user_api_key_dict=user_api_key_dict, + data=data, + max_parallel_requests=None, + precise_minute=precise_minute, + tpm_limit=user_api_key_dict.user_tpm_limit, + rpm_limit=user_api_key_dict.user_rpm_limit, + rate_limit_type="user", + ) + ) + elif user_api_key_dict.team_id is not None: + tasks.append( + self.check_key_in_limits_v2( + user_api_key_dict=user_api_key_dict, + data=data, + max_parallel_requests=None, + precise_minute=precise_minute, + tpm_limit=user_api_key_dict.team_tpm_limit, + rpm_limit=user_api_key_dict.team_rpm_limit, + rate_limit_type="team", + ) + ) + await asyncio.gather(*tasks) return @@ -260,16 +298,18 @@ class _PROXY_MaxParallelRequestsHandler(BaseRoutingStrategy, CustomLogger): "rpm": 0, } - for group in ["request_count", "rpm", "tpm"]: - key = self._get_current_usage_key( - user_api_key_dict=user_api_key_dict, - precise_minute=precise_minute, - model=model, - rate_limit_type="key", - group=cast(RateLimitGroups, group), - ) + rate_limit_types = ["key", "user", "customer", "team"] + for rate_limit_type in rate_limit_types: + for group in ["request_count", "rpm", "tpm"]: + key = self._get_current_usage_key( + user_api_key_dict=user_api_key_dict, + precise_minute=precise_minute, + model=model, + rate_limit_type=cast(RateLimitTypes, rate_limit_type), + group=cast(RateLimitGroups, group), + ) - increment_list.append((key, increment_value_by_group[group])) + increment_list.append((key, increment_value_by_group[group])) if increment_list: # Only call if we have values to increment await self._increment_value_list_in_current_window( diff --git a/tests/enterprise/enterprise_hooks/test_parallel_request_limiter_v2.py b/tests/enterprise/enterprise_hooks/test_parallel_request_limiter_v2.py index ef5014f2dd..0a472db638 100644 --- a/tests/enterprise/enterprise_hooks/test_parallel_request_limiter_v2.py +++ b/tests/enterprise/enterprise_hooks/test_parallel_request_limiter_v2.py @@ -7,14 +7,18 @@ import sys from datetime import datetime import pytest +from fastapi import HTTPException import litellm +from enterprise.enterprise_hooks.parallel_request_limiter_v2 import ( + _PROXY_MaxParallelRequestsHandler, +) from litellm import Router from litellm.caching.caching import DualCache from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token -from enterprise.enterprise_hooks.parallel_request_limiter_v2 import _PROXY_MaxParallelRequestsHandler -from fastapi import HTTPException + + @pytest.mark.asyncio async def test_normal_router_call_v2(monkeypatch): """ @@ -52,7 +56,9 @@ async def test_normal_router_call_v2(monkeypatch): _api_key = hash_token(_api_key) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=1) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) monkeypatch.setattr(litellm, "callbacks", [parallel_request_handler]) await parallel_request_handler.async_pre_call_hook( @@ -96,8 +102,19 @@ async def test_normal_router_call_v2(monkeypatch): ) +@pytest.mark.parametrize( + "rate_limit_object", + [ + "key", + # "model_per_key", + "user", + # "customer", + "team", + ], +) +@pytest.mark.flaky(reruns=3) @pytest.mark.asyncio -async def test_normal_router_call_tpm(monkeypatch): +async def test_normal_router_call_tpm(monkeypatch, rate_limit_object): """ Test normal router call with parallel request limiter v2 """ @@ -131,9 +148,16 @@ async def test_normal_router_call_tpm(monkeypatch): _api_key = "sk-12345" _api_key = hash_token(_api_key) - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, tpm_limit=10) + if rate_limit_object == "key": + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, tpm_limit=10) + elif rate_limit_object == "user": + user_api_key_dict = UserAPIKeyAuth(user_id="12345", user_tpm_limit=10) + elif rate_limit_object == "team": + user_api_key_dict = UserAPIKeyAuth(team_id="12345", team_tpm_limit=10) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) monkeypatch.setattr(litellm, "callbacks", [parallel_request_handler]) await parallel_request_handler.async_pre_call_hook( @@ -148,7 +172,7 @@ async def test_normal_router_call_tpm(monkeypatch): user_api_key_dict=user_api_key_dict, precise_minute=precise_minute, model=None, - rate_limit_type="key", + rate_limit_type=rate_limit_object, group="tpm", ) await asyncio.sleep(1) @@ -163,7 +187,11 @@ async def test_normal_router_call_tpm(monkeypatch): response = await router.acompletion( model="azure-model", messages=[{"role": "user", "content": "Hey, how's it going?"}], - metadata={"user_api_key": _api_key}, + metadata={ + "user_api_key": _api_key, + "user_api_key_user_id": user_api_key_dict.user_id, + "user_api_key_team_id": user_api_key_dict.team_id, + }, mock_response="hello", ) await asyncio.sleep(1) # success is done in a separate thread @@ -177,8 +205,20 @@ async def test_normal_router_call_tpm(monkeypatch): == response.usage.total_tokens ) + +@pytest.mark.parametrize( + "rate_limit_object", + [ + "key", + # "model_per_key", + "user", + # "customer", + "team", + ], +) +@pytest.mark.flaky(reruns=3) @pytest.mark.asyncio -async def test_normal_router_call_rpm(monkeypatch): +async def test_normal_router_call_rpm(monkeypatch, rate_limit_object): """ Test normal router call with parallel request limiter v2 """ @@ -212,9 +252,16 @@ async def test_normal_router_call_rpm(monkeypatch): _api_key = "sk-12345" _api_key = hash_token(_api_key) - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, tpm_limit=10) + if rate_limit_object == "key": + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, rpm_limit=1) + elif rate_limit_object == "user": + user_api_key_dict = UserAPIKeyAuth(user_id="12345", user_rpm_limit=1) + elif rate_limit_object == "team": + user_api_key_dict = UserAPIKeyAuth(team_id="12345", team_rpm_limit=1) local_cache = DualCache() - parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) monkeypatch.setattr(litellm, "callbacks", [parallel_request_handler]) await parallel_request_handler.async_pre_call_hook( @@ -229,7 +276,7 @@ async def test_normal_router_call_rpm(monkeypatch): user_api_key_dict=user_api_key_dict, precise_minute=precise_minute, model=None, - rate_limit_type="key", + rate_limit_type=rate_limit_object, group="rpm", ) await asyncio.sleep(1) @@ -244,7 +291,11 @@ async def test_normal_router_call_rpm(monkeypatch): response = await router.acompletion( model="azure-model", messages=[{"role": "user", "content": "Hey, how's it going?"}], - metadata={"user_api_key": _api_key}, + metadata={ + "user_api_key": _api_key, + "user_api_key_user_id": user_api_key_dict.user_id, + "user_api_key_team_id": user_api_key_dict.team_id, + }, mock_response="hello", ) await asyncio.sleep(1) # success is done in a separate thread @@ -260,11 +311,13 @@ async def test_normal_router_call_rpm(monkeypatch): with pytest.raises(HTTPException): await parallel_request_handler.async_pre_call_hook( - user_api_key_dict=user_api_key_dict, cache=local_cache, data={}, call_type="" + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={}, + call_type="", ) - @pytest.mark.asyncio async def test_streaming_router_call_v2(monkeypatch): """ @@ -303,9 +356,11 @@ async def test_streaming_router_call_v2(monkeypatch): _api_key = hash_token(_api_key) user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=1) local_cache = DualCache() - + print(f"litellm callbacks pre-set: {litellm.callbacks}") - parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) monkeypatch.setattr(litellm, "callbacks", [parallel_request_handler]) await parallel_request_handler.async_pre_call_hook( @@ -351,8 +406,20 @@ async def test_streaming_router_call_v2(monkeypatch): == 0 ) + +@pytest.mark.parametrize( + "rate_limit_object", + [ + "key", + # "model_per_key", + "user", + # "customer", + "team", + ], +) +@pytest.mark.flaky(reruns=3) @pytest.mark.asyncio -async def test_bad_router_call_v2(monkeypatch): +async def test_bad_router_call_v2(monkeypatch, rate_limit_object): """ Test bad router call with parallel request limiter v2 """ @@ -386,10 +453,17 @@ async def test_bad_router_call_v2(monkeypatch): _api_key = "sk-12345" _api_key = hash_token(_api_key) - user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, max_parallel_requests=1) + if rate_limit_object == "key": + user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, rpm_limit=1) + elif rate_limit_object == "user": + user_api_key_dict = UserAPIKeyAuth(user_id="12345", user_rpm_limit=1) + elif rate_limit_object == "team": + user_api_key_dict = UserAPIKeyAuth(team_id="12345", team_rpm_limit=1) local_cache = DualCache() - - parallel_request_handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(local_cache)) + + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) monkeypatch.setattr(litellm, "callbacks", [parallel_request_handler]) await parallel_request_handler.async_pre_call_hook( @@ -404,8 +478,8 @@ async def test_bad_router_call_v2(monkeypatch): user_api_key_dict=user_api_key_dict, precise_minute=precise_minute, model=None, - rate_limit_type="key", - group="request_count", + rate_limit_type=rate_limit_object, + group="rpm", ) await asyncio.sleep(1) assert ( @@ -421,10 +495,10 @@ async def test_bad_router_call_v2(monkeypatch): original_exception=Exception("test"), user_api_key_dict=user_api_key_dict, ) - + assert ( parallel_request_handler.internal_usage_cache.get_cache( key=request_count_api_key ) - == 0 + == 1 )