diff --git a/redisvl/index/index.py b/redisvl/index/index.py index 3daa69d3..7a4f2522 100644 --- a/redisvl/index/index.py +++ b/redisvl/index/index.py @@ -964,9 +964,15 @@ def connect(self, redis_url: str | None = None, **kwargs): ModuleNotFoundError: If required Redis modules are not installed. """ self.invalidate_sql_schema_cache() - self.__redis_client = RedisConnectionFactory.get_redis_connection( + new_client = RedisConnectionFactory.get_redis_connection( redis_url=redis_url, **kwargs ) + if self._owns_redis_client: + self._detach_client_finalizer() + if self.__redis_client is not None: + self.__redis_client.close() + self.__redis_client = new_client + self._owns_redis_client = True self._register_client_finalizer(self.__redis_client) @deprecated_function("set_client", "Pass connection parameters in __init__.") @@ -986,7 +992,12 @@ def set_client(self, redis_client: SyncRedisClient, **kwargs): """ RedisConnectionFactory.validate_sync_redis(redis_client) self.invalidate_sql_schema_cache() + if self._owns_redis_client: + self._detach_client_finalizer() + if self.__redis_client is not None: + self.__redis_client.close() self.__redis_client = redis_client + self._owns_redis_client = False self._register_client_finalizer(redis_client) return self diff --git a/tests/unit/test_connection_normalization.py b/tests/unit/test_connection_normalization.py index 98622eb1..e6738003 100644 --- a/tests/unit/test_connection_normalization.py +++ b/tests/unit/test_connection_normalization.py @@ -2,11 +2,13 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from redis.exceptions import ConnectionError from redisvl.extensions.cache.embeddings import EmbeddingsCache from redisvl.extensions.router.semantic import SemanticRouter from redisvl.index import AsyncSearchIndex, SearchIndex from redisvl.query.sql import SQLQuery +from redisvl.schema import IndexSchema def _schema_dict(name: str = "idx") -> dict: @@ -59,6 +61,70 @@ def test_search_index_from_existing_prefers_provided_client(): assert index.client is provided_client +def test_search_index_set_client_releases_owned_client_without_taking_ownership(): + owned_client = MagicMock() + provided_client = MagicMock() + with patch( + "redisvl.index.index.RedisConnectionFactory.get_redis_connection", + return_value=owned_client, + ): + index = SearchIndex( + IndexSchema.from_dict(_schema_dict()), redis_url="redis://localhost:6379" + ) + assert index._redis_client is owned_client + + with patch("redisvl.index.index.RedisConnectionFactory.validate_sync_redis"): + index.set_client(provided_client) + + owned_client.close.assert_called_once_with() + assert index.client is provided_client + assert index._owns_redis_client is False + index.disconnect() + provided_client.close.assert_not_called() + + +def test_search_index_connect_keeps_ownership_after_replacing_owned_client(): + first_client = MagicMock() + second_client = MagicMock() + with patch( + "redisvl.index.index.RedisConnectionFactory.get_redis_connection", + side_effect=[first_client, second_client], + ): + index = SearchIndex( + IndexSchema.from_dict(_schema_dict()), redis_url="redis://first:6379" + ) + assert index._redis_client is first_client + index.connect("redis://second:6379") + + first_client.close.assert_called_once_with() + assert index._redis_client is second_client + assert index._owns_redis_client is True + index.disconnect() + second_client.close.assert_called_once_with() + + +def test_search_index_connect_preserves_live_client_when_replacement_fails(): + first_client = MagicMock() + connection_error = ConnectionError("replacement unavailable") + with patch( + "redisvl.index.index.RedisConnectionFactory.get_redis_connection", + side_effect=[first_client, connection_error], + ): + index = SearchIndex( + IndexSchema.from_dict(_schema_dict()), redis_url="redis://first:6379" + ) + assert index._redis_client is first_client + + with pytest.raises(ConnectionError, match="replacement unavailable"): + index.connect("redis://unavailable:6379") + + assert index.client is first_client + assert index._owns_redis_client is True + first_client.close.assert_not_called() + index.disconnect() + first_client.close.assert_called_once_with() + + def test_search_index_from_existing_owns_factory_created_client(): """Reuse a single sync client created internally from redis_url.""" created_client = MagicMock() @@ -140,6 +206,25 @@ async def test_async_search_index_from_existing_prefers_provided_client(): assert index.client is provided_client +@pytest.mark.asyncio +async def test_async_search_index_set_client_does_not_take_ownership(): + provided_client = AsyncMock() + index = AsyncSearchIndex( + IndexSchema.from_dict(_schema_dict()), redis_client=provided_client + ) + + replacement = AsyncMock() + with patch.object( + index, "_validate_client", new=AsyncMock(return_value=replacement) + ): + await index.set_client(replacement) + + assert index.client is replacement + assert index._owns_redis_client is False + provided_client.aclose.assert_not_awaited() + replacement.aclose.assert_not_awaited() + + @pytest.mark.asyncio async def test_async_search_index_from_existing_owns_factory_created_client(): """Reuse a single async client created internally from redis_url."""