fix(caching): set the counter TTL atomically in async_increment

RedisCache.async_increment sent INCRBYFLOAT and then EXPIRE (or TTL plus
EXPIRE) as separate awaits. A task cancelled between the two, which happens
whenever the client disconnects mid request, left the counter without any
expiry, so a spend or rate limit counter could live forever in Redis

The increment and the TTL decision now run in one Lua call, following the
pattern async_increment_with_floor and async_set_max already use in this
file. The refresh_ttl semantics are unchanged: the flag re-arms the TTL on
every increment, otherwise only a key with no TTL gets one

Co-authored-by: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com>
Co-authored-by: songkuan-zheng <songkuan-zheng@users.noreply.github.com>
This commit is contained in:
songkuan-zheng 2026-09-11 10:55:13 +00:00
parent 04fa760bf2
commit 6fcc2f602a
2 changed files with 90 additions and 18 deletions

View file

@ -103,7 +103,16 @@ _INCREMENT_WITH_FLOOR_LUA: Final = (
"return count"
)
_INCREMENT_WITH_TTL_LUA: Final = (
"local value = redis.call('INCRBYFLOAT', KEYS[1], ARGV[1]) "
"local ttl = tonumber(ARGV[2]) "
"if ttl > 0 and (ARGV[3] == '1' or redis.call('TTL', KEYS[1]) == -1) then "
"redis.call('EXPIRE', KEYS[1], ttl) end "
"return value"
)
_LUA_COUNT: Final = TypeAdapter(int)
_LUA_FLOAT: Final = TypeAdapter(float)
_OPTIONAL_COUNTS: Final = TypeAdapter(tuple[int | None, ...])
@ -1401,19 +1410,10 @@ class RedisCache(BaseCache):
async def _incrbyfloat_with_ttl(
_redis_client: "Redis", key: str, value: float, ttl: int | None, refresh_ttl: bool
) -> float:
"""INCRBYFLOAT plus its TTL command in one round trip; a third only when an unexpiring key needs an EXPIRE."""
if ttl is None:
return await _redis_client.incrbyfloat(name=key, amount=value)
async with _redis_client.pipeline(transaction=False) as pipe:
pipe.incrbyfloat(name=key, amount=value)
if refresh_ttl:
pipe.expire(key, ttl)
else:
pipe.ttl(key)
result, ttl_or_expire = await pipe.execute()
if not refresh_ttl and ttl_or_expire == -1:
await _redis_client.expire(key, ttl)
return float(result)
raw_value: Final = await _redis_client.eval(
_INCREMENT_WITH_TTL_LUA, 1, key, value, ttl or 0, "1" if refresh_ttl else "0"
)
return _LUA_FLOAT.validate_python(raw_value)
@_redis_circuit_breaker_guard
async def async_increment(

View file

@ -1303,7 +1303,7 @@ async def test_write_path_timeouts_inside_the_interval_stay_at_debug(call_method
client = MagicMock()
client.pipeline.return_value.__aenter__.side_effect = timeout
client.sadd = AsyncMock(side_effect=timeout)
client.incrbyfloat = AsyncMock(side_effect=timeout)
client.eval = AsyncMock(side_effect=timeout)
client.rpush = AsyncMock(side_effect=timeout)
client.lpop = AsyncMock(side_effect=timeout)
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
@ -1400,6 +1400,13 @@ class _RoundTripCountingRedis:
self.ttls[name] = time
return True
async def eval(self, script: str, numkeys: int, key: str, amount: object, ttl_arg: int, refresh: str) -> bytes:
self.round_trips += 1
value = self._incr(key, float(amount)) # pyright: ignore[reportArgumentType] # fake receives the raw float
if ttl_arg > 0 and (refresh == "1" or self.ttls.get(key) == -1):
self.ttls[key] = ttl_arg
return str(value).encode()
def _incr(self, name: str, amount: float) -> float:
self.values[name] = self.values.get(name, 0.0) + amount
self.ttls.setdefault(name, self._initial_ttl)
@ -1446,12 +1453,12 @@ class _RoundTripCountingRedis:
@pytest.mark.parametrize(
("refresh_ttl", "existing_ttl", "expected_round_trips", "expected_ttl"),
[
pytest.param(True, 100, 2, 60, id="refresh_ttl: INCRBYFLOAT+EXPIRE in one round trip each"),
pytest.param(False, 100, 2, 100, id="keep ttl: INCRBYFLOAT+TTL in one round trip each, no EXPIRE"),
pytest.param(False, -1, 3, 60, id="unexpiring key: INCRBYFLOAT+TTL then EXPIRE once, 1 trip after"),
pytest.param(True, 100, 2, 60, id="refresh_ttl: one EVAL per increment, TTL re-armed"),
pytest.param(False, 100, 2, 100, id="keep ttl: one EVAL per increment, existing TTL kept"),
pytest.param(False, -1, 2, 60, id="unexpiring key: one EVAL per increment, TTL armed in the same call"),
],
)
async def test_async_increment_pipelines_the_ttl_command(
async def test_async_increment_sets_the_ttl_in_one_round_trip(
monkeypatch, redis_no_ping, refresh_ttl, existing_ttl, expected_round_trips, expected_ttl
):
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
@ -1556,3 +1563,68 @@ async def test_async_rpush_and_trim_runs_push_and_trim_in_one_transaction(monkey
assert pushed_len == 4
assert rows == ["b", "c", "d"]
assert pipe.queued == [("rpush", "ns:buf", "c", "d"), ("ltrim", "ns:buf", "-3", "-1")]
class _SpyRedisCommands:
def __init__(self, eval_result: object) -> None:
self.commands: list[str] = []
self.eval_calls: list[tuple[object, ...]] = []
self._eval_result = eval_result
async def eval(self, script: str, numkeys: int, *keys_and_args: object) -> object:
self.commands.append("eval")
self.eval_calls.append((script, numkeys, *keys_and_args))
return self._eval_result
async def incrbyfloat(self, name: str, amount: float) -> float:
self.commands.append("incrbyfloat")
return amount
async def ttl(self, name: str) -> int:
self.commands.append("ttl")
return -1
async def expire(self, name: str, time: int) -> bool:
self.commands.append("expire")
return True
class _SpyRedisCache(RedisCache):
def __init__(self, spy: _SpyRedisCommands, **kwargs: object) -> None:
super().__init__(**kwargs)
self.spy = spy
def init_async_client(self, *args: object, **kwargs: object) -> object:
return self.spy
@pytest.mark.parametrize(("namespace", "expected_key"), [(None, "spend:key:abc"), ("ns", "ns:spend:key:abc")])
@pytest.mark.asyncio
async def test_redis_cache_async_increment_arms_ttl_in_the_same_command(
namespace, expected_key, monkeypatch, redis_no_ping
):
"""The increment and its TTL reach Redis as one server-side step."""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
spy = _SpyRedisCommands(eval_result=b"1.25")
redis_cache = _SpyRedisCache(spy, namespace=namespace)
result = await redis_cache.async_increment(key="spend:key:abc", value=0.25, ttl=30)
assert result == 1.25
assert spy.commands == ["eval"]
script, numkeys, key, amount, ttl, refresh = spy.eval_calls[0]
assert "INCRBYFLOAT" in script and "EXPIRE" in script and "TTL" in script
assert (numkeys, key, amount, ttl, refresh) == (1, expected_key, 0.25, 30, "0")
@pytest.mark.asyncio
async def test_redis_cache_async_increment_refresh_ttl_sends_refresh_flag(monkeypatch, redis_no_ping):
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
spy = _SpyRedisCommands(eval_result="2.5")
redis_cache = _SpyRedisCache(spy)
result = await redis_cache.async_increment(key="spend:team:t1", value=0.5, refresh_ttl=True)
assert result == 2.5
assert spy.commands == ["eval"]
assert spy.eval_calls[0][3:] == (0.5, 60, "1")