fix: remove all graph references from tests/test_main.py
Strip graph_memory patches, GraphStoreFactory mocks, graph_store config, enable_graph parametrize, _add_to_graph assertions, and "relations" assertions from all test fixtures and test functions. Graph store has been removed from the codebase. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
+12
-83
@@ -22,20 +22,13 @@ def memory_instance():
|
||||
patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store,
|
||||
patch("mem0.utils.factory.LlmFactory") as mock_llm,
|
||||
patch("mem0.memory.telemetry.capture_event"),
|
||||
patch("mem0.memory.graph_memory.MemoryGraph"),
|
||||
patch("mem0.memory.main.GraphStoreFactory") as mock_graph_store,
|
||||
):
|
||||
mock_embedder.create.return_value = Mock()
|
||||
mock_vector_store.create.return_value = Mock()
|
||||
mock_vector_store.create.return_value.search.return_value = []
|
||||
mock_llm.create.return_value = Mock()
|
||||
|
||||
# Create a mock instance that won't try to access config attributes
|
||||
mock_graph_instance = Mock()
|
||||
mock_graph_store.create.return_value = mock_graph_instance
|
||||
|
||||
config = MemoryConfig(version="v1.1")
|
||||
config.graph_store.config = {"some_config": "value"}
|
||||
return Memory(config)
|
||||
|
||||
|
||||
@@ -46,54 +39,33 @@ def memory_custom_instance():
|
||||
patch("mem0.memory.main.VectorStoreFactory") as mock_vector_store,
|
||||
patch("mem0.utils.factory.LlmFactory") as mock_llm,
|
||||
patch("mem0.memory.telemetry.capture_event"),
|
||||
patch("mem0.memory.graph_memory.MemoryGraph"),
|
||||
patch("mem0.memory.main.GraphStoreFactory") as mock_graph_store,
|
||||
):
|
||||
mock_embedder.create.return_value = Mock()
|
||||
mock_vector_store.create.return_value = Mock()
|
||||
mock_vector_store.create.return_value.search.return_value = []
|
||||
mock_llm.create.return_value = Mock()
|
||||
|
||||
# Create a mock instance that won't try to access config attributes
|
||||
mock_graph_instance = Mock()
|
||||
mock_graph_store.create.return_value = mock_graph_instance
|
||||
|
||||
config = MemoryConfig(
|
||||
version="v1.1",
|
||||
custom_instructions="custom prompt extracting memory in json format",
|
||||
)
|
||||
config.graph_store.config = {"some_config": "value"}
|
||||
return Memory(config)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version, enable_graph", [("v1.0", False), ("v1.1", True)])
|
||||
def test_add(memory_instance, version, enable_graph):
|
||||
@pytest.mark.parametrize("version", ["v1.0", "v1.1"])
|
||||
def test_add(memory_instance, version):
|
||||
memory_instance.config.version = version
|
||||
if not enable_graph:
|
||||
memory_instance.graph = None
|
||||
memory_instance._add_to_vector_store = Mock(return_value=[{"memory": "Test memory", "event": "ADD"}])
|
||||
memory_instance._add_to_graph = Mock(return_value=[])
|
||||
|
||||
result = memory_instance.add(messages=[{"role": "user", "content": "Test message"}], user_id="test_user")
|
||||
|
||||
if enable_graph:
|
||||
assert "results" in result
|
||||
assert result["results"] == [{"memory": "Test memory", "event": "ADD"}]
|
||||
assert "relations" in result
|
||||
assert result["relations"] == []
|
||||
else:
|
||||
assert "results" in result
|
||||
assert result["results"] == [{"memory": "Test memory", "event": "ADD"}]
|
||||
assert "results" in result
|
||||
assert result["results"] == [{"memory": "Test memory", "event": "ADD"}]
|
||||
|
||||
memory_instance._add_to_vector_store.assert_called_once_with(
|
||||
[{"role": "user", "content": "Test message"}], {"user_id": "test_user"}, {"user_id": "test_user"}, True
|
||||
)
|
||||
|
||||
# Remove the conditional assertion for _add_to_graph
|
||||
memory_instance._add_to_graph.assert_called_once_with(
|
||||
[{"role": "user", "content": "Test message"}], {"user_id": "test_user"}
|
||||
)
|
||||
|
||||
|
||||
def test_get(memory_instance):
|
||||
mock_memory = Mock(
|
||||
@@ -120,11 +92,9 @@ def test_get(memory_instance):
|
||||
assert result["metadata"] == {"extra_field": "extra_value"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version, enable_graph", [("v1.0", False), ("v1.1", True)])
|
||||
def test_search(memory_instance, version, enable_graph):
|
||||
@pytest.mark.parametrize("version", ["v1.0", "v1.1"])
|
||||
def test_search(memory_instance, version):
|
||||
memory_instance.config.version = version
|
||||
if not enable_graph:
|
||||
memory_instance.graph = None
|
||||
mock_memories = [
|
||||
Mock(id="1", payload={"data": "Memory 1", "user_id": "test_user"}, score=0.9),
|
||||
Mock(id="2", payload={"data": "Memory 2", "user_id": "test_user"}, score=0.8),
|
||||
@@ -132,8 +102,6 @@ def test_search(memory_instance, version, enable_graph):
|
||||
memory_instance.vector_store.search = Mock(return_value=mock_memories)
|
||||
memory_instance.vector_store.keyword_search = Mock(return_value=None) # No BM25
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
if memory_instance.graph:
|
||||
memory_instance.graph.search = Mock(return_value=[{"relation": "test_relation"}])
|
||||
|
||||
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test query"), \
|
||||
patch("mem0.memory.main.extract_entities", return_value=[]):
|
||||
@@ -148,20 +116,11 @@ def test_search(memory_instance, version, enable_graph):
|
||||
assert result["results"][0]["score"] == pytest.approx(0.9)
|
||||
assert "score_breakdown" in result["results"][0]
|
||||
|
||||
if enable_graph:
|
||||
assert "relations" in result
|
||||
assert result["relations"] == [{"relation": "test_relation"}]
|
||||
else:
|
||||
assert "relations" not in result
|
||||
|
||||
# Hybrid pipeline over-fetches: max(100*4, 60) = 400
|
||||
memory_instance.vector_store.search.assert_called_once_with(
|
||||
query="test query", vectors=[0.1, 0.2, 0.3], top_k=400, filters={"user_id": "test_user"}
|
||||
)
|
||||
|
||||
if enable_graph:
|
||||
memory_instance.graph.search.assert_called_once_with("test query", {"user_id": "test_user"}, 100)
|
||||
|
||||
|
||||
def test_update(memory_instance):
|
||||
memory_instance.embedding_model = Mock()
|
||||
@@ -218,17 +177,13 @@ def test_delete(memory_instance):
|
||||
assert result["message"] == "Memory deleted successfully!"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("version, enable_graph", [("v1.0", False), ("v1.1", True)])
|
||||
def test_delete_all(memory_instance, version, enable_graph):
|
||||
@pytest.mark.parametrize("version", ["v1.0", "v1.1"])
|
||||
def test_delete_all(memory_instance, version):
|
||||
memory_instance.config.version = version
|
||||
if not enable_graph:
|
||||
memory_instance.graph = None
|
||||
mock_memories = [Mock(id="1"), Mock(id="2")]
|
||||
memory_instance.vector_store.list = Mock(return_value=(mock_memories, None))
|
||||
memory_instance.vector_store.reset = Mock()
|
||||
memory_instance._delete_memory = Mock()
|
||||
if memory_instance.graph:
|
||||
memory_instance.graph.delete_all = Mock()
|
||||
|
||||
result = memory_instance.delete_all(user_id="test_user")
|
||||
|
||||
@@ -236,37 +191,20 @@ def test_delete_all(memory_instance, version, enable_graph):
|
||||
# Ensure the collection is NOT dropped — only matched memories should be removed
|
||||
memory_instance.vector_store.reset.assert_not_called()
|
||||
|
||||
if enable_graph:
|
||||
memory_instance.graph.delete_all.assert_called_once_with({"user_id": "test_user"})
|
||||
|
||||
assert result["message"] == "Memories deleted successfully!"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"version, enable_graph, expected_result",
|
||||
"version, expected_result",
|
||||
[
|
||||
("v1.0", False, {"results": [{"id": "1", "memory": "Memory 1", "user_id": "test_user"}]}),
|
||||
("v1.1", False, {"results": [{"id": "1", "memory": "Memory 1", "user_id": "test_user"}]}),
|
||||
(
|
||||
"v1.1",
|
||||
True,
|
||||
{
|
||||
"results": [{"id": "1", "memory": "Memory 1", "user_id": "test_user"}],
|
||||
"relations": [{"source": "entity1", "relationship": "rel", "target": "entity2"}],
|
||||
},
|
||||
),
|
||||
("v1.0", {"results": [{"id": "1", "memory": "Memory 1", "user_id": "test_user"}]}),
|
||||
("v1.1", {"results": [{"id": "1", "memory": "Memory 1", "user_id": "test_user"}]}),
|
||||
],
|
||||
)
|
||||
def test_get_all(memory_instance, version, enable_graph, expected_result):
|
||||
def test_get_all(memory_instance, version, expected_result):
|
||||
memory_instance.config.version = version
|
||||
if not enable_graph:
|
||||
memory_instance.graph = None
|
||||
mock_memories = [Mock(id="1", payload={"data": "Memory 1", "user_id": "test_user"})]
|
||||
memory_instance.vector_store.list = Mock(return_value=(mock_memories, None))
|
||||
if memory_instance.graph:
|
||||
memory_instance.graph.get_all = Mock(
|
||||
return_value=[{"source": "entity1", "relationship": "rel", "target": "entity2"}]
|
||||
)
|
||||
|
||||
result = memory_instance.get_all(user_id="test_user")
|
||||
|
||||
@@ -279,17 +217,8 @@ def test_get_all(memory_instance, version, enable_graph, expected_result):
|
||||
assert result_item["memory"] == expected_item["memory"]
|
||||
assert result_item["user_id"] == expected_item["user_id"]
|
||||
|
||||
if enable_graph:
|
||||
assert "relations" in result
|
||||
assert result["relations"] == expected_result["relations"]
|
||||
else:
|
||||
assert "relations" not in result
|
||||
|
||||
memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=100)
|
||||
|
||||
if enable_graph:
|
||||
memory_instance.graph.get_all.assert_called_once_with({"user_id": "test_user"}, 100)
|
||||
|
||||
|
||||
def test_no_telemetry_vector_store_when_disabled():
|
||||
"""VectorStoreFactory should only be called once (for user data) when telemetry is disabled."""
|
||||
|
||||
Reference in New Issue
Block a user