fix(proxy): refresh a concurrency key's Redis TTL on every admission, not just its first

TAG_RL_CHECK_AND_INCR_SCRIPT only called EXPIRE when a key had no TTL at all,
so a concurrency counter's expiry was fixed from its first admission and never
pushed out by later ones. A concurrency bucket isn't epoch-windowed like
requests/tokens/dollars -- its TTL exists purely as a crash-safety net for a
reservation whose explicit release never runs -- so a still-active bucket
under sustained traffic would expire mid-flight, silently admitting past the
cap and letting a later release decrement an unrelated, newer cohort's
counter. Adds a refresh_ttl script argument, true only for the concurrency
caller, and verified against a real Redis instance since the in-memory
fallback (which already refreshes unconditionally) can't reproduce this.
bugbot caught this on review.
This commit is contained in:
Deepanshu 2026-08-26 15:08:56 -04:00
parent 54a7f6f7b3
commit 513c273750
3 changed files with 88 additions and 23 deletions

View file

@ -934,12 +934,16 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
return built
async def _check_and_increment_one(
self, cache: InternalUsageCache, key: str, limit: float, increment: float, ttl: int
self, cache: InternalUsageCache, key: str, limit: float, increment: float, ttl: int, refresh_ttl: bool
) -> tuple[bool, float]:
"""Single-key atomic check-and-increment. Always one key per Lua
call -- see TAG_RL_CHECK_AND_INCR_SCRIPT's module docstring for why."""
call -- see TAG_RL_CHECK_AND_INCR_SCRIPT's module docstring for why,
and for why `refresh_ttl` must be True for a concurrency key and
False for a requests key."""
if self._check_and_incr_script is not None:
raw: Final = await self._check_and_incr_script(keys=(key,), args=(limit, increment, ttl))
raw: Final = await self._check_and_incr_script(
keys=(key,), args=(limit, increment, ttl, 1 if refresh_ttl else 0)
)
return bool(raw[0]), float(raw[1])
async with self._lock:
@ -962,7 +966,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
async def _atomic_check_and_increment(
self,
checks: Sequence[tuple[InternalUsageCache, str, float, float, int]],
checks: Sequence[tuple[InternalUsageCache, str, float, float, int, bool]],
) -> tuple[int | None, tuple[float, ...]]:
"""
All-or-nothing across every (cache, key, limit, increment, ttl) in
@ -1019,10 +1023,10 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
# accumulated so far in favor of refunding and returning early, so
# this can't be expressed as a one-shot comprehension.
admitted_values: Final = [] # mutable-ok: sequential async accumulator, discardable on early rejection; see comment above
for index, (cache, key, limit, increment, ttl) in enumerate(checks):
for index, (cache, key, limit, increment, ttl, refresh_ttl) in enumerate(checks):
admitted = False
try:
admitted, value = await self._check_and_increment_one(cache, key, limit, increment, ttl)
admitted, value = await self._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl)
finally:
# Runs on a normal rejection (admitted stays False) and on
# any exception/cancellation from the awaited call above
@ -1040,10 +1044,10 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
return None, tuple(admitted_values)
async def _refund_admitted(
self, checks: Sequence[tuple[InternalUsageCache, str, float, float, int]], up_to_index: int
self, checks: Sequence[tuple[InternalUsageCache, str, float, float, int, bool]], up_to_index: int
) -> None:
for refund_index in range(up_to_index):
refund_cache, refund_key, _limit, refund_increment, _ttl = checks[refund_index]
refund_cache, refund_key, _limit, refund_increment, _ttl, _refresh_ttl = checks[refund_index]
try:
await self._decrement_floor_zero(refund_cache, refund_key, -refund_increment)
except Exception as e: # noqa: BLE001 - one failed refund must not block refunding the rest
@ -1150,6 +1154,7 @@ class _PROXY_ModelBasedTagRateLimitsHook( # pyright: ignore[reportUnusedClass]
# with nothing to replace it.
0.0 if configured_limit.unit == "requests" and key in stale_request_keys else 1.0,
self._ttl_for(configured_limit),
configured_limit.unit == "concurrency",
)
for partition, (configured_limit, _tag_value, key) in zip(atomic_partitions, atomic_checks)
)

View file

@ -93,19 +93,36 @@ CONCURRENCY_MIN_SAFETY_TTL_SECONDS: Final = 3600
# `atomic_check_and_increment_by_n` in parallel_request_limiter_v3.py, applied
# per-key instead of per-descriptor since each key already is one hash-tag
# group by construction.
#
# refresh_ttl (ARGV[4]) distinguishes the two callers of this script:
# "requests" is an epoch-bucketed fixed window, whose TTL must be set once
# (at first write) and never extended, or the bucket outlives the epoch it's
# meant to reset at. "concurrency" is not windowed at all -- its TTL exists
# purely as a crash-safety net for a reservation whose explicit release never
# runs -- so a still-active bucket must keep pushing that TTL out on every
# admission, or a long-lived burst of continuous traffic expires the whole
# counter mid-flight (silently admitting past the cap, and letting a release
# for a since-reset counter decrement an unrelated, newer cohort).
TAG_RL_CHECK_AND_INCR_SCRIPT: Final = """
local key = KEYS[1]
local limit = tonumber(ARGV[1])
local increment = tonumber(ARGV[2])
local ttl = tonumber(ARGV[3])
local refresh_ttl = tonumber(ARGV[4])
local current = tonumber(redis.call('GET', key) or 0)
if current + increment > limit then
return { 0, current }
end
local new_value = redis.call('INCRBY', key, increment)
local current_ttl = redis.call('TTL', key)
if current_ttl == -1 and ttl > 0 then
redis.call('EXPIRE', key, ttl)
if ttl > 0 then
if refresh_ttl == 1 then
redis.call('EXPIRE', key, ttl)
else
local current_ttl = redis.call('TTL', key)
if current_ttl == -1 then
redis.call('EXPIRE', key, ttl)
end
end
end
return { 1, new_value }
"""

