From e87240d4d9a18da68d014a8d7abfc8f5a5c184a5 Mon Sep 17 00:00:00 2001 From: Utkarsh Date: Fri, 27 Mar 2026 21:50:54 +0530 Subject: [PATCH] fix(vector_stores): handle vector=None in Milvus and Qdrant update methods (#4568) Co-authored-by: utkarsh240799 Co-authored-by: Claude Opus 4.6 (1M context) --- mem0/vector_stores/milvus.py | 11 +++++ mem0/vector_stores/qdrant.py | 18 +++++++- tests/vector_stores/test_milvus.py | 69 ++++++++++++++++++++++++++++++ tests/vector_stores/test_qdrant.py | 38 ++++++++++++++++ 4 files changed, 134 insertions(+), 2 deletions(-) diff --git a/mem0/vector_stores/milvus.py b/mem0/vector_stores/milvus.py index 09e49a954..4a0cd7961 100644 --- a/mem0/vector_stores/milvus.py +++ b/mem0/vector_stores/milvus.py @@ -181,6 +181,17 @@ class MilvusDB(VectorStoreBase): vector (List[float], optional): Updated vector. payload (Dict, optional): Updated payload. """ + if vector is None or payload is None: + existing = self.client.get(collection_name=self.collection_name, ids=vector_id) + if not existing: + raise ValueError(f"Vector with id {vector_id} not found in collection {self.collection_name}") + if vector is None: + vector = existing[0].get("vectors") + if vector is None: + raise ValueError(f"Existing record {vector_id} has no vector data") + if payload is None: + payload = existing[0].get("metadata") + schema = {"id": vector_id, "vectors": vector, "metadata": payload} self.client.upsert(collection_name=self.collection_name, data=schema) diff --git a/mem0/vector_stores/qdrant.py b/mem0/vector_stores/qdrant.py index c15e4caab..06ab7edae 100644 --- a/mem0/vector_stores/qdrant.py +++ b/mem0/vector_stores/qdrant.py @@ -12,6 +12,7 @@ from qdrant_client.models import ( MatchValue, PointIdsList, PointStruct, + PointVectors, Range, VectorParams, ) @@ -326,8 +327,21 @@ class Qdrant(VectorStoreBase): vector (list, optional): Updated vector. Defaults to None. payload (dict, optional): Updated payload. Defaults to None. """ - point = PointStruct(id=vector_id, vector=vector, payload=payload) - self.client.upsert(collection_name=self.collection_name, points=[point]) + if vector is not None and payload is not None: + point = PointStruct(id=vector_id, vector=vector, payload=payload) + self.client.upsert(collection_name=self.collection_name, points=[point]) + else: + if payload is not None: + self.client.set_payload( + collection_name=self.collection_name, + payload=payload, + points=[vector_id], + ) + if vector is not None: + self.client.update_vectors( + collection_name=self.collection_name, + points=[PointVectors(id=vector_id, vector=vector)], + ) def get(self, vector_id: int) -> dict: """ diff --git a/tests/vector_stores/test_milvus.py b/tests/vector_stores/test_milvus.py index d296620be..976056a25 100644 --- a/tests/vector_stores/test_milvus.py +++ b/tests/vector_stores/test_milvus.py @@ -233,6 +233,75 @@ class TestMilvusDB: assert parsed[1].id == "mem2" assert parsed[1].score == 0.85 + def test_update_with_none_vector_fetches_existing(self, milvus_db, mock_milvus_client): + """Test that update with vector=None fetches the existing vector (fixes #3708).""" + vector_id = "test_id" + existing_vector = [0.5] * 1536 + payload = {"user_id": "alice", "data": "Updated memory"} + + mock_milvus_client.get.return_value = [ + {"id": vector_id, "vectors": existing_vector, "metadata": {"user_id": "alice"}} + ] + + milvus_db.update(vector_id=vector_id, vector=None, payload=payload) + + mock_milvus_client.get.assert_called_once_with( + collection_name="test_collection", ids=vector_id + ) + call_args = mock_milvus_client.upsert.call_args + assert call_args[1]['data']['vectors'] == existing_vector + assert call_args[1]['data']['metadata'] == payload + + def test_update_with_none_payload_fetches_existing(self, milvus_db, mock_milvus_client): + """Test that update with payload=None fetches the existing metadata.""" + vector_id = "test_id" + vector = [0.1] * 1536 + existing_metadata = {"user_id": "alice", "data": "Original"} + + mock_milvus_client.get.return_value = [ + {"id": vector_id, "vectors": [0.5] * 1536, "metadata": existing_metadata} + ] + + milvus_db.update(vector_id=vector_id, vector=vector, payload=None) + + call_args = mock_milvus_client.upsert.call_args + assert call_args[1]['data']['vectors'] == vector + assert call_args[1]['data']['metadata'] == existing_metadata + + def test_update_with_both_none_fetches_existing(self, milvus_db, mock_milvus_client): + """Test that update with both vector=None and payload=None fetches existing data.""" + vector_id = "test_id" + existing_vector = [0.5] * 1536 + existing_metadata = {"user_id": "alice"} + + mock_milvus_client.get.return_value = [ + {"id": vector_id, "vectors": existing_vector, "metadata": existing_metadata} + ] + + milvus_db.update(vector_id=vector_id, vector=None, payload=None) + + # Should only call get once even though both are None + assert mock_milvus_client.get.call_count == 1 + call_args = mock_milvus_client.upsert.call_args + assert call_args[1]['data']['vectors'] == existing_vector + assert call_args[1]['data']['metadata'] == existing_metadata + + def test_update_with_none_vector_raises_on_missing_record(self, milvus_db, mock_milvus_client): + """Test that update raises ValueError when the record doesn't exist.""" + mock_milvus_client.get.return_value = [] + + with pytest.raises(ValueError, match="not found"): + milvus_db.update(vector_id="nonexistent", vector=None, payload={"data": "test"}) + + def test_update_with_none_vector_raises_on_missing_vector_data(self, milvus_db, mock_milvus_client): + """Test that update raises ValueError when existing record has no vector.""" + mock_milvus_client.get.return_value = [ + {"id": "test_id", "vectors": None, "metadata": {"user_id": "alice"}} + ] + + with pytest.raises(ValueError, match="no vector data"): + milvus_db.update(vector_id="test_id", vector=None, payload={"data": "test"}) + def test_collection_already_exists(self, mock_milvus_client): """Test that existing collection is not recreated.""" mock_milvus_client.has_collection.return_value = True diff --git a/tests/vector_stores/test_qdrant.py b/tests/vector_stores/test_qdrant.py index 59a1c3e24..7e08e515d 100644 --- a/tests/vector_stores/test_qdrant.py +++ b/tests/vector_stores/test_qdrant.py @@ -15,6 +15,7 @@ from qdrant_client.models import ( MatchValue, PointIdsList, PointStruct, + PointVectors, Range, VectorParams, ) @@ -232,6 +233,43 @@ class TestQdrant(unittest.TestCase): self.assertEqual(point.vector, updated_vector) self.assertEqual(point.payload, updated_payload) + def test_update_with_none_vector_uses_set_payload(self): + """Test that update with vector=None uses set_payload instead of upsert (fixes #3708).""" + vector_id = str(uuid.uuid4()) + updated_payload = {"key": "updated_value"} + + self.qdrant.update(vector_id=vector_id, vector=None, payload=updated_payload) + + self.client_mock.upsert.assert_not_called() + self.client_mock.set_payload.assert_called_once_with( + collection_name="test_collection", + payload=updated_payload, + points=[vector_id], + ) + + def test_update_with_none_payload_uses_update_vectors(self): + """Test that update with payload=None uses update_vectors instead of upsert.""" + vector_id = str(uuid.uuid4()) + updated_vector = [0.2, 0.3] + + self.qdrant.update(vector_id=vector_id, vector=updated_vector, payload=None) + + self.client_mock.upsert.assert_not_called() + self.client_mock.update_vectors.assert_called_once_with( + collection_name="test_collection", + points=[PointVectors(id=vector_id, vector=updated_vector)], + ) + + def test_update_with_both_none_is_noop(self): + """Test that update with both vector=None and payload=None is a no-op.""" + vector_id = str(uuid.uuid4()) + + self.qdrant.update(vector_id=vector_id, vector=None, payload=None) + + self.client_mock.upsert.assert_not_called() + self.client_mock.set_payload.assert_not_called() + self.client_mock.update_vectors.assert_not_called() + def test_get(self): vector_id = str(uuid.uuid4()) self.client_mock.retrieve.return_value = [{"id": vector_id, "payload": {"key": "value"}}]