From 6fcc2f602a96ebeb137030caa284f2f91e840dc3 Mon Sep 17 00:00:00 2001 From: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com> Date: Fri, 11 Sep 2026 10:55:13 +0000 Subject: [PATCH 1/3] 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 --- litellm/caching/redis_cache.py | 26 ++++---- tests/unit/caching/test_redis_cache.py | 82 ++++++++++++++++++++++++-- 2 files changed, 90 insertions(+), 18 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 29e390b1d9a..0e0284ade5b 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -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( diff --git a/tests/unit/caching/test_redis_cache.py b/tests/unit/caching/test_redis_cache.py index 5f83be7c7bc..43c7d56072c 100644 --- a/tests/unit/caching/test_redis_cache.py +++ b/tests/unit/caching/test_redis_cache.py @@ -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") From 30ee14620d7fb43bac0ccde38ca7a8f721b985cc Mon Sep 17 00:00:00 2001 From: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com> Date: Fri, 11 Sep 2026 11:50:39 +0000 Subject: [PATCH 2/3] fix(caching): keep the exact TTL semantics in the atomic increment The Lua script skipped EXPIRE for any non-positive TTL, while the old code only skipped it when get_ttl returned None and otherwise passed the value through, so EXPIRE 0 still deleted the key. The TTL now travels as an empty string for None and as the literal value otherwise, and the script only branches on that emptiness Co-authored-by: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com> Co-authored-by: songkuan-zheng --- litellm/caching/redis_cache.py | 8 ++++---- tests/unit/caching/test_redis_cache.py | 28 +++++++++++++++++++++----- 2 files changed, 27 insertions(+), 9 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 0e0284ade5b..49f00790978 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -105,9 +105,8 @@ _INCREMENT_WITH_FLOOR_LUA: Final = ( _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 " + "if ARGV[2] ~= '' and (ARGV[3] == '1' or redis.call('TTL', KEYS[1]) == -1) then " + "redis.call('EXPIRE', KEYS[1], tonumber(ARGV[2])) end " "return value" ) @@ -1410,8 +1409,9 @@ class RedisCache(BaseCache): async def _incrbyfloat_with_ttl( _redis_client: "Redis", key: str, value: float, ttl: int | None, refresh_ttl: bool ) -> float: + ttl_arg: Final = "" if ttl is None else str(ttl) raw_value: Final = await _redis_client.eval( - _INCREMENT_WITH_TTL_LUA, 1, key, value, ttl or 0, "1" if refresh_ttl else "0" + _INCREMENT_WITH_TTL_LUA, 1, key, value, ttl_arg, "1" if refresh_ttl else "0" ) return _LUA_FLOAT.validate_python(raw_value) diff --git a/tests/unit/caching/test_redis_cache.py b/tests/unit/caching/test_redis_cache.py index 43c7d56072c..9378a770bd1 100644 --- a/tests/unit/caching/test_redis_cache.py +++ b/tests/unit/caching/test_redis_cache.py @@ -1400,11 +1400,11 @@ 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: + async def eval(self, script: str, numkeys: int, key: str, amount: object, ttl_arg: str, 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 + if ttl_arg != "" and (refresh == "1" or self.ttls.get(key) == -1): + self.ttls[key] = int(ttl_arg) return str(value).encode() def _incr(self, name: str, amount: float) -> float: @@ -1614,7 +1614,7 @@ async def test_redis_cache_async_increment_arms_ttl_in_the_same_command( 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") + assert (numkeys, key, amount, ttl, refresh) == (1, expected_key, 0.25, "30", "0") @pytest.mark.asyncio @@ -1627,4 +1627,22 @@ async def test_redis_cache_async_increment_refresh_ttl_sends_refresh_flag(monkey assert result == 2.5 assert spy.commands == ["eval"] - assert spy.eval_calls[0][3:] == (0.5, 60, "1") + assert spy.eval_calls[0][3:] == (0.5, "60", "1") + + +@pytest.mark.parametrize( + ("ttl", "default_ttl", "expected_ttl_arg"), [(None, None, ""), (0, None, "0"), (None, 15, "15")] +) +@pytest.mark.asyncio +async def test_redis_cache_async_increment_forwards_ttl_exactly( + ttl, default_ttl, expected_ttl_arg, monkeypatch, redis_no_ping +): + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + spy = _SpyRedisCommands(eval_result=b"0.75") + redis_cache = _SpyRedisCache(spy) + redis_cache.default_ttl = default_ttl + + result = await redis_cache.async_increment(key="spend:key:abc", value=0.75, ttl=ttl) + + assert result == 0.75 + assert spy.eval_calls[0][4] == expected_ttl_arg From 28cb1070ec8f6b53c128c5da4ceb6a1e209814b1 Mon Sep 17 00:00:00 2001 From: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com> Date: Sun, 13 Sep 2026 06:58:08 +0000 Subject: [PATCH 3/3] test(caching): cancel an increment mid-reply against a real Redis The test relays Redis through a local socket that delays every reply, cancels the increment task while the reply is in flight, and asserts the counter still carries its TTL. It fails against the previous two-step implementation, where the same cancellation leaves TTL -1 Co-authored-by: songkuan-zheng <252822057+songkuan-zheng@users.noreply.github.com> Co-authored-by: songkuan-zheng --- tests/local_testing/test_caching.py | 52 +++++++++++++++++++++++++++++ 1 file changed, 52 insertions(+) diff --git a/tests/local_testing/test_caching.py b/tests/local_testing/test_caching.py index 3e96896f47f..ae400edbe57 100644 --- a/tests/local_testing/test_caching.py +++ b/tests/local_testing/test_caching.py @@ -3018,3 +3018,55 @@ async def test_cache_key_in_hidden_params_acompletion(): assert response1.id == response2.id litellm.cache = None + + +@pytest.mark.asyncio +async def test_redis_async_increment_keeps_ttl_when_client_cancels_mid_reply(): + """A counter increment cancelled while Redis is replying still carries its TTL.""" + from redis.asyncio import Redis + + from litellm.caching.redis_cache import RedisCache + + redis_host = os.environ["REDIS_HOST"] + redis_port = int(os.environ.get("REDIS_PORT", "6379")) + redis_password = os.environ.get("REDIS_PASSWORD") + + async def pump(src: asyncio.StreamReader, dst: asyncio.StreamWriter, delay: float) -> None: + while data := await src.read(65536): + await asyncio.sleep(delay) + dst.write(data) + await dst.drain() + + async def relay(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + up_reader, up_writer = await asyncio.open_connection(redis_host, redis_port) + pumps = [ + asyncio.ensure_future(pump(reader, up_writer, 0.0)), + asyncio.ensure_future(pump(up_reader, writer, 0.3)), + ] + await asyncio.wait(pumps, return_when=asyncio.FIRST_COMPLETED) + for pending in pumps: + pending.cancel() + up_writer.close() + writer.close() + + server = await asyncio.start_server(relay, "127.0.0.1", 0) + relay_port = server.sockets[0].getsockname()[1] + cache = RedisCache(host="127.0.0.1", port=relay_port, password=redis_password, namespace="cancel-probe") + key = f"spend:key:{uuid.uuid4()}" + admin = Redis(host=redis_host, port=redis_port, password=redis_password) + try: + await cache.init_async_client().ping() + task = asyncio.ensure_future(cache.async_increment(key=key, value=0.25, ttl=60, refresh_ttl=False)) + await asyncio.sleep(0.15) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + await asyncio.sleep(0.5) + + assert await admin.get(f"cancel-probe:{key}") == b"0.25" + assert 0 < await admin.ttl(f"cancel-probe:{key}") <= 60 + finally: + await admin.delete(f"cancel-probe:{key}") + await admin.aclose() + await cache.init_async_client().aclose() + server.close()