From cf530c4becf8855b7e089566b4f76abb0e2637b8 Mon Sep 17 00:00:00 2001 From: Swarnaprakash Udayakumar Date: Mon, 13 Apr 2026 07:37:11 -0700 Subject: [PATCH] feat(valkey): add cluster mode enabled (CME) support (#4759) --- docs/components/vectordbs/dbs/valkey.mdx | 21 ++++ mem0/configs/vector_stores/valkey.py | 1 + mem0/vector_stores/valkey.py | 16 ++- tests/vector_stores/test_valkey.py | 126 +++++++++++++++++++++-- 4 files changed, 155 insertions(+), 9 deletions(-) diff --git a/docs/components/vectordbs/dbs/valkey.mdx b/docs/components/vectordbs/dbs/valkey.mdx index 8cd63df8a..af1e28219 100644 --- a/docs/components/vectordbs/dbs/valkey.mdx +++ b/docs/components/vectordbs/dbs/valkey.mdx @@ -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. diff --git a/mem0/configs/vector_stores/valkey.py b/mem0/configs/vector_stores/valkey.py index 595110882..0d2bdc4dc 100644 --- a/mem0/configs/vector_stores/valkey.py +++ b/mem0/configs/vector_stores/valkey.py @@ -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 diff --git a/mem0/vector_stores/valkey.py b/mem0/vector_stores/valkey.py index 4dc0dc78d..e3bd843fa 100644 --- a/mem0/vector_stores/valkey.py +++ b/mem0/vector_stores/valkey.py @@ -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. diff --git a/tests/vector_stores/test_valkey.py b/tests/vector_stores/test_valkey.py index 7269830c8..bd30214d5 100644 --- a/tests/vector_stores/test_valkey.py +++ b/tests/vector_stores/test_valkey.py @@ -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")