fix(parallel_request_limiter_v2.py): update tpm tracking to use slot key logic

This commit is contained in:
Krrish Dholakia 2025-05-31 00:13:35 -07:00
parent 732dffa6b1
commit 808d7cc31c
3 changed files with 71 additions and 32 deletions

View file

@ -431,8 +431,12 @@ class _PROXY_MaxParallelRequestsHandler_v2(BaseRoutingStrategy, CustomLogger):
rate_limit_types = ["key", "user", "customer", "team", "model_per_key"]
current_time = datetime.now()
current_slot = (current_time.minute * 60 + current_time.second) // 15
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_time.hour:02d}-{current_slot}"
current_hour = current_time.hour
current_minute = current_time.minute
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot}"
for rate_limit_type in rate_limit_types:
for group in ["request_count", "rpm", "tpm"]:
key = self._get_current_usage_key(

View file

@ -2729,7 +2729,7 @@ class ProxyConfig:
"""
await self._init_guardrails_in_db(prisma_client=prisma_client)
await self._init_vector_stores_in_db(prisma_client=prisma_client)
await self._init_mcp_servers_in_db()
# await self._init_mcp_servers_in_db()
async def _init_guardrails_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.guardrails.guardrail_registry import (

View file

@ -66,13 +66,17 @@ async def test_normal_router_call_v2(monkeypatch):
user_api_key_dict=user_api_key_dict, cache=local_cache, data={}, call_type=""
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
current_time = datetime.now()
current_hour = current_time.hour
current_minute = current_time.minute
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot}"
print(f"slot_key: {slot_key}")
request_count_api_key = parallel_request_handler._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=precise_minute,
precise_minute=slot_key,
model=None,
rate_limit_type="key",
group="request_count",
@ -175,17 +179,22 @@ async def test_normal_router_call_tpm(monkeypatch, rate_limit_object):
call_type="",
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
current_time = datetime.now()
current_hour = current_time.hour
current_minute = current_time.minute
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot}"
print(f"slot_key: {slot_key}")
request_count_api_key = parallel_request_handler._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=precise_minute,
precise_minute=slot_key,
model="azure-model",
rate_limit_type=rate_limit_object,
group="tpm",
)
print(f"request_count_api_key: {request_count_api_key}")
await asyncio.sleep(1)
assert (
parallel_request_handler.internal_usage_cache.get_cache(
@ -210,11 +219,26 @@ async def test_normal_router_call_tpm(monkeypatch, rate_limit_object):
print(f"request_count_api_key: {request_count_api_key}")
next_slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot + 1 if current_slot < 3 else 0}"
request_count_api_key_next_slot = parallel_request_handler._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=next_slot_key,
model="azure-model",
rate_limit_type=rate_limit_object,
group="tpm",
)
## check if current slot matches response.usage.total_tokens else next slot
current_slot_get_cache = parallel_request_handler.internal_usage_cache.get_cache(
key=request_count_api_key
)
next_slot_get_cache = parallel_request_handler.internal_usage_cache.get_cache(
key=request_count_api_key_next_slot
)
assert (
parallel_request_handler.internal_usage_cache.get_cache(
key=request_count_api_key
)
== response.usage.total_tokens
current_slot_get_cache == response.usage.total_tokens
or next_slot_get_cache == response.usage.total_tokens
)
@ -290,18 +314,22 @@ async def test_normal_router_call_rpm(monkeypatch, rate_limit_object):
call_type="",
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
current_time = datetime.now()
current_hour = current_time.hour
current_minute = current_time.minute
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot}"
request_count_api_key = parallel_request_handler._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=precise_minute,
precise_minute=slot_key,
model="azure-model",
rate_limit_type=rate_limit_object,
group="rpm",
)
await asyncio.sleep(1)
assert (
parallel_request_handler.internal_usage_cache.get_cache(
key=request_count_api_key
@ -391,13 +419,17 @@ async def test_streaming_router_call_v2(monkeypatch):
user_api_key_dict=user_api_key_dict, cache=local_cache, data={}, call_type=""
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
current_time = datetime.now()
current_hour = current_time.hour
current_minute = current_time.minute
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot}"
request_count_api_key = parallel_request_handler._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=precise_minute,
precise_minute=slot_key,
model=None,
rate_limit_type="key",
group="request_count",
@ -494,13 +526,16 @@ async def test_bad_router_call_v2(monkeypatch, rate_limit_object):
user_api_key_dict=user_api_key_dict, cache=local_cache, data={}, call_type=""
)
current_date = datetime.now().strftime("%Y-%m-%d")
current_hour = datetime.now().strftime("%H")
current_minute = datetime.now().strftime("%M")
precise_minute = f"{current_date}-{current_hour}-{current_minute}"
current_time = datetime.now()
current_hour = current_time.hour
current_minute = current_time.minute
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{current_hour:02d}-{current_minute:02d}-{current_slot}"
request_count_api_key = parallel_request_handler._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=precise_minute,
precise_minute=slot_key,
model=None,
rate_limit_type=rate_limit_object,
group="rpm",