fix: async delete_all race condition corrupts entity store linked_memory_ids (#5553)

This commit is contained in:
Hrushikesh Yadav
2026-06-18 16:57:03 +05:30
committed by GitHub
parent 3e2ae734e7
commit 466249113c
3 changed files with 80 additions and 5 deletions
+28 -5
View File
@@ -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
+1
View File
@@ -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
+51
View File
@@ -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