diff --git a/mem0/vector_stores/weaviate.py b/mem0/vector_stores/weaviate.py index c341b8a45..2c56567eb 100644 --- a/mem0/vector_stores/weaviate.py +++ b/mem0/vector_stores/weaviate.py @@ -1,6 +1,6 @@ import logging import uuid -from typing import Dict, List, Mapping, Optional +from typing import Dict, List, Optional from urllib.parse import urlparse from pydantic import BaseModel @@ -297,13 +297,7 @@ class Weaviate(VectorStoreBase): collection.data.update(uuid=vector_id, properties=payload) if vector: - existing_data = self.get(vector_id) - if existing_data: - existing_data = dict(existing_data) - if "id" in existing_data: - del existing_data["id"] - existing_payload: Mapping[str, str] = existing_data - collection.data.update(uuid=vector_id, properties=existing_payload, vector=vector) + collection.data.update(uuid=vector_id, vector=vector) def get(self, vector_id) -> Optional[OutputData]: """ diff --git a/tests/vector_stores/test_weaviate.py b/tests/vector_stores/test_weaviate.py index d246e08df..bdcc7e79e 100644 --- a/tests/vector_stores/test_weaviate.py +++ b/tests/vector_stores/test_weaviate.py @@ -220,6 +220,50 @@ class TestWeaviateDB(unittest.TestCase): self.client_mock.collections.delete.assert_called_once_with("test_collection") self.client_mock.collections.create.assert_called_once() + def test_update_preserves_properties_and_does_not_write_model_fields(self): + # Regression: the vector branch of update() previously resent + # dict(OutputData) as properties, i.e. the model field names + # {"id", "score", "payload"}, corrupting the stored object instead of + # preserving its real properties. + valid_uuid = str(uuid.uuid4()) + collection = self.client_mock.collections.get.return_value + + # Mock get() so that, on the buggy code path, the vector branch can run + # to completion and actually write its (wrong) properties. + mock_response = MagicMock() + mock_response.properties = {"data": "existing", "hash": "abc123"} + mock_response.uuid = valid_uuid + collection.query.fetch_object_by_id.return_value = mock_response + + new_payload = { + "data": "updated memory", + "hash": "def456", + "user_id": "user_123", + } + + self.weaviate_db.update( + vector_id=valid_uuid, + vector=[0.1] * 1536, + payload=new_payload, + ) + + update_calls = collection.data.update.call_args_list + # The payload branch updates the real properties. + property_updates = [c for c in update_calls if "properties" in c.kwargs] + self.assertEqual(len(property_updates), 1) + self.assertEqual(property_updates[0].kwargs["properties"], new_payload) + + # No update call may write the OutputData model field names as properties. + for call in update_calls: + props = call.kwargs.get("properties", {}) + self.assertNotIn("score", props) + self.assertNotIn("payload", props) + + # The vector branch updates the vector without touching properties. + vector_updates = [c for c in update_calls if "vector" in c.kwargs] + self.assertEqual(len(vector_updates), 1) + self.assertNotIn("properties", vector_updates[0].kwargs) + if __name__ == "__main__": unittest.main()