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:
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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"):
|
||||
|
||||
@@ -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"})
|
||||
@@ -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
|
||||
@@ -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
@@ -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!"
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user