fix(proxy): preserve rate limit windows during migration

This commit is contained in:
Sisyphus 2026-08-30 16:20:35 +08:00
parent d806d3b776
commit d15740e29b
2 changed files with 118 additions and 11 deletions

View file

@ -408,6 +408,11 @@ def descriptor_window_key(descriptor_key: str, descriptor_value: str, rate_limit
return f"{{{descriptor_key}:{descriptor_value}}}:window:{rate_limit_type}"
def legacy_descriptor_window_key(descriptor_key: str, descriptor_value: str) -> str:
"""Return the shared window key used before RPM and TPM were separated."""
return f"{{{descriptor_key}:{descriptor_value}}}:window"
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
requests_per_unit: int | None
tokens_per_unit: int | None
@ -1069,6 +1074,71 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return results
async def _backfill_legacy_window_keys(
self,
window_keys: Sequence[str],
parent_otel_span: Span | None = None,
) -> None:
"""Seed new per-type windows from the pre-split shared window key.
The counter keys are intentionally unchanged, so copying the old
window start preserves the active window during a rolling upgrade.
Redis uses ``NX`` so concurrent replicas cannot overwrite a window
that another replica has already initialized. The legacy key remains
readable until its normal TTL expires.
"""
redis_cache = self.internal_usage_cache.dual_cache.redis_cache
for window_key in window_keys:
if not window_key.endswith((":window:requests", ":window:tokens")):
continue
new_window_value = await self.internal_usage_cache.async_get_cache(
key=window_key,
litellm_parent_otel_span=parent_otel_span,
local_only=False,
)
if new_window_value is not None:
continue
legacy_window_key = f"{window_key.rsplit(':window:', 1)[0]}:window"
legacy_window_value = await self.internal_usage_cache.async_get_cache(
key=legacy_window_key,
litellm_parent_otel_span=parent_otel_span,
local_only=False,
)
if legacy_window_value is None:
continue
if redis_cache is None:
await self.internal_usage_cache.async_set_cache(
key=window_key,
value=legacy_window_value,
ttl=self.window_size,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
continue
inserted = await redis_cache.async_set_cache(
key=window_key,
value=legacy_window_value,
ttl=self.window_size,
nx=True,
parent_otel_span=parent_otel_span,
)
current_window_value = legacy_window_value if inserted else await redis_cache.async_get_cache(
key=window_key,
parent_otel_span=parent_otel_span,
)
if current_window_value is not None:
await self.internal_usage_cache.async_set_cache(
key=window_key,
value=current_window_value,
ttl=self.window_size,
litellm_parent_otel_span=parent_otel_span,
local_only=True,
)
def create_rate_limit_keys(
self,
key: str,
@ -1106,10 +1176,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
item_code = "OK"
window_key = keys_to_fetch[i]
counter_key = keys_to_fetch[i + 1]
counter_value = cache_values[i + 1]
counter_value: CacheCounterValue | None = cache_values[i + 1]
requests_limit = key_metadata[window_key]["requests_limit"]
tokens_limit = key_metadata[window_key]["tokens_limit"]
window_expired = False
if now_int is not None and counter_value is not None:
window_start = cache_values[i]
if window_start is not None:
@ -1117,11 +1188,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
window_expired = (now_int - int(window_start)) >= key_metadata[window_key]["window_size"]
except (TypeError, ValueError):
window_expired = False
if window_expired:
# This counter belongs to a window that has already
# rolled over; it must not reject the request that
# starts the new window.
counter_value = 0
# This counter belongs to a window that has already rolled over;
# it must not reject the request that starts the new window.
effective_counter_value: Final[CacheCounterValue | None] = 0 if window_expired else counter_value
# Determine which limit to use for current_limit and limit_remaining
current_limit: int | None = None
@ -1136,12 +1206,16 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if current_limit is None or rate_limit_type is None:
continue
if counter_value is not None and int(counter_value) > current_limit:
if effective_counter_value is not None and int(effective_counter_value) > current_limit:
overall_code = "OVER_LIMIT"
item_code = "OVER_LIMIT"
# Only compute limit_remaining if current_limit is not None
limit_remaining = current_limit - int(counter_value) if counter_value is not None else current_limit
limit_remaining = (
current_limit - int(effective_counter_value)
if effective_counter_value is not None
else current_limit
)
statuses.append(
{
@ -1317,6 +1391,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
windowed_response = RateLimitResponse(overall_code="OK", statuses=[])
if keys_to_fetch:
await self._backfill_legacy_window_keys(
window_keys=keys_to_fetch[::2],
parent_otel_span=parent_otel_span,
)
## CHECK IN-MEMORY CACHE
cache_values = await self._batch_get_counter_values( # rebind-ok: refreshed by the Redis read below when the in-memory pass is under limit
keys=keys_to_fetch,
@ -1343,7 +1421,8 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# For keys that don't exist yet, set them to 0
if cache_values is None:
cache_values = [ # rebind-ok: missing keys default to a zeroed window snapshot
str(now_int) if ":window" in key else 0 for key in keys_to_fetch
str(now_int) if key.endswith((":window:requests", ":window:tokens")) else 0
for key in keys_to_fetch
]
elif self.batch_rate_limiter_script is not None:
# NORMAL MODE: Increment counters in Redis
@ -1447,11 +1526,13 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
if requests_limit is not None:
rpm_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, "requests")
keys_to_fetch.extend((requests_window_key, rpm_key))
if requests_window_key is not None:
keys_to_fetch.extend((requests_window_key, rpm_key))
rate_limit_set = True
if tokens_limit is not None:
tpm_key = self.create_rate_limit_keys(descriptor_key, descriptor_value, "tokens")
keys_to_fetch.extend((tokens_window_key, tpm_key))
if tokens_window_key is not None:
keys_to_fetch.extend((tokens_window_key, tpm_key))
rate_limit_set = True
if not rate_limit_set:
@ -1760,6 +1841,11 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
if not descriptor_groups:
return RateLimitResponse(overall_code="OK", statuses=[])
await self._backfill_legacy_window_keys(
window_keys=[meta["window_key"] for _keys, _args, group_meta in descriptor_groups for meta in group_meta],
parent_otel_span=parent_otel_span,
)
# Multi-process atomicity via Redis Lua, per descriptor for slot
# co-location. Single-process atomicity falls back to the
# asyncio.Lock + in-memory sliding window below — there are no

View file

@ -29,6 +29,8 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
)
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
descriptor_window_key,
legacy_descriptor_window_key,
)
from litellm.proxy.utils import InternalUsageCache, ProxyLogging, hash_token
from litellm.types.caching import RedisPipelineIncrementOperation
@ -6406,3 +6408,22 @@ def test_is_cache_list_over_limit_ignores_counter_from_expired_window(time_contr
# Without now_int the legacy (window-blind) behavior is preserved.
response = parallel_request_handler.is_cache_list_over_limit(keys_to_fetch, ["0", 3], key_metadata)
assert response["overall_code"] == "OVER_LIMIT"
@pytest.mark.asyncio
async def test_new_window_keys_backfill_active_legacy_window(time_controller):
"""A rolling upgrade preserves an active legacy window in the new keys."""
local_cache = DualCache()
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache),
time_provider=time_controller.now,
)
descriptor_key = "api_key"
descriptor_value = "sk-24677-migration"
legacy_key = legacy_descriptor_window_key(descriptor_key, descriptor_value)
new_key = descriptor_window_key(descriptor_key, descriptor_value, "requests")
await local_cache.async_set_cache(key=legacy_key, value="100", ttl=60)
await parallel_request_handler._backfill_legacy_window_keys([new_key])
assert await local_cache.async_get_cache(key=new_key) == 100