From c5e748b46905b29d8d017ed6b79d055bd30a57d4 Mon Sep 17 00:00:00 2001 From: Soumil Rathi Date: Mon, 13 Apr 2026 13:14:40 -0700 Subject: [PATCH] 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) --- tests/test_main.py | 95 ++++++---------------------------------------- 1 file changed, 12 insertions(+), 83 deletions(-) diff --git a/tests/test_main.py b/tests/test_main.py index ff4a97b2e..06fa8d9a4 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -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."""