Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion redisvl/index/index.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Comment thread
mikemikimike marked this conversation as resolved.
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__.")
Expand All @@ -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

Expand Down
85 changes: 85 additions & 0 deletions tests/unit/test_connection_normalization.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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."""
Expand Down