From 5ffb099ee169d014e7c0992e6e6e83dd8b07aae6 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Sat, 9 May 2026 01:53:00 +0000 Subject: [PATCH] Enhance RedisCache async_increment method to use EXPIRE NX for improved performance. Added fallback mechanism for compatibility with Redis clients that do not support the nx argument. Updated tests to verify new behavior and fallback logic. --- litellm/caching/redis_cache.py | 13 +++-- .../test_litellm/caching/test_redis_cache.py | 49 +++++++++++++++++-- 2 files changed, 54 insertions(+), 8 deletions(-) diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index cb9ce475d30..a94bc6637e4 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -846,9 +846,16 @@ class RedisCache(BaseCache): if refresh_ttl: await _redis_client.expire(key, _used_ttl) else: - current_ttl = await _redis_client.ttl(key) - if current_ttl == -1: - await _redis_client.expire(key, _used_ttl) + # Prefer EXPIRE NX to avoid an extra TTL round trip on the + # hot increment path while preserving window semantics. + try: + await _redis_client.expire(key, _used_ttl, nx=True) + except TypeError: + # Backward-compatible fallback if the Redis client does + # not support the "nx" kwarg on expire(). + current_ttl = await _redis_client.ttl(key) + if current_ttl == -1: + await _redis_client.expire(key, _used_ttl) ## LOGGING ## end_time = time.time() diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index 78192400fb0..488e8ce0daa 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -71,27 +71,66 @@ async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl( ) mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60) + mock_redis_instance.ttl.assert_not_awaited() @pytest.mark.asyncio -async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl( +async def test_redis_cache_async_increment_default_uses_expire_nx( monkeypatch, redis_no_ping ): - """Default (refresh_ttl=False) preserves window-style semantics: TTL is - set only on first creation, never refreshed (used by rate-limit windows).""" + """Default (refresh_ttl=False) uses EXPIRE NX to preserve window-style + semantics in a single RTT (no explicit TTL read).""" monkeypatch.setenv("REDIS_HOST", "https://my-test-host") redis_cache = RedisCache() mock_redis_instance = AsyncMock() mock_redis_instance.__aenter__.return_value = mock_redis_instance mock_redis_instance.__aexit__.return_value = None - mock_redis_instance.ttl.return_value = 42 # key already has ~42s left with patch.object( redis_cache, "init_async_client", return_value=mock_redis_instance ): await redis_cache.async_increment(key="rate_limit:window", value=1) - mock_redis_instance.expire.assert_not_awaited() + mock_redis_instance.expire.assert_awaited_once_with( + "rate_limit:window", 60, nx=True + ) + mock_redis_instance.ttl.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_redis_cache_async_increment_default_falls_back_when_expire_nx_unsupported( + monkeypatch, redis_no_ping +): + """If expire(nx=True) is unsupported by the client, fallback to + TTL-check + conditional EXPIRE for compatibility.""" + monkeypatch.setenv("REDIS_HOST", "https://my-test-host") + redis_cache = RedisCache() + mock_redis_instance = AsyncMock() + mock_redis_instance.__aenter__.return_value = mock_redis_instance + mock_redis_instance.__aexit__.return_value = None + mock_redis_instance.expire.side_effect = [ + TypeError("unexpected keyword argument 'nx'"), + True, + ] + mock_redis_instance.ttl.return_value = -1 + + with patch.object( + redis_cache, "init_async_client", return_value=mock_redis_instance + ): + await redis_cache.async_increment(key="rate_limit:window", value=1) + + assert mock_redis_instance.expire.await_count == 2 + assert mock_redis_instance.expire.await_args_list[0].args == ( + "rate_limit:window", + 60, + ) + assert mock_redis_instance.expire.await_args_list[0].kwargs == {"nx": True} + assert mock_redis_instance.expire.await_args_list[1].args == ( + "rate_limit:window", + 60, + ) + assert mock_redis_instance.expire.await_args_list[1].kwargs == {} + mock_redis_instance.ttl.assert_awaited_once_with("rate_limit:window") @pytest.mark.asyncio