diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index a94bc6637e4..29ff72f39c4 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -94,6 +94,24 @@ def _get_call_stack_info(num_frames: int = 2) -> str: return "unknown" +def _is_expire_nx_unsupported_error(error: Exception) -> bool: + """ + Return True when EXPIRE ... NX is unsupported and we should fallback. + + Compatibility cases: + - Older redis-py clients that do not accept `nx` -> TypeError + - Redis < 7 servers that reject `EXPIRE key ttl NX` -> ResponseError syntax/arity errors + """ + if isinstance(error, TypeError): + return True + + if error.__class__.__name__ != "ResponseError": + return False + + error_message = str(error).lower() + return "syntax error" in error_message or "wrong number of arguments" in error_message + + class RedisCircuitBreaker: """ Tracks Redis health for a RedisCache instance. @@ -850,9 +868,12 @@ class RedisCache(BaseCache): # hot increment path while preserving window semantics. try: await _redis_client.expire(key, _used_ttl, nx=True) - except TypeError: + except Exception as e: + if not _is_expire_nx_unsupported_error(e): + raise # Backward-compatible fallback if the Redis client does - # not support the "nx" kwarg on expire(). + # not support the "nx" kwarg on expire(), or if an + # older Redis server rejects EXPIRE ... NX. current_ttl = await _redis_client.ttl(key) if current_ttl == -1: await _redis_client.expire(key, _used_ttl) diff --git a/tests/test_litellm/caching/test_redis_cache.py b/tests/test_litellm/caching/test_redis_cache.py index 7283d8307d4..968029a649b 100644 --- a/tests/test_litellm/caching/test_redis_cache.py +++ b/tests/test_litellm/caching/test_redis_cache.py @@ -163,6 +163,62 @@ async def test_redis_cache_async_increment_default_fallback_existing_ttl_skips_s mock_redis_instance.ttl.assert_awaited_once_with("rate_limit:window") +@pytest.mark.asyncio +async def test_redis_cache_async_increment_default_falls_back_on_redis6_expire_nx_syntax_error( + monkeypatch, redis_no_ping +): + """If Redis server rejects EXPIRE ... NX, fallback to TTL-check + conditional EXPIRE.""" + 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 + + class ResponseError(Exception): + pass + + mock_redis_instance.expire.side_effect = [ + ResponseError("ERR syntax error"), + 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].kwargs == {"nx": True} + assert mock_redis_instance.expire.await_args_list[1].kwargs == {} + mock_redis_instance.ttl.assert_awaited_once_with("rate_limit:window") + + +@pytest.mark.asyncio +async def test_redis_cache_async_increment_default_raises_non_compat_expire_error( + monkeypatch, redis_no_ping +): + """Non-compatibility expire errors should still propagate.""" + 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 + + class ResponseError(Exception): + pass + + mock_redis_instance.expire.side_effect = ResponseError("READONLY You can't write") + + with patch.object( + redis_cache, "init_async_client", return_value=mock_redis_instance + ): + with pytest.raises(ResponseError, match="READONLY"): + await redis_cache.async_increment(key="rate_limit:window", value=1) + + mock_redis_instance.ttl.assert_not_awaited() + + @pytest.mark.asyncio async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping): monkeypatch.setenv("REDIS_HOST", "my-fake-host")