fix: async delete_all race condition corrupts entity store linked_memory_ids (#5553)
This commit is contained in:
+28
-5
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user