mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge 28cb1070ec into 6f123b7083
This commit is contained in:
commit
3818b7faab
3 changed files with 160 additions and 18 deletions
|
|
@ -103,7 +103,15 @@ _INCREMENT_WITH_FLOOR_LUA: Final = (
|
|||
"return count"
|
||||
)
|
||||
|
||||
_INCREMENT_WITH_TTL_LUA: Final = (
|
||||
"local value = redis.call('INCRBYFLOAT', KEYS[1], ARGV[1]) "
|
||||
"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"
|
||||
)
|
||||
|
||||
_LUA_COUNT: Final = TypeAdapter(int)
|
||||
_LUA_FLOAT: Final = TypeAdapter(float)
|
||||
_OPTIONAL_COUNTS: Final = TypeAdapter(tuple[int | None, ...])
|
||||
|
||||
|
||||
|
|
@ -1401,19 +1409,11 @@ 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)
|
||||
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_arg, "1" if refresh_ttl else "0"
|
||||
)
|
||||
return _LUA_FLOAT.validate_python(raw_value)
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_increment(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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: 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 != "" 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:
|
||||
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,86 @@ 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")
|
||||
|
||||
|
||||
@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