mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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 <songkuan-zheng@users.noreply.github.com>
This commit is contained in:
parent
6fcc2f602a
commit
30ee14620d
2 changed files with 27 additions and 9 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue