diff --git a/mem0/graphs/neptune/base.py b/mem0/graphs/neptune/base.py index 458c51755..2732ceab3 100644 --- a/mem0/graphs/neptune/base.py +++ b/mem0/graphs/neptune/base.py @@ -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) diff --git a/mem0/memory/apache_age_memory.py b/mem0/memory/apache_age_memory.py index 74c51b555..e3c69f1a4 100644 --- a/mem0/memory/apache_age_memory.py +++ b/mem0/memory/apache_age_memory.py @@ -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"] diff --git a/mem0/memory/graph_memory.py b/mem0/memory/graph_memory.py index a0c89cf76..86e49099e 100644 --- a/mem0/memory/graph_memory.py +++ b/mem0/memory/graph_memory.py @@ -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"] diff --git a/mem0/memory/kuzu_memory.py b/mem0/memory/kuzu_memory.py index a025aec36..1ed769fb4 100644 --- a/mem0/memory/kuzu_memory.py +++ b/mem0/memory/kuzu_memory.py @@ -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"] diff --git a/mem0/memory/main.py b/mem0/memory/main.py index 84b760790..c9b67c75c 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -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) diff --git a/mem0/memory/memgraph_memory.py b/mem0/memory/memgraph_memory.py index 84b6d0ca8..c7f52ad79 100644 --- a/mem0/memory/memgraph_memory.py +++ b/mem0/memory/memgraph_memory.py @@ -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"): diff --git a/tests/test_graph_delete.py b/tests/test_graph_delete.py new file mode 100644 index 000000000..a54d6d4f6 --- /dev/null +++ b/tests/test_graph_delete.py @@ -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"}) diff --git a/tests/test_graph_delete_docker.py b/tests/test_graph_delete_docker.py new file mode 100644 index 000000000..091f32c22 --- /dev/null +++ b/tests/test_graph_delete_docker.py @@ -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 diff --git a/tests/test_graph_delete_e2e.py b/tests/test_graph_delete_e2e.py new file mode 100644 index 000000000..0169d2ce0 --- /dev/null +++ b/tests/test_graph_delete_e2e.py @@ -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 diff --git a/tests/test_main.py b/tests/test_main.py index e179f3715..96bcb77fc 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -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!"