fix: clean up graph store data on Memory.delete() (#4505)

Co-authored-by: utkarsh240799 <utkarsh240799@users.noreply.github.com>
Co-authored-by: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Utkarsh
2026-03-23 19:27:38 +05:30
committed by GitHub
parent d8a6960b4a
commit 5332741961
10 changed files with 2237 additions and 9 deletions
+22
View File
@@ -409,6 +409,28 @@ class NeptuneBase(ABC):
"""
pass
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters["user_id"])
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
cypher, params = self._delete_all_cypher(filters)
self.graph.query(cypher, params=params)
+22
View File
@@ -246,6 +246,28 @@ class MemoryGraph:
logger.info(f"Returned {len(search_results)} search results")
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
"""Delete all nodes and relationships for a user or specific agent."""
where_parts = ["n.user_id = %s"]
+22
View File
@@ -129,6 +129,28 @@ class MemoryGraph:
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then soft-deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
# Build node properties for filtering
node_props = ["user_id: $user_id"]
+22
View File
@@ -149,6 +149,28 @@ class MemoryGraph:
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id, run_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
# Build node properties for filtering
node_props = ["user_id: $user_id"]
+50 -8
View File
@@ -1101,7 +1101,27 @@ class Memory(MemoryBase):
memory_id (str): ID of the memory to delete.
"""
capture_event("mem0.delete", self, {"memory_id": memory_id, "sync_type": "sync"})
self._delete_memory(memory_id)
existing_memory = self.vector_store.get(vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
# Clean up graph entities before deleting from vector store
if self.enable_graph:
try:
memory_text = existing_memory.payload.get("data", "")
if memory_text:
filters = {}
for key in ("user_id", "agent_id", "run_id"):
val = existing_memory.payload.get(key)
if val:
filters[key] = val
if filters.get("user_id"):
self.graph.delete(memory_text, filters)
except Exception as e:
logger.error(f"Error cleaning up graph for memory {memory_id}: {e}")
self._delete_memory(memory_id, existing_memory)
return {"message": "Memory deleted successfully!"}
def delete_all(self, user_id: Optional[str] = None, agent_id: Optional[str] = None, run_id: Optional[str] = None):
@@ -1277,11 +1297,12 @@ class Memory(MemoryBase):
)
return memory_id
def _delete_memory(self, memory_id):
def _delete_memory(self, memory_id, existing_memory=None):
logger.info(f"Deleting memory with {memory_id=}")
existing_memory = self.vector_store.get(vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
existing_memory = self.vector_store.get(vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
prev_value = existing_memory.payload.get("data", "")
self.vector_store.delete(vector_id=memory_id)
self.db.add_history(
@@ -2174,7 +2195,27 @@ class AsyncMemory(MemoryBase):
memory_id (str): ID of the memory to delete.
"""
capture_event("mem0.delete", self, {"memory_id": memory_id, "sync_type": "async"})
await self._delete_memory(memory_id)
existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
# Clean up graph entities before deleting from vector store
if self.enable_graph:
try:
memory_text = existing_memory.payload.get("data", "")
if memory_text:
filters = {}
for key in ("user_id", "agent_id", "run_id"):
val = existing_memory.payload.get(key)
if val:
filters[key] = val
if filters.get("user_id"):
await asyncio.to_thread(self.graph.delete, memory_text, filters)
except Exception as e:
logger.error(f"Error cleaning up graph for memory {memory_id}: {e}")
await self._delete_memory(memory_id, existing_memory)
return {"message": "Memory deleted successfully!"}
async def delete_all(self, user_id=None, agent_id=None, run_id=None):
@@ -2375,11 +2416,12 @@ class AsyncMemory(MemoryBase):
)
return memory_id
async def _delete_memory(self, memory_id):
async def _delete_memory(self, memory_id, existing_memory=None):
logger.info(f"Deleting memory with {memory_id=}")
existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
existing_memory = await asyncio.to_thread(self.vector_store.get, vector_id=memory_id)
if existing_memory is None:
raise ValueError(f"Memory with id {memory_id} not found")
prev_value = existing_memory.payload.get("data", "")
await asyncio.to_thread(self.vector_store.delete, vector_id=memory_id)
+22
View File
@@ -134,6 +134,28 @@ class MemoryGraph:
return search_results
def delete(self, data, filters):
"""
Delete graph entities associated with the given memory text.
Extracts entities and relationships from the memory text using the same
pipeline as add(), then deletes the matching relationships in the graph.
Args:
data (str): The memory text whose graph entities should be removed.
filters (dict): Scope filters (user_id, agent_id).
"""
try:
entity_type_map = self._retrieve_nodes_from_data(data, filters)
if not entity_type_map:
logger.debug("No entities found in memory text, skipping graph cleanup")
return
to_be_deleted = self._establish_nodes_relations_from_data(data, filters, entity_type_map)
if to_be_deleted:
self._delete_entities(to_be_deleted, filters)
except Exception as e:
logger.error(f"Error during graph cleanup for memory delete: {e}")
def delete_all(self, filters):
"""Delete all nodes and relationships for a user or specific agent."""
if filters.get("agent_id"):
+517
View File
@@ -0,0 +1,517 @@
"""Tests for graph cleanup on memory deletion (issue #3245)."""
from unittest.mock import MagicMock, patch
import pytest
from mem0.configs.base import MemoryConfig
class MockVectorMemory:
def __init__(self, memory_id, payload, score=0.8):
self.id = memory_id
self.payload = payload
self.score = score
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_calls_graph_cleanup_when_graph_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""When graph is enabled, delete() should call graph.delete() with memory text and filters."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
# Enable graph with a mock
memory.enable_graph = True
memory.graph = MagicMock()
# Set up vector store to return a memory with graph-relevant data
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1",
{
"data": "Alice likes Bob",
"user_id": "user-1",
"agent_id": "agent-1",
"hash": "abc",
},
)
memory.delete("mem-1")
# graph.delete should have been called with the memory text and filters
memory.graph.delete.assert_called_once_with(
"Alice likes Bob", {"user_id": "user-1", "agent_id": "agent-1"}
)
# _delete_memory should still have been called (vector store + history cleanup)
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_skips_graph_when_not_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""When graph is not enabled, delete() should not attempt graph cleanup."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
assert memory.enable_graph is False
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_continues_if_graph_cleanup_fails(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""If graph cleanup raises an exception, delete() should still succeed."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
memory.graph.delete.side_effect = RuntimeError("Neo4j connection lost")
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
# Should not raise
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
# Vector store deletion should still proceed
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_skips_graph_when_no_user_id(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Graph cleanup should be skipped if the memory has no user_id."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
# Memory with no user_id
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Some data", "hash": "abc"}
)
memory.delete("mem-1")
# graph.delete should NOT have been called since there's no user_id
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_skips_graph_when_no_memory_text(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Graph cleanup should be skipped if the memory has no text data."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"user_id": "user-1", "hash": "abc"}
)
memory.delete("mem-1")
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_passes_all_filters_to_graph(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Graph cleanup should include all available filters (user_id, agent_id, run_id)."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1",
{
"data": "Alice likes Bob",
"user_id": "user-1",
"agent_id": "agent-1",
"run_id": "run-1",
"hash": "abc",
},
)
memory.delete("mem-1")
memory.graph.delete.assert_called_once_with(
"Alice likes Bob",
{"user_id": "user-1", "agent_id": "agent-1", "run_id": "run-1"},
)
@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_calls_graph_cleanup(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Async delete() should also perform graph cleanup."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1",
{
"data": "Alice likes Bob",
"user_id": "user-1",
"hash": "abc",
},
)
result = await memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
memory.graph.delete.assert_called_once_with("Alice likes Bob", {"user_id": "user-1"})
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@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_continues_if_graph_cleanup_fails(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Async delete() should continue even if graph cleanup fails."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.enable_graph = True
memory.graph = MagicMock()
memory.graph.delete.side_effect = RuntimeError("Graph error")
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
result = await memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_raises_for_nonexistent_memory_with_graph_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""delete() should raise ValueError for non-existent memory even with graph enabled."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = None
with pytest.raises(ValueError, match="Memory with id non-existent not found"):
memory.delete("non-existent")
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_not_called()
@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_raises_for_nonexistent_memory_with_graph_enabled(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""Async delete() should raise ValueError for non-existent memory even with graph enabled."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_store.get.return_value = None
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import AsyncMemory
config = MemoryConfig()
memory = AsyncMemory(config)
memory.enable_graph = True
memory.graph = MagicMock()
with pytest.raises(ValueError, match="Memory with id non-existent not found"):
await memory.delete("non-existent")
memory.graph.delete.assert_not_called()
mock_vector_store.delete.assert_not_called()
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_delete_all_does_not_trigger_per_memory_graph_cleanup(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""delete_all() should use graph.delete_all(), not per-memory graph.delete()."""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mem1 = MockVectorMemory("mem-1", {"data": "Alice likes Bob", "user_id": "user-1"})
mem2 = MockVectorMemory("mem-2", {"data": "Bob likes Charlie", "user_id": "user-1"})
mock_vector_store.list.return_value = ([mem1, mem2], 2)
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1"}
)
memory.delete_all(user_id="user-1")
# graph.delete (per-memory) should NOT be called
memory.graph.delete.assert_not_called()
# graph.delete_all (bulk) SHOULD be called
memory.graph.delete_all.assert_called_once_with({"user_id": "user-1"})
@patch("mem0.utils.factory.EmbedderFactory.create")
@patch("mem0.utils.factory.VectorStoreFactory.create")
@patch("mem0.utils.factory.LlmFactory.create")
@patch("mem0.memory.storage.SQLiteManager")
def test_internal_delete_memory_does_not_trigger_graph_cleanup(
mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory
):
"""_delete_memory() should NOT call graph.delete() — only the public delete() does.
This ensures that the DELETE branch inside _add_to_vector_store() (which calls
_delete_memory directly) does not interfere with the parallel graph pipeline
running in _add_to_graph().
"""
mock_embedder_factory.return_value = MagicMock()
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm_factory.return_value = MagicMock()
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
memory.enable_graph = True
memory.graph = MagicMock()
mock_vector_store.get.return_value = MockVectorMemory(
"mem-1", {"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"}
)
# Call _delete_memory directly (as _add_to_vector_store does for DELETE events)
memory._delete_memory("mem-1")
# graph.delete should NOT have been called — graph cleanup is only in delete()
memory.graph.delete.assert_not_called()
# But vector store deletion should proceed
mock_vector_store.delete.assert_called_once_with(vector_id="mem-1")
def test_graph_memory_delete_calls_internal_methods():
"""Test that MemoryGraph.delete() calls the expected internal pipeline methods."""
from unittest.mock import patch as _patch
# We need to mock the Neo4j import
with _patch.dict("sys.modules", {"langchain_neo4j": MagicMock(), "rank_bm25": MagicMock()}):
from mem0.memory.graph_memory import MemoryGraph
with _patch.object(MemoryGraph, "__init__", return_value=None):
graph = MemoryGraph.__new__(MemoryGraph)
# Mock the internal methods
graph._retrieve_nodes_from_data = MagicMock(
return_value={"alice": "person", "bob": "person"}
)
graph._establish_nodes_relations_from_data = MagicMock(
return_value=[
{"source": "alice", "destination": "bob", "relationship": "likes"}
]
)
graph._delete_entities = MagicMock(return_value=[])
filters = {"user_id": "user-1"}
graph.delete("Alice likes Bob", filters)
graph._retrieve_nodes_from_data.assert_called_once_with("Alice likes Bob", filters)
graph._establish_nodes_relations_from_data.assert_called_once_with(
"Alice likes Bob", filters, {"alice": "person", "bob": "person"}
)
graph._delete_entities.assert_called_once_with(
[{"source": "alice", "destination": "bob", "relationship": "likes"}],
filters,
)
def test_graph_memory_delete_skips_when_no_entities():
"""Test that MemoryGraph.delete() does nothing when no entities are extracted."""
from unittest.mock import patch as _patch
with _patch.dict("sys.modules", {"langchain_neo4j": MagicMock(), "rank_bm25": MagicMock()}):
from mem0.memory.graph_memory import MemoryGraph
with _patch.object(MemoryGraph, "__init__", return_value=None):
graph = MemoryGraph.__new__(MemoryGraph)
graph._retrieve_nodes_from_data = MagicMock(return_value={})
graph._establish_nodes_relations_from_data = MagicMock()
graph._delete_entities = MagicMock()
graph.delete("Some text", {"user_id": "user-1"})
graph._retrieve_nodes_from_data.assert_called_once()
graph._establish_nodes_relations_from_data.assert_not_called()
graph._delete_entities.assert_not_called()
def test_graph_memory_delete_handles_exception():
"""Test that MemoryGraph.delete() catches exceptions without raising."""
from unittest.mock import patch as _patch
with _patch.dict("sys.modules", {"langchain_neo4j": MagicMock(), "rank_bm25": MagicMock()}):
from mem0.memory.graph_memory import MemoryGraph
with _patch.object(MemoryGraph, "__init__", return_value=None):
graph = MemoryGraph.__new__(MemoryGraph)
graph._retrieve_nodes_from_data = MagicMock(
side_effect=RuntimeError("LLM error")
)
# Should not raise
graph.delete("Some text", {"user_id": "user-1"})
+821
View File
@@ -0,0 +1,821 @@
"""
End-to-end tests for graph cleanup on memory deletion against real
Neo4j, Memgraph, and Apache AGE instances running in Docker.
Requires:
docker run -d --name mem0-neo4j-test -p 7687:7687 -e NEO4J_AUTH=neo4j/testpassword neo4j:5.23
docker run -d --name mem0-memgraph-test -p 7688:7687 memgraph/memgraph:latest
docker run -d --name mem0-age-test -p 5432:5432 -e POSTGRES_USER=postgres -e POSTGRES_PASSWORD=testpassword -e POSTGRES_DB=testdb apache/age:latest
Tests are skipped automatically if the databases or required Python
packages are not available.
"""
import hashlib
import sys
import warnings
from unittest.mock import MagicMock
import pytest
warnings.filterwarnings("ignore")
EMBEDDING_DIMS = 64
# ---------------------------------------------------------------------------
# Deterministic embedding helper (shared across backends)
# ---------------------------------------------------------------------------
def _make_deterministic_embedder():
cache = {}
counter = [0]
def embed(text, *args, **kwargs):
t = text.lower().strip()
if t not in cache:
vec = [0.0] * EMBEDDING_DIMS
idx = counter[0] % EMBEDDING_DIMS
vec[idx] = 1.0
h = hashlib.sha256(t.encode()).digest()
for i in range(EMBEDDING_DIMS):
vec[i] += float(h[i % len(h)]) / 25500.0
norm = sum(v * v for v in vec) ** 0.5
cache[t] = [v / norm for v in vec]
counter[0] += 1
return cache[t]
mock = MagicMock()
mock.embed.side_effect = embed
mock.config.embedding_dims = EMBEDDING_DIMS
return mock
def _make_mock_llm(entities, relations):
"""Create an LLM mock that returns specific entities and relations."""
mock = MagicMock()
def generate_response(messages, tools):
tool_names = []
for t in tools:
if isinstance(t, dict):
fn = t.get("function", t)
tool_names.append(fn.get("name", ""))
else:
tool_names.append(getattr(t, "name", str(t)))
if any("extract_entities" in n for n in tool_names):
return {
"tool_calls": [
{"name": "extract_entities", "arguments": {"entities": entities}}
]
}
elif any("establish" in n or "relation" in n for n in tool_names):
return {
"tool_calls": [
{"name": "establish_nodes_relations", "arguments": {"entities": relations}}
]
}
elif any("delete" in n for n in tool_names):
return {"tool_calls": []}
return {"tool_calls": []}
mock.generate_response.side_effect = generate_response
return mock
# ===========================================================================
# NEO4J
# ===========================================================================
def _port_open(host, port, timeout=1):
"""Quick TCP check — avoids slow driver-level timeouts."""
import socket
try:
with socket.create_connection((host, port), timeout=timeout):
return True
except OSError:
return False
def _neo4j_available():
if not _port_open("localhost", 7687):
return False
try:
from langchain_neo4j import Neo4jGraph
g = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
g.query("RETURN 1")
return True
except Exception:
return False
requires_neo4j = pytest.mark.skipif(not _neo4j_available(), reason="Neo4j not available")
@pytest.fixture
def neo4j_graph():
"""Create a Neo4j-backed MemoryGraph with mocked LLM/embedder."""
from langchain_neo4j import Neo4jGraph
from mem0.memory.graph_memory import MemoryGraph
mg = MemoryGraph.__new__(MemoryGraph)
mg.graph = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
mg.graph.query("MATCH (n) DETACH DELETE n")
mg.node_label = ":`__Entity__`"
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
mg.config.graph_store.config.base_label = True
yield mg
mg.graph.query("MATCH (n) DETACH DELETE n")
@requires_neo4j
class TestNeo4jDeleteE2E:
def _node_count(self, mg):
return mg.graph.query("MATCH (n) RETURN count(n) AS cnt")[0]["cnt"]
def _valid_edge_count(self, mg):
return mg.graph.query(
"MATCH ()-[r]->() WHERE r.valid IS NULL OR r.valid = true RETURN count(r) AS cnt"
)[0]["cnt"]
def _invalid_edge_count(self, mg):
return mg.graph.query(
"MATCH ()-[r]->() WHERE r.valid = false RETURN count(r) AS cnt"
)[0]["cnt"]
def test_add_creates_graph_data(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
assert self._valid_edge_count(mg) == 1
def test_delete_soft_deletes_relationships(self, neo4j_graph):
"""Neo4j delete() should set r.valid=false (soft-delete), not hard-delete."""
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
assert self._invalid_edge_count(mg) == 0
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 0
assert self._invalid_edge_count(mg) == 1 # soft-deleted, not removed
assert self._node_count(mg) == 2 # nodes preserved
def test_delete_preserves_other_relationships(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Charlie", "entity_type": "person"}],
[{"source": "Alice", "destination": "Charlie", "relationship": "knows"}],
)
mg.add("Alice knows Charlie", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 2
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
assert self._invalid_edge_count(mg) == 1
def test_delete_user_isolation(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.add("Alice likes Bob", {"user_id": "u2"})
assert self._valid_edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
def test_delete_all_hard_deletes(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert self._node_count(mg) == 0
def test_add_delete_add_cycle(self, neo4j_graph):
mg = neo4j_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 0
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._valid_edge_count(mg) == 1
# ===========================================================================
# MEMGRAPH
# ===========================================================================
def _memgraph_available():
if not _port_open("localhost", 7688):
return False
try:
from langchain_memgraph.graphs.memgraph import Memgraph
g = Memgraph("bolt://localhost:7688", "memgraph", "memgraph")
g.query("RETURN 1")
return True
except Exception:
return False
requires_memgraph = pytest.mark.skipif(
not _memgraph_available(), reason="Memgraph not available"
)
@pytest.fixture
def memgraph_graph():
"""Create a Memgraph-backed MemoryGraph with mocked LLM/embedder."""
from langchain_memgraph.graphs.memgraph import Memgraph
from mem0.memory.memgraph_memory import MemoryGraph
mg = MemoryGraph.__new__(MemoryGraph)
mg.graph = Memgraph("bolt://localhost:7688", "memgraph", "memgraph")
mg.graph.query("MATCH (n) DETACH DELETE n")
try:
mg.graph.query("DROP VECTOR INDEX memzero;")
except Exception:
pass
mg.graph.query(
f"CREATE VECTOR INDEX memzero ON :Entity(embedding) "
f"WITH CONFIG {{'dimension': {EMBEDDING_DIMS}, 'capacity': 1000, 'metric': 'cos'}};"
)
try:
mg.graph.query("CREATE INDEX ON :Entity(user_id);")
except Exception:
pass
try:
mg.graph.query("CREATE INDEX ON :Entity;")
except Exception:
pass
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
mg.config.embedder.config = {"embedding_dims": EMBEDDING_DIMS}
yield mg
mg.graph.query("MATCH (n) DETACH DELETE n")
@requires_memgraph
class TestMemgraphDeleteE2E:
def _node_count(self, mg):
return mg.graph.query("MATCH (n:Entity) RETURN count(n) AS cnt")[0]["cnt"]
def _edge_count(self, mg):
return mg.graph.query("MATCH ()-[r]->() RETURN count(r) AS cnt")[0]["cnt"]
def test_add_creates_graph_data(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
assert self._edge_count(mg) == 1
def test_delete_hard_deletes_relationships(self, memgraph_graph):
"""Memgraph delete() should hard-delete the relationship (DELETE r)."""
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 0
assert self._node_count(mg) == 2
def test_delete_preserves_other_relationships(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Charlie", "entity_type": "person"}],
[{"source": "Alice", "destination": "Charlie", "relationship": "knows"}],
)
mg.add("Alice knows Charlie", {"user_id": "u1"})
assert self._edge_count(mg) == 2
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
def test_delete_user_isolation(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.add("Alice likes Bob", {"user_id": "u2"})
assert self._edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
def test_delete_all_hard_deletes(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert self._node_count(mg) == 0
assert self._edge_count(mg) == 0
def test_add_delete_add_cycle(self, memgraph_graph):
mg = memgraph_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 0
mg.add("Alice likes Bob", {"user_id": "u1"})
assert self._edge_count(mg) == 1
# ===========================================================================
# APACHE AGE
# ===========================================================================
def _age_available():
if not _port_open("localhost", 5432):
return False
try:
import age
ag = age.connect(
host="localhost",
port=5432,
dbname="testdb",
user="postgres",
password="testpassword",
)
with ag.connection.cursor() as cur:
cur.execute("CREATE EXTENSION IF NOT EXISTS age;")
cur.execute("SET search_path = ag_catalog, '$user', public;")
ag.connection.commit()
ag.close()
return True
except Exception:
return False
requires_age = pytest.mark.skipif(not _age_available(), reason="Apache AGE not available")
@pytest.fixture
def age_graph():
"""Create an Apache AGE-backed MemoryGraph with mocked LLM/embedder."""
import age
from mem0.memory.apache_age_memory import MemoryGraph
graph_name = "mem0_test_delete"
ag = age.connect(
graph=graph_name,
host="localhost",
port=5432,
dbname="testdb",
user="postgres",
password="testpassword",
)
with ag.connection.cursor() as cur:
cur.execute("CREATE EXTENSION IF NOT EXISTS age;")
cur.execute("SET search_path = ag_catalog, '$user', public;")
ag.connection.commit()
age.setUpAge(ag.connection, graph_name)
ag.connection.commit()
mg = MemoryGraph.__new__(MemoryGraph)
mg.ag = ag
mg.graph_name = graph_name
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
try:
ag.execCypher("MATCH (n) DETACH DELETE n")
ag.commit()
except Exception:
ag.rollback()
yield mg
try:
ag.execCypher("MATCH (n) DETACH DELETE n")
ag.commit()
except Exception:
ag.rollback()
ag.close()
def _age_node_count(mg):
cursor = mg.ag.execCypher("MATCH (n) RETURN count(n)", cols=["cnt"])
rows = cursor.fetchall()
return rows[0][0] if rows else 0
def _age_edge_count(mg):
cursor = mg.ag.execCypher("MATCH ()-[r]->() RETURN count(r)", cols=["cnt"])
rows = cursor.fetchall()
return rows[0][0] if rows else 0
@requires_age
class TestApacheAgeDeleteE2E:
def test_add_creates_graph_data(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_node_count(mg) == 2
assert _age_edge_count(mg) == 1
def test_delete_hard_deletes_relationships(self, age_graph):
"""Apache AGE delete() should hard-delete the relationship."""
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 0
assert _age_node_count(mg) == 2
def test_delete_preserves_other_relationships(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Charlie", "entity_type": "person"}],
[{"source": "Alice", "destination": "Charlie", "relationship": "knows"}],
)
mg.add("Alice knows Charlie", {"user_id": "u1"})
assert _age_edge_count(mg) == 2
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
def test_delete_user_isolation(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
mg.add("Alice likes Bob", {"user_id": "u2"})
assert _age_edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
def test_delete_all_hard_deletes(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert _age_node_count(mg) == 0
assert _age_edge_count(mg) == 0
def test_add_delete_add_cycle(self, age_graph):
mg = age_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 0
mg.add("Alice likes Bob", {"user_id": "u1"})
assert _age_edge_count(mg) == 1
# ===========================================================================
# NEPTUNE (tested via Neo4j OpenCypher — same query language)
# ===========================================================================
def _neptune_test_available():
"""Neptune uses OpenCypher — we test NeptuneBase.delete() against Neo4j."""
if not _port_open("localhost", 7687):
return False
try:
# Mock langchain_aws so NeptuneBase can be imported without AWS deps
sys.modules.setdefault("langchain_aws", MagicMock())
sys.modules.setdefault("botocore", MagicMock())
sys.modules.setdefault("botocore.config", MagicMock())
from mem0.graphs.neptune.base import NeptuneBase # noqa: F401
from langchain_neo4j import Neo4jGraph
g = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
g.query("RETURN 1")
return True
except Exception:
return False
requires_neptune_test = pytest.mark.skipif(
not _neptune_test_available(),
reason="Neo4j not available (used as OpenCypher backend for Neptune tests)",
)
def _make_concrete_neptune_subclass():
"""Create a concrete NeptuneBase subclass for testing, backed by Neo4j."""
# Ensure mocks are in place for import
sys.modules.setdefault("langchain_aws", MagicMock())
sys.modules.setdefault("botocore", MagicMock())
sys.modules.setdefault("botocore.config", MagicMock())
from mem0.graphs.neptune.base import NeptuneBase
class TestableNeptune(NeptuneBase):
def __init__(self):
pass
def _delete_entities_cypher(self, source, destination, relationship, user_id):
cypher = f"""
MATCH (n:`__Entity__` {{name: $source_name, user_id: $user_id}})
-[r:{relationship}]->
(m:`__Entity__` {{name: $dest_name, user_id: $user_id}})
DELETE r
RETURN n.name AS source, m.name AS target, type(r) AS relationship
"""
return cypher, {"source_name": source, "dest_name": destination, "user_id": user_id}
def _delete_all_cypher(self, filters):
return (
"MATCH (n:`__Entity__` {user_id: $user_id}) DETACH DELETE n",
{"user_id": filters["user_id"]},
)
# Stubs for abstract methods not used in delete path
def _add_entities_by_source_cypher(self, *a, **kw): pass
def _add_entities_by_destination_cypher(self, *a, **kw): pass
def _add_relationship_entities_cypher(self, *a, **kw): pass
def _add_new_entities_cypher(self, *a, **kw): pass
def _search_source_node_cypher(self, *a, **kw): pass
def _search_destination_node_cypher(self, *a, **kw): pass
def _get_all_cypher(self, *a, **kw): pass
def _search_graph_db_cypher(self, *a, **kw): pass
return TestableNeptune
@pytest.fixture
def neptune_graph():
"""NeptuneBase subclass backed by a real Neo4j container."""
from langchain_neo4j import Neo4jGraph
cls = _make_concrete_neptune_subclass()
mg = cls()
mg.graph = Neo4jGraph(
url="bolt://localhost:7687",
username="neo4j",
password="testpassword",
refresh_schema=False,
driver_config={"notifications_min_severity": "OFF"},
)
mg.graph.query("MATCH (n) DETACH DELETE n")
mg.node_label = ":`__Entity__`"
mg.llm_provider = "openai"
mg.user_id = None
mg.threshold = 0.99
mg.embedding_model = _make_deterministic_embedder()
mg.llm = MagicMock()
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
yield mg
mg.graph.query("MATCH (n) DETACH DELETE n")
def _neptune_node_count(mg):
return mg.graph.query("MATCH (n) RETURN count(n) AS cnt")[0]["cnt"]
def _neptune_edge_count(mg):
return mg.graph.query("MATCH ()-[r]->() RETURN count(r) AS cnt")[0]["cnt"]
def _neptune_create_entities(mg, user_id):
"""Create test entities directly via Cypher."""
mg.graph.query(f"""
CREATE (a:`__Entity__` {{name: 'alice', user_id: '{user_id}'}})
CREATE (b:`__Entity__` {{name: 'bob', user_id: '{user_id}'}})
CREATE (a)-[:likes]->(b)
""")
@requires_neptune_test
class TestNeptuneDeleteE2E:
"""Test NeptuneBase.delete() using Neo4j as the OpenCypher backend.
Neptune uses standard OpenCypher, the same query language as Neo4j.
This validates that:
- NeptuneBase.delete() correctly calls _delete_entities(to_be_deleted, user_id) with a string
- The generated Cypher from _delete_entities_cypher runs correctly
- User isolation works
"""
def test_delete_removes_relationship(self, neptune_graph):
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
_neptune_create_entities(mg, "u1")
assert _neptune_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _neptune_edge_count(mg) == 0
assert _neptune_node_count(mg) == 2
def test_delete_user_isolation(self, neptune_graph):
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
_neptune_create_entities(mg, "u1")
_neptune_create_entities(mg, "u2")
assert _neptune_edge_count(mg) == 2
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _neptune_edge_count(mg) == 1
def test_delete_passes_user_id_string_not_dict(self, neptune_graph):
"""Verify NeptuneBase.delete() passes filters['user_id'] (string) to _delete_entities."""
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
original = mg._delete_entities
call_args = []
def spy(to_be_deleted, user_id):
call_args.append(("to_be_deleted", to_be_deleted, "user_id", user_id))
return original(to_be_deleted, user_id)
mg._delete_entities = spy
_neptune_create_entities(mg, "u1")
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert len(call_args) == 1
assert call_args[0][3] == "u1"
assert isinstance(call_args[0][3], str)
def test_delete_all(self, neptune_graph):
mg = neptune_graph
_neptune_create_entities(mg, "u1")
assert _neptune_node_count(mg) == 2
mg.delete_all({"user_id": "u1"})
assert _neptune_node_count(mg) == 0
def test_add_delete_add_cycle(self, neptune_graph):
mg = neptune_graph
mg.llm = _make_mock_llm(
[{"entity": "Alice", "entity_type": "person"}, {"entity": "Bob", "entity_type": "person"}],
[{"source": "Alice", "destination": "Bob", "relationship": "likes"}],
)
_neptune_create_entities(mg, "u1")
assert _neptune_edge_count(mg) == 1
mg.delete("Alice likes Bob", {"user_id": "u1"})
assert _neptune_edge_count(mg) == 0
_neptune_create_entities(mg, "u1")
assert _neptune_edge_count(mg) == 1
+736
View File
@@ -0,0 +1,736 @@
"""
End-to-end tests for graph cleanup on memory deletion (issue #3245).
Uses a real Kuzu embedded database to verify that graph entities are
correctly cleaned up when memories are deleted. LLM and embedding calls
are mocked to provide deterministic entity extraction.
Tests are skipped automatically if kuzu is not installed.
"""
import shutil
import tempfile
from unittest.mock import MagicMock, patch
import pytest
from mem0.configs.base import MemoryConfig
try:
import kuzu # noqa: F401
_kuzu_available = True
except ImportError:
_kuzu_available = False
requires_kuzu = pytest.mark.skipif(not _kuzu_available, reason="kuzu is not installed")
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _node_count(kuzu_graph):
"""Return total node count in the Kuzu graph."""
result = kuzu_graph.execute("MATCH (n:Entity) RETURN count(n) AS cnt")
rows = list(result.rows_as_dict())
return int(rows[0]["cnt"])
def _edge_count(kuzu_graph):
"""Return total edge count in the Kuzu graph."""
result = kuzu_graph.execute("MATCH ()-[r:CONNECTED_TO]->() RETURN count(r) AS cnt")
rows = list(result.rows_as_dict())
return int(rows[0]["cnt"])
def _get_edges(kuzu_graph):
"""Return all edges as list of (source, relationship, destination) tuples."""
result = kuzu_graph.execute(
"MATCH (s:Entity)-[r:CONNECTED_TO]->(d:Entity) "
"RETURN s.name AS src, r.name AS rel, d.name AS dst"
)
return [(row["src"], row["rel"], row["dst"]) for row in result.rows_as_dict()]
def _get_nodes(kuzu_graph):
"""Return all node names."""
result = kuzu_graph.execute("MATCH (n:Entity) RETURN n.name AS name, n.user_id AS uid")
return [(row["name"], row["uid"]) for row in result.rows_as_dict()]
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
class MockVectorMemory:
"""Mimics the object returned by vector_store.get()."""
def __init__(self, memory_id, payload, score=0.8):
self.id = memory_id
self.payload = payload
self.score = score
@pytest.fixture
def kuzu_graph_memory():
"""
Create a real Kuzu-backed MemoryGraph with mocked LLM and embedder.
Yields (graph_memory_instance, kuzu_connection) then cleans up.
"""
import os
import kuzu
tmpdir = tempfile.mkdtemp()
db_path = os.path.join(tmpdir, "test.kuzu")
db = kuzu.Database(db_path)
conn = kuzu.Connection(db)
# We'll construct the MemoryGraph by bypassing __init__ and setting up manually
from mem0.memory.kuzu_memory import MemoryGraph
mg = MemoryGraph.__new__(MemoryGraph)
# Real Kuzu connection
mg.db = db
mg.graph = conn
mg.node_label = ":Entity"
mg.rel_label = ":CONNECTED_TO"
mg.kuzu_create_schema()
# Deterministic embedding: use one-hot-style vectors per entity name
# to avoid accidental cosine similarity matches between different entities
embedding_dims = 64
mg.embedding_dims = embedding_dims
_embed_cache = {}
_embed_counter = [0]
def deterministic_embed(text):
"""Generate a deterministic, near-orthogonal embedding for each unique text."""
text_lower = text.lower().strip()
if text_lower not in _embed_cache:
# Create a sparse vector — set a unique dimension to 1.0
vec = [0.0] * embedding_dims
idx = _embed_counter[0] % embedding_dims
vec[idx] = 1.0
# Add small noise to other dims so it's not exactly zero
import hashlib
h = hashlib.sha256(text_lower.encode()).digest()
for i in range(embedding_dims):
vec[i] += float(h[i % len(h)]) / 25500.0 # tiny noise
norm = sum(v * v for v in vec) ** 0.5
_embed_cache[text_lower] = [v / norm for v in vec]
_embed_counter[0] += 1
return _embed_cache[text_lower]
mock_embedder = MagicMock()
mock_embedder.embed.side_effect = deterministic_embed
mock_embedder.config.embedding_dims = embedding_dims
mg.embedding_model = mock_embedder
# Mock LLM — configured per-test via mock_embedder
mg.llm = MagicMock()
mg.llm_provider = "openai"
mg.user_id = None
# High threshold so only identical entity names merge, not similar ones
mg.threshold = 0.99
mg.config = MagicMock()
mg.config.graph_store.custom_prompt = None
yield mg, conn
# Cleanup
conn.close()
shutil.rmtree(tmpdir, ignore_errors=True)
def _setup_llm_for_entities(mg, entities, relations):
"""
Configure the mock LLM to return specific entities and relations.
entities: list of {"entity": str, "entity_type": str}
relations: list of {"source": str, "destination": str, "relationship": str}
"""
def generate_response(messages, tools):
# Detect which tool is being called based on tool definition names
tool_names = []
for t in tools:
if isinstance(t, dict):
fn = t.get("function", t)
tool_names.append(fn.get("name", ""))
else:
tool_names.append(getattr(t, "name", str(t)))
if any("extract_entities" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "extract_entities",
"arguments": {"entities": entities},
}
]
}
elif any("establish" in n or "relation" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "establish_nodes_relations",
"arguments": {"entities": relations},
}
]
}
elif any("delete" in n for n in tool_names):
# For _get_delete_entities_from_search_output during add() — return nothing to delete
return {"tool_calls": []}
return {"tool_calls": []}
mg.llm.generate_response.side_effect = generate_response
# ---------------------------------------------------------------------------
# End-to-end tests
# ---------------------------------------------------------------------------
@requires_kuzu
class TestKuzuGraphDeleteE2E:
"""End-to-end tests using a real Kuzu database."""
def test_add_creates_nodes_and_edges(self, kuzu_graph_memory):
"""Baseline: verify add() actually creates graph data."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
assert _node_count(conn) == 2
assert _edge_count(conn) == 1
edges = _get_edges(conn)
assert ("alice", "likes", "bob") in edges
def test_delete_removes_edges_created_by_add(self, kuzu_graph_memory):
"""Core test: delete() should remove the relationships that add() created."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
# Now delete using the same text — should remove the relationship
mg.delete("Alice likes Bob", filters)
assert _edge_count(conn) == 0
# Nodes remain (we don't delete nodes on single memory delete)
assert _node_count(conn) == 2
def test_delete_only_removes_matching_edges(self, kuzu_graph_memory):
"""delete() should only remove edges matching the extracted relationships."""
mg, conn = kuzu_graph_memory
# First add: Alice likes Bob
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
# Second add: Alice knows Charlie
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Charlie", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Charlie", "relationship": "knows"},
],
)
mg.add("Alice knows Charlie", filters)
assert _edge_count(conn) == 2
# Delete only the "Alice likes Bob" memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
mg.delete("Alice likes Bob", filters)
assert _edge_count(conn) == 1
edges = _get_edges(conn)
assert ("alice", "knows", "charlie") in edges
assert ("alice", "likes", "bob") not in edges
def test_delete_with_different_user_id_does_not_affect_other_users(self, kuzu_graph_memory):
"""delete() scoped to user_id should not touch another user's graph data."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
# Add for user1
mg.add("Alice likes Bob", {"user_id": "user1"})
# Add same data for user2
mg.add("Alice likes Bob", {"user_id": "user2"})
assert _edge_count(conn) == 2
# Delete only user1's data
mg.delete("Alice likes Bob", {"user_id": "user1"})
assert _edge_count(conn) == 1
# Remaining edge belongs to user2
nodes = _get_nodes(conn)
user2_nodes = [n for n in nodes if n[1] == "user2"]
assert len(user2_nodes) == 2
def test_delete_nonexistent_relationship_is_safe(self, kuzu_graph_memory):
"""delete() on data that doesn't exist in the graph should be a no-op."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "hates"},
],
)
filters = {"user_id": "test_user"}
# Nothing in the graph yet
assert _edge_count(conn) == 0
assert _node_count(conn) == 0
# Should not raise
mg.delete("Alice hates Bob", filters)
assert _edge_count(conn) == 0
assert _node_count(conn) == 0
def test_delete_with_llm_failure_does_not_raise(self, kuzu_graph_memory):
"""If LLM fails during entity extraction, delete() should not raise."""
mg, conn = kuzu_graph_memory
# Make LLM raise
mg.llm.generate_response.side_effect = RuntimeError("LLM service down")
filters = {"user_id": "test_user"}
# Should not raise
mg.delete("Alice likes Bob", filters)
def test_delete_with_empty_entity_extraction(self, kuzu_graph_memory):
"""If LLM returns no entities, delete() should be a no-op."""
mg, conn = kuzu_graph_memory
# Add real data
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
# Now delete but LLM returns no entities
_setup_llm_for_entities(mg, entities=[], relations=[])
mg.delete("some text", filters)
# Data should still be there
assert _edge_count(conn) == 1
def test_delete_all_removes_everything_for_user(self, kuzu_graph_memory):
"""delete_all() should remove all nodes/edges for a user (baseline behavior)."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
mg.add("Alice likes Bob", filters)
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Bob", "entity_type": "person"},
{"entity": "Charlie", "entity_type": "person"},
],
relations=[
{"source": "Bob", "destination": "Charlie", "relationship": "knows"},
],
)
mg.add("Bob knows Charlie", filters)
assert _node_count(conn) >= 3
assert _edge_count(conn) == 2
mg.delete_all(filters)
assert _node_count(conn) == 0
assert _edge_count(conn) == 0
def test_add_delete_add_cycle(self, kuzu_graph_memory):
"""Verify that add → delete → re-add works correctly."""
mg, conn = kuzu_graph_memory
_setup_llm_for_entities(
mg,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
filters = {"user_id": "test_user"}
# Add
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
# Delete
mg.delete("Alice likes Bob", filters)
assert _edge_count(conn) == 0
# Re-add
mg.add("Alice likes Bob", filters)
assert _edge_count(conn) == 1
edges = _get_edges(conn)
assert ("alice", "likes", "bob") in edges
@requires_kuzu
class TestMemoryDeleteWithGraphE2E:
"""
End-to-end tests for Memory.delete() with graph enabled.
Uses a real Kuzu database for the graph store and mocks for
the vector store, LLM, and embedder.
"""
@pytest.fixture
def memory_with_graph(self):
"""Create a Memory instance with a real Kuzu graph backend."""
import os
import kuzu
tmpdir = tempfile.mkdtemp()
with (
patch("mem0.utils.factory.EmbedderFactory.create") as mock_embedder_factory,
patch("mem0.utils.factory.VectorStoreFactory.create") as mock_vector_factory,
patch("mem0.utils.factory.LlmFactory.create") as mock_llm_factory,
patch("mem0.memory.storage.SQLiteManager") as mock_sqlite,
):
_mem_embed_cache = {}
_mem_embed_counter = [0]
def _mem_deterministic_embed(text, *args, **kwargs):
text_lower = text.lower().strip()
if text_lower not in _mem_embed_cache:
import hashlib
vec = [0.0] * 64
idx = _mem_embed_counter[0] % 64
vec[idx] = 1.0
h = hashlib.sha256(text_lower.encode()).digest()
for i in range(64):
vec[i] += float(h[i % len(h)]) / 25500.0
norm = sum(v * v for v in vec) ** 0.5
_mem_embed_cache[text_lower] = [v / norm for v in vec]
_mem_embed_counter[0] += 1
return _mem_embed_cache[text_lower]
mock_embedder = MagicMock()
mock_embedder.embed.side_effect = _mem_deterministic_embed
mock_embedder.config.embedding_dims = 64
mock_embedder_factory.return_value = mock_embedder
mock_vector_store = MagicMock()
mock_vector_factory.return_value = mock_vector_store
mock_llm = MagicMock()
mock_llm_factory.return_value = mock_llm
mock_sqlite.return_value = MagicMock()
from mem0.memory.main import Memory
config = MemoryConfig()
memory = Memory(config)
# Now wire up a real Kuzu graph
db_path = os.path.join(tmpdir, "test.kuzu")
db = kuzu.Database(db_path)
conn = kuzu.Connection(db)
from mem0.memory.kuzu_memory import MemoryGraph as KuzuMemoryGraph
graph = KuzuMemoryGraph.__new__(KuzuMemoryGraph)
graph.db = db
graph.graph = conn
graph.node_label = ":Entity"
graph.rel_label = ":CONNECTED_TO"
graph.kuzu_create_schema()
graph.embedding_dims = 64
graph.embedding_model = mock_embedder
graph.llm = mock_llm
graph.llm_provider = "openai"
graph.user_id = None
graph.threshold = 0.99
graph.config = MagicMock()
graph.config.graph_store.custom_prompt = None
memory.graph = graph
memory.enable_graph = True
yield memory, mock_vector_store, mock_llm, conn
conn.close()
shutil.rmtree(tmpdir, ignore_errors=True)
def test_memory_delete_triggers_graph_cleanup(self, memory_with_graph):
"""
Full integration: Memory.delete() should clean up both vector store and graph.
"""
memory, mock_vs, mock_llm, conn = memory_with_graph
# 1. Manually add entities to the graph (simulating what add() would do)
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
assert _edge_count(conn) == 1
# 2. Set up mock vector store to return this memory
mock_vs.get.return_value = MockVectorMemory(
"mem-1",
{"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"},
)
# 3. Delete the memory
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
# 4. Verify graph was cleaned up
assert _edge_count(conn) == 0
# 5. Verify vector store was also cleaned up
mock_vs.delete.assert_called_once_with(vector_id="mem-1")
def test_memory_delete_with_graph_preserves_other_users_data(self, memory_with_graph):
"""Deleting user1's memory should not affect user2's graph data."""
memory, mock_vs, mock_llm, conn = memory_with_graph
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
# Add data for two users
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
memory.graph.add("Alice likes Bob", {"user_id": "user-2"})
assert _edge_count(conn) == 2
# Delete only user-1's memory
mock_vs.get.return_value = MockVectorMemory(
"mem-1",
{"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"},
)
memory.delete("mem-1")
# user-2's data should be intact
assert _edge_count(conn) == 1
nodes = _get_nodes(conn)
remaining_user_ids = set(uid for _, uid in nodes)
assert "user-2" in remaining_user_ids
def test_memory_delete_graph_failure_still_deletes_vector(self, memory_with_graph):
"""If graph cleanup fails, vector store deletion should still proceed."""
memory, mock_vs, mock_llm, conn = memory_with_graph
# Make LLM raise during entity extraction (graph cleanup will fail)
mock_llm.generate_response.side_effect = RuntimeError("LLM exploded")
mock_vs.get.return_value = MockVectorMemory(
"mem-1",
{"data": "Alice likes Bob", "user_id": "user-1", "hash": "abc"},
)
result = memory.delete("mem-1")
assert result == {"message": "Memory deleted successfully!"}
mock_vs.delete.assert_called_once_with(vector_id="mem-1")
def test_memory_delete_all_uses_bulk_not_per_memory(self, memory_with_graph):
"""delete_all() should use delete_all() on graph, not per-memory delete()."""
memory, mock_vs, mock_llm, conn = memory_with_graph
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
assert _edge_count(conn) == 1
# Set up vector store to return memories for deletion
mem1 = MockVectorMemory("mem-1", {"data": "Alice likes Bob", "user_id": "user-1"})
mock_vs.list.return_value = ([mem1], 1)
mock_vs.get.return_value = mem1
memory.delete_all(user_id="user-1")
# After delete_all, graph should be empty (via graph.delete_all)
assert _edge_count(conn) == 0
assert _node_count(conn) == 0
def test_memory_delete_nonexistent_raises_without_graph_side_effects(self, memory_with_graph):
"""Deleting a non-existent memory should raise ValueError without touching graph."""
memory, mock_vs, mock_llm, conn = memory_with_graph
# Add some graph data that should NOT be affected
_setup_llm_for_memory_graph(
mock_llm,
entities=[
{"entity": "Alice", "entity_type": "person"},
{"entity": "Bob", "entity_type": "person"},
],
relations=[
{"source": "Alice", "destination": "Bob", "relationship": "likes"},
],
)
memory.graph.add("Alice likes Bob", {"user_id": "user-1"})
assert _edge_count(conn) == 1
# Memory doesn't exist in vector store
mock_vs.get.return_value = None
with pytest.raises(ValueError, match="Memory with id non-existent not found"):
memory.delete("non-existent")
# Graph data should be untouched
assert _edge_count(conn) == 1
def _setup_llm_for_memory_graph(mock_llm, entities, relations):
"""Configure mock LLM for the Memory-level graph operations."""
def generate_response(messages, tools):
tool_names = []
for t in tools:
if isinstance(t, dict):
fn = t.get("function", t)
tool_names.append(fn.get("name", ""))
else:
tool_names.append(getattr(t, "name", str(t)))
if any("extract_entities" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "extract_entities",
"arguments": {"entities": entities},
}
]
}
elif any("establish" in n or "relation" in n for n in tool_names):
return {
"tool_calls": [
{
"name": "establish_nodes_relations",
"arguments": {"entities": relations},
}
]
}
elif any("delete" in n for n in tool_names):
return {"tool_calls": []}
return {"tool_calls": []}
mock_llm.generate_response.side_effect = generate_response
+3 -1
View File
@@ -186,7 +186,9 @@ def test_delete(memory_instance):
result = memory_instance.delete("test_id")
memory_instance._delete_memory.assert_called_once_with("test_id")
# delete() now fetches the memory first and passes it to _delete_memory
existing_memory = memory_instance.vector_store.get.return_value
memory_instance._delete_memory.assert_called_once_with("test_id", existing_memory)
assert result["message"] == "Memory deleted successfully!"