View file

@ -3340,6 +3340,49 @@ async def test_redis_backed_token_admission_sees_increments_the_in_memory_cache_
await redis_cache.async_delete_cache(key=token_key)
@pytest.mark.asyncio
async def test_redis_backed_concurrency_ttl_refreshes_on_every_admission(time_controller):
"""
Bugbot finding: TAG_RL_CHECK_AND_INCR_SCRIPT only ran EXPIRE when a key
had no TTL at all, so a concurrency counter's expiry was fixed from its
first admission and never pushed out by later ones. A concurrency bucket
isn't epoch-windowed like requests/tokens/dollars -- its TTL exists only
as a crash-safety net for a reservation whose explicit release never
runs -- so a still-active bucket receiving continuous admissions must
keep extending that TTL, or it expires mid-flight under sustained
traffic, silently admitting past the cap.
"""
limiter, redis_cache = _redis_limiter(time_controller)
try:
await redis_cache.ping()
except Exception as e:
pytest.skip(f"Redis connection failed: {e!s}")
key = f"{{tag_rl:test:ttl-refresh:{uuid.uuid4().hex}}}:inflight"
cache = limiter.internal_usage_cache
try:
# A short, fixed ttl (bypassing _ttl_for's 3600s safety floor, which
# would make a real-time before/after comparison too slow to assert
# on deterministically) with refresh_ttl=True, matching how a
# concurrency check is actually admitted.
admitted, _ = await limiter._check_and_increment_one(cache, key, limit=100, increment=1.0, ttl=3, refresh_ttl=True)
assert admitted
ttl_after_first_admission = await redis_cache.redis_async_client.ttl(key)
assert ttl_after_first_admission > 0
await asyncio.sleep(2)
# A second admission on the same still-live key, most of the way
# through the first admission's ttl, must push the ttl back out to
# the full window again, not leave it counting down toward zero.
admitted, _ = await limiter._check_and_increment_one(cache, key, limit=100, increment=1.0, ttl=3, refresh_ttl=True)
assert admitted
ttl_after_second_admission = await redis_cache.redis_async_client.ttl(key)
assert ttl_after_second_admission >= 2
finally:
await redis_cache.async_delete_cache(key=key)
# ---------------------------------------------------------------------------
# team_public_model_name alias -- index lookup must not miss
# ---------------------------------------------------------------------------
@ -3843,9 +3886,9 @@ async def test_refund_failure_on_one_key_does_not_block_others_or_raise(time_con
failing_index, values = await flaky._atomic_check_and_increment(
[
(flaky.internal_usage_cache, failing_key, 10.0, 1.0, 60),
(flaky.internal_usage_cache, other_key, 10.0, 1.0, 60),
(flaky.internal_usage_cache, rejecting_key, 0.0, 1.0, 60),
(flaky.internal_usage_cache, failing_key, 10.0, 1.0, 60, False),
(flaky.internal_usage_cache, other_key, 10.0, 1.0, 60, False),
(flaky.internal_usage_cache, rejecting_key, 0.0, 1.0, 60, False),
]
)
@ -3870,18 +3913,18 @@ async def test_exception_mid_batch_refunds_every_earlier_admission_before_propag
raising_key = "{tag_rl:test:exception-refund:b}:requests"
class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook):
async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int):
async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int, refresh_ttl: bool):
if key == raising_key:
raise RuntimeError("simulated transient redis failure")
return await super()._check_and_increment_one(cache, key, limit, increment, ttl)
return await super()._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl)
flaky = _FlakyLimiter(internal_usage_cache=DualCache(), time_provider=time_controller.now)
with pytest.raises(RuntimeError):
await flaky._atomic_check_and_increment(
[
(flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60),
(flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60),
(flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60, False),
(flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60, False),
]
)
@ -3908,22 +3951,22 @@ async def test_a_raising_keys_own_ambiguous_outcome_is_never_refunded(time_contr
raising_key = "{tag_rl:test:ambiguous-no-refund:b}:requests"
class _FlakyLimiter(_PROXY_ModelBasedTagRateLimitsHook):
async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int):
async def _check_and_increment_one(self, cache, key: str, limit: float, increment: float, ttl: int, refresh_ttl: bool):
if key == raising_key:
# Simulate Redis committing the increment before the
# response is lost: the write actually happens...
await super()._check_and_increment_one(cache, key, limit, increment, ttl)
await super()._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl)
# ...but the caller never finds out.
raise RuntimeError("simulated lost response after a committed redis write")
return await super()._check_and_increment_one(cache, key, limit, increment, ttl)
return await super()._check_and_increment_one(cache, key, limit, increment, ttl, refresh_ttl)
flaky = _FlakyLimiter(internal_usage_cache=DualCache(), time_provider=time_controller.now)
with pytest.raises(RuntimeError):
await flaky._atomic_check_and_increment(
[
(flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60),
(flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60),
(flaky.internal_usage_cache, admitted_key, 10.0, 1.0, 60, False),
(flaky.internal_usage_cache, raising_key, 10.0, 1.0, 60, False),
]
)