mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(parallel_request_limiter_v2.py): update tpm tracking to use slot key logic
This commit is contained in:
parent
732dffa6b1
commit
808d7cc31c
3 changed files with 71 additions and 32 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue