feat(parallel_request_limiter_v2.py): add sliding window logic

allows rate limiting to work across minutes
This commit is contained in:
Krrish Dholakia 2025-05-30 16:26:04 -07:00
parent c9062ac044
commit f757489033
3 changed files with 135 additions and 17 deletions

File diff suppressed because one or more lines are too long

View file

@ -132,42 +132,95 @@ class _PROXY_MaxParallelRequestsHandler_v2(BaseRoutingStrategy, CustomLogger):
):
## INCREMENT CURRENT USAGE
increment_list: List[Tuple[str, int]] = []
slots_to_check: List[str] = []
increment_value_by_group = {
"request_count": 1,
"tpm": 0,
"rpm": 1,
}
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=data.get("model", None),
rate_limit_type=rate_limit_type,
group=cast(RateLimitGroups, group),
)
if key is None:
continue
increment_list.append((key, increment_value_by_group[group]))
# Get current time and calculate the last 4 15s slots
current_time = datetime.now()
current_slot = (
current_time.second // 15
) # This gives us 0-3 for the current 15s slot
slots_to_check = []
slot_cache_keys = []
# Calculate the last 4 slots, handling minute boundaries
for i in range(4):
slot_number = (current_slot - i) % 4 # This ensures we wrap around properly
minute = current_time.minute
hour = current_time.hour
# If we need to look at previous minute
if current_slot - i < 0:
if minute == 0:
# If we're at minute 0, go to previous hour
hour = (current_time.hour - 1) % 24
minute = 59
else:
minute = current_time.minute - 1
slot_key = f"{current_time.strftime('%Y-%m-%d')}-{hour:02d}-{minute:02d}-{slot_number}"
slots_to_check.append(slot_key)
# For each slot, create keys for all rate limit groups
for slot_key in slots_to_check:
for group in ["request_count", "rpm", "tpm"]:
key = self._get_current_usage_key(
user_api_key_dict=user_api_key_dict,
precise_minute=slot_key,
model=data.get("model", None),
rate_limit_type=rate_limit_type,
group=cast(RateLimitGroups, group),
)
if key is None:
continue
# Only increment the current slot
if slot_key == slots_to_check[0]:
increment_list.append((key, increment_value_by_group[group]))
slot_cache_keys.append(key)
if (
not max_parallel_requests and not rpm_limit and not tpm_limit
): # no rate limits
return
# Use the existing atomic increment-and-check functionality
results = await self._increment_value_list_in_current_window(
increment_list=increment_list,
ttl=60,
)
# Get the current values for all slots to check limits
current_values = await self.internal_usage_cache.async_batch_get_cache(
slot_cache_keys
)
if current_values is None:
current_values = [None] * len(slot_cache_keys)
# Calculate totals across all slots, handling None values
total_requests = sum(
v if v is not None else 0 for v in current_values[::3]
) # Every 3rd value is request_count
total_rpm = sum(
v if v is not None else 0 for v in current_values[1::3]
) # Every 3rd value is rpm
total_tpm = sum(
v if v is not None else 0 for v in current_values[2::3]
) # Every 3rd value is tpm
should_raise_error = False
if max_parallel_requests is not None:
should_raise_error = results[0] > max_parallel_requests
should_raise_error = total_requests > max_parallel_requests
if rpm_limit is not None:
should_raise_error = should_raise_error or results[1] > rpm_limit
should_raise_error = should_raise_error or total_rpm > rpm_limit
if tpm_limit is not None:
should_raise_error = should_raise_error or results[2] > tpm_limit
should_raise_error = should_raise_error or total_tpm > 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 usage: max_parallel_requests: {results[0]}, current_rpm: {results[1]}, current_tpm: {results[2]}. Current limits: max_parallel_requests: {max_parallel_requests}, rpm_limit: {rpm_limit}, tpm_limit: {tpm_limit}."
additional_details=f"{CommonProxyErrors.max_parallel_request_limit_reached.value}. Hit limit for {rate_limit_type}. Current usage: max_parallel_requests: {total_requests}, current_rpm: {total_rpm}, current_tpm: {total_tpm}. Current limits: max_parallel_requests: {max_parallel_requests}, rpm_limit: {rpm_limit}, tpm_limit: {tpm_limit}."
)
def time_to_next_minute(self) -> float:
@ -356,11 +409,14 @@ 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}"
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,
precise_minute=slot_key,
model=model,
rate_limit_type=cast(RateLimitTypes, rate_limit_type),
group=cast(RateLimitGroups, group),

View file

@ -526,3 +526,66 @@ async def test_bad_router_call_v2(monkeypatch, rate_limit_object):
)
== 1
)
@pytest.mark.asyncio
async def test_check_key_in_limits_v2_sliding_window():
"""
Test the check_key_in_limits_v2 function with sliding window logic
"""
print("Starting test")
_api_key = "sk-12345"
_api_key = hash_token(_api_key)
user_api_key_dict = UserAPIKeyAuth(api_key=_api_key, rpm_limit=2)
local_cache = DualCache()
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
print("Created handler")
# Get current time and calculate slots
current_time = datetime.now()
current_slot = (current_time.minute * 60 + current_time.second) // 15
current_slot_key = (
f"{current_time.strftime('%Y-%m-%d')}-{current_time.hour:02d}-{current_slot}"
)
print(f"Current slot key: {current_slot_key}")
print("Making first request")
# Test 1: First request should succeed
await parallel_request_handler.check_key_in_limits_v2(
user_api_key_dict=user_api_key_dict,
data={},
max_parallel_requests=None,
precise_minute=current_slot_key,
tpm_limit=None,
rpm_limit=3,
rate_limit_type="key",
)
print("First request completed")
print("Making second request")
# Test 2: Second request should succeed
await parallel_request_handler.check_key_in_limits_v2(
user_api_key_dict=user_api_key_dict,
data={},
max_parallel_requests=None,
precise_minute=current_slot_key,
tpm_limit=None,
rpm_limit=3,
rate_limit_type="key",
)
print("Second request completed")
print("Verifying cache")
# Make third request - should fail
with pytest.raises(HTTPException):
await parallel_request_handler.check_key_in_limits_v2(
user_api_key_dict=user_api_key_dict,
data={},
max_parallel_requests=None,
precise_minute=current_slot_key,
tpm_limit=None,
rpm_limit=2,
rate_limit_type="key",
)