From 466249113ced32243a494df3bedd60aadabb4079 Mon Sep 17 00:00:00 2001 From: Hrushikesh Yadav <136978914+HrushiYadav@users.noreply.github.com> Date: Thu, 18 Jun 2026 16:57:03 +0530 Subject: [PATCH] fix: async delete_all race condition corrupts entity store linked_memory_ids (#5553) --- mem0/memory/main.py | 33 +++++++++++++--- tests/memory/test_decay_usage_notice.py | 1 + tests/test_memory.py | 51 +++++++++++++++++++++++++ 3 files changed, 80 insertions(+), 5 deletions(-) diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 9f2ef2b69..c4cf3a626 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -2076,6 +2076,27 @@ class AsyncMemory(MemoryBase): except Exception as e: logger.warning(f"Entity upsert failed for '{entity_text}' (async): {e}") + async def _bulk_clear_entity_store(self, filters): + """Delete all entity records matching the given scope filters. + + Used by delete_all to avoid the race condition that occurs when + concurrent _delete_memory coroutines each try to read-modify-write + the same entity rows' linked_memory_ids lists. + """ + if self._entity_store is None: + return + search_filters = {k: v for k, v in filters.items() if k in ("user_id", "agent_id", "run_id") and v} + try: + listed = await asyncio.to_thread(self.entity_store.list, filters=search_filters, top_k=10000) + rows = listed[0] if isinstance(listed, (list, tuple)) and listed and isinstance(listed[0], list) else listed + for row in rows or []: + try: + await asyncio.to_thread(self.entity_store.delete, vector_id=row.id) + except Exception as e: + logger.debug(f"Bulk entity delete failed for id={row.id}: {e}") + except Exception as e: + logger.warning(f"Bulk entity store cleanup failed: {e}") + async def _remove_memory_from_entity_store(self, memory_id, filters): """Async variant of `Memory._remove_memory_from_entity_store`.""" if self._entity_store is None: @@ -3233,10 +3254,13 @@ class AsyncMemory(MemoryBase): delete_tasks = [] for memory in memories[0]: - delete_tasks.append(self._delete_memory(memory.id)) + delete_tasks.append(self._delete_memory(memory.id, skip_entity_cleanup=True)) results = await asyncio.gather(*delete_tasks, return_exceptions=True) + if self._entity_store is not None: + await self._bulk_clear_entity_store(filters) + errors = [r for r in results if isinstance(r, BaseException)] if errors: logger.warning("Failed to delete %d out of %d memories", len(errors), len(results)) @@ -3418,7 +3442,7 @@ class AsyncMemory(MemoryBase): return memory_id - async def _delete_memory(self, memory_id, existing_memory=None): + async def _delete_memory(self, memory_id, existing_memory=None, skip_entity_cleanup=False): logger.info(f"Deleting memory with {memory_id=}") if existing_memory is None: existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id) @@ -3444,9 +3468,8 @@ class AsyncMemory(MemoryBase): is_deleted=1, ) - # Entity-store cleanup: strip this memory's id from any entity records - # that linked to it. Non-fatal — the helper swallows errors. - await self._remove_memory_from_entity_store(memory_id, session_filters) + if not skip_entity_cleanup: + await self._remove_memory_from_entity_store(memory_id, session_filters) return memory_id diff --git a/tests/memory/test_decay_usage_notice.py b/tests/memory/test_decay_usage_notice.py index 817f2d738..c65a5520a 100644 --- a/tests/memory/test_decay_usage_notice.py +++ b/tests/memory/test_decay_usage_notice.py @@ -18,6 +18,7 @@ def make_async_memory(): memory = AsyncMemory.__new__(AsyncMemory) memory.vector_store = MagicMock() memory._delete_memory = AsyncMock() + memory._entity_store = None return memory diff --git a/tests/test_memory.py b/tests/test_memory.py index 7ef17c81c..53f4bcca4 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -1256,3 +1256,54 @@ class TestPreserveCustomMetadata: assert payload["source"] == "chat" assert payload["data"] == "I love swimming" assert payload["user_id"] == "user_1" + + +class TestAsyncDeleteAllEntityRace: + """Tests for async delete_all entity store race condition fix.""" + + @pytest.mark.asyncio + @patch('mem0.utils.factory.EmbedderFactory.create') + @patch('mem0.utils.factory.VectorStoreFactory.create') + @patch('mem0.utils.factory.LlmFactory.create') + @patch('mem0.memory.storage.SQLiteManager') + async def test_async_delete_all_bulk_clears_entity_store(self, mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory): + """ + Verify that async delete_all bulk-clears entity records after + concurrent memory deletes complete, preventing both the + read-modify-write race and entity orphaning on partial failures. + """ + mock_embedder_factory.return_value = MagicMock() + mock_llm_factory.return_value = MagicMock() + mock_sqlite.return_value = MagicMock() + + mock_vector_store = MagicMock() + mem_a = MagicMock() + mem_a.id = "mem-a" + mem_a.payload = {"data": "Alice likes Python", "user_id": "alice"} + mem_b = MagicMock() + mem_b.id = "mem-b" + mem_b.payload = {"data": "Alice works at Acme", "user_id": "alice"} + mock_vector_store.list.return_value = ([mem_a, mem_b],) + mock_vector_store.get.side_effect = lambda vector_id: {"mem-a": mem_a, "mem-b": mem_b}[vector_id] + mock_vector_factory.return_value = mock_vector_store + + mock_entity_store = MagicMock() + entity_row = MagicMock() + entity_row.id = "entity-alice" + entity_row.payload = { + "data": "alice", + "user_id": "alice", + "linked_memory_ids": ["mem-a", "mem-b"], + } + mock_entity_store.list.return_value = ([entity_row],) + + from mem0.memory.main import AsyncMemory + config = MemoryConfig() + memory = AsyncMemory(config) + memory._entity_store = mock_entity_store + + await memory.delete_all(user_id="alice") + + mock_entity_store.delete.assert_called_once_with(vector_id="entity-alice") + + assert mock_vector_store.delete.call_count == 2