feat(valkey): add cluster mode enabled (CME) support (#4759)

This commit is contained in:
Swarnaprakash Udayakumar
2026-04-13 07:37:11 -07:00
committed by GitHub
parent 92b958c1cc
commit cf530c4bec
4 changed files with 155 additions and 9 deletions
+21
View File
@@ -50,4 +50,25 @@ Here are the parameters available for configuring Valkey:
| `hnsw_m` | Number of bi-directional links for HNSW | `16` |
| `hnsw_ef_construction` | Size of dynamic candidate list for HNSW | `200` |
| `hnsw_ef_runtime` | Size of dynamic candidate list for search | `10` |
| `cluster_mode` | Enable cluster mode for Valkey cluster (CME) deployments | `false` |
| `distance_metric` | Distance metric for vector similarity | `cosine` |
## Cluster Mode
To use Valkey with cluster mode enabled (CME), set `cluster_mode` to `true`:
```python
config = {
"vector_store": {
"provider": "valkey",
"config": {
"collection_name": "memories",
"valkey_url": "valkey://cluster-endpoint:6379",
"embedding_model_dims": 1536,
"cluster_mode": True
}
}
}
```
When cluster mode is enabled, the connector uses `ValkeyCluster` instead of the standalone client, which handles `MOVED`/`ASK` redirections automatically. Search queries are coordinated across all shards by the valkey-search module's built-in coordinator. See the [valkey-search documentation](https://github.com/valkey-io/valkey-search) for details on cluster mode behavior.
+1
View File
@@ -14,6 +14,7 @@ class ValkeyConfig(BaseModel):
hnsw_m: int = Field(16, description="HNSW: number of connections per layer")
hnsw_ef_construction: int = Field(200, description="HNSW: search width during index construction")
hnsw_ef_runtime: int = Field(10, description="HNSW: search width during queries")
cluster_mode: bool = Field(False, description="Enable cluster mode for Valkey cluster (CME) deployments")
@model_validator(mode="before")
@classmethod
+13 -3
View File
@@ -52,6 +52,7 @@ class ValkeyDB(VectorStoreBase):
hnsw_m: int = 16,
hnsw_ef_construction: int = 200,
hnsw_ef_runtime: int = 10,
cluster_mode: bool = False,
):
"""
Initialize the Valkey vector store.
@@ -65,6 +66,7 @@ class ValkeyDB(VectorStoreBase):
hnsw_m (int, optional): HNSW M parameter (connections per node). Defaults to 16.
hnsw_ef_construction (int, optional): HNSW ef_construction parameter. Defaults to 200.
hnsw_ef_runtime (int, optional): HNSW ef_runtime parameter. Defaults to 10.
cluster_mode (bool, optional): Enable cluster mode for Valkey cluster (CME) deployments. Defaults to False.
"""
self.embedding_model_dims = embedding_model_dims
self.collection_name = collection_name
@@ -74,6 +76,7 @@ class ValkeyDB(VectorStoreBase):
self.hnsw_m = hnsw_m
self.hnsw_ef_construction = hnsw_ef_construction
self.hnsw_ef_runtime = hnsw_ef_runtime
self.cluster_mode = cluster_mode
# Validate index type
if self.index_type not in ["hnsw", "flat"]:
@@ -81,8 +84,13 @@ class ValkeyDB(VectorStoreBase):
# Connect to Valkey
try:
self.client = valkey.from_url(valkey_url)
logger.debug(f"Successfully connected to Valkey at {valkey_url}")
if self.cluster_mode:
from valkey.cluster import ValkeyCluster
self.client = ValkeyCluster.from_url(valkey_url)
else:
self.client = valkey.from_url(valkey_url)
logger.debug(f"Successfully connected to Valkey at {valkey_url} (cluster_mode={cluster_mode})")
except Exception as e:
logger.exception(f"Failed to connect to Valkey at {valkey_url}: {e}")
raise
@@ -185,7 +193,6 @@ class ValkeyDB(VectorStoreBase):
"""
# Check if the search module is available
try:
# Try to execute a search command
self.client.execute_command("FT._LIST")
except ResponseError as e:
if "unknown command" in str(e).lower():
@@ -353,6 +360,9 @@ class ValkeyDB(VectorStoreBase):
"""
Execute a search query.
In cluster mode, the valkey-search module's built-in coordinator handles
fan-out across all shards and aggregates results server-side.
Args:
query (str): The search query to execute.
params (dict): The query parameters.
+120 -6
View File
@@ -222,9 +222,7 @@ def test_update_with_none_vector_preserves_embedding(valkey_db, mock_valkey_clie
mock_valkey_client.hset.assert_called_once()
args, kwargs = mock_valkey_client.hset.call_args
assert "embedding" not in kwargs["mapping"], (
"embedding should not be in hash_data when vector is None"
)
assert "embedding" not in kwargs["mapping"], "embedding should not be in hash_data when vector is None"
assert kwargs["mapping"]["memory_id"] == "test_id"
assert kwargs["mapping"]["memory"] == "updated_data"
@@ -243,9 +241,7 @@ def test_update_with_vector_includes_embedding(valkey_db, mock_valkey_client):
mock_valkey_client.hset.assert_called_once()
args, kwargs = mock_valkey_client.hset.call_args
assert "embedding" in kwargs["mapping"], (
"embedding should be in hash_data when vector is provided"
)
assert "embedding" in kwargs["mapping"], "embedding should be in hash_data when vector is provided"
expected_bytes = np.array(vector, dtype=np.float32).tobytes()
assert kwargs["mapping"]["embedding"] == expected_bytes
@@ -906,3 +902,121 @@ def test_list_with_missing_fields_and_defaults(valkey_db, mock_valkey_client):
assert result.id == "fallback_id"
assert "hash" in result.payload
assert "data" in result.payload # memory is renamed to data
# Cluster mode tests
@pytest.fixture
def mock_valkey_cluster_client():
"""Create a mock ValkeyCluster client."""
with patch("valkey.cluster.ValkeyCluster.from_url") as mock_from_url:
mock_client = MagicMock()
mock_ft = MagicMock()
mock_client.ft = MagicMock(return_value=mock_ft)
mock_client.execute_command = MagicMock()
mock_client.hset = MagicMock()
mock_client.hgetall = MagicMock()
mock_client.delete = MagicMock()
mock_from_url.return_value = mock_client
yield mock_client
@pytest.fixture
def valkey_db_cluster(mock_valkey_cluster_client):
"""Create a ValkeyDB instance in cluster mode with a mock client."""
valkey_db = ValkeyDB(
valkey_url="valkey://localhost:7000",
collection_name="test_cluster",
embedding_model_dims=1536,
cluster_mode=True,
)
valkey_db.client = mock_valkey_cluster_client
return valkey_db
def test_cluster_mode_init(mock_valkey_cluster_client):
"""Test that cluster_mode=True uses ValkeyCluster client."""
db = ValkeyDB(
valkey_url="valkey://localhost:7000",
collection_name="test_cluster",
embedding_model_dims=1536,
cluster_mode=True,
)
assert db.cluster_mode is True
def test_cluster_mode_create_index(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that index creation works in cluster mode (server handles propagation)."""
mock_valkey_cluster_client.execute_command.reset_mock()
mock_valkey_cluster_client.ft.return_value.info.side_effect = ResponseError("not found")
valkey_db_cluster._create_index(1536)
call_args = mock_valkey_cluster_client.execute_command.call_args
assert call_args is not None
assert "FT.CREATE" in call_args.args
def test_cluster_mode_drop_index(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that dropping index works in cluster mode."""
mock_valkey_cluster_client.execute_command.reset_mock()
valkey_db_cluster._drop_index("test_cluster")
call_args = mock_valkey_cluster_client.execute_command.call_args
assert "FT.DROPINDEX" in call_args.args
def test_cluster_mode_search(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that search in cluster mode uses ft().search() (server handles cross-shard fan-out)."""
ts = str(int(datetime.now().timestamp()))
mock_doc = MagicMock()
mock_doc.memory_id = "id1"
mock_doc.hash = "h1"
mock_doc.memory = "data1"
mock_doc.created_at = ts
mock_doc.metadata = "{}"
mock_doc.vector_score = "0.1"
mock_results = MagicMock()
mock_results.docs = [mock_doc]
mock_valkey_cluster_client.ft.return_value.search.return_value = mock_results
results = valkey_db_cluster.search("test", np.random.rand(1536).tolist(), limit=5)
assert len(results) == 1
assert results[0].id == "id1"
mock_valkey_cluster_client.ft.return_value.search.assert_called_once()
def test_cluster_mode_insert(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that insert works in cluster mode (ValkeyCluster handles routing)."""
vectors = [np.random.rand(1536).tolist()]
payloads = [{"hash": "h1", "data": "test", "user_id": "u1"}]
ids = ["id1"]
valkey_db_cluster.insert(vectors=vectors, payloads=payloads, ids=ids)
mock_valkey_cluster_client.hset.assert_called_once()
def test_cluster_mode_get(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that get works in cluster mode."""
mock_valkey_cluster_client.hgetall.return_value = {
"memory_id": "id1",
"hash": "h1",
"memory": "test_data",
"created_at": str(int(datetime.now().timestamp())),
"metadata": "{}",
}
result = valkey_db_cluster.get("id1")
assert result.id == "id1"
assert result.payload["data"] == "test_data"
def test_cluster_mode_delete(valkey_db_cluster, mock_valkey_cluster_client):
"""Test that delete works in cluster mode."""
valkey_db_cluster.delete("id1")
mock_valkey_cluster_client.delete.assert_called_once_with("mem0:test_cluster:id1")