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
This commit is contained in:
Krish Dholakia 2025-05-01 22:09:26 -07:00 committed by GitHub
parent a4c96d5224
commit 132bdb1380
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 165 additions and 51 deletions

View File

@ -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(

View File

@ -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
